Aller au contenu principal

Module 8 — Points de contrôle et reprise d'entraînement

Un entraînement peut durer des heures ou des jours. Une coupure de courant, un plantage du pilote GPU, un noyau relancé par l'OS suffisent à tout perdre. Ce module met en place la seule protection réaliste contre ces incidents : des points de contrôle écrits à intervalles réguliers et un mécanisme de reprise qui redémarre exactement là où on s'était arrêté.

Ce qu'un point de contrôle doit contenir

Un point de contrôle ne se limite pas au modèle. Reprendre à l'identique demande cinq éléments.

  • L'état du modèle (state_dict du nn.Module).
  • L'état de l'optimiseur (moments de Adam, vitesses de SGD).
  • L'état du planificateur de taux (compteur d'époques, dernier lr).
  • Le numéro de l'époque atteinte, indispensable pour redémarrer la boucle à la bonne itération.
  • L'état des générateurs pseudo-aléatoires (Python, NumPy, PyTorch) si l'on veut une reprise strictement reproductible.

En pratique, on ajoute aussi les métriques observées pour ne pas les recalculer, et le scaler.state_dict() quand on utilise la précision mixte du module 7.

Sauver au bon format

torch.save sérialise un dictionnaire arbitraire en un fichier binaire. On y met state_dict, jamais l'objet Python entier.

import torch

def sauver_checkpoint(chemin, modele, optimiseur, planificateur, epoque, meilleure_val):
torch.save({
"epoque": epoque,
"modele_state": modele.state_dict(),
"optimiseur_state": optimiseur.state_dict(),
"planificateur_state": planificateur.state_dict(),
"meilleure_val": meilleure_val,
"torch_rng": torch.get_rng_state(),
"cuda_rng": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,
}, chemin)

L'extension conventionnelle est .pt ou .pth. Le nom du fichier indique en général l'époque ou le rôle : checkpoint_dernier.pt, meilleur.pt.

Ne jamais sauver l'objet Python

torch.save(modele, ...) fonctionne, mais crée une dépendance au chemin exact de la classe (__main__.ReseauFashionMNIST). Déplacer le fichier .py, renommer un dossier, importer depuis un module différent cassent le chargement. Sauver state_dict uniquement.

Charger avec map_location

Un point de contrôle enregistre le device d'origine de chaque tenseur. Charger un state_dict sauvé sur GPU dans un environnement sans GPU échoue si l'on n'indique pas où placer les tenseurs.

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
checkpoint = torch.load("checkpoint_dernier.pt", map_location=device)

map_location=device remappe systématiquement les tenseurs vers le device cible. C'est la seule ligne qui permet de reprendre sur un serveur différent, ou de servir sur CPU un modèle entraîné sur GPU.

La reprise complète, cinq étapes

def reprendre_entrainement(chemin, modele, optimiseur, planificateur, device):
checkpoint = torch.load(chemin, map_location=device)

modele.load_state_dict(checkpoint["modele_state"])
optimiseur.load_state_dict(checkpoint["optimiseur_state"])
planificateur.load_state_dict(checkpoint["planificateur_state"])

torch.set_rng_state(checkpoint["torch_rng"])
if checkpoint["cuda_rng"] is not None and torch.cuda.is_available():
torch.cuda.set_rng_state_all(checkpoint["cuda_rng"])

epoque_depart = checkpoint["epoque"] + 1 # reprendre à la suivante
meilleure_val = checkpoint["meilleure_val"]
return epoque_depart, meilleure_val

epoque + 1 évite de rejouer l'époque déjà sauvegardée. On l'oublie souvent : la boucle redémarre alors à l'époque déjà entraînée et ré-écrase le meilleur checkpoint avec une version dégradée du même état.

Meilleure époque contre dernière époque : deux fichiers distincts

Le point de contrôle « meilleur » n'est pas nécessairement le dernier. En cas de surapprentissage, la meilleure exactitude de validation survient au milieu de l'entraînement, puis la perte remonte pendant que la perte d'entraînement continue à baisser.

def entrainer_avec_checkpoints(modele, train_loader, val_loader, dossier_ckpt,
nb_epoques=15, epoque_depart=0, meilleure_val=float("inf")):
from pathlib import Path
Path(dossier_ckpt).mkdir(parents=True, exist_ok=True)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
modele = modele.to(device)
criterion = torch.nn.CrossEntropyLoss()
optimiseur = torch.optim.AdamW(modele.parameters(), lr=1e-3)
planificateur = torch.optim.lr_scheduler.CosineAnnealingLR(optimiseur, T_max=nb_epoques)

for epoque in range(epoque_depart, nb_epoques):
entrainer_une_epoque(modele, train_loader, criterion, optimiseur)
perte_val, exactitude = evaluer(modele, val_loader, criterion)
planificateur.step()

# Dernier : écrasé à chaque époque
sauver_checkpoint(f"{dossier_ckpt}/dernier.pt",
modele, optimiseur, planificateur, epoque, meilleure_val)

# Meilleur : uniquement en cas d'amélioration
if perte_val < meilleure_val:
meilleure_val = perte_val
sauver_checkpoint(f"{dossier_ckpt}/meilleur.pt",
modele, optimiseur, planificateur, epoque, meilleure_val)

Deux fichiers, deux usages :

  • dernier.pt sert à reprendre après interruption.
  • meilleur.pt sert à évaluer, exporter, servir. C'est l'équivalent PyTorch du restore_best_weights=True de Keras.
Rotation des points de contrôle

Sur un long entraînement, garder aussi une copie toutes les N époques (checkpoint_epoque_05.pt, _10.pt, _15.pt) permet de revenir en arrière en cas de choix malheureux d'hyperparamètre ajusté en cours de route. Un checkpoint pèse rarement plus de quelques centaines de mégaoctets pour un modèle Fashion-MNIST ; sur un ResNet, c'est de l'ordre du gigaoctet.

Le point de contrôle et la précision mixte

Quand la boucle utilise GradScaler du module 7, il faut aussi le sauver et le recharger, sans quoi le facteur d'échelle repart à sa valeur initiale au redémarrage.

torch.save({
...,
"scaler_state": scaler.state_dict(),
}, chemin)

# ...
scaler.load_state_dict(checkpoint["scaler_state"])

L'oubli n'empêche pas la reprise, mais provoque quelques itérations « bruyantes » le temps que le scaler se cale à nouveau.

Compatibilité de versions

Un state_dict sauvegardé avec PyTorch 2.0 se recharge normalement avec PyTorch 2.4, l'inverse est rarement garanti. Deux règles.

  1. Épingler la version de PyTorch utilisée pour un entraînement de production dans un fichier de dépendances (requirements.txt, pyproject.toml). Un torch>=2.0 mène tôt ou tard à un checkpoint irrelisable dans deux ans.
  2. Écrire un petit test de rechargement en fin de sauvegarde : recharger le fichier tout juste écrit, comparer un lot d'entrée passe avant. Cela détecte immédiatement un état incompatible ou une corruption d'écriture.
def verifier_checkpoint(chemin, modele_reference, device, exemple):
modele_test = type(modele_reference)().to(device)
checkpoint = torch.load(chemin, map_location=device)
modele_test.load_state_dict(checkpoint["modele_state"])
modele_test.eval()
modele_reference.eval()
with torch.no_grad():
assert torch.allclose(modele_reference(exemple), modele_test(exemple), atol=1e-5)

L'arrêt anticipé, un cas particulier utile

Combiner « garder le meilleur » et arrêter quand la perte de validation stagne donne l'équivalent de Keras EarlyStopping.

def entrainer_avec_arret(modele, train_loader, val_loader, patience=5, nb_max=100):
meilleure_val, epoques_sans_progres = float("inf"), 0
for epoque in range(nb_max):
entrainer_une_epoque(modele, train_loader, criterion, optimiseur)
perte_val, _ = evaluer(modele, val_loader, criterion)
if perte_val < meilleure_val - 1e-4:
meilleure_val = perte_val
epoques_sans_progres = 0
sauver_checkpoint("meilleur.pt", modele, optimiseur, planificateur, epoque, meilleure_val)
else:
epoques_sans_progres += 1
if epoques_sans_progres >= patience:
print(f"Arrêt anticipé à l'époque {epoque}")
break

patience=5 signifie qu'on tolère cinq époques consécutives sans amélioration avant de renoncer. Un seuil de progression (-1e-4) évite qu'un bruit minuscule ne compte comme une amélioration.

En résumé

  • Un point de contrôle contient cinq états : modèle, optimiseur, planificateur, époque, générateurs pseudo-aléatoires ; ajouter scaler en précision mixte.
  • Sauver state_dict, jamais l'objet Python ; le second couple le fichier au chemin de la classe et casse au premier déplacement.
  • map_location est la clé pour recharger un checkpoint sur un matériel différent, y compris de GPU vers CPU.
  • Meilleur ≠ dernier : deux fichiers distincts pour deux usages — reprendre, ou évaluer et servir.

Le module suivant remplace le petit réseau maison par un ResNet18 préentraîné de torchvision, et introduit l'apprentissage par transfert qui atteint 92 % d'exactitude sur Fashion-MNIST en quelques minutes.