Skip to content

Latest commit

 

History

4 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

HOPE: Nested Learning in PyTorch

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

CI arXiv PyTorch Python License

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 loss curve: HOPE 30.6M on FineWeb-Edu, 0.3B tokens, converging to 1.721 bits per byte

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.

Results

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):

8K extrapolation chart, bits per byte for each 1K chunk with the 2K training window shaded: 1.66, 1.81, then 1.85, 1.83, 1.86, 2.07, 1.79, 1.96 — flat within ~0.3 bpb of the training window, no blow-up, no horizon.

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.

Quickstart

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 -q

Memory-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).

What HOPE is (and what this repo implements)

HOPE as a stack of learners that differ only in when they write: SMT fast memory writes every token via the gated delta rule; CMS MLP levels pour an accumulated task-gradient bucket every 1, 2 and 4 chunks; AdamW writes after the sequence and teaches all lanes below it. Fast lanes remember this context; slow lanes remember the world.

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.

The part that makes it fast: exact chunk parallelism

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.grad call,
  • 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.py verifies 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.

How we trained and tested

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=off consumes 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.

BPE pretraining loss curve: HOPE 161.5M on FineWeb-Edu, Llama-2 tokenizer, 1B tokens

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.

How we tested

  1. 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.
  2. 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.
  3. 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.
  4. 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.
  5. 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.

Faithfulness: what we chose, what we fixed, what to check

  • 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 = NaN under multiplicative masking, and more. Each entry has symptom, mechanism, fix, and a diagnosis recipe.

A checklist for any HOPE implementation (including this one)

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

What this is not

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

Citations

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

License

MIT.

About

HOPE: Nested Learning in PyTorch (arXiv:2512.24695) — attention-free language model with self-modifying Titans and a Continuum Memory System

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages