"""Run the portable GLiNER2.5 Multi package with CompiledModel on the CPU. With --seq auto, --model selects the storage family (wfp16 or fp32); the smallest fitting sibling graph is selected using both processed text slots and encoded token counts. No input is silently truncated. """ import argparse import contextlib import json import os from pathlib import Path import re import sys os.environ.setdefault("HF_HUB_OFFLINE", "1") os.environ.setdefault("TRANSFORMERS_OFFLINE", "1") os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") sys.dont_write_bytecode = True PACKAGE = Path(__file__).resolve().parents[1] sys.path.insert(0, str(PACKAGE / "host_assets/runtime")) import numpy as np import torch from ai_edge_litert import schema_py_generated as schema from ai_edge_litert.compiled_model import CompiledModel, CpuOptions, HardwareAccelerator, Options from host_runtime import HostRuntime def graph_path(model, seq): if model in {"fp32", "wfp16"}: storage = model parent = PACKAGE / "fp32" if storage == "fp32" else PACKAGE else: candidate = Path(model) if not candidate.is_absolute(): candidate = PACKAGE / candidate match = re.fullmatch(r"gliner25_multi_s(?:128|256|512)_(fp32|wfp16)\.tflite", candidate.name) if match is None: raise ValueError("--model must be wfp16, fp32, or a supplied graph filename") storage, parent = match.group(1), candidate.parent path = parent / f"gliner25_multi_s{seq}_{storage}.tflite" if not path.is_file(): raise FileNotFoundError(f"Selected graph is absent: {path.name}") return path def run_graph(path, inputs): # Signature order is buffer order. Map by the five distinct input shapes, # rather than assuming alphabetical tensor names or tensor-index order. blob = path.read_bytes() flat = schema.Model.GetRootAsModel(blob, 0) signature = flat.SignatureDefs(0) graph = flat.Subgraphs(signature.SubgraphIndex()) actual_shapes = [tuple(int(v) for v in x.shape) for x in inputs] order = [actual_shapes.index(tuple(int(v) for v in graph.Tensors(signature.Inputs(i).TensorIndex()).ShapeAsNumpy())) for i in range(signature.InputsLength())] assert len(order) == 5 and len(set(order)) == 5 assert signature.OutputsLength() == 1 output_shape = tuple(int(v) for v in graph.Tensors(signature.Outputs(0).TensorIndex()).ShapeAsNumpy()) del graph, signature, flat, blob model = CompiledModel.from_file(str(path), options=Options( hardware_accelerators=HardwareAccelerator.CPU, cpu_options=CpuOptions(num_threads=4))) ins, outs = [], [] try: ins = model.create_input_buffers(0) outs = model.create_output_buffers(0) for buffer, index in zip(ins, order): buffer.write(np.ascontiguousarray(inputs[index], dtype=np.float32)) model.run_by_index(0, ins, outs) packed = outs[0].read(int(np.prod(output_shape)), np.float32).reshape(output_shape).copy() if not np.isfinite(packed).all(): raise RuntimeError("The graph produced NaN or Inf") return packed finally: for buffer in ins + outs: buffer.destroy() model.close() def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", default="wfp16", help="wfp16, fp32, or a graph path within the package") parser.add_argument("--text", required=True) parser.add_argument("--splitter", required=True, choices=("whitespace", "char")) parser.add_argument("--table", default="fp16", choices=("fp16", "fp32")) parser.add_argument("--seq", default="auto", choices=("auto", "128", "256", "512")) args = parser.parse_args() torch.set_num_threads(4) # Keep stdout machine-readable even if an upstream constructor prints. with contextlib.redirect_stdout(sys.stderr), torch.inference_mode(): host = HostRuntime(PACKAGE / "host_assets", table=args.table) sizes = (128, 256, 512) if args.seq == "auto" else (int(args.seq),) for seq in sizes: try: inputs, captured = host.prepare(args.text, seq, args.splitter) break except ValueError as error: if not str(error).startswith("Input exceeds encoded/text capacity:"): raise if seq == sizes[-1]: raise ValueError("Text exceeds the supplied window capacities; split it into shorter inputs") from error path = graph_path(args.model, seq) packed = run_graph(path, inputs) decoded = host.decode(captured, packed, inputs) output = {"model": path.relative_to(PACKAGE).as_posix(), "seq": seq, "splitter": args.splitter, "table": args.table, "text_slots": int(captured["batch"].text_word_indices.shape[1]), "encoded_tokens": int(captured["batch"].input_ids.shape[1]), "all_finite": True, **decoded} print(json.dumps(output, ensure_ascii=False, indent=2, allow_nan=False)) if __name__ == "__main__": main()