Passer à la navigation

Fonctions de perte dans NeMo RL

Afficher en Markdown

Les fonctions de perte dans NeMo RL sont spécialement conçues pour garantir que l’entraînement par lot complet est équivalent à l’entraînement avec accumulation de gradient. Pour comprendre pourquoi un soin particulier doit être apporté ici, considérez l’exemple suivant d’une fonction de perte simple qui prend la moyenne de quelques pertes par jeton sur tous les jetons dans un micro-lot, puis moyenne la perte sur les micro-lots.

Supposons que nous ayons un lot global avec 16 jetons non masqués. Les 10 premiers jetons non masqués proviennent de la première moitié des échantillons du lot, et les 6 derniers proviennent de la seconde moitié. Si l’on s’entraîne avec un lot global,

L=∑t=116Lt16.L = \frac{\sum_{t=1}^{16} L_t}{16}.

Mais si l’on s’entraîne avec deux micro-lots,

L=∑t=110Lt10+∑t=1016Lt62,L = \frac{\frac{\sum_{t=1}^{10} L_t}{10} + \frac{\sum_{t={10}}^{16} L_t}{6}}{2},

ce qui n’est généralement pas équivalent à la perte du lot complet. Pour résoudre ce problème, nous devons faire en sorte que chaque micro-lot ait des informations sur le nombre de jetons dans les autres micro-lots du lot global.

Dans NeMo RL, ces informations sont transmises directement à la fonction de perte. Chaque fonction de perte devrait appartenir à l’une des deux catégories, au niveau du jeton ou au niveau de la séquence, qui est un attribut de la fonction de perte elle-même (voir loss_functions.py pour quelques exemples). La politique utilise ensuite ces informations pour calculer le facteur de normalisation global en utilisant le lot complet (pour les pertes au niveau du jeton, c’est le nombre total de jetons dans le lot. Pour les pertes au niveau de la séquence, c’est le nombre de séquences valides dans le lot). Le facteur de normalisation est ensuite transmis à la fonction de perte, qui l’utilise pour normaliser la perte du micro-lot. Pour obtenir la perte du lot global, la politique additionne simplement les pertes de tous les micro-lots.

Pour notre exemple simple ci-dessus, cela ressemblerait à :

import torch
from nemo_rl.algorithms.interfaces import LossFunction
from nemo_rl.algorithms.loss_functions import LossType
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
class SimpleAverageLoss(LossFunction):
"""Fonction de perte moyenne simple qui démontre la gestion correcte des micro-lots.
REMARQUE : Pour simplifier, nous supposons que les pertes par jeton sont transmises directement à cette fonction de perte.
Ce n'est pas le cas en pratique !
"""
loss_type = LossType.TOKEN_LEVEL
def __call__(
self,
next_token_losses: torch.Tensor,
data: BatchedDataDict,
total_valid_tokens_or_seqs: torch.Tensor,
) -> tuple[torch.Tensor, dict]:
"""Calculer la perte moyenne simple avec une gestion correcte des micro-lots."""
token_mask = data["token_mask"] ## masque de jeton pour ce micro-lot
sample_mask = data["sample_mask"] ## masque d'échantillon pour ce micro-lot
# mask.sum() sera 10 pour le micro-lot 1, 6 pour le micro-lot 2
mask = token_mask * sample_mask.unsqueeze(-1)
# total_valid_tokens_or_seqs sera 16 dans notre exemple car il y a 16 jetons dans le lot global
# comme nous avons spécifié qu'il s'agit d'une perte au niveau du jeton, la politique
# nous donnera automatiquement le bon facteur de normalisation.
loss = (next_token_losses * mask).sum() / (total_valid_tokens_or_seqs + 1e-8)
return loss
## tester la fonction de perte
import torch
## dans cet exemple, nous avons un lot de taille 2 avec une longueur de séquence de 16
batch_size = 2
seq_len = 16
next_token_losses = torch.randn((batch_size, seq_len))
sample_data = {
"token_mask": torch.tensor(
[
[1] * 10 + [0] * 6,
[1] * 6 + [0] * 10,
]
),
"sample_mask": torch.ones(2)
}
total_valid_tokens_or_seqs = torch.sum(sample_data["token_mask"] * sample_data["sample_mask"].unsqueeze(-1))
loss_fn = SimpleAverageLoss()
loss_no_microbatching = loss_fn(next_token_losses, sample_data, total_valid_tokens_or_seqs)
microbatch_1_data = {
"token_mask": sample_data["token_mask"][:1],
"sample_mask": sample_data["sample_mask"][:1],
}
microbatch_2_data = {
"token_mask": sample_data["token_mask"][1:],
"sample_mask": sample_data["sample_mask"][1:],
}
loss_with_microbatching = (
loss_fn(next_token_losses[:1], microbatch_1_data, total_valid_tokens_or_seqs)
+ loss_fn(next_token_losses[1:], microbatch_2_data, total_valid_tokens_or_seqs)
)
torch.testing.assert_close(loss_no_microbatching, loss_with_microbatching)
:hide: