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