File size: 6,666 Bytes
48c8658 | 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 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | """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()
|