psikosen's picture
Update to v5: Prefix Sliding KV-cache, SMELT scaling, CMA checkpoint, sPTC tool caller
2976bd9 verified
Raw History Blame Contribute Delete
4.17 kB
"""
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