File size: 5,926 Bytes
bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 ed2bbfb bbbd564 e4a7eb7 bbbd564 e4a7eb7 ed2bbfb bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 e4a7eb7 bbbd564 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | ---
library_name: kernels
license: apache-2.0
---
# quant-matmul
W4A8, W8A8, and FP16-storage linear layers for quantized LLM inference on
CPUs, aarch64 first, loadable through `kernels`. Int4 weights in
GPTQ/AWQ-compatible group quantization run against per-token INT8
activations. The reference baseline is a dequantized fp32 product, matched
within 1e-2. Companions:
[bitnet-cpu](https://huggingface.co/kernels/phanerozoic/bitnet-cpu)
(ternary),
[cpu-attn](https://huggingface.co/kernels/phanerozoic/cpu-attn) (int8 KV
attention),
[decode-ops](https://huggingface.co/kernels/phanerozoic/decode-ops) (fused
glue).
A 4-bit checkpoint is small enough for a single-board computer, but the
Python stack has no CPU kernel that multiplies the packed nibbles directly:
loaders dequantize to float first, which discards the size advantage and the
speed with it. This kernel keeps the weights packed, decodes nibbles in
registers against an even/odd split of the activations, and picks the widest
integer dot-product instruction the CPU actually has, so a 4096-square
decode layer runs in about a millisecond on a Raspberry Pi where fp32 takes
forty.

*Measured on the boards: 1.10 ms for a 4096-square W4A8 decode layer on a
Pi 5's Cortex-A76 against 40.6 ms in fp32 (37x) with weights down from 64 MB
to 8 MB, 3.44 ms for an 11008-wide FFN layer (31x), and the same kernel at
8.76 and 22.68 ms on a Pi 4's A72, within 1e-2 of a dequantized fp32
reference throughout.*
## Usage
```python
import torch
from kernels import get_kernel
qm = get_kernel("phanerozoic/quant-matmul", version=1, trust_remote_code=True)
# From a GPTQ checkpoint's tensors:
packed, s, z, gs, perm = qm.from_gptq(qweight, qzeros, scales, g_idx)
y = qm.w4a8_linear(x, packed, s, z, gs, perm) # [..., N]
# Or as a module:
layer = qm.QuantLinear.from_gptq_tensors(qweight, qzeros, scales, g_idx)
y = layer(x)
```
`version` selects the release branch; `trust_remote_code` is required by
`kernels` for publishers without the trusted-publisher mark.
## API
| Symbol | Purpose |
|---|---|
| `pack_w4(intw, scales, zeros, group_size)` | raw nibbles `[N, K]` -> native packing |
| `from_gptq(qweight, qzeros, scales, g_idx=None, zero_offset=1)` | GPTQ tensors -> native; act-order returns an activation `perm` |
| `from_awq(qweight, qzeros, scales)` | AutoAWQ GEMM tensors -> native |
| `dequant_w4(packed, scales, zeros, group_size)` | reference f32 dequantization |
| `quantize_activation(x)` | per-token absmax INT8 |
| `w4a8_linear(x, packed, scales, zeros, group_size, perm)` | int4 weights x int8 activations |
| `w8a8_linear(x, w_int8, scale_wt)` | int8 x int8, per-row weight scale |
| `fp16_gemv(x, w_f16)` | f16-storage weights widened to f32 in registers |
| `QuantLinear` | `nn.Module` wrapper |
`zero_offset=1` follows the AutoGPTQ storage convention; pass 0 for
checkpoints without it.
## Method
Per group `g` of `group_size` weights along K:
```
y[m, n] = sa[m] * sum_g s[n, g] * (dot_g(a_m, w_n) - z[n, g] * asum_g[m])
```
Group dots are exact in int32 (nibble <= 15, |a| <= 127); groups accumulate
in fp32 in fixed order, so every dispatch tier produces bit-identical f32
output. Nibbles are shifted and masked in registers against a one-time
even/odd split of the activations; no unpack buffer exists.
| path | instruction | selected when |
|---|---|---|
| AVX-VNNI | `vpdpbusd` (256-bit) | Intel 12th-gen+, AMD Zen 4+ |
| AVX2 | `vpmaddubsw` + `vpmaddwd` | Haswell and newer |
| 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; `QM_CPU_ISA` demotes the
selection for A/B runs and never promotes.
## Measured
Raspberry Pi 5 (Cortex-A76, SDOT) and Pi 4 Model B (Cortex-A72, NEON),
torch 2.13 CPU, group_size 128, against torch fp32 on the same board:
| op, shape (M, N, K) | torch fp32 | Pi 5 | Pi 4 | weights |
|---|---|---|---|---|
| w4a8 1, 4096, 4096 | 40.6 ms | 1.10 ms | 8.76 ms | 64 -> 8 MB |
| w4a8 1, 11008, 4096 | 107.4 ms | 3.44 ms | 22.68 ms | 172 -> 22 MB |
| w8a8 1, 4096, 4096 | 39.6 ms | 1.48 ms | 11.45 ms | 64 -> 16 MB |
Further measurements at group_size 128 on the same boards: w4a8
`1, 4096, 11008` runs 3.57 ms (Pi 5) and 10.4 ms (Pi 4); `128, 4096, 4096`
runs 23.2 and 209 ms; fp16-storage `1, 4096, 4096` runs 2.80 and 13.5 ms.
Decode streams packed weights at 8-12 GB/s on the Pi 5, the board's memory
ceiling; M=128 runs 185 GOPS on the SDOT tier. On x86-64 (12th-gen Intel
mobile, AVX-VNNI), W4A8 M=1 4096-square runs 0.16 ms (1.6x its scalar path)
and the fp16-storage GEMV 0.25 ms (7x).
## Correctness
Output is within 1e-2 relative of a dequantized fp32 reference across the
measured shapes on both boards, which is the int8 activation quantization
step. Group dots are exact in int32 and groups accumulate in a fixed order,
so every dispatch tier gives bit-identical f32 output; this is verified by
demoting through `QM_CPU_ISA`.
## Requirements and limits
- `K` divisible by `group_size`; `group_size` a multiple of 32 (64 and 128
tested).
- bf16 or f32 activations; output follows the input dtype.
- GPTQ act-order (`g_idx`) is supported through an activation permutation
applied per forward.
- Fast paths cover aarch64 (NEON; SDOT with dotprod) and x86-64 (AVX2;
AVX-VNNI where present); anything else uses the correct scalar fallback.
- llama.cpp remains faster as a standalone runtime; this is for models
composed in Python.
## References
Frantar et al., "GPTQ" (2023); Lin et al., "AWQ" (2023); Arm `sdot`; Intel
AVX-VNNI `vpdpbusd`.
## License
Apache-2.0.
|