Skip to content

About

Reversible 20M-param LLM trained on 50M FineWeb-Edu tokens on a 6GB RTX 3060 Laptop: leapfrog-midpoint reversibility is accuracy-neutral, 2.2x less activation memory, 3.4x larger max batch, ~1.5x slower

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Reversible 20M-parameter LLM — 50M tokens on a 6 GB laptop GPU

Does making a transformer reversible — recomputing activations in the backward pass instead of storing them — buy enough memory to be worth it, and what does it cost in loss and in wall-clock?

This repo answers that with three parameter-matched, token-matched runs on one NVIDIA RTX 3060 Laptop, 6 GB, each at its own tuned learning rate:

# run arch batch steps LR val loss ppl tok/s peak alloc peak reserved wall
1 baseline, fixed batch standard pre-norm 16 6103 7.5e-4 3.783 43.9 40.5k 2130 MB 2348 MB 20.6 min
2 + reversibility, same batch rev_mid (leapfrog) 16 6103 7.5e-4 3.797 44.6 28.8k 1157 MB 1482 MB 28.9 min
3 + reversibility, max batch rev_mid 192 508 1.3e-3 4.254 70.4 35.2k 4502 MB 5118 MB 23.6 min

All three consumed ~50M tokens (49.9–50.0M) of FineWeb-Edu at 23.08M parameters.

Headline results

  • Accuracy: reversibility is a wash. 3.797 vs 3.783 val loss — a 0.014-nat gap (1.4% perplexity) at matched parameters, matched token budget, each at its own tuned LR.
  • Memory: 1.84x less peak allocated (1157 vs 2130 MB); counting only activations — subtracting the 352 MB of params + grads + Adam moments that no architecture can avoid — it is 2.2x less.
  • Speed: ~1.5x slower, measured twice in opposite thermal regimes (1.53x heat-soaked, 1.52x in boost) so the number is not a thermal artifact.
  • Max batch: 56 → 192 (3.4x) on the same 6 GB card.
  • But max batch is the wrong way to spend a fixed token budget: 508 optimizer steps instead of 6103 costs 0.46 nats, and the LR was swept to prove that is not mistuning.
  • Tuning mattered ~26x more than the architecture. Re-tuning the baseline's LR alone moved val loss 0.37 nats — 26x the entire baseline-vs-reversible gap.

Notebooks

notebook what it is outputs committed?
notebooks/01_colab_train_50M.ipynb Self-contained training notebook. Writes model.py / train.py / data_prep.py into the runtime, pulls the FineWeb-Edu shard from the Hub, probes the largest batch that fits that GPU, then runs all three configs. Adds an fp16 + GradScaler path so it also works on a T4, where bf16 is unsupported. no — it is meant to be run
notebooks/02_results_and_analysis.ipynb The full report, re-derived from the raw artifacts. Every table, figure and ratio is computed from runs/*/metrics.json, runs/*/curve.npy and runs/clocks_*.log — nothing hard-coded. Needs only numpy/pandas/matplotlib, no GPU. yes — executed, 4 figures
notebooks/03_gradient_correctness.ipynb Is the reversible backward actually correct? Runs the fp32/CPU gradient check live, shows the autocast fix in model.py, the superseded runs from when it was broken, and the GPU/bf16 reconstruction measurements with trained weights. yes — executed on the 3060

Colab: https://colab.research.google.com/drive/1fpLLKz5KpJwemOxmJm6s811W-V_MWPgD — notebook 01 is generated by make_colab_notebook.py from the same source files the local runs used, so the Colab copy and the committed .py files cannot drift.

Provenance note. Every result in this README and in RESULTS.md comes from the local 3060 — all 46 committed metrics.json files record "gpu": "NVIDIA GeForce RTX 3060 Laptop GPU". The Colab notebook is the portable path for reproducing the experiment on a different GPU.

Full write-up with the reasoning behind every choice: REPORT.md. Auto-generated tables for all 46 runs: RESULTS.md.

loss curves


What is actually implemented

One model file, three memory strategies, identical parameter counts — d_model 512, 6 blocks, 8 heads, MLP 4x, RoPE, tied embeddings, vocab 8192, seq 512 → 23.08M params (18.88M non-embedding

  • 4.19M embedding).

baseline

Standard pre-norm transformer. Autograd stores every activation.

rev_euler — two-stream coupling (RevNet / Reformer)

a ← a + f(b)          f = attention block
b ← b + g(a)          g = MLP block

Exactly invertible from the final (a, b) pair, so nothing in between is stored. Structurally it is a semi-implicit Euler step of a 2D ODE. Note what the coupling costs: attention only ever reads stream b, and the MLP only reads a.

rev_mid — single-stream explicit midpoint / leapfrog ← the one that won

x_{i+1} = x_{i-1} + f_i(x_i),    x_{-1} = 0

Inverse: x_{i-1} = x_{i+1} − f_i(x_i) from any adjacent pair. Symmetric and 2nd-order, and — the practical point — it reuses the baseline's own full-width block verbatim. Only the skip connection moves, so there is no channel halving and no restriction on what attention can see.

rev_mid_seed / rev_euler_seed

Same, but the backward takes x_0 from the saved input instead of reconstructing it (§Reconstruction below). One extra stored activation; removes the block-0 gradient corruption.

Also in model.py: a fused chunked cross-entropy (FusedCE) that keeps logits memory O(chunk) in both directions — without it the 8192-vocab logits dominate the memory an activation saving is supposed to free.

Why midpoint beat Euler

10M-token selection grid, batch 16, val loss. Every architecture is bracketed on both sides of its optimum, so this is best-vs-best rather than best-vs-mistuned:

LR baseline rev_mid rev_euler
2e-4 5.169 — 5.147
4e-4 4.912 4.853 4.921
7.5e-4 4.754 4.672 4.980
1.5e-3 4.952 4.846 5.081
3e-3 5.181 — 5.269
6e-3 — — 5.894

4.672 vs 4.921 — and rev_mid is the only variant that also beats the baseline at this horizon. The likely reason is structural, not numerical: leapfrog keeps the baseline's block intact, while the two-stream coupling forces attention to read a stream that accumulates only MLP outputs.


Reproducing

Requirements

Python 3.11+, a CUDA GPU, and a CUDA build of PyTorch (2.5+; developed on 2.10.0+cu126 — install it from https://pytorch.org for your CUDA version, then pip install -r requirements.txt). ~2.5 GB of disk for the corpus. Only the trained tokenizer (data/bpe_8192.json) is committed, so tokenization is reproducible bit-for-bit; the 2.1 GB parquet shard and the .bin memmaps are gitignored and regenerated.

Data

data_prep.py expects one FineWeb-Edu parquet shard at data/fineweb-edu-00000.parquet (the Colab notebook downloads it for you):

python data_prep.py       # byte-level BPE to 8192 tokens, then train.bin (52M) / val.bin (5M) as uint16

The experiment

python gradcheck.py                              # fp32/CPU exactness of both reversible stacks + FusedCE
python revcheck.py --arch rev_mid --steps 300    # GPU/bf16 drift with TRAINED weights

python train.py --arch baseline --batch 16  --lr 7.5e-4 --tokens 50e6 --out runs/r1b   # run 1
python train.py --arch rev_mid  --batch 16  --lr 7.5e-4 --tokens 50e6 --out runs/r2    # run 2
python train.py --arch rev_mid  --batch 192 --lr 1.3e-3 --tokens 50e6 --out runs/r3    # run 3

python speed_ab.py --batch 16                    # interleaved A/B throughput
python thermal_bench.py                          # steady-state throughput + clock telemetry
python repeat_guarded.py                         # boost-regime cross-check (train.py --thermal-guard)
python report.py && python plot_curves.py        # regenerate RESULTS.md and curves_main.png

Find the largest batch for your card before run 3 — 192 is specific to 6 GB:

python train.py --arch rev_mid --probe-batch 192   # 3 steps, prints peak allocated/reserved, exits

Repo layout

path what
model.py all four archs, both reversible autograd.Functions, FusedCE, the autocast plumbing
train.py trainer: cosine schedule to 10% of peak, fused AdamW, grad-clip 1.0, bf16 autocast, --probe-batch, --thermal-guard
data_prep.py parquet → BPE tokenizer → train.bin / val.bin
gradcheck.py fp32/CPU gradient equivalence vs a store-everything reference
revcheck.py GPU/bf16 reconstruction + gradient error, at init and with trained weights
probe_mem.py, bench_batch.py memory probe; throughput-vs-batch sweep that detects driver spill
speed_ab.py, thermal_bench.py, repeat_guarded.py the three throughput measurements (naive interleaved, heat-soaked, boost-guarded)
report.py, plot_curves.py regenerate RESULTS.md and curves_main.png
driver.py, final_runs.sh, finals*.sh, sweep*.sh the actual chains that produced the runs, kept as a record
make_colab_notebook.py builds notebook 01 from model.py / train.py / data_prep.py
runs/ all 46 runs: metrics.json, curve.npy, stdout logs, clock telemetry, revcheck_*.json — including the superseded buggy_* runs
notebooks/ the three notebooks above

Findings worth reading before you trust reversibility

A silent, catastrophic bug: the autocast weight cache kills reversible gradients

This invalidated an entire first round of results. Under torch.autocast the bf16 copy of each weight is cached. A reversible backward recomputes blocks inside torch.autograd.grad, but the forward ran under torch.no_grad() — so the cached bf16 copy carries no grad_fn, the recompute cannot reach the fp32 parameter, and allow_unused=True returns None for every Linear weight in every block.

Result: all attention and MLP matrices were frozen at initialization. Only LayerNorm gains and the embedding trained. The loss still fell from 7.5 to 4.9, so nothing looked wrong — and gradcheck.py passed at 6e-7 the whole time, because it runs fp32 on CPU where there is no weight cache.

Correctness tests must run in the precision and on the device you actually train in.

The fix is one context manager, re-entering autocast with the cache disabled inside the backward:

with torch.autocast("cuda", dtype=dtype, cache_enabled=False), torch.enable_grad():
    ...

Best reversible result at 10M tokens, before and after the fix (best over both variants and the whole LR grid):

best val @10M best LR LR response
broken 5.267 (rev_euler) 6e-3 nearly flat — 0.02 nats over a 4x LR range
fixed 4.672 (rev_mid) 7.5e-4 sharp interior minimum — 0.18 nats over the same range

A flat LR response was the tell — a model that barely cares about its learning rate over a 4x range is usually not training. model.py now raises rather than silently accepting a None gradient for a block parameter. The bug also flipped the variant verdict: while broken, rev_euler appeared to beat rev_mid.

Reconstruction of the first activation is catastrophically ill-conditioned

Measured on GPU, in bf16, with trained weights (runs/revcheck_mid_trained.json): x_1…x_4 come back to ~1e-4–1e-3 relative error, and then x_0 fails outright — 87 relative, with the minimum gradient cosine similarity dropping to 0.15 (0.9997 at init, where everything is still small). REPORT.md §5b localises it: block 0 at cosine 0.054, blocks 1–5 at 0.998–1.000.

The cause is catastrophic cancellation, not rounding: x_0 is embedding-scale, but it is recovered by subtracting six blocks' worth of accumulated residual. Re-running the check in fp32 only improves block 0 to ~0.24 — precision does not fix an ill-conditioned subtraction.

It is cheap to fix, because x_0 is already an input to the reversible Function and can simply be stored (rev_mid_seed):

block-0 cosine max grad rel err val @10M peak alloc (b=16) max batch
rev_mid 0.054 1.20 4.6716 1157 MB 192
rev_mid_seed 0.9999 0.031 4.6731 1176 MB 176

And the surprise is that fixing it changes nothing. Restoring block 0's gradient from near-random to exact moves val loss by 0.0015 nats — noise — while costing ~179 MB at batch 192, enough to drop the max batch below 192. At 6 layers the first block does little enough work that a garbage gradient there is harmless. I would expect that to stop being true with depth.

The maximum batch is not how you spend a fixed token budget

Reversibility raises the largest batch that fits from 56 to 192, but at 50M tokens that means 508 optimizer steps instead of 6103, and the loss regresses 0.46 nats. The LR was properly bracketed:

batch 192, LR 6.5e-4 1.3e-3 2.6e-3 5.2e-3
val loss 4.469 4.254 4.402 4.989

The optimum is only ~1.7x the batch-16 LR for a 12x larger batch — square-root scaling (2.6e-3) overshoots, because at 508 steps the run is nowhere near convergence and a large LR mostly costs early stability. Big batches are worth having when memory is the binding constraint (longer sequences, a bigger model, avoiding gradient accumulation), not as a way to spend tokens faster.

Throughput on a laptop GPU needs a thermal method, or it means nothing

Under sustained load this chassis reaches 88–89 °C and the driver's SW thermal slowdown governor drops the SM clock from ~1880 MHz to ~800–1000 MHz, roughly halving throughput mid-run. Identical configs measured anywhere from 28.3k to 40.5k tok/s depending on how warm the card was at launch — so the tok/s column in the headline table is not an architecture comparison.

Two independent methods, from opposite ends of the thermal envelope:

method regime baseline / rev_mid
interleaved heat-soak (thermal_bench.py) throttled, pinned 1387 MHz 1.53x
boost-regime guard (repeat_guarded.py) unthrottled, ≥1702 MHz 1.52x

They agree to within 1%, so the ~1.5x cost is a property of the code, not of when a run started. It is more than the ~1.33x "one extra forward" theory predicts, because the backward re-runs each block under torch.autograd.grad with a fresh autocast cast per block, which does not fuse as well as a plain backward.

The guard also settles whether the 50M-token experiment could be run unthrottled here: it cannot. The boost regime lasts 23–27 s — under 4% of one run. Sustaining it means dissipating 114 W against the ~59 W this chassis holds at its 89 °C ceiling, a ~2x cooling deficit that a cooling pad's typical 3–8 °C does not bridge. The throttled numbers are the real sustained performance of this hardware. (--thermal-guard is a measurement-regime guard, not a hardware safety device; the card's own protection sits near 93 °C and was never approached. Its loss values at 1–2M tokens are meaningless.)

Two Windows-specific memory traps

  • There is no OOM at the ceiling. Past ~5120 MB reserved the driver falls back to shared system memory instead of raising. Batch 256 "succeeded" at 6382 MB reserved with throughput down ~20% — check max_memory_reserved against torch.cuda.mem_get_info() rather than trusting the absence of an exception.
  • expandable_segments:True is a no-op on Windows. It is the usual remedy for the ~1.2x gap between allocated and reserved; here it changed reserved by exactly 0 MB at every batch. That fragmentation gap is what actually caps the batch at 192 — allocated is only 4502 MB.

Caveats

  • One seed (1234) per configuration. The 0.014-nat baseline-vs-reversible gap is smaller than seed noise almost certainly is at this scale; read it as "no measurable difference", not as a ranking.
  • 6 layers, 23M params, 50M tokens. Two of the findings are explicitly depth-dependent — the harmless block-0 gradient, and the ~1.5x recompute overhead, which is amortized differently in deeper stacks.
  • The boost-regime runs are short (12–239 steps), so their tok/s carries some warm-up amortization; the batch-192 entry timed only 7 steps and is the noisiest number in the repo.
  • speed_ab_b16.txt is tail-truncated (the .json next to it captured a crash from a missing nvidia-ml-py); notebook 02 parses what survived, which is the per-round data that matters.
  • Deterministic clocks need an admin shell: nvidia-smi -lgc 1000,1000 to pin, nvidia-smi -rgc to restore. Everything here was measured without that, which is why the thermal methodology exists.

About

Reversible 20M-param LLM trained on 50M FineWeb-Edu tokens on a 6GB RTX 3060 Laptop: leapfrog-midpoint reversibility is accuracy-neutral, 2.2x less activation memory, 3.4x larger max batch, ~1.5x slower

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages