fsi-anomaly / model /quant.py
FerrellSyntheticIntelligence's picture
backup all: 100 files (batch)
8b8e59d verified
Raw
History Blame Contribute Delete
1.73 kB
"""Our own Q8 quantization (GGUF Q8_0-style): int8 storage, fp32 compute.
torch's built-in dynamic quantization has no kernels on this ARM build
('unknown architecure'), so we implement the standard scheme used by GGUF
Q8_0: per-channel symmetric int8 weights with fp32 scales. Storage drops to
~1/4 of fp32; compute dequantizes on load, so behavior is near-lossless.
Usage:
qstate = quantize_q8(model) # {key: {"scale": [out], "q": int8}}
load_q8(model, qstate) # dequant into float weights in place
"""
import torch
@torch.no_grad()
def quantize_tensor(t: torch.Tensor) -> dict:
"""Per-channel symmetric int8 quantization of a [out, in] linear weight."""
t = t.float().contiguous()
out_dim = t.shape[0]
amax = t.abs().amax(dim=1, keepdim=True).clamp(min=1e-8)
scale = (amax / 127.0).squeeze(1)
q = torch.round(t / scale.view(-1, 1)).clamp(-127, 127).to(torch.int8)
return {"q": q, "scale": scale}
@torch.no_grad()
def dequantize_tensor(qstate: dict) -> torch.Tensor:
return qstate["q"].float() * qstate["scale"].view(-1, 1)
def quantize_q8(model: torch.nn.Module) -> dict:
qs = {}
for name, mod in model.named_modules():
if isinstance(mod, torch.nn.Linear):
prefix = name + ".weight"
qs[prefix] = quantize_tensor(mod.weight.detach())
return qs
def load_q8(model: torch.nn.Module, qstate: dict):
"""Dequantize Q8 states into the model's float weights (in place)."""
with torch.no_grad():
for name, mod in model.named_modules():
if isinstance(mod, torch.nn.Linear) and name + ".weight" in qstate:
mod.weight.copy_(dequantize_tensor(qstate[name + ".weight"]))
return model