laya-onnx / scripts /export.py
techtheist's picture
Laya ONNX exports: en int4/int8, multilingual int8, dynamic sequence length
7622ba3 verified
Raw History Blame Contribute Delete
2.63 kB
"""Export a Laya checkpoint to ONNX with a truly dynamic sequence length,
then check parity against PyTorch at several lengths and batch widths."""
import sys, time, json, os
import numpy as np, torch
from laya.agent import Agent
from laya.common import build_sequence, collate_items, QTYPES
ckpt, out = sys.argv[1], sys.argv[2]
agent = Agent(ckpt, compile=False, device="cpu")
model = agent.model.float().eval()
tok = agent.tok
def batch_for(texts, q):
items = []
for s in texts:
seq, mk = build_sequence(tok, s, q, agent.cfg["max_len"], agent.cfg["head_max_len"])
items.append({"ids": seq, "markers": mk, "qtype": QTYPES[q["t"]]})
b = collate_items([items], tok.pad_token_id)
return (b["input_ids"], b["attention_mask"], b["marker_pos"], b["marker_mask"], b["qtype"])
Q = {"t": "choice", "ins": "What is the relationship between `premise` and `hypothesis`?",
"crit": {"entailment": "the premise implies the hypothesis is true",
"neutral": "the premise neither implies nor contradicts the hypothesis",
"contradiction": "the premise implies the hypothesis is false"}}
ex = batch_for([{"premise": "a b c", "hypothesis": "d e f"}, {"premise": "x " * 40, "hypothesis": "y"}], Q)
B, L, K = torch.export.Dim("batch", max=64), torch.export.Dim("seq", min=8, max=8192), torch.export.Dim("markers", max=255)
t0 = time.time()
with torch.no_grad():
prog = torch.onnx.export(
model, ex, dynamo=True, opset_version=18,
input_names=["input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype"],
output_names=["logits", "act_logits"],
dynamic_shapes=({0: B, 1: L}, {0: B, 1: L}, {0: B, 1: K}, {0: B, 1: K}, {0: B}),
)
prog.optimize()
prog.save(out, external_data=True)
print("exported in %.0fs" % (time.time() - t0))
import onnxruntime as ort
sess = ort.InferenceSession(out, providers=["CPUExecutionProvider"])
worst = 0.0
for texts in ([{"premise": "the cache uses 7 retries", "hypothesis": "the cache uses 19 retries"}],
[{"premise": "word " * 150, "hypothesis": "other " * 3}, {"premise": "short", "hypothesis": "x"}],
[{"premise": "Engram uses Rust " * 30, "hypothesis": "TepinDB " * 60}]):
ins = batch_for(texts, Q)
with torch.no_grad():
ref = model(*ins)[0].numpy()
got = sess.run(["logits"], {n: t.numpy() for n, t in zip(
["input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype"], ins)})[0]
m = ins[3].numpy()
d = np.abs(ref - got)[m].max()
worst = max(worst, d)
print("seq", ins[0].shape, "max |dlogit|", d)
print("WORST", worst)