Download scripts/export.py from techtheist/laya-onnx: direct link, hf CLI and curl.
- Browser
- Download file 2.63 kB
-
https://huggingface.co/techtheist/laya-onnx/resolve/main/scripts/export.py
- Command line
-
hf download hf://techtheist/laya-onnx/scripts/export.py
-
curl -L -o export.py https://huggingface.co/techtheist/laya-onnx/resolve/main/scripts/export.py
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) | |