raya / training /export_onnx.py
cderinbogaz's picture
Add training kit: train your own System-1 model
48c8658 verified
Raw History Blame Contribute Delete
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()