editlens_roberta_modelkit / examples /onnx_inference.py
CoderBak's picture
Publish attributed EditLens model kit: FP32 default, FP16, experimental INT8
f7cb4b0 verified
Raw History Blame Contribute Delete
2.95 kB
"""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="<pad>")
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()