cpu-attn / README.md
phanerozoic's picture
Card: standardized form with hero
02440c4 verified
|
Raw
History Blame
4.75 kB
---
library_name: kernels
license: apache-2.0
---
# cpu-attn
Single-token decode attention over an INT8-quantized KV cache, aarch64
first, loadable through `kernels`. The reference baseline is fp32 attention
on the dequantized cache, matched within 8e-3. Companions:
[bitnet-cpu](https://huggingface.co/kernels/phanerozoic/bitnet-cpu) and
[quant-matmul](https://huggingface.co/kernels/phanerozoic/quant-matmul)
(linears),
[decode-ops](https://huggingface.co/kernels/phanerozoic/decode-ops) (fused
glue and sampling).
Once the linear layers of a CPU inference stack are quantized, the KV cache
is what remains: on an 8 GB single-board computer it is both the memory that
runs out and the bandwidth that sets decode latency. Storing keys and values
at int8 halves both, and this kernel does the decode step directly against
that quantized cache, with exact integer logit dots, so long-context decode
on a Raspberry Pi stops being the first thing that dies.
![Per-context latency bars for Pi 5 and Pi 4 with the cache footprint halved at every length](https://huggingface.co/kernels/phanerozoic/cpu-attn/resolve/main/media/hero.gif)
*Measured on the boards, 8 heads at head dimension 128: 0.10, 0.42, and
1.80 ms at 512, 2,048, and 8,192 tokens on a Pi 5's Cortex-A76, against
1.18, 4.30, and 10.51 ms on a Pi 4's A72, with the cache down from 32 MB to
16 MB at the longest context and output within 8e-3 of fp32 attention on the
dequantized cache.*
## Usage
```python
import torch
from kernels import get_kernel
ca = get_kernel("phanerozoic/cpu-attn", version=1, trust_remote_code=True)
cache = ca.Int8KVCache(n_kv_heads=8, max_seq=4096, head_dim=128)
for k, v in kv_stream: # [Hkv, D] f32 per token
cache.append(k, v)
out = cache.decode(q) # q [H, D] -> [H, D] f32
```
`version` selects the release branch; `trust_remote_code` is required by
`kernels` for publishers without the trusted-publisher mark. GQA follows
from shapes: query head `h` reads kv head `h // (H // Hkv)`.
## API
| Symbol | Purpose |
|---|---|
| `kv_append(k_cache, k_scale, v_cache, v_scale, k_new, v_new, pos)` | quantize one token into the caches |
| `attn_decode(q, k_cache, k_scale, v_cache, v_scale, seq_len, scale)` | decode attention for one token |
| `Int8KVCache(n_kv_heads, max_seq, head_dim)` | allocating wrapper with `append` / `decode` / `dequant_k` / `dequant_v` |
## Method
K and V quantize per token per head (absmax int8) at append time; decode
runs integer q.k logits, an f32 softmax, and f32 accumulation over the int8
V stream:
```
logits[s] = sdot(q_i8, k_i8[s]) * q_scale * k_scale[s] * softmax_scale
out[d] = sum_s softmax(logits)[s] * v_scale[s] * v_i8[s, d]
```
Logit dots are exact int32, accumulation order is fixed, and the NEON
polynomial exp satisfies `exp(0) == 1` exactly, so uniform-logit
constructions are bit-exact end to end.
| path | instruction | selected when |
|---|---|---|
| AVX2 | 16-bit widening `vpmaddwd` + Cephes exp | x86-64 with AVX2 and FMA |
| SDOT | `sdot` (dotprod) | aarch64 with `asimddp`: Cortex-A76+, Neoverse, Apple silicon |
| NEON | `smull` + `sadalp` | any other aarch64 (Pi 4, Pi Zero 2) |
| scalar | portable C++ | everything else |
Chosen once at runtime from CPUID / HWCAP; `CA_CPU_ISA` demotes for A/B
runs. The int8 q.k dot is signed x signed, which AVX-VNNI does not
accelerate, so both x86 tiers share the AVX2 path.
## Measured
Raspberry Pi 5 (4x Cortex-A76 2.4 GHz, SDOT) and Raspberry Pi 4 Model B (4x
Cortex-A72 1.8 GHz, NEON), 64-bit Raspberry Pi OS, torch 2.13 CPU, D = 128:
| H, Hkv, S | Pi 5 | Pi 4 | cache bf16 -> int8 |
|---|---|---|---|
| 8, 8, 512 | 0.10 ms | 1.18 ms | 2 -> 1 MB |
| 8, 8, 2048 | 0.42 ms | 4.30 ms | 8 -> 4 MB |
| 8, 8, 8192 | 1.80 ms | 10.51 ms | 32 -> 16 MB |
| 32, 8, 2048 | 1.12 ms | 3.88 ms | 8 -> 4 MB |
The Pi 5 streams KV at 9-12 GB/s, the board's memory ceiling; an fp16 cache
would move twice the bytes for the same context.
## Correctness
Output matches fp32 attention on the dequantized cache within 8e-3 relative
across the measured contexts on both boards, which separates kernel
correctness from quantization error. Head dim a multiple of 16 up to 1024;
`H % Hkv == 0`.
## Requirements and limits
- f32 queries; head dim a multiple of 16, at most 1024; `H % Hkv == 0`.
- This kernel is the decode path; prefill attention stays in torch SDPA.
- Fast paths cover aarch64 (NEON; SDOT with dotprod) and x86-64 (AVX2 +
FMA); anything else uses the correct scalar fallback.
- Per-token per-head absmax only, no group-wise scaling within a row.
## References
Dao et al., "FlashAttention-2" (2023); Hooper et al., "KVQuant" (2024);
Ainslie et al., "GQA" (2023); Arm `sdot`.
## License
Apache-2.0.