import numpy as np, onnxruntime as ort, torch, time, json D = "/Users/shreyas/work/rnd/gliner2/work/export/onnx" ref = torch.load("/Users/shreyas/work/rnd/gliner2/work/export/ref.pt") # real prompts captured from the Python processor (ids from parity.py cases), re-tokenised here to avoid loading torch model from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained("/Users/shreyas/work/rnd/gliner2/work/export") def build(qs, text): import re pat = re.compile(r"(?:https?://[^\s]+|www\.[^\s]+)|[a-z0-9._%+-]+@[a-z0-9.-]+\.[a-z]{2,}|@[a-z0-9_]+|\w+(?:[-_]\w+)*|\S", re.I) pieces, markers, groups = [], [], [] def add(s): pieces.extend(tok.tokenize(s)) for qi, (task, labels, descs) in enumerate(qs): if qi: add("[SEP_STRUCT]") prompt = task + "".join(f" [DESCRIPTION] {l}: {d}" for l, d in (descs or {}).items() if l in labels) add("("); add("[P]"); add(prompt); add("(") g = [] for l in labels: g.append(len(pieces)); add("[L]"); add(l) groups.append(g); add(")"); add(")") add("[SEP_TEXT]") for m in pat.finditer(text): add(m.group().lower()) ids = tok.convert_tokens_to_ids(pieces) return np.array([ids], np.int64), groups cases = json.load(open("/Users/shreyas/work/rnd/gliner2/work/export/conversion/cases.json")) base = ort.InferenceSession(D + "/model.onnx", providers=["CPUExecutionProvider"]) for variant in ["model_fp16.onnx", "model_q4.onnx", "model_q4f16.onnx"]: s = ort.InferenceSession(D + "/" + variant, providers=["CPUExecutionProvider"]) worst, agree, n, ms = 0.0, 0, 0, [] for text, qs in cases: ids, groups = build([(q[0], q[1], q[2]) for q in qs], text) flat = [p for g in groups for p in g] feed = {"input_ids": ids, "attention_mask": np.ones_like(ids), "marker_positions": np.array([flat], np.int64)} lb = base.run(["logits"], feed)[0][0] t0 = time.time(); lv = s.run(["logits"], feed)[0][0]; ms.append((time.time()-t0)*1000) off = 0 for g in groups: a, b = lb[off:off+len(g)], lv[off:off+len(g)]; off += len(g) pa = np.exp(a-a.max()); pa /= pa.sum(); pb = np.exp(b-b.max()); pb /= pb.sum() worst = max(worst, float(np.abs(pa-pb).max())); agree += int(pa.argmax()==pb.argmax()); n += 1 print(f"{variant:18s} argmax agreement {agree}/{n} worst |dp| {worst:.4f} cpu median {np.median(ms):.0f} ms")