"""Calibrate INT4 heads from pristine BF16 tensors and publish a local fast variant.""" import argparse import collections import copy import json import math import os import sys from pathlib import Path import numpy as np import torch from safetensors import safe_open from safetensors.torch import save_file from compressed_tensors.compressors.pack_quantized.base import pack_to_int32 from common import ROOT, RUN, SOURCE, BASELINE, FAST, read_json, write_json, records, sha256, stamp from checkpoint import clone, audit sys.path.insert(0, str(ROOT / "drafter")) from gptq_utils import accumulate_hessian, gptq_quantize, rtn_quantize, dequant def original(key): index = read_json(SOURCE / "model.safetensors.index.json")["weight_map"] with safe_open(SOURCE / index[key], "pt") as f: return f.get_tensor(key) def vocabulary(): from transformers import AutoTokenizer rows = records(RUN / "calibration/gen.jsonl") counts, held = collections.Counter(), collections.Counter() for r in rows: (held if r["split"] == "holdout" else counts).update(r["output_ids"]) assert counts and held tok = AutoTokenizer.from_pretrained(BASELINE) special = set(tok.all_special_ids) total = sum(counts.values()) selected = set(special) mass = sum(counts[t] for t in selected) for token, count in counts.most_common(): if token not in selected: selected.add(token) mass += count if mass / total >= .998 and len(selected) >= 16384: break if len(selected) >= 65536: break target_size = min(65536, max(16384, math.ceil(len(selected) / 128) * 128)) for token in list(counts) + torch.load(BASELINE / "mtp_draft_vocab_ids.pt", weights_only=True).tolist(): if len(selected) >= target_size: break selected.add(token) assert len(selected) == target_size ids = sorted(selected) report = {"size": len(ids), "calibration_tokens": total, "calibration_coverage": sum(counts[t] for t in ids) / total, "holdout_tokens": sum(held.values()), "holdout_coverage": sum(held[t] for t in ids) / sum(held.values()), "selection": "99.8% calibration output-token coverage, min 16384 and max 65536; rounded to 128 rows; all special ids included; reference IDs only pad if calibration vocabulary is too small"} write_json(RUN / "calibration/draft_vocab_ids.json", ids) write_json(RUN / "calibration/vocabulary-report.json", report) print(json.dumps(report), flush=True) def replace_tensors(model, replacements, bits): index = read_json(model / "model.safetensors.index.json") grouped = collections.defaultdict(dict) for module, tensors in replacements.items(): shard = index["weight_map"][module + ".weight_packed"] grouped[shard][module] = tensors for shard, modules in grouped.items(): # Never open a shared inode for writing: replace via a new file. with safe_open(model / shard, "pt") as f: data = {k: f.get_tensor(k) for k in f.keys() if not any(k.startswith(m + ".") for m in modules)} for m, tensors in modules.items(): data.update({m + "." + k: v for k, v in tensors.items()}) temporary = model / (shard + ".tmp") save_file(data, temporary, metadata={"format": "pt"}) os.replace(temporary, model / shard) del data config = read_json(model / "config.json") for group in bits: config["quantization_config"]["config_groups"][group]["weights"]["num_bits"] = bits[group] write_json(model / "config.json", config) if (model / "quantization_config.json").exists(): write_json(model / "quantization_config.json", config["quantization_config"]) def packed(q, scale): return {"weight_packed": pack_to_int32(q.cpu(), 4, packed_dim=1).contiguous(), "weight_scale": scale.cpu().to(torch.float16).contiguous(), "weight_shape": torch.tensor(list(q.shape), dtype=torch.int64)} def output_head(calib_rows): clone(BASELINE, FAST) hiddens = np.load(RUN / "calibration/hidden.npy", mmap_mode="r") seqs = read_json(RUN / "calibration/seqs.json") rng = np.random.default_rng(15027) pools = {"calibration": [], "holdout": []} for seq in seqs: # Include output-producing states; exclude prompt body and final nonpredicting row. start = seq["off"] + max(0, seq["n_prompt"] - 1) end = seq["off"] + seq["n"] - 1 if end > start: pools[seq["split"]].append(np.arange(start, end, dtype=np.int64)) selected = {} for split, size in [("calibration", calib_rows), ("holdout", 1024)]: pool = np.concatenate(pools[split]) selected[split] = np.sort(rng.choice(pool, size=min(size, len(pool)), replace=False)) assert not np.intersect1d(selected["calibration"], selected["holdout"]).size W = original("lm_head.weight").cuda() H = torch.zeros(W.shape[1], W.shape[1], device="cuda") seen = 0 for start in range(0, len(selected["calibration"]), 8192): x = torch.from_numpy(np.array(hiddens[selected["calibration"][start:start+8192]])).view(torch.bfloat16).cuda() H, seen = accumulate_hessian(H, x, seen) del x held = torch.from_numpy(np.array(hiddens[selected["holdout"]])).view(torch.bfloat16).cuda() @torch.no_grad() def kl(q, scale): dq = dequant(q.cuda(), scale.cuda()).to(torch.bfloat16) total = 0. for start in range(0, len(held), 64): x = held[start:start+64] p = (x @ W.t()).float().log_softmax(-1) lp = (x @ dq.t()).float().log_softmax(-1) total += (p.exp() * (p-lp)).sum().item() del dq return total / len(held) q, s = rtn_quantize(W, 4, 128) rtn_kl = kl(q, s) del q, s torch.cuda.empty_cache() qs, scales = [], [] for start in range(0, W.shape[0], 8192): q, s = gptq_quantize(W[start:start+8192], H, bits=4, group=128) qs.append(q.cpu()); scales.append(s.cpu()) print("lm_head rows", start + len(q), "/", W.shape[0], flush=True) del q, s q, s = torch.cat(qs), torch.cat(scales) gptq_kl = kl(q, s) assert math.isfinite(gptq_kl), gptq_kl report = {"calibration_rows": seen, "holdout_rows": len(held), "rtn_int4_kl": rtn_kl, "gptq_int4_kl": gptq_kl, "original": "pristine source lm_head BF16", "row_split": "whole examples before generation"} print(json.dumps(report), flush=True) # A worse result must be reviewed rather than silently promoted. if gptq_kl > rtn_kl: raise RuntimeError("GPTQ failed to improve held-out KL over round-to-nearest") replace_tensors(FAST, {"lm_head": packed(q, s)}, {"group_1": 4}) write_json(RUN / "calibration/lm-head-report.json", report) def mtp(): hessians = torch.load(RUN / "calibration/mtp_hessians.pt", weights_only=True, map_location="cpu") replacements, report = {}, {} for module, hessian in hessians.items(): W = original(module + ".weight").cuda() q, scale = gptq_quantize(W, hessian.cuda(), bits=4, group=128) error = ((dequant(q, scale) - W.float()).norm() / W.float().norm()).item() assert math.isfinite(error) report[module] = {"relative_weight_error": error, "shape": list(W.shape)} replacements[module] = packed(q, scale) print(module, error, flush=True) del W, q, scale torch.cuda.empty_cache() assert len(replacements) == 8, list(replacements) replace_tensors(FAST, replacements, {"group_3": 4}) write_json(RUN / "calibration/mtp-report.json", report) if __name__ == "__main__": ap = argparse.ArgumentParser() ap.add_argument("stage", choices=["vocabulary", "lm-head", "mtp"]) ap.add_argument("--calib-rows", type=int, default=300000) args = ap.parse_args() if args.stage == "vocabulary": vocabulary() elif args.stage == "lm-head": output_head(args.calib_rows) else: mtp()