training / tinychess /quant.py
cazyundee's picture
tinychess: self-play research substrate (phase 1-4)
3495881 verified
Raw History Blame Contribute Delete
2.78 kB
"""
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