"""Local inference without PyTorch. Download the selected files before running. License: CC-BY-NC-SA-4.0. See LICENSE and NOTICE. """ import argparse import json from pathlib import Path import numpy as np import onnxruntime as ort from tokenizers import Tokenizer def main(): ort.disable_telemetry_events() p = argparse.ArgumentParser(description=__doc__) p.add_argument("text", nargs="+", help="One or more texts; each gets its own output.") p.add_argument("--model-dir", type=Path, default=Path(__file__).resolve().parents[1]) p.add_argument("--variant", choices=["fp32", "fp16", "int8"], default="fp32", help="INT8 is experimental and failed the published parity check; FP32 is recommended.") p.add_argument("--provider", choices=["cpu", "cuda", "coreml"], default="cpu", help="Only CPU was validated for this release. GPU options require a compatible runtime and testing.") args = p.parse_args() config = json.loads((args.model_dir / "config.json").read_text()) tok = Tokenizer.from_file(str(args.model_dir / "tokenizer.json")) tok.enable_truncation(max_length=512) tok.enable_padding(pad_id=config["pad_token_id"], pad_token="") enc = tok.encode_batch(args.text) feeds = {"input_ids": np.array([e.ids for e in enc], dtype=np.int64), "attention_mask": np.array([e.attention_mask for e in enc], dtype=np.int64)} files = {"fp32": "model.onnx", "fp16": "model_fp16.onnx", "int8": "model_int8.onnx"} provider = {"cpu": "CPUExecutionProvider", "cuda": "CUDAExecutionProvider", "coreml": "CoreMLExecutionProvider"}[args.provider] if provider not in ort.get_available_providers(): raise SystemExit("Requested provider is unavailable in this runtime: " + provider) options = ort.SessionOptions() options.intra_op_num_threads = 4 options.inter_op_num_threads = 1 providers = [provider] if args.provider == "cpu" else [provider, "CPUExecutionProvider"] session = ort.InferenceSession(str(args.model_dir / "onnx" / files[args.variant]), sess_options=options, providers=providers) logits = session.run(["logits"], feeds)[0].astype(np.float64) exp = np.exp(logits - logits.max(axis=1, keepdims=True)) probs = exp / exp.sum(axis=1, keepdims=True) rows = [{"label": config["id2label"][str(int(row.argmax()))], "probabilities": row.tolist()} for row in probs] print(json.dumps({"variant": args.variant, "experimental": args.variant == "int8", "requested_provider": provider, "configured_providers": session.get_providers(), "note": "Configured providers do not prove every operator ran on an accelerator. Class probabilities are not a percentage of AI-written words.", "results": rows}, indent=2)) if __name__ == "__main__": main()