Download evaluation/code/swift15/quantize.py from daavidhauser/Swift-1.5-Qwen3.8-27B-W4A16-HyperQwen: direct link, hf CLI and curl.
- Browser
- Download file 8.1 kB
-
https://huggingface.co/daavidhauser/Swift-1.5-Qwen3.8-27B-W4A16-HyperQwen/resolve/main/evaluation/code/swift15/quantize.py
- Command line
-
hf download hf://daavidhauser/Swift-1.5-Qwen3.8-27B-W4A16-HyperQwen/evaluation/code/swift15/quantize.py
-
curl -L -o quantize.py https://huggingface.co/daavidhauser/Swift-1.5-Qwen3.8-27B-W4A16-HyperQwen/resolve/main/evaluation/code/swift15/quantize.py
8.1 kB
| """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() | |
| 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() | |