MCPcopy Create free account
hub / github.com/dhakalnirajan/LLaMA-BitNet / BitLinear

Class BitLinear

utils.py:88–146  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

86# ──────────────────────────────────────────────────────────────────────────────
87
88class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected