bitnet-cpu / README.md
phanerozoic's picture
Card: standardized form with hero
587d142 verified
|
Raw
History Blame
6.6 kB
---
library_name: kernels
license: apache-2.0
---
# bitnet-cpu
Ternary x INT8 GEMM for BitNet b1.58 (W1.58 A8) on CPUs, loadable through
`kernels`. The CUDA member of this stack is
[bitnet-tc](https://huggingface.co/kernels/phanerozoic/bitnet-tc); both
share the same 2-bit packing (compatible with Microsoft's bitnet.cpp) and
the same Python API, so code written against one runs against the other.
The reference baseline is an fp32 reference product, matched within bf16
output rounding.
A ternary-weight model is small enough to fit a single-board computer, and
useless there if the matmul runs at framework speed: a 2B model's LM head
alone takes 0.7 seconds per token in fp32 on a Raspberry Pi. This kernel
multiplies the packed 2-bit weights directly against per-token INT8
activations, picking the widest instruction the CPU actually has at load
time, so the same binary runs a 27B-geometry layer in 0.4 ms on a Pi 5's
Cortex-A76 and still works correctly on a Pi 4, an old x86, or anything
else.
![Per-layer bars where the fp32 path dwarfs the kernel, and a three-rung dispatch ladder with Pi 5 and Pi 4 side by side](https://huggingface.co/kernels/phanerozoic/bitnet-cpu/resolve/main/media/hero.gif)
*Measured on the boards: decode layers at 72x to 98x the fp32 path with
weights down from 1,252 MB to 79 MB on the LM head, and the runtime
dispatch ladder from the same binary, 2.06 ms scalar, 0.58 NEON, 0.42 SDOT
on a Pi 5's Cortex-A76, against 14.3 / 3.62 / 3.95 on a Pi 4's A72, which
has no dotprod and correctly stays on base NEON.*
## Usage
```python
import torch
from kernels import get_kernel
bitnet = get_kernel("phanerozoic/bitnet-cpu", version=1, trust_remote_code=True)
N, K = 4096, 4096
W = torch.randint(-1, 2, (N, K), dtype=torch.int8)
w_packed = bitnet.pack_weights(W) # [N, K//4] uint8
scale_wt = torch.ones(N, dtype=torch.bfloat16)
x = torch.randn(1, K, dtype=torch.bfloat16)
y = bitnet.bitnet_linear(x, w_packed, scale_wt) # [1, N] bf16
```
`version` selects the release branch; `trust_remote_code` is required by
`kernels` for publishers without the trusted-publisher mark. `BitLinear` is
the drop-in `nn.Module`; `BitLinearKernel` is the `kernelize` layer for the
transformers BitNet modules.
## API
| Symbol | Purpose |
|---|---|
| `pack_weights(W)` | ternary `{-1,0,+1}` int8 `[N,K]` -> packed uint8 `[N,K//4]` |
| `quantize_activation(x)` | bf16/f32 `[..,K]` -> (`int8`, per-row bf16 scale) |
| `bitnet_gemm(x_int8, w_packed, scale_act, scale_wt)` | INT8 activations x packed ternary weights -> bf16 |
| `bitnet_gemv_fused(x, w_packed, scale_wt)` | fused quantize + GEMV, M < 16 |
| `bitnet_linear(x, w_packed, scale_wt)` | one-shot forward, auto-dispatch |
| `BitLinear(in, out)` | `nn.Module` wrapper |
| `BitLinearKernel` | `kernelize` layer for transformers BitNet modules |
Weight encoding: ternary `{-1, 0, +1}` -> 2-bit codes `{1, 2, 3}` packed
four per byte (decode is `byte - 2`), the packing used by bitnet-tc and
Microsoft's bitnet.cpp.
## Method
On x86-64 the inner product uses `dot(w, a) = dot(w + 1, a) - sum(a)`,
mapping onto unsigned x signed multiply-accumulate; on aarch64 the packed
codes multiply directly as signed int8 via
`dot(w, a) = dot(w + 2, a) - 2 * sum(a)`, and the decode path (M <= 16)
reads weights straight from the packed bytes with no unpack buffer.
| path | instruction | selected when |
|---|---|---|
| AVX-512 VNNI | `vpdpbusd` (512-bit) | AMD Zen 4+, Intel server cores |
| AVX-VNNI | `vpdpbusd` (256-bit) | Intel 12th-gen+ client cores |
| 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 |
The path is chosen once at runtime from cpuid / HWCAP; no flags needed.
`BITNET_CPU_ISA` (`scalar`, `neon`, `avx2`, ...) demotes the selection for
A/B runs and never promotes.
## Measured
Raspberry Pi 5 (4x Cortex-A76 2.4 GHz, SDOT path), 27B-class layer
geometry, against torch fp32 on the same board:
| layer (M, N, K) | torch fp32 | bitnet-cpu | speedup | weights |
|---|---|---|---|---|
| attn QKV 1, 6912, 2560 | 35.7 ms | 0.40 ms | 89x | 68 -> 4 MB |
| MLP down 1, 2560, 6912 | 32.4 ms | 0.45 ms | 72x | 68 -> 4 MB |
| LM head 1, 128256, 2560 | 684.3 ms | 7.00 ms | 98x | 1,252 -> 79 MB |
| prefill 64, 6912, 2560 | 46.1 ms | 8.43 ms | 5.5x | 68 -> 4 MB |
Dispatch ladder, same binary, 1 x 6912 x 2560:
| tier | Pi 5 (A76) | Pi 4 (A72) |
|---|---|---|
| scalar | 2.06 ms | 14.31 ms |
| NEON | 0.58 ms | 3.62 ms |
| SDOT / auto | 0.42 ms | 3.95 ms |
The A76 takes SDOT; the A72 has no dotprod, so auto correctly selects base
NEON and the SDOT row is unavailable to it.
On a 16 vCPU x86-64 host (AVX2, no VNNI), torch 2.12 CPU, median of 15:
| shape (M, N, K) | torch bf16 matmul | bitnet-cpu | speedup |
|---|---|---|---|
| 1, 6912, 2560 | 0.24 ms | 0.37 ms | 0.7x |
| 1, 11008, 4096 | 0.52 ms | 0.42 ms | 1.2x |
| 1, 128256, 2560 | 9.86 ms | 2.01 ms | 4.9x |
| 128, 6912, 2560 | 22.83 ms | 1.66 ms | 13.8x |
| 512, 11008, 4096 | 351.6 ms | 13.25 ms | 26.5x |
End-to-end `microsoft/bitnet-b1.58-2B-4T-bf16` through transformers with all
210 BitLinear layers routed to the kernel: decode 0.78 -> 3.86 tok/s (5.0x)
on the x86 host, 0.035 -> 0.94 tok/s (27x) on the Pi 5, with identical
greedy output over the measured window.
## Correctness
The integer path is exact; output differs from an fp32 reference only by
bf16 output rounding (max relative error ~1e-2 across dispatch paths on the
Pi measurements above, ~4e-3 on the x86 suite). Every tier produces the same
result within that bound, verified by demoting through `BITNET_CPU_ISA` on
both a Cortex-A76 and a Cortex-A72 board.
## Requirements and limits
- `K` divisible by 32; bf16 or f32 activations; bf16 output; uint8 packed
weights.
- Fast paths cover x86-64 (AVX2 or newer) and aarch64 (NEON; SDOT where the
CPU has dotprod). Anything else uses the scalar fallback, which is correct
and roughly 5x slower than NEON.
- Prefill (large M) narrows to about 5x; the win is concentrated in decode.
- Microsoft's bitnet.cpp remains faster on CPU via lookup-table kernels and
is the right choice when a dedicated llama.cpp-style runtime is
acceptable; this kernel is for `get_kernel` and the transformers
ecosystem.
## References
Ma et al., "The Era of 1-bit LLMs" (BitNet b1.58, 2024); Microsoft
bitnet.cpp (the shared packing); Arm `sdot` and NEON integer dot products.
## License
Apache-2.0.