""" Adaptive numerical precision as a modular execution feature. We implement *fake quantisation* (quantise -> dequantise) with a straight-through estimator, so the FP32 reference and the reduced-precision path share one code path and can be compared numerically. Precision is assigned per named site (e.g. 'attn', 'mlp'), which makes it a routable resource rather than a fixed compile-time property. """ from __future__ import annotations from dataclasses import dataclass, field from typing import Optional import torch LEVELS = {"int2": 2, "int3": 3, "int4": 4, "int8": 8, "fp16": 16, "bf16": 16, "fp32": 32} def quantise(x: torch.Tensor, bits: int, per: str = "tensor") -> torch.Tensor: """Symmetric uniform fake-quant with straight-through gradient.""" if bits >= 32: return x qmax = 2 ** (bits - 1) - 1 if per == "row": s = x.abs().amax(dim=-1, keepdim=True).clamp_min(1e-8) / qmax else: s = x.abs().amax().clamp_min(1e-8) / qmax q = torch.clamp(torch.round(x / s), -qmax - 1, qmax) * s return x + (q - x).detach() def cast_float(x: torch.Tensor, kind: str) -> torch.Tensor: if kind == "fp16": return x.half().float() if kind == "bf16": return x.bfloat16().float() return x @dataclass class QuantPolicy: """Maps a computation site -> precision level. `default` applies to any site not explicitly listed. `sensitive` sites are kept at higher precision. A policy is data-independent unless `difficulty_fn` is supplied, in which case precision may vary per position. """ default: str = "fp32" sites: dict = field(default_factory=dict) per: str = "row" enabled: bool = True def level_for(self, site: str) -> str: return self.sites.get(site, self.default) def describe(self) -> dict: return {"default": self.default, "sites": dict(self.sites), "per": self.per} def maybe_quant(x: torch.Tensor, qp: Optional[QuantPolicy], site: str) -> torch.Tensor: if qp is None or not qp.enabled: return x lvl = qp.level_for(site) if lvl in ("fp16", "bf16"): return cast_float(x, lvl) if lvl == "fp32": return x return quantise(x, LEVELS[lvl], qp.per) def quantise_model_weights(model, bits: int, skip=("norm", "mem0", "step_emb")) -> dict: """In-place fake-quantise weights. Returns per-tensor relative error.""" err = {} with torch.no_grad(): for name, p in model.named_parameters(): if any(s in name for s in skip) or p.dim() < 2: continue ref = p.detach().clone() q = quantise(p.data, bits, "row") p.data.copy_(q) err[name] = float((q - ref).norm() / ref.norm().clamp_min(1e-9)) return err