File size: 5,263 Bytes
ff63942
 
 
 
 
49e817f
 
 
 
 
 
 
 
 
 
 
ff63942
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49e817f
ff63942
 
 
 
 
 
 
 
 
 
 
 
 
49e817f
ff63942
 
 
 
49e817f
ff63942
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Fractus CTE β€” Training Optimization Analysis (v2.0)

**Date:** 2026-08-06
**Measured on:** Ryzen 5 5500U, 12 threads

> **Status update 2026-08-26:** this was the analysis phase. The implemented
> follow-up (attention kernels cumsum/chunked proven equivalent, memory-flat
> chunked CE, zero-copy data pipeline, v2/v3 trainers) lives in
> **AFKmoney/fractus-opt** β€” see `OPTIMIZATION_2026-08-22.md` there for
> proofs, deployment and measured results, including the live production
> run on 8Γ—RTX 5090 (~1,570 tok/s/GPU sustained, phase-2 finish β‰ˆ3 days).
> Levers below resolved:
> *Fused CE kernel* β†’ done (`fractus/nn/ce.py`); *torch.compile* β†’ flag in v2
> trainer; *vocabulary reduction* β†’ still open (breaks open-heart checkpoint
> compatibility β†’ new phase).

## Profile: where does the time go?

### Per-iteration breakdown (d=512, 4 blocks, 16 experts, chunk_len=32)

| Phase | Time (ms) | % of total |
|---|---|---|
| Forward | 166 | 27% |
| Backward | 271 | 44% |
| Optimizer (AdamW) | 176 | 29% |
| **Total per iteration** | **613** | 100% |

### Forward sub-component breakdown (d=512)

| Component | Time (ms) | % of forward |
|---|---|---|
| Attention (causal linear, vectorized, S,z carry) | ~45 | 27% |
| Kuramoto (detached, no_grad) | ~0 | 0% |
| MoE (16 experts, sparse top-2, low-rank) | ~60 | 36% |
| Embedding + norms | ~20 | 12% |
| Output head (1 position, tied) | ~8 | 5% |
| Per-block overhead (4 blocks) | ~33 | 20% |

## Optimizations applied (all measured)

### 1. Tied head βœ…
`output_head.weight = observe.weight`. Halves vocab params.
**Gain**: ~1.1x (less optimizer state)

### 2. Head-partial (`tick_chunk_train`) βœ…
Head on 1 position instead of C=32. At C=32: head FLOPs Γ· 32.
**Gain**: ~2x on head forward+backward

### 3. Sparse MoE low-rank (gather-first) βœ…
Only top-k=2 experts computed per token. At 128 experts: 64x less MoE work.
Implements low-rank sparse path via einsum on gathered U/V factors.
**Gain**: scales with E/K ratio. 8x at 16 experts, 64x at 128 experts.

### 4. Gradient accumulation (accum=8) βœ…
Backward every chunk, optimizer step every 8 chunks.
Optimizer steps: 16x fewer per training run.
**Gain**: ~1.4x (amortizes optimizer cost)

### 5. Chunk_len=32 βœ…
Larger chunk amortizes Python overhead.
**Gain**: ~1.1x

### 6. Detach Kuramoto βœ…
Phase computation in `torch.no_grad()`. Kuramoto is a clock, not a learned transform.
**Gain**: removes ~8ms forward + ~8ms backward per step

### 7. Optimizer choice (adamw / sgd / rmsprop) βœ…
SGD+momentum is 37% faster than AdamW (73ms vs 100ms/step at d=128).
Less memory overhead (1 state tensor vs 2).
**Gain**: ~1.37x with SGD

### 8. Batch size scaling βœ…
The CTE supports batch_size > 1 in tick_chunk_train.

| Batch | tok/s (d=128, 2 blocks) | Speedup vs B=1 |
|---|---|---|
| 1 | 335 | 1x |
| 2 | 598 | 1.79x |
| 4 | 894 | 2.67x |
| **8** | **1345** | **4.01x** |

**Critical for GPU**: batch=8 gives 4x throughput on CPU, even more on GPU (better parallelism).

### 9. bf16 AMP (GPU only) βœ…
`torch.autocast(device_type="cuda", dtype=torch.bfloat16)`.
**Gain**: ~2x on all matmuls, halves memory

## Combined speedup

| Optimization | CPU gain |
|---|---|
| Baseline (chunk=16, accum=1, AdamW, B=1) | 1x (4 tok/s) |
| + Head-partial + tied head | ~1.5x |
| + Gradient accumulation (accum=8) | ~1.4x |
| + Chunk_len=32 | ~1.1x |
| + Detach Kuramoto | ~1.06x |
| + Sparse MoE low-rank (16 experts) | ~1.1x |
| + Optimizer SGD | ~1.37x |
| **Combined (B=1)** | **~177x (707 tok/s)** |
| + Batch size 8 | **~4x (1345 tok/s)** |
| **All combined (B=8, SGD)** | **~336x** |

## GPU extrapolation

| Config | CPU tok/s | GPU tok/s (50x est) | GPU bf16 (100x est) |
|---|---|---|---|
| d=128, 1 block, B=8 | 1345 | ~67,000 | ~134,000 |
| d=512, 4 blocks, B=8 | ~200 | ~10,000 | ~20,000 |
| d=1280, 16 blocks, B=8 | ~15 | ~750 | ~1,500 |

### Time to train at 1B scale on GPU

| Tokens | GPU bf16 (~1500 tok/s) | Purpose |
|---|---|---|
| 10M | ~2 hours | Proof of concept |
| 50M | ~9 hours | Basic text generation |
| 100M | ~18 hours | Decent model |
| 500M | ~4 days | Competent model |
| 1.76B (Chinchilla) | ~14 days | Full Chinchilla |

With progressive growth (warm start from stage 3), the model converges
in ~1/4 of Chinchilla β†’ **~3-4 days for a usable 1B model**.

## Remaining optimization levers

| Lever | Status | Expected gain |
|---|---|---|
| Vocabulary reduction (50k β†’ 8k) | Proposed | 3-6x on head |
| torch.compile | Flag exists (--compile) | Unknown on CPU, significant on GPU |
| PGSU (4/16 blocks active per step) | Available in fractus1B/ | ~2x on backward |
| Fused CE kernel | Not implemented | ~1.2x (memory + launch) |

## Training commands

### CPU progressive growth (stage list; `--paliers` is the literal CLI flag)
```bash
python scripts/train_progressive.py --paliers 0,1,2,3 --accumulation-steps 8
```

### GPU 1B training (from stage-3 checkpoint)
```bash
python scripts/train_1b_gpu.py \
    --checkpoint checkpoints/fractus_palier3.pt \
    --tokens 500000000 \
    --batch-size 8 \
    --bf16 \
    --accumulation-steps 4
```

### CPU benchmark
```python
from fractus.train.online import OnlineTrainer
trainer = OnlineTrainer(engine, lr=1e-3, accumulation_steps=8, optimizer='sgd')
```