Download tools/official_int4_reference.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 8.11 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_int4_reference.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/official_int4_reference.py
-
curl -L -o official_int4_reference.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_int4_reference.py
8.11 kB
| #!/usr/bin/env python3 | |
| """Evaluate released INT4 codes/scales without a second weight quantization. | |
| Compute is the same CUDA BF16 teacher-forcing reference as previous diagnostics, | |
| not native compressed INT4 execution and not an RK3588 speed measurement. | |
| """ | |
| import argparse | |
| import hashlib | |
| import importlib.metadata | |
| import importlib.util | |
| import json | |
| from contextlib import ExitStack | |
| from pathlib import Path | |
| import time | |
| import torch | |
| from safetensors import safe_open | |
| from transformers import AutoTokenizer | |
| from quant_loss_reference import load_model, evaluate | |
| REVISION = "65a6d1d71e01f73ba01e572992bbd69ea92c865f" | |
| def unpack4(packed, shape): | |
| assert packed.dtype == torch.int32 and len(shape) == 2 | |
| shifts = torch.arange(0, 32, 4, dtype=torch.int32, device=packed.device) | |
| return (((packed.unsqueeze(-1) >> shifts) & 15).reshape(shape[0], -1) | |
| [:, :shape[1]] - 8).to(torch.int8) | |
| def verify_unpack(): | |
| # Load the dependency's standalone helper without importing optional runtimes. | |
| dist = importlib.metadata.distribution("compressed-tensors") | |
| path = dist.locate_file("compressed_tensors/compressors/pack_quantized/helpers.py") | |
| spec = importlib.util.spec_from_file_location("ct_pack_helper", path) | |
| helper = importlib.util.module_from_spec(spec) | |
| spec.loader.exec_module(helper) | |
| codes = torch.arange(-8, 8, dtype=torch.int8).repeat(8).reshape(4, 32) | |
| packed = helper.pack_to_int32(codes, 4) | |
| assert torch.equal(unpack4(packed, codes.shape), codes) | |
| assert torch.equal(helper.unpack_from_int32(packed, 4, codes.shape), codes) | |
| return {"package_version": dist.version, "helper_sha256": hashlib.sha256(path.read_bytes()).hexdigest()}, helper | |
| def replace_weights(model, official, helper): | |
| index = json.loads((official / "model.safetensors.index.json").read_text())["weight_map"] | |
| cfg = json.loads((official / "config.json").read_text())["quantization_config"] | |
| weights = cfg["config_groups"]["group_0"]["weights"] | |
| assert cfg["format"] == "pack-quantized" | |
| assert weights["group_size"] == 32 and weights["num_bits"] == 4 and weights["symmetric"] | |
| assert not any(k.endswith(("weight_zero_point", "weight_g_idx")) for k in index) | |
| state = model.state_dict(keep_vars=True) | |
| projected = {k.removesuffix("_packed") if k.endswith("weight_packed") else k | |
| for k in index if not k.endswith(("weight_scale", "weight_shape"))} | |
| if projected != set(state): | |
| raise RuntimeError(f"tensor coverage mismatch: missing={set(state)-projected}, extra={projected-set(state)}") | |
| stats = {"revision": REVISION, "quantized_matrices": 0, "quantized_elements": 0, | |
| "raw_tensors": 0, "raw_elements": 0, "raw_differences": [], "raw_compute_casts": [], | |
| "original_weight_squared_sum": 0., "weight_error_squared_sum": 0., | |
| "bf16_reconstruction_rounding_squared_sum": 0., "scale_fp16_not_exact": 0, | |
| "compute": "original BF16 implementation; released codes * BF16 group32 scales, cast to BF16; no A8"} | |
| with ExitStack() as stack, torch.no_grad(): | |
| handles = {s: stack.enter_context(safe_open(official / s, framework="pt", device="cpu")) | |
| for s in sorted(set(index.values()))} | |
| def get(name): | |
| return handles[index[name]].get_tensor(name) | |
| for i, (name, param) in enumerate(state.items()): | |
| if name in index: | |
| value = get(name).to(param.device) | |
| if value.dtype != param.dtype: | |
| # Previous BF16 reference casts non-parameter router-bias buffers. | |
| # Match that reference arithmetic rather than silently changing it. | |
| if not name.endswith(".mlp.gate.expert_bias"): | |
| raise RuntimeError(f"raw dtype mismatch {name}: {value.dtype} vs {param.dtype}") | |
| stats["raw_compute_casts"].append({"name": name, "source": str(value.dtype), "compute": str(param.dtype)}) | |
| value = value.to(param.dtype) | |
| if not torch.equal(value, param): | |
| stats["raw_differences"].append(name) | |
| param.copy_(value) | |
| stats["raw_tensors"] += 1 | |
| stats["raw_elements"] += param.numel() | |
| else: | |
| assert ".mlp.experts." in name, name | |
| shape = tuple(get(name + "_shape").tolist()) | |
| assert shape == tuple(param.shape) | |
| packed = get(name + "_packed") | |
| if stats["quantized_matrices"] == 0: | |
| assert torch.equal(unpack4(packed, shape), helper.unpack_from_int32(packed, 4, shape)) | |
| codes = unpack4(packed.to(param.device), shape) | |
| scale = get(name + "_scale").to(param.device) | |
| assert scale.shape == (shape[0], shape[1] // 32) | |
| assert torch.isfinite(scale).all() and (scale >= 0).all() | |
| stats["scale_fp16_not_exact"] += int((scale.half().bfloat16() != scale).sum()) | |
| exact = (codes.reshape(shape[0], -1, 32).float() * scale.float().unsqueeze(-1)).reshape(shape) | |
| original = param.float() | |
| stats["original_weight_squared_sum"] += original.double().square().sum().item() | |
| stats["weight_error_squared_sum"] += (original-exact).double().square().sum().item() | |
| stats["bf16_reconstruction_rounding_squared_sum"] += (exact-exact.bfloat16().float()).double().square().sum().item() | |
| param.copy_(exact) | |
| stats["quantized_matrices"] += 1 | |
| stats["quantized_elements"] += param.numel() | |
| if (i+1) % 1000 == 0: | |
| print(json.dumps({"loaded_tensors": i+1, "total": len(state)}), flush=True) | |
| stats["expert_relative_weight_rmse"] = (stats["weight_error_squared_sum"] / stats["original_weight_squared_sum"]) ** .5 | |
| return stats | |
| def main(): | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument("--source", type=Path, required=True) | |
| p.add_argument("--official", type=Path, required=True) | |
| p.add_argument("--suite", type=Path, required=True) | |
| p.add_argument("--output-dir", type=Path, required=True) | |
| args = p.parse_args() | |
| args.output_dir.mkdir(parents=True, exist_ok=True) | |
| torch.set_num_threads(4) | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| checks, helper = verify_unpack() | |
| suite = json.loads(args.suite.read_text()) | |
| tok = AutoTokenizer.from_pretrained(args.official, trust_remote_code=True, local_files_only=True) | |
| for case in suite["cases"]: | |
| assert tok.encode(case["prompt_text"], add_special_tokens=False) == case["prompt_ids"] | |
| assert tok.encode(case["target_text"], add_special_tokens=False) == case["teacher_ids"] | |
| source_config = json.loads((args.source / "config.json").read_text()) | |
| official_config = json.loads((args.official / "config.json").read_text()) | |
| differences = {k: [source_config.get(k), official_config.get(k)] for k in set(source_config) | set(official_config) | |
| if source_config.get(k) != official_config.get(k)} | |
| allowed = {"quantization_config", "num_nextn_predict_layers", "model_type", "auto_map", "torch_dtype", "transformers_version", "_name_or_path"} | |
| if set(differences) - allowed: | |
| raise RuntimeError(f"unreviewed config differences: {differences}") | |
| print(json.dumps({"config_differences": differences, "unpack_check": checks}), flush=True) | |
| start = time.monotonic() | |
| model = load_model(args.source) | |
| stats = replace_weights(model, args.official, helper) | |
| stats.update({"unpack_check": checks, "config_differences": differences, "load_seconds": time.monotonic()-start, | |
| "suite_sha256": hashlib.sha256(args.suite.read_bytes()).hexdigest()}) | |
| (args.output_dir / "weight-audit.json").write_text(json.dumps(stats, indent=2) + "\n") | |
| evaluate(model, suite["cases"], args.output_dir, "official_int4_bf16") | |
| print("PASS: official INT4 weight reconstruction and paired likelihood evaluation", flush=True) | |
| if __name__ == "__main__": | |
| main() | |