File size: 8,098 Bytes
2bc6021 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """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()
|