Passer à la navigation

Entraînement du Modèle de Récompense dans NeMo RL

Afficher en Markdown

Ce document explique comment entraîner des modèles de récompense (RM) dans NeMo RL. Actuellement, seuls les modèles de récompense Bradley-Terry sont pris en charge sur le backend DTensor. Le support du backend Megatron est suivi ici.

Lancer un Travail d’Entraînement

Le script, examples/run_rm.py, est utilisé pour entraîner un modèle de récompense Bradley-Terry. Ce script peut être lancé localement ou via Slurm. Pour plus de détails sur la configuration de Ray et le lancement d’un travail avec Slurm, reportez-vous à la documentation du cluster.

Veillez à lancer le travail en utilisant uv. La commande pour lancer un travail d’entraînement est la suivante :

uv run examples/run_rm.py
# Vous pouvez également ajouter des remplacements en CLI, comme modifier la configuration ou le modèle
uv run examples/run_rm.py --config examples/configs/rm.yaml policy.model_name=Qwen/Qwen2.5-1.5B

La configuration YAML par défaut partage le même modèle de base que la configuration SFT mais inclut une nouvelle section reward_model_cfg avec enabled: true pour charger le modèle en tant que Modèle de Récompense. Vous pouvez trouver un exemple de fichier de configuration RM à examples/configs/rm.yaml.

Rappel : Définissez votre HF_HOME, WANDB_API_KEY, et HF_DATASETS_CACHE (si nécessaire). Assurez-vous de vous connecter en utilisant huggingface-cli si vous travaillez avec des modèles Llama.

Jeux de Données

Chaque classe de jeu de données RM est censée avoir les attributs suivants :

  1. formatted_ds : Le dictionnaire des jeux de données formatés, où chaque jeu de données doit être formaté comme
{
"context": [], // liste de dictionnaires - Le message de prompt (incluant les tours précédents, le cas échéant)
"completions": [ // liste de dictionnaires — La liste des complétions
{
"rank": 0, // entier — Le rang de la complétion (rang inférieur est préféré)
"completion": [] // liste de dictionnaires — Le(s) message(s) de complétion
},
{
"rank": 1, // entier — Le rang de la complétion (rang inférieur est préféré)
"completion": [] // liste de dictionnaires — Le(s) message(s) de complétion
}
]
}
  1. task_spec : Le TaskDataSpec pour ce jeu de données. Cela doit spécifier le nom que vous choisissez pour ce jeu de données.

Actuellement, l’entraînement RM ne prend en charge que deux complétions (où le rang le plus bas est préféré et le plus élevé est rejeté), chaque complétion étant une seule réponse. Par exemple :

{
"context": [
{
"role": "user",
"content": "Quelle est la capitale de la France ?"
},
{
"role": "assistant",
"content": "La capitale de la France est Paris."
},
{
"role": "user",
"content": "Merci ! Et quelle est la capitale de l'Allemagne ?"
}
],
"completions": [
{
"rank": 0,
"completion": [
{
"role": "assistant",
"content": "La capitale de l'Allemagne est Berlin."
}
]
},
{
"rank": 1,
"completion": [
{
"role": "assistant",
"content": "La capitale de l'Allemagne est Munich."
}
]
}
]
}

NeMo RL fournit une implémentation compatible avec RM du jeu de données HelpSteer3 comme exemple. Ce jeu de données est téléchargé depuis Hugging Face et prétraité à la volée, il n’est donc pas nécessaire de fournir un chemin vers des jeux de données sur disque.

Nous fournissons également une classe PreferenceDataset compatible avec les jeux de données de préférence au format JSONL. Vous pouvez modifier votre configuration comme suit pour utiliser un jeu de données de préférence personnalisé :

data:
dataset_name: PreferenceDataset
train_data_path: <CheminLocalVersJeuDeDonnéesEntraînement>
val_data_paths:
<NomDuJeuDeDonnéesValidation>: <CheminLocalVersJeuDeDonnéesValidation>

avec la prise en charge de plusieurs jeux de validation :

data:
dataset_name: PreferenceDataset
train_data_path: <CheminLocalVersJeuDeDonnéesEntraînement>
val_data_paths:
<NomDuJeuDeDonnéesValidation1>: <CheminLocalVersJeuDeDonnéesValidation1>
<NomDuJeuDeDonnéesValidation2>: <CheminLocalVersJeuDeDonnéesValidation2>

Veuillez noter :

  • Si vous utilisez un logger, le préfixe utilisé pour chaque jeu de validation sera validation-<NomDuJeuDeDonnéesValidation>. Le temps de validation total, sommé sur tous les jeux de validation, est rapporté sous timing/validation/total_validation_time.
  • Si vous effectuez des points de contrôle, la valeur metric_name dans votre configuration de points de contrôle doit refléter la métrique et le jeu de validation à suivre. Par exemple, validation-<NomDuJeuDeDonnéesValidation1>_loss.