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)