Drop-in replacement for nn.Linear implementing BitNet b1.58 training. Forward pass (training): 1. Normalise input with internal parameter-free RMSNorm. 2. Quantise normalised input activations to 8-bit per token (STE). 3. Quantise weight matrix to ternary {-1, 0, +1} via
| 86 | # ────────────────────────────────────────────────────────────────────────────── |
| 87 | |
| 88 | class BitLinear(nn.Linear): |
| 89 | """ |
| 90 | Drop-in replacement for nn.Linear implementing BitNet b1.58 training. |
| 91 | |
| 92 | Forward pass (training): |
| 93 | 1. Normalise input with internal parameter-free RMSNorm. |
| 94 | 2. Quantise normalised input activations to 8-bit per token (STE). |
| 95 | 3. Quantise weight matrix to ternary {-1, 0, +1} via absmean (STE). |
| 96 | 4. Compute F.linear(x_q, w_q) using the regular fp matmul kernel |
| 97 | (specialised ternary kernels only needed at inference time). |
| 98 | 5. Rescale output: multiply by (weight_scale × activation_scale). |
| 99 | |
| 100 | Bias is always disabled as per the paper. |
| 101 | """ |
| 102 | |
| 103 | def __init__( |
| 104 | self, |
| 105 | in_features: int, |
| 106 | out_features: int, |
| 107 | bias: bool = False, # paper: no bias |
| 108 | dtype: torch.dtype | None = None, |
| 109 | ) -> None: |
| 110 | super().__init__(in_features, out_features, bias=bias, dtype=dtype) |
| 111 | self.norm = _RMSNorm(in_features) |
| 112 | |
| 113 | # ------------------------------------------------------------------ |
| 114 | def forward(self, x: torch.Tensor) -> torch.Tensor: |
| 115 | # 1. Internal RMSNorm |
| 116 | x_normed = self.norm(x) |
| 117 | |
| 118 | # 2. Quantise activations (STE, keeps gradients) |
| 119 | x_q, act_scale = _quantize_activations_int8(x_normed) |
| 120 | |
| 121 | # 3. Quantise weights (STE, keeps gradients through round) |
| 122 | w_q, w_scale = _quantize_weights_ternary(self.weight) |
| 123 | |
| 124 | # 4. Linear using the (still fp) quantised values |
| 125 | out = F.linear(x_q, w_q, self.bias) |
| 126 | |
| 127 | # 5. Rescale: broadcast act_scale [..., 1] and scalar w_scale |
| 128 | out = out * act_scale * w_scale |
| 129 | |
| 130 | return out |
| 131 | |
| 132 | # ------------------------------------------------------------------ |
| 133 | @classmethod |
| 134 | def from_linear(cls, linear: nn.Linear) -> "BitLinear": |
| 135 | """Construct a BitLinear from an existing nn.Linear, copying its weight.""" |
| 136 | bl = cls( |
| 137 | in_features=linear.in_features, |
| 138 | out_features=linear.out_features, |
| 139 | bias=linear.bias is not None, |
| 140 | dtype=linear.weight.dtype, |
| 141 | ) |
| 142 | with torch.no_grad(): |
| 143 | bl.weight.copy_(linear.weight) |
| 144 | if linear.bias is not None and bl.bias is not None: |
| 145 | bl.bias.copy_(linear.bias) |
nothing calls this directly
no outgoing calls
no test coverage detected