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