"""BitNet 1.58 — ternary quantization primitives (weights {-1, 0, +1}). Pure primitive module: tensor-level quantization functions only. No nn.Module wrapper classes — the unified QuantizedModule (base.py) is the single place for quantized layer wrappers. BitNet b1.58: each weight is quantized to one of three values: {-1, 0, +1}. A single per-tensor (or per-channel) scale factor reconstructs the weight: W ~= ternary_weight * scale Quantization (data-free, from weights only): scale = mean(abs(W)) # per-tensor (BitNet original) scale = mean(abs(W), dim=1) # per-channel (better for uneven layers) ternary = round(W / scale) # clamped to {-1, 0, +1} W_approx = ternary * scale STE (Straight-Through Estimator): Forward: ternary = round(W_latent / scale) # non-differentiable Backward: grad flows to W_latent as identity # STE bypass Inference: dequant = ternary.to(float) * scale, then normal matmul/conv. Storage: ternary weights packed 2 values per int8 byte (2 bits each + padding), scale is float32 per-tensor or per-channel. """ import torch from agiws_neural_quant.training.ste import STEQuantize # --------------------------------------------------------------------------- # Ternary quantization core (data-free, from weight tensor) # --------------------------------------------------------------------------- def ternarize_tensor( w: torch.Tensor, scale_mode: str = "per-channel", ) -> tuple[torch.Tensor, torch.Tensor]: """Quantize a weight tensor to ternary {-1, 0, +1} + scale. Args: w: float weight tensor. For Linear: [out, in]. For Conv3d: [out, in, kT, kH, kW]. scale_mode: 'per-tensor' (single scale) or 'per-channel' (one scale per output channel). Returns: (ternary, scale) where: ternary: int8 tensor with values {-1, 0, +1}, same shape as w scale: float32 tensor — scalar (per-tensor) or [out] (per-channel) """ w = w.detach().float() if scale_mode == "per-tensor": scale = w.abs().mean().clamp(min=1e-8) ternary = torch.clamp(torch.round(w / scale), min=-1, max=1).to(torch.int8) return ternary, scale.reshape(1) elif scale_mode == "per-channel": # Per output channel (dim 0 for both Linear [out,in] and Conv [out,in,kT,kH,kW]) reduce_dims = tuple(d for d in range(1, w.dim())) scale = w.abs().mean(dim=reduce_dims).clamp(min=1e-8) # [out] # Reshape scale for broadcasting reshape = [1] * w.dim() reshape[0] = w.shape[0] ternary = torch.clamp(torch.round(w / scale.reshape(reshape)), min=-1, max=1).to(torch.int8) return ternary, scale else: raise ValueError(f"scale_mode must be 'per-tensor' or 'per-channel', got {scale_mode!r}") def ternary_dequantize(ternary: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: """Reconstruct float weight from ternary + scale. Args: ternary: int8 tensor with values {-1, 0, +1} scale: scalar tensor (per-tensor) or [out] tensor (per-channel) Returns: float32 dequantized weight, same shape as ternary """ t = ternary.to(torch.float32) if scale.dim() == 0 or scale.numel() == 1: return t * scale.to(torch.float32) else: # Per-channel: reshape scale for broadcasting over [out, ...] reshape = [1] * ternary.dim() reshape[0] = ternary.shape[0] return t * scale.to(torch.float32).reshape(reshape) def fake_ternarize(w: torch.Tensor, scale_mode: str = "per-channel") -> torch.Tensor: """STE fake ternarization for training (QAT / block-wise distillation). Forward: ternarize {-1,0,+1} and dequantize back (simulates quantization error). Backward: gradient flows to w as identity (STE). Args: w: latent float weight (nn.Parameter, requires_grad=True) scale_mode: 'per-tensor' or 'per-channel' Returns: Fake-quantized weight (float, same shape as w), gradient-connected to w via STE. """ w = w.float() if scale_mode == "per-tensor": scale = w.abs().mean().clamp(min=1e-8).detach() # STE: symmetric=True for ternary {-1,0,+1} (clamp to [-1, +1], not [-1, 0]) return STEQuantize.apply(w, scale.unsqueeze(0), 1, True) elif scale_mode == "per-channel": reduce_dims = tuple(d for d in range(1, w.dim())) scale = w.abs().mean(dim=reduce_dims).clamp(min=1e-8).detach() reshape = [1] * w.dim() reshape[0] = w.shape[0] # STE: symmetric=True for ternary {-1,0,+1} return STEQuantize.apply(w, scale.reshape(reshape), 1, True) else: raise ValueError(f"scale_mode must be 'per-tensor' or 'per-channel', got {scale_mode!r}") __all__ = [ "ternarize_tensor", "ternary_dequantize", "fake_ternarize", ]