Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 60 additions & 1 deletion CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,20 @@ set(protobuf_BUILD_TESTS OFF CACHE INTERNAL "Disable tests")
set(protobuf_MSVC_STATIC_RUNTIME OFF CACHE INTERNAL "Force build static runtime")
FetchContent_MakeAvailable(protobuf)

# Cross-compilation shim: protoc must run on the host arch at build time.
# When cross-compiling (e.g. x64 host → ARM64 target), FetchContent builds
# protoc for the target arch, which can't execute on the host. Pass
# -DPROTOC_EXECUTABLE=<path-to-host-protoc> to inject a pre-built host-arch
# binary as the protobuf::protoc imported target and set
# -Dprotobuf_BUILD_PROTOC_BINARIES=OFF to skip building the target-arch one.
if(DEFINED PROTOC_EXECUTABLE AND NOT TARGET protobuf::protoc)
add_executable(protobuf::protoc IMPORTED GLOBAL)
set_target_properties(protobuf::protoc PROPERTIES
IMPORTED_LOCATION "${PROTOC_EXECUTABLE}"
)
message(STATUS "Using pre-built protoc at ${PROTOC_EXECUTABLE}")
endif()

# Add ONNX
FetchContent_Declare(
onnx
Expand Down Expand Up @@ -227,25 +241,38 @@ if(onnxruntime_ep_tensorrt_BUILD_TESTS)
FetchContent_MakeAvailable(googletest)

add_executable(trt_ep_tests
tests/test_main.cc
tests/cuda_graph_test.cc
tests/tensorrt_basic_test.cc
tests/tensorrt_ep_dla_options_test.cc
)
target_compile_definitions(trt_ep_tests PRIVATE
-DONNX_NAMESPACE=onnx
-DONNX_ML
-DNOMINMAX
-DORT_API_MANUAL_INIT
-DEP_LIB_PATH="$<TARGET_FILE:onnxruntime_ep_tensorrt>"
)
target_include_directories(trt_ep_tests PRIVATE
"$<BUILD_INTERFACE:${ORT_INCLUDE_DIR}>"
"$<BUILD_INTERFACE:${PROJECT_SOURCE_DIR}/src>"
"$<BUILD_INTERFACE:${PROJECT_SOURCE_DIR}/tests>"
"$<BUILD_INTERFACE:${CUDAToolkit_INCLUDE_DIRS}>"
)
target_link_libraries(trt_ep_tests PRIVATE
${ORT_LIBS}
CUDA::cudart
onnx
protobuf::libprotobuf
GTest::gtest_main
GTest::gtest
)

# Copy the EP DLL next to the test binary so resolve_ep_lib() finds it first.
add_custom_command(TARGET trt_ep_tests POST_BUILD
COMMAND ${CMAKE_COMMAND} -E copy_if_different
"$<TARGET_FILE:onnxruntime_ep_tensorrt>"
"$<TARGET_FILE_DIR:trt_ep_tests>/"
COMMENT "Copying EP library next to test binary"
)

# Copy testdata to the build directory so tests can find it
Expand All @@ -265,6 +292,38 @@ if(onnxruntime_ep_tensorrt_BUILD_TESTS)
)
endif()

# DLA graph transform support is Windows-only for now: the transforms library is
# distributed as a static .lib (per MSFT request) and has not been validated on Linux.
option(onnxruntime_ep_tensorrt_DLA_TRANSFORMS "Enable DLA graph transforms (Windows only)" OFF)

if(onnxruntime_ep_tensorrt_DLA_TRANSFORMS)
if(NOT WIN32)
message(FATAL_ERROR
"onnxruntime_ep_tensorrt_DLA_TRANSFORMS=ON is only supported on Windows. "
"The DLA transforms library ships as a static .lib (per MSFT request) "
"and has not been validated on Linux.")
endif()
if(NOT DEFINED DLA_TRANSFORMS_ROOT)
message(FATAL_ERROR "DLA_TRANSFORMS_ROOT must be set when onnxruntime_ep_tensorrt_DLA_TRANSFORMS is ON")
endif()
target_compile_definitions(onnxruntime_ep_tensorrt PRIVATE USE_DLA_TRANSFORMS)
target_include_directories(onnxruntime_ep_tensorrt PRIVATE "${DLA_TRANSFORMS_ROOT}/include")
target_link_libraries(onnxruntime_ep_tensorrt PRIVATE "${DLA_TRANSFORMS_ROOT}/lib/dla_transforms.lib")

if(onnxruntime_ep_tensorrt_BUILD_TESTS)
target_sources(trt_ep_tests PRIVATE tests/tensorrt_ep_dla_transforms_test.cc)
target_compile_definitions(trt_ep_tests PRIVATE USE_DLA_TRANSFORMS)
target_include_directories(trt_ep_tests PRIVATE "${DLA_TRANSFORMS_ROOT}/include")
# Copy dla_transforms.dll next to the test binary so resolve_dla_transforms_dll() finds it.
add_custom_command(TARGET trt_ep_tests POST_BUILD
COMMAND ${CMAKE_COMMAND} -E copy_if_different
"${DLA_TRANSFORMS_ROOT}/bin/dla_transforms.dll"
"$<TARGET_FILE_DIR:trt_ep_tests>/"
COMMENT "Copying dla_transforms.dll next to test binary"
)
endif()
endif()

if(onnxruntime_ep_tensorrt_INSTALL)
# Installation target
include(GNUInstallDirs)
Expand Down
85 changes: 85 additions & 0 deletions python/tests/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

import os
from dataclasses import dataclass, field
from types import ModuleType

import pytest


@dataclass
class RegisteredEp:
ort: ModuleType
ep_name: str
library_path: str
registered_by_fixture: bool = field(default=False)


@pytest.fixture(scope="session")
def registered_ep() -> RegisteredEp:
"""Register the TRT EP library for the test session and unregister on teardown."""
ort = pytest.importorskip("onnxruntime")
ep_pkg = pytest.importorskip("onnxruntime_ep_tensorrt")

lib = ep_pkg.get_library_path()
ep_name = ep_pkg.get_ep_names()[0]

if not ep_name:
pytest.skip("onnxruntime_ep_tensorrt.get_ep_names() returned an empty name")
if not os.path.isfile(lib):
pytest.skip(f"TRT EP library not found: {lib}")
if not hasattr(ort, "register_execution_provider_library"):
pytest.skip("onnxruntime build does not expose register_execution_provider_library")
if not hasattr(ort, "get_ep_devices"):
pytest.skip("onnxruntime build does not expose get_ep_devices")

registered_by_fixture = False
try:
ort.register_execution_provider_library(ep_name, lib)
registered_by_fixture = True
except Exception as exc:
# EP may already be registered (e.g. test re-run in the same process).
# Continue only if a device is actually visible; otherwise the session is unusable.
devices = [d for d in getattr(ort, "get_ep_devices", lambda: [])()
if getattr(d, "ep_name", None) == ep_name]
if not devices:
pytest.skip(f"Failed to register TRT EP library: {exc}")

yield RegisteredEp(ort=ort, ep_name=ep_name, library_path=lib,
registered_by_fixture=registered_by_fixture)

if registered_by_fixture and hasattr(ort, "unregister_execution_provider_library"):
try:
ort.unregister_execution_provider_library(ep_name)
except Exception:
pass


@pytest.fixture(scope="session")
def has_dla(registered_ep: RegisteredEp) -> None:
"""Skip the test if DLA hardware is not available."""
if not os.environ.get("TRT_EP_HAS_DLA"):
pytest.skip("TRT_EP_HAS_DLA not set — no DLA hardware available")


@pytest.fixture(scope="session")
def has_dla_transforms(registered_ep: RegisteredEp) -> dict:
"""Skip if EP not built with DLA transforms; return the provider option to enable them."""
if os.environ.get("TRT_EP_HAS_DLA_TRANSFORMS") != "1":
pytest.skip("TRT_EP_HAS_DLA_TRANSFORMS not set — EP not built with USE_DLA_TRANSFORMS")
return {"trt_dla_transform_enable": "1"}


@pytest.fixture(scope="session")
def has_two_dla_cores(has_dla) -> None:
"""Skip if fewer than 2 DLA cores are available.

Set TRT_EP_DLA_CORE_COUNT to the number of DLA cores on the device.
Defaults to 1 if not set.
"""
count = int(os.environ.get("TRT_EP_DLA_CORE_COUNT", "1"))
if count < 2:
pytest.skip("TRT_EP_DLA_CORE_COUNT < 2 — only one DLA core available")
132 changes: 132 additions & 0 deletions python/tests/ort_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

import threading
from pathlib import Path
from typing import Callable, Mapping, TypeVar

import numpy as np
import pytest

_T = TypeVar("_T")


def run_in_large_stack(fn: Callable[[], _T], stack_mb: int = 64) -> _T:
"""Run fn() in a thread with a larger stack and return its result.

TRT/DLA engine compilation can exhaust the default Python thread stack
(~1 MB on ARM64 Windows). This helper runs the callable in a new thread
with a 64 MB stack and re-raises any exception so pytest.raises() works.
"""
result: list = [None]
exc: list = [None]

def _worker() -> None:
try:
result[0] = fn()
except BaseException as e: # noqa: BLE001
exc[0] = e

old_size = threading.stack_size(stack_mb * 1024 * 1024)
t = threading.Thread(target=_worker, daemon=True)
threading.stack_size(old_size)
t.start()
t.join()
if exc[0] is not None:
raise exc[0]
return result[0]

from conftest import RegisteredEp


def get_trt_ep_devices(registered_ep: RegisteredEp):
"""Return OrtEpDevice entries matching the registered TRT EP name."""
devices = [
d for d in registered_ep.ort.get_ep_devices()
if getattr(d, "ep_name", None) == registered_ep.ep_name
]
if not devices:
pytest.skip(f"No OrtEpDevice found for {registered_ep.ep_name}")
return devices


def make_session_options(
registered_ep: RegisteredEp,
provider_options: Mapping[str, str] | None = None,
session_config: Mapping[str, str] | None = None,
):
"""Build SessionOptions with the TRT EP appended via add_provider_for_devices."""
ort = registered_ep.ort
so = ort.SessionOptions()

for key, value in (session_config or {}).items():
so.add_session_config_entry(str(key), str(value))

if not hasattr(so, "add_provider_for_devices"):
pytest.skip("onnxruntime.SessionOptions does not expose add_provider_for_devices")

so.add_provider_for_devices(
get_trt_ep_devices(registered_ep),
{str(k): str(v) for k, v in (provider_options or {}).items()},
)
return so


def create_session(
registered_ep: RegisteredEp,
model_path_or_bytes,
provider_options: Mapping[str, str] | None = None,
session_config: Mapping[str, str] | None = None,
):
"""Create an InferenceSession with the TRT EP and the given options."""
so = make_session_options(registered_ep, provider_options=provider_options,
session_config=session_config)
return registered_ep.ort.InferenceSession(model_path_or_bytes, sess_options=so)


def _concrete_shape(shape: list, override: list | None = None) -> list[int]:
if override is not None:
return list(override)
return [d if isinstance(d, int) and d > 0 else 1 for d in shape]


def _numpy_dtype(ort_type: str):
mapping = {
"tensor(float)": np.float32,
"tensor(float16)": np.float16,
"tensor(double)": np.float64,
"tensor(int64)": np.int64,
"tensor(int32)": np.int32,
"tensor(int8)": np.int8,
"tensor(uint8)": np.uint8,
"tensor(uint16)": np.uint16,
"tensor(int16)": np.int16,
"tensor(uint32)": np.uint32,
"tensor(uint64)": np.uint64,
"tensor(bool)": np.bool_,
}
dt = mapping.get(ort_type)
if dt is None:
pytest.skip(f"No numpy mapping for ORT type {ort_type}")
return dt


def make_zero_feeds(
session,
shape_overrides: Mapping[str, list[int]] | None = None,
) -> dict:
"""Build a dict of zeroed numpy arrays matching the session's input signatures."""
feeds = {}
for inp in session.get_inputs():
shape = _concrete_shape(inp.shape, (shape_overrides or {}).get(inp.name))
dtype = _numpy_dtype(inp.type)
feeds[inp.name] = np.zeros(shape, dtype=dtype)
return feeds


def run_session_once(session, feeds=None, shape_overrides=None):
if feeds is None:
feeds = make_zero_feeds(session, shape_overrides)
return session.run(None, feeds)
72 changes: 72 additions & 0 deletions python/tests/test_basic_inference.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Basic inference test for the TensorRT plugin EP.

Model: single Conv op, FP16, 4D tensors
input "input" FLOAT16 [1, 4, 8, 8]
weight "weight" FLOAT16 [8, 4, 3, 3] (zero initializer)
output "output" FLOAT16 [1, 8, 6, 6]

Environment variables:
TRT_EP_HAS_DLA set to "1" to run with DLA-specific provider options;
otherwise runs on GPU.
"""

from __future__ import annotations

import os

import numpy as np
import pytest

from conftest import RegisteredEp
from ort_helpers import make_session_options


def _conv_fp16_model() -> bytes:
import onnx
import numpy as np
from onnx import TensorProto, helper, numpy_helper

inp = helper.make_tensor_value_info("input", TensorProto.FLOAT16, [1, 4, 8, 8])
output = helper.make_tensor_value_info("output", TensorProto.FLOAT16, [1, 8, 6, 6])
weight = numpy_helper.from_array(
np.ones((8, 4, 3, 3), dtype=np.float16), name="weight"
)
node = helper.make_node("Conv", inputs=["input", "weight"], outputs=["output"])
graph = helper.make_graph([node], "conv_fp16", [inp], [output], initializer=[weight])
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)])
model.ir_version = onnx.IR_VERSION
return model.SerializeToString()


def test_conv_fp16_inference(registered_ep: RegisteredEp):
dla = os.environ.get("TRT_EP_HAS_DLA") == "1"

provider_options: dict[str, str] = {"trt_fp16_enable": "1"}
if dla:
provider_options.update({
"trt_dla_enable": "1",
"trt_dla_core": "0",
"trt_dla_gpu_fallback_enable": "0",
"trt_dla_adjust_for_dla": "1",
"trt_dla_enable_uint8_asymmetric_quantization": "1",
})

so = make_session_options(registered_ep, provider_options=provider_options)
session = registered_ep.ort.InferenceSession(_conv_fp16_model(), sess_options=so)

# Non-zero input ensures data flows through the conv.
feeds = {"input": np.ones((1, 4, 8, 8), dtype=np.float16)}
outputs = session.run(None, feeds)

assert len(outputs) == 1, f"Expected 1 output, got {len(outputs)}"

# Shape check
assert outputs[0].shape == (1, 8, 6, 6), f"Unexpected output shape: {outputs[0].shape}"

# Accuracy check: ones input * ones weight → each output = C_in×kH×kW = 4×3×3 = 36.0
expected = np.full((1, 8, 6, 6), 36.0, dtype=np.float16)
np.testing.assert_allclose(outputs[0], expected, atol=1.0,
err_msg="Conv output does not match expected value of 36.0")
Loading