opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
Browse files- docs/OPTIMIZATION_2026-08-22.md +200 -0
docs/OPTIMIZATION_2026-08-22.md
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Optimisation de l'entraînement Fractus-1B — 2026-08-22
|
| 2 |
+
|
| 3 |
+
**Principe cardinal : chirurgie à cœur ouvert.**
|
| 4 |
+
Toute modification du code doit être **mathématiquement équivalente** à
|
| 5 |
+
l'implémentation de référence, prouvée par tests d'équivalence numérique,
|
| 6 |
+
sans changer la moindre forme de paramètre. Le cerveau (`.pt`) ne bouge pas ;
|
| 7 |
+
seul le corps (le code) est opéré. La reprise au token exact via les manifests
|
| 8 |
+
reste valide — un pod peut passer à ce code **en cours de run**, sans jeter la
|
| 9 |
+
digestion.
|
| 10 |
+
|
| 11 |
+
---
|
| 12 |
+
|
| 13 |
+
## 1. Diagnostic : pourquoi ~1000 tok/s/GPU alors que la 5090 peut faire bien plus
|
| 14 |
+
|
| 15 |
+
Config 1B par bloc : 20 têtes × d_head=64, 2 niveaux, chunk C=128.
|
| 16 |
+
Le moteur aplatit `(B, niv, H)` en `G = B·niv·H = B·40` groupes.
|
| 17 |
+
|
| 18 |
+
### Goulot n°1 — le cumsum déguisé en matmul (`attention.py`, ancien code)
|
| 19 |
+
|
| 20 |
+
```python
|
| 21 |
+
S = torch.einsum("tj,bjpq->btpq", mask_tril, outer) # Σ_{j≤t} outer[j]
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
C'est un cumsum causal écrit comme une multiplication par un masque
|
| 25 |
+
triangulaire inférieur :
|
| 26 |
+
|
| 27 |
+
| Forme | Opérations | À C=128, D=64 |
|
| 28 |
+
|---|---|---|
|
| 29 |
+
| masque (ancienne) | O(C²·D²) MACs + rétro O(C²·D²) | ≈ 67M MACs/groupe |
|
| 30 |
+
| cumsum (nouvelle) | O(C·D²) additions | ≈ 0.5M adds/groupe → **~128× moins** |
|
| 31 |
+
|
| 32 |
+
La rétropropagation profite pareil : dérivée d'un cumsum = cumsum inverse,
|
| 33 |
+
pas un nouveau matmul masqué géant.
|
| 34 |
+
|
| 35 |
+
### Goulot n°2 — matérialisation des tenseurs `(G, C, D, D)`
|
| 36 |
+
|
| 37 |
+
`outer` et `S` font `G·128·64·64` éléments. À B=4 : G=160 → **84M éléments
|
| 38 |
+
(~168 Mo en bf16)** chacun, retenus pour la backward, **par bloc ×16 blocs**.
|
| 39 |
+
C'est LA pression VRAM qui force `BATCH=2-4` et interdit `torch.compile`
|
| 40 |
+
(commentaire "compile disabled for VRAM" dans fast4gpu_boost.py).
|
| 41 |
+
|
| 42 |
+
Forme *chunked* (bloc=64) : on ne matérialise que `(G, 64, 64)` (scores)
|
| 43 |
+
+ l'état courant `(G, D, D)` → **~32× moins d'activations dominantes**.
|
| 44 |
+
|
| 45 |
+
### Goulot n°3 — la tête liée 50257
|
| 46 |
+
|
| 47 |
+
Les logits `(B, C, 50257)` fp32 (~103 Mo à B=4) + softmax backward sont
|
| 48 |
+
matérialisés deux fois avec SS. La tête ≈ **la moitié des FLOPs actifs**
|
| 49 |
+
(64.3M des ~192M MACs/token). CE par morceaux checkpointée → mémoire
|
| 50 |
+
transitoire plate (`ce_chunk·vocab`), sans toucher aux FLOPs.
|
| 51 |
+
|
| 52 |
+
### Goulot n°4 — pipeline données
|
| 53 |
+
|
| 54 |
+
`np.load(mmap).to(torch.int64)` copiait **tout le shard en RAM** :
|
| 55 |
+
430M tokens × 8 octets = **3,4 Go/process × 8 process**. Le slicing
|
| 56 |
+
par chunk int32→long supprime ça (et la page cache fait le reste).
|
| 57 |
+
|
| 58 |
+
### Estimation d'utilisation GPU
|
| 59 |
+
|
| 60 |
+
FLOPs actifs/token ≈ 1.15 GFLOP (fwd+bwd) → à 1000 tok/s ≈ 1.15 TFLOPS
|
| 61 |
+
sur une 5090 (~200 TFLOPS bf16 crête) : **<1% d'utilisation**. Le modèle est
|
| 62 |
+
borné par la mémoire/les petits batches/l'overhead Python, pas par le calcul.
|
| 63 |
+
Tout l'espace d'optimisation est là.
|
| 64 |
+
|
| 65 |
+
---
|
| 66 |
+
|
| 67 |
+
## 2. Les optimisations (toutes prouvées équivalentes)
|
| 68 |
+
|
| 69 |
+
### Opt 1 — noyau cumsum (`fractus/nn/attention.py`)
|
| 70 |
+
- `_linear_attention_causal_einsum` : **référence conservée** (vérité terrain).
|
| 71 |
+
- `_linear_attention_causal_cumsum` : production par défaut.
|
| 72 |
+
- `_linear_attention_causal_vectorized` : dispatcher inchangé pour les appelants.
|
| 73 |
+
- Preuves : `tests/test_attention_equivalence.py::test_forward_matches_einsum_*`,
|
| 74 |
+
`test_gradients_match_with_carry`, `test_both_match_looped_reference`.
|
| 75 |
+
|
| 76 |
+
### Opt 2 — forme chunked mémoire-plate (`_linear_attention_causal_chunked`)
|
| 77 |
+
Intra-bloc via matrice de scores `(bloc×bloc)` masquée + inter-bloc via l'état
|
| 78 |
+
courant `(S_run, z_run)` ; mise à jour inclusive APRÈS lecture (= sémantique
|
| 79 |
+
exacte S_t incluant le token t). Sélection : `set_attention_impl('chunked')`
|
| 80 |
+
ou env `FRACTUS_ATTN_IMPL=chunked`. Ragged → fallback cumsum exact.
|
| 81 |
+
- Recommandation pod : **chunked sur GPU** (VRAM → compile + gros batch),
|
| 82 |
+
cumsum reste parfait sur CPU.
|
| 83 |
+
|
| 84 |
+
### Opt 3 — CE par morceaux checkpointée (`fractus/nn/ce.py`)
|
| 85 |
+
`chunked_cross_entropy(h, W, targets, ce_chunk)` : perte moyenne identique,
|
| 86 |
+
grads identiques (tolérance ordre de sommation fp32), pics mémoire plats grâce
|
| 87 |
+
à `torch.utils.checkpoint` par morceau (recompute en backward).
|
| 88 |
+
Nouvelle entrée moteur : `tick_chunk_train_ce(obs, targets[, return_hidden])`.
|
| 89 |
+
`sample_tokens_chunked(...)` remplace le multinomial dense pour le SS
|
| 90 |
+
(distribution identique ; flux RNG consommé par morceaux → tirages non
|
| 91 |
+
bit-identiques, statistiquement équivalents).
|
| 92 |
+
|
| 93 |
+
### Opt 4 — sémantique SS préservée + accumulation optionnelle
|
| 94 |
+
`scripts/fast4gpu_boost_v2.py` réplique EXACTEMENT v1 quand `ACCUM=1`
|
| 95 |
+
(défaut) : step TF puis step SS séparés. `ACCUM>1` = capability nouvelle,
|
| 96 |
+
déviation documentée (grads TF+SS accumulés, un clip+step tous les N lots).
|
| 97 |
+
|
| 98 |
+
### Opt 5 — pipeline données zéro-copie
|
| 99 |
+
Slice memmap int32 par chunk → transfert → `.long()` sur GPU.
|
| 100 |
+
Un seul fetch `(B·SEQ+1)` fournit chunk ET target. RAM économisée :
|
| 101 |
+
~27 Go sur un pod 8×5090.
|
| 102 |
+
|
| 103 |
+
---
|
| 104 |
+
|
| 105 |
+
## 3. Comment déployer sur le pod (cœur ouvert)
|
| 106 |
+
|
| 107 |
+
```bash
|
| 108 |
+
# 1. Sauvegarder l'état (rien à faire de plus : HF = source de vérité)
|
| 109 |
+
# Les checkpoints gpu*.pt + RESUME_MANIFEST restent valides tels quels.
|
| 110 |
+
|
| 111 |
+
# 2. Remplacer le corps :
|
| 112 |
+
# fractus/nn/attention.py, fractus/nn/ce.py (nouveau),
|
| 113 |
+
# fractus/continuous_engine.py (méthode ajoutée, rien retiré),
|
| 114 |
+
# scripts/fast4gpu_boost_v2.py (nouveau).
|
| 115 |
+
|
| 116 |
+
# 3. Relancer chaque GPU avec les MÊMES offsets qu'avant l'arrêt :
|
| 117 |
+
CUDA_VISIBLE_DEVICES=$i GPU_ID=$i \
|
| 118 |
+
START_TOKEN=$(python -c "import json;print(json.load(open('checkpoints/RESUME_MANIFEST_8GPU.json'))['gpu_$i']['start_token'])") \
|
| 119 |
+
BATCH=8 CE_CHUNK=2048 FRACTUS_ATTN_IMPL=chunked COMPILE=1 \
|
| 120 |
+
python -u scripts/fast4gpu_boost_v2.py
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
Ordre de montée en puissance recommandé (valider à chaque étape) :
|
| 124 |
+
1. Code swap, mêmes réglages que v1 (`BATCH=4 CE_CHUNK=0`) → vérifier ema_tf
|
| 125 |
+
continue exactement sa courbe (équivalence en conditions réelles).
|
| 126 |
+
2. `CE_CHUNK=2048` → perte identique, VRAM ↓.
|
| 127 |
+
3. `FRACTUS_ATTN_IMPL=chunked` → VRAM ↓↓ (mesuré §4 : ×15–25 vs référence,
|
| 128 |
+
mémoire plate là où la référence explose).
|
| 129 |
+
4. `COMPILE=1` puis monter `BATCH` (8, 16…) — surveiller tok/s et VRAM.
|
| 130 |
+
5. Optionnel `ACCUM` si on veut un lot effectif plus grand sans OOM.
|
| 131 |
+
|
| 132 |
+
Critères de non-régression (cf. docs/TRUSTED_LOSS.md) : ema_tf continue de
|
| 133 |
+
descendre sans NaN, lb stable ~14, sondes gén inchangées de comportement.
|
| 134 |
+
|
| 135 |
+
---
|
| 136 |
+
|
| 137 |
+
## 4. Résultats mesurés (CPU local, torch 2.9.1+cpu, 2026-08-23)
|
| 138 |
+
|
| 139 |
+
Micro-bench attention formes 1B réelles (G=B·40 groupes, C=128, dH=64),
|
| 140 |
+
carry actif ; médiane sur 8 itérations. Chaque cellule tourne dans son
|
| 141 |
+
propre sous-processus : un crash natif d'une cellule n'emporte pas la table.
|
| 142 |
+
|
| 143 |
+
| Impl | B | G | fwd ms | fwd+bwd ms | RSS Δ MB |
|
| 144 |
+
|---|---|---|---|---|---|
|
| 145 |
+
| einsum (réf) | 2 | 80 | 267.4 | 572.1 | +3 |
|
| 146 |
+
| einsum (réf) | 4 | 160 | 584.7 | 1126.6 | 0 |
|
| 147 |
+
| einsum (réf) | 8 | 320 | — | — | crash natif (commit) |
|
| 148 |
+
| cumsum | 2 | 80 | 322.4 | 799.1 | −8 |
|
| 149 |
+
| cumsum | 4 | 160 | — | — | crash natif (commit) |
|
| 150 |
+
| chunked | 2 | 80 | **12.3** | **31.9** | 0 |
|
| 151 |
+
| chunked | 4 | 160 | **23.0** | **72.4** | 0 |
|
| 152 |
+
| chunked | 8 | 320 | **40.5** | **150.0** | −10 |
|
| 153 |
+
|
| 154 |
+
Lecture honnête de ces chiffres :
|
| 155 |
+
|
| 156 |
+
- **Chunked vs einsum à B égal** : ×21.7 (fwd) et ×17.9 (fwd+bwd) à B=2 ;
|
| 157 |
+
×25.4 / ×15.6 à B=4.
|
| 158 |
+
- **Mémoire** : einsum et cumsum matérialisent le tenseur (G, C, dH, dH)
|
| 159 |
+
(~0.34 GB à G=160, ~1.3 GB à G=320, multiplié par les buffers de
|
| 160 |
+
backward). Sur cette machine au commit mémoire limité ils segfaultent
|
| 161 |
+
(rc=3221225477) dès G≥160–320 ; chunked est memory-flat ((G, block²+dH²))
|
| 162 |
+
et traverse toutes les tailles testées. C'est exactement la propriété qui
|
| 163 |
+
déverrouille BATCH≥8 + torch.compile sur pod.
|
| 164 |
+
- **Cumsum sur CPU est plus lent qu'einsum** : l'einsum masqué descend en
|
| 165 |
+
bmm BLAS très optimisé, tandis que le scan élément-par-élément est
|
| 166 |
+
borné par la bande passante. La réduction de FLOPs O(C²·dH²)→O(C·dH²)
|
| 167 |
+
se paie réellement sur GPU/compile, pas sur ce CPU — c'est pour ça que
|
| 168 |
+
le défaut repo reste `cumsum` (mathématiquement prouvé, sémantique
|
| 169 |
+
simple) et que la recommandation pod passe directement à `chunked`.
|
| 170 |
+
- Les timings absolus CPU ne se transfèrent pas à CUDA ; ce qui se
|
| 171 |
+
transfère : l'ordre relatif des kernels, le profil mémoire, et le fait
|
| 172 |
+
mesuré que chunked scale là où la référence explose. Bench GPU à faire
|
| 173 |
+
sur pod (même script).
|
| 174 |
+
|
| 175 |
+
End-to-end moteur CPU (bench_engine.py, d=128, 2 blocs, E8, B=8, SEQ=128,
|
| 176 |
+
20 steps) :
|
| 177 |
+
|
| 178 |
+
| Kernel | tok/s |
|
| 179 |
+
|---|---|
|
| 180 |
+
| cumsum (défaut) | 703 |
|
| 181 |
+
| chunked | 764 (+8.7 %) |
|
| 182 |
+
|
| 183 |
+
## 5. Tests
|
| 184 |
+
|
| 185 |
+
```bash
|
| 186 |
+
py -m pytest tests/ -q # 44 passed = suite repo (28) + équivalences (14) + smoke v2 (2)
|
| 187 |
+
py benchmarks/bench_attention.py --iters 8
|
| 188 |
+
py benchmarks/bench_engine.py --steps 20
|
| 189 |
+
```
|
| 190 |
+
|
| 191 |
+
Statut : **44/44 passent** (2026-08-23). Les deux preuves moteur
|
| 192 |
+
(`test_engine_tick_chunk_train_ce_matches_train`,
|
| 193 |
+
`test_engine_end_to_end_chunk_equivalence`) exigent le clonage explicite des
|
| 194 |
+
poids (`load_state_dict`) entre instances comparées — deux constructions sous
|
| 195 |
+
un même seed ont des poids DIFFÉRENTS, piège documenté dans les docstrings.
|
| 196 |
+
|
| 197 |
+
---
|
| 198 |
+
|
| 199 |
+
*Document généré pendant la session d'optimisation 2026-08-22/23.
|
| 200 |
+
Règle : ne jamais merger dans fractus-cte sans que la section 4 soit remplie.*
|