From e8e336c8b706083debc7b8db1cf8c52375cbcdc9 Mon Sep 17 00:00:00 2001 From: yifeif <277870278+yifeif-nv@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:04:40 -0700 Subject: [PATCH 1/2] refactor(cli): migrate ten vision families (batch 3) Move build and runtime command declarations and handlers into eight timm families, DINOv3, and DETR. Preserve legacy build compatibility and the existing numerical acceptance criteria. Signed-off-by: yifeif <277870278+yifeif-nv@users.noreply.github.com> --- families/detr/cli.json | 135 +++++++++++++++ families/detr/cli.py | 101 ++++++++++++ families/detr/model.py | 23 +-- families/detr/runtime/CMakeLists.txt | 23 +++ families/detr/runtime/cli.cpp | 99 +++++++++++ families/detr/tests/cpp/test_cli.cpp | 142 ++++++++++++++++ .../detr/tests/manifests/detr-resnet-50.json | 3 +- families/detr/tests/test_cli.py | 142 ++++++++++++++++ families/detr/tests/test_e2e.py | 20 +-- families/dinov3/cli.json | 121 ++++++++++++++ families/dinov3/cli.py | 98 +++++++++++ families/dinov3/model.py | 40 +---- families/dinov3/runtime/CMakeLists.txt | 23 +++ families/dinov3/runtime/cli.cpp | 92 +++++++++++ families/dinov3/tests/cpp/test_cli.cpp | 141 ++++++++++++++++ ...inov3-convnext-tiny-pretrain-lvd1689m.json | 3 +- .../dinov3-vits16-pretrain-lvd1689m.json | 3 +- .../manifests/dinov3-vits16-timm-l0.json | 3 +- families/dinov3/tests/test_cli.py | 135 +++++++++++++++ families/dinov3/tests/test_e2e.py | 22 +-- families/timm_resnet/cli.json | 123 ++++++++++++++ families/timm_resnet/cli.py | 131 +++++++++++++++ families/timm_resnet/model.py | 28 +--- families/timm_resnet/runtime/CMakeLists.txt | 36 ++++ families/timm_resnet/runtime/cli.cpp | 110 +++++++++++++ families/timm_resnet/tests/cpp/test_cli.cpp | 155 ++++++++++++++++++ .../tests/manifests/resnet50-a1-in1k.json | 3 +- families/timm_resnet/tests/test_cli.py | 127 ++++++++++++++ families/timm_resnet/tests/test_e2e.py | 22 +-- families/timm_senet/cli.json | 123 ++++++++++++++ families/timm_senet/cli.py | 131 +++++++++++++++ families/timm_senet/model.py | 27 +-- families/timm_senet/runtime/CMakeLists.txt | 36 ++++ families/timm_senet/runtime/cli.cpp | 79 +++++++++ families/timm_senet/tests/cpp/test_cli.cpp | 138 ++++++++++++++++ .../tests/manifests/senet154-gluon-in1k.json | 3 +- families/timm_senet/tests/test_cli.py | 127 ++++++++++++++ families/timm_senet/tests/test_e2e.py | 16 +- families/timm_seresnet/cli.json | 123 ++++++++++++++ families/timm_seresnet/cli.py | 131 +++++++++++++++ families/timm_seresnet/model.py | 27 +-- families/timm_seresnet/runtime/CMakeLists.txt | 36 ++++ families/timm_seresnet/runtime/cli.cpp | 79 +++++++++ families/timm_seresnet/tests/cpp/test_cli.cpp | 138 ++++++++++++++++ .../tests/manifests/seresnet50-a1-in1k.json | 3 +- families/timm_seresnet/tests/test_cli.py | 127 ++++++++++++++ families/timm_seresnet/tests/test_e2e.py | 16 +- families/timm_swin/cli.json | 123 ++++++++++++++ families/timm_swin/cli.py | 131 +++++++++++++++ families/timm_swin/model.py | 27 +-- families/timm_swin/runtime/CMakeLists.txt | 36 ++++ families/timm_swin/runtime/cli.cpp | 110 +++++++++++++ families/timm_swin/tests/cpp/test_cli.cpp | 155 ++++++++++++++++++ .../swin-tiny-patch4-window7-224-ms-in1k.json | 3 +- families/timm_swin/tests/test_cli.py | 127 ++++++++++++++ families/timm_swin/tests/test_e2e.py | 16 +- families/timm_vgg/cli.json | 123 ++++++++++++++ families/timm_vgg/cli.py | 123 ++++++++++++++ families/timm_vgg/model.py | 28 +--- families/timm_vgg/runtime/CMakeLists.txt | 36 ++++ families/timm_vgg/runtime/cli.cpp | 79 +++++++++ families/timm_vgg/tests/cpp/test_cli.cpp | 138 ++++++++++++++++ .../tests/manifests/vgg16-tv-in1k.json | 3 +- families/timm_vgg/tests/test_cli.py | 126 ++++++++++++++ families/timm_vgg/tests/test_e2e.py | 22 +-- families/timm_vit/cli.json | 137 ++++++++++++++++ families/timm_vit/cli.py | 150 +++++++++++++++++ families/timm_vit/model.py | 29 +--- families/timm_vit/runtime/CMakeLists.txt | 36 ++++ families/timm_vit/runtime/cli.cpp | 110 +++++++++++++ families/timm_vit/tests/cpp/test_cli.cpp | 155 ++++++++++++++++++ ...base-p16-224-augreg-in21k-ft-in1k-tp4.json | 3 +- ...vit-base-p16-224-augreg-in21k-ft-in1k.json | 3 +- families/timm_vit/tests/test_cli.py | 150 +++++++++++++++++ families/timm_vit/tests/test_e2e.py | 19 +-- families/timm_xception/cli.json | 123 ++++++++++++++ families/timm_xception/cli.py | 131 +++++++++++++++ families/timm_xception/model.py | 27 +-- families/timm_xception/runtime/CMakeLists.txt | 36 ++++ families/timm_xception/runtime/cli.cpp | 79 +++++++++ families/timm_xception/tests/cpp/test_cli.cpp | 138 ++++++++++++++++ .../tests/manifests/xception41-tf-in1k.json | 3 +- families/timm_xception/tests/test_cli.py | 127 ++++++++++++++ families/timm_xception/tests/test_e2e.py | 16 +- families/timm_xcit/cli.json | 123 ++++++++++++++ families/timm_xcit/cli.py | 123 ++++++++++++++ families/timm_xcit/model.py | 28 +--- families/timm_xcit/runtime/CMakeLists.txt | 36 ++++ families/timm_xcit/runtime/cli.cpp | 110 +++++++++++++ families/timm_xcit/tests/cpp/test_cli.cpp | 155 ++++++++++++++++++ .../xcit-nano-12-p16-224-fb-in1k.json | 3 +- .../xcit-tiny-12-p16-224-fb-in1k.json | 3 +- families/timm_xcit/tests/test_cli.py | 126 ++++++++++++++ families/timm_xcit/tests/test_e2e.py | 16 +- 94 files changed, 6701 insertions(+), 364 deletions(-) create mode 100644 families/detr/cli.json create mode 100644 families/detr/cli.py create mode 100644 families/detr/runtime/cli.cpp create mode 100644 families/detr/tests/cpp/test_cli.cpp create mode 100644 families/detr/tests/test_cli.py create mode 100644 families/dinov3/cli.json create mode 100644 families/dinov3/cli.py create mode 100644 families/dinov3/runtime/cli.cpp create mode 100644 families/dinov3/tests/cpp/test_cli.cpp create mode 100644 families/dinov3/tests/test_cli.py create mode 100644 families/timm_resnet/cli.json create mode 100644 families/timm_resnet/cli.py create mode 100644 families/timm_resnet/runtime/cli.cpp create mode 100644 families/timm_resnet/tests/cpp/test_cli.cpp create mode 100644 families/timm_resnet/tests/test_cli.py create mode 100644 families/timm_senet/cli.json create mode 100644 families/timm_senet/cli.py create mode 100644 families/timm_senet/runtime/cli.cpp create mode 100644 families/timm_senet/tests/cpp/test_cli.cpp create mode 100644 families/timm_senet/tests/test_cli.py create mode 100644 families/timm_seresnet/cli.json create mode 100644 families/timm_seresnet/cli.py create mode 100644 families/timm_seresnet/runtime/cli.cpp create mode 100644 families/timm_seresnet/tests/cpp/test_cli.cpp create mode 100644 families/timm_seresnet/tests/test_cli.py create mode 100644 families/timm_swin/cli.json create mode 100644 families/timm_swin/cli.py create mode 100644 families/timm_swin/runtime/cli.cpp create mode 100644 families/timm_swin/tests/cpp/test_cli.cpp create mode 100644 families/timm_swin/tests/test_cli.py create mode 100644 families/timm_vgg/cli.json create mode 100644 families/timm_vgg/cli.py create mode 100644 families/timm_vgg/runtime/cli.cpp create mode 100644 families/timm_vgg/tests/cpp/test_cli.cpp create mode 100644 families/timm_vgg/tests/test_cli.py create mode 100644 families/timm_vit/cli.json create mode 100644 families/timm_vit/cli.py create mode 100644 families/timm_vit/runtime/cli.cpp create mode 100644 families/timm_vit/tests/cpp/test_cli.cpp create mode 100644 families/timm_vit/tests/test_cli.py create mode 100644 families/timm_xception/cli.json create mode 100644 families/timm_xception/cli.py create mode 100644 families/timm_xception/runtime/cli.cpp create mode 100644 families/timm_xception/tests/cpp/test_cli.cpp create mode 100644 families/timm_xception/tests/test_cli.py create mode 100644 families/timm_xcit/cli.json create mode 100644 families/timm_xcit/cli.py create mode 100644 families/timm_xcit/runtime/cli.cpp create mode 100644 families/timm_xcit/tests/cpp/test_cli.cpp create mode 100644 families/timm_xcit/tests/test_cli.py diff --git a/families/detr/cli.json b/families/detr/cli.json new file mode 100644 index 0000000000..beffaf9ae7 --- /dev/null +++ b/families/detr/cli.json @@ -0,0 +1,135 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one detr TensorRT bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Hugging Face model ID or local snapshot" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "object_detection" + ], + "default": "object_detection" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + }, + { + "name": "image_height", + "flags": [ + "--image-height" + ], + "type": "int" + }, + { + "name": "image_width", + "flags": [ + "--image-width" + ], + "type": "int" + } + ] + }, + { + "name": "detect", + "help": "Run detr on an image", + "executor": "native", + "handler": "detect", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + } + ] +} diff --git a/families/detr/cli.py b/families/detr/cli.py new file mode 100644 index 0000000000..421c1b4992 --- /dev/null +++ b/families/detr/cli.py @@ -0,0 +1,101 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""detr CLI handlers and build inputs; model imports stay lazy.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "object_detection" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + image_height: int | None = None + image_width: int | None = None + + def __post_init__(self) -> None: + if self.task != "object_detection": + raise ValueError("detr supports only task=object_detection") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("detr supports only fp16 or fp32") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + for name in ("image_height", "image_width"): + value = getattr(self, name) + if value is not None and (type(value) is not int or value < 1): + raise ValueError(f"{name} must be a positive integer") + + +def _legacy_common_options(request: object) -> None: + if request.task != "object_detection": + raise ValueError("detr supports only task=object_detection") + if request.backend not in {"trt", "trt_rtx"}: + raise ValueError("detr supports only backend=trt") + if request.dynamic_kv_cache: + raise NotImplementedError("detr does not support dynamic_kv_cache") + if request.max_sequence_length not in {None, 1}: + raise NotImplementedError("detr supports only max_sequence_length=1") + if request.max_batch_size != 1: + raise NotImplementedError("detr does not support max_batch_size") + + +def _legacy_model_options(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("detr does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("detr does not support context parallelism") + if request.video_num_frames is not None: + raise NotImplementedError("detr does not support video_num_frames") + if request.quantization not in {None, "none"}: + raise NotImplementedError("detr does not support quantization") + if request.fp32_layers: + raise NotImplementedError("detr does not support mixed-precision layers") + + +def coerce_request(request: object) -> BuildRequest: + """Preserve supported legacy inputs and reject unsupported nondefaults.""" + if isinstance(request, BuildRequest): + return request + _legacy_common_options(request) + _legacy_model_options(request) + fields = BuildRequest.__dataclass_fields__ + legacy = {"family", "output_path", "graph_transform", "dynamic_kv_cache", "image_height", + "image_width", "video_num_frames", "max_batch_size", "context_parallel_size", + "quantization", "fp32_layers", "tensor_parallel_size", "max_sequence_length"} + if unknown := set(vars(request)) - set(fields) - legacy: + raise ValueError(f"unknown detr build inputs: {sorted(unknown)}") + return BuildRequest(**{name: getattr(request, name) for name in fields}) + + +def build_bundle(request: BuildRequest, output: Path, *, transform: GraphTransform | None = None) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build(*, model: str, output: Path, revision: str | None = None, + task: str = "object_detection", precision: str = "fp32", backend: str = "trt", + verbose: bool = False, image_height: int | None = None, image_width: int | None = None) -> int: + request = BuildRequest(resolve_model(model, revision), task=task, precision=precision, + backend=backend, verbose=verbose, image_height=image_height, image_width=image_width) + build_bundle(request, output) + return 0 diff --git a/families/detr/model.py b/families/detr/model.py index 0aa752d035..47a9d69399 100644 --- a/families/detr/model.py +++ b/families/detr/model.py @@ -440,26 +440,9 @@ def get_bundle_config_overrides(self, config: ModelConfig) -> dict: def build(request, writer) -> None: """Build one DETR object-detection bundle.""" - if request.task != "object_detection": - raise ValueError("detr supports only task=object_detection") - if request.backend not in {"trt", "trt_rtx"}: - raise ValueError("detr supports only backend=trt") - if request.dynamic_kv_cache: - raise NotImplementedError("detr does not support dynamic_kv_cache") - if request.max_sequence_length not in {None, 1}: - raise NotImplementedError("detr supports only max_sequence_length=1") - if request.max_batch_size != 1: - raise NotImplementedError("detr does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("detr does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("detr does not support context parallelism") - if request.video_num_frames is not None: - raise NotImplementedError("detr does not support video_num_frames") - if request.quantization not in {None, "none"}: - raise NotImplementedError("detr does not support quantization") - if request.fp32_layers: - raise NotImplementedError("detr does not support mixed-precision layers") + from .cli import coerce_request + + request = coerce_request(request) model_dir = Path(request.model_dir) config = ModelConfig.from_dir(model_dir) diff --git a/families/detr/runtime/CMakeLists.txt b/families/detr/runtime/CMakeLists.txt index 50e7154066..1a13dbd791 100644 --- a/families/detr/runtime/CMakeLists.txt +++ b/families/detr/runtime/CMakeLists.txt @@ -58,3 +58,26 @@ if(TRTMC_BUILD_TESTS) COMMAND test_detr_image_preprocess ) endif() + +# The CLI adapter owns image I/O and uses the existing family runtime contract. +add_library(trtmc_cli_detr SHARED cli.cpp) +target_include_directories(trtmc_cli_detr PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_detr SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_detr PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_detr PRIVATE -Wall -Wextra -Wpedantic) +set_target_properties(trtmc_cli_detr PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_detr LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_executable(test_detr_cli ${PROJECT_SOURCE_DIR}/families/detr/tests/cpp/test_cli.cpp cli.cpp) + target_include_directories(test_detr_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_include_directories(test_detr_cli SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) + target_link_libraries(test_detr_cli PRIVATE trtmc_core nlohmann_json::nlohmann_json) + target_compile_options(test_detr_cli PRIVATE -Wall -Wextra -Wpedantic) + add_test(NAME detr_cli COMMAND test_detr_cli) + set_tests_properties(detr_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/detr/runtime/cli.cpp b/families/detr/runtime/cli.cpp new file mode 100644 index 0000000000..8c91f25f35 --- /dev/null +++ b/families/detr/runtime/cli.cpp @@ -0,0 +1,99 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include +#include +#include +#include +#include +#include +#include + +namespace { +struct Image { + std::vector pixels; + int width = 0, height = 0; +}; + +Image read_image(const std::string& path) { + Image result; + int channels = 0; + std::unique_ptr bytes( + stbi_load(path.c_str(), &result.width, &result.height, &channels, 3), stbi_image_free); + if (!bytes || result.width <= 0 || result.height <= 0) + throw std::runtime_error("unable to decode image: " + path); + const auto count = static_cast(result.width) * result.height * 3U; + result.pixels.resize(count); + for (std::size_t i = 0; i < count; ++i) + result.pixels[i] = static_cast(bytes.get()[i]) / 255.0F; + return result; +} + +void require_finite(const std::vector& values) { + for (const auto value : values) { + if (!std::isfinite(value)) + throw std::runtime_error("detr returned a non-finite result"); + } +} + +nlohmann::json execute(const std::string& handler, const nlohmann::json& values, + const char* default_runtime_root) { + if (handler != "detect") + throw std::invalid_argument("unknown detr CLI handler: " + handler); + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "detr") + throw std::invalid_argument("detr CLI requires a detr bundle"); + const auto image = read_image(values.at("image").get()); + auto task = trtmc::load_task( + reader, values.value("runtime_root", std::string(default_runtime_root)), 0, + values.value("runtime_cache", std::string()), values.value("cuda_graphs", false)); + auto* model = dynamic_cast(task.get()); + if (!model) + throw std::invalid_argument("detr bundle does not implement IObjectDetection"); + const auto result = model->detect(image.pixels.data(), image.height, image.width); + std::vector boxes, scores; + std::vector classes; + for (const auto& box : result.boxes) { + boxes.insert(boxes.end(), {box.x_min, box.y_min, box.x_max, box.y_max}); + scores.push_back(box.score); + classes.push_back(box.class_id); + } + require_finite(boxes); + require_finite(scores); + return {{"boxes", boxes}, + {"scores", scores}, + {"classes", classes}, + {"image_height", result.image_height}, + {"image_width", result.image_width}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + const auto result = + execute(handler, nlohmann::json::parse(values_json), default_runtime_root).dump() + + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } catch (...) { + const std::string message = "Error: detr CLI failed\n"; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/detr/tests/cpp/test_cli.cpp b/families/detr/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..af7ea3fcb1 --- /dev/null +++ b/families/detr/tests/cpp/test_cli.cpp @@ -0,0 +1,142 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0, task_loads = 0; +bool bad_output = false, loaded_graphs = false; +std::string loaded_runtime, loaded_cache; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +class FakeModel final : public trtmc::IObjectDetection { + public: + trtmc::ObjectDetectionResult detect(const float* pixels, std::int32_t height, + std::int32_t width) override { + check(height == 1 && width == 2, "owner passes image dimensions"); + check(pixels[0] == 1.0F && pixels[1] == 128.0F / 255.0F && pixels[2] == 0.0F, + "owner passes normalized HWC RGB"); + return {{{0.0F, 0.0F, 2.0F, 1.0F, + bad_output ? std::numeric_limits::infinity() : 0.9F, 4}}, + height, + width}; + } +}; +void write_bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "object_detection"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream out(path, std::ios::binary); + out.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + out.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + out.write(header.data(), static_cast(header.size())); + out.write("PLAN", 4); +} +struct Capture { + int status; + std::string output, error; +}; +void output(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void error(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +Capture invoke(const char* handler, const Json& values) { + Capture result{}; + result.status = trtmc_family_cli_v1(handler, values.dump().c_str(), "/installed/runtime", + &result, output, error); + return result; +} +void contract(const fs::path& root) { + const auto bundle = root / "model.bundle", image = root / "input.ppm"; + write_bundle(bundle, "detr"); + { + std::ofstream out(image, std::ios::binary); + out << "P6\n2 1\n255\n"; + const std::array rgb{255, 128, 0, 0, 255, 64}; + out.write(reinterpret_cast(rgb.data()), rgb.size()); + } + Json values{{"bundle", bundle.string()}, {"image", image.string()}}; + const auto first = invoke("detect", values); + check(first.status == 0 && first.error.empty(), "valid owner request succeeds"); + const auto output = Json::parse(first.output); + check(output.at("boxes") == Json({0.0, 0.0, 2.0, 1.0}), "detection coordinates are preserved"); + check(output.at("scores").at(0).get() == 0.9F && output.at("classes") == Json({4}), + "detection scores and classes are preserved"); + check(output.at("image_height") == 1 && output.at("image_width") == 2, + "image dimensions are preserved"); + + check(loaded_runtime == "/installed/runtime" && loaded_cache.empty() && !loaded_graphs, + "default load options are preserved"); + values.update({{"runtime_root", "/override/runtime"}, + {"runtime_cache", "runtime.cache"}, + {"cuda_graphs", true}}); + check(invoke("detect", values).status == 0, "explicit runtime options succeed"); + check(loaded_runtime == "/override/runtime" && loaded_cache == "runtime.cache" && loaded_graphs, + "runtime options reach the loader"); + bad_output = true; + const auto invalid_output = invoke("detect", values); + check(invalid_output.status != 0 && invalid_output.output.empty(), + "non-finite runtime results fail closed"); + bad_output = false; + const auto before = task_loads; + values["image"] = (root / "missing.png").string(); + check(invoke("detect", values).status != 0 && task_loads == before, + "invalid image fails before model loading"); + values["image"] = image.string(); + write_bundle(bundle, "another_family"); + check(invoke("detect", values).status != 0 && task_loads == before, + "wrong-family bundle fails before model loading"); + check(invoke("unknown", values).status != 0 && task_loads == before, + "unknown handler never loads a task"); +} +} // namespace + +namespace trtmc { +std::unique_ptr load_task(const BundleReader&, const std::string& runtime_root, + std::uint64_t, const std::string& runtime_cache, + bool cuda_graphs) { + ++task_loads; + loaded_runtime = runtime_root; + loaded_cache = runtime_cache; + loaded_graphs = cuda_graphs; + return std::make_unique(); +} +} // namespace trtmc +int main() { + const auto root = fs::temp_directory_path() / ("trtmc-detr-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(root); + } catch (const std::exception& exception) { + std::cerr << exception.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures == 0 ? 0 : 1; +} diff --git a/families/detr/tests/manifests/detr-resnet-50.json b/families/detr/tests/manifests/detr-resnet-50.json index ec9ac5b010..bf84b8d44a 100644 --- a/families/detr/tests/manifests/detr-resnet-50.json +++ b/families/detr/tests/manifests/detr-resnet-50.json @@ -16,6 +16,5 @@ "score_threshold": 0.5 } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/detr/tests/test_cli.py b/families/detr/tests/test_cli.py new file mode 100644 index 0000000000..40b49f1cfe --- /dev/null +++ b/families/detr/tests/test_cli.py @@ -0,0 +1,142 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import replace +import json +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest + +from families.detr import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli +from trtmc_benchmark.catalog import ManifestCatalog +from trtmc_benchmark.types import BenchmarkError + + +def test_declared_build_uses_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr(cli, "resolve_model", lambda model, revision: tmp_path) + monkeypatch.setattr(cli, "build_bundle", lambda request, output: calls.append((request, output))) + output = tmp_path / "out.bundle" + assert family_cli.main(["detr", "build", "checkpoint", "-o", str(output)]) == 0 + legacy = LegacyRequest(tmp_path, output, "detr", "object_detection", "fp32") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "dynamic_kv_cache") + assert not hasattr(calls[0][0], "tensor_parallel_size") + assert not hasattr(calls[0][0], "output_path") + + +@pytest.mark.parametrize("changes", [ + {"dynamic_kv_cache": True}, {"video_num_frames": 2}, {"max_batch_size": 2}, + {"tensor_parallel_size": 2}, {"context_parallel_size": 2}, + {"quantization": "fp8"}, {"fp32_layers": (1,)}, +]) +def test_legacy_rejects_unsupported_nondefaults(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", "detr", "object_detection", "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_legacy_fields_are_not_dropped(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", "detr", "object_detection", "fp32") + with pytest.raises(ValueError, match="unknown detr build inputs"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unrecognized_option=1)) + + +def test_help_and_rejected_options_do_not_import_handlers(monkeypatch, capsys): + discover_namespace = family_cli.importlib.import_module + def reject_import(name, *args, **kwargs): + if name == "families": + return discover_namespace(name, *args, **kwargs) + raise AssertionError(f"unexpected lazy import: {name}") + monkeypatch.setattr(family_cli.importlib, "import_module", reject_import) + for command in ("build", "detect"): + with pytest.raises(SystemExit) as result: + family_cli.main(["detr", command, "--help"]) + assert result.value.code == 0 + assert "--runtime-root" in capsys.readouterr().out + with pytest.raises(SystemExit) as result: + family_cli.main(["detr", "build", "checkpoint", "-o", "out.bundle", "--tensor-parallel-size", "2"]) + assert result.value.code == 2 + + +def test_bundle_selects_backend_before_import_and_publishes_atomically(monkeypatch, tmp_path): + events = [] + output = tmp_path / "model.bundle" + def build(request, writer): + events.append("build") + writer.set_header(family="detr", task="object_detection", backend=request.backend) + writer.add_bytes("engine.plan", b"fixture") + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, "families.detr.model", SimpleNamespace(build=build)) + monkeypatch.setattr(cli, "select_backend", select) + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), output) + assert events == ["trt_rtx", "build"] + published = output.read_bytes() + assert published.startswith(b"BUNDLE") + def fail(request, writer): + writer.set_header(family="detr", task="object_detection", backend=request.backend) + writer.add_bytes("engine.plan", b"partial") + raise RuntimeError("owner build failed") + monkeypatch.setattr(cli, "select_backend", lambda backend: monkeypatch.setitem( + sys.modules, "families.detr.model", SimpleNamespace(build=fail))) + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path), output) + assert output.read_bytes() == published + assert list(tmp_path.iterdir()) == [output] + + +def test_real_manifests_serialize_through_the_owner_handler(monkeypatch, tmp_path): + captured = [] + monkeypatch.setattr(cli, "resolve_model", lambda model, revision: tmp_path) + monkeypatch.setattr(cli, "build_bundle", lambda request, output: captured.append((request, output))) + spec = family_cli.load_family_cli("detr")["commands"][0] + manifests = sorted((Path(__file__).parent / "manifests").glob("*.json")) + assert manifests + for path in manifests: + model = ManifestCatalog._load(path) + values = dict(model.build_settings) + values.update(model=model.hf_id, output=str(tmp_path / model.bundle_name), task=model.task, precision=model.precision) + argv = family_cli.serialize_arguments(spec, values) + assert family_cli.main(["detr", "build", *argv]) == 0 + request, output = captured[-1] + assert type(request) is cli.BuildRequest and request.precision == model.precision + assert request.task == model.task and output == tmp_path / model.bundle_name + raw = json.loads(manifests[0].read_text()) + raw["build"] = {"unknown_owner_option": 1} + invalid = tmp_path / "invalid.json" + invalid.write_text(json.dumps(raw)) + with pytest.raises(BenchmarkError, match="undeclared build fields"): + ManifestCatalog._load(invalid) + + +def test_image_dimensions_are_owner_build_inputs(monkeypatch, tmp_path): + captured = [] + monkeypatch.setattr(cli, "resolve_model", lambda model, revision: tmp_path) + monkeypatch.setattr(cli, "build_bundle", lambda request, output: captured.append(request)) + assert family_cli.main(["detr", "build", "checkpoint", "-o", str(tmp_path / "out"), + "--image-height", "384", "--image-width", "640"]) == 0 + assert (captured[0].image_height, captured[0].image_width) == (384, 640) + legacy = LegacyRequest(tmp_path, tmp_path / "out", "detr", "object_detection", "fp32", image_height=384, image_width=640) + assert cli.coerce_request(legacy) == captured[0] + with pytest.raises(NotImplementedError): + cli.coerce_request(replace(legacy, max_sequence_length=2)) + with pytest.raises(ValueError, match="positive integer"): + cli.BuildRequest(tmp_path, image_height=0) + + +def test_e2e_build_preserves_single_device_rejection(monkeypatch, tmp_path): + from families.detr.tests import test_e2e + + calls = [] + monkeypatch.setattr(test_e2e, "build_bundle", lambda *args: calls.append(args)) + manifest = json.loads(next((Path(__file__).parent / "manifests").glob("*.json")).read_text()) + manifest["tensor_parallel_size"] = 2 + with pytest.raises(NotImplementedError, match="tensor parallelism"): + test_e2e._build(tmp_path, tmp_path / "out.bundle", manifest) + assert not calls diff --git a/families/detr/tests/test_e2e.py b/families/detr/tests/test_e2e.py index cf222d5751..14c88ea508 100644 --- a/families/detr/tests/test_e2e.py +++ b/families/detr/tests/test_e2e.py @@ -12,7 +12,7 @@ import pytest -from tensorrt_model_connect import BuildRequest, build +from families.detr.cli import BuildRequest, build_bundle FAMILY = "detr" TASKS = frozenset({"object_detection"}) @@ -114,6 +114,8 @@ def _runtime() -> tuple[Path, Path]: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise FileNotFoundError(f"selected {FAMILY} E2E requires its native CLI adapter") import torch assert torch.cuda.is_available(), f"selected {FAMILY} E2E requires CUDA" @@ -122,22 +124,17 @@ def _runtime() -> tuple[Path, Path]: def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError("detr does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=manifest.get("max_sequence_length"), image_height=manifest.get("image_height"), image_width=manifest.get("image_width"), - video_num_frames=manifest.get("video_num_frames"), - max_batch_size=int(manifest.get("max_batch_size", 1)), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - quantization=manifest.get("quantization"), - fp32_layers=tuple((int(layer) for layer in manifest.get("fp32_layers", ()))), - ) + ), + bundle, ) @@ -151,6 +148,7 @@ def _run_json( ) -> dict: invocation = [ str(binary), + FAMILY, command, str(bundle), "--runtime-root", diff --git a/families/dinov3/cli.json b/families/dinov3/cli.json new file mode 100644 index 0000000000..4f1e657432 --- /dev/null +++ b/families/dinov3/cli.json @@ -0,0 +1,121 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one dinov3 TensorRT bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Hugging Face model ID or local snapshot" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "image_features" + ], + "default": "image_features" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "extract-features", + "help": "Run dinov3 on an image", + "executor": "native", + "handler": "extract_features", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + } + ] +} diff --git a/families/dinov3/cli.py b/families/dinov3/cli.py new file mode 100644 index 0000000000..1493264df4 --- /dev/null +++ b/families/dinov3/cli.py @@ -0,0 +1,98 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""dinov3 CLI handlers and build inputs; model imports stay lazy.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "image_features" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "image_features": + raise ValueError("dinov3 supports only task=image_features") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("dinov3 supports only fp16 or fp32") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _legacy_common_options(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("dinov3 does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("dinov3 does not support image_height") + if request.image_width is not None: + raise NotImplementedError("dinov3 does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("dinov3 does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("dinov3 does not support max_batch_size") + + +def _legacy_model_options(request: object) -> None: + if request.context_parallel_size != 1: + raise ValueError("this family does not support context parallelism") + if request.task != "image_features": + raise ValueError("dinov3 supports only task=image_features") + if request.quantization not in {None, "none"}: + raise NotImplementedError("DINOv3 does not support quantization") + if request.fp32_layers: + raise NotImplementedError("DINOv3 does not support mixed-precision layers") + if request.tensor_parallel_size != 1: + raise NotImplementedError("DINOv3 does not support tensor parallelism") + sequence_length = request.max_sequence_length or 1 + if isinstance(sequence_length, bool) or int(sequence_length) < 1: + raise ValueError("max_sequence_length must be a positive integer") + + +def coerce_request(request: object) -> BuildRequest: + """Preserve supported legacy inputs and reject unsupported nondefaults.""" + if isinstance(request, BuildRequest): + return request + _legacy_common_options(request) + _legacy_model_options(request) + fields = BuildRequest.__dataclass_fields__ + legacy = {"family", "output_path", "graph_transform", "dynamic_kv_cache", "image_height", + "image_width", "video_num_frames", "max_batch_size", "context_parallel_size", + "quantization", "fp32_layers", "tensor_parallel_size", "max_sequence_length"} + if unknown := set(vars(request)) - set(fields) - legacy: + raise ValueError(f"unknown dinov3 build inputs: {sorted(unknown)}") + return BuildRequest(**{name: getattr(request, name) for name in fields}) + + +def build_bundle(request: BuildRequest, output: Path, *, transform: GraphTransform | None = None) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build(*, model: str, output: Path, revision: str | None = None, + task: str = "image_features", precision: str = "fp32", backend: str = "trt", + verbose: bool = False) -> int: + request = BuildRequest(resolve_model(model, revision), task=task, precision=precision, + backend=backend, verbose=verbose) + build_bundle(request, output) + return 0 diff --git a/families/dinov3/model.py b/families/dinov3/model.py index bb5ca50fe9..6b49f717bc 100644 --- a/families/dinov3/model.py +++ b/families/dinov3/model.py @@ -523,7 +523,7 @@ def add_residual(residual, tensor, scale: str): if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest + from .cli import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -637,37 +637,12 @@ def get_bundle_config_overrides(self, config: ModelConfig) -> dict: } -def _positive_int(value: object, name: str) -> int: - if isinstance(value, bool): - raise ValueError(f"{name} must be a positive integer") - result = int(value) - if result < 1: - raise ValueError(f"{name} must be a positive integer") - return result - - def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one DINOv3 bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("dinov3 does not support dynamic_kv_cache") - - if request.image_height is not None: - raise NotImplementedError("dinov3 does not support image_height") - - if request.image_width is not None: - raise NotImplementedError("dinov3 does not support image_width") - - if request.video_num_frames is not None: - raise NotImplementedError("dinov3 does not support video_num_frames") - - if request.max_batch_size != 1: - raise NotImplementedError("dinov3 does not support max_batch_size") + from .cli import coerce_request - if request.context_parallel_size != 1: - raise ValueError("this family does not support context parallelism") + request = coerce_request(request) - if request.task != "image_features": - raise ValueError("dinov3 supports only task=image_features") model_dir = Path(request.model_dir) config = ModelConfig.from_dir(model_dir) if str(config.model_type).lower() not in { @@ -677,20 +652,13 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: }: raise ValueError(f"DINOv3 does not support model_type={config.model_type!r}") precision = str(request.precision).lower() - max_sequence_length = _positive_int(request.max_sequence_length or 1, "max_sequence_length") - if request.quantization not in {None, "none"}: - raise NotImplementedError("DINOv3 does not support quantization") - if request.fp32_layers: - raise NotImplementedError("DINOv3 does not support mixed-precision layers") - if request.tensor_parallel_size != 1: - raise NotImplementedError("DINOv3 does not support tensor parallelism") model = _Dinov3Model() weights = model.load_weights(str(model_dir), config, precision=precision) writer.set_header(family="dinov3", task=request.task, backend=request.backend) plan = model.build_engine( config, weights, - max_sequence_length, + 1, precision=precision, quant_ctx=None, verbose=bool(request.verbose), diff --git a/families/dinov3/runtime/CMakeLists.txt b/families/dinov3/runtime/CMakeLists.txt index e0a486a13d..c496c1d4da 100644 --- a/families/dinov3/runtime/CMakeLists.txt +++ b/families/dinov3/runtime/CMakeLists.txt @@ -65,3 +65,26 @@ if(TRTMC_BUILD_TESTS) target_compile_options(test_dinov3_pipeline PRIVATE -Wall -Wextra -Wpedantic) add_test(NAME test_dinov3_pipeline COMMAND test_dinov3_pipeline) endif() + +# The CLI adapter owns image I/O and uses the existing family runtime contract. +add_library(trtmc_cli_dinov3 SHARED cli.cpp) +target_include_directories(trtmc_cli_dinov3 PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_dinov3 SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_dinov3 PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_dinov3 PRIVATE -Wall -Wextra -Wpedantic) +set_target_properties(trtmc_cli_dinov3 PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_dinov3 LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_executable(test_dinov3_cli ${PROJECT_SOURCE_DIR}/families/dinov3/tests/cpp/test_cli.cpp cli.cpp) + target_include_directories(test_dinov3_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_include_directories(test_dinov3_cli SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) + target_link_libraries(test_dinov3_cli PRIVATE trtmc_core nlohmann_json::nlohmann_json) + target_compile_options(test_dinov3_cli PRIVATE -Wall -Wextra -Wpedantic) + add_test(NAME dinov3_cli COMMAND test_dinov3_cli) + set_tests_properties(dinov3_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/dinov3/runtime/cli.cpp b/families/dinov3/runtime/cli.cpp new file mode 100644 index 0000000000..5bf41184d5 --- /dev/null +++ b/families/dinov3/runtime/cli.cpp @@ -0,0 +1,92 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include +#include +#include +#include +#include +#include +#include + +namespace { +struct Image { + std::vector pixels; + int width = 0, height = 0; +}; + +Image read_image(const std::string& path) { + Image result; + int channels = 0; + std::unique_ptr bytes( + stbi_load(path.c_str(), &result.width, &result.height, &channels, 3), stbi_image_free); + if (!bytes || result.width <= 0 || result.height <= 0) + throw std::runtime_error("unable to decode image: " + path); + const auto count = static_cast(result.width) * result.height * 3U; + result.pixels.resize(count); + for (std::size_t i = 0; i < count; ++i) + result.pixels[i] = static_cast(bytes.get()[i]) / 255.0F; + return result; +} + +void require_finite(const std::vector& values) { + for (const auto value : values) { + if (!std::isfinite(value)) + throw std::runtime_error("dinov3 returned a non-finite result"); + } +} + +nlohmann::json execute(const std::string& handler, const nlohmann::json& values, + const char* default_runtime_root) { + if (handler != "extract_features") + throw std::invalid_argument("unknown dinov3 CLI handler: " + handler); + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "dinov3") + throw std::invalid_argument("dinov3 CLI requires a dinov3 bundle"); + const auto image = read_image(values.at("image").get()); + auto task = trtmc::load_task( + reader, values.value("runtime_root", std::string(default_runtime_root)), 0, + values.value("runtime_cache", std::string()), values.value("cuda_graphs", false)); + auto* model = dynamic_cast(task.get()); + if (!model) + throw std::invalid_argument("dinov3 bundle does not implement IImageFeatureExtractor"); + const auto result = + model->extract_image_features(image.pixels.data(), image.height, image.width); + require_finite(result.last_hidden_state); + require_finite(result.pooler_output); + return {{"last_hidden_state", result.last_hidden_state}, + {"last_hidden_state_shape", result.last_hidden_state_shape}, + {"pooler_output", result.pooler_output}, + {"pooler_output_shape", result.pooler_output_shape}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + const auto result = + execute(handler, nlohmann::json::parse(values_json), default_runtime_root).dump() + + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } catch (...) { + const std::string message = "Error: dinov3 CLI failed\n"; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/dinov3/tests/cpp/test_cli.cpp b/families/dinov3/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..3bae60aa9f --- /dev/null +++ b/families/dinov3/tests/cpp/test_cli.cpp @@ -0,0 +1,141 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0, task_loads = 0; +bool bad_output = false, loaded_graphs = false; +std::string loaded_runtime, loaded_cache; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +class FakeModel final : public trtmc::IImageFeatureExtractor { + public: + trtmc::ImageFeaturesResult extract_image_features(const float* pixels, std::int32_t height, + std::int32_t width) override { + check(height == 1 && width == 2, "owner passes image dimensions"); + check(pixels[0] == 1.0F && pixels[1] == 128.0F / 255.0F && pixels[2] == 0.0F, + "owner passes normalized HWC RGB"); + return {{pixels, pixels + 6}, + {1, 2, 3}, + {bad_output ? std::numeric_limits::infinity() : pixels[0], pixels[5]}, + {1, 2}}; + } +}; +void write_bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "image_features"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream out(path, std::ios::binary); + out.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + out.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + out.write(header.data(), static_cast(header.size())); + out.write("PLAN", 4); +} +struct Capture { + int status; + std::string output, error; +}; +void output(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void error(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +Capture invoke(const char* handler, const Json& values) { + Capture result{}; + result.status = trtmc_family_cli_v1(handler, values.dump().c_str(), "/installed/runtime", + &result, output, error); + return result; +} +void contract(const fs::path& root) { + const auto bundle = root / "model.bundle", image = root / "input.ppm"; + write_bundle(bundle, "dinov3"); + { + std::ofstream out(image, std::ios::binary); + out << "P6\n2 1\n255\n"; + const std::array rgb{255, 128, 0, 0, 255, 64}; + out.write(reinterpret_cast(rgb.data()), rgb.size()); + } + Json values{{"bundle", bundle.string()}, {"image", image.string()}}; + const auto first = invoke("extract_features", values); + check(first.status == 0 && first.error.empty(), "valid owner request succeeds"); + const auto output = Json::parse(first.output); + check(output.at("last_hidden_state_shape") == Json({1, 2, 3}), "feature shape is preserved"); + check(output.at("pooler_output_shape") == Json({1, 2}), "pooler shape is preserved"); + check(output.at("pooler_output").at(1).get() == 64.0F / 255.0F, + "feature output preserves normalized RGB"); + + check(loaded_runtime == "/installed/runtime" && loaded_cache.empty() && !loaded_graphs, + "default load options are preserved"); + values.update({{"runtime_root", "/override/runtime"}, + {"runtime_cache", "runtime.cache"}, + {"cuda_graphs", true}}); + check(invoke("extract_features", values).status == 0, "explicit runtime options succeed"); + check(loaded_runtime == "/override/runtime" && loaded_cache == "runtime.cache" && loaded_graphs, + "runtime options reach the loader"); + bad_output = true; + const auto invalid_output = invoke("extract_features", values); + check(invalid_output.status != 0 && invalid_output.output.empty(), + "non-finite runtime results fail closed"); + bad_output = false; + const auto before = task_loads; + values["image"] = (root / "missing.png").string(); + check(invoke("extract_features", values).status != 0 && task_loads == before, + "invalid image fails before model loading"); + values["image"] = image.string(); + write_bundle(bundle, "another_family"); + check(invoke("extract_features", values).status != 0 && task_loads == before, + "wrong-family bundle fails before model loading"); + check(invoke("unknown", values).status != 0 && task_loads == before, + "unknown handler never loads a task"); +} +} // namespace + +namespace trtmc { +std::unique_ptr load_task(const BundleReader&, const std::string& runtime_root, + std::uint64_t, const std::string& runtime_cache, + bool cuda_graphs) { + ++task_loads; + loaded_runtime = runtime_root; + loaded_cache = runtime_cache; + loaded_graphs = cuda_graphs; + return std::make_unique(); +} +} // namespace trtmc +int main() { + const auto root = fs::temp_directory_path() / ("trtmc-dinov3-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(root); + } catch (const std::exception& exception) { + std::cerr << exception.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures == 0 ? 0 : 1; +} diff --git a/families/dinov3/tests/manifests/dinov3-convnext-tiny-pretrain-lvd1689m.json b/families/dinov3/tests/manifests/dinov3-convnext-tiny-pretrain-lvd1689m.json index f562053634..49684ffce9 100644 --- a/families/dinov3/tests/manifests/dinov3-convnext-tiny-pretrain-lvd1689m.json +++ b/families/dinov3/tests/manifests/dinov3-convnext-tiny-pretrain-lvd1689m.json @@ -13,6 +13,5 @@ "num_register_tokens": 0 } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/dinov3/tests/manifests/dinov3-vits16-pretrain-lvd1689m.json b/families/dinov3/tests/manifests/dinov3-vits16-pretrain-lvd1689m.json index c41b4be3f7..7a76540e35 100644 --- a/families/dinov3/tests/manifests/dinov3-vits16-pretrain-lvd1689m.json +++ b/families/dinov3/tests/manifests/dinov3-vits16-pretrain-lvd1689m.json @@ -13,6 +13,5 @@ "num_register_tokens": 4 } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/dinov3/tests/manifests/dinov3-vits16-timm-l0.json b/families/dinov3/tests/manifests/dinov3-vits16-timm-l0.json index f897ddde39..ce25f98c29 100644 --- a/families/dinov3/tests/manifests/dinov3-vits16-timm-l0.json +++ b/families/dinov3/tests/manifests/dinov3-vits16-timm-l0.json @@ -14,6 +14,5 @@ "num_register_tokens": 4 } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/dinov3/tests/test_cli.py b/families/dinov3/tests/test_cli.py new file mode 100644 index 0000000000..013633d37b --- /dev/null +++ b/families/dinov3/tests/test_cli.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import replace +import json +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest + +from families.dinov3 import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli +from trtmc_benchmark.catalog import ManifestCatalog +from trtmc_benchmark.types import BenchmarkError + + +def test_declared_build_uses_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr(cli, "resolve_model", lambda model, revision: tmp_path) + monkeypatch.setattr(cli, "build_bundle", lambda request, output: calls.append((request, output))) + output = tmp_path / "out.bundle" + assert family_cli.main(["dinov3", "build", "checkpoint", "-o", str(output)]) == 0 + legacy = LegacyRequest(tmp_path, output, "dinov3", "image_features", "fp32") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "dynamic_kv_cache") + assert not hasattr(calls[0][0], "tensor_parallel_size") + assert not hasattr(calls[0][0], "output_path") + + +@pytest.mark.parametrize("changes", [ + {"dynamic_kv_cache": True}, {"video_num_frames": 2}, {"max_batch_size": 2}, + {"tensor_parallel_size": 2}, {"context_parallel_size": 2}, + {"quantization": "fp8"}, {"fp32_layers": (1,)}, +]) +def test_legacy_rejects_unsupported_nondefaults(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", "dinov3", "image_features", "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_legacy_fields_are_not_dropped(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", "dinov3", "image_features", "fp32") + with pytest.raises(ValueError, match="unknown dinov3 build inputs"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unrecognized_option=1)) + + +def test_help_and_rejected_options_do_not_import_handlers(monkeypatch, capsys): + discover_namespace = family_cli.importlib.import_module + def reject_import(name, *args, **kwargs): + if name == "families": + return discover_namespace(name, *args, **kwargs) + raise AssertionError(f"unexpected lazy import: {name}") + monkeypatch.setattr(family_cli.importlib, "import_module", reject_import) + for command in ("build", "extract-features"): + with pytest.raises(SystemExit) as result: + family_cli.main(["dinov3", command, "--help"]) + assert result.value.code == 0 + assert "--runtime-root" in capsys.readouterr().out + with pytest.raises(SystemExit) as result: + family_cli.main(["dinov3", "build", "checkpoint", "-o", "out.bundle", "--tensor-parallel-size", "2"]) + assert result.value.code == 2 + + +def test_bundle_selects_backend_before_import_and_publishes_atomically(monkeypatch, tmp_path): + events = [] + output = tmp_path / "model.bundle" + def build(request, writer): + events.append("build") + writer.set_header(family="dinov3", task="image_features", backend=request.backend) + writer.add_bytes("engine.plan", b"fixture") + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, "families.dinov3.model", SimpleNamespace(build=build)) + monkeypatch.setattr(cli, "select_backend", select) + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), output) + assert events == ["trt_rtx", "build"] + published = output.read_bytes() + assert published.startswith(b"BUNDLE") + def fail(request, writer): + writer.set_header(family="dinov3", task="image_features", backend=request.backend) + writer.add_bytes("engine.plan", b"partial") + raise RuntimeError("owner build failed") + monkeypatch.setattr(cli, "select_backend", lambda backend: monkeypatch.setitem( + sys.modules, "families.dinov3.model", SimpleNamespace(build=fail))) + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path), output) + assert output.read_bytes() == published + assert list(tmp_path.iterdir()) == [output] + + +def test_real_manifests_serialize_through_the_owner_handler(monkeypatch, tmp_path): + captured = [] + monkeypatch.setattr(cli, "resolve_model", lambda model, revision: tmp_path) + monkeypatch.setattr(cli, "build_bundle", lambda request, output: captured.append((request, output))) + spec = family_cli.load_family_cli("dinov3")["commands"][0] + manifests = sorted((Path(__file__).parent / "manifests").glob("*.json")) + assert manifests + for path in manifests: + model = ManifestCatalog._load(path) + values = dict(model.build_settings) + values.update(model=model.hf_id, output=str(tmp_path / model.bundle_name), task=model.task, precision=model.precision) + argv = family_cli.serialize_arguments(spec, values) + assert family_cli.main(["dinov3", "build", *argv]) == 0 + request, output = captured[-1] + assert type(request) is cli.BuildRequest and request.precision == model.precision + assert request.task == model.task and output == tmp_path / model.bundle_name + raw = json.loads(manifests[0].read_text()) + raw["build"] = {"unknown_owner_option": 1} + invalid = tmp_path / "invalid.json" + invalid.write_text(json.dumps(raw)) + with pytest.raises(BenchmarkError, match="undeclared build fields"): + ManifestCatalog._load(invalid) + + +def test_legacy_sequence_limit_remains_accepted_but_unexposed(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", "dinov3", "image_features", "fp32", max_sequence_length=512) + assert cli.coerce_request(legacy) == cli.BuildRequest(tmp_path) + for field in ("image_height", "image_width"): + with pytest.raises(NotImplementedError): + cli.coerce_request(replace(legacy, **{field: 32})) + + +def test_e2e_build_preserves_single_device_rejection(monkeypatch, tmp_path): + from families.dinov3.tests import test_e2e + + calls = [] + monkeypatch.setattr(test_e2e, "build_bundle", lambda *args: calls.append(args)) + manifest = json.loads(next((Path(__file__).parent / "manifests").glob("*.json")).read_text()) + manifest["tensor_parallel_size"] = 2 + with pytest.raises(NotImplementedError, match="tensor parallelism"): + test_e2e._build(tmp_path, tmp_path / "out.bundle", manifest) + assert not calls diff --git a/families/dinov3/tests/test_e2e.py b/families/dinov3/tests/test_e2e.py index 93d5982ef6..2fd274dba9 100644 --- a/families/dinov3/tests/test_e2e.py +++ b/families/dinov3/tests/test_e2e.py @@ -14,7 +14,7 @@ from pathlib import Path import pytest import numpy as np -from tensorrt_model_connect import BuildRequest, build +from families.dinov3.cli import BuildRequest, build_bundle FAMILY = "dinov3" TASKS = frozenset({"image_features"}) @@ -117,6 +117,8 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise FileNotFoundError(f"selected {FAMILY} E2E requires its native CLI adapter") import torch required_gpus = int(manifest["tensor_parallel_size"]) @@ -128,22 +130,15 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError("dinov3 does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=manifest.get("max_sequence_length"), - image_height=manifest.get("image_height"), - image_width=manifest.get("image_width"), - video_num_frames=manifest.get("video_num_frames"), - max_batch_size=int(manifest.get("max_batch_size", 1)), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - quantization=manifest.get("quantization"), - fp32_layers=tuple((int(layer) for layer in manifest.get("fp32_layers", ()))), - ) + ), + bundle, ) @@ -158,6 +153,7 @@ def _run_json( ) -> dict: invocation = [ str(binary), + FAMILY, command, str(bundle), "--runtime-root", diff --git a/families/timm_resnet/cli.json b/families/timm_resnet/cli.json new file mode 100644 index 0000000000..5aad3de104 --- /dev/null +++ b/families/timm_resnet/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_resnet bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "image_to_class_scores" + ], + "default": "image_to_class_scores" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_resnet bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_resnet/cli.py b/families/timm_resnet/cli.py new file mode 100644 index 0000000000..4a626264f3 --- /dev/null +++ b/families/timm_resnet/cli.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_resnet commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "image_to_class_scores" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "image_to_class_scores": + raise ValueError("timm_resnet supports only task=image_to_class_scores") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_resnet supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_resnet does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_resnet does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_resnet does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_resnet does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_resnet does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_resnet does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_resnet does not support context parallelism") + if request.task != "image_to_class_scores": + raise ValueError("timm_resnet supports only task=image_to_class_scores") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_resnet does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_resnet does not support mixed-precision layers") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_resnet build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "image_to_class_scores", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_resnet/model.py b/families/timm_resnet/model.py index a7ca23eaf9..42987ff6af 100644 --- a/families/timm_resnet/model.py +++ b/families/timm_resnet/model.py @@ -36,8 +36,10 @@ from .config import ModelConfig +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter # timm's ResNet BatchNorm2d layers use the PyTorch default epsilon; it is not @@ -375,27 +377,7 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm ResNet image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_resnet does not support dynamic_kv_cache") - - if request.image_height is not None: - raise NotImplementedError("timm_resnet does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_resnet does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_resnet does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_resnet does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_resnet does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_resnet does not support context parallelism") - if request.task != "image_to_class_scores": - raise ValueError("timm_resnet supports only task=image_to_class_scores") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_resnet does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_resnet does not support mixed-precision layers") + request = coerce_request(request) model_dir = Path(request.model_dir) config = ModelConfig.from_dir(model_dir) @@ -403,7 +385,7 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: if not model_type.startswith(("resnet", "resnext", "wide_resnet")): raise ValueError(f"timm ResNet does not support model_type={config.model_type!r}") precision = str(request.precision).lower() - max_sequence_length = _positive_int(request.max_sequence_length or 1, "max_sequence_length") + max_sequence_length = 1 model = _TimmResnetModel() weights = model.load_weights(str(model_dir), config, precision=precision) plan = model.build_engine( diff --git a/families/timm_resnet/runtime/CMakeLists.txt b/families/timm_resnet/runtime/CMakeLists.txt index 255731af27..15ecf2ed7e 100644 --- a/families/timm_resnet/runtime/CMakeLists.txt +++ b/families/timm_resnet/runtime/CMakeLists.txt @@ -93,3 +93,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_resnet_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_resnet SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_resnet PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_resnet SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_resnet PRIVATE trtmc_c trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_resnet PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_resnet PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_resnet LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_resnet_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_resnet/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_resnet_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_resnet_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_resnet_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_resnet_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_resnet + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_resnet-cli" + ) + add_dependencies(timm_resnet_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_resnet_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_resnet_cli ${PROJECT_SOURCE_DIR}/families/timm_resnet/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_resnet_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_resnet_cli PRIVATE trtmc_cli_timm_resnet nlohmann_json::nlohmann_json) + target_compile_options(test_timm_resnet_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_resnet_cli timm_resnet_cli_fixture) + add_test(NAME timm_resnet_cli COMMAND test_timm_resnet_cli $) + set_tests_properties(timm_resnet_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_resnet/runtime/cli.cpp b/families/timm_resnet/runtime/cli.cpp new file mode 100644 index 0000000000..82eff6a3ae --- /dev/null +++ b/families/timm_resnet/runtime/cli.cpp @@ -0,0 +1,110 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/core.hpp" +#include "trtmc/features.hpp" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_resnet") + throw std::invalid_argument("timm_resnet CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + const auto model = + trtmc::Model::load(reader.path(), {runtime_root, 0, runtime_cache, cuda_graphs}); + const auto task = model.task(); + const auto result = task.run( + {trtmc::ImageInput{trtmc::Span{pixels.data(), pixels.size()}, + static_cast(height), static_cast(width)}}, + {}); + std::vector scores(result.scores().begin(), result.scores().end()); + for (const auto value : scores) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned a non-finite score"); + } + const char* kind = nullptr; + switch (result.kind()) { + case TRTMC_SCORE_LOGIT: + kind = "logit"; + break; + case TRTMC_SCORE_PROBABILITY: + kind = "probability"; + break; + case TRTMC_SCORE_UNBOUNDED: + kind = "unbounded"; + break; + default: + throw std::runtime_error("classification returned an unknown score kind"); + } + auto labels = Json::array(); + for (const auto label : result.labels()) + labels.push_back(std::string(label)); + Json output{{"scores", scores}, + {"score_kind", kind}, + {"labels", labels}, + {"vocabulary_id", std::string(result.vocabulary_id())}, + {"task", trtmc::ImageToClassScores::kTask}}; + if (result.kind() == TRTMC_SCORE_LOGIT) + output["logits"] = scores; + if (scores.empty()) { + output["top_class"] = -1; + output["top_score"] = nullptr; + } else { + const auto best = std::max_element(scores.begin(), scores.end()); + output["top_class"] = best - scores.begin(); + output["top_score"] = *best; + } + return output; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_resnet CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_resnet/tests/cpp/test_cli.cpp b/families/timm_resnet/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..f14022bfc5 --- /dev/null +++ b/families/timm_resnet/tests/cpp/test_cli.cpp @@ -0,0 +1,155 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/internal/features.h" +#include "trtmc/internal/model.h" +#include "trtmc/runtime/family_factory.h" +namespace { +class Fixture final : public trtmc::internal::IModel, public trtmc::internal::IImageToClassScores { + public: + const char* task() const noexcept override { return "image_to_class_scores"; } + std::vector task_bindings() override { + return {trtmc::internal::bind(*this)}; + } + trtmc::internal::LabelScoresResult + run(const trtmc::internal::ImageToClassScoresRequest& request, + trtmc::internal::ConfigView config) override { + const auto& image = request.image; + if (!config.empty() || image.width != 2 || image.height != 1 || image.channels != 3 || + image.format != trtmc::internal::ImageFormat::Float32 || image.byte_size != 24) + throw std::invalid_argument("fixture expects one decoded 2x1 RGB float32 image"); + const auto* pixels = static_cast(image.data); + if (pixels[0] != 1.0F || pixels[1] != 0.0F || pixels[2] != 128.0F / 255.0F || + pixels[4] != 1.0F) + throw std::invalid_argument("owner changed pixel order or range"); + return {{-2.0F, 4.0F, 0.5F}, + {"first", "second", "third"}, + trtmc::internal::ScoreKind::Logit, + "fixture:classes"}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "image_to_class_scores"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_resnet"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.at("scores") == actual.at("logits") && actual.at("score_kind") == "logit" && + actual.at("labels") == Json({"first", "second", "third"}) && + actual.at("vocabulary_id") == "fixture:classes" && + actual.at("task") == "image_to_class_scores", + "SDK class identity and score semantics are preserved"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_resnet"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_resnet-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_resnet/tests/manifests/resnet50-a1-in1k.json b/families/timm_resnet/tests/manifests/resnet50-a1-in1k.json index 900e02d962..a5127854b5 100644 --- a/families/timm_resnet/tests/manifests/resnet50-a1-in1k.json +++ b/families/timm_resnet/tests/manifests/resnet50-a1-in1k.json @@ -12,6 +12,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_resnet/tests/test_cli.py b/families/timm_resnet/tests/test_cli.py new file mode 100644 index 0000000000..6c3885b9f3 --- /dev/null +++ b/families/timm_resnet/tests/test_cli.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_resnet command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_resnet import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_resnet" +TASK = "image_to_class_scores" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_resnet/tests/test_e2e.py b/families/timm_resnet/tests/test_e2e.py index f504f0fc37..2b43144671 100644 --- a/families/timm_resnet/tests/test_e2e.py +++ b/families/timm_resnet/tests/test_e2e.py @@ -14,7 +14,7 @@ from pathlib import Path import pytest import numpy as np -from tensorrt_model_connect import BuildRequest, build +from families.timm_resnet.cli import BuildRequest, build_bundle FAMILY = "timm_resnet" TASKS = frozenset({"image_to_class_scores"}) @@ -116,6 +116,8 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") import torch required_gpus = int(manifest["tensor_parallel_size"]) @@ -127,22 +129,15 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=manifest.get("max_sequence_length"), - image_height=manifest.get("image_height"), - image_width=manifest.get("image_width"), - video_num_frames=manifest.get("video_num_frames"), - max_batch_size=int(manifest.get("max_batch_size", 1)), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - quantization=manifest.get("quantization"), - fp32_layers=tuple((int(layer) for layer in manifest.get("fp32_layers", ()))), - ) + ), + bundle, ) @@ -157,6 +152,7 @@ def _run_json( ) -> dict: invocation = [ str(binary), + FAMILY, command, str(bundle), "--runtime-root", diff --git a/families/timm_senet/cli.json b/families/timm_senet/cli.json new file mode 100644 index 0000000000..27fc9e93cc --- /dev/null +++ b/families/timm_senet/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_senet bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "classification" + ], + "default": "classification" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_senet bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_senet/cli.py b/families/timm_senet/cli.py new file mode 100644 index 0000000000..a77122286c --- /dev/null +++ b/families/timm_senet/cli.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_senet commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "classification" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "classification": + raise ValueError("timm_senet supports only task=classification") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_senet supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_senet does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_senet does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_senet does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_senet does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_senet does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_senet does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_senet does not support context parallelism") + if request.task != "classification": + raise ValueError("timm_senet supports only task=classification") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_senet does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_senet does not support mixed-precision layers") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_senet build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "classification", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_senet/model.py b/families/timm_senet/model.py index 54cef0bac6..17a4045979 100644 --- a/families/timm_senet/model.py +++ b/families/timm_senet/model.py @@ -28,8 +28,10 @@ from .checkpoint import Checkpoint +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -341,27 +343,8 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm SENet image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_senet does not support dynamic_kv_cache") - if request.image_height is not None: - raise NotImplementedError("timm_senet does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_senet does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_senet does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_senet does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_senet does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_senet does not support context parallelism") - if request.task != "classification": - raise ValueError("timm_senet supports only task=classification") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_senet does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_senet does not support mixed-precision layers") - _positive_int(request.max_sequence_length or 1, "max_sequence_length") + request = coerce_request(request) + model_dir = Path(request.model_dir) raw = _read_config(model_dir) plan, runtime = _build_engine( diff --git a/families/timm_senet/runtime/CMakeLists.txt b/families/timm_senet/runtime/CMakeLists.txt index e8ce9a840b..5d9766d279 100644 --- a/families/timm_senet/runtime/CMakeLists.txt +++ b/families/timm_senet/runtime/CMakeLists.txt @@ -55,3 +55,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_senet_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_senet SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_senet PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_senet SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_senet PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_senet PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_senet PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_senet LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_senet_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_senet/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_senet_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_senet_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_senet_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_senet_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_senet + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_senet-cli" + ) + add_dependencies(timm_senet_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_senet_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_senet_cli ${PROJECT_SOURCE_DIR}/families/timm_senet/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_senet_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_senet_cli PRIVATE trtmc_cli_timm_senet nlohmann_json::nlohmann_json) + target_compile_options(test_timm_senet_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_senet_cli timm_senet_cli_fixture) + add_test(NAME timm_senet_cli COMMAND test_timm_senet_cli $) + set_tests_properties(timm_senet_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_senet/runtime/cli.cpp b/families/timm_senet/runtime/cli.cpp new file mode 100644 index 0000000000..67d80381b2 --- /dev/null +++ b/families/timm_senet/runtime/cli.cpp @@ -0,0 +1,79 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_senet") + throw std::invalid_argument("timm_senet CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + auto task = trtmc::load_task(reader, runtime_root, 0, runtime_cache, cuda_graphs); + auto* classifier = dynamic_cast(task.get()); + if (!classifier) + throw std::invalid_argument("bundle does not implement image classification"); + const auto result = classifier->classify(pixels.data(), height, width); + for (const auto value : result.logits) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned non-finite logits"); + } + if (!std::isfinite(result.top_score)) + throw std::runtime_error("classification returned a non-finite top score"); + return {{"logits", result.logits}, + {"top_class", result.top_class}, + {"top_score", result.top_score}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_senet CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_senet/tests/cpp/test_cli.cpp b/families/timm_senet/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..94c1481274 --- /dev/null +++ b/families/timm_senet/tests/cpp/test_cli.cpp @@ -0,0 +1,138 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/runtime/family_factory.h" +#include "trtmc/task.h" +namespace { +class Fixture final : public trtmc::IImageClassification { + public: + const char* task() const noexcept override { return "classification"; } + trtmc::ClassificationResult classify(const float* pixels, std::int32_t height, + std::int32_t width) override { + if (width != 2 || height != 1 || pixels[0] != 1.0F || pixels[1] != 0.0F || + pixels[2] != 128.0F / 255.0F || pixels[4] != 1.0F) + throw std::invalid_argument("owner changed decoded image shape, order or range"); + return {{-2.0F, 4.0F, 0.5F}, 1, 4.0F}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "classification"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_senet"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.size() == 3, "legacy classification output shape remains unchanged"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_senet"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_senet-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_senet/tests/manifests/senet154-gluon-in1k.json b/families/timm_senet/tests/manifests/senet154-gluon-in1k.json index ddb9d56f64..b3398b1a81 100644 --- a/families/timm_senet/tests/manifests/senet154-gluon-in1k.json +++ b/families/timm_senet/tests/manifests/senet154-gluon-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_senet/tests/test_cli.py b/families/timm_senet/tests/test_cli.py new file mode 100644 index 0000000000..4856b0326b --- /dev/null +++ b/families/timm_senet/tests/test_cli.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_senet command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_senet import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_senet" +TASK = "classification" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_senet/tests/test_e2e.py b/families/timm_senet/tests/test_e2e.py index bd97e3ebfa..b3bcbe57fe 100644 --- a/families/timm_senet/tests/test_e2e.py +++ b/families/timm_senet/tests/test_e2e.py @@ -15,7 +15,7 @@ import numpy as np import pytest -from tensorrt_model_connect import BuildRequest, build +from families.timm_senet.cli import BuildRequest, build_bundle FAMILY = "timm_senet" @@ -115,25 +115,27 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: assert (runtime_root / "libtrtmc_backend_trt.so").is_file() with evidence_stage("setup"): assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") model_dir = _model_dir(manifest) record_evidence("checkpoint", {"model_dir": str(model_dir), "hf_id": manifest.get("hf_id"), "hf_revision": manifest.get("hf_revision")}) bundle = tmp_path / manifest["bundle"] with evidence_stage("build"): - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=int(manifest["max_sequence_length"]), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - ) + ), + bundle, ) with evidence_stage("native"): completed = subprocess.run( [ str(binary), + FAMILY, "classify", str(bundle), "--runtime-root", diff --git a/families/timm_seresnet/cli.json b/families/timm_seresnet/cli.json new file mode 100644 index 0000000000..2b96a0a76d --- /dev/null +++ b/families/timm_seresnet/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_seresnet bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "classification" + ], + "default": "classification" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_seresnet bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_seresnet/cli.py b/families/timm_seresnet/cli.py new file mode 100644 index 0000000000..20364751f9 --- /dev/null +++ b/families/timm_seresnet/cli.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_seresnet commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "classification" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "classification": + raise ValueError("timm_seresnet supports only task=classification") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_seresnet supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_seresnet does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_seresnet does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_seresnet does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_seresnet does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_seresnet does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_seresnet does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_seresnet does not support context parallelism") + if request.task != "classification": + raise ValueError("timm_seresnet supports only task=classification") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_seresnet does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_seresnet does not support mixed-precision layers") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_seresnet build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "classification", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_seresnet/model.py b/families/timm_seresnet/model.py index 601a29b18f..050101b8eb 100644 --- a/families/timm_seresnet/model.py +++ b/families/timm_seresnet/model.py @@ -28,8 +28,10 @@ from .checkpoint import Checkpoint +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -322,27 +324,8 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm SE-ResNet image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_seresnet does not support dynamic_kv_cache") - if request.image_height is not None: - raise NotImplementedError("timm_seresnet does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_seresnet does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_seresnet does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_seresnet does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_seresnet does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_seresnet does not support context parallelism") - if request.task != "classification": - raise ValueError("timm_seresnet supports only task=classification") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_seresnet does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_seresnet does not support mixed-precision layers") - _positive_int(request.max_sequence_length or 1, "max_sequence_length") + request = coerce_request(request) + model_dir = Path(request.model_dir) raw = _read_config(model_dir) plan, runtime = _build_engine( diff --git a/families/timm_seresnet/runtime/CMakeLists.txt b/families/timm_seresnet/runtime/CMakeLists.txt index eaa905494c..8503062660 100644 --- a/families/timm_seresnet/runtime/CMakeLists.txt +++ b/families/timm_seresnet/runtime/CMakeLists.txt @@ -55,3 +55,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_seresnet_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_seresnet SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_seresnet PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_seresnet SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_seresnet PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_seresnet PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_seresnet PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_seresnet LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_seresnet_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_seresnet/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_seresnet_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_seresnet_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_seresnet_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_seresnet_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_seresnet + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_seresnet-cli" + ) + add_dependencies(timm_seresnet_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_seresnet_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_seresnet_cli ${PROJECT_SOURCE_DIR}/families/timm_seresnet/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_seresnet_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_seresnet_cli PRIVATE trtmc_cli_timm_seresnet nlohmann_json::nlohmann_json) + target_compile_options(test_timm_seresnet_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_seresnet_cli timm_seresnet_cli_fixture) + add_test(NAME timm_seresnet_cli COMMAND test_timm_seresnet_cli $) + set_tests_properties(timm_seresnet_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_seresnet/runtime/cli.cpp b/families/timm_seresnet/runtime/cli.cpp new file mode 100644 index 0000000000..cb4fa5e3d2 --- /dev/null +++ b/families/timm_seresnet/runtime/cli.cpp @@ -0,0 +1,79 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_seresnet") + throw std::invalid_argument("timm_seresnet CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + auto task = trtmc::load_task(reader, runtime_root, 0, runtime_cache, cuda_graphs); + auto* classifier = dynamic_cast(task.get()); + if (!classifier) + throw std::invalid_argument("bundle does not implement image classification"); + const auto result = classifier->classify(pixels.data(), height, width); + for (const auto value : result.logits) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned non-finite logits"); + } + if (!std::isfinite(result.top_score)) + throw std::runtime_error("classification returned a non-finite top score"); + return {{"logits", result.logits}, + {"top_class", result.top_class}, + {"top_score", result.top_score}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_seresnet CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_seresnet/tests/cpp/test_cli.cpp b/families/timm_seresnet/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..b1700cc5b7 --- /dev/null +++ b/families/timm_seresnet/tests/cpp/test_cli.cpp @@ -0,0 +1,138 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/runtime/family_factory.h" +#include "trtmc/task.h" +namespace { +class Fixture final : public trtmc::IImageClassification { + public: + const char* task() const noexcept override { return "classification"; } + trtmc::ClassificationResult classify(const float* pixels, std::int32_t height, + std::int32_t width) override { + if (width != 2 || height != 1 || pixels[0] != 1.0F || pixels[1] != 0.0F || + pixels[2] != 128.0F / 255.0F || pixels[4] != 1.0F) + throw std::invalid_argument("owner changed decoded image shape, order or range"); + return {{-2.0F, 4.0F, 0.5F}, 1, 4.0F}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "classification"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_seresnet"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.size() == 3, "legacy classification output shape remains unchanged"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_seresnet"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_seresnet-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_seresnet/tests/manifests/seresnet50-a1-in1k.json b/families/timm_seresnet/tests/manifests/seresnet50-a1-in1k.json index 352981ea5e..3207a71745 100644 --- a/families/timm_seresnet/tests/manifests/seresnet50-a1-in1k.json +++ b/families/timm_seresnet/tests/manifests/seresnet50-a1-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_seresnet/tests/test_cli.py b/families/timm_seresnet/tests/test_cli.py new file mode 100644 index 0000000000..991a81aab6 --- /dev/null +++ b/families/timm_seresnet/tests/test_cli.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_seresnet command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_seresnet import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_seresnet" +TASK = "classification" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_seresnet/tests/test_e2e.py b/families/timm_seresnet/tests/test_e2e.py index e0c657b53a..9683542392 100644 --- a/families/timm_seresnet/tests/test_e2e.py +++ b/families/timm_seresnet/tests/test_e2e.py @@ -15,7 +15,7 @@ import numpy as np import pytest -from tensorrt_model_connect import BuildRequest, build +from families.timm_seresnet.cli import BuildRequest, build_bundle FAMILY = "timm_seresnet" @@ -115,25 +115,27 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: assert (runtime_root / "libtrtmc_backend_trt.so").is_file() with evidence_stage("setup"): assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") model_dir = _model_dir(manifest) record_evidence("checkpoint", {"model_dir": str(model_dir), "hf_id": manifest.get("hf_id"), "hf_revision": manifest.get("hf_revision")}) bundle = tmp_path / manifest["bundle"] with evidence_stage("build"): - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=int(manifest["max_sequence_length"]), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - ) + ), + bundle, ) with evidence_stage("native"): completed = subprocess.run( [ str(binary), + FAMILY, "classify", str(bundle), "--runtime-root", diff --git a/families/timm_swin/cli.json b/families/timm_swin/cli.json new file mode 100644 index 0000000000..615cd69842 --- /dev/null +++ b/families/timm_swin/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_swin bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "image_to_class_scores" + ], + "default": "image_to_class_scores" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_swin bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_swin/cli.py b/families/timm_swin/cli.py new file mode 100644 index 0000000000..b2417ba180 --- /dev/null +++ b/families/timm_swin/cli.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_swin commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "image_to_class_scores" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "image_to_class_scores": + raise ValueError("timm_swin supports only task=image_to_class_scores") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_swin supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_swin does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_swin does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_swin does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_swin does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_swin does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_swin does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_swin does not support context parallelism") + if request.task != "image_to_class_scores": + raise ValueError("timm_swin supports only task=image_to_class_scores") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_swin does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_swin does not support mixed-precision layers") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_swin build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "image_to_class_scores", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_swin/model.py b/families/timm_swin/model.py index a902205405..38ee8fdb64 100644 --- a/families/timm_swin/model.py +++ b/families/timm_swin/model.py @@ -28,8 +28,10 @@ from .checkpoint import Checkpoint +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -402,27 +404,8 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm Swin Transformer image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_swin does not support dynamic_kv_cache") - if request.image_height is not None: - raise NotImplementedError("timm_swin does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_swin does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_swin does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_swin does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_swin does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_swin does not support context parallelism") - if request.task != "image_to_class_scores": - raise ValueError("timm_swin supports only task=image_to_class_scores") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_swin does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_swin does not support mixed-precision layers") - _positive_int(request.max_sequence_length or 1, "max_sequence_length") + request = coerce_request(request) + model_dir = Path(request.model_dir) raw = _read_config(model_dir) plan, runtime = _build_engine( diff --git a/families/timm_swin/runtime/CMakeLists.txt b/families/timm_swin/runtime/CMakeLists.txt index 8ec4833c3a..fe5a0ef2fc 100644 --- a/families/timm_swin/runtime/CMakeLists.txt +++ b/families/timm_swin/runtime/CMakeLists.txt @@ -90,3 +90,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_swin_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_swin SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_swin PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_swin SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_swin PRIVATE trtmc_c trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_swin PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_swin PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_swin LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_swin_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_swin/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_swin_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_swin_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_swin_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_swin_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_swin + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_swin-cli" + ) + add_dependencies(timm_swin_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_swin_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_swin_cli ${PROJECT_SOURCE_DIR}/families/timm_swin/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_swin_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_swin_cli PRIVATE trtmc_cli_timm_swin nlohmann_json::nlohmann_json) + target_compile_options(test_timm_swin_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_swin_cli timm_swin_cli_fixture) + add_test(NAME timm_swin_cli COMMAND test_timm_swin_cli $) + set_tests_properties(timm_swin_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_swin/runtime/cli.cpp b/families/timm_swin/runtime/cli.cpp new file mode 100644 index 0000000000..cbd2f8bd92 --- /dev/null +++ b/families/timm_swin/runtime/cli.cpp @@ -0,0 +1,110 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/core.hpp" +#include "trtmc/features.hpp" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_swin") + throw std::invalid_argument("timm_swin CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + const auto model = + trtmc::Model::load(reader.path(), {runtime_root, 0, runtime_cache, cuda_graphs}); + const auto task = model.task(); + const auto result = task.run( + {trtmc::ImageInput{trtmc::Span{pixels.data(), pixels.size()}, + static_cast(height), static_cast(width)}}, + {}); + std::vector scores(result.scores().begin(), result.scores().end()); + for (const auto value : scores) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned a non-finite score"); + } + const char* kind = nullptr; + switch (result.kind()) { + case TRTMC_SCORE_LOGIT: + kind = "logit"; + break; + case TRTMC_SCORE_PROBABILITY: + kind = "probability"; + break; + case TRTMC_SCORE_UNBOUNDED: + kind = "unbounded"; + break; + default: + throw std::runtime_error("classification returned an unknown score kind"); + } + auto labels = Json::array(); + for (const auto label : result.labels()) + labels.push_back(std::string(label)); + Json output{{"scores", scores}, + {"score_kind", kind}, + {"labels", labels}, + {"vocabulary_id", std::string(result.vocabulary_id())}, + {"task", trtmc::ImageToClassScores::kTask}}; + if (result.kind() == TRTMC_SCORE_LOGIT) + output["logits"] = scores; + if (scores.empty()) { + output["top_class"] = -1; + output["top_score"] = nullptr; + } else { + const auto best = std::max_element(scores.begin(), scores.end()); + output["top_class"] = best - scores.begin(); + output["top_score"] = *best; + } + return output; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_swin CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_swin/tests/cpp/test_cli.cpp b/families/timm_swin/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..691bb83b6b --- /dev/null +++ b/families/timm_swin/tests/cpp/test_cli.cpp @@ -0,0 +1,155 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/internal/features.h" +#include "trtmc/internal/model.h" +#include "trtmc/runtime/family_factory.h" +namespace { +class Fixture final : public trtmc::internal::IModel, public trtmc::internal::IImageToClassScores { + public: + const char* task() const noexcept override { return "image_to_class_scores"; } + std::vector task_bindings() override { + return {trtmc::internal::bind(*this)}; + } + trtmc::internal::LabelScoresResult + run(const trtmc::internal::ImageToClassScoresRequest& request, + trtmc::internal::ConfigView config) override { + const auto& image = request.image; + if (!config.empty() || image.width != 2 || image.height != 1 || image.channels != 3 || + image.format != trtmc::internal::ImageFormat::Float32 || image.byte_size != 24) + throw std::invalid_argument("fixture expects one decoded 2x1 RGB float32 image"); + const auto* pixels = static_cast(image.data); + if (pixels[0] != 1.0F || pixels[1] != 0.0F || pixels[2] != 128.0F / 255.0F || + pixels[4] != 1.0F) + throw std::invalid_argument("owner changed pixel order or range"); + return {{-2.0F, 4.0F, 0.5F}, + {"first", "second", "third"}, + trtmc::internal::ScoreKind::Logit, + "fixture:classes"}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "image_to_class_scores"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_swin"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.at("scores") == actual.at("logits") && actual.at("score_kind") == "logit" && + actual.at("labels") == Json({"first", "second", "third"}) && + actual.at("vocabulary_id") == "fixture:classes" && + actual.at("task") == "image_to_class_scores", + "SDK class identity and score semantics are preserved"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_swin"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_swin-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_swin/tests/manifests/swin-tiny-patch4-window7-224-ms-in1k.json b/families/timm_swin/tests/manifests/swin-tiny-patch4-window7-224-ms-in1k.json index 82619c4de0..4fd726f294 100644 --- a/families/timm_swin/tests/manifests/swin-tiny-patch4-window7-224-ms-in1k.json +++ b/families/timm_swin/tests/manifests/swin-tiny-patch4-window7-224-ms-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_swin/tests/test_cli.py b/families/timm_swin/tests/test_cli.py new file mode 100644 index 0000000000..b97f3f9b61 --- /dev/null +++ b/families/timm_swin/tests/test_cli.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_swin command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_swin import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_swin" +TASK = "image_to_class_scores" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_swin/tests/test_e2e.py b/families/timm_swin/tests/test_e2e.py index 7546b9ab4c..7f1057d6bb 100644 --- a/families/timm_swin/tests/test_e2e.py +++ b/families/timm_swin/tests/test_e2e.py @@ -14,7 +14,7 @@ import numpy as np import pytest -from tensorrt_model_connect import BuildRequest, build +from families.timm_swin.cli import BuildRequest, build_bundle from tools.e2e_evidence import evidence_stage, record_evidence @@ -166,6 +166,8 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") model_dir = _model_dir(manifest) record_evidence( "checkpoint", @@ -173,21 +175,21 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: ) bundle = tmp_path / manifest["bundle"] with evidence_stage("build"): - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=int(manifest["max_sequence_length"]), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - ) + ), + bundle, ) with evidence_stage("native"): completed = subprocess.run( [ str(binary), + FAMILY, "classify", str(bundle), "--runtime-root", diff --git a/families/timm_vgg/cli.json b/families/timm_vgg/cli.json new file mode 100644 index 0000000000..1dc13aa7ea --- /dev/null +++ b/families/timm_vgg/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_vgg bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "classification" + ], + "default": "classification" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_vgg bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_vgg/cli.py b/families/timm_vgg/cli.py new file mode 100644 index 0000000000..eb5536a1f1 --- /dev/null +++ b/families/timm_vgg/cli.py @@ -0,0 +1,123 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_vgg commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "classification" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "classification": + raise ValueError("timm_vgg supports only task=classification") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_vgg supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_vgg does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_vgg does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_vgg does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_vgg does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_vgg does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_vgg does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_vgg does not support context parallelism") + if request.task != "classification": + raise ValueError("timm_vgg supports only task=classification") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_vgg does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_vgg does not support mixed-precision layers") + if request.max_sequence_length not in {None, 1}: + raise NotImplementedError("timm_vgg supports only max_sequence_length=1") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_vgg build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "classification", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_vgg/model.py b/families/timm_vgg/model.py index 156411d43a..0cf842ce1f 100644 --- a/families/timm_vgg/model.py +++ b/families/timm_vgg/model.py @@ -34,8 +34,10 @@ from .config import ModelConfig +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter # timm's VGG stacks 3x3 convolutions with padding 1 and halves with 2x2 max @@ -275,29 +277,7 @@ def get_bundle_config_overrides(self, config: ModelConfig) -> dict: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm VGG image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_vgg does not support dynamic_kv_cache") - - if request.image_height is not None: - raise NotImplementedError("timm_vgg does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_vgg does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_vgg does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_vgg does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_vgg does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_vgg does not support context parallelism") - if request.task != "classification": - raise ValueError("timm_vgg supports only task=classification") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_vgg does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_vgg does not support mixed-precision layers") - if request.max_sequence_length not in {None, 1}: - raise NotImplementedError("timm_vgg supports only max_sequence_length=1") + request = coerce_request(request) model_dir = Path(request.model_dir) config = ModelConfig.from_dir(model_dir) diff --git a/families/timm_vgg/runtime/CMakeLists.txt b/families/timm_vgg/runtime/CMakeLists.txt index 8154dcdb4a..d0dbfbd06b 100644 --- a/families/timm_vgg/runtime/CMakeLists.txt +++ b/families/timm_vgg/runtime/CMakeLists.txt @@ -58,3 +58,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_vgg_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_vgg SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_vgg PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_vgg SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_vgg PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_vgg PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_vgg PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_vgg LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_vgg_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_vgg/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_vgg_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_vgg_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_vgg_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_vgg_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_vgg + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_vgg-cli" + ) + add_dependencies(timm_vgg_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_vgg_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_vgg_cli ${PROJECT_SOURCE_DIR}/families/timm_vgg/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_vgg_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_vgg_cli PRIVATE trtmc_cli_timm_vgg nlohmann_json::nlohmann_json) + target_compile_options(test_timm_vgg_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_vgg_cli timm_vgg_cli_fixture) + add_test(NAME timm_vgg_cli COMMAND test_timm_vgg_cli $) + set_tests_properties(timm_vgg_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_vgg/runtime/cli.cpp b/families/timm_vgg/runtime/cli.cpp new file mode 100644 index 0000000000..fe44818480 --- /dev/null +++ b/families/timm_vgg/runtime/cli.cpp @@ -0,0 +1,79 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_vgg") + throw std::invalid_argument("timm_vgg CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + auto task = trtmc::load_task(reader, runtime_root, 0, runtime_cache, cuda_graphs); + auto* classifier = dynamic_cast(task.get()); + if (!classifier) + throw std::invalid_argument("bundle does not implement image classification"); + const auto result = classifier->classify(pixels.data(), height, width); + for (const auto value : result.logits) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned non-finite logits"); + } + if (!std::isfinite(result.top_score)) + throw std::runtime_error("classification returned a non-finite top score"); + return {{"logits", result.logits}, + {"top_class", result.top_class}, + {"top_score", result.top_score}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_vgg CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_vgg/tests/cpp/test_cli.cpp b/families/timm_vgg/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..b763ba70b4 --- /dev/null +++ b/families/timm_vgg/tests/cpp/test_cli.cpp @@ -0,0 +1,138 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/runtime/family_factory.h" +#include "trtmc/task.h" +namespace { +class Fixture final : public trtmc::IImageClassification { + public: + const char* task() const noexcept override { return "classification"; } + trtmc::ClassificationResult classify(const float* pixels, std::int32_t height, + std::int32_t width) override { + if (width != 2 || height != 1 || pixels[0] != 1.0F || pixels[1] != 0.0F || + pixels[2] != 128.0F / 255.0F || pixels[4] != 1.0F) + throw std::invalid_argument("owner changed decoded image shape, order or range"); + return {{-2.0F, 4.0F, 0.5F}, 1, 4.0F}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "classification"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_vgg"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.size() == 3, "legacy classification output shape remains unchanged"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_vgg"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_vgg-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_vgg/tests/manifests/vgg16-tv-in1k.json b/families/timm_vgg/tests/manifests/vgg16-tv-in1k.json index 6e94e1192e..cad875666b 100644 --- a/families/timm_vgg/tests/manifests/vgg16-tv-in1k.json +++ b/families/timm_vgg/tests/manifests/vgg16-tv-in1k.json @@ -12,6 +12,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_vgg/tests/test_cli.py b/families/timm_vgg/tests/test_cli.py new file mode 100644 index 0000000000..8ed6914950 --- /dev/null +++ b/families/timm_vgg/tests/test_cli.py @@ -0,0 +1,126 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_vgg command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_vgg import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_vgg" +TASK = "classification" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + with pytest.raises(NotImplementedError, match="max_sequence_length"): + cli.coerce_request(replace(legacy, max_sequence_length=2)) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_vgg/tests/test_e2e.py b/families/timm_vgg/tests/test_e2e.py index 7c4f1fa2ee..0de8e7a57b 100644 --- a/families/timm_vgg/tests/test_e2e.py +++ b/families/timm_vgg/tests/test_e2e.py @@ -12,7 +12,7 @@ from pathlib import Path import pytest import numpy as np -from tensorrt_model_connect import BuildRequest, build +from families.timm_vgg.cli import BuildRequest, build_bundle FAMILY = "timm_vgg" TASKS = frozenset({"classification"}) @@ -114,6 +114,8 @@ def _runtime() -> tuple[Path, Path]: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") import torch assert torch.cuda.is_available(), f"selected {FAMILY} E2E requires CUDA" @@ -122,22 +124,15 @@ def _runtime() -> tuple[Path, Path]: def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=manifest.get("max_sequence_length"), - image_height=manifest.get("image_height"), - image_width=manifest.get("image_width"), - video_num_frames=manifest.get("video_num_frames"), - max_batch_size=int(manifest.get("max_batch_size", 1)), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - quantization=manifest.get("quantization"), - fp32_layers=tuple((int(layer) for layer in manifest.get("fp32_layers", ()))), - ) + ), + bundle, ) @@ -151,6 +146,7 @@ def _run_json( ) -> dict: invocation = [ str(binary), + FAMILY, command, str(bundle), "--runtime-root", diff --git a/families/timm_vit/cli.json b/families/timm_vit/cli.json new file mode 100644 index 0000000000..4f0f059208 --- /dev/null +++ b/families/timm_vit/cli.json @@ -0,0 +1,137 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_vit bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "image_to_class_scores" + ], + "default": "image_to_class_scores" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "tensor_parallel_size", + "flags": [ + "--tensor-parallel-size" + ], + "type": "int", + "choices": [ + 1, + 2, + 4, + 8 + ], + "default": 1 + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_vit bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_vit/cli.py b/families/timm_vit/cli.py new file mode 100644 index 0000000000..f409220dd4 --- /dev/null +++ b/families/timm_vit/cli.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_vit commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "image_to_class_scores" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + tensor_parallel_size: int = 1 + + def __post_init__(self) -> None: + if self.task != "image_to_class_scores": + raise ValueError("timm_vit supports only task=image_to_class_scores") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_vit supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + if type(self.tensor_parallel_size) is not int or self.tensor_parallel_size not in { + 1, + 2, + 4, + 8, + }: + raise ValueError("timm ViT tensor_parallel_size must be one of 1, 2, 4, 8") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_vit does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_vit does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_vit does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_vit does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_vit does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.context_parallel_size != 1: + raise ValueError("this family does not support context parallelism") + if request.task != "image_to_class_scores": + raise ValueError("timm_vit supports only task=image_to_class_scores") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm ViT does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm ViT does not support mixed-precision layers") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_vit build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, + request.task, + request.precision, + request.backend, + request.verbose, + tensor_parallel_size=_positive_int(request.tensor_parallel_size, "tensor_parallel_size"), + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "image_to_class_scores", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, + tensor_parallel_size: int = 1, +) -> int: + request = BuildRequest( + resolve_model(model, revision), + task, + precision, + backend, + verbose, + tensor_parallel_size=tensor_parallel_size, + ) + build_bundle(request, output) + return 0 diff --git a/families/timm_vit/model.py b/families/timm_vit/model.py index 2cca55d4d3..0f8b29f015 100644 --- a/families/timm_vit/model.py +++ b/families/timm_vit/model.py @@ -30,6 +30,7 @@ _target_np_dtype, _transpose_2d, ) +from .cli import BuildRequest, coerce_request from .config import ModelConfig from .parallel import ParallelConfig from .parallel import normalize_parallel_config @@ -69,7 +70,6 @@ def _resolve_vit_config(raw: dict) -> dict: if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -482,26 +482,8 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm ViT bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_vit does not support dynamic_kv_cache") + request = coerce_request(request) - if request.image_height is not None: - raise NotImplementedError("timm_vit does not support image_height") - - if request.image_width is not None: - raise NotImplementedError("timm_vit does not support image_width") - - if request.video_num_frames is not None: - raise NotImplementedError("timm_vit does not support video_num_frames") - - if request.max_batch_size != 1: - raise NotImplementedError("timm_vit does not support max_batch_size") - - if request.context_parallel_size != 1: - raise ValueError("this family does not support context parallelism") - - if request.task != "image_to_class_scores": - raise ValueError("timm_vit supports only task=image_to_class_scores") model_dir = Path(request.model_dir) config = ModelConfig.from_dir(model_dir) if not ( @@ -510,11 +492,8 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: ): raise ValueError(f"timm ViT does not support model_type={config.model_type!r}") precision = str(request.precision).lower() - max_sequence_length = _positive_int(request.max_sequence_length or 1, "max_sequence_length") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm ViT does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm ViT does not support mixed-precision layers") + max_sequence_length = 1 + parallel = ParallelConfig( tp_size=_positive_int(request.tensor_parallel_size, "tensor_parallel_size") ) diff --git a/families/timm_vit/runtime/CMakeLists.txt b/families/timm_vit/runtime/CMakeLists.txt index a9312314d1..50d5decee2 100644 --- a/families/timm_vit/runtime/CMakeLists.txt +++ b/families/timm_vit/runtime/CMakeLists.txt @@ -86,3 +86,39 @@ if(TRTMC_BUILD_TESTS) target_compile_options(test_timm_vit_image_preprocess PRIVATE -Wall -Wextra -Wpedantic) add_test(NAME timm_vit_image_preprocess COMMAND test_timm_vit_image_preprocess) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_vit SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_vit PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_vit SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_vit PRIVATE trtmc_c trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_vit PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_vit PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_vit LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_vit_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_vit/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_vit_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_vit_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_vit_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_vit_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_vit + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_vit-cli" + ) + add_dependencies(timm_vit_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_vit_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_vit_cli ${PROJECT_SOURCE_DIR}/families/timm_vit/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_vit_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_vit_cli PRIVATE trtmc_cli_timm_vit nlohmann_json::nlohmann_json) + target_compile_options(test_timm_vit_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_vit_cli timm_vit_cli_fixture) + add_test(NAME timm_vit_cli COMMAND test_timm_vit_cli $) + set_tests_properties(timm_vit_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_vit/runtime/cli.cpp b/families/timm_vit/runtime/cli.cpp new file mode 100644 index 0000000000..7463519ae3 --- /dev/null +++ b/families/timm_vit/runtime/cli.cpp @@ -0,0 +1,110 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/core.hpp" +#include "trtmc/features.hpp" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_vit") + throw std::invalid_argument("timm_vit CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + const auto model = + trtmc::Model::load(reader.path(), {runtime_root, 0, runtime_cache, cuda_graphs}); + const auto task = model.task(); + const auto result = task.run( + {trtmc::ImageInput{trtmc::Span{pixels.data(), pixels.size()}, + static_cast(height), static_cast(width)}}, + {}); + std::vector scores(result.scores().begin(), result.scores().end()); + for (const auto value : scores) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned a non-finite score"); + } + const char* kind = nullptr; + switch (result.kind()) { + case TRTMC_SCORE_LOGIT: + kind = "logit"; + break; + case TRTMC_SCORE_PROBABILITY: + kind = "probability"; + break; + case TRTMC_SCORE_UNBOUNDED: + kind = "unbounded"; + break; + default: + throw std::runtime_error("classification returned an unknown score kind"); + } + auto labels = Json::array(); + for (const auto label : result.labels()) + labels.push_back(std::string(label)); + Json output{{"scores", scores}, + {"score_kind", kind}, + {"labels", labels}, + {"vocabulary_id", std::string(result.vocabulary_id())}, + {"task", trtmc::ImageToClassScores::kTask}}; + if (result.kind() == TRTMC_SCORE_LOGIT) + output["logits"] = scores; + if (scores.empty()) { + output["top_class"] = -1; + output["top_score"] = nullptr; + } else { + const auto best = std::max_element(scores.begin(), scores.end()); + output["top_class"] = best - scores.begin(); + output["top_score"] = *best; + } + return output; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_vit CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_vit/tests/cpp/test_cli.cpp b/families/timm_vit/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..1662dc3c93 --- /dev/null +++ b/families/timm_vit/tests/cpp/test_cli.cpp @@ -0,0 +1,155 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/internal/features.h" +#include "trtmc/internal/model.h" +#include "trtmc/runtime/family_factory.h" +namespace { +class Fixture final : public trtmc::internal::IModel, public trtmc::internal::IImageToClassScores { + public: + const char* task() const noexcept override { return "image_to_class_scores"; } + std::vector task_bindings() override { + return {trtmc::internal::bind(*this)}; + } + trtmc::internal::LabelScoresResult + run(const trtmc::internal::ImageToClassScoresRequest& request, + trtmc::internal::ConfigView config) override { + const auto& image = request.image; + if (!config.empty() || image.width != 2 || image.height != 1 || image.channels != 3 || + image.format != trtmc::internal::ImageFormat::Float32 || image.byte_size != 24) + throw std::invalid_argument("fixture expects one decoded 2x1 RGB float32 image"); + const auto* pixels = static_cast(image.data); + if (pixels[0] != 1.0F || pixels[1] != 0.0F || pixels[2] != 128.0F / 255.0F || + pixels[4] != 1.0F) + throw std::invalid_argument("owner changed pixel order or range"); + return {{-2.0F, 4.0F, 0.5F}, + {"first", "second", "third"}, + trtmc::internal::ScoreKind::Logit, + "fixture:classes"}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "image_to_class_scores"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_vit"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.at("scores") == actual.at("logits") && actual.at("score_kind") == "logit" && + actual.at("labels") == Json({"first", "second", "third"}) && + actual.at("vocabulary_id") == "fixture:classes" && + actual.at("task") == "image_to_class_scores", + "SDK class identity and score semantics are preserved"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_vit"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_vit-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k-tp4.json b/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k-tp4.json index df778df742..43761be6e7 100644 --- a/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k-tp4.json +++ b/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k-tp4.json @@ -11,6 +11,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 4, - "max_sequence_length": 1 + "tensor_parallel_size": 4 } diff --git a/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k.json b/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k.json index c192d93ee9..c8c56cb0c0 100644 --- a/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k.json +++ b/families/timm_vit/tests/manifests/timm-vit-base-p16-224-augreg-in21k-ft-in1k.json @@ -12,6 +12,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_vit/tests/test_cli.py b/families/timm_vit/tests/test_cli.py new file mode 100644 index 0000000000..b4e365521d --- /dev/null +++ b/families/timm_vit/tests/test_cli.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_vit command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_vit import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_vit" +TASK = "image_to_class_scores" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] + + +def test_tensor_parallel_build_is_family_owned(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr(cli, "build_bundle", lambda request, output: calls.append(request)) + assert ( + family_cli.main( + [ + FAMILY, + "build", + str(tmp_path), + "-o", + str(tmp_path / "out"), + "--tensor-parallel-size", + "4", + "--precision", + "fp32", + ] + ) + == 0 + ) + assert calls[0].tensor_parallel_size == 4 + with pytest.raises(ValueError, match="one of"): + cli.BuildRequest(tmp_path, tensor_parallel_size=3) diff --git a/families/timm_vit/tests/test_e2e.py b/families/timm_vit/tests/test_e2e.py index 2136aaef7f..4a28637bbf 100644 --- a/families/timm_vit/tests/test_e2e.py +++ b/families/timm_vit/tests/test_e2e.py @@ -14,7 +14,7 @@ from pathlib import Path import pytest import numpy as np -from tensorrt_model_connect import BuildRequest, build +from families.timm_vit.cli import BuildRequest, build_bundle FAMILY = "timm_vit" TASKS = frozenset({"image_to_class_scores"}) @@ -116,6 +116,8 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") import torch required_gpus = int(manifest["tensor_parallel_size"]) @@ -127,22 +129,14 @@ def _runtime(manifest: dict) -> tuple[Path, Path]: def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: - build( + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=manifest.get("max_sequence_length"), - image_height=manifest.get("image_height"), - image_width=manifest.get("image_width"), - video_num_frames=manifest.get("video_num_frames"), - max_batch_size=int(manifest.get("max_batch_size", 1)), tensor_parallel_size=int(manifest["tensor_parallel_size"]), - quantization=manifest.get("quantization"), - fp32_layers=tuple((int(layer) for layer in manifest.get("fp32_layers", ()))), - ) + ), + bundle, ) @@ -157,6 +151,7 @@ def _run_json( ) -> dict: invocation = [ str(binary), + FAMILY, command, str(bundle), "--runtime-root", diff --git a/families/timm_xception/cli.json b/families/timm_xception/cli.json new file mode 100644 index 0000000000..a8c15ce9f3 --- /dev/null +++ b/families/timm_xception/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_xception bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "classification" + ], + "default": "classification" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_xception bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_xception/cli.py b/families/timm_xception/cli.py new file mode 100644 index 0000000000..7a6e7153f6 --- /dev/null +++ b/families/timm_xception/cli.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_xception commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "classification" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "classification": + raise ValueError("timm_xception supports only task=classification") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_xception supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool): + raise ValueError(f"{name} must be a positive integer") + result = int(value) + if result < 1: + raise ValueError(f"{name} must be a positive integer") + return result + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_xception does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_xception does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_xception does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_xception does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_xception does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_xception does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_xception does not support context parallelism") + if request.task != "classification": + raise ValueError("timm_xception supports only task=classification") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_xception does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_xception does not support mixed-precision layers") + _positive_int(request.max_sequence_length or 1, "max_sequence_length") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_xception build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "classification", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_xception/model.py b/families/timm_xception/model.py index 2a6c1eaaf0..0d20e1f46a 100644 --- a/families/timm_xception/model.py +++ b/families/timm_xception/model.py @@ -29,8 +29,10 @@ from .checkpoint import Checkpoint +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -339,27 +341,8 @@ def _positive_int(value: object, name: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm Xception image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_xception does not support dynamic_kv_cache") - if request.image_height is not None: - raise NotImplementedError("timm_xception does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_xception does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_xception does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_xception does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_xception does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_xception does not support context parallelism") - if request.task != "classification": - raise ValueError("timm_xception supports only task=classification") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_xception does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_xception does not support mixed-precision layers") - _positive_int(request.max_sequence_length or 1, "max_sequence_length") + request = coerce_request(request) + model_dir = Path(request.model_dir) raw = _read_config(model_dir) plan, runtime = _build_engine( diff --git a/families/timm_xception/runtime/CMakeLists.txt b/families/timm_xception/runtime/CMakeLists.txt index 5a5b46223d..f04faa6fdb 100644 --- a/families/timm_xception/runtime/CMakeLists.txt +++ b/families/timm_xception/runtime/CMakeLists.txt @@ -55,3 +55,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_xception_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_xception SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_xception PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_xception SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_xception PRIVATE trtmc_runtime trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_xception PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_xception PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_xception LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_xception_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_xception/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_xception_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_xception_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_xception_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_xception_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_xception + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_xception-cli" + ) + add_dependencies(timm_xception_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_xception_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_xception_cli ${PROJECT_SOURCE_DIR}/families/timm_xception/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_xception_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_xception_cli PRIVATE trtmc_cli_timm_xception nlohmann_json::nlohmann_json) + target_compile_options(test_timm_xception_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_xception_cli timm_xception_cli_fixture) + add_test(NAME timm_xception_cli COMMAND test_timm_xception_cli $) + set_tests_properties(timm_xception_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_xception/runtime/cli.cpp b/families/timm_xception/runtime/cli.cpp new file mode 100644 index 0000000000..db691ad192 --- /dev/null +++ b/families/timm_xception/runtime/cli.cpp @@ -0,0 +1,79 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/runtime/family_loader.h" +#include "trtmc/task.h" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_xception") + throw std::invalid_argument("timm_xception CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + auto task = trtmc::load_task(reader, runtime_root, 0, runtime_cache, cuda_graphs); + auto* classifier = dynamic_cast(task.get()); + if (!classifier) + throw std::invalid_argument("bundle does not implement image classification"); + const auto result = classifier->classify(pixels.data(), height, width); + for (const auto value : result.logits) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned non-finite logits"); + } + if (!std::isfinite(result.top_score)) + throw std::runtime_error("classification returned a non-finite top score"); + return {{"logits", result.logits}, + {"top_class", result.top_class}, + {"top_score", result.top_score}}; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_xception CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_xception/tests/cpp/test_cli.cpp b/families/timm_xception/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..03f2a4429d --- /dev/null +++ b/families/timm_xception/tests/cpp/test_cli.cpp @@ -0,0 +1,138 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/runtime/family_factory.h" +#include "trtmc/task.h" +namespace { +class Fixture final : public trtmc::IImageClassification { + public: + const char* task() const noexcept override { return "classification"; } + trtmc::ClassificationResult classify(const float* pixels, std::int32_t height, + std::int32_t width) override { + if (width != 2 || height != 1 || pixels[0] != 1.0F || pixels[1] != 0.0F || + pixels[2] != 128.0F / 255.0F || pixels[4] != 1.0F) + throw std::invalid_argument("owner changed decoded image shape, order or range"); + return {{-2.0F, 4.0F, 0.5F}, 1, 4.0F}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "classification"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_xception"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.size() == 3, "legacy classification output shape remains unchanged"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_xception"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_xception-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_xception/tests/manifests/xception41-tf-in1k.json b/families/timm_xception/tests/manifests/xception41-tf-in1k.json index 132901f93d..eb90727ac9 100644 --- a/families/timm_xception/tests/manifests/xception41-tf-in1k.json +++ b/families/timm_xception/tests/manifests/xception41-tf-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_xception/tests/test_cli.py b/families/timm_xception/tests/test_cli.py new file mode 100644 index 0000000000..37e2823b92 --- /dev/null +++ b/families/timm_xception/tests/test_cli.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_xception command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_xception import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_xception" +TASK = "classification" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + assert cli.coerce_request(replace(legacy, max_sequence_length=7)) == request + with pytest.raises(ValueError, match="positive integer"): + cli.coerce_request(SimpleNamespace(**{**vars(legacy), "max_sequence_length": -1})) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_xception/tests/test_e2e.py b/families/timm_xception/tests/test_e2e.py index 91e2da600d..4496db42e2 100644 --- a/families/timm_xception/tests/test_e2e.py +++ b/families/timm_xception/tests/test_e2e.py @@ -15,7 +15,7 @@ import numpy as np import pytest -from tensorrt_model_connect import BuildRequest, build +from families.timm_xception.cli import BuildRequest, build_bundle FAMILY = "timm_xception" @@ -115,25 +115,27 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: assert (runtime_root / "libtrtmc_backend_trt.so").is_file() with evidence_stage("setup"): assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") model_dir = _model_dir(manifest) record_evidence("checkpoint", {"model_dir": str(model_dir), "hf_id": manifest.get("hf_id"), "hf_revision": manifest.get("hf_revision")}) bundle = tmp_path / manifest["bundle"] with evidence_stage("build"): - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=int(manifest["max_sequence_length"]), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - ) + ), + bundle, ) with evidence_stage("native"): completed = subprocess.run( [ str(binary), + FAMILY, "classify", str(bundle), "--runtime-root", diff --git a/families/timm_xcit/cli.json b/families/timm_xcit/cli.json new file mode 100644 index 0000000000..3f464c4e8e --- /dev/null +++ b/families/timm_xcit/cli.json @@ -0,0 +1,123 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one timm_xcit bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Model ID or local checkpoint directory" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "image_to_class_scores" + ], + "default": "image_to_class_scores" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp16", + "fp32" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + } + ] + }, + { + "name": "classify", + "help": "Classify an image with a timm_xcit bundle", + "executor": "native", + "handler": "classify", + "arguments": [ + { + "name": "bundle", + "type": "path" + }, + { + "name": "image", + "flags": [ + "--image" + ], + "type": "path", + "required": true + }, + { + "name": "runtime_root", + "flags": [ + "--runtime-root" + ], + "type": "path", + "help": "Override the installed runtime directory" + }, + { + "name": "runtime_cache", + "flags": [ + "--runtime-cache" + ], + "type": "path", + "help": "TensorRT-RTX runtime cache" + }, + { + "name": "cuda_graphs", + "flags": [ + "--cuda-graphs" + ], + "type": "bool", + "action": "store_true", + "help": "TensorRT-RTX CUDA graph capture" + } + ] + } + ] +} diff --git a/families/timm_xcit/cli.py b/families/timm_xcit/cli.py new file mode 100644 index 0000000000..72e3f1a8f0 --- /dev/null +++ b/families/timm_xcit/cli.py @@ -0,0 +1,123 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned timm_xcit commands with lazy builder imports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import GraphTransform, graph_transform +from tensorrt_model_connect.model_support import resolve_model + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = "image_to_class_scores" + precision: str = "fp32" + backend: str = "trt" + verbose: bool = False + + def __post_init__(self) -> None: + if self.task != "image_to_class_scores": + raise ValueError("timm_xcit supports only task=image_to_class_scores") + if str(self.precision).lower() not in {"fp16", "fp32"}: + raise ValueError("timm_xcit supports only fp16 and fp32 precision") + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + + +def _validate_legacy_inputs(request: object) -> None: + if request.dynamic_kv_cache: + raise NotImplementedError("timm_xcit does not support dynamic_kv_cache") + if request.image_height is not None: + raise NotImplementedError("timm_xcit does not support image_height") + if request.image_width is not None: + raise NotImplementedError("timm_xcit does not support image_width") + if request.video_num_frames is not None: + raise NotImplementedError("timm_xcit does not support video_num_frames") + if request.max_batch_size != 1: + raise NotImplementedError("timm_xcit does not support max_batch_size") + + +def _validate_legacy_build(request: object) -> None: + if request.tensor_parallel_size != 1: + raise NotImplementedError("timm_xcit does not support tensor parallelism") + if request.context_parallel_size != 1: + raise NotImplementedError("timm_xcit does not support context parallelism") + if request.task != "image_to_class_scores": + raise ValueError("timm_xcit supports only task=image_to_class_scores") + if request.quantization not in {None, "none"}: + raise NotImplementedError("timm_xcit does not support quantization") + if request.fp32_layers: + raise NotImplementedError("timm_xcit does not support mixed-precision layers") + if request.max_sequence_length not in {None, 1}: + raise NotImplementedError("timm_xcit supports only max_sequence_length=1") + + +def coerce_request(request: object) -> BuildRequest: + """Retain legacy rejection behavior without retaining its shared request union.""" + if isinstance(request, BuildRequest): + return request + _validate_legacy_inputs(request) + _validate_legacy_build(request) + allowed = { + "model_dir", + "task", + "precision", + "backend", + "verbose", + "family", + "output_path", + "graph_transform", + "dynamic_kv_cache", + "image_height", + "image_width", + "video_num_frames", + "max_batch_size", + "tensor_parallel_size", + "context_parallel_size", + "quantization", + "fp32_layers", + "max_sequence_length", + } + if unknown := set(vars(request)) - allowed: + raise ValueError(f"unknown timm_xcit build inputs: {sorted(unknown)}") + return BuildRequest( + request.model_dir, request.task, request.precision, request.backend, request.verbose + ) + + +def build_bundle( + request: BuildRequest, output: Path, *, transform: GraphTransform | None = None +) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build( + *, + model: str, + output: Path, + revision: str | None = None, + task: str = "image_to_class_scores", + precision: str = "fp32", + backend: str = "trt", + verbose: bool = False, +) -> int: + request = BuildRequest(resolve_model(model, revision), task, precision, backend, verbose) + build_bundle(request, output) + return 0 diff --git a/families/timm_xcit/model.py b/families/timm_xcit/model.py index ba6281f651..5d3a9befc1 100644 --- a/families/timm_xcit/model.py +++ b/families/timm_xcit/model.py @@ -36,8 +36,10 @@ from .checkpoint import Checkpoint +from .cli import BuildRequest, coerce_request + + if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -668,28 +670,8 @@ def _build_engine( def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one timm XCiT image-classification bundle.""" - if request.dynamic_kv_cache: - raise NotImplementedError("timm_xcit does not support dynamic_kv_cache") - if request.image_height is not None: - raise NotImplementedError("timm_xcit does not support image_height") - if request.image_width is not None: - raise NotImplementedError("timm_xcit does not support image_width") - if request.video_num_frames is not None: - raise NotImplementedError("timm_xcit does not support video_num_frames") - if request.max_batch_size != 1: - raise NotImplementedError("timm_xcit does not support max_batch_size") - if request.tensor_parallel_size != 1: - raise NotImplementedError("timm_xcit does not support tensor parallelism") - if request.context_parallel_size != 1: - raise NotImplementedError("timm_xcit does not support context parallelism") - if request.task != "image_to_class_scores": - raise ValueError("timm_xcit supports only task=image_to_class_scores") - if request.quantization not in {None, "none"}: - raise NotImplementedError("timm_xcit does not support quantization") - if request.fp32_layers: - raise NotImplementedError("timm_xcit does not support mixed-precision layers") - if request.max_sequence_length not in {None, 1}: - raise NotImplementedError("timm_xcit supports only max_sequence_length=1") + request = coerce_request(request) + model_dir = Path(request.model_dir) raw = _read_config(model_dir) plan, runtime = _build_engine( diff --git a/families/timm_xcit/runtime/CMakeLists.txt b/families/timm_xcit/runtime/CMakeLists.txt index b199080c8d..be17ba3ab3 100644 --- a/families/timm_xcit/runtime/CMakeLists.txt +++ b/families/timm_xcit/runtime/CMakeLists.txt @@ -90,3 +90,39 @@ if(TRTMC_BUILD_TESTS) COMMAND test_timm_xcit_image_preprocess ) endif() + +# The family CLI is an application adapter; the model does not depend on its loader. +add_library(trtmc_cli_timm_xcit SHARED cli.cpp) +target_include_directories(trtmc_cli_timm_xcit PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) +target_include_directories(trtmc_cli_timm_xcit SYSTEM PRIVATE ${PROJECT_SOURCE_DIR}/third_party/stb) +target_link_libraries(trtmc_cli_timm_xcit PRIVATE trtmc_c trtmc_core nlohmann_json::nlohmann_json) +target_compile_options(trtmc_cli_timm_xcit PRIVATE -Wall -Wextra -Wno-unused-function) +set_target_properties(trtmc_cli_timm_xcit PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_cli_timm_xcit LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR}) + +if(TRTMC_BUILD_TESTS) + add_library(timm_xcit_cli_fixture SHARED ${PROJECT_SOURCE_DIR}/families/timm_xcit/tests/cpp/test_cli.cpp) + target_compile_definitions(timm_xcit_cli_fixture PRIVATE TRTMC_FAMILY_CLI_FIXTURE) + target_include_directories(timm_xcit_cli_fixture PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(timm_xcit_cli_fixture PRIVATE trtmc_core nlohmann_json::nlohmann_json) + set_target_properties(timm_xcit_cli_fixture PROPERTIES + OUTPUT_NAME trtmc_model_timm_xcit + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/timm_xcit-cli" + ) + add_dependencies(timm_xcit_cli_fixture trtmc_test_backend_fake) + add_custom_command(TARGET timm_xcit_cli_fixture POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different $ + $ + ) + add_executable(test_timm_xcit_cli ${PROJECT_SOURCE_DIR}/families/timm_xcit/tests/cpp/test_cli.cpp) + target_include_directories(test_timm_xcit_cli PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) + target_link_libraries(test_timm_xcit_cli PRIVATE trtmc_cli_timm_xcit nlohmann_json::nlohmann_json) + target_compile_options(test_timm_xcit_cli PRIVATE -Wall -Wextra -Wpedantic) + add_dependencies(test_timm_xcit_cli timm_xcit_cli_fixture) + add_test(NAME timm_xcit_cli COMMAND test_timm_xcit_cli $) + set_tests_properties(timm_xcit_cli PROPERTIES LABELS cpu) +endif() diff --git a/families/timm_xcit/runtime/cli.cpp b/families/timm_xcit/runtime/cli.cpp new file mode 100644 index 0000000000..0fe5d834d7 --- /dev/null +++ b/families/timm_xcit/runtime/cli.cpp @@ -0,0 +1,110 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/internal/cli.h" + +#include "trtmc/bundle.h" +#include "trtmc/core.hpp" +#include "trtmc/features.hpp" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include "stb_image.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { +using Json = nlohmann::json; + +std::vector read_image(const std::string& path, int& width, int& height) { + int channels = 0; + std::unique_ptr image( + stbi_load(path.c_str(), &width, &height, &channels, 3), stbi_image_free); + if (!image || width <= 0 || height <= 0) + throw std::invalid_argument("unable to decode classification image"); + std::vector pixels(static_cast(width) * height * 3); + std::transform(image.get(), image.get() + pixels.size(), pixels.begin(), + [](stbi_uc value) { return value / 255.0F; }); + return pixels; +} + +Json classify(const Json& values, const char* default_runtime_root) { + const trtmc::BundleReader reader(values.at("bundle").get()); + if (reader.info().family != "timm_xcit") + throw std::invalid_argument("timm_xcit CLI requires its own family bundle"); + int width = 0, height = 0; + const auto pixels = read_image(values.at("image").get(), width, height); + const auto runtime_root = values.value("runtime_root", std::string(default_runtime_root)); + const auto runtime_cache = values.value("runtime_cache", std::string{}); + const auto cuda_graphs = values.value("cuda_graphs", false); + const auto model = + trtmc::Model::load(reader.path(), {runtime_root, 0, runtime_cache, cuda_graphs}); + const auto task = model.task(); + const auto result = task.run( + {trtmc::ImageInput{trtmc::Span{pixels.data(), pixels.size()}, + static_cast(height), static_cast(width)}}, + {}); + std::vector scores(result.scores().begin(), result.scores().end()); + for (const auto value : scores) { + if (!std::isfinite(value)) + throw std::runtime_error("classification returned a non-finite score"); + } + const char* kind = nullptr; + switch (result.kind()) { + case TRTMC_SCORE_LOGIT: + kind = "logit"; + break; + case TRTMC_SCORE_PROBABILITY: + kind = "probability"; + break; + case TRTMC_SCORE_UNBOUNDED: + kind = "unbounded"; + break; + default: + throw std::runtime_error("classification returned an unknown score kind"); + } + auto labels = Json::array(); + for (const auto label : result.labels()) + labels.push_back(std::string(label)); + Json output{{"scores", scores}, + {"score_kind", kind}, + {"labels", labels}, + {"vocabulary_id", std::string(result.vocabulary_id())}, + {"task", trtmc::ImageToClassScores::kTask}}; + if (result.kind() == TRTMC_SCORE_LOGIT) + output["logits"] = scores; + if (scores.empty()) { + output["top_class"] = -1; + output["top_score"] = nullptr; + } else { + const auto best = std::max_element(scores.begin(), scores.end()); + output["top_class"] = best - scores.begin(); + output["top_score"] = *best; + } + return output; +} +} // namespace + +extern "C" int trtmc_family_cli_v1(const char* handler, const char* values_json, + const char* default_runtime_root, void* context, + trtmc_cli_write_v1 output, trtmc_cli_write_v1 error) { + try { + if (std::string(handler) != "classify") + throw std::invalid_argument("unknown timm_xcit CLI handler"); + const auto result = classify(Json::parse(values_json), default_runtime_root).dump() + '\n'; + output(context, result.data(), result.size()); + return 0; + } catch (const std::exception& exception) { + const auto message = std::string("Error: ") + exception.what() + '\n'; + error(context, message.data(), message.size()); + return 1; + } +} diff --git a/families/timm_xcit/tests/cpp/test_cli.cpp b/families/timm_xcit/tests/cpp/test_cli.cpp new file mode 100644 index 0000000000..12f0e86336 --- /dev/null +++ b/families/timm_xcit/tests/cpp/test_cli.cpp @@ -0,0 +1,155 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/bundle.h" +#include "trtmc/internal/cli.h" + +#include +#include +#include + +#ifdef TRTMC_FAMILY_CLI_FIXTURE +#include "trtmc/internal/features.h" +#include "trtmc/internal/model.h" +#include "trtmc/runtime/family_factory.h" +namespace { +class Fixture final : public trtmc::internal::IModel, public trtmc::internal::IImageToClassScores { + public: + const char* task() const noexcept override { return "image_to_class_scores"; } + std::vector task_bindings() override { + return {trtmc::internal::bind(*this)}; + } + trtmc::internal::LabelScoresResult + run(const trtmc::internal::ImageToClassScoresRequest& request, + trtmc::internal::ConfigView config) override { + const auto& image = request.image; + if (!config.empty() || image.width != 2 || image.height != 1 || image.channels != 3 || + image.format != trtmc::internal::ImageFormat::Float32 || image.byte_size != 24) + throw std::invalid_argument("fixture expects one decoded 2x1 RGB float32 image"); + const auto* pixels = static_cast(image.data); + if (pixels[0] != 1.0F || pixels[1] != 0.0F || pixels[2] != 128.0F / 255.0F || + pixels[4] != 1.0F) + throw std::invalid_argument("owner changed pixel order or range"); + return {{-2.0F, 4.0F, 0.5F}, + {"first", "second", "third"}, + trtmc::internal::ScoreKind::Logit, + "fixture:classes"}; + } +}; +} // namespace +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext&) { + return new Fixture(); +} +#else +#include +#include +#include +#include + +namespace { +namespace fs = std::filesystem; +using Json = nlohmann::json; +int failures = 0; +void check(bool condition, const char* message) { + if (!condition) { + std::cerr << "FAIL: " << message << '\n'; + ++failures; + } +} +void bundle(const fs::path& path, const std::string& family) { + const auto header = Json{{"format", 1}, + {"family", family}, + {"task", "image_to_class_scores"}, + {"backend", "fake"}, + {"sections", {{"engine.plan", {{"offset", 0}, {"length", 4}}}}}} + .dump(); + std::ofstream output(path, std::ios::binary); + output.write("BUNDLE\x01\x00", 8); + for (unsigned shift = 0; shift < 64; shift += 8) + output.put(static_cast((static_cast(header.size()) >> shift) & 255U)); + output.write(header.data(), static_cast(header.size())); + output.write("PLAN", 4); +} +struct Capture { + std::string output, error; +}; +void emit(void* context, const char* data, std::size_t size) { + static_cast(context)->output.append(data, size); +} +void reject(void* context, const char* data, std::size_t size) { + static_cast(context)->error.append(data, size); +} +void contract(const fs::path& runtime_root, const fs::path& root) { + const auto path = root / "model.bundle"; + const auto image = root / "image.ppm"; + bundle(path, "timm_xcit"); + { + std::ofstream output(image, std::ios::binary); + output << "P6\n2 1\n255\n"; + const unsigned char pixels[] = {255, 0, 128, 0, 255, 0}; + output.write(reinterpret_cast(pixels), sizeof(pixels)); + } + auto invoke = [&](Json values, const char* handler = "classify", + const std::string& fallback = "") { + Capture captured; + const auto root_value = fallback.empty() ? runtime_root.string() : fallback; + const auto status = trtmc_family_cli_v1(handler, values.dump().c_str(), root_value.c_str(), + &captured, emit, reject); + return std::pair{status, captured}; + }; + Json values{{"bundle", path.string()}, {"image", image.string()}}; + const auto result = invoke(values); + check(result.first == 0, "owner classification succeeds through native callback and runtime"); + if (result.first == 0) { + const auto actual = Json::parse(result.second.output); + check(actual.at("logits") == Json({-2.0, 4.0, 0.5}) && actual.at("top_class") == 1 && + actual.at("top_score") == 4.0, + "classification preserves logits and argmax without softmax"); + check(actual.at("scores") == actual.at("logits") && actual.at("score_kind") == "logit" && + actual.at("labels") == Json({"first", "second", "third"}) && + actual.at("vocabulary_id") == "fixture:classes" && + actual.at("task") == "image_to_class_scores", + "SDK class identity and score semantics are preserved"); + } + auto explicit_root = values; + explicit_root["runtime_root"] = runtime_root.string(); + check(invoke(explicit_root, "classify", (root / "missing").string()).first == 0, + "explicit runtime root overrides the installed default"); + explicit_root["runtime_root"] = (root / "missing").string(); + check(invoke(explicit_root).first != 0, "invalid explicit root never retries the default"); + auto rtx = values; + rtx["cuda_graphs"] = true; + check(invoke(rtx).first != 0, + "RTX graph flag reaches loader validation instead of being ignored"); + rtx = values; + rtx["runtime_cache"] = (root / "cache").string(); + check(invoke(rtx).first != 0, + "RTX cache flag reaches loader validation instead of being ignored"); + check(invoke(values, "unknown").first != 0, "unknown owner handler is rejected"); + bundle(path, "another_family"); + check(invoke(values).first != 0, "wrong-family bundle is rejected"); + bundle(path, "timm_xcit"); + std::ofstream(image) << "invalid image"; + const auto invalid_image = invoke(values); + check(invalid_image.first != 0 && + invalid_image.second.error.find("decode") != std::string::npos, + "invalid image fails in owner decoding before task execution"); +} +} // namespace +int main(int argc, char** argv) { + if (argc != 2) + return 2; + const auto root = fs::temp_directory_path() / ("timm_xcit-cli-" + std::to_string(getpid())); + fs::create_directories(root); + try { + contract(argv[1], root); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + ++failures; + } + fs::remove_all(root); + return failures ? 1 : 0; +} +#endif diff --git a/families/timm_xcit/tests/manifests/xcit-nano-12-p16-224-fb-in1k.json b/families/timm_xcit/tests/manifests/xcit-nano-12-p16-224-fb-in1k.json index 1ea43fe5d8..6757d4ecd5 100644 --- a/families/timm_xcit/tests/manifests/xcit-nano-12-p16-224-fb-in1k.json +++ b/families/timm_xcit/tests/manifests/xcit-nano-12-p16-224-fb-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_xcit/tests/manifests/xcit-tiny-12-p16-224-fb-in1k.json b/families/timm_xcit/tests/manifests/xcit-tiny-12-p16-224-fb-in1k.json index 3047b773bf..c3cf3d9bf8 100644 --- a/families/timm_xcit/tests/manifests/xcit-tiny-12-p16-224-fb-in1k.json +++ b/families/timm_xcit/tests/manifests/xcit-tiny-12-p16-224-fb-in1k.json @@ -13,6 +13,5 @@ "test_image": "data/test_img.jpeg" } ], - "tensor_parallel_size": 1, - "max_sequence_length": 1 + "tensor_parallel_size": 1 } diff --git a/families/timm_xcit/tests/test_cli.py b/families/timm_xcit/tests/test_cli.py new file mode 100644 index 0000000000..12549b20bb --- /dev/null +++ b/families/timm_xcit/tests/test_cli.py @@ -0,0 +1,126 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU contracts for the timm_xcit command boundary.""" + +from dataclasses import replace +import sys +from types import SimpleNamespace + +import pytest + +from families.timm_xcit import cli +from tensorrt_model_connect import BuildRequest as LegacyRequest +from tensorrt_model_connect import family_cli + +FAMILY = "timm_xcit" +TASK = "image_to_class_scores" + + +def test_build_arguments_reach_the_narrow_owner_request(monkeypatch, tmp_path): + calls = [] + monkeypatch.setattr( + cli, "build_bundle", lambda request, output: calls.append((request, output)) + ) + output = tmp_path / "model.bundle" + assert ( + family_cli.main([FAMILY, "build", str(tmp_path), "-o", str(output), "--precision", "fp16"]) + == 0 + ) + legacy = LegacyRequest(tmp_path, output, FAMILY, TASK, "fp16") + assert calls == [(cli.coerce_request(legacy), output)] + assert type(calls[0][0]) is cli.BuildRequest + assert not hasattr(calls[0][0], "image_height") + assert not hasattr(calls[0][0], "max_sequence_length") + + +def test_ignored_legacy_settings_do_not_become_new_cli_options(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + request = cli.coerce_request(legacy) + assert ( + cli.coerce_request(replace(legacy, max_sequence_length=1, quantization="none")) == request + ) + with pytest.raises(NotImplementedError, match="max_sequence_length"): + cli.coerce_request(replace(legacy, max_sequence_length=2)) + + +@pytest.mark.parametrize( + "changes", + [ + {"dynamic_kv_cache": True}, + {"image_height": 2}, + {"image_width": 2}, + {"video_num_frames": 2}, + {"max_batch_size": 2}, + {"context_parallel_size": 2}, + {"quantization": "fp8"}, + {"fp32_layers": (0,)}, + {"tensor_parallel_size": 2}, + ], +) +def test_legacy_unsupported_options_remain_rejected(tmp_path, changes): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises((ValueError, NotImplementedError)): + cli.coerce_request(replace(legacy, **changes)) + + +def test_unknown_python_inputs_are_rejected(tmp_path): + legacy = LegacyRequest(tmp_path, tmp_path / "out", FAMILY, TASK, "fp32") + with pytest.raises(ValueError, match="unknown"): + cli.coerce_request(SimpleNamespace(**vars(legacy), unknown_owner_input=1)) + + +def test_help_and_rejection_do_not_import_heavy_builder(monkeypatch, capsys): + original = family_cli.importlib.import_module + + def guarded(name, *args, **kwargs): + assert name not in { + f"families.{FAMILY}.model", + "tensorrt", + "tensorrt_rtx", + "huggingface_hub", + } + return original(name, *args, **kwargs) + + monkeypatch.setattr(family_cli.importlib, "import_module", guarded) + with pytest.raises(SystemExit) as caught: + family_cli.main([FAMILY, "build", "--help"]) + assert caught.value.code == 0 + assert "--max-sequence-length" not in capsys.readouterr().out + with pytest.raises(SystemExit) as rejected: + family_cli.main([FAMILY, "build", "checkpoint", "-o", "out", "--max-sequence-length", "1"]) + assert rejected.value.code == 2 + with pytest.raises(SystemExit) as classify: + family_cli.main([FAMILY, "classify", "--help"]) + assert classify.value.code == 0 + help_text = capsys.readouterr().out + assert "--runtime-cache" in help_text and "--cuda-graphs" in help_text + + +@pytest.mark.parametrize("fail", [False, True]) +def test_backend_selection_and_bundle_lifecycle(monkeypatch, tmp_path, fail): + events = [] + + def run(request, writer): + events.append("build") + if fail: + raise RuntimeError("owner build failed") + + def select(backend): + events.append(backend) + monkeypatch.setitem(sys.modules, f"families.{FAMILY}.model", SimpleNamespace(build=run)) + + monkeypatch.setattr(cli, "select_backend", select) + monkeypatch.setattr( + cli, + "BundleWriter", + lambda output: SimpleNamespace( + finish=lambda: events.append("finish"), abort=lambda: events.append("abort") + ), + ) + if fail: + with pytest.raises(RuntimeError, match="owner build failed"): + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + else: + cli.build_bundle(cli.BuildRequest(tmp_path, backend="trt_rtx"), tmp_path / "out") + assert events == ["trt_rtx", "build", "abort" if fail else "finish"] diff --git a/families/timm_xcit/tests/test_e2e.py b/families/timm_xcit/tests/test_e2e.py index da0e88d0cf..894ca587e3 100644 --- a/families/timm_xcit/tests/test_e2e.py +++ b/families/timm_xcit/tests/test_e2e.py @@ -14,7 +14,7 @@ import numpy as np import pytest -from tensorrt_model_connect import BuildRequest, build +from families.timm_xcit.cli import BuildRequest, build_bundle from tools.e2e_evidence import record_evidence @@ -166,23 +166,25 @@ def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") assert (runtime_root / "libtrtmc_backend_trt.so").is_file() assert (runtime_root / f"libtrtmc_model_{FAMILY}.so").is_file() + if not (runtime_root / f"libtrtmc_cli_{FAMILY}.so").is_file(): + raise AssertionError(f"selected {FAMILY} E2E requires its native CLI adapter") model_dir = _model_dir(manifest) record_evidence("checkpoint", {"hf_id": manifest["hf_id"], "hf_revision": manifest["hf_revision"]}) bundle = tmp_path / manifest["bundle"] - build( + if int(manifest["tensor_parallel_size"]) != 1: + raise NotImplementedError(f"{FAMILY} does not support tensor parallelism") + build_bundle( BuildRequest( model_dir=model_dir, - output_path=bundle, - family=FAMILY, task=manifest["task"], precision=manifest["precision"], - max_sequence_length=int(manifest["max_sequence_length"]), - tensor_parallel_size=int(manifest["tensor_parallel_size"]), - ) + ), + bundle, ) completed = subprocess.run( [ str(binary), + FAMILY, "classify", str(bundle), "--runtime-root", From f4b4679676e8ed4f8548428c17bd43a527bdde09 Mon Sep 17 00:00:00 2001 From: yifeif <277870278+yifeif-nv@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:32:39 -0700 Subject: [PATCH 2/2] fix(detr): name both supported backends Keep the unsupported-backend diagnostic consistent with the owner's request validator, which accepts trt and trt_rtx. Signed-off-by: yifeif <277870278+yifeif-nv@users.noreply.github.com> --- families/detr/cli.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/families/detr/cli.py b/families/detr/cli.py index 421c1b4992..b7ffcb1638 100644 --- a/families/detr/cli.py +++ b/families/detr/cli.py @@ -41,7 +41,7 @@ def _legacy_common_options(request: object) -> None: if request.task != "object_detection": raise ValueError("detr supports only task=object_detection") if request.backend not in {"trt", "trt_rtx"}: - raise ValueError("detr supports only backend=trt") + raise ValueError("backend must be 'trt' or 'trt_rtx'") if request.dynamic_kv_cache: raise NotImplementedError("detr does not support dynamic_kv_cache") if request.max_sequence_length not in {None, 1}: