File size: 8,502 Bytes
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
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
"""bitnet-cpu: ternary x INT8 GEMM for BitNet b1.58 (W1.58 A8) on CPUs.

Same 2-bit weight packing and API as phanerozoic/bitnet-tc (the CUDA member
of the stack); AVX-512 VNNI / AVX-VNNI / AVX2 paths selected at runtime,
with a portable scalar fallback.

High-level API:
    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(...)             int8 activations x packed ternary weights -> bf16
    bitnet_gemv_fused(...)       bf16/f32 activations (fused quantize) -> bf16, M<16
    bitnet_linear(x, w, sw)      one-shot forward, auto-dispatches by M
    BitLinear                    nn.Module wrapper
    BitLinearKernel              kernelize layer for transformers BitNet modules
"""
from typing import Optional

import torch
import torch.nn as nn

from ._ops import ops

_OK_DTYPES = (torch.bfloat16, torch.float32)


def pack_weights(W: torch.Tensor) -> torch.Tensor:
    """Pack ternary weights {-1, 0, +1} into 2-bit codes, 4 per byte.

    Encoding: -1 -> 1, 0 -> 2, +1 -> 3 (the kernel decodes via byte - 2).
    Identical to the phanerozoic/bitnet-tc and Microsoft bitnet.cpp packing.

    Args:
        W: [N, K] int8 with values in {-1, 0, +1}.
    Returns:
        [N, K // 4] uint8.
    """
    assert W.dtype == torch.int8, f"W must be int8, got {W.dtype}"
    assert W.dim() == 2, "W must be 2D"
    N, K = W.shape
    assert K % 4 == 0, f"K={K} must be a multiple of 4"
    assert ((W >= -1) & (W <= 1)).all(), "W values must be in {-1, 0, +1}"
    codes = (W.to(torch.int32) + 2).to(torch.uint8).view(N, K // 4, 4)
    packed = (
        codes[..., 0]
        | (codes[..., 1] << 2)
        | (codes[..., 2] << 4)
        | (codes[..., 3] << 6)
    ).contiguous()
    return packed


def quantize_activation(x: torch.Tensor, eps: float = 1e-5):
    """Per-token absmax INT8 quantization.

    Args:
        x: [..., K] bf16 or f32 tensor.
    Returns:
        (x_int8 [M, K], scale [M] bf16) where M is the flattened leading dim.
    """
    x_flat = x.reshape(-1, x.shape[-1]).contiguous()
    M, K = x_flat.shape
    if x_flat.dtype in _OK_DTYPES and x_flat.device.type == "cpu":
        x_int8 = torch.empty((M, K), dtype=torch.int8)
        scale = torch.empty((M,), dtype=torch.bfloat16)
        ops.quantize_act(x_int8, scale, x_flat)
        return x_int8, scale
    absmax = x_flat.abs().amax(dim=-1, keepdim=True).clamp(min=eps)
    scale = absmax / 127.0
    x_int8 = (x_flat / scale).round().clamp(-127, 127).to(torch.int8)
    return x_int8, scale.squeeze(-1).to(torch.bfloat16)


def bitnet_gemm(
    x_int8: torch.Tensor,
    w_packed: torch.Tensor,
    scale_act: torch.Tensor,
    scale_wt: torch.Tensor,
) -> torch.Tensor:
    """Ternary x INT8 GEMM. Returns [M, N] bf16."""
    M, K = x_int8.shape
    N = w_packed.shape[0]
    out = torch.empty((M, N), dtype=torch.bfloat16)
    ops.bitnet_gemm(out, x_int8, w_packed, scale_act, scale_wt, None)
    return out


def bitnet_gemv_fused(
    x: torch.Tensor,
    w_packed: torch.Tensor,
    scale_wt: torch.Tensor,
) -> torch.Tensor:
    """Fused quantize + ternary GEMV for M < 16. Returns [M, N] bf16."""
    M, K = x.shape
    N = w_packed.shape[0]
    out = torch.empty((M, N), dtype=torch.bfloat16)
    ops.bitnet_gemv_fused(out, x, w_packed, scale_wt)
    return out


def bitnet_linear(
    x: torch.Tensor,
    w_packed: torch.Tensor,
    scale_wt: torch.Tensor,
) -> torch.Tensor:
    """One-shot BitLinear forward. Auto-dispatches fused (M<16) vs split path.

    Args:
        x: [..., K] bf16 or f32.
        w_packed: [N, K//4] uint8 packed ternary weights.
        scale_wt: [N] bf16 per-output weight scale.
    Returns:
        [..., N] bf16.
    """
    assert x.dtype in _OK_DTYPES, f"x must be bf16 or f32, got {x.dtype}"
    orig = x.shape
    K = orig[-1]
    x_flat = x.reshape(-1, K).contiguous()
    M = x_flat.shape[0]
    N = w_packed.shape[0]
    if M < 16:
        y = bitnet_gemv_fused(x_flat, w_packed, scale_wt)
    else:
        x_int8, scale_act = quantize_activation(x_flat)
        y = bitnet_gemm(x_int8, w_packed, scale_act, scale_wt)
    return y.view(*orig[:-1], N)


class BitLinear(nn.Module):
    """Inference-only BitLinear: ternary weights, per-token INT8 activations."""

    def __init__(self, in_features: int, out_features: int, bias: bool = False):
        super().__init__()
        assert in_features % 32 == 0, f"in_features={in_features} must be % 32"
        self.in_features = in_features
        self.out_features = out_features
        self.register_buffer(
            "w_packed", torch.zeros(out_features, in_features // 4, dtype=torch.uint8)
        )
        self.register_buffer(
            "scale_wt", torch.ones(out_features, dtype=torch.bfloat16)
        )
        if bias:
            self.bias = nn.Parameter(torch.zeros(out_features, dtype=torch.bfloat16))
        else:
            self.bias = None

    @classmethod
    def from_dense(cls, lin: nn.Linear) -> "BitLinear":
        """Quantize a dense Linear via absmean ternarization (sanity-test only;
        a real BitNet model is trained with QAT, not post-hoc quantized)."""
        out_features, in_features = lin.weight.shape
        bl = cls(in_features, out_features, bias=lin.bias is not None)
        with torch.no_grad():
            W = lin.weight.detach().float()
            gamma = W.abs().mean(dim=-1, keepdim=True).clamp(min=1e-8)
            Wq = (W / gamma).round().clamp(-1, 1).to(torch.int8)
            bl.w_packed.copy_(pack_weights(Wq.cpu()))
            bl.scale_wt.copy_(gamma.squeeze(-1).to(torch.bfloat16))
            if bl.bias is not None:
                bl.bias.copy_(lin.bias.detach().to(torch.bfloat16))
        return bl

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        y = bitnet_linear(x, self.w_packed, self.scale_wt)
        if self.bias is not None:
            y = y + self.bias
        return y


class BitLinearKernel(nn.Module):
    """kernelize layer for the transformers BitNet linear modules.

    kernelize binds this forward onto the host BitLinear / AutoBitLinear module,
    so self is that module and the forward reads its attributes (weight,
    weight_scale, rms_norm, bias). The ternary weight is converted to the
    kernel's 2-bit layout once and cached; each call runs per-token INT8
    quantization and the VNNI/AVX2 ternary GEMM.
    """

    can_torch_compile: bool = False

    def forward(self, input: torch.Tensor) -> torch.Tensor:
        dt = input.dtype
        rms = getattr(self, "rms_norm", None)
        if rms is not None:
            input = rms(input)

        if not hasattr(self, "_bitnet_cpu_packed"):
            w = self.weight
            if w.dtype == torch.uint8:
                # transformers stores ternary as int8 viewed through uint8 (255 -> -1).
                tern = w.contiguous().view(torch.int8)
                if tern.shape[0] != self.out_features:        # bit-packed (out//4, in)
                    from transformers.integrations.bitnet import unpack_weights
                    tern = unpack_weights(w, dtype=torch.bfloat16)
                tern = tern.round().clamp(-1, 1).to(torch.int8)
            else:
                g = w.detach().float().abs().mean().clamp(min=1e-5)
                tern = (w.detach().float() / g).round().clamp(-1, 1).to(torch.int8)
            self._bitnet_cpu_n = tern.shape[0]
            self._bitnet_cpu_packed = pack_weights(tern.contiguous().cpu())
            ws = getattr(self, "weight_scale", None)
            if ws is not None:
                s = ws.detach().float().reshape(-1)
                s = s.expand(self._bitnet_cpu_n) if s.numel() == 1 else s
            else:
                s = w.detach().float().abs().mean().clamp(min=1e-5).reshape(1).expand(self._bitnet_cpu_n)
            self._bitnet_cpu_sw = s.to(torch.bfloat16).contiguous().cpu()

        shp = input.shape
        x = input.reshape(-1, shp[-1])
        if x.dtype not in _OK_DTYPES:
            x = x.to(torch.float32)
        x = x.contiguous()
        out = bitnet_linear(x, self._bitnet_cpu_packed, self._bitnet_cpu_sw)
        if getattr(self, "bias", None) is not None:
            out = out + self.bias
        return out.reshape(*shp[:-1], self._bitnet_cpu_n).to(dt)


__all__ = [
    "pack_weights",
    "quantize_activation",
    "bitnet_gemm",
    "bitnet_gemv_fused",
    "bitnet_linear",
    "BitLinear",
    "BitLinearKernel",
]