"""Export a trained checkpoint to ONNX for fast CPU serving, and check it against PyTorch. python export_onnx.py --model my-router --task task.example.json --data data/example.jsonl Writes into ``/onnx/``: * ``model.onnx``: fp32, matches PyTorch on any CPU (the safe default); * ``model-int8-blockwise.onnx``: 8-bit block-wise encoder weights (ONNX Runtime MatMulNBits, block 32). It kept Raya's accuracy and ran ~10-15% faster on CPUs with VNNI int8 instructions (Intel Cascade Lake/Alder Lake+, AMD Zen 4+), but slower without VNNI. Skip with --fp32-only. Serve either with Laya's ONNXAgent (same answers and format as laya.Agent): from laya.onnx_agent import ONNXAgent agent = ONNXAgent("my-router", onnx_path="my-router/onnx/model.onnx"); agent.cfg["max_len"] = 512 Every exported file is checked on up to 20 rows of --data: the choice must not change and probabilities must stay within 0.001 (fp32) or 0.05 (int8) of PyTorch. """ from __future__ import annotations import argparse import json import os from pathlib import Path os.environ.setdefault("USE_TF", "0") import numpy as np import onnx import torch from laya import Agent from laya.onnx_agent import ONNXAgent from common import answer_probs, load_rows, load_task, row_state GRAPH_INPUTS = ("input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype") def capture_inputs(agent: Agent, state, question: dict) -> dict[str, np.ndarray]: """The network inputs Laya builds for one decision (used as the export example).""" captured: dict[str, np.ndarray] = {} model = agent.model class Capture(torch.nn.Module): def forward(self, *tensors): captured.update({n: t.numpy() for n, t in zip(GRAPH_INPUTS, tensors)}) return model(*tensors) agent.model = Capture() try: agent.system_one(state, {"q": question}) finally: agent.model = model return captured def export_fp32(agent: Agent, state, question: dict, path: Path) -> None: example = capture_inputs(agent, state, question) seq = torch.export.Dim("seq", min=8, max=4096) opts = torch.export.Dim("opts", min=1, max=16) torch.onnx.export(agent.model.eval(), tuple(torch.from_numpy(example[k]) for k in GRAPH_INPUTS), str(path), input_names=list(GRAPH_INPUTS), output_names=["logits", "act_logits"], dynamic_shapes=({1: seq}, {1: seq}, {1: opts}, {1: opts}, None), dynamo=True, external_data=False) # Shapes recorded by the exporter contradict ONNX shape inference during quantization; # drop them (ONNX Runtime re-infers them at load). model = onnx.load(str(path)) del model.graph.value_info[:] onnx.save(model, str(path)) def encoder_weight_matmuls(path: Path, n_layers: int) -> list[str] | None: """Names of the encoder's weight matmuls (4 per ModernBERT/mmBERT layer), or None if not found.""" model = onnx.load(str(path), load_external_data=False) constants = {i.name for i in model.graph.initializer} producers = {o: n for n in model.graph.node for o in n.output} def is_constant(name: str) -> bool: p = producers.get(name) return name in constants or (p is not None and p.op_type in ("Transpose", "Cast", "Identity") and all(is_constant(i) for i in p.input)) names = [node.name for node in model.graph.node if node.op_type == "MatMul" and is_constant(node.input[1]) and "encoder.layers." in {p.key: p.value for p in node.metadata_props}.get("pkg.torch.onnx.name_scopes", "")] return names if len(names) == 4 * n_layers else None def quantize_blockwise(src: Path, dst: Path, nodes: list[str]) -> None: from onnxruntime.quantization.matmul_nbits_quantizer import MatMulNBitsQuantizer quantizer = MatMulNBitsQuantizer(onnx.load(str(src)), bits=8, block_size=32, is_symmetric=True, accuracy_level=4, nodes_to_include=nodes) quantizer.process() quantizer.model.save_model_to_file(str(dst), use_external_data_format=False) def check(agent: Agent, model_dir: str, onnx_path: Path, rows, question, labels, max_tokens, tolerance) -> float: served = ONNXAgent(model_dir, onnx_path=str(onnx_path)) served.cfg["max_len"] = max_tokens worst = 0.0 for row in rows: state = row_state(row) want = answer_probs(agent.system_one(state, {"q": question})["answers"]["q"], question, labels) got = answer_probs(served.system_one(state, {"q": question})["answers"]["q"], question, labels) delta = max(abs(a - b) for a, b in zip(want, got)) worst = max(worst, delta) if int(np.argmax(want)) != int(np.argmax(got)) or delta > tolerance: raise SystemExit(f"{onnx_path.name} disagrees with PyTorch on row {row['id']} (delta {delta:.4f})") return worst def main() -> None: ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--model", required=True, help="trained checkpoint directory") ap.add_argument("--task", required=True) ap.add_argument("--data", required=True, help="JSONL rows used to check the export (labels not needed)") ap.add_argument("--max-tokens", type=int, default=512, help="serving token budget (match training)") ap.add_argument("--fp32-only", action="store_true") args = ap.parse_args() task = load_task(args.task) labels, question = task["labels"], task["questions"][0] rows = load_rows(args.data, labels, require_labels=False)[:20] agent = Agent(args.model, device="cpu") agent.cfg["max_len"] = args.max_tokens out_dir = Path(args.model) / "onnx" out_dir.mkdir(exist_ok=True) fp32 = out_dir / "model.onnx" export_fp32(agent, row_state(rows[0]), question, fp32) print(f"{fp32}: max probability difference vs PyTorch " f"{check(agent, args.model, fp32, rows, question, labels, args.max_tokens, 1e-3):.5f}") if args.fp32_only: return n_layers = json.loads((Path(args.model) / "encoder" / "config.json").read_text())["num_hidden_layers"] nodes = encoder_weight_matmuls(fp32, n_layers) if nodes is None: print("int8: encoder is not ModernBERT/mmBERT-shaped; skipping block-wise quantization") return int8 = out_dir / "model-int8-blockwise.onnx" quantize_blockwise(fp32, int8, nodes) print(f"{int8}: max probability difference vs PyTorch " f"{check(agent, args.model, int8, rows, question, labels, args.max_tokens, 0.05):.5f}") if __name__ == "__main__": main()