GLiNER2.5-Multi-LiteRT / examples /run_example.py
mlboydaisuke's picture
GLiNER2.5 Multi for LiteRT: s128/s256/s512 wfp16 + fp32 graphs, fp16/fp32 host embedding tables, sparse decoder, Python host runtime, host contract, card
3bcde0b verified
Raw History Blame Contribute Delete
5.12 kB
"""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()