nemo_rl.algorithms.interfaces

Afficher en Markdown

Contenu du module

Classes

NomDescription
LossTypeCréer une collection de paires nom/valeur.
LossFunctionSignature pour les fonctions de perte utilisées dans les algorithmes d’apprentissage par renforcement.

API

class nemo_rl.algorithms.interfaces.LossType(*args, **kwds)

Bases: enum.Enum

TOKEN_LEVEL

Valeur: token_level

SEQUENCE_LEVEL

Valeur: sequence_level

class nemo_rl.algorithms.interfaces.LossFunction

Bases: typing.Protocol

Signature pour les fonctions de perte utilisées dans les algorithmes d’apprentissage par renforcement.

Les fonctions de perte calculent une valeur de perte scalaire et des métriques associées à partir des log-probabilités du modèle et d’autres données contenues dans un BatchedDataDict.

loss_type: nemo_rl.algorithms.interfaces.LossType

Valeur: None

next_token_logits: torch.Tensor,
data: nemo_rl.distributed.batched_data_dict.BatchedDataDict,
global_valid_seqs: torch.Tensor,
global_valid_toks: torch.Tensor
) -> tuple[torch.Tensor, dict[str, typing.Any]]

Calculer la perte et les métriques à partir des log-probabilités et d’autres données.

Args: next_token_logits: Logits du modèle, généralement avec la forme [batch_size, seq_len, vocab_size]. Pour chaque position (b, i), contient la distribution de logits sur l’ensemble du vocabulaire pour prédire le prochain jeton (à la position i+1). Par exemple, lors du traitement de “The cat sat on”, next_token_logits[b, 3] contiendrait les logits pour prédire le mot qui suit “on”. data: Dictionnaire contenant toutes les données pertinentes pour le calcul de la perte telles que les récompenses, les valeurs, les actions, les avantages, les masques et autres informations spécifiques à l’algorithme nécessaires pour le calcul particulier de la perte. global_valid_seqs: torch.Tensor ce tenseur doit contenir le nombre de séquences valides dans le micro-lot. Il est utilisé pour la normalisation globale des pertes/métriques qui sont calculées au niveau de la séquence et doivent être agrégées sur tous les micro-lots. global_valid_toks: torch.Tensor Ce tenseur doit contenir le nombre de jetons valides dans le micro-lot. Il est utilisé pour la normalisation globale des pertes/métriques qui sont calculées au niveau du jeton et doivent être agrégées sur tous les micro-lots.

Returns: tuple: (perte, métriques)

  • perte : Un tenseur scalaire représentant la valeur de perte à minimiser pendant l’entraînement
  • métriques : Un dictionnaire de métriques liées au calcul de la perte, qui peut inclure des pertes composantes, des statistiques sur les gradients/récompenses et autres informations de diagnostic