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

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

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.

Downloads last month
-
apache-2.0
OS
linux
Arch
x86_64aarch64
Kernel Builder
19aaa64