Fonctions de perte dans NeMo RL
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,
Mais si l’on s’entraîne avec deux micro-lots,
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 à :