GLiNER2.5-Decide-ONNX / conversion /verify_parity.py
shreyask's picture
GLiNER2.5-Decide classification path, ONNX fp32/fp16/q4/q4f16
2a9b872 verified
Raw History Blame Contribute Delete
2.81 kB
import json, numpy as np, onnxruntime as ort, torch, time
from gliner2 import AutoExtractor
OUT = "/Users/shreyas/work/rnd/gliner2/work/export"
m = AutoExtractor.from_pretrained("fastino/GLiNER2.5-Decide"); m.eval()
proc, tok = m.processor, m.processor.tokenizer
captured = {}
orig = proc.collate_fn_inference
def hook(batch, *a, **k):
out = orig(batch, *a, **k); captured['out'] = out; return out
proc.collate_fn_inference = hook
cases = [
("My payouts have failed three times this week and nobody replied to my emails.",
[("team", ["payments","account","other"], None), ("urgent", ["yes","no"], None)]),
("Hi, I was charged twice for my Pro subscription this month, please fix.",
[("Which team should handle this?", ["billing team","technical support","sales"],
{"billing team":"charges, refunds, invoices","technical support":"bugs, outages","sales":"pricing, upgrades"}),
("Is the customer angry?", ["yes","no"], None)]),
("The export button crashes in Safari but works in Chrome. Not blocking, we told users to switch browsers.",
[("How severe is this bug?", ["cosmetic","degraded but there is a workaround","blocking"], None),
("Is the bug browser-specific?", ["no","yes"], None),
("Should this be escalated?", ["no","yes"], None)]),
]
sess = ort.InferenceSession(OUT + "/onnx/model.onnx", providers=["CPUExecutionProvider"])
print("onnx inputs:", [i.name for i in sess.get_inputs()], "outputs:", [o.name for o in sess.get_outputs()])
worst = 0.0
for text, qs in cases:
schema = m.create_schema()
for task, labels, descs in qs:
schema.classification(task, labels, label_descriptions=descs) if descs else schema.classification(task, labels)
py = m.extract(text, schema, include_confidence=True)
b = captured['out']
ids = b.input_ids.numpy().astype(np.int64); mask = b.attention_mask.numpy().astype(np.int64)
groups = [idx[1:] for idx in b.schema_special_indices[0]] # drop the [P] position
flat = [p for g in groups for p in g]
t0 = time.time()
logits = sess.run(["logits"], {"input_ids": ids, "attention_mask": mask, "marker_positions": np.array([flat], dtype=np.int64)})[0][0]
dt = (time.time()-t0)*1000
off = 0
for (task, labels, _), g in zip(qs, groups):
lg = logits[off:off+len(g)]; off += len(g)
probs = np.exp(lg - lg.max()); probs /= probs.sum()
best = int(probs.argmax())
py_label, py_conf = py[task]["label"], py[task]["confidence"]
diff = abs(probs[best] - py_conf) if labels[best] == py_label else 1.0
worst = max(worst, diff)
print(f" {task[:32]:32s} onnx={labels[best]:>34s} {probs[best]:.4f} | py={py_label:>34s} {py_conf:.4f} | diff={diff:.2e}")
print(f" seq={ids.shape[1]} tokens, onnx cpu {dt:.0f} ms")
print("WORST ABS DIFF:", worst)