Passer à la navigation

nemo_rl.models.huggingface.common

Afficher en Markdown

Contenu du Module

Classes

NomDescription
FlashAttentionKwargsClasse de données pour stocker les kwargs de FlashAttention v2.
ModelFlagEnum qui définit des drapeaux spéciaux pour les comportements spécifiques des modèles.

Fonctions

NomDescription
is_gemma_modelAucun
group_and_cat_tensorsRegroupe et concatène des tenseurs selon les tailles de groupe, puis les remplit pour former un tenseur 2D.
pack_sequencesEmballe des séquences en lignes où chaque ligne concatène plusieurs séquences.
unpack_tensorDéballe un tenseur emballé en séquences individuelles remplies à la même longueur.
get_flash_attention_kwargsRenvoie les kwargs requis pour les fonctions avant FlashAttention v2.

Données

Tensor

API

(reste du document traduit de la même manière, en conservant la structure MDX et les éléments de code)