Download training/export_onnx.py from TextCortex/raya: direct link, hf CLI and curl.
- Browser
- Download file 6.67 kB
-
https://huggingface.co/TextCortex/raya/resolve/main/training/export_onnx.py
- Command line
-
hf download hf://TextCortex/raya/training/export_onnx.py
-
curl -L -o export_onnx.py https://huggingface.co/TextCortex/raya/resolve/main/training/export_onnx.py
6.67 kB
| """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 ``<model>/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() | |