Systèmes de pré-formation et optimisation

Relier objectif causal, entropie croisée, rétropropagation, optimiseur, précision mixte, parallélisme et checkpoints.

Applied AI · advanced · Session 11

Contrat de la séance

  • Calculer une perte d’entropie croisée.
  • Suivre gradient, optimiseur et mise à jour.
  • Expliquer la reprise fiable depuis un checkpoint.

Carte du mécanisme

  lot (B,T) ──▶ ┌──────────────┐ logits z (B,T,V)
                │  passe avant │────────────────┐
                └──────┬───────┘                ▼
                       │ activations     ┌────────────────────┐
                       │ sauvegardées    │ CE = −z_y + logΣe^z│  L = 0,408
                       │ (mémoire)       └─────────┬──────────┘
                       ▼                           │
                ┌──────────────┐  règle de chaîne  │
                │ rétropropag. │◀──────────────────┘
                └──────┬───────┘  g = ∂L/∂θ
                       ▼
                ┌──────────────┐  m, v par paramètre
                │    AdamW     │  θ ← θ − lr·m̂/(√v̂+ε) − lr·λ·θ
                └──────┬───────┘  calcul BF16 │ maître FP32
                       ▼
 CHECKPOINT = θ + (m,v) + scheduler + scaler + position données + états RNG
              └─ poids seuls = reprise APPROXIMATIVE, pas exacte ─┘

1. Objectif causal

Le modèle maximise la vraisemblance du prochain token à chaque position autorisée par le masque causal. La perte moyenne agrège les positions et les exemples valides.

L = −Σ log p(x_t | x_<t)

Contrôle — Objectif causal

L = −Σ log p(x_t | x_<t) sur un lot B = 2, T = 4 dont 3 positions sont padées. Sur combien de termes divisez-vous, et que devient la courbe si vous divisez par 8 ?

2. Entropie croisée depuis les logits

La log-softmax stabilisée soustrait le log-sum-exp. La perte choisit ensuite le log-probabilité de la cible. Des logits plus grands ne sont utiles que relativement aux autres.

CE(z,y)=−z_y+log Σ exp(z_j)

Contrôle — Entropie croisée depuis les logits

CE(z,y) = −z_y + log Σ exp(z_j). Ajoutez +10 à chacun des logits [2,1,0] et recalculez : que vaut la perte, et quelle propriété de la log-softmax venez-vous de démontrer ?

3. Rétropropagation

La règle de chaîne calcule comment chaque paramètre a contribué à la perte. Les activations sauvegardées coûtent de la mémoire ; le checkpointing d’activations échange du recalcul contre de la mémoire.

Contrôle — Rétropropagation

Le checkpointing d’activations libère la mémoire des activations sauvegardées entre la passe avant et la rétropropagation. Que payez-vous en échange, et sur quel axe mesurez-vous le gain net ?

4. Optimiseur

AdamW combine moments des gradients, taux d’apprentissage et décroissance des poids. Le clipping peut borner des gradients extrêmes, mais ne répare pas une donnée ou une architecture défectueuse.

Contrôle — Optimiseur

AdamW conserve m et v par paramètre. Pour 7 milliards de paramètres en FP32, chiffrez la mémoire des seuls états d’optimiseur, et dites pourquoi compter les poids seuls sous-estime le besoin d’un facteur important.

5. Précision et parallélisme

BF16 réduit la mémoire des tenseurs sans représenter tous les états en pleine précision. Le parallélisme de données réplique les poids ; tensor/pipeline parallel répartissent d’autres dimensions avec communication.

Contrôle — Précision et parallélisme

En BF16 le calcul est en 16 bits mais une copie maître reste en FP32. Pourquoi ne pas tout passer en BF16 ? Nommez ce qui se dégrade en premier : la passe avant ou l’accumulation des mises à jour.

6. Checkpoint complet

Une reprise exacte exige poids, état de l’optimiseur, scheduler, scaler éventuel, position dans les données et états aléatoires. Un fichier de poids seul n’est pas un checkpoint d’entraînement complet.

Contrôle — Checkpoint complet

Vous reprenez l’entraînement à l’étape 12 000 depuis un fichier de poids seul. Citez trois champs manquants du checkpoint et l’effet observable de chacun sur la courbe de perte des 200 pas suivants.

Cas guidé — données

Pour logits [2,1,0] et cible 0, softmax ≈ [0,665;0,245;0,090], donc CE ≈ 0,408. Le laboratoire modifie taux, gradient et poids, puis montre les champs nécessaires à une reprise.

Cas guidé — trace complète

logits z = [2, 1, 0], cible y = 0

exp : e² = 7,389   e¹ = 2,718   e⁰ = 1,000   Σ = 11,107
softmax        ≈ [0,665 ; 0,245 ; 0,090]     (somme = 1,000 ✅)
CE = −z_y + log Σ e^z = −2 + log(11,107) = −2 + 2,408 = 0,408   ✅

INVARIANCE PAR DÉCALAGE : z + 10 = [12, 11, 10]
  −12 + log(e¹² + e¹¹ + e¹⁰) = −12 + 12,408 = 0,408   ✅ perte identique

SOFTMAX NAÏVE, z = [1000, 999, 998]
  exp(1000) → inf ;  inf/inf → NaN                     ❌ perte détruite
  version stable : soustraire max(z) → [0,−1,−2] → 0,408 ✅

REPRISE DEPUIS POIDS SEULS (étape 12 000)
  m = 0, v = 0 → le premier pas AdamW est mal calibré
  lr relancé au début du scheduler → pic de perte transitoire
  (p. ex. 0,4 → ~1,9 dans le laboratoire), puis retour  ❌

CONTRÔLE DE FORME : z (B,T,V) et y (B,T). La moyenne divise par le nombre de
positions NON masquées, pas par B×T.

Leviers mémoire/débit en pré-formation : ce que chacun coûte

Levier Ce qu’il gagne Ce qu’il coûte réellement
BF16 (calcul) Tenseurs deux fois plus légers, matmuls plus rapides Copie maître FP32 conservée ; accumulations sensibles
Checkpointing d’activations Mémoire d’activations fortement réduite Recalcul de la passe avant : ~30 % de temps en plus
Parallélisme de données Débit quasi linéaire en nombre de GPU Poids et états d’optimiseur répliqués sur chaque rang
Clipping de gradient Borne les pics de norme, évite les NaN Ne corrige ni une donnée sale ni une architecture instable

Laboratoire causal

Prédire → modifier une variable → exécuter → expliquer l’écart

/interactives/curriculum/optimizer-checkpoint.html?lang=fr

Erreur fréquente 1

« Des logits plus grands signifient une perte plus faible. »

La trace le réfute en une ligne : [2,1,0] et [12,11,10] donnent exactement 0,408. Seuls les écarts entre logits comptent ; l’échelle absolue disparaît dans le log-sum-exp.

Erreur fréquente 2

« Le fichier .safetensors est mon checkpoint d’entraînement. »

C’est un checkpoint d’inférence. Sans m, v, scheduler, scaler, position dans les données et états RNG, la reprise à l’étape 12 000 repart avec m = v = 0 et produit un pic de perte, plus une répétition ou un saut d’échantillons.

Frontière de validité

Les performances distribuées dépendent du matériel, du réseau, de la taille du modèle et de l’implémentation ; aucune estimation du laboratoire n’est un benchmark.

Statut de preuve

Mécanismes établis ; les simplifications numériques sont pédagogiques.

Sources à vérifier

  • Kingma & Ba, “Adam: A Method for Stochastic Optimization”, ICLR (2015).
  • Loshchilov & Hutter, “Decoupled Weight Decay Regularization” (AdamW), ICLR (2019).
  • Micikevicius et al., “Mixed Precision Training”, ICLR (2018).
  • Dossier source fourni par le propriétaire; les détails sur des produits nommés restent attribués à cette source jusqu’à vérification primaire.

Défi de transfert

Votre run de 7 B s’arrête à l’étape 12 000. Deux reprises possibles : A) fichier de poids seul, au bon pas ; B) checkpoint complet, mais vieux de 2 000 pas.

  1. Prédisez la courbe de perte des 200 premiers pas pour chaque option.
  2. Chiffrez le coût de chaque option (taille du checkpoint complet ; pas à rejouer).
  3. Choisissez, et fixez le seuil qui vous ferait changer d’avis.

Synthèse

  • Objectif causal
  • Entropie croisée depuis les logits
  • Rétropropagation
  • Optimiseur
  • Précision et parallélisme
  • Checkpoint complet

Ticket de sortie

Mécanisme · trace · observation · frontière · preuve · prochaine expérience

Notes formateur: Poser le problème avant de nommer le mécanisme. Recueillir une prédiction initiale et la conserver pour le ticket de sortie.

Notes formateur: Ces objectifs sont observables : trace, calcul, comparaison. Une définition récitée ne clôt aucun objectif.

Notes formateur: Faire relier chaque étape à la suivante par un verbe causal. Signaler toute flèche purement décorative.

Notes formateur: Commencer par faire écrire les formes (B,T,V) et (B,T) au tableau, puis demander le dénominateur de la moyenne. Ceux qui répondent B×T tiennent le bug de comptage le plus fréquent en production — le laisser vivre une minute.

Notes formateur: Réponse : diviser par 5 — les positions non masquées — jamais par 8. Diviser par 8 aplatit la courbe de façon optimiste sans qu’aucune prédiction ne s’améliore. Erreur attendue : B×T = 8. Rappel utile : même bug de dénominateur qu’en session intermédiaire 11 (18 contre 13).

Notes formateur: Faire calculer 0,408 à la main puis relancer avec z + 10 sans annoncer le résultat attendu. La surprise de l’égalité vaut mieux qu’une démonstration algébrique donnée d’avance.

Notes formateur: Réponse : la perte reste exactement 0,408 — invariance par décalage de la log-softmax, seuls les écarts entre logits comptent. Erreur attendue : « des logits plus grands, donc perte plus faible ». Faire refaire la ligne de trace −12 + 12,408 avant de corriger.

Notes formateur: Faire dessiner la mémoire d’activations comme une pile qui grandit pendant la passe avant. Demander où couper la pile : le checkpointing devient un choix qu’ils font, pas une option de bibliothèque.

Notes formateur: Réponse : on paie un recalcul de la passe avant, environ 30 % de temps en plus ; le gain net se mesure en mémoire d’activations libérée, donc en taille de lot ou de modèle rendue possible. Erreur attendue : « c’est gratuit puisque la mémoire baisse ».

Notes formateur: Faire chiffrer en salle le budget mémoire complet de 7 B en AdamW avant de montrer le tableau. L’écart entre leur estimation et le total réel est le vrai contenu de cette diapositive.

Notes formateur: Réponse : m + v en FP32 = 2 × 7 × 10⁹ × 4 octets ≈ 56 Go — deux fois les poids FP32 (28 Go) ; poids + états ≈ 84 Go, soit 3× les poids seuls. Erreurs attendues : compter un seul moment, ou oublier les 4 octets du FP32. L’écart entre leur estimation et 56 Go est le contenu du slide.

Notes formateur: Poser le parallélisme comme un arbitrage de communication : demander ce qui traverse le réseau à chaque pas dans chacun des trois schémas. Refuser toute réponse en « c’est plus rapide ».

Notes formateur: Réponse : l’accumulation des mises à jour se dégrade en premier — les petits incréments disparaissent dans la mantisse courte du BF16 ; la copie maître FP32 les préserve. Erreur attendue : « la passe avant devient fausse » — elle tolère bien mieux la précision réduite.

Notes formateur: Distribuer une liste de champs de checkpoint dont trois manquent et faire diagnostiquer le symptôme attendu pour chacun. C’est l’exercice qui transfère le mieux vers un incident réel.

Notes formateur: Réponse — trois champs et leurs symptômes : m,v absents → premier pas AdamW mal calibré ; scheduler absent → lr relancé, pic de perte transitoire ; position de données absente → échantillons répétés ou sautés ; états RNG absents → run non reproductible. Exiger l’effet observable sur 200 pas pour chaque champ cité.

Notes formateur: Masquer le résultat final. Faire annoncer le signe, la forme et l’ordre de grandeur avant chaque opération.

Notes formateur: Dérouler ligne par ligne. Une incohérence se localise à la première étape fautive, pas seulement sur la dernière ligne.

Notes formateur: Faire remplir la dernière ligne par les apprenants avant de la révéler : c’est le compromis qui décide en production.

Notes formateur: Conserver les valeurs initiales et finales. Interdire les changements simultanés qui rendent l’écart impossible à attribuer.

Notes formateur: Faire produire le contre-exemple minimal par le groupe avant de donner la correction.

Notes formateur: Faire produire le contre-exemple minimal par le groupe avant de donner la correction.

Notes formateur: La limite n’est pas une note de bas de page : elle définit les cas où le mécanisme ne suffit plus.

Notes formateur: Séparer mécanisme vérifiable, choix d’implémentation rapporté et résultat expérimental. La précision de la preuve doit suivre celle de l’affirmation.

Notes formateur: Attribuer chaque choix de produit au dossier fourni et conserver le statut rapporté tant qu’aucune source primaire indépendante ne le confirme.

Notes formateur: Attendus : A repart au bon pas mais m = v = 0 et scheduler relancé → pic transitoire puis récupération ; B rejoue 2 000 pas (coût GPU pur) avec une courbe saine. Chiffrage : checkpoint complet ≈ 28 Go de poids FP32 + 56 Go de moments ; B coûte 2 000 × coût/pas. Idée fausse à récolter : « les poids suffisent, l’optimiseur se recalibre vite » — vrai seulement à petit lr. Basculer vers B dès que le pic dépasse ~2× la perte courante. 8 minutes, binômes.

Notes formateur: Faire reconstruire la chaîne sans regarder les slides. Réouvrir uniquement le premier point de rupture.

Notes formateur: Six lignes maximum. Comparer à la prédiction initiale et nommer ce qui a réellement changé. Citer une trace conservée qui permet à un pair de vérifier la conclusion.