Attention-free language model with test-time learning: Self-Modifying Titans + Continuum Memory System (CMS), from "Nested Learning: The Illusion of Deep Learning Architectures" (arXiv:2512.24695).
Quickstart · Results · How it works · Deviations · Stability
Nested Learning: The Illusion of Deep Learning Architectures Ali Behrouz, Meisam Razaviyayn, Peilin Zhong, Vahab Mirrokni (Google Research) Paper: https://arxiv.org/abs/2512.24695
No official code for HOPE has been released, so this implementation was built from what is available: the paper's mechanisms, turned into a checklist with the help of the existing public reproductions; the prior work the paper builds on (Titans, DeltaNet, TTT) for pieces it references but does not spell out; and our own experiments for what no paper specifies — each such choice is recorded in DEVIATIONS.md.
Byte-level pretraining on FineWeb-Edu, 0.3B bytes, 30.6M parameters, fully attention-free. Line = train bits/byte (100-step mean); dots = held-out validation, final 1.721 bpb. This is not a historical curve: it was produced by running scripts/train_byte.py in this repository end to end on a single GPU (~7 h), and it lands within 0.001 bpb of our original research run.
All numbers from the released recipe in this repo (scripts/train_byte.py, single GPU):
| Metric | What the task is | HOPE, 30.6M, byte-level, 0.3B tokens |
|---|---|---|
| Held-out validation | next-byte prediction on 2,048-byte sequences (the training context length) sampled from a held-out 2M-byte split | 1.721 bits/byte |
| Needle recall at 1K / 6.6K bytes | plant a fact once in the stream; when it recurs 1K or 6.6K bytes later, how much better does the model predict it than a control that never saw it? (tests what test-time writes actually stored) | +2.59 / +2.10 bpb (n=40, s.e. 0.04/0.07) |
| 8K-byte extrapolation (trained at 2K) | stream 4x the training length with the memories still updating; does loss stay flat past the training window, or blow up? | flat, no horizon (chart below) |
Recall is measured by scripts/eval_recall.py on the checkpoint the released
training recipe produces (see Quickstart). It plants a sentence carrying a high-entropy code in a
stream of held-out bytes, repeats the same sentence distance bytes later, and
scores how well the model predicts the code at its second occurrence. The
control stream is byte-identical except that the first occurrence is absent, so
the gain (control bpb minus treatment bpb on the code) isolates one thing:
whether the model still holds what it read distance bytes ago. A model with
no long-range memory scores ~0. The two
distances are the near/far ends of what the setup admits: 1K sits inside the 2K
training window, and 6.6K places the probe near the top of the 8K evaluation
stream (600-byte plant position + 6,600 + sentence + margin).
8K extrapolation, bpb per 1K chunk, no guards (the model keeps updating its memories through the whole stream, 4x past its training length):
No blow-up, no horizon — the curve stays within ~0.3 bpb of the training window out to 4x while the memories keep updating. This is a property of the paper's design: with no positional embedding, nothing in the model is tied to the training length.
Requirements: Linux, Python 3.10+, PyTorch 2.1+ with CUDA. The released runs used a single NVIDIA H200 (PyTorch 2.13, CUDA 13.0, bf16); any GPU with ~16 GB reproduces the byte recipe, and the exactness tests run on CPU.
git clone https://github.com/smallhours19/nested-learning-hope && cd nested-learning-hope
pip install -r requirements.txt
# 1) data: FineWeb-Edu sample shards (~2.2 GB each; 3 are plenty)
python scripts/download_data.py --out data/fineweb-edu --shards 3
# 2) sanity: exactness + causality tests (CPU, seconds)
pytest tests/test_exactness.py -q
# 3) train the released recipe: 30.6M params, 0.3B bytes, one GPU (~14 GB)
python scripts/train_byte.py --data data/fineweb-edu --out runs/hope-byte
# 4) evaluate: validation bpb + unguarded 8K extrapolation
python scripts/evaluate.py --ckpt runs/hope-byte/ckpt.pt --data data/fineweb-edu
# optional: 10-minute smoke — verifies your run tracks our logged trajectory
pytest tests/test_smoke.py -qMemory-form flag (scripts/train_byte.py):
| Flag | What it changes |
|---|---|
--deep-mem-learnable |
two-layer residual deep memory, all weights updated in-context (paper Eq. 91) — default is the matrix memory of Eq. 88-90, whose intra-chunk state has the exact closed form used here |
Both training scripts — train_byte.py above and the BPE scale-up
train_bpe.py (multi-GPU DDP, see "How we trained") — support --resume <ckpt>.
Checkpoints carry optimizer moments and RNG state; the BPE trainer additionally
writes per-rank RNG files (rng_rank{i}.pt) and fast-forwards the streaming
dataloader on resume, so a resumed run reproduces the uninterrupted run's losses
to bf16 noise. Checkpoint saves use no collectives (see the module docstring in
scripts/train_bpe.py for why that matters under DDP).
Every block stacks two mechanisms, both of which learn at test time — weights change during the forward pass, in eval as much as in training:
Self-Modifying Titans (SMT) (paper Eq. 83-93). An associative memory whose key/value/learning-rate/forget-gate projections are themselves small memories, updated online by gradient descent on a self-generated objective (k -> M(v), the model writes its own targets). Update strength is the gradient of the inner loss — the "surprise" signal of Titans (arXiv:2501.00663): tokens the memory already predicts write almost nothing; unexpected tokens write hard. Per-token gates alpha_t (retention) and eta_t (write strength) follow Eq. 90:
W <- W (alpha_t I - eta_t k_t k_t^T) - eta_t grad_W l_t
Continuum Memory System (paper Eq. 70-71). A chain of MLP levels where level l accumulates the task-loss gradient and applies it every C^(l) chunks — a spectrum of update frequencies between "context" and "weights". The paper is explicit that the teaching signal is the task objective ("for language modeling it is next token prediction"), not a local reconstruction loss; this distinction is easy to miss and changes what the levels learn.
Naive test-time learning runs a Python loop per token. The paper prescribes chunk-wise parallelism (§8); we implement it with the WY representation of products of (alpha I - eta k k^T) factors and decayed lower-triangular solves (the "UT transform" of Gated DeltaNet, arXiv:2412.06464), so that within a chunk of 512 tokens:
- all four self-modifying memories are updated with one
autograd.gradcall, - the main matrix memory's retrieval is per-token exact: token t reads a state containing all writes s < t of the same chunk, with zero lag —
tests/test_exactness.pyverifies the parallel path against a token-by-token sequential reference to ~1e-16 (fp64 machine precision).
This is not an optimization detail: a per-token Python loop is orders of magnitude slower, which in practice means such an implementation never gets validated at any real token count.
Data: FineWeb-Edu sample/100BT shards (Penedo et al., 2024) — the corpus
family the Titans/ATLAS/HOPE papers pretrain on. The byte run concatenates raw
UTF-8 with a NUL document separator (vocab 256, no tokenizer); the BPE run uses
the Llama-2 32K tokenizer with EOS separators.
Every value below is a script default: running the commands in the Quickstart reproduces these runs exactly. Both are single-node.
byte run (scripts/train_byte.py) |
BPE run (scripts/train_bpe.py) |
|
|---|---|---|
| Model | ||
| parameters | 30.6M | 161.5M |
| layers / width | 4 / 512 | 8 / 768 |
| vocabulary | 256 (bytes) | 32,000 (Llama-2 BPE) |
| chunk length | 512 | 1024 |
CMS periods C^(l) |
(1, 2, 4) chunks | (1, 2, 4) chunks |
| memory hidden mult. | 2 | 2 |
| positional embedding | none | none |
| Optimization | ||
| optimizer | AdamW, betas (0.9, 0.95) | AdamW, betas (0.9, 0.95) |
| peak learning rate | 6e-4 | 4e-4 |
| schedule | warmup 300, cosine to 0 | warmup 250, cosine |
| weight decay | 0.1 | 0.1 |
| gradient clip | 1.0 (non-finite steps skipped) | 1.0 (non-finite steps skipped) |
| precision | bf16 autocast, fp32 master | bf16 autocast, fp32 master |
| Batching | ||
| sequence length | 2048 | 4096 |
| micro-batch x accum x GPUs | 4 x 2 x 1 | 4 x 16 x 2 |
| tokens / step | 16,384 | 524,288 |
| steps / total tokens | 18,310 / 0.3B | 1,907 / 1.0B |
| Online update settings | ||
update_clip (relative) |
0.5 | 0.5 |
delta_retention (Eq. 90) |
on | on |
norm_target / norm_out |
on / on | on / on |
cms_second_order |
off (first-order consumption) | off |
max_updates guard |
none (unguarded) | none |
| Cost | ~7 h, 13.8 GB peak, 1 GPU | ~31 h at 9.0k tok/s, 2 GPUs |
| Result | 1.721 bpb held-out | 4.395 nats held-out (ppl 81.0) |
Provenance of the BPE result: 4.395 nats was measured on the released final
checkpoint (step 1907) with the training-time validation protocol (last shard,
20 batches x 4 x 4096); the committed assets/reference_run_bpe.log shows the
trajectory, whose last inline validation is 4.408 at step 1750.
Notes on two entries that mattered more than expected:
- The BPE run enables activation checkpointing and
skip_logits(logits are materialized only for the loss) — at 32K vocab the full logit tensor otherwise dominates memory. cms_second_order=offconsumes CMS task gradients first-order. It cuts peak memory ~3x at these scales with no measurable loss penalty; turn it on to let outer training shape the CMS update path through second-order terms.- We also learned the hard way that the surrounding recipe is not incidental: an earlier release candidate that differed only in lr (4e-4), schedule floor, and data sampling landed at 1.893 bpb — 0.17 worse — with an identical model.
The BPE scale-up run (scripts/train_bpe.py, Llama-2 32K tokenizer, 4K context, 2-GPU DDP). At equal tokens mid-run (0.52B) it tied a hybrid (attention + neural memory) baseline we trained identically (4.557 vs 4.553 nats) — the paper's competitive-with-hybrids claim, reproduced. In the byte regime the same hybrid leads by ~0.3 bpb on mean loss; under BPE that deficit vanishes — tokenization changes which architecture family wins.
- Exactness (
tests/test_exactness.py): parallel state/retrieval vs sequential Eq. 88-90 reference, fp64, asserts < 1e-10 (measured ~1e-16); causality under future-token perturbation (asserts bit-exact past logits); outer-loop gradient reaches the gate memories. - Trajectory smoke (
tests/test_smoke.py): 200 training steps on real data must land within +/-0.5 bpb of the original-run trajectory (5.083 @ step 100, 3.673 @ step 200). The committed re-run log (assets/reference_run.log) lands at 5.083 / 3.676 — step 100 to three decimals, step 200 within 0.003. This band catches every silent failure mode we met: a frozen memory, a 1/T-scaled write, a missing retention term. - End-to-end reproduction: the full 0.3B recipe was re-run from this repository as a release gate; final held-out 1.721 bpb vs 1.722 in the original research run, and the unguarded 8K extrapolation table above comes from that checkpoint via
scripts/evaluate.py. - Step-count gates during development: every change was compared against baselines at step 2000 on identical data; regressions were treated as implementation bugs, not model properties.
- Adversarial stability: NaN incidents were reproduced from dumped checkpoints with the triggering inputs (binary blobs, quasi-random bytes) and each fix verified 10/10 on the reproducer. The full catalog: docs/STABILITY.md.
- docs/DEVIATIONS.md — every point where we had to choose beyond the paper's text (the paper elides normalizations "for the sake of clarity", never defines
eta^(l)or the momentum coefficient, and leaves the deep-memory retention distribution undefined), what we chose, why, and which prior work each choice comes from (DeltaNet, Gated DeltaNet, TTT, Titans, RMSNorm, LoRA-style zero-init). - docs/STABILITY.md — 14 stability devices that appear in no paper but without which this architecture silently fails or NaNs: the zero-init clip lockout, the alpha-floor / fp32-overflow chain, beta <= 1 on entropy-tail inputs,
Inf * 0 = NaNunder multiplicative masking, and more. Each entry has symptom, mechanism, fix, and a diagnosis recipe.
Before we wrote a line of code we read the public HOPE/Nested-Learning reproductions — all volunteer work, and we are grateful for them; several of the questions below only became visible because someone else had already written the code that made them concrete. What we found is that the paper has a handful of mechanisms that are easy to transcribe and easy to get subtly wrong, and that a model can train, converge, and look healthy with any of them broken. So here is the checklist we ended up using. Run it against this repo too.
| Mechanism (paper) | this repo | how to check it in any implementation |
|---|---|---|
| All self-modifying memories actually updated (Eq. 83-86) | yes | are M_k, M_v, M_eta, M_alpha in the update dispatch, or only declared? |
| Self-generated value targets (Eq. 84-85) | yes | is the target M(v) produced by the memory itself, with stop-gradient? |
Queries non-adaptive (prose; the paper's equations include q in the update set — see DEVIATIONS.md D11) |
yes (prose reading) | decide which reading you implement, and say so |
Delta retention W(aI - eta kk^T) (Eq. 88-90) |
yes | is aI present, applied to the right factor, and is eta large enough to matter? A doubled decay term or eta ~ 1e-3 silently degrades this to scalar decay |
Token-wise gates alpha_t, eta_t |
yes | per-token, or one constant per head/layer? |
| Momentum on updates | yes | implemented and enabled by default? |
CMS periodic accumulation C^(l) (Eq. 71) |
yes | is the gradient accumulated across the period, or discarded and recomputed? |
| CMS taught by the task loss | yes | autograd.grad(task_loss, level_params) — not a local reconstruction/Hebbian objective |
| Two-layer residual memory (Eq. 91) | yes* | is it x + W1 sigma(W2 x), and — more importantly — are those weights learned by outer pretraining? If so, the memory starts each sequence full of general knowledge and test-time writes are marginal edits on top of it |
| conv-4 + l2-normalized q/k | yes | normalization is often implemented but switched off in configs |
| Exact chunk-parallel dual form (§8) | yes | a per-token Python loop is ~100x slower — which usually means the implementation was never validated at a real token count |
| No positional embeddings | yes | order must come from the recurrence (plus the conv-4), not from a position table |
Memory updates inside forward() |
yes | if updates are injected by the training loop instead, the trained weights are useless on their own — check by evaluating a checkpoint with no adaptation |
| Published training curves | yes | the ultimate check: does the repo show what its own code produces? |
* This is the one row where we deviate on purpose, and we ran the experiment
that tests it. The paper's Eq. 91 memory is a two-layer residual MLP; we default
to a plain, zero-initialized matrix memory M(x) = x + W x and ship the paper's
form behind deep_mem_learnable=True. Trained under identical conditions, the
two land here: matrix 1.721 bits/byte vs deep 1.743 (matrix ahead, while deep
carries 14% more parameters), but on needle recall deep wins — +2.94 vs +2.59 at 1K and
+2.48 vs +2.10 at 6.6K, several standard errors apart. We had predicted the
opposite and were wrong. Matrix stays the default because Eq. 88-90 are derived
in that setting and our exactness guarantee holds only there — not because it is
the better memory. DEVIATIONS.md D1 has the full table and
the caveats.
The last two rows are the ones we would check first, and they are related. The most dangerous failure mode of this architecture is a memory update that lives in the training loop rather than in the model's forward pass: training loss looks healthy, nothing crashes, and the standalone checkpoint has no working memory at all. That failure mode is why this repository leads with reproducible curves rather than with claims — and why the smoke test asserts a trajectory, not just "loss goes down".
- Not official. Independent reproduction; no affiliation with the authors. Where our reading of an under-specified detail differs from their intent, DEVIATIONS.md is the contract.
- Not a state-of-the-art claim. 30.6M and 161.5M parameters are reproduction scales, chosen so the runs fit on hardware you are likely to have. The paper's own experiments go to 1.3B.
- Not a general-purpose win over attention. Averaged over bytes, this attention-free model trails a hybrid baseline we trained identically by ~0.3 bpb; it wins on long-range recall and on length extrapolation. Pick it for what it is good at.
- No pretrained weights are distributed (yet). Everything in the results table is reproducible from this repo in a few GPU-hours; if that changes we will link releases here.
- Not tuned. No hyperparameter search was run — the recipe follows the paper family's published conventions, and the numbers are what that recipe gives.
Third-party reproductions are the most valuable contribution. If you run
scripts/train_byte.py and land outside the smoke-test band, or disagree with a
call in DEVIATIONS.md, open an issue with your log — mismatches are how the
stability catalog got written in the first place.
If you use this code, please cite the paper and this repository:
@article{behrouz2025nested,
title = {Nested Learning: The Illusion of Deep Learning Architectures},
author = {Behrouz, Ali and Razaviyayn, Meisam and Zhong, Peilin and
Mirrokni, Vahab},
journal = {arXiv preprint arXiv:2512.24695},
year = {2025}
}
@software{hope_nested_learning,
author = {smallhours19},
title = {HOPE: Nested Learning in PyTorch},
year = {2026},
url = {https://github.com/smallhours19/nested-learning-hope},
note = {Unofficial reproduction of arXiv:2512.24695 at reduced scale}
}Design choices in this implementation additionally build on: Titans (arXiv:2501.00663) · DeltaNet parallelization (arXiv:2406.06484) · Gated DeltaNet (arXiv:2412.06464) · Test-Time Training / dual form (arXiv:2407.04620) · RMSNorm (arXiv:1910.07467) · FineWeb-Edu (arXiv:2406.17557).
MIT.

