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)