| |
| """Greedy TinyReceiptVQA inference using only NumPy, Pillow, and ONNX Runtime.""" |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import re |
| import unicodedata |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
| from PIL import Image |
|
|
| try: |
| from byte_bpe import ( |
| bpe1536_contract, |
| load_bpe1536_vocab, |
| ) |
| except ModuleNotFoundError: |
| from tiny_receipt_vqa.byte_bpe import ( |
| bpe1536_contract, |
| load_bpe1536_vocab, |
| ) |
|
|
|
|
| FAMILY_NAMES = [ |
| "phone", |
| "address", |
| "store", |
| "item_row", |
| "item_math", |
| "item_lookup", |
| "math", |
| "other", |
| ] |
|
|
|
|
| def read_json(path: Path) -> dict[str, Any]: |
| return json.loads(path.read_text(encoding="utf-8")) |
|
|
|
|
| def clean_text(value: object) -> str: |
| return unicodedata.normalize( |
| "NFC", |
| re.sub( |
| r"\s+", |
| " ", |
| str(value if value is not None else "").replace("\n", " "), |
| ).strip(), |
| ) |
|
|
|
|
| def parse_answer(text: str) -> tuple[str, bool]: |
| match = re.search(r"<answer>(.*?)</answer>", text) |
| if match is None: |
| return "", False |
| return clean_text(match.group(1)), True |
|
|
|
|
| def extract_answer(text: str) -> str: |
| return parse_answer(text)[0] |
|
|
|
|
| class BPE1536Tokenizer: |
| def __init__(self, data: dict[str, Any]): |
| self.bpe = load_bpe1536_vocab(data) |
| self.contract = bpe1536_contract(data) |
| self.itos = self.bpe.itos |
| self.stoi = self.bpe.stoi |
| self.pad = self.bpe.pad |
| self.bos = self.bpe.bos |
| self.eos = self.bpe.eos |
| self.unk = self.bpe.unk |
|
|
| def encode(self, text: str, max_length: int) -> np.ndarray: |
| ids = self.bpe.encode(text, add_eos=True, max_len=max_length) |
| return np.asarray(ids, dtype=np.int64)[None, :] |
|
|
| def decode(self, ids: np.ndarray) -> str: |
| return self.bpe.decode(ids.tolist()) |
|
|
|
|
| def preprocess_image(path: Path) -> np.ndarray: |
| image = Image.open(path).convert("L").resize((672, 320), Image.Resampling.BILINEAR) |
| pixels = np.asarray(image, dtype=np.float32) / 255.0 |
| pixels = (pixels - 0.5) / 0.5 |
| return np.ascontiguousarray(pixels[None, None, :, :]) |
|
|
|
|
| def family_input(value: str) -> tuple[np.ndarray, str]: |
| family = clean_text(value).lower() |
| if not family or family == "auto": |
| return np.asarray([-1], dtype=np.int64), "auto" |
| if family not in FAMILY_NAMES: |
| raise ValueError( |
| f"unknown family {value!r}; expected auto or one of {FAMILY_NAMES}" |
| ) |
| return np.asarray([FAMILY_NAMES.index(family)], dtype=np.int64), family |
|
|
|
|
| def select_providers(ort: Any, requested: str) -> list[str]: |
| available = list(ort.get_available_providers()) |
| if requested == "auto": |
| preferred = [ |
| "CUDAExecutionProvider", |
| "CoreMLExecutionProvider", |
| "CPUExecutionProvider", |
| ] |
| selected = [provider for provider in preferred if provider in available] |
| return selected or available |
| if requested not in available: |
| raise RuntimeError( |
| f"ONNX Runtime provider {requested!r} is unavailable; available={available}" |
| ) |
| providers = [requested] |
| if requested != "CPUExecutionProvider" and "CPUExecutionProvider" in available: |
| providers.append("CPUExecutionProvider") |
| return providers |
|
|
|
|
| def model_files( |
| manifest: dict[str, Any], |
| precision: str, |
| ) -> tuple[str, str, str]: |
| if precision == "fp32": |
| return ( |
| str(manifest["files"]["encoder"]), |
| str(manifest["files"]["decoder"]), |
| "fp32", |
| ) |
| variants = manifest.get("variants") or {} |
| variant_name = "int8_w8a8" |
| variant = variants.get(variant_name) |
| if not isinstance(variant, dict): |
| raise RuntimeError( |
| "manifest does not contain an INT8 ONNX variant" |
| ) |
| return ( |
| str(variant["encoder"]), |
| str(variant["decoder"]), |
| variant_name, |
| ) |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser(description="Ask an ONNX TinyReceiptVQA model.") |
| parser.add_argument( |
| "--model-dir", |
| default=".", |
| help="directory containing manifest.json, config.json, vocab.json, and ONNX files", |
| ) |
| parser.add_argument("--image", required=True) |
| parser.add_argument("--question", required=True) |
| parser.add_argument( |
| "--family", |
| default="auto", |
| help="auto uses the learned router; otherwise select an explicit adapter family", |
| ) |
| parser.add_argument("--max-len", type=int, default=0) |
| parser.add_argument( |
| "--provider", |
| default="auto", |
| help="auto or an ONNX Runtime execution provider name", |
| ) |
| parser.add_argument( |
| "--precision", |
| choices=("fp32", "int8"), |
| default="fp32", |
| help="int8 selects the static W8A8 U8S8 QDQ ONNX variant", |
| ) |
| parser.add_argument("--intra-op-threads", type=int, default=0) |
| args = parser.parse_args() |
|
|
| try: |
| import onnxruntime as ort |
| except ImportError as exc: |
| raise SystemExit("install runtime dependencies with: pip install onnxruntime numpy Pillow") from exc |
|
|
| model_dir = Path(args.model_dir) |
| image_path = Path(args.image) |
| if not image_path.is_file(): |
| raise FileNotFoundError(f"missing image: {image_path}") |
| manifest = read_json(model_dir / "manifest.json") |
| config = read_json(model_dir / manifest["files"]["config"]) |
| vocab = BPE1536Tokenizer( |
| read_json(model_dir / manifest["files"]["vocab"]) |
| ) |
| if int(config.get("vocab_size", 0)) != len(vocab.itos): |
| raise RuntimeError("config.json does not declare the BPE1536 vocabulary") |
| if manifest.get("tokenizer") != vocab.contract: |
| raise RuntimeError("manifest tokenizer contract does not match vocab.json") |
| question = clean_text(args.question) |
| if not question: |
| parser.error("--question must not be empty") |
| maximum_length = int(args.max_len or config["max_out_len"]) |
| if not 2 <= maximum_length <= int(config["max_out_len"]): |
| parser.error( |
| f"--max-len must be between 2 and {int(config['max_out_len'])}" |
| ) |
|
|
| session_options = ort.SessionOptions() |
| if args.intra_op_threads > 0: |
| session_options.intra_op_num_threads = args.intra_op_threads |
| providers = select_providers(ort, args.provider) |
| encoder_file, decoder_file, model_variant = model_files( |
| manifest, |
| args.precision, |
| ) |
| encoder = ort.InferenceSession( |
| str(model_dir / encoder_file), |
| sess_options=session_options, |
| providers=providers, |
| ) |
| decoder = ort.InferenceSession( |
| str(model_dir / decoder_file), |
| sess_options=session_options, |
| providers=providers, |
| ) |
|
|
| image = preprocess_image(image_path) |
| question_ids = vocab.encode(question, int(config["max_q_len"])) |
| requested_family_ids, requested_family = family_input(args.family) |
| memory, memory_padding_mask, router_logits, selected_family_ids = encoder.run( |
| None, |
| { |
| "image": image, |
| "question_ids": question_ids, |
| "family_ids": requested_family_ids, |
| }, |
| ) |
|
|
| generated_ids = np.asarray([[vocab.bos]], dtype=np.int64) |
| for _ in range(maximum_length - 1): |
| logits = decoder.run( |
| ["logits"], |
| { |
| "decoder_input_ids": generated_ids, |
| "memory": memory, |
| "memory_padding_mask": memory_padding_mask, |
| "family_ids": selected_family_ids, |
| }, |
| )[0] |
| next_ids = np.argmax(logits[:, -1, :], axis=-1).astype(np.int64) |
| generated_ids = np.concatenate([generated_ids, next_ids[:, None]], axis=1) |
| if np.all(next_ids == vocab.eos): |
| break |
|
|
| generated = vocab.decode(generated_ids[0]) |
| answer, well_formed = parse_answer(generated) |
| selected_id = int(selected_family_ids[0]) |
| result = { |
| "format": manifest["format"], |
| "model_variant": model_variant, |
| "providers": encoder.get_providers(), |
| "image": str(image_path), |
| "question": question, |
| "requested_family": requested_family, |
| "selected_family": ( |
| FAMILY_NAMES[selected_id] |
| if 0 <= selected_id < len(FAMILY_NAMES) |
| else str(selected_id) |
| ), |
| "router_logits": [float(value) for value in router_logits[0]], |
| "generated": generated, |
| "answer": answer, |
| "well_formed": well_formed, |
| } |
| print(json.dumps(result, ensure_ascii=False, indent=2)) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|