#!/usr/bin/env python3 """Quantize the Laya ONNX graph to 8-bit weights (MatMulNBits) and check it on tests/vectors.json. # 1. Quantize (needs the ORT 1.30 quantization tools): uv run --no-project --python 3.12 --with onnxruntime==1.30.0 --with onnx==1.23.1 \ --with onnx-ir==1.0.0 python scripts/quantize_laya_nbits.py # 2. Check the file on ONNX Runtime 1.23.0: uv run --no-project --python 3.12 --with onnxruntime==1.23.0 --with numpy \ python scripts/quantize_laya_nbits.py --check-only The source is `onnx/laya-fp32.onnx` from scripts/export_laya_onnx.py (checkpoint convaiinnovations/laya at 55cf4c4e...). The output is `onnx/laya-q8.onnx`. Every MatMul weight becomes an asymmetric 8-bit block (block size 32) with fp32 compute (accuracy level 0). Embeddings, norms, and the rest stay fp32. The check prints the max probability error, the argmax changes, and the cases whose top probability changes side of 0.40, against the `laya` package. Other settings were worse on a private set of 67 routing inputs (max probability error): symmetric 8-bit (0.018), 8-bit with int8 compute (0.017), and every 4-bit setting (0.12 to 0.31). This setting gave 0.0086 there. """ from __future__ import annotations import argparse import hashlib import json import sys import time from pathlib import Path ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT / "tests")) from check import compare, load_vectors, run_onnx # noqa: E402 SOURCE = ROOT / "onnx/laya-fp32.onnx" OUT = ROOT / "onnx/laya-q8.onnx" THRESHOLD = 0.40 def sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as f: for chunk in iter(lambda: f.read(1 << 20), b""): digest.update(chunk) return digest.hexdigest() def quantize(source: Path, out: Path) -> None: import onnx from onnxruntime.quantization import matmul_nbits_quantizer as nbits settings = dict(block_size=32, is_symmetric=False, accuracy_level=None, bits=8, op_types_to_quantize=("MatMul",), quant_axes=(("MatMul", 0),)) config = nbits.DefaultWeightOnlyQuantConfig(**settings) quantizer = nbits.MatMulNBitsQuantizer(onnx.load(str(source)), algo_config=config, **settings) quantizer.process() onnx.save_model(quantizer.model.model, str(out)) def check(path: Path) -> dict: import onnxruntime as ort data, cases = load_vectors() probabilities = run_onnx(path, cases, data["temperature"], len(data["question"]["options"])) return {"file": path.name, "onnxruntime": ort.__version__, "bytes": path.stat().st_size, "sha256": sha256(path), **compare(cases, probabilities, THRESHOLD)} def main() -> None: parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) parser.add_argument("--source", type=Path, default=SOURCE) parser.add_argument("--out", type=Path, default=OUT) parser.add_argument("--check-only", action="store_true") args = parser.parse_args() if not args.check_only: start = time.perf_counter() quantize(args.source, args.out) print(f"quantized {args.out} in {time.perf_counter() - start:.0f} s " f"(source sha256 {sha256(args.source)})") print(json.dumps(check(args.out), indent=2)) if __name__ == "__main__": main()