fractus-cte / docs /OPTIMIZATION_2026-08-22.md
thefinalboss's picture
opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
397bf25 verified
|
Raw History Blame
9.29 kB

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)

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)

# 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

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.