Instructions to use techtheist/laya-onnx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Laya
How to use techtheist/laya-onnx with Laya:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
File size: 2,628 Bytes
7622ba3 | 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 | """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)
|