Chunking, causalité et préfill parallèle

Réconcilier état récurrent et parallélisme GPU par calcul en blocs causaux.

Applied AI · advanced · Session 15

Carte du mécanisme

PROMPT : 8 tokens, chunks de 4          PRÉFILL (parallèle)
┌──────── CHUNK 1 : t1..t4 ────────┐   ┌──────── CHUNK 2 : t5..t8 ────────┐
│  M = tril(Q Kᵀ)   4×4            │   │  M = tril(Q Kᵀ)   4×4            │
│    ┌1 . . .┐   « . » = masqué    │   │    ┌1 . . .┐                     │
│    │0 1 . .│     (futur)         │   │    │0 1 . .│                     │
│    │1 0 1 .│   « 0 » = q·k nul   │   │    │1 0 1 .│                     │
│    └0 1 0 1┘     (autorisé)      │   │    └0 1 0 1┘                     │
│  O = M·V + K·S₀                  │   │  O = M·V + K·S₄                  │
└───────────────┬──────────────────┘   └───────────────┬──────────────────┘
   S₀ = 0 ──────┤                            S₄ ───────┤
                ▼  S₄ = S₀ + K₁ᵀV₁                     ▼  S₈ = S₄ + K₂ᵀV₂
          S₄ = [[3,4],          ───SÉQUENTIEL───▶ S₈ = [[4,5],
                [5,5]]           (1 seule passe)        [7,7]]

DÉCODAGE : 1 token → M est 1×1, il ne reste que O = qᵀS  (aucun triangle)

Le problème — Préfill et décodage

Un prompt de 8 000 tokens arrive d’un coup ; la réponse sort ensuite token par token. Si le moteur traite le prompt au rythme de la génération — un token à la fois — l’utilisateur attend le premier mot pendant des secondes entières.

L’idée — Préfill et décodage

Deux phases, deux régimes : le préfill voit tous les tokens du prompt en même temps (travail massivement parallélisable) ; le décodage n’ajoute qu’un token par pas (travail intrinsèquement séquentiel). Même mécanisme, profils d’exécution opposés.

Pourquoi / à quel prix — Préfill et décodage

Séparer les deux permet d’optimiser chacun — latence du premier token d’un côté, débit de génération de l’autre. Le prix : deux chemins de code pour un seul mécanisme, qui doivent produire exactement les mêmes nombres.

Contrôle : Le préfill des 8 tokens construit un triangle 4×4 par chunk. Au décodage du token 9, quelle est la taille de ce triangle, et quelle partie du calcul disparaît complètement ?

Le problème — Récurrence naïve

La mémoire récurrente semble condamner le préfill : S₅ exige S₄, qui exige S₃… Exécuter 8 000 mises à jour l’une après l’autre laisse un GPU — conçu pour des matrices entières — presque vide à chaque pas.

L’idée — Récurrence naïve

Le diagnostic précis : la DÉPENDANCE est séquentielle (chaque S_t dépend de S_{t−1}), mais la majorité du CALCUL par token — produits q·k locaux, écritures k vᵀ — ne l’est pas. La récurrence naïve sérialise tout parce qu’elle ne sépare pas les deux.

Pourquoi / à quel prix — Récurrence naïve

Ce constat ouvre la porte du chunking : ne sérialiser que ce qui doit l’être. Le prix d’en rester à la version naïve se mesure directement : des unités matricielles facturées à l’heure qui exécutent des produits vecteur-matrice.

Contrôle : La récurrence exige S₈ après S₄. Pourquoi ne peut-on pas lancer les 8 mises à jour en parallèle, et pourquoi le GPU est-il malgré tout sous-utilisé si on les fait strictement une par une ?

Support visuel — Récurrence naïve

récurrence naïve : 8 pas séquentiels
t1 ▶ t2 ▶ t3 ▶ t4 ▶ t5 ▶ t6 ▶ t7 ▶ t8    GPU : ~1 token utile/pas

chunking C=4 : 2 pas séquentiels
┌ t1 t2 t3 t4 ┐ ──S₄──▶ ┌ t5 t6 t7 t8 ┐   GPU : 4 tokens/pas
└ en parallèle┘         └ en parallèle┘

la dépendance demeure (S₈ après S₄) —
mais elle ne porte plus chaque token

Le problème — Découper en chunks

Comment donner au GPU des blocs matriciels pleins sans violer l’ordre ? Il faut un découpage où l’intérieur d’un bloc se calcule en parallèle et où le passé lointain arrive compressé — sans double compte ni oubli.

L’idée — Découper en chunks

Un chunk de C tokens calcule d’un coup ses interactions internes autorisées (matrice C×C triangulaire) et lit le passé antérieur via l’état entrant : O = M·V + K·S_entrant. Sur la trace, le chunk 1 produit S₄, le chunk 2 le consomme — et o₅..o₈ sont exactement ceux de la récurrence.

Pourquoi / à quel prix — Découper en chunks

Exactitude algébrique, parallélisme retrouvé. Le prix : une mémoire temporaire en C² pour le triangle, et une complexité de code réelle — deux termes à additionner, donc deux occasions de se tromper : ce sont précisément les deux bugs de la trace.

Contrôle : Passer de C=4 à C=2 sur ces 8 tokens : combien d’états intermédiaires faut-il transmettre, et les sorties o₅..o₈ changent-elles de valeur ? Justifiez avec les chiffres de la trace.

Le problème — Triangle causal

À l’intérieur d’un chunk, tous les tokens se calculent ensemble — y compris t5 avec t7, qui est son futur. Sans garde-fou, le préfill apprendrait des dépendances que le décodage ne pourra jamais reproduire.

L’idée — Triangle causal

Une matrice triangulaire inférieure matérialise la règle « i ne lit que j ≤ i » : les cases au-dessus de la diagonale sont interdites par construction. Piège de lecture : un 0 sous la diagonale est un produit scalaire nul (autorisé) ; un « . » au-dessus est la causalité.

Pourquoi / à quel prix — Triangle causal

Le triangle rend la contrainte vérifiable d’un coup d’œil et gratuite à appliquer. Le prix d’un masque faux est vicieux : la perplexité s’améliore « trop bien » en préfill et rien ne casse — jusqu’au décodage, qui ne peut pas tricher.

Contrôle : Dans le masque du chunk 1, l’entrée (ligne 3, colonne 4) vaut 0 alors que k₃·k₄ = 0 aussi. Ces deux zéros ont-ils la même cause ? Que se passerait-il si k₃·k₄ valait 1 ?

Support visuel — Triangle causal

        colonne j (token lu)
          1  2  3  4
   ligne ┌1  .  .  .┐    « . » au-dessus de la diagonale :
   i     │0  1  .  .│    interdit — c’est le futur
  (token │1  0  1  .│
   lisant└0  1  0  1┘    « 0 » en dessous : autorisé,
                          produit scalaire nul

  deux zéros visuellement proches, deux causes distinctes

Le problème — État entrant et sortant

Le chunk 2 ne doit revoir aucun token du chunk 1 — sinon le parallélisme s’effondre — mais o₅ dépend de v₁ et v₃. Comment transmettre « tout le passé utile » sans transmettre le passé ?

L’idée — État entrant et sortant

Par l’état : le chunk 1 émet S₄ = K₁ᵀV₁, un résumé de taille fixe ; le chunk 2 le lit par q_tᵀS₄ et ajoute ses termes locaux. Sur la trace : o₅ = [3,4] (hérité) + [0,1] (local) = [3,5]. Les frontières transportent l’ordre et la causalité, pas les tokens.

Pourquoi / à quel prix — État entrant et sortant

Un seul objet à passer entre blocs, de taille fixe. Le prix : le chunk 2 ne peut plus décomposer [3,4] entre v₁ et v₃ — la compression de la session 13 s’applique aux frontières. Et oublier le terme entrant (bug 1) ampute tout le prompt.

Contrôle : Le chunk 2 reçoit S₄ = [[3,4],[5,5]] et rien d’autre du passé. Reconstituez o₅ = [3,5] en séparant le terme entrant du terme local, puis dites ce que le chunk 2 ne peut plus savoir sur v₁ et v₃.

Le problème — Taille du chunk

C = 1 redonne la récurrence lente ; C = longueur totale fait exploser le triangle en C². Entre les deux, qui décide ? Le même code peut tourner plusieurs fois plus lentement avec un C mal choisi pour le GPU.

L’idée — Taille du chunk

C arbitre deux coûts opposés : transitions séquentielles en n/C contre mémoire temporaire en C². Doubler C divise les transitions par 2 et multiplie le triangle par 4 — l’optimum se trouve là où le triangle sature juste la mémoire rapide (SRAM).

Pourquoi / à quel prix — Taille du chunk

Un C bien choisi sature le matériel. Le prix : le bon C n’est pas transférable d’un GPU à l’autre — c’est un paramètre d’exécution à re-mesurer, pas une constante du modèle. Et il ne change jamais les résultats, seulement leur coût.

Contrôle : Vous doublez C de 64 à 128 sur un GPU dont la SRAM est déjà pleine. Le triangle temporaire est en C² : par quel facteur augmente-t-il, et pourquoi le débit peut-il baisser alors qu’il y a moins de transitions ?

Support visuel — Taille du chunk

   C      transitions (8/C)     triangle C²
   1             8                    1
   2             4                    4
   4             2                   16
   8             1                   64

 transitions ÷2  ⇔  triangle ×4
 l’optimum : le plus grand C dont le triangle tient en SRAM

Cas guidé — trace complète

Pour 8 tokens en chunks de 4, le premier calcule un triangle 4×4 puis transmet S₄. Le second reçoit S₄, calcule son triangle local et produit S₈. Aucun token du premier bloc ne peut lire le second.

8 tokens, d=2, chunk C=4, q_t = k_t, S₀ = [[0,0],[0,0]]
  chunk 1 : k=[1,0],[0,1],[1,0],[0,1]   v=[2,3],[5,1],[1,1],[0,4]
  chunk 2 : k=[1,0],[0,1],[1,0],[0,1]   v=[0,1],[2,2],[1,0],[0,0]

── CHUNK 1 ───────────────────────────────────────────────────────
  Q Kᵀ masqué en triangle inférieur (i lit j ≤ i) :
      ┌1 0 0 0┐
      │0 1 0 0│      les zéros au-dessus de la diagonale
      │1 0 1 0│      sont la causalité, pas une valeur nulle
      └0 1 0 1┘
  O = M·V + K·S₀ ,  S₀=0 :
    o₁=[2,3]  o₂=[5,1]  o₃=v₁+v₃=[3,4]  o₄=v₂+v₄=[5,5]
  état sortant  S₄ = K₁ᵀV₁ = [[3,4],[5,5]]

── CHUNK 2 (reçoit S₄, ne voit jamais v₁..v₄ individuellement) ────
  terme inter-chunk : q₅ᵀS₄ = [3,4]   ;  q₆ᵀS₄ = [5,5]
    o₅ = [3,4] + v₅        = [3,4]+[0,1] = [3,5]
    o₆ = [5,5] + v₆        = [5,5]+[2,2] = [7,7]
    o₇ = [3,4] + v₅+v₇     = [3,4]+[1,1] = [4,5]
    o₈ = [5,5] + v₆+v₈     = [5,5]+[2,2] = [7,7]
  S₈ = S₄ + K₂ᵀV₂ = [[3+1,4+1],[5+2,5+2]] = [[4,5],[7,7]]
  ✅ o₅..o₈ sont EXACTEMENT ceux de la récurrence token par token

── ❌ ERREUR 1 : oublier le terme d’état entrant ──────────────────
  o₅ = v₅ = [0,1]  au lieu de [3,5] : le chunk 2 a perdu tout le prompt

── ❌ ERREUR 2 : masque plein au lieu du triangle ─────────────────
  o₅ inclurait v₇ : [3,4]+[0,1]+[1,0] = [4,5] ≠ [3,5]
  le token 5 lit le token 7 → fuite du futur, invisible en préfill,
  impossible à reproduire en décodage : la perplexité chute « trop bien »

CONTRÔLE DE FORME : K(4×2)·Kᵀ(2×4) → M(4×4) ; M(4×4)·V(4×2) → O(4×2) ;
K(4×2)·S(2×2) → (4×2). Mémoire temporaire du triangle = C² = 16, pas 4.

Choisir C : le même calcul, deux régimes matériels

Critère (8 tokens) Petit chunk, C=1 Grand chunk, C=8
Transitions d’état séquentielles 8 (aucun parallélisme) 1 (un seul passage d’état)
Triangle temporaire en mémoire C² = 1 par chunk C² = 64 par chunk
Forme des matmuls envoyées au GPU vecteur × matrice, unités inactives matrice × matrice, unités saturées
Régime où c’est le bon choix décodage : 1 token à la fois préfill : prompt entier disponible

Laboratoire causal

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

/interactives/curriculum/chunk-gate-memory.html?lang=fr

Erreurs fréquentes

« Le chunking est une approximation : on gagne de la vitesse en perdant un peu de précision. »

La trace le contredit chiffre pour chiffre : o₅..o₈ valent [3,5],[7,7],[4,5],[7,7] avec chunks comme sans. C’est une réorganisation algébrique exacte. Ce qui change est le coût mémoire (C²) et le nombre de transitions, jamais le résultat.

« Un chunk plus grand est toujours plus rapide, puisqu’il y a moins d’étapes séquentielles. »

Le triangle temporaire croît en C² : de C=64 à C=128, il est multiplié par 4. Dès qu’il déborde de la SRAM, le débit s’effondre malgré la réduction du nombre de transitions. L’optimum est matériel, pas mathématique.

Frontière, preuve et sources

Le chunking améliore l’exécution ; il ne change pas automatiquement la capacité informationnelle de l’état.

Statut de preuve : Mécanismes établis ; les simplifications numériques sont pédagogiques.

  • Dossier de cours bilingue fourni par le propriétaire, chapitre 10.
  • Katharopoulos et al., “Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention”, ICML (2020).
  • Yang et al., “Gated Linear Attention Transformers with Hardware-Efficient Training” (chunkwise parallel form), ICML (2024).
  • 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

Refaites la trace complète avec C = 2 : chunks {t1,t2}{t3,t4}{t5,t6}{t7,t8}.

  1. Avant calcul : combien de passages d’état, et de quelle taille chacun ?
  2. Calculez S₂, S₄, S₆, puis vérifiez o₅..o₈ contre la trace en C = 4.
  3. Concluez : qu’est-ce qui a changé, et qu’est-ce qui n’a pas le droit de changer ?

Synthèse et ticket de sortie

  • Préfill et décodage
  • Récurrence naïve
  • Découper en chunks
  • Triangle causal
  • État entrant et sortant
  • Taille du chunk

Ticket : 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: Faire relier chaque étape à la suivante par un verbe causal. Signaler toute flèche purement décorative.

Notes formateur: Chronométrer un vrai chat à main levée : « le premier mot met deux secondes, les suivants trente millisecondes — pourquoi ? ». La question du vécu utilisateur amorce mieux le beat que le vocabulaire préfill/décodage.

Notes formateur: Faire chronométrer mentalement les deux phases : « combien de tokens en parallèle au préfill ? combien au décodage ? ». Tant que cette asymétrie n’est pas dite par un apprenant, le reste de la séance sonne comme une optimisation gratuite.

Notes formateur: Réponse : le triangle devient 1×1, c’est-à-dire trivial — il ne reste que la lecture d’état O = qᵀS plus le terme local. Toute la machinerie chunkwise est une optimisation du préfill ; le décodage n’en profite pas. Erreur attendue : « un 4×4 avec du padding ».

Notes formateur: Faire mimer la chaîne : huit apprenants, chacun ne calcule qu’après avoir reçu le papier du précédent. Pointer les sept qui attendent : « voilà le GPU ». L’image reste toute la séance.

Notes formateur: Demander à quelqu’un d’exécuter les 8 mises à jour à voix haute, une par une. La lenteur ressentie fait le travail pédagogique mieux qu’un graphique de débit.

Notes formateur: Réponse : les mises à jour d’état forment une chaîne — les paralléliser telles quelles changerait le résultat ; mais une par une, chaque pas n’offre qu’un produit vecteur-matrice et les unités matricielles restent inoccupées. Nuance à exiger : le GPU n’est pas « lent », il est inoccupé.

Notes formateur: Faire compter les flèches ▶ de chaque régime : 7 contre 1. Puis demander ce que devient ce compte pour 8 000 tokens et C = 64 — la division du chemin critique est la seule chose que le chunking achète.

Notes formateur: Demander : « que faudrait-il pour que la moitié droite du tableau travaille en même temps que la gauche ? ». Les propositions spontanées — copier le passé ? tout renvoyer ? — préparent la valeur d’un état unique S₄.

Notes formateur: Découper le tableau physiquement en deux zones et interdire à la moitié droite de regarder la gauche autrement que par un post-it « S₄ ». Cette contrainte spatiale est ce que le code fait vraiment.

Notes formateur: Réponse : chunks {1-2}{3-4}{5-6}{7-8} → transmettre S₂, S₄, S₆, soit trois passages au lieu d’un ; o₅..o₈ = [3,5],[7,7],[4,5],[7,7], inchangés — le chunking est exact, seule l’exécution change. Erreur attendue : « retrouver » des écarts d’arrondi imaginaires.

Notes formateur: Poser le piège avant la solution : « dans un bloc calculé d’un coup, qu’est-ce qui empêche t5 de lire t7 ? ». Réponse honnête : rien — sauf le masque. Laisser le malaise s’installer avant de montrer le triangle.

Notes formateur: Effacer un zéro du triangle et faire chercher le bug par le groupe. La bonne réponse attendue n’est pas « le nombre est faux » mais « le token 5 a lu le futur, et l’entraînement s’en félicitera ».

Notes formateur: Réponse : causes différentes — (3,4) est au-dessus de la diagonale, interdite par la causalité ; k₃·k₄ = 0 est un zéro CALCULÉ qui avait le droit d’être non nul. Si k₃·k₄ valait 1, la case (4,3) — sous la diagonale — deviendrait 1, mais (3,4) resterait masquée. Erreur attendue : « c’est zéro partout, même chose ».

Notes formateur: Faire colorier les deux familles de zéros en deux couleurs avant d’énoncer la règle. Le contrôle « Ces deux zéros ont-ils la même cause ? » revient exactement là-dessus — cette diapositive est sa préparation.

Notes formateur: Écrire S₄ = [[3,4],[5,5]] au tableau et demander : « v₁ était [2,3] — où est-il ? ». Le silence est la leçon : il est dedans, mais plus séparable. Rappel explicite de la superposition de la session 13.

Notes formateur: Faire écrire S₄ sur un papier, retourner le tableau du chunk 1, puis demander de calculer o₅. Ce qu’ils réclameront spontanément est exactement l’information que l’état doit transporter.

Notes formateur: Réponse : o₅ = q₅ᵀS₄ + v₅ = [3,4] + [0,1] = [3,5] — terme hérité plus terme local. Le chunk 2 ne peut plus savoir COMMENT [3,4] se décompose entre v₁ et v₃ : l’attribution est perdue à la frontière (superposition, session 13). Erreur attendue : croire v₁ et v₃ récupérables « quelque part » dans S₄.

Notes formateur: Sonder : « chunk de 4, de 64, de 4 096 : qui dit mieux ? ». Faire voter pour un C avant d’exposer le compromis — voter force chacun à choisir un critère, et les critères divergents font le débat du beat.

Notes formateur: Faire calculer C² pour C = 16, 64, 128 et comparer à un budget SRAM annoncé. Terminer sur « il n’y a pas de bon C, il y a un bon C pour ce GPU » — c’est la seule conclusion honnête.

Notes formateur: Réponse : C² passe de 4 096 à 16 384 — ×4 pour un doublement. Si la SRAM déborde, le triangle migre vers la mémoire lente et chaque accès coûte un ordre de grandeur de plus : le débit chute malgré deux fois moins de transitions. Erreur attendue : raisonner en FLOPs seuls, sans hiérarchie mémoire.

Notes formateur: Faire prolonger la table jusqu’à C = 128 et donner un budget SRAM fictif (p. ex. 8 192 cases) : le groupe doit trouver seul le C qui déborde. La règle finale s’énonce alors sans aide.

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: Pour chaque affirmation, faire produire le contre-exemple minimal par le groupe avant de donner la correction.

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: Réponses : trois passages (S₂, S₄, S₆), chacun 2×2 — la taille de l’état ne dépend pas de C. S₂ = [[2,3],[5,1]] ; S₄ = [[3,4],[5,5]] (identique à la trace) ; S₆ = [[3,5],[7,7]]. o₅..o₈ = [3,5],[7,7],[4,5],[7,7] — inchangés : le chunking est exact (erreur fréquente 1). Ce qui change : le nombre de transitions et la taille des triangles (quatre 2×2 au lieu de deux 4×4). Idée fausse attendue : chercher des écarts d’arrondi. 12 minutes, binômes.

Notes formateur: Faire reconstruire la chaîne sans regarder les slides, puis remplir le ticket en six lignes maximum. Comparer à la prédiction initiale du premier slide et nommer ce qui a réellement changé.