nemo_rl.algorithms.interfaces
nemo_rl.algorithms.interfaces
Contenu du module
Classes
API
Bases: enum.Enum
Valeur: token_level
Valeur: sequence_level
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.
Valeur: None
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