Passer à la navigation

nemo_rl.environments.metrics

Afficher en Markdown

Contenu du module

Fonctions

NomDescription
calculate_pass_rate_per_promptFonction pour calculer la fraction de prompts ayant au moins une réponse correcte (récompense > 0).

API

nemo_rl.environments.metrics.calculate_pass_rate_per_prompt(
prompts: torch.Tensor, is_correct: torch.Tensor
) -> float

Fonction pour calculer la fraction de prompts ayant au moins une réponse correcte (récompense > 0).

prompts: tensor (b, s) Tenseur des prompts utilisés par le modèle. Peut être sur n’importe quel appareil is_correct: tensor (b,) étiquette booléenne. Peut être sur n’importe quel appareil

Retourne : pass rate : float