#!/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"] + "") 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()