sparse-memory-lm B-16M
A small Llama-style language model with a product-key memory: a table of 16.8 million learned vectors (6.4B parameters), of which the model reads a few hundred per token. Per token it uses 33M parameters. It's the B-16M model from github.com/re133/sparse-memory-lm, where the code, the training setup and all the measurements are.
The point of uploading it: the table doesn't have to be in GPU memory. On my Radeon RX 9070 the model writes about 140 tokens/s with the table memory-mapped from an NVMe SSD, using 0.42 GiB of VRAM.
Files
| File | What | Size |
|---|---|---|
values_q4.bin |
the table, 4 bit, 16,777,216 rows x 384, two values per byte | 3.2 GB |
scales_q4.bin |
one fp16 scale per row | 34 MB |
rest.pt |
everything else (PyTorch state dict + config), fp32 | 284 MB |
hot_rows.npy |
rows sorted by how often they were read in training, for the RAM cache | 134 MB |
meta.json |
shapes, the 4-bit format, sha256 of every file |
The full fp32 table isn't uploaded (25.8 GB). The 4-bit table loses almost nothing: validation PPL 19.98 against 19.96.
Running it
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 # or the CUDA wheel
.venv/bin/pip install -r requirements.txt huggingface_hub
.venv/bin/python scripts/demo_generate.py --download --table nvme # or ram / vram, -i for your own prompts
--table nvme reads the table from the SSD, ram keeps it in RAM, vram on the GPU (~200 tokens/s on the
RX 9070). All three write the same text (with Triton 3.5 the logits are even bit-identical; with Triton 3.8 they can
differ in the last bits). It needs the Triton kernels from the repo, so it doesn't load with
transformers.
Training
- Data: 500M tokens of English Wikipedia (
wikimedia/wikipedia, 20231101.en), GPT-2 tokenizer, every article asTitle\n\nText. - Model: d=384, 12 layers, 6 heads, SwiGLU, RoPE, 1024-token context. Layers 3, 7 and 11 have a memory layer instead of the FFN (4 heads, top 32 of 4096² product keys, one table shared by the three layers).
- Hardware: one NVIDIA H200, ~3 hours, about 101 GB of GPU memory.
Results
| Val PPL (Wikipedia) | |
|---|---|
| Same model without the table (21M) | 25.67 |
| B-16M | 19.96 (fp32 table) / 19.98 (this 4-bit table) |
That's about as good as a dense model with ~114M parameters trained on the same data (between 106M and 123M).
Limitations
- Small model: it writes fluent Wikipedia-style English, but most facts in it are made up. It's a research model for measuring what such a table costs and brings, not for actually using it.
- No instruction tuning, no safety tuning. It just continues Wikipedia articles.
- Training data: Wikipedia text, licensed CC BY-SA 4.0. The weights themselves are Apache 2.0.