File size: 3,327 Bytes
d8b837f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 | #!/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()
|