File size: 6,823 Bytes
3fd1a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Small paired think/no-think diagnostic, official BF16 and W4-rebuilt weights.

This is fixed-text likelihood/implementation validation, not a benchmark score.
Original checkpoints are read-only. W4 reconstruction is in GPU memory only.
"""
import argparse
import hashlib
import json
from pathlib import Path
import time

import torch
from safetensors import safe_open
from transformers import AutoModelForCausalLM, AutoTokenizer


def prepare(source, cases_path, output):
    tokenizer = AutoTokenizer.from_pretrained(source, trust_remote_code=True, local_files_only=True)
    cases = json.loads(cases_path.read_text())
    result = []
    for case in cases:
        for thinking in (False, True):
            prompt = tokenizer.apply_chat_template(
                [{"role": "user", "content": case["prompt"]}], tokenize=False,
                add_generation_prompt=True, enable_thinking=thinking)
            prefix = (case["reasoning"] + "</think>") if thinking else ""
            target = prefix + case["answer"] + tokenizer.eos_token
            ids = tokenizer.encode(prompt, add_special_tokens=False)
            teacher = tokenizer.encode(target, add_special_tokens=False)
            assert tokenizer.encode(prompt + target, add_special_tokens=False) == ids + teacher
            item = {"id": case["id"] + ("_think" if thinking else "_nothink"),
                    "domain": case["domain"], "thinking": thinking,
                    "prompt_text": prompt, "target_text": target,
                    "prompt_ids": ids, "teacher_ids": teacher,
                    "answer_start": len(tokenizer.encode(prefix, add_special_tokens=False))}
            result.append(item)
    output.mkdir(parents=True, exist_ok=True)
    manifest = {"source": str(source.resolve()), "cases_sha256": hashlib.sha256(cases_path.read_bytes()).hexdigest(),
                "description": "synthetic fixed correct continuations; paired official templates; not task accuracy",
                "cases": result}
    (output / "suite.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2) + "\n")
    for item in result:
        print(json.dumps({"prepared": item["id"], "prompt_tokens": len(item["prompt_ids"]),
                          "teacher_tokens": len(item["teacher_ids"]), "answer_start": item["answer_start"]}), flush=True)
    return result


def load_model(source):
    model = AutoModelForCausalLM.from_pretrained(source, trust_remote_code=True,
        local_files_only=True, torch_dtype=torch.bfloat16, device_map={"": "cuda"},
        attn_implementation="eager").eval()
    index = json.loads((source / "model.safetensors.index.json").read_text())["weight_map"]
    parameters = dict(model.named_parameters())
    for shard in sorted(set(index.values())):
        with safe_open(source / shard, framework="pt", device="cpu") as handle:
            for name in handle.keys():
                if name in parameters and handle.get_slice(name).get_dtype() == "F32":
                    parameters[name].data = handle.get_tensor(name).to("cuda")
    return model


def evaluate(model, cases, output, name):
    with torch.inference_mode():
        for case in cases:
            start = time.monotonic()
            ids = torch.tensor([case["prompt_ids"] + case["teacher_ids"][:-1]], device="cuda")
            # Official full-batch teacher forcing. No sampled feedback and no old KV.
            result = model(input_ids=ids, use_cache=False)
            logits = result.logits[0, len(case["prompt_ids"]) - 1:].float().cpu()
            assert logits.shape == (len(case["teacher_ids"]), model.config.vocab_size)
            if not torch.isfinite(logits).all():
                raise RuntimeError("nonfinite official logits: " + case["id"])
            logits.numpy().tofile(output / (case["id"] + "." + name + ".f32"))
            print(json.dumps({"stage": name, "case": case["id"], "steps": len(logits),
                              "seconds": time.monotonic()-start}), flush=True)
            del result, logits, ids


def rebuild_w4(model):
    stats = {"method": "per-output-channel max(max_positive/7, max_negative/8), round-to-nearest-even, int4[-8,7]",
             "compute": "reconstructed weights cast to BF16; activations NOT A8", "matrices": 0,
             "elements": 0, "sum_squared_weights": 0.0, "sum_squared_error": 0.0}
    with torch.no_grad():
        for name, parameter in model.named_parameters():
            # Exactly the matrix families selected by tools/quantize_model.py.
            if parameter.ndim != 2 or name == "model.word_embeddings.weight" or name.endswith(".mlp.gate.weight"):
                continue
            if not (name.endswith(".weight") and (".attention." in name or ".mlp." in name or name == "lm_head.weight")):
                raise RuntimeError("unclassified matrix " + name)
            for offset in range(0, parameter.shape[0], 1024):
                original = parameter[offset:offset+1024].float()
                scales = torch.maximum(original.amax(1).clamp_min(0)/7,
                                       (-original.amin(1)).clamp_min(0)/8)
                scales = torch.where(scales == 0, torch.ones_like(scales), scales)
                reconstructed = torch.round(original/scales[:, None]).clamp(-8, 7)*scales[:, None]
                stats["elements"] += original.numel()
                stats["sum_squared_weights"] += original.double().square().sum().item()
                stats["sum_squared_error"] += (original.double()-reconstructed.double()).square().sum().item()
                parameter[offset:offset+1024].copy_(reconstructed.to(parameter.dtype))
            stats["matrices"] += 1
            if stats["matrices"] % 1000 == 0:
                print(json.dumps({"w4_rebuilt_matrices": stats["matrices"]}), flush=True)
    stats["relative_weight_rmse"] = (stats["sum_squared_error"]/stats["sum_squared_weights"])**.5
    return stats


def main():
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--source", type=Path, required=True)
    p.add_argument("--cases", type=Path, required=True)
    p.add_argument("--output-dir", type=Path, required=True)
    p.add_argument("--prepare-only", action="store_true")
    a = p.parse_args()
    torch.set_num_threads(4)
    torch.backends.cuda.matmul.allow_tf32 = False
    cases = prepare(a.source, a.cases, a.output_dir)
    if a.prepare_only:
        return
    model = load_model(a.source)
    evaluate(model, cases, a.output_dir, "bf16")
    stats = rebuild_w4(model)
    (a.output_dir / "weight-reconstruction.json").write_text(json.dumps(stats, indent=2) + "\n")
    evaluate(model, cases, a.output_dir, "w4_bf16")
    print("PASS: paired official BF16/W4-rebuilt reference complete", flush=True)


if __name__ == "__main__":
    main()