Download tinychess/quant.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 2.78 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/quant.py
- Command line
-
hf download hf://spaces/cazyundee/training/tinychess/quant.py
-
curl -L -o quant.py https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/quant.py
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 | |
| 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 | |