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 as Title\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.
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train fechyy/sparse-memory-lm-B-16M