nemo_rl.models.megatron.common
nemo_rl.models.megatron.common
Module Contents
Fonctions
API
Empaqueter des séquences pour le traitement du modèle Megatron avec une parallélisation de contexte facultative.
Arguments : input_ids: ID de jetons d’entrée [batch_size, seq_length] seq_lengths: Longueurs de séquence réelles pour chaque échantillon [batch_size] pad_individual_seqs_to_multiple_of: Bourrer les séquences individuelles à un multiple de cette valeur pad_packed_seq_to: Bourrer les séquences empaquetées jusqu’à cette valeur (avant CP) cp_size: Taille de la parallélisation de contexte
Retourne : Tuple de :
- packed_input_ids: Tenseur d’entrée empaqueté [1, T]
- input_ids_cp_sharded: Tenseur d’entrée fragmenté [cp_size, T // cp_size]
- packed_seq_params: Objet PackedSeqParams
- cu_seqlens: Longueurs de séquence cumulatives
- cu_seqlens_padded: Longueurs de séquence cumulatives bourrées
Désempaqueter des séquences à partir du format de sortie Megatron.
Arguments : output_tensor: Tenseur de sortie empaqueté [1, T, vocab_size] seq_lengths: Longueurs de séquence réelles pour chaque échantillon cu_seqlens: Longueurs de séquence cumulatives cu_seqlens_padded: Longueurs de séquence cumulatives bourrées (si CP était utilisé) original_batch_size: Taille de batch originale original_seq_length: Longueur de séquence maximale originale
Retourne : Tenseur de sortie désempaqueté [batch_size, seq_length, vocab_size]
Étape d’entraînement avant avec prise en charge des séquences empaquetées et de la parallélisation de contexte.
Arguments : state (GlobalState): État global de l’exécution global_valid_seqs: Nombre global de séquences valides global_valid_toks: Nombre global de jetons valides data_iterator: Itérateur de données d’entrée model (GPTModel): Le modèle GPT loss_fn (LossFunction): Fonction de perte à appliquer pack_sequences (bool): Indique s’il faut empaqueter les séquences pour l’efficacité seq_length_key (Optional[str]): Clé dans data_dict contenant les longueurs de séquence réelles cp_normalize (bool): Indique s’il faut normaliser la perte par cp_size policy_cfg (Optional[dict]): Configuration de stratégie contenant les paramètres de génération
Notes sur les séquences empaquetées avec parallélisation de contexte (CP) :
- Lorsque CP > 1, chaque séquence est bourrée à un multiple de (cp_size * 2)
- Le facteur 2 assure un équilibrage de charge pour l’attention causale
- cu_seqlens suit les limites de séquence réelles
- cu_seqlens_padded suit les limites de séquence bourrées pour CP
- Nécessite TransformerEngine >= 1.10 pour la prise en charge CP
Diffuse un tenseur de src_rank à tous les rangs du groupe à l’aide de broadcast_object_list pour les métadonnées.
Gère le cas où le tenseur d’entrée peut être None sur des rangs non sources. Si le tenseur d’entrée est fourni sur des rangs non sources, il doit avoir la forme et le type de données correspondant au tenseur sur le rang source.
Arguments : tensor: Le tenseur à diffuser sur le rang source. Peut être None sur des rangs non sources (sera créé avec la forme/dtype correcte). Si ce n’est pas None sur des rangs non sources, il est utilisé comme tampon pour la diffusion et doit correspondre aux métadonnées du tenseur source. src_rank (int): Le rang global du processus source. group: Le groupe de processus pour la communication.
Retourne : torch.Tensor: Le tenseur diffusé. Sur des rangs non sources, ce sera le tenseur reçu de la source.
Lève : ValueError: Si le tenseur est None sur le rang source, ou si un tenseur fourni sur un rang non source a une forme/dtype/device incompatible. TypeError: Si la diffusion des métadonnées échoue (par exemple, en raison de problèmes de sérialisation).