File size: 5,119 Bytes
3bcde0b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
"""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()