Download scripts/quantize_laya_nbits.py from charioteer/laya-mobile: direct link, hf CLI and curl.
- Browser
- Download file 3.33 kB
-
https://huggingface.co/charioteer/laya-mobile/resolve/main/scripts/quantize_laya_nbits.py
- Command line
-
hf download hf://charioteer/laya-mobile/scripts/quantize_laya_nbits.py
-
curl -L -o quantize_laya_nbits.py https://huggingface.co/charioteer/laya-mobile/resolve/main/scripts/quantize_laya_nbits.py
3.33 kB
| #!/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() | |