Passer à la navigation

nemo_rl.utils.checkpoint

Afficher en Markdown

Utilitaires de gestion des points de contrôle pour la boucle de l’algorithme RL.

Il gère la logique au niveau de l’algorithme. Chaque Acteur RL est censé avoir sa propre fonction de sauvegarde de point de contrôle (appelée par la boucle de l’algorithme).

Contenu du Module

Classes

NomDescription
CheckpointingConfigConfiguration pour la gestion des points de contrôle.
CheckpointManagerGère les points de contrôle du modèle pendant l’entraînement.

Fonctions

NomDescription
_load_checkpoint_historyCharge l’historique des points de contrôle et leurs métriques.

Données

PathLike

API

nemo_rl.utils.checkpoint.PathLike

Valeur: None

class nemo_rl.utils.checkpoint.CheckpointingConfig

Bases: typing.TypedDict

Configuration pour la gestion des points de contrôle.

Attributs : enabled (bool): Indique si les points de contrôle sont activés. checkpoint_dir (PathLike): Répertoire où les points de contrôle seront sauvegardés. metric_name (str | None): Nom de la métrique à utiliser pour déterminer les meilleurs points de contrôle. higher_is_better (bool): Indique si des valeurs plus élevées de la métrique indiquent de meilleures performances. keep_top_k (Optional[int]): Nombre de meilleurs points de contrôle à conserver. Si None, tous les points de contrôle sont conservés.

enabled: bool

Valeur: None

checkpoint_dir: nemo_rl.utils.checkpoint.PathLike

Valeur: None

metric_name: str | None

Valeur: None

higher_is_better: bool

Valeur: None

save_period: int

Valeur: None

keep_top_k: typing.NotRequired[int]

Valeur: None

checkpoint_must_save_by: typing.NotRequired[str | None]

Valeur: None

class nemo_rl.utils.checkpoint.CheckpointManager(config: nemo_rl.utils.checkpoint.CheckpointingConfig)

Gère les points de contrôle du modèle pendant l’entraînement.

Cette classe gère la création de répertoires de points de contrôle, la sauvegarde des informations d’entraînement et des configurations. Elle fournit également des utilitaires pour conserver uniquement les k meilleurs points de contrôle. La structure des points de contrôle ressemble à ceci :

checkpoint_dir/
step_0/
training_info.json
config.yaml
policy.py (à la charge de la boucle de l'algorithme de sauvegarder ici)
policy_optimizer.py (à la charge de la boucle de l'algorithme de sauvegarder ici)
...
step_1/
...

Attributs : Dérivés de CheckpointingConfig.

step: int,
training_info: typing.Mapping[str, typing.Any],
run_config: typing.Optional[typing.Mapping[str, typing.Any]] = None
) -> nemo_rl.utils.checkpoint.PathLike

Initialise un répertoire de point de contrôle temporaire.

Crée un répertoire temporaire pour un nouveau point de contrôle et sauvegarde les informations d’entraînement et la configuration. Le répertoire est nommé ‘tmp_step_{step}’ et sera renommé ‘step_{step}’ lorsque le point de contrôle sera terminé. Nous procédons ainsi pour permettre à la boucle de l’algorithme de sauvegarder les fichiers qu’elle souhaite dans un répertoire temporaire sûr.

Args : step (int) : Le numéro de l’étape d’entraînement. training_info (dict[str, Any]) : Dictionnaire contenant les métriques et informations d’entraînement. run_config (Optional[dict[str, Any]]) : Configuration optionnelle pour l’exécution de l’entraînement.

Returns : PathLike : Chemin vers le répertoire de point de contrôle temporaire.

finalize_checkpoint(checkpoint_path: nemo_rl.utils.checkpoint.PathLike) -> None

Termine un point de contrôle en le déplaçant de l’emplacement temporaire à l’emplacement permanent.

Si un point de contrôle à l’emplacement cible existe déjà (par exemple lors de la reprise de l’entraînement), nous remplaçons l’ancien. Déclenche également le nettoyage des anciens points de contrôle en fonction du paramètre keep_top_k.

Args : checkpoint_path (PathLike) : Chemin vers le répertoire de point de contrôle temporaire.

remove_old_checkpoints(exclude_latest: bool = True) -> None

Supprime les points de contrôle qui ne sont pas dans les k meilleurs ou les plus récents en fonction de la métrique (optionnelle).

Si keep_top_k est défini, cette méthode supprime tous les points de contrôle à l’exception des k meilleurs. Les points de contrôle “meilleurs” sont déterminés par :

  • Si une métrique est fournie : la valeur de la métrique donnée et le paramètre higher_is_better. Lorsque plusieurs points de contrôle ont la même valeur de métrique, les points de contrôle plus récents (numéros d’étapes plus élevés) sont privilégiés.
  • Si aucune métrique n’est fournie : le numéro de l’étape. Les k points de contrôle les plus récents sont conservés.

Args : exclude_latest (bool) : Indique s’il faut exclure le dernier point de contrôle de la suppression. (peut résulter en K+1 points de contrôle)

get_best_checkpoint_path() -> typing.Optional[str]

Obtient le chemin vers le meilleur point de contrôle en fonction de la métrique.

Renvoie le chemin vers le point de contrôle avec la meilleure valeur de métrique. Si aucun point de contrôle n’existe, renvoie None. Si la métrique n’est pas trouvée, un avertissement est émis et le dernier point de contrôle est renvoyé.

Returns : Optional[str] : Chemin vers le meilleur point de contrôle, ou None si aucun point de contrôle valide n’existe.

get_latest_checkpoint_path() -> typing.Optional[str]

Obtient le chemin vers le dernier point de contrôle.

Renvoie le chemin vers le point de contrôle avec le numéro d’étape le plus élevé.

Returns : Optional[str] : Chemin vers le dernier point de contrôle, ou None si aucun point de contrôle n’existe.

load_training_info(checkpoint_path: typing.Optional[nemo_rl.utils.checkpoint.PathLike] = None) -> typing.Optional[dict[str, typing.Any]]

Charge les informations d’entraînement à partir d’un point de contrôle.

Args : checkpoint_path (Optional[PathLike]) : Chemin vers le point de contrôle. Si None, renvoie None.

Returns : Optional[dict[str, Any]] : Dictionnaire contenant les informations d’entraînement, ou None si checkpoint_path est None.

nemo_rl.utils.checkpoint._load_checkpoint_history(checkpoint_dir: pathlib.Path) -> list[tuple[int, nemo_rl.utils.checkpoint.PathLike, dict[str, typing.Any]]]

Charge l’historique des points de contrôle et leurs métriques.

Args : checkpoint_dir (Path) : Répertoire contenant les points de contrôle.

Returns : list[tuple[int, PathLike, dict[str, Any]]] : Liste de tuples contenant (numéro_d’étape, chemin_du_point_de_contrôle, informations) pour chaque point de contrôle.