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()