I set out to implement TurboQuant (PolarQuant + QJL) for Gemma 4 31B's KV cache — a 31 billion parameter model running on a single Mac. It doesn't work on this model. What I built instead is faster.
The TurboQuant paper (ICLR 2026) proposes PolarQuant + QJL for KV cache compression. I implemented both as Metal compute shaders and tested them on the full Gemma 4 31B (4-bit, 17.4 GB) across 177 experiments.
PolarQuant encodes key vectors as recursive polar coordinates (radius + angles). On most models this achieves ~7x compression with high quality. On Gemma 4, it produces gibberish. The reason: Gemma 4 uses attn_scale=1.0 (the Q/K norms handle magnitude instead), which means attention scores are not dampened — a 4% directional error from PolarQuant gets amplified through softmax and compounds across 60 layers.
QJL (1-bit quantized Johnson-Lindenstrauss) is supposed to correct PolarQuant's residual error. On Gemma 4, it makes quality worse — +6.5% perplexity. The 1-bit variance adds noise to an already sensitive attention mechanism.
MLX's native int4 quantization (per-group scale + bias) works fine at 6.4x compression — close to PolarQuant's 7.1x — because it preserves the exact direction of key vectors better than angular encoding.
Since int4 KV cache works but the standard attention path wastes bandwidth dequantizing keys every token, I wrote a fused Metal kernel (sdpa_int4) that computes attention scores directly from the int4 packed data. No dequantization. No intermediate float32 matrices. Everything happens in GPU registers.
Standard: dequantize(K_int4) → K_float32 [780MB temporary] → SDPA → output
Ours: sdpa_int4(Q, K_int4, V_int4) → output [0 bytes temporary]
The kernel uses MLX's own qdot vectorization pattern: pre-divide query values by {1, 16, 256, 4096}, multiply against uint16 masks {0xF, 0xF0, 0xF00, 0xF000}, accumulate with online softmax. One Metal dispatch per attention head group, with SIMD-parallel reduction across simdgroups.
The baseline slows down because dequantizing longer key sequences costs more bandwidth. The fused kernel reads int4 data directly — its cost barely changes with context length.
| Context | Fused Kernel | Baseline (deq + SDPA) | Speedup |
|---|---|---|---|
| 33 tokens | 10.4 tok/s | 9.6 tok/s | +8% |
| 423 tokens | 10.0 tok/s | 7.4 tok/s | +34% |
| 786 tokens | 10.0 tok/s | 7.3 tok/s | +37% |
Every attention layer in the baseline allocates a temporary float32 matrix for the dequantized keys. Across 50 sliding-attention layers at 950 tokens, that's 780MB of temporaries. The fused kernel allocates nothing.
Gemma 4 31B at 4-bit weights needs 17.4 GB. Whether 256K context fits depends on KV cache format:
| Approach | Result | Why |
|---|---|---|
| Fused sdpa_int4 kernel | +37% decode | Dequantize in registers, zero temporaries |
| PolarQuant (Metal shader) | Broken output | attn_scale=1.0 amplifies angular error |
| QJL residual correction | +6.5% perplexity | 1-bit variance hurts softmax |
| int4 KV > 950 tokens | Fails | Compound error across 60 layers |
| int8 KV > 1000 tokens | Fails | Same fundamental limit (attn_scale=1.0) |
| Speculative decode (E2B→31B) | 12 tok/s (slower) | 25% draft acceptance rate |
| FP16 intermediates | Slower | quantized_matmul has cast overhead |
| async_eval pipelining | Neutral | Mutable KV cache prevents overlap |
| Chunked layer eval | Slower | Sync overhead exceeds graph overhead |
- macOS with Apple Silicon (M1/M2/M3/M4)
- Xcode Command Line Tools (
xcode-select --install) - CMake 3.27+ (
brew install cmake) - Python 3.12 (
brew install python@3.12) — only needed for model download and test runner - ~20 GB free disk space (for model weights)
git clone https://github.com/your-org/turboquant.git
cd turboquantTurboQuant links against MLX's C++ library. Build it once:
git clone --depth 1 https://github.com/ml-explore/mlx.git /tmp/mlx-source
mkdir /tmp/mlx-source/build && cd /tmp/mlx-source/build
cmake .. -DCMAKE_BUILD_TYPE=Release -DMLX_BUILD_TESTS=OFF -DMLX_BUILD_PYTHON_BINDINGS=OFF
make -j8
cd -mkdir -p build && cd build
cmake .. -DCMAKE_BUILD_TYPE=Release
make -j8
cd ..This produces two binaries in build/:
gemma4— the inference enginetest_sdpa_int4— kernel correctness test
The model uses mlx-community/gemma-4-31b-it-4bit — a 4-bit quantized version of Google's Gemma 4 31B-IT (31 billion parameters, 17.4 GB on disk).
pip3.12 install huggingface_hub
python3.12 -c "
from huggingface_hub import snapshot_download
snapshot_download('mlx-community/gemma-4-31b-it-4bit', local_dir_links=True)
"The engine expects weights at ~/.cache/huggingface/hub/models--mlx-community--gemma-4-31b-it-4bit/. If your path differs, update the model_dir string in engine/gemma4_multilayer.cpp.
# Verify the build
./build/test_sdpa_int4 # should print PASS
bash engine/run_tests.sh # 5/5 should passpython3.12 engine/chat_repl.pyThis starts a streaming multi-turn chat. The model remembers the full conversation. Tokens appear in real-time as they're generated (~10 tok/s).
Gemma 4 31B — TurboQuant Fused int4 SDPA
Type 'quit' to exit, 'clear' to reset memory
You: What is the capital of France?
Gemma: The capital of France is Paris.
[17 tokens in 1.6s (10.4 tok/s)]
You: And Germany?
Gemma: The capital of Germany is Berlin.
[15 tokens in 1.4s (10.2 tok/s)]
You: clear
[Memory cleared]
turboquant/
├── lib/ # Fused int4 SDPA library
│ ├── turboquant.h # C++ API
│ ├── turboquant.cpp # Metal kernel dispatch (MLX Primitive)
│ └── turboquant.metal # Metal compute shaders (sdpa_int4 + PolarQuant)
├── engine/ # Gemma 4 31B inference engine
│ ├── gemma4_multilayer.cpp # 60-layer forward pass
│ ├── chat.py / chat_repl.py # Python wrappers
│ └── run_tests.sh # Regression tests (5/5)
├── tests/
│ └── test_sdpa_int4.cpp # Kernel correctness (max error < 0.00001)
└── assets/ # Benchmark charts
- TurboQuant (ICLR 2026) — the paper I attempted to implement
- QJL (AAAI) — 1-bit quantized JL transform
- PolarQuant (AISTATS 2026) — recursive polar quantization



