File size: 6,602 Bytes
ce8679e
 
 
 
 
 
 
587d142
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ce8679e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0d07041
587d142
 
 
ce8679e
 
 
587d142
ce8679e
 
 
 
 
 
 
587d142
ce8679e
587d142
 
 
ce8679e
587d142
b02a470
587d142
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b02a470
 
 
 
 
 
 
 
 
 
587d142
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ce8679e
 
 
 
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
147
148
149
150
151
152
153
154
155
156
157
158
159
---
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.