Passer à la navigation

nemo_rl.algorithms.utils

Afficher en Markdown

Contenu du Module

Fonctions

NomDescription
calculate_kl_penalty_joschu2020Calcule une estimation par jeton de la divergence KL entre deux log_probs.
calculate_baseline_and_std_per_promptFonction pour calculer une référence pour chaque paire (prompt, réponse) du lot.
surpress_user_warningsAucun
masked_meanCalcule la moyenne d’un micro-lot, en utilisant une statistique globale comme facteur de normalisation.
set_seedDéfinit la graine pour python, numpy et pytorch.
get_tokenizerObtenir le tokenizer et définir le jeton de remplissage sur le jeton de fin de séquence s’il n’est pas déjà défini.
maybe_pad_last_batchRemplit le lot donné de sorte que sa taille soit divisible par (mbs * dp_size).

API

nemo_rl.algorithms.utils.calculate_kl_penalty_joschu2020(
logprobs_policy: torch.Tensor, logprobs_reference: torch.Tensor
) -> torch.Tensor

Calcule une estimation par jeton de la divergence KL entre deux log_probs.

D’après Schulman 2020, toujours positif.

logprobs_policy: torch.Tensor (b, s) logprobs_reference: torch.Tensor (b, s)

nemo_rl.algorithms.utils.calculate_baseline_and_std_per_prompt(
prompts: torch.Tensor,
rewards: torch.Tensor,
valid_mask: torch.Tensor,
leave_one_out_baseline: bool = True
) -> tuple[torch.Tensor, torch.Tensor]

Fonction pour calculer une référence pour chaque paire (prompt, réponse) du lot.

La même référence est calculée pour chaque prompt. Les échantillons définis à 0 dans ‘valid_mask’ ne sont pas inclus dans le calcul de la référence.

prompts: tensor (b, s) Tenseur des prompts utilisés par le modèle. Peut être sur n’importe quel appareil rewards: tensor (b,) Récompenses de type flottant. Peut être sur n’importe quel appareil valid_mask: tensor (b,) Vecteur de 0/1, où 0 est à ignorer et 1 est à conserver leave_one_out_baseline: bool Calculer une référence non biaisée en excluant l’échantillon pour lequel la référence est calculée (de RLOO https://arxiv.org/abs/2402.14740)

Retourne : tensor (b,), tensor (b,) de références et d’écart-type sur le même appareil que ‘rewards’

nemo_rl.algorithms.utils.surpress_user_warnings(f)
nemo_rl.algorithms.utils.masked_mean(
values: torch.Tensor,
mask: torch.Tensor,
dim: typing.Optional[int] = None,
global_normalization_factor: typing.Optional[torch.Tensor | float] = None
)

Calcule la moyenne d’un micro-lot, en utilisant une statistique globale comme facteur de normalisation.

nemo_rl.algorithms.utils.set_seed(seed: int) -> None

Définit la graine pour python, numpy et pytorch.

nemo_rl.algorithms.utils.get_tokenizer(
tokenizer_config: nemo_rl.models.policy.TokenizerConfig,
get_processor: bool = False
) -> transformers.PreTrainedTokenizerBase

Obtenir le tokenizer et définir le jeton de remplissage sur le jeton de fin de séquence s’il n’est pas déjà défini.

Cette fonction initialise un tokenizer de la bibliothèque Hugging Face transformers et le configure avec des modèles de chat et des jetons de remplissage appropriés.

Args : tokenizer_config : Un dictionnaire contenant la configuration du tokenizer. Clés requises :

  • name : Le nom ou le chemin du tokenizer préentraîné Clés optionnelles :
  • chat_template : Le modèle de chat à utiliser. Peut être :
    • None : Utilise un modèle de passage qui renvoie simplement le contenu du message
    • “default” : Utilise le modèle par défaut du tokenizer
    • Une chaîne de modèle jinja2 personnalisée Si non spécifié, le modèle par défaut du tokenizer sera utilisé. get_processor : Indique s’il faut renvoyer un processeur (via AutoProcessor) au lieu d’un tokenizer.

Retourne : PreTrainedTokenizerBase : L’instance de tokenizer configurée

Exemples :

>>> from transformers import AutoTokenizer
>>> from nemo_rl.algorithms.utils import get_tokenizer
>>> # ne pas spécifier de modèle de chat utilise le modèle par défaut du tokenizer
>>> config = {"name": "meta-llama/Llama-3.2-1B-Instruct"}
>>> tokenizer = get_tokenizer(config)
Aucun modèle de chat fourni, utilisation du modèle par défaut du tokenizer
>>> messages = [
... {"role": "system", "content": "Vous êtes un assistant IA utile."},
... {"role": "user", "content": "Bonjour !"}
... ]
>>> formatted = tokenizer.apply_chat_template(messages, tokenize=False)
>>> assert formatted == AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B-Instruct").apply_chat_template(messages, tokenize=False)
>>> # Utilisation d'un modèle de passage
>>> config = {
... "name": "meta-llama/Llama-3.2-1B-Instruct",
... "chat_template": None
... }
>>> tokenizer = get_tokenizer(config)
Utilisation du modèle de chat de passage
>>> formatted = tokenizer.apply_chat_template(messages, tokenize=False)
>>> assert formatted == "".join(msg["content"] for msg in messages)
>>> # Utilisation d'un modèle personnalisé
>>> config = {
... "name": "meta-llama/Llama-3.2-1B-Instruct",
... "chat_template": "{% for message in messages %}{{ ' DÉBUT : ' + message['content'] + ' FIN.' }}{% endfor %}"
... }
>>> tokenizer = get_tokenizer(config)
Utilisation du modèle de chat personnalisé
>>> formatted = tokenizer.apply_chat_template(messages, tokenize=False)
>>> assert formatted == " DÉBUT : Vous êtes un assistant IA utile. FIN. DÉBUT : Bonjour ! FIN."
>>> # Demande d'un processeur (pour les modèles multimodaux comme Qwen-VL)
>>> config = {"name": "Qwen/Qwen2.5-VL-3B-Instruct"}
>>> processor = get_tokenizer(config, get_processor=True)
Aucun modèle de chat fourni, utilisation du modèle par défaut du tokenizer
>>> messages = [
... {"role": "system", "content": "Vous êtes un assistant IA utile."},
... {"role": "user", "content": "Bonjour !"}
... ]
>>> formatted = processor.tokenizer.apply_chat_template(messages, tokenize=False)
>>> assert formatted == AutoTokenizer.from_pretrained(
... "Qwen/Qwen2.5-VL-3B-Instruct", trust_remote_code=True
... ).apply_chat_template(messages, tokenize=False)
>>> assert processor.pad_token_id == processor.tokenizer.pad_token_id
>>>
nemo_rl.algorithms.utils.maybe_pad_last_batch(
batch: dict,
dp_size: int,
mbs: int
) -> dict

Remplit le lot donné de sorte que sa taille soit divisible par (mbs * dp_size).

Args : batch (dict) : Le lot à remplir. dp_size (int) : Taille du parallélisme des données. mbs (int) : Taille du micro-lot.

Retourne : dict : Le lot rempli.