""" BitNet b1.58 Ternary (W1.58A8) and Binary (W1A8) Quantization with Straight-Through Estimator (STE). """ import torch import torch.nn as nn import torch.nn.functional as F from typing import Optional class WeightQuantSTE(torch.autograd.Function): """Straight-Through Estimator for Ternary and Binary Weight Quantization.""" @staticmethod def forward(ctx, weight: torch.Tensor, bits: float = 1.58, eps: float = 1e-7) -> torch.Tensor: scale = weight.abs().mean().clamp(min=eps) if bits == 1.58: # Ternary: {-1, 0, +1} * scale scaled_w = weight / scale quant_w = torch.clamp(torch.round(scaled_w), -1.0, 1.0) * scale elif bits == 1.0: # Binary: {-1, +1} * scale sign_w = torch.where(weight >= 0, torch.ones_like(weight), -torch.ones_like(weight)) quant_w = sign_w * scale else: quant_w = weight return quant_w @staticmethod def backward(ctx, grad_output: torch.Tensor): # Straight-through: pass gradients directly to FP32/BF16 master weights return grad_output, None, None class ActivationQuantSTE(torch.autograd.Function): """Per-token 8-bit activation quantization with STE.""" @staticmethod def forward(ctx, x: torch.Tensor, eps: float = 1e-7) -> torch.Tensor: # Per-token maximum along last dimension gamma = 127.0 / (x.abs().amax(dim=-1, keepdim=True).clamp(min=eps)) scaled_x = x * gamma quant_x = torch.clamp(torch.round(scaled_x), -128.0, 127.0) / gamma return quant_x @staticmethod def backward(ctx, grad_output: torch.Tensor): # Straight-through return grad_output, None def quantize_weight(weight: torch.Tensor, bits: float = 1.58) -> torch.Tensor: return WeightQuantSTE.apply(weight, bits) def quantize_activation(x: torch.Tensor) -> torch.Tensor: return ActivationQuantSTE.apply(x) class BitLinear(nn.Linear): """ BitLinear layer implementing BitNet b1.58 ternary (W1.58A8) or binary (W1A8) weights with per-token 8-bit activations and Straight-Through Estimator. """ def __init__( self, in_features: int, out_features: int, bias: bool = False, bits: float = 1.58, device=None, dtype=None, ): super().__init__(in_features, out_features, bias=bias, device=device, dtype=dtype) self.bits = bits self.active = True def forward(self, x: torch.Tensor) -> torch.Tensor: if not self.active or self.bits > 2.0: return F.linear(x, self.weight, self.bias) # 1. Per-token 8-bit activation quantization x_quant = quantize_activation(x) # 2. Weight quantization (ternary or binary) w_quant = quantize_weight(self.weight, self.bits) # 3. Linear projection with quantized tensors return F.linear(x_quant, w_quant, self.bias) def convert_to_bitlinear( module: nn.Module, bits: float = 1.58, preserve_modules: tuple = ("router", "thought_bus", "visit_adapter", "lm_head", "norm"), ) -> nn.Module: """ Recursively replaces standard nn.Linear modules with BitLinear, preserving sensitive components (router, norms, Thought Bus, and visit adapters). """ for name, child in module.named_children(): if any(p in name for p in preserve_modules): continue if isinstance(child, nn.Linear) and not isinstance(child, BitLinear): bit_linear = BitLinear( in_features=child.in_features, out_features=child.out_features, bias=child.bias is not None, bits=bits, device=child.weight.device, dtype=child.weight.dtype, ) with torch.no_grad(): bit_linear.weight.copy_(child.weight) if child.bias is not None: bit_linear.bias.copy_(child.bias) setattr(module, name, bit_linear) else: convert_to_bitlinear(child, bits=bits, preserve_modules=preserve_modules) return module