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_dictdunn.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.
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.ptsert à reprendre après interruption.meilleur.ptsert à évaluer, exporter, servir. C'est l'équivalent PyTorch durestore_best_weights=Truede Keras.
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.
- Épingler la version de PyTorch utilisée pour un entraînement de
production dans un fichier de dépendances (
requirements.txt,pyproject.toml). Untorch>=2.0mène tôt ou tard à un checkpoint irrelisable dans deux ans. - É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
scaleren 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_locationest 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.