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.*