#!/usr/bin/env python3 """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 ( # type: ignore[no-redef] 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"(.*?)", 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())