import os, sys, json, gc, time, torch, torch.nn as nn OUT = "/Users/shreyas/work/rnd/gliner2/work/export"; os.makedirs(OUT + "/onnx", exist_ok=True) from gliner2 import AutoExtractor m = AutoExtractor.from_pretrained("fastino/GLiNER2.5-Decide"); m.eval() enc, clf, proc = m.encoder, m.classifier, m.processor # keep only what the classification path needs for name in list(vars(m).get("_modules", {})): if name not in ("encoder", "classifier"): delattr(m, name) proc.tokenizer.save_pretrained(OUT) cfg = json.loads(enc.config.to_json_string()); cfg["architectures"] = ["DebertaV2Model"] cfg["gliner2"] = {"description": "GLiNER2.5-Decide classification path: schema-conditioned DeBERTa-v3-large, one logit per [L] label marker from a 1024-2048-1 MLP head.", "special_tokens": proc.SPECIAL_TOKENS, "max_len": 512, "temperature": 1.0, "inputs": "input_ids, attention_mask, marker_positions (index of each [L] token)", "output": "logits (batch, markers); softmax within each question's markers", "source": "fastino/GLiNER2.5-Decide"} json.dump(cfg, open(OUT + "/config.json", "w"), indent=2) special = proc.SPECIAL_TOKENS del m, proc; gc.collect() class DecideGraph(nn.Module): def __init__(self, e, c): super().__init__(); self.encoder, self.classifier = e, c def forward(self, input_ids, attention_mask, marker_positions): h = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state g = torch.gather(h, 1, marker_positions.unsqueeze(-1).expand(-1, -1, h.size(-1))) return self.classifier(g).squeeze(-1) graph = DecideGraph(enc, clf).eval() ids = torch.tensor([[128004, 5, 6, 128007, 7, 128007, 8, 128002, 9, 10, 11]]); mask = torch.ones_like(ids); pos = torch.tensor([[3, 5]]) with torch.no_grad(): ref = graph(ids, mask, pos) torch.save({"ids": ids, "mask": mask, "pos": pos, "logits": ref}, OUT + "/ref.pt") print("ref", ref.tolist(), flush=True) t0 = time.time() mode = sys.argv[1] if len(sys.argv) > 1 else "dynamo" if mode == "dynamo": torch.onnx.export(graph, (ids, mask, pos), OUT + "/onnx/model.onnx", input_names=["input_ids", "attention_mask", "marker_positions"], output_names=["logits"], dynamic_shapes={"input_ids": {0: "batch", 1: "sequence"}, "attention_mask": {0: "batch", 1: "sequence"}, "marker_positions": {0: "batch", 1: "markers"}}, opset_version=17, dynamo=True, external_data=True, optimize=True) else: torch.onnx.export(graph, (ids, mask, pos), OUT + "/onnx/model.onnx", input_names=["input_ids", "attention_mask", "marker_positions"], output_names=["logits"], dynamic_axes={"input_ids": {0: "batch", 1: "sequence"}, "attention_mask": {0: "batch", 1: "sequence"}, "marker_positions": {0: "batch", 1: "markers"}, "logits": {0: "batch", 1: "markers"}}, opset_version=17, do_constant_folding=False, dynamo=False) print("exported", mode, "in", round(time.time() - t0, 1), "s", flush=True)