ivere27's picture
Initial release
ac68cef
Raw
History Blame Contribute Delete
8.72 kB
#!/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"<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())