File size: 5,696 Bytes
31ec85e
 
 
 
 
 
 
 
ca7b082
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31ec85e
 
 
 
 
 
 
 
 
 
 
 
 
 
ca7b082
 
31ec85e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ca7b082
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31ec85e
 
 
 
ca7b082
 
 
 
 
31ec85e
 
 
ca7b082
 
 
31ec85e
 
 
 
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
---
library_name: kernels
license: apache-2.0
---

# binary-gemm

W1A8 matrix products for binary-weight language models, loadable through
`kernels`. Each weight is one bit meaning `+1` or `-1`, with one fp16 scale
per group of `G` weights: at `G = 128` that is 1.125 effective bits per
weight, against 2 bits for a ternary layer and 16 for bf16. The reference
baseline is the dequantized product; the integer path is `torch.equal`-exact
where bf16 can represent the accumulator. Ternary siblings:
[bitnet-tc](https://huggingface.co/kernels/phanerozoic/bitnet-tc) (CUDA) and
[bitnet-cpu](https://huggingface.co/kernels/phanerozoic/bitnet-cpu) (CPU).

At one bit per weight a model's memory is its scales plus raw sign bits, and
decode, which reads every weight per token, becomes almost pure bit traffic.
This kernel multiplies the packed bits directly: a 16-entry constant table
expands a nibble of the pattern into four packed `±1` int8 values feeding
`__dp4a`, so the weight is never materialized as a number and the decode wall
drops with the bytes.

![A real sign-bit field under a scanline next to per-layer decode bars where the binary path runs a fraction of bf16](https://huggingface.co/kernels/phanerozoic/binary-gemm/resolve/main/media/hero.gif)

*Actual packed sign bits of a layer (cyan +1, magenta -1), and decode
measured with weights rotated past L2 in both paths: 2.4x on attention QKV,
2.8x on the MLP gate, 3.8x on the 248,320-row LM head.*

## Usage

```python
import torch
from kernels import get_kernel

bg = get_kernel("phanerozoic/binary-gemm", version=1, trust_remote_code=True)

W = torch.where(torch.randn(N, K, device="cuda") >= 0, 1, -1).to(torch.int8)
wq = bg.pack_weights(W)                       # [N, K//32] int32
ws = group_scales.to(torch.float16)           # [N, K//128] fp16

y = bg.binary_linear(x, wq, ws, group_size=128)   # bf16 in, bf16 out

layer = bg.BinaryLinear.from_dense(nn.Linear(K, N))   # sign + group mean-abs scale
```

`version` selects the release branch; `trust_remote_code` is required by
`kernels` for publishers without the trusted-publisher mark.

## API

| Symbol | Purpose |
|---|---|
| `pack_weights(W)` | `{-1,+1}` int8 `[N,K]` -> bit-packed int32 `[N,K//32]` |
| `unpack_weights(wq, K)` | inverse, for tests and reference paths |
| `quantize_activation(x)` | bf16 `[...,K]` -> (int8 `[M,K]`, fp32 per-token scale) |
| `binary_gemm(act_q, act_scale, wq, wscale, bias, group_size)` | packed product -> bf16 |
| `binary_linear(x, wq, wscale, bias, group_size)` | one-shot forward |
| `BinaryLinear(in, out, bias, group_size)` | `nn.Module`; `from_dense` converts a dense layer |

## Method

Binary weights collapse the inner product: writing the stored bit as
`b = (w + 1) / 2`,

```
dot(w, a) = 2 * sum_{i : b_i = 1} a_i  -  sum_i a_i
```

so the weight never has to exist as a number. The kernel keeps the int8
activation path and resolves four weights at a time through a 16-entry
constant table that expands a nibble directly into four packed int8 `±1`
values for `__dp4a`; the table is the whole decode. One warp owns each output
column: lanes stride the weight groups, accumulate an int32 dot per group,
apply that group's fp16 scale in fp32, and the warp reduces. The batch
dimension is tiled at compile time so a weight word is fetched once per tile.

## Measured

Bonsai 27B shapes, hidden 5120, `G = 128`, against cuBLAS bf16, weights
rotated over enough copies to exceed twice the card's L2 on both paths (an
84 MB bf16 matrix inside a 96 MB L2 otherwise measures cache bandwidth):

| layer | N | cuBLAS bf16 | binary | speedup |
|---|---|---|---|---|
| attention QKV | 8,192 | 0.1124 ms | 0.0404 ms | 2.78x |
| attention out | 5,120 | 0.0733 ms | 0.0276 ms | 2.65x |
| MLP gate | 17,408 | 0.2425 ms | 0.0756 ms | 3.21x |
| MLP down | 5,120 | 0.0738 ms | 0.0272 ms | 2.71x |
| LM head | 248,320 | 3.3622 ms | 0.9637 ms | 3.49x |

Weight footprint at hidden 5120:

| layer | bf16 | ternary (2-bit) | binary (1.125-bit) | vs ternary |
|---|---|---|---|---|
| attention QKV | 80.0 MB | 10.6 MB | 5.6 MB | 1.89x |
| attention out | 50.0 MB | 6.6 MB | 3.5 MB | 1.89x |
| LM head | 2,425.0 MB | 322.1 MB | 170.5 MB | 1.89x |

1.89x rather than the nominal 1.78x because the fp16 group scales are a
fixed overhead both formats pay.

## Correctness

- Packing is lossless: `unpack(pack(W))` is `torch.equal` to `W` for every
  shape tested. Zeros are treated as `-1`, so a ternary tensor deliberately
  does not round-trip.
- The integer path is exact: with unit group scales the output is
  `torch.equal` to an exact int32 reference wherever bf16 represents the
  accumulator (`|acc| < 256`).
- The full path matches a dequantize-and-fp32 reference to below `2^-8`
  normwise, the bf16 output rounding step.
- Activation quantization is per-token absmax with reconstruction error
  under 1% of the row maximum; group size verified at 32, 64, 128, 256.
- Deterministic: repeated products are bitwise identical; fixed reduction
  tree, no atomics.

## Requirements and limits

- NVIDIA GPU with compute capability 8.0+ (`__dp4a`).
- `K` a multiple of 32 and of `group_size`; bf16 activations and output.
- Decode only: a warp-per-column GEMV built for `M = 1`. Above that it loses
  to cuBLAS (0.40x at `M = 2`, 0.02x at `M = 256`); route prefill and batch
  elsewhere.
- Weights must already be binary; `from_dense` is a sign quantizer for
  testing, not a compression method.

## References

Rastegari et al., "XNOR-Net" (2016); Courbariaux et al., "BinaryConnect"
(2015); Ma et al., "The Era of 1-bit LLMs" (2024); group-wise scaling with
1-bit weights as shipped in the Bonsai family (PrismML, 2026).

## License

Apache-2.0.