Rembourrage dans NeMo RL

Afficher en Markdown

Ce document explique le rembourrage dans NeMo RL et pourquoi un rembourrage cohérent est critique pour le framework.

Approche de Rembourrage

NeMo RL utilise un rembourrage à droite pour toutes les opérations de tenseurs, où les jetons de rembourrage sont ajoutés à droite/à la fin des séquences :

[101, 2054, 2003, 0, 0] # Longueur 3
[101, 2054, 2003, 2001, 1996] # Longueur 5 (pas de rembourrage nécessaire)
[101, 2054, 0, 0, 0] # Longueur 2

Cette approche :

  1. S’aligne naturellement avec le traitement LLM : Les jetons sont traités de gauche à droite.
  2. Garde les jetons significatifs contigus : Tous les jetons valides apparaissent au début des tenseurs.
  3. Simplifie l’indexation et les opérations : Les limites des jetons valides sont facilement définies avec une seule valeur de longueur.

Exemple de Génération Rembourrée à Droite

Entrée (rembourrée à droite) → Génération → Finale (rembourrée à droite) :

[101, 2054, 2003, 0, 0] # Entrée originale (longueur 3)
↓
[101, 2054, 2003, 2001, 1996, 4568, 7899, 0] # Après génération
|-- entrée --| |----- génération -----| |rembourrage|

Logprobs correspondants :

[ 0, 0, 0, -1.2, -0.8, -1.5, -2.1, 0]
|-- zéros pour entrée --| |- logprobs génération -| |rembourrage|

Vérifier le Rembourrage à Droite

NeMo RL fournit des utilitaires pour vérifier le rembourrage correct. Par exemple :

import torch
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
from nemo_rl.models.generation.interfaces import verify_right_padding
# Pour les données d'entrée (BatchedDataDict contenant input_ids et input_lengths)
input_data = BatchedDataDict({
"input_ids": torch.tensor([
[101, 2054, 2003, 0, 0], # Exemple de séquence d'entrée
[101, 2054, 0, 0, 0] # Autre séquence d'entrée
]),
"input_lengths": torch.tensor([3, 2]) # Longueur de chaque séquence
})
# Vérifier si les données d'entrée sont correctement rembourrées à droite
is_right_padded, error_msg = verify_right_padding(input_data, pad_value=0)
# Pour les données de sortie de génération (BatchedDataDict contenant output_ids et generation_lengths)
output_data = BatchedDataDict({
"output_ids": torch.tensor([
[101, 2054, 2003, 2001, 1996, 0, 0], # Exemple de séquence de sortie
[101, 2054, 2001, 4568, 0, 0, 0] # Autre séquence de sortie
]),
"generation_lengths": torch.tensor([2, 2]), # Longueur de la réponse générée
"unpadded_sequence_lengths": torch.tensor([5, 4]) # Nombre total de jetons valides
})
# Vérifier si les données de sortie sont correctement rembourrées à droite
is_right_padded, error_msg = verify_right_padding(output_data, pad_value=0)
if not is_right_padded:
print(f"Erreur de rembourrage : {error_msg}")
:hide:

La fonction verify_right_padding() vérifie que :

  1. Tout rembourrage (zéros ou jeton de rembourrage fourni par l’utilisateur) apparaît après les jetons valides.
  2. Le rembourrage commence à la position spécifiée par le tenseur de longueur.

La fonction détecte automatiquement si vous passez des données d’entrée ou de sortie :

  • Pour les données d’entrée : Nécessite les champs input_ids et input_lengths.
  • Pour les données de sortie : Nécessite output_ids et soit generation_lengths, soit unpadded_sequence_lengths.

Bonnes Pratiques

  1. Toujours Utiliser le Rembourrage à Droite : Tous les composants s’attendent à ce format.

  2. Suivre les Tenseurs de Longueur : Inclure les tenseurs de longueur appropriés avec vos données.

  3. Vérifier le Rembourrage : Utiliser verify_right_padding() en cas de doute.

  4. Masquer le Rembourrage dans les Opérations : Utiliser les longueurs pour exclure les jetons de rembourrage des calculs de perte.