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..b7ffcb1638 --- /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("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}: + 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",