Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

24 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Fused int4 Attention for Gemma 4 31B on Apple Silicon

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.

Background: Why TurboQuant Fails on Gemma 4

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.

Paper vs Implementation

What I Built Instead

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.

Performance

Decode speed stays constant as context grows

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.

Throughput vs Context

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%

Peak memory: 780MB less at 950 tokens

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.

Memory Savings

Hardware

Gemma 4 31B at 4-bit weights needs 17.4 GB. Whether 256K context fits depends on KV cache format:

Hardware Requirements

What I Tried (177 experiments)

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

Getting Started

Prerequisites

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

Step 1: Clone and set up

git clone https://github.com/your-org/turboquant.git
cd turboquant

Step 2: Build MLX from source

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

Step 3: Build TurboQuant

mkdir -p build && cd build
cmake .. -DCMAKE_BUILD_TYPE=Release
make -j8
cd ..

This produces two binaries in build/:

  • gemma4 — the inference engine
  • test_sdpa_int4 — kernel correctness test

Step 4: Download Gemma 4 31B weights

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.

Step 5: Run

# Verify the build
./build/test_sdpa_int4          # should print PASS
bash engine/run_tests.sh        # 5/5 should pass

Step 6: Chat

python3.12 engine/chat_repl.py

This 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]

File Structure

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

References

  • TurboQuant (ICLR 2026) — the paper I attempted to implement
  • QJL (AAAI) — 1-bit quantized JL transform
  • PolarQuant (AISTATS 2026) — recursive polar quantization

About

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.

Topics

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages