A 21M-parameter LM with a 16.8M-row product-key memory: as good as a 114M dense model and runs with the table on an SSD. Triton kernels for ROCm and CUDA.
Python
1
112 commits
updated Oct 6, 2026
What happens if you give a tiny language model a really big lookup table?
I trained a small Llama-style model (21M parameters) and added a product-key memory to it: a table with up to 16.8 million learned vectors, of which the model only reads a few hundred per token. Then I measured what that is actually worth, what it costs, and whether the table even has to sit in GPU memory. Most of this ran on my own PC (Radeon RX 9070). The big runs were on rented cloud GPUs, about 70 dollars in total.
Before every run I wrote down what would count as a success. Some things worked out, some didn't. Both are in here. The full lab notebook with every criterion, every number and every mishap is REPORT.md (in German).
Look inside the table · Model on Hugging Face · Try it on your GPU · I'm looking for bigger GPUs
Look inside the table: on this page you can click any word of a Wikipedia text and see which of the 1M entries of B-1M the model reads for it, browse a map of the whole table, and look at single entries: where they get read and what their 384 numbers look like.
About five minutes, most of it is the PyTorch download. Nothing else to download: the tests and the benchmark use
random data. Checked from a fresh clone on my RX 9070 with Python 3.14 and the official rocm7.2 wheel:
git clone https://github.com/re133/sparse-memory-lm.git && cd sparse-memory-lm
python3 -m venv .venv
.venv/bin/pip install torch --index-url https://download.pytorch.org/whl/rocm7.2
.venv/bin/pip install -r requirements.txt
.venv/bin/python -m pytest -q tests # ~1 min
.venv/bin/python scripts/kernel_speedup.py # ~40 s
What the last command printed on my card (the first run also shows a few gcc warnings while Triton compiles its launchers, those are harmless):
AMD Radeon RX 9070 (gfx1201), 16 GiB | PyTorch 2.14.1+rocm7.2 (HIP 7.2.53211) | Triton 3.8.0 | Python 3.14.7
train step train tok/s prefill tok/s decode tok/s peak GiB
no table (A) 333 ms 98,341 438,762 199 7.9
B-1M-sparse, PyTorch 787 ms 41,652 168,342 161 10.5
B-1M-sparse, Triton kernels 512 ms 64,027 295,060 203 10.5
kernels vs PyTorch: training 1.54x, prefill 1.75x, decode 1.26x
with the kernels the table model trains at 65% of the speed of the model without a table
rocm7.2 wheel. With rocm7.1 one test aborts on my card
(docs/rocm-issues). On the MI350X I used rocm7.1 with Python 3.12, and there
everything passed.python-pytorch-rocm), create
the venv with --system-site-packages and skip the torch line.pip install torch. The tests passed on an H200; the benchmark script itself I've only run on
the RX 9070 so far.The trained B-16M model is on Hugging Face: the 16.8M-row table in 4 bit (3.2 GB) plus the rest of the model. The 4-bit table scores 19.98 validation PPL, against 19.96 for the full fp32 one.
.venv/bin/pip install huggingface_hub
.venv/bin/python scripts/demo_generate.py --download --table nvme # 3.7 GB into data/tables/B-16M
.venv/bin/python scripts/demo_generate.py --table vram -i # table in VRAM, your own prompts
With the rocm7.2 wheel from above (other PyTorch builds sample a different text):
loaded in 1.5 s, table in nvme: 0.27 GiB VRAM, 1.8 GiB RAM + 3.8 GiB mapped files
Isaac Newton
Sir Isaac Newton was born on 14 November 1803, the son of the Rev. Samuel Newton of Basing, Middlesex, and his
wife Elizabeth, daughter of William Wilberforce of Westmorland. He was educated at Harrow and Trinity College,
Cambridge. He matriculated at Magdalen College, Oxford, on 11 June 1841. He was ordained in 1844. [...]
[200 tokens in 1.48 s = 135 tok/s, table in nvme, peak 0.42 GiB VRAM, 2.5 GiB RAM + 3.9 GiB mapped files]
--table ram and --table vram (~200 tok/s) write the same text. With the pip wheel
the logits can differ in the last bits between the modes (see above), so a sampled text could in rare cases go
its own way."Title\n\nFirst words".All models saw the same 500M Wikipedia tokens and are scored on the same held-out Wikipedia articles. The memory models share one table between three memory layers.
| Model | Params used per token | Table | Val PPL | As good as a dense model with |
|---|---|---|---|---|
| A (no table) | 21M | 25.67 | ||
| B-1M | 23M | 0.4B | 21.84 | ~60M params |
| B-4M | 26M | 1.6B | 20.80 | ~83M params |
| B-16M | 33M | 6.4B | 19.96 | ~114M params (106 to 123M) |
The "as good as" column comes from dense models with 50M, 100M, 200M and 400M parameters (non-embedding) that I trained on exactly the same data, interpolated on a log-log curve. The range in brackets shows how far the value moves if every perplexity is off by ±0.4% (the seed noise I measured on the 1M model). It's a sensitivity range, not a confidence interval.

Catches:
B-16M with the table in VRAM, in RAM, or as a 4-bit file on an NVMe SSD (Samsung 990 PRO, memory-mapped, with a RAM cache for the rows that get read most):
| Table in | Writing (batch 1) | Reading a long prompt | VRAM used |
|---|---|---|---|
| VRAM | 212 to 216 tok/s | 108,000 tok/s | 3.6 GiB (4-bit) / 13.1 GiB (bf16) |
| RAM | 139 to 154 tok/s | 19,000 to 62,000 tok/s | 0.5 GiB |
| NVMe | 114 to 138 tok/s | 1,500 to 6,500 tok/s | 0.5 GiB |
All three compute the same: same perplexity down to the last digit, and with my PyTorch build (Triton 3.5) the
logits are bit-identical too (check). With the pip wheel (Triton 3.8) the
kernel sums in a slightly different order for the small staging table than for the full one, so logits can differ
in the last bits. Both are equally close to an exact fp64 sum, and the generated tokens stayed the same
(check). VRAM is the whole card as rocm-smi reports it, desktop
included. When writing, the slow part isn't
the SSD but the three round trips between GPU and CPU per token: even with no cache at all the NVMe version is only
18% slower than RAM. When reading a prompt every token needs ~270 different rows, and every missed row costs a
whole 4 KB page from the SSD. That's where it falls apart.

Setup:
| Q | Q + table | Q + dense | |
|---|---|---|---|
| PPL on new, held-out articles | 12.98 | 10.09 | 10.01 |
| PPL on the training articles | 13.36 | 5.64 | 9.38 |
| Fact test, training articles (exact fill-in) | 4.8% | 10.2% | 7.8% |
| Fact test, articles never seen | 4.0% | 10.2% | 8.0% |
| MMLU | 49.7 | 47.0 | 49.6 |
What came out:
Base model: Llama-style decoder, d=384, 12 layers, 6 heads, SwiGLU, RoPE, RMSNorm, GPT-2 tokenizer
(smlm/model.py).
Memory layers: in layers 3, 7 and 11 the FFN is replaced by a memory layer (smlm/pkm.py), following Lample et
al. 2019 and Meta's Memory Layers at Scale.
Training the table: row-sparse gradients and a lazy Adam that only touches the rows read in a step
(smlm/sparse_values.py). Same quality as dense Adam on the table, faster, and much less memory.
Triton kernels (smlm/kernels.py), on ROCm and CUDA, matching the PyTorch reference up to float rounding:
On the RX 9070 this made training 1.47x faster and brought batch-1 decoding to the speed of the plain model (the table model replays its memory layers as graphs, the plain model runs without graphs). What each kernel does, how much it brings and where its limits are: docs/kernels.md.
Table outside the GPU: smlm/offload.py.
Qwen add-on: smlm/qwen_memory.py.
Set up the venv as in Try it, then:
.venv/bin/python cloud/fetch_data.py # WikiText-103 + Wikipedia at pinned revisions, sha256-checked
# plain model and B-1M on 500M Wikipedia tokens
.venv/bin/python -m smlm.train --model A --out_dir runs/A --data wikipedia --tokens 500e6 --extra_val wikitext103
.venv/bin/python -m smlm.train --model B-1M-sparse --mem_impl triton --value_lr 2.4e-3 --out_dir runs/B-1M \
--data wikipedia --tokens 500e6 --extra_val wikitext103
Bigger runs:
B-4M-sparse and B-16M-sparse need ~29 GB and ~101 GB of GPU memory.D-50M … D-400M.cloud/ and
scripts/run_*.py (notes in German: docs/notes/CLOUD.md).Table outside the GPU: scripts/convert_table.py turns a checkpoint into bf16 / 4-bit table files, and
scripts/bench_offload.py runs the measurements above.
Qwen experiment:
pip install -r requirements-qwen.txt.scripts/prepare_qwen_data.py. Fact test: scripts/make_fact_cloze.py. Training:
scripts/train_qwen_memory.py. Evaluation: scripts/eval_*.py, scripts/qwen_step3_eval.py.models/Qwen3.5-0.8B or QWEN_DIR, data in data/qwen_wiki or QWEN_DATA.I also ran the whole thing on an AMD Instinct MI350X (288 GB, rented for ~3 dollars):
Details: REPORT.md, section "AMD Instinct MI350X".

Everything was built and tested on an RX 9070 (gfx1201, ROCm 7.2, Triton 3.5) and gives the same results on an MI350X (gfx950, ROCm 7.1, Triton 3.7), an H100 and an H200. Two things to know on consumer AMD cards:
tl.atomic_add on fp32 compiles to the native instruction, but it's about 8x slower than a
plain store (repro). Sorting by row and only using atomics at program
borders took one kernel from 17 ms to 2.7 ms.More in docs/rocm-issues.
Everything here ran on one gaming GPU and about 70 dollars of rented cloud time. The question I'd really like to answer is whether the table still pays off at a size people actually use: a model around 1B parameters, a table with tens of millions of rows, trained on tens of billions of tokens. That's beyond what I can rent. It needs a multi-GPU machine for a good while, and the table with its optimizer state doesn't fit on one GPU anymore (B-16M already needed 101 GB).
What I'd do with more compute:
If you have GPUs to spare, AMD Instinct or anything else, I'd love to hear from you. The kernels already run on MI350X, H100 and H200 without changes. I'd run it the same way as here: success criteria written down before each run, and the results published whatever they turn out to be.
Contact: fechner.leon [at] protonmail.com, Discord fechyyyyy, or open an issue in this repo.
A 21M-parameter LM with a 16.8M-row product-key memory: as good as a 114M dense model and runs with the table on an SSD. Triton kernels for ROCm and CUDA.
Python
1
112 commits
updated Oct 6, 2026
What happens if you give a tiny language model a really big lookup table?
I trained a small Llama-style model (21M parameters) and added a product-key memory to it: a table with up to 16.8 million learned vectors, of which the model only reads a few hundred per token. Then I measured what that is actually worth, what it costs, and whether the table even has to sit in GPU memory. Most of this ran on my own PC (Radeon RX 9070). The big runs were on rented cloud GPUs, about 70 dollars in total.
Before every run I wrote down what would count as a success. Some things worked out, some didn't. Both are in here. The full lab notebook with every criterion, every number and every mishap is REPORT.md (in German).
Look inside the table · Model on Hugging Face · Try it on your GPU · I'm looking for bigger GPUs
Look inside the table: on this page you can click any word of a Wikipedia text and see which of the 1M entries of B-1M the model reads for it, browse a map of the whole table, and look at single entries: where they get read and what their 384 numbers look like.
About five minutes, most of it is the PyTorch download. Nothing else to download: the tests and the benchmark use
random data. Checked from a fresh clone on my RX 9070 with Python 3.14 and the official rocm7.2 wheel:
git clone https://github.com/re133/sparse-memory-lm.git && cd sparse-memory-lm
python3 -m venv .venv
.venv/bin/pip install torch --index-url https://download.pytorch.org/whl/rocm7.2
.venv/bin/pip install -r requirements.txt
.venv/bin/python -m pytest -q tests # ~1 min
.venv/bin/python scripts/kernel_speedup.py # ~40 s
What the last command printed on my card (the first run also shows a few gcc warnings while Triton compiles its launchers, those are harmless):
AMD Radeon RX 9070 (gfx1201), 16 GiB | PyTorch 2.14.1+rocm7.2 (HIP 7.2.53211) | Triton 3.8.0 | Python 3.14.7
train step train tok/s prefill tok/s decode tok/s peak GiB
no table (A) 333 ms 98,341 438,762 199 7.9
B-1M-sparse, PyTorch 787 ms 41,652 168,342 161 10.5
B-1M-sparse, Triton kernels 512 ms 64,027 295,060 203 10.5
kernels vs PyTorch: training 1.54x, prefill 1.75x, decode 1.26x
with the kernels the table model trains at 65% of the speed of the model without a table
rocm7.2 wheel. With rocm7.1 one test aborts on my card
(docs/rocm-issues). On the MI350X I used rocm7.1 with Python 3.12, and there
everything passed.python-pytorch-rocm), create
the venv with --system-site-packages and skip the torch line.pip install torch. The tests passed on an H200; the benchmark script itself I've only run on
the RX 9070 so far.The trained B-16M model is on Hugging Face: the 16.8M-row table in 4 bit (3.2 GB) plus the rest of the model. The 4-bit table scores 19.98 validation PPL, against 19.96 for the full fp32 one.
.venv/bin/pip install huggingface_hub
.venv/bin/python scripts/demo_generate.py --download --table nvme # 3.7 GB into data/tables/B-16M
.venv/bin/python scripts/demo_generate.py --table vram -i # table in VRAM, your own prompts
With the rocm7.2 wheel from above (other PyTorch builds sample a different text):
loaded in 1.5 s, table in nvme: 0.27 GiB VRAM, 1.8 GiB RAM + 3.8 GiB mapped files
Isaac Newton
Sir Isaac Newton was born on 14 November 1803, the son of the Rev. Samuel Newton of Basing, Middlesex, and his
wife Elizabeth, daughter of William Wilberforce of Westmorland. He was educated at Harrow and Trinity College,
Cambridge. He matriculated at Magdalen College, Oxford, on 11 June 1841. He was ordained in 1844. [...]
[200 tokens in 1.48 s = 135 tok/s, table in nvme, peak 0.42 GiB VRAM, 2.5 GiB RAM + 3.9 GiB mapped files]
--table ram and --table vram (~200 tok/s) write the same text. With the pip wheel
the logits can differ in the last bits between the modes (see above), so a sampled text could in rare cases go
its own way."Title\n\nFirst words".All models saw the same 500M Wikipedia tokens and are scored on the same held-out Wikipedia articles. The memory models share one table between three memory layers.
| Model | Params used per token | Table | Val PPL | As good as a dense model with |
|---|---|---|---|---|
| A (no table) | 21M | 25.67 | ||
| B-1M | 23M | 0.4B | 21.84 | ~60M params |
| B-4M | 26M | 1.6B | 20.80 | ~83M params |
| B-16M | 33M | 6.4B | 19.96 | ~114M params (106 to 123M) |
The "as good as" column comes from dense models with 50M, 100M, 200M and 400M parameters (non-embedding) that I trained on exactly the same data, interpolated on a log-log curve. The range in brackets shows how far the value moves if every perplexity is off by ±0.4% (the seed noise I measured on the 1M model). It's a sensitivity range, not a confidence interval.

Catches:
B-16M with the table in VRAM, in RAM, or as a 4-bit file on an NVMe SSD (Samsung 990 PRO, memory-mapped, with a RAM cache for the rows that get read most):
| Table in | Writing (batch 1) | Reading a long prompt | VRAM used |
|---|---|---|---|
| VRAM | 212 to 216 tok/s | 108,000 tok/s | 3.6 GiB (4-bit) / 13.1 GiB (bf16) |
| RAM | 139 to 154 tok/s | 19,000 to 62,000 tok/s | 0.5 GiB |
| NVMe | 114 to 138 tok/s | 1,500 to 6,500 tok/s | 0.5 GiB |
All three compute the same: same perplexity down to the last digit, and with my PyTorch build (Triton 3.5) the
logits are bit-identical too (check). With the pip wheel (Triton 3.8) the
kernel sums in a slightly different order for the small staging table than for the full one, so logits can differ
in the last bits. Both are equally close to an exact fp64 sum, and the generated tokens stayed the same
(check). VRAM is the whole card as rocm-smi reports it, desktop
included. When writing, the slow part isn't
the SSD but the three round trips between GPU and CPU per token: even with no cache at all the NVMe version is only
18% slower than RAM. When reading a prompt every token needs ~270 different rows, and every missed row costs a
whole 4 KB page from the SSD. That's where it falls apart.

Setup:
| Q | Q + table | Q + dense | |
|---|---|---|---|
| PPL on new, held-out articles | 12.98 | 10.09 | 10.01 |
| PPL on the training articles | 13.36 | 5.64 | 9.38 |
| Fact test, training articles (exact fill-in) | 4.8% | 10.2% | 7.8% |
| Fact test, articles never seen | 4.0% | 10.2% | 8.0% |
| MMLU | 49.7 | 47.0 | 49.6 |
What came out:
Base model: Llama-style decoder, d=384, 12 layers, 6 heads, SwiGLU, RoPE, RMSNorm, GPT-2 tokenizer
(smlm/model.py).
Memory layers: in layers 3, 7 and 11 the FFN is replaced by a memory layer (smlm/pkm.py), following Lample et
al. 2019 and Meta's Memory Layers at Scale.
Training the table: row-sparse gradients and a lazy Adam that only touches the rows read in a step
(smlm/sparse_values.py). Same quality as dense Adam on the table, faster, and much less memory.
Triton kernels (smlm/kernels.py), on ROCm and CUDA, matching the PyTorch reference up to float rounding:
On the RX 9070 this made training 1.47x faster and brought batch-1 decoding to the speed of the plain model (the table model replays its memory layers as graphs, the plain model runs without graphs). What each kernel does, how much it brings and where its limits are: docs/kernels.md.
Table outside the GPU: smlm/offload.py.
Qwen add-on: smlm/qwen_memory.py.
Set up the venv as in Try it, then:
.venv/bin/python cloud/fetch_data.py # WikiText-103 + Wikipedia at pinned revisions, sha256-checked
# plain model and B-1M on 500M Wikipedia tokens
.venv/bin/python -m smlm.train --model A --out_dir runs/A --data wikipedia --tokens 500e6 --extra_val wikitext103
.venv/bin/python -m smlm.train --model B-1M-sparse --mem_impl triton --value_lr 2.4e-3 --out_dir runs/B-1M \
--data wikipedia --tokens 500e6 --extra_val wikitext103
Bigger runs:
B-4M-sparse and B-16M-sparse need ~29 GB and ~101 GB of GPU memory.D-50M … D-400M.cloud/ and
scripts/run_*.py (notes in German: docs/notes/CLOUD.md).Table outside the GPU: scripts/convert_table.py turns a checkpoint into bf16 / 4-bit table files, and
scripts/bench_offload.py runs the measurements above.
Qwen experiment:
pip install -r requirements-qwen.txt.scripts/prepare_qwen_data.py. Fact test: scripts/make_fact_cloze.py. Training:
scripts/train_qwen_memory.py. Evaluation: scripts/eval_*.py, scripts/qwen_step3_eval.py.models/Qwen3.5-0.8B or QWEN_DIR, data in data/qwen_wiki or QWEN_DATA.I also ran the whole thing on an AMD Instinct MI350X (288 GB, rented for ~3 dollars):
Details: REPORT.md, section "AMD Instinct MI350X".

Everything was built and tested on an RX 9070 (gfx1201, ROCm 7.2, Triton 3.5) and gives the same results on an MI350X (gfx950, ROCm 7.1, Triton 3.7), an H100 and an H200. Two things to know on consumer AMD cards:
tl.atomic_add on fp32 compiles to the native instruction, but it's about 8x slower than a
plain store (repro). Sorting by row and only using atomics at program
borders took one kernel from 17 ms to 2.7 ms.More in docs/rocm-issues.
Everything here ran on one gaming GPU and about 70 dollars of rented cloud time. The question I'd really like to answer is whether the table still pays off at a size people actually use: a model around 1B parameters, a table with tens of millions of rows, trained on tens of billions of tokens. That's beyond what I can rent. It needs a multi-GPU machine for a good while, and the table with its optimizer state doesn't fit on one GPU anymore (B-16M already needed 101 GB).
What I'd do with more compute:
If you have GPUs to spare, AMD Instinct or anything else, I'd love to hear from you. The kernels already run on MI350X, H100 and H200 without changes. I'd run it the same way as here: success criteria written down before each run, and the results published whatever they turn out to be.
Contact: fechner.leon [at] protonmail.com, Discord fechyyyyy, or open an issue in this repo.