onnx_inference_demo.py

script

← Back to skill

Content hash: 2d3580c42bba5bf0569c3d8b732f25e15e1e7e94af41121e873d264741abddf3
#!/usr/bin/env python3
"""ONNX Runtime inference demo: providers, optimization levels, IOBinding.

Self-contained example that works with any ONNX model or a trivial dummy.
Requires: pip install onnx onnxruntime numpy
"""

from __future__ import annotations

import sys
import time

import numpy as np

try:
    import onnxruntime as ort
except ImportError:
    print("Install: pip install onnxruntime-gpu  # or onnxruntime for CPU-only", file=sys.stderr)
    sys.exit(1)


def create_dummy_session():
    """Create a small ONNX model in memory if no file given."""
    import onnx
    from onnx import helper

    # Minimal model: y = x * W (matmul with a 2x2 identity)
    matmul = helper.make_node("MatMul", ["input", "weight"], ["output"])
    graph = helper.make_graph(
        [matmul], "dummy",
        [helper.make_tensor_value_info("input", onnx.TensorProto.FLOAT, [1, 2])],
        [helper.make_tensor_value_info("output", onnx.TensorProto.FLOAT, [1, 2])],
    )
    weight = np.array([[1.0, 0.0], [0.0, 1.0]], dtype=np.float32)
    model = helper.make_model(graph)
    model.ir_version = 7
    model.opset_import[0].version = 18

    init = model.graph.initializer.add()
    init.name = "weight"
    init.CopyFrom(helper.make_tensor("weight", onnx.TensorProto.FLOAT, [2, 2], weight.tobytes(), raw=True))
    return model.SerializeToString()


def setup_session(model_path: str | None = None, use_gpu: bool = True) -> ort.InferenceSession:
    """Configure an optimized ONNX Runtime session."""

    # Provider ordering: fastest first, fall back silently
    providers = []
    if use_gpu:
        providers.append("CUDAExecutionProvider")
    providers.append("CPUExecutionProvider")

    so = ort.SessionOptions()
    so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL

    if model_path:
        session = ort.InferenceSession(model_path, providers=providers, sess_options=so)
    else:
        model_bytes = create_dummy_session()
        session = ort.InferenceSession(model_bytes, providers=providers, sess_options=so)

    # CRITICAL: confirm the expected provider actually loaded
    actual = session.get_providers()
    print(f"Providers: {actual}")
    if actual[0] != providers[0]:
        print(f"  WARNING: {providers[0]} not loaded, using {actual[0]}", file=sys.stderr)

    return session


def benchmark(session, input_data: dict, warmup: int = 3, runs: int = 100) -> float:
    """Warm up, then measure mean latency."""
    # Warm-up: first run pays session init + graph optimization
    for _ in range(warmup):
        session.run(None, input_data)

    times: list[float] = []
    for _ in range(runs):
        t0 = time.perf_counter()
        session.run(None, input_data)
        times.append(time.perf_counter() - t0)

    mean_ms = np.mean(times) * 1000
    p95_ms = np.percentile(times, 95) * 1000
    print(f"  mean: {mean_ms:.3f}ms, p95: {p95_ms:.3f}ms (over {runs} runs)")
    return mean_ms


def main() -> None:
    model_path = sys.argv[1] if len(sys.argv) > 1 else None

    # CPU baseline
    print("=== CPU ===")
    cpu_session = setup_session(model_path, use_gpu=False)
    input_data = {"input": np.array([[3.0, 4.0]], dtype=np.float32)}
    output = cpu_session.run(None, input_data)
    print(f"  Output: {output[0].tolist()}")
    benchmark(cpu_session, input_data)

    # GPU attempt (if available)
    print("\n=== GPU ===")
    gpu_session = setup_session(model_path, use_gpu=True)
    output = gpu_session.run(None, input_data)
    print(f"  Output: {output[0].tolist()}")
    benchmark(gpu_session, input_data)

    print("\nProduction notes:")
    print("  1. Always print get_providers() - silent CPU fallback is common")
    print("  2. Reuse session (create once, warm once)")
    print("  3. IOBinding to skip CPU<->GPU copies in decoding loops")
    print("  4. For INT8 quantization: use onnxruntime.quantization module")
    print("  5. For TensorRT: disable ORT graph optims, TensorRTProvider first")


if __name__ == "__main__":
    main()