Ling-3.0-tiny-RKNN / tools /quant_loss_reference.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
6.82 kB
#!/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()