From 60101f390340fc7f6e1dc2254113f83a17e1a46a Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Tue, 22 Sep 2026 22:00:58 -0700 Subject: [PATCH 01/15] feat(runtime): build and load the native runtime on Windows The native runtime, CLI, and family loader were ELF-only: they used dlopen/dladdr and /proc/self/exe directly, CMake passed GCC-only flags and a linker version script, and Conan packaging assumed patchelf. Multi-rank launches also required OpenMPI's mpirun. Add a model-agnostic platform layer in trtmc_core (trtmc/runtime/dynamic_library.h): LoadLibraryExW/GetProcAddress on Windows and dlopen/dlsym on ELF platforms, platform library file names, module and executable path lookup, and the shared NCCL library selection (TRTMC_NCCL_LIBRARY, else nccl.dll or libnccl.so.2). Load and missing-symbol errors name the purpose, the library, and the symbol. The family loader and the C API runtime-root lookup use it; the CLI dispatcher keeps its own small #ifdef so it stays free of trtmc/ headers. CMake gains an MSVC block: exported DLL symbols, one output directory for the executable and every DLL (the runtime root), /EHs so extern "C" plugin entry points may throw, and translation of the inline GCC warning flags. TRTMC_FAMILIES optionally restricts which families are built. On Windows, Conan provides nlohmann_json and packages the DLLs. tools/launch_ranks.py starts N local ranks with the same contract the family runtimes read under mpirun (OMPI_COMM_WORLD_* variables, one CUDA_VISIBLE_DEVICES list, a fresh TRTMC_NCCL_RENDEZVOUS file per launch) and tags rank output like mpirun --tag-output. Linux launches through mpirun are unchanged. The architecture tests accept either NCCL loader form, so each family can move to the portable loader independently. Signed-off-by: Peter Kisfaludi (cherry picked from commit 305579a832b938afdce8618602748f6cb055b4c1) --- CMakeLists.txt | 108 ++++++- apps/cli/family_cli.cpp | 116 ++++++- apps/cli/family_cli.h | 2 +- conanfile.py | 50 ++- core/api/runtime/api.cpp | 9 +- .../include/trtmc/runtime/dynamic_library.h | 78 +++++ core/runtime/loader/family_loader.cpp | 63 ++-- core/runtime/primitives/dynamic_library.cpp | 251 ++++++++++++++ core/runtime/tests/fake_partial_nccl.cpp | 22 ++ core/runtime/tests/test_dynamic_library.cpp | 150 +++++++++ tools/launch_ranks.py | 305 ++++++++++++++++++ tools/tests/test_architecture.py | 23 +- tools/tests/test_launch_ranks.py | 232 +++++++++++++ 13 files changed, 1337 insertions(+), 72 deletions(-) create mode 100644 core/runtime/include/trtmc/runtime/dynamic_library.h create mode 100644 core/runtime/primitives/dynamic_library.cpp create mode 100644 core/runtime/tests/fake_partial_nccl.cpp create mode 100644 core/runtime/tests/test_dynamic_library.cpp create mode 100644 tools/launch_ranks.py create mode 100644 tools/tests/test_launch_ranks.py diff --git a/CMakeLists.txt b/CMakeLists.txt index 92e11c3ee0..3173657c8e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -17,6 +17,32 @@ set(CMAKE_POSITION_INDEPENDENT_CODE ON) set(CMAKE_BUILD_RPATH_USE_ORIGIN TRUE) include(GNUInstallDirs) + +# MSVC (Windows) builds. Sources stay portable C++17; this block only adapts +# toolchain mechanics: DLL symbol export, runtime layout, and warning flags. +if(MSVC) + # Family, backend, and runtime DLLs are loaded at run time and must export + # their entry points (trtmc_create_family, trtmc_create_backend, ...) and + # the C++ runtime API, as ELF shared objects do by default. + set(CMAKE_WINDOWS_EXPORT_ALL_SYMBOLS ON) + add_compile_definitions(_CRT_SECURE_NO_WARNINGS NOMINMAX WIN32_LEAN_AND_MEAN _USE_MATH_DEFINES) + # /EHs without "c": the extern "C" plugin entry points (trtmc_create_family, + # trtmc_create_backend) report failures by throwing C++ exceptions to the + # loader, so MSVC must not assume extern "C" functions never throw. + string(REPLACE "/EHsc" "/EHs" CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS}") + string(REPLACE "/EHsc" "/EHs" CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS}") + add_compile_options( + "$<$:/utf-8>" + "$<$:/Zc:__cplusplus;/bigobj;/permissive->" + "$<$:-Xcompiler=/utf-8,/Zc:__cplusplus,/bigobj>" + ) + # Windows resolves a DLL's dependencies from the application directory and + # PATH (there is no RUNPATH), so executables and every DLL share one + # directory. The runtime root is that directory. + set(CMAKE_RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}") + set(CMAKE_LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}") + set(CMAKE_ARCHIVE_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/lib") +endif() find_package(CUDAToolkit REQUIRED) find_package(nlohmann_json 3.11 REQUIRED) option(TRTMC_BUILD_SERVER "Build the optional text-generation server application" ON) @@ -39,7 +65,7 @@ find_path(TRTMC_TRT_INCLUDE_DIR REQUIRED ) find_library(TRTMC_TRT_LIBRARY - NAMES nvinfer libnvinfer.so.11 + NAMES nvinfer libnvinfer.so.11 nvinfer_11 HINTS ${_trtmc_dependency_roots} PATH_SUFFIXES lib lib64 lib/aarch64-linux-gnu lib/x86_64-linux-gnu REQUIRED @@ -90,6 +116,7 @@ set(TRTMC_CUDART_LIBRARY CUDA::cudart) add_library(trtmc_core SHARED core/runtime/bundle/bundle_format.cpp core/runtime/primitives/cuda_common.cpp + core/runtime/primitives/dynamic_library.cpp core/runtime/primitives/device_tensor.cpp core/runtime/primitives/trt_common.cpp ) @@ -106,6 +133,7 @@ target_link_libraries(trtmc_core CUDA::cudart PRIVATE nlohmann_json::nlohmann_json + ${CMAKE_DL_LIBS} ) target_compile_options(trtmc_core PRIVATE -Wall -Wextra -Wpedantic) set_target_properties(trtmc_core PROPERTIES @@ -167,9 +195,11 @@ target_compile_definitions(trtmc_c PRIVATE TRTMC_VERSION_STRING="${PROJECT_VERSION}" ) target_compile_options(trtmc_c PRIVATE -Wall -Wextra -Wpedantic) -target_link_options(trtmc_c PRIVATE - "-Wl,--version-script=${PROJECT_SOURCE_DIR}/core/api/runtime/exports.map" -) +if(NOT WIN32) + target_link_options(trtmc_c PRIVATE + "-Wl,--version-script=${PROJECT_SOURCE_DIR}/core/api/runtime/exports.map" + ) +endif() set_target_properties(trtmc_c PROPERTIES EXPORT_NAME c VERSION 1 @@ -304,7 +334,10 @@ endif() # A family owns its target, sources, dependencies, warnings, and installation. # The root knows only the directory convention; adding a family never changes -# a central source or target list. +# a central source or target list. TRTMC_FAMILIES optionally restricts the +# build to a list of family directory names (default: every family). +set(TRTMC_FAMILIES "" CACHE STRING + "Semicolon-separated family names to build; empty builds every family") file(GLOB _trtmc_family_runtime_cmake CONFIGURE_DEPENDS "${PROJECT_SOURCE_DIR}/families/*/runtime/CMakeLists.txt" ) @@ -312,6 +345,9 @@ foreach(_trtmc_runtime_cmake IN LISTS _trtmc_family_runtime_cmake) get_filename_component(_trtmc_runtime_dir "${_trtmc_runtime_cmake}" DIRECTORY) get_filename_component(_trtmc_family_dir "${_trtmc_runtime_dir}" DIRECTORY) get_filename_component(_trtmc_family "${_trtmc_family_dir}" NAME) + if(TRTMC_FAMILIES AND NOT _trtmc_family IN_LIST TRTMC_FAMILIES) + continue() + endif() add_subdirectory( "${_trtmc_runtime_dir}" "${CMAKE_BINARY_DIR}/families/${_trtmc_family}" @@ -370,9 +406,11 @@ target_link_libraries(trtmc_cli ) target_compile_definitions(trtmc_cli PRIVATE TRTMC_VERSION_STRING="${PROJECT_VERSION}") target_compile_options(trtmc_cli PRIVATE -Wall -Wextra -Wpedantic) -set_source_files_properties(apps/cli/io.cpp PROPERTIES - COMPILE_OPTIONS "-Wno-missing-field-initializers;-Wno-pedantic" -) +if(NOT MSVC) + set_source_files_properties(apps/cli/io.cpp PROPERTIES + COMPILE_OPTIONS "-Wno-missing-field-initializers;-Wno-pedantic" + ) +endif() add_executable(trtmc apps/cli/main.cpp) target_include_directories(trtmc PRIVATE ${PROJECT_SOURCE_DIR}/apps) @@ -532,6 +570,20 @@ if(TRTMC_BUILD_TESTS) add_test(NAME family_cli COMMAND test_family_cli $) set_tests_properties(family_cli PROPERTIES LABELS cpu) + add_library(trtmc_test_fake_partial_nccl SHARED core/runtime/tests/fake_partial_nccl.cpp) + set_target_properties(trtmc_test_fake_partial_nccl PROPERTIES + OUTPUT_NAME trtmc_fake_partial_nccl + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/dynamic-library" + RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/tests/dynamic-library" + ) + add_executable(test_dynamic_library core/runtime/tests/test_dynamic_library.cpp) + target_link_libraries(test_dynamic_library PRIVATE trtmc_core) + target_compile_options(test_dynamic_library PRIVATE -Wall -Wextra -Wpedantic -Werror) + add_dependencies(test_dynamic_library trtmc_test_fake_partial_nccl) + add_test(NAME dynamic_library + COMMAND test_dynamic_library $) + set_tests_properties(dynamic_library PROPERTIES LABELS cpu) + set(_trtmc_test_runtime_root "${CMAKE_BINARY_DIR}/tests/runtime") add_library(trtmc_test_backend_fake SHARED core/runtime/tests/fake_backend.cpp) target_include_directories(trtmc_test_backend_fake PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) @@ -1074,6 +1126,7 @@ install(FILES ) install(FILES core/runtime/include/trtmc/runtime/device_tensor.h + core/runtime/include/trtmc/runtime/dynamic_library.h core/runtime/include/trtmc/runtime/family_factory.h core/runtime/include/trtmc/runtime/family_loader.h core/runtime/include/trtmc/runtime/tensor.h @@ -1125,3 +1178,42 @@ install(FILES DESTINATION ${CMAKE_INSTALL_DATADIR}/cmake/trtmc COMPONENT sdk ) + +# GCC/Clang warning flags are spelled inline on each target. MSVC does not +# understand them (and nvcc forwards them to cl.exe), so translate every +# target's options once, after all targets exist. +if(MSVC) + function(_trtmc_msvc_translate_warning_flags directory) + get_property(_targets DIRECTORY "${directory}" PROPERTY BUILDSYSTEM_TARGETS) + foreach(_target IN LISTS _targets) + get_target_property(_type ${_target} TYPE) + if(_type STREQUAL "INTERFACE_LIBRARY" OR _type STREQUAL "UTILITY") + continue() + endif() + get_target_property(_options ${_target} COMPILE_OPTIONS) + if(NOT _options) + continue() + endif() + set(_translated) + foreach(_option IN LISTS _options) + if(_option MATCHES "^-W" OR _option MATCHES "^-Werror") + continue() + endif() + # Generator expressions such as $<$:-Wall;-Wextra> + # arrive split on ';'; drop GCC flags inside them and keep the rest. + string(REGEX REPLACE "-W[A-Za-z0-9=_-]+" "" _option "${_option}") + if(_option MATCHES "^\\$<\\$:>?$" OR _option STREQUAL "" OR _option STREQUAL ">") + continue() + endif() + list(APPEND _translated "${_option}") + endforeach() + list(APPEND _translated "$<$:/W3>") + set_property(TARGET ${_target} PROPERTY COMPILE_OPTIONS "${_translated}") + endforeach() + get_property(_subdirectories DIRECTORY "${directory}" PROPERTY SUBDIRECTORIES) + foreach(_subdirectory IN LISTS _subdirectories) + _trtmc_msvc_translate_warning_flags("${_subdirectory}") + endforeach() + endfunction() + _trtmc_msvc_translate_warning_flags("${PROJECT_SOURCE_DIR}") +endif() diff --git a/apps/cli/family_cli.cpp b/apps/cli/family_cli.cpp index 1d24435862..e25c24b3a4 100644 --- a/apps/cli/family_cli.cpp +++ b/apps/cli/family_cli.cpp @@ -12,7 +12,6 @@ #include #include #include -#include #include #include #include @@ -21,9 +20,16 @@ #include #include #include -#include #include +#if defined(_WIN32) +#include +#include +#else +#include +#include +#endif + namespace trtmc::cli { namespace { namespace fs = std::filesystem; @@ -371,9 +377,92 @@ void write_error(void* context, const char* data, std::size_t size) { } } +// Family CLI adapters are shared libraries beside the executable: +// libtrtmc_cli_.so on ELF platforms, trtmc_cli_.dll on Windows. +std::string cli_library_name(const std::string& family) { +#if defined(_WIN32) + return "trtmc_cli_" + family + ".dll"; +#else + return "libtrtmc_cli_" + family + ".so"; +#endif +} + +class CliLibrary { + public: + explicit CliLibrary(const fs::path& path) { +#if defined(_WIN32) + handle_ = LoadLibraryExW(path.wstring().c_str(), nullptr, LOAD_WITH_ALTERED_SEARCH_PATH); + if (handle_ == nullptr) { + throw std::runtime_error("cannot load family CLI: " + path.string() + + " (Windows error " + std::to_string(GetLastError()) + ")"); + } +#else + handle_ = dlopen(path.c_str(), RTLD_NOW | RTLD_LOCAL); + if (handle_ == nullptr) + throw std::runtime_error("cannot load family CLI: " + std::string(dlerror())); +#endif + } + CliLibrary(const CliLibrary&) = delete; + CliLibrary& operator=(const CliLibrary&) = delete; + ~CliLibrary() { +#if defined(_WIN32) + FreeLibrary(handle_); +#else + dlclose(handle_); +#endif + } + + void* symbol(const char* name) const { +#if defined(_WIN32) + return reinterpret_cast(GetProcAddress(handle_, name)); +#else + dlerror(); + void* result = dlsym(handle_, name); + return dlerror() == nullptr ? result : nullptr; +#endif + } + + private: +#if defined(_WIN32) + HMODULE handle_{nullptr}; +#else + void* handle_{nullptr}; +#endif +}; + +fs::path running_executable() { +#if defined(_WIN32) + std::wstring buffer(32768, L'\0'); + const DWORD length = + GetModuleFileNameW(nullptr, buffer.data(), static_cast(buffer.size())); + if (length == 0 || length >= buffer.size()) + throw std::runtime_error("cannot resolve the trtmc executable path"); + buffer.resize(length); + return fs::path(buffer); +#else + return fs::read_symlink("/proc/self/exe"); +#endif +} + int invoke(const fs::path& executable, const std::string& family, const Json& command, const Json& values, int argc, char** argv, std::ostream& output, std::ostream& error) { if (command.at("executor") == "python") { +#if defined(_WIN32) + // Windows has no exec: run the Python command as a child process and + // return its exit status. The launcher is "python" on Windows. + std::vector arguments{"python", "-m", "tensorrt_model_connect"}; + for (int i = 1; i < argc; ++i) + arguments.push_back(argv[i]); + arguments.push_back(nullptr); + output.flush(); + error.flush(); + const intptr_t status = _spawnvp(_P_WAIT, arguments.front(), arguments.data()); + if (status == -1) { + throw std::runtime_error("cannot execute Python family command: " + + std::string(std::strerror(errno))); + } + return static_cast(status); +#else std::vector arguments{const_cast("python3"), const_cast("-m"), const_cast("tensorrt_model_connect")}; for (int i = 1; i < argc; ++i) @@ -382,11 +471,13 @@ int invoke(const fs::path& executable, const std::string& family, const Json& co execvp(arguments.front(), arguments.data()); throw std::runtime_error("cannot execute Python family command: " + std::string(std::strerror(errno))); +#endif } const auto directory = executable.parent_path(); + const auto file_name = cli_library_name(family); fs::path library; for (const auto& root : {directory, directory / "../lib", directory / "../lib64"}) { - auto candidate = root / ("libtrtmc_cli_" + family + ".so"); + auto candidate = root / file_name; if (fs::is_regular_file(candidate)) { library = fs::absolute(candidate).lexically_normal(); break; @@ -394,21 +485,14 @@ int invoke(const fs::path& executable, const std::string& family, const Json& co } if (library.empty()) throw std::runtime_error("family CLI library is not installed: " + family); - void* handle = dlopen(library.c_str(), RTLD_NOW | RTLD_LOCAL); - if (handle == nullptr) - throw std::runtime_error("cannot load family CLI: " + std::string(dlerror())); - struct Close { - void* handle; - ~Close() { dlclose(handle); } - } close{handle}; - dlerror(); - auto dispatch = reinterpret_cast(dlsym(handle, "trtmc_family_cli_v1")); - if (const char* reason = dlerror(); reason != nullptr || dispatch == nullptr) + const CliLibrary handle(library); + auto dispatch = reinterpret_cast(handle.symbol("trtmc_family_cli_v1")); + if (dispatch == nullptr) throw std::runtime_error("family does not provide trtmc_family_cli_v1: " + family); Sinks sinks{output, error}; const auto result = dispatch(command.at("handler").get().c_str(), values.dump().c_str(), - library.parent_path().c_str(), &sinks, write_output, write_error); + library.parent_path().string().c_str(), &sinks, write_output, write_error); if (sinks.failed || !output || !error) throw std::runtime_error("failed to write family CLI output"); return result; @@ -418,8 +502,8 @@ int invoke(const fs::path& executable, const std::string& family, const Json& co std::optional run_family_cli(int argc, char** argv, std::ostream& output, std::ostream& error, const fs::path& executable_override) { try { - const auto executable = executable_override.empty() ? fs::read_symlink("/proc/self/exe") - : fs::absolute(executable_override); + const auto executable = + executable_override.empty() ? running_executable() : fs::absolute(executable_override); const std::string family = argc < 2 ? "--help" : argv[1]; if (family == "--help" || family == "-h" || family == "help") { std::map declarations; diff --git a/apps/cli/family_cli.h b/apps/cli/family_cli.h index e99df6a061..39d7d0cd39 100644 --- a/apps/cli/family_cli.h +++ b/apps/cli/family_cli.h @@ -12,7 +12,7 @@ namespace trtmc::cli { // Returns no value only when the invocation does not select a declared family. -// The explicit executable path is a test seam; production resolves /proc/self/exe. +// The explicit executable path is a test seam; production resolves the running executable. std::optional run_family_cli(int argc, char** argv, std::ostream& output, std::ostream& error, const std::filesystem::path& executable = {}); diff --git a/conanfile.py b/conanfile.py index 24c715a709..37069316d9 100644 --- a/conanfile.py +++ b/conanfile.py @@ -11,7 +11,7 @@ from conan import ConanFile from conan.errors import ConanException -from conan.tools.cmake import CMake, CMakeToolchain, cmake_layout +from conan.tools.cmake import CMake, CMakeDeps, CMakeToolchain, cmake_layout from conan.tools.files import copy @@ -51,21 +51,36 @@ class TensorRTModelConnectConan(ConanFile): settings = "os", "compiler", "build_type", "arch" + def _windows(self) -> bool: + return str(self.settings.os) == "Windows" + def layout(self) -> None: cmake_layout(self) # CMakeToolchain derives install directories from the package layout. self.cpp.package.libdirs = ["bin"] + def requirements(self) -> None: + # Linux images provide nlohmann-json3-dev; MSVC builds take it from Conan. + if self._windows(): + self.requires("nlohmann_json/3.11.3") + def generate(self) -> None: toolchain = CMakeToolchain(self) toolchain.cache_variables["TRTMC_BUILD_TESTS"] = False for name in ( "TRT_ROOT", "CMAKE_CUDA_ARCHITECTURES", + "TRTMC_FAMILIES", ): value = os.environ.get(name) if value: toolchain.cache_variables[name] = value + if self._windows(): + # The Windows port covers the native runtime, CLI, and model + # families; the server, BYOK bridge, and examples stay ELF-only. + for option in ("TRTMC_BUILD_SERVER", "TRTMC_ENABLE_BYOK", "TRTMC_BUILD_EXAMPLES"): + toolchain.cache_variables[option] = False + CMakeDeps(self).generate() toolchain.generate() def build(self) -> None: @@ -73,7 +88,40 @@ def build(self) -> None: cmake.configure() cmake.build() + def _package_windows(self) -> None: + source = Path(self.source_folder) + build = Path(self.build_folder) + module_bin = Path(self.package_folder) / "tensorrt_model_connect" / "bin" + # Windows has no RUNPATH: the executable, runtime DLLs, backend, and + # family DLLs share one directory, which is also the runtime root. + copy(self, "trtmc.exe", src=str(build), dst=str(module_bin), keep_path=False) + copy(self, "*.dll", src=str(build), dst=str(module_bin), keep_path=False) + selected = [name for name in os.environ.get("TRTMC_FAMILIES", "").split(";") if name] + expected = set(selected) or { + path.parent.name for path in (source / "families").glob("*/model.py") + } + packaged = { + path.stem.removeprefix("trtmc_model_") for path in module_bin.glob("trtmc_model_*.dll") + } + required = ( + "trtmc.exe", + "trtmc_core.dll", + "trtmc_runtime.dll", + "trtmc_c.dll", + "trtmc_backend_trt.dll", + ) + if not all((module_bin / name).is_file() for name in required): + raise ConanException("native Windows runtime package is incomplete") + if packaged != expected: + raise ConanException( + f"family DLL set does not match: missing={sorted(expected - packaged)}, " + f"extra={sorted(packaged - expected)}" + ) + def package(self) -> None: + if self._windows(): + self._package_windows() + return source = Path(self.source_folder) build = Path(self.build_folder) package = Path(self.package_folder) diff --git a/core/api/runtime/api.cpp b/core/api/runtime/api.cpp index e1d824d18c..1e13d1a6cc 100644 --- a/core/api/runtime/api.cpp +++ b/core/api/runtime/api.cpp @@ -5,11 +5,11 @@ #include "api_internal.h" #include "trtmc/bundle.h" +#include "trtmc/runtime/dynamic_library.h" #include "trtmc/runtime/family_loader.h" #include #include -#include #include #include #include @@ -252,10 +252,11 @@ ConvertedConfig::ConvertedConfig(const trtmc_config_view_v1* config) { std::string default_runtime_root() { static const unsigned char library_location = 0; - Dl_info info{}; - if (dladdr(&library_location, &info) == 0 || info.dli_fname == nullptr) + try { + return platform::module_path_containing(&library_location).parent_path().string(); + } catch (const std::exception&) { throw ApiFailure{TRTMC_INTERNAL_ERROR, "cannot locate the runtime library directory"}; - return std::filesystem::absolute(info.dli_fname).parent_path().string(); + } } struct TaskSnapshot { diff --git a/core/runtime/include/trtmc/runtime/dynamic_library.h b/core/runtime/include/trtmc/runtime/dynamic_library.h new file mode 100644 index 0000000000..d029d65cb6 --- /dev/null +++ b/core/runtime/include/trtmc/runtime/dynamic_library.h @@ -0,0 +1,78 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +// Model-agnostic platform mechanics for runtime-loaded shared libraries: +// dlopen/dlsym on ELF platforms, LoadLibraryExW/GetProcAddress on Windows. + +#include +#include +#include + +namespace trtmc::platform { + +// Platform file name of a shared library built from CMake target output name +// `stem`: "lib.so" on ELF platforms and ".dll" on Windows. +std::string shared_library_filename(std::string_view stem); + +// An owned handle to a loaded shared library. Construction throws +// std::runtime_error with `purpose`, the requested library, and the platform +// loader error when the library cannot be loaded. +// +// `name_or_path` is either a bare file name, resolved through the platform +// search order (LD_LIBRARY_PATH / PATH), or a path. On Windows, a path also +// makes the loader search that library's directory for its own dependencies. +class DynamicLibrary { + public: + DynamicLibrary(const std::string& name_or_path, std::string purpose); + ~DynamicLibrary(); + + DynamicLibrary(const DynamicLibrary&) = delete; + DynamicLibrary& operator=(const DynamicLibrary&) = delete; + + // nullptr when the symbol is absent. + void* find_symbol(const char* name) const noexcept; + + // Throws std::runtime_error naming the purpose, library, and symbol when + // the symbol is absent. + void* require_symbol(const char* name) const; + + template + Function require(const char* name) const { + return reinterpret_cast(require_symbol(name)); + } + + // The name or path passed to the constructor. + const std::string& name() const noexcept { return name_; } + + // The file the platform loader actually mapped, or name() when the + // platform cannot report it. + std::string loaded_path() const; + + private: + std::string name_; + std::string purpose_; + void* handle_{nullptr}; +}; + +// Absolute path of the executable or shared library that contains `address`. +// Throws std::runtime_error when the platform cannot resolve it. +std::filesystem::path module_path_containing(const void* address); + +// Absolute path of the running executable. +std::filesystem::path current_executable_path(); + +// Environment variable that overrides the NCCL shared library for every +// runtime that loads NCCL at run time. +inline constexpr const char* kNcclLibraryEnv = "TRTMC_NCCL_LIBRARY"; + +// "nccl.dll" on Windows, "libnccl.so.2" on ELF platforms. +const char* default_nccl_library(); + +// TRTMC_NCCL_LIBRARY when set and non-empty, otherwise default_nccl_library(). +std::string nccl_library(); + +} // namespace trtmc::platform diff --git a/core/runtime/loader/family_loader.cpp b/core/runtime/loader/family_loader.cpp index e1a2ca6ec6..4892055780 100644 --- a/core/runtime/loader/family_loader.cpp +++ b/core/runtime/loader/family_loader.cpp @@ -6,10 +6,10 @@ #include "trtmc/runtime/family_loader.h" #include "runtime/bundle/bundle_format.h" +#include "trtmc/runtime/dynamic_library.h" #include "trtmc/runtime/family_factory.h" #include "trtmc/runtime/trt_backend.h" -#include #include #include #include @@ -62,13 +62,14 @@ fs::path resolve_runtime_root(const std::string& runtime_root) { return absolute_path(runtime_root, "runtime_root"); static const char runtime_library_anchor = 0; - Dl_info info{}; - if (dladdr(&runtime_library_anchor, &info) == 0 || info.dli_fname == nullptr || - info.dli_fname[0] == '\0') { - throw std::runtime_error("Unable to locate the loaded libtrtmc_runtime shared library; " - "specify runtime_root explicitly"); + fs::path library; + try { + library = platform::module_path_containing(&runtime_library_anchor); + } catch (const std::exception& error) { + throw std::runtime_error(std::string("Unable to locate the loaded trtmc_runtime shared " + "library; specify runtime_root explicitly: ") + + error.what()); } - const fs::path library = absolute_path(info.dli_fname, "loaded runtime library"); if (library.parent_path().empty()) throw std::runtime_error("Loaded runtime library has no parent directory: '" + library.string() + "'"); @@ -77,44 +78,29 @@ fs::path resolve_runtime_root(const std::string& runtime_root) { class SharedLibrary { public: - explicit SharedLibrary(const fs::path& path) : path_(path.string()) { - dlerror(); - handle_ = dlopen(path_.c_str(), RTLD_NOW | RTLD_LOCAL); - if (handle_ == nullptr) { - const char* error = dlerror(); - throw std::runtime_error("Unable to load '" + path_ + - "': " + (error != nullptr ? error : "unknown dlopen error")); - } - } + explicit SharedLibrary(const fs::path& path) : library_(path.string(), "trtmc runtime") {} SharedLibrary(const SharedLibrary&) = delete; SharedLibrary& operator=(const SharedLibrary&) = delete; - ~SharedLibrary() { - if (handle_ != nullptr) - dlclose(handle_); - } - - void* require_symbol(const char* name) const { - dlerror(); - void* symbol = dlsym(handle_, name); - const char* error = dlerror(); - if (error != nullptr || symbol == nullptr) { - throw std::runtime_error("Library '" + path_ + "' is missing required symbol '" + name + - "'"); - } - return symbol; - } + void* require_symbol(const char* name) const { return library_.require_symbol(name); } private: - std::string path_; - void* handle_{nullptr}; + platform::DynamicLibrary library_; }; +fs::path backend_library_path(const fs::path& runtime_root, const std::string& backend_id) { + return runtime_root / platform::shared_library_filename("trtmc_backend_" + backend_id); +} + +fs::path family_library_path(const fs::path& runtime_root, const std::string& family_id) { + return runtime_root / platform::shared_library_filename("trtmc_model_" + family_id); +} + class BackendLibrary { public: BackendLibrary(const fs::path& runtime_root, const std::string& backend_id) - : library_(runtime_root / ("libtrtmc_backend_" + backend_id + ".so")) { + : library_(backend_library_path(runtime_root, backend_id)) { const auto create = reinterpret_cast(library_.require_symbol("trtmc_create_backend")); destroy_ = @@ -152,7 +138,7 @@ class BackendLibrary { class FamilyLibrary { public: FamilyLibrary(const fs::path& runtime_root, const std::string& family_id) - : library_(runtime_root / ("libtrtmc_model_" + family_id + ".so")), + : library_(family_library_path(runtime_root, family_id)), create_(reinterpret_cast(library_.require_symbol(kCreateFamilySymbol))) {} FamilyLibrary(const FamilyLibrary&) = delete; @@ -245,7 +231,7 @@ RuntimeLibraryCache& runtime_library_cache() { } IBackend& cached_backend(const fs::path& runtime_root, const std::string& backend_id) { - const std::string path = (runtime_root / ("libtrtmc_backend_" + backend_id + ".so")).string(); + const std::string path = backend_library_path(runtime_root, backend_id).string(); auto& cache = runtime_library_cache(); std::lock_guard lock(cache.mutex); const auto found = cache.backends.find(path); @@ -275,7 +261,7 @@ IBackend& cached_configured_backend(IBackend& backend, const std::string& runtim } FamilyLibrary& cached_family(const fs::path& runtime_root, const std::string& family_id) { - const std::string path = (runtime_root / ("libtrtmc_model_" + family_id + ".so")).string(); + const std::string path = family_library_path(runtime_root, family_id).string(); auto& cache = runtime_library_cache(); std::lock_guard lock(cache.mutex); const auto found = cache.families.find(path); @@ -305,7 +291,8 @@ void preload_byok_kernel(const std::string& runtime_root, const std::string& lib if (library.empty() || function.empty() || kernel_name.empty()) throw std::invalid_argument("BYOK library, function, and kernel name must be non-empty"); using LoadKernel = const char* (*)(const char*, const char*, const char*) noexcept; - const auto path = resolve_runtime_root(runtime_root) / "libtrtmc_byok_tvm_ffi.so"; + const auto path = resolve_runtime_root(runtime_root) / + platform::shared_library_filename("trtmc_byok_tvm_ffi"); LoadKernel load; { auto& cache = runtime_library_cache(); diff --git a/core/runtime/primitives/dynamic_library.cpp b/core/runtime/primitives/dynamic_library.cpp new file mode 100644 index 0000000000..a50cee4207 --- /dev/null +++ b/core/runtime/primitives/dynamic_library.cpp @@ -0,0 +1,251 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/runtime/dynamic_library.h" + +#include +#include +#include +#include + +#if defined(_WIN32) +#ifndef WIN32_LEAN_AND_MEAN +#define WIN32_LEAN_AND_MEAN +#endif +#ifndef NOMINMAX +#define NOMINMAX +#endif +#include +#else +#include +#if defined(__GLIBC__) +#include +#endif +#endif + +namespace trtmc::platform { +namespace { + +#if defined(_WIN32) +std::wstring widen(const std::string& value) { + if (value.empty()) + return {}; + const int size = + MultiByteToWideChar(CP_UTF8, 0, value.data(), static_cast(value.size()), nullptr, 0); + if (size <= 0) + throw std::runtime_error("Library name is not valid UTF-8: " + value); + std::wstring result(static_cast(size), L'\0'); + MultiByteToWideChar(CP_UTF8, 0, value.data(), static_cast(value.size()), result.data(), + size); + return result; +} + +std::string narrow(const std::wstring& value) { + if (value.empty()) + return {}; + const int size = WideCharToMultiByte(CP_UTF8, 0, value.data(), static_cast(value.size()), + nullptr, 0, nullptr, nullptr); + if (size <= 0) + return {}; + std::string result(static_cast(size), '\0'); + WideCharToMultiByte(CP_UTF8, 0, value.data(), static_cast(value.size()), result.data(), + size, nullptr, nullptr); + return result; +} + +// The system text for a Windows error code, without the trailing period and line break. +std::string windows_system_message(DWORD code) { + LPWSTR buffer = nullptr; + const DWORD length = FormatMessageW( + FORMAT_MESSAGE_ALLOCATE_BUFFER | FORMAT_MESSAGE_FROM_SYSTEM | FORMAT_MESSAGE_IGNORE_INSERTS, + nullptr, code, MAKELANGID(LANG_NEUTRAL, SUBLANG_DEFAULT), reinterpret_cast(&buffer), + 0, nullptr); + std::string message; + if (length != 0 && buffer != nullptr) + message = narrow(std::wstring(buffer, length)); + if (buffer != nullptr) + LocalFree(buffer); + message.erase(message.find_last_not_of("\r\n .") + 1); + return message.empty() ? std::string("unknown error") : message; +} + +// What usually causes the loader errors that a missing or mismatched DLL produces. +const char* windows_loader_hint(DWORD code) { + switch (code) { + case ERROR_MOD_NOT_FOUND: + return "; the library or one of its dependent DLLs was not found on the DLL search " + "path (application directory, the library's directory for a path, PATH)"; + case ERROR_PROC_NOT_FOUND: + return "; a dependent DLL is missing an imported function (version mismatch)"; + case ERROR_BAD_EXE_FORMAT: + return "; the file is not a 64-bit Windows DLL"; + default: + return ""; + } +} + +std::string windows_error_message(DWORD code) { + return windows_system_message(code) + " (Windows error " + std::to_string(code) + ")" + + windows_loader_hint(code); +} + +std::filesystem::path module_file_name(HMODULE module) { + std::wstring buffer(MAX_PATH, L'\0'); + for (;;) { + const DWORD length = + GetModuleFileNameW(module, buffer.data(), static_cast(buffer.size())); + if (length == 0) { + throw std::runtime_error("GetModuleFileNameW failed: " + + windows_error_message(GetLastError())); + } + if (length < buffer.size()) { + buffer.resize(length); + return std::filesystem::path(buffer); + } + buffer.resize(buffer.size() * 2); + } +} + +bool has_directory(const std::string& value) { + return value.find_first_of("\\/") != std::string::npos; +} +#endif + +} // namespace + +std::string shared_library_filename(std::string_view stem) { +#if defined(_WIN32) + return std::string(stem) + ".dll"; +#else + return "lib" + std::string(stem) + ".so"; +#endif +} + +DynamicLibrary::DynamicLibrary(const std::string& name_or_path, std::string purpose) + : name_(name_or_path), purpose_(std::move(purpose)) { + if (name_.empty()) + throw std::runtime_error(purpose_ + ": empty shared library name"); +#if defined(_WIN32) + const std::wstring wide = widen(name_); + // A path loads exactly that file and searches its directory for the + // library's own dependencies; a bare name uses the standard DLL search + // order (application directory, system directories, PATH). + const DWORD flags = has_directory(name_) ? LOAD_WITH_ALTERED_SEARCH_PATH : 0; + // Report a missing DLL as an exception instead of a modal dialog. + const UINT previous_mode = SetErrorMode(SEM_FAILCRITICALERRORS | SEM_NOOPENFILEERRORBOX); + HMODULE module = LoadLibraryExW(wide.c_str(), nullptr, flags); + const DWORD error = module == nullptr ? GetLastError() : ERROR_SUCCESS; + SetErrorMode(previous_mode); + if (module == nullptr) { + throw std::runtime_error(purpose_ + ": unable to load '" + name_ + + "': " + windows_error_message(error)); + } + handle_ = module; +#else + dlerror(); + handle_ = dlopen(name_.c_str(), RTLD_NOW | RTLD_LOCAL); + if (handle_ == nullptr) { + const char* error = dlerror(); + throw std::runtime_error(purpose_ + ": unable to load '" + name_ + + "': " + (error != nullptr ? error : "unknown dlopen error")); + } +#endif +} + +DynamicLibrary::~DynamicLibrary() { + if (handle_ == nullptr) + return; +#if defined(_WIN32) + FreeLibrary(static_cast(handle_)); +#else + dlclose(handle_); +#endif +} + +void* DynamicLibrary::find_symbol(const char* name) const noexcept { + if (handle_ == nullptr || name == nullptr) + return nullptr; +#if defined(_WIN32) + return reinterpret_cast(GetProcAddress(static_cast(handle_), name)); +#else + dlerror(); + void* symbol = dlsym(handle_, name); + if (dlerror() != nullptr) + return nullptr; + return symbol; +#endif +} + +void* DynamicLibrary::require_symbol(const char* name) const { + void* symbol = find_symbol(name); + if (symbol == nullptr) { + throw std::runtime_error(purpose_ + ": library '" + loaded_path() + + "' is missing required symbol '" + + (name != nullptr ? name : "") + "'"); + } + return symbol; +} + +std::string DynamicLibrary::loaded_path() const { +#if defined(_WIN32) + try { + return narrow(module_file_name(static_cast(handle_)).wstring()); + } catch (const std::exception&) { + return name_; + } +#elif defined(__GLIBC__) + struct link_map* map = nullptr; + if (dlinfo(handle_, RTLD_DI_LINKMAP, &map) == 0 && map != nullptr && map->l_name != nullptr && + map->l_name[0] != '\0') + return map->l_name; + return name_; +#else + return name_; +#endif +} + +std::filesystem::path module_path_containing(const void* address) { +#if defined(_WIN32) + HMODULE module = nullptr; + if (!GetModuleHandleExW(GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS | + GET_MODULE_HANDLE_EX_FLAG_UNCHANGED_REFCOUNT, + static_cast(address), &module) || + module == nullptr) { + throw std::runtime_error("Unable to locate the module containing an address: " + + windows_error_message(GetLastError())); + } + return std::filesystem::absolute(module_file_name(module)).lexically_normal(); +#else + Dl_info info{}; + if (dladdr(address, &info) == 0 || info.dli_fname == nullptr || info.dli_fname[0] == '\0') + throw std::runtime_error("Unable to locate the shared library containing an address"); + return std::filesystem::absolute(info.dli_fname).lexically_normal(); +#endif +} + +std::filesystem::path current_executable_path() { +#if defined(_WIN32) + return module_file_name(nullptr); +#else + return std::filesystem::read_symlink("/proc/self/exe"); +#endif +} + +const char* default_nccl_library() { +#if defined(_WIN32) + return "nccl.dll"; +#else + return "libnccl.so.2"; +#endif +} + +std::string nccl_library() { + const char* configured = std::getenv(kNcclLibraryEnv); + if (configured != nullptr && *configured != '\0') + return configured; + return default_nccl_library(); +} + +} // namespace trtmc::platform diff --git a/core/runtime/tests/fake_partial_nccl.cpp b/core/runtime/tests/fake_partial_nccl.cpp new file mode 100644 index 0000000000..99d94258c9 --- /dev/null +++ b/core/runtime/tests/fake_partial_nccl.cpp @@ -0,0 +1,22 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +// A shared library that exports only part of the NCCL entry points the +// family runtimes resolve, to test the missing-symbol diagnostics. + +#if defined(_WIN32) +#define TRTMC_FAKE_EXPORT __declspec(dllexport) +#else +#define TRTMC_FAKE_EXPORT __attribute__((visibility("default"))) +#endif + +extern "C" TRTMC_FAKE_EXPORT int ncclGetVersion(int* version) { + *version = 23007; + return 0; +} + +extern "C" TRTMC_FAKE_EXPORT int ncclGetUniqueId(void* id) { + return id == nullptr ? 4 : 0; +} diff --git a/core/runtime/tests/test_dynamic_library.cpp b/core/runtime/tests/test_dynamic_library.cpp new file mode 100644 index 0000000000..94314a044c --- /dev/null +++ b/core/runtime/tests/test_dynamic_library.cpp @@ -0,0 +1,150 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "trtmc/runtime/dynamic_library.h" + +#include +#include +#include +#include +#include + +namespace { + +namespace fs = std::filesystem; +using trtmc::platform::DynamicLibrary; + +int failures = 0; + +void check(bool condition, const std::string& name) { + if (!condition) { + std::cerr << "FAIL: " << name << '\n'; + ++failures; + } +} + +bool contains(const std::string& text, const std::string& part) { + return text.find(part) != std::string::npos; +} + +void set_env(const char* name, const char* value) { +#if defined(_WIN32) + _putenv_s(name, value == nullptr ? "" : value); +#else + if (value == nullptr) + unsetenv(name); + else + setenv(name, value, 1); +#endif +} + +template +std::string error_of(Function&& function) { + try { + function(); + } catch (const std::runtime_error& error) { + return error.what(); + } + return {}; +} + +using GetVersionFn = int (*)(int*); + +void test_file_names() { +#if defined(_WIN32) + check(trtmc::platform::shared_library_filename("trtmc_model_flux") == "trtmc_model_flux.dll", + "windows family file name"); + check(std::string(trtmc::platform::default_nccl_library()) == "nccl.dll", "windows NCCL name"); +#else + check(trtmc::platform::shared_library_filename("trtmc_model_flux") == "libtrtmc_model_flux.so", + "ELF family file name"); + check(std::string(trtmc::platform::default_nccl_library()) == "libnccl.so.2", "ELF NCCL name"); +#endif +} + +void test_nccl_library_override() { + set_env(trtmc::platform::kNcclLibraryEnv, nullptr); + check(trtmc::platform::nccl_library() == trtmc::platform::default_nccl_library(), + "NCCL default without override"); + set_env(trtmc::platform::kNcclLibraryEnv, ""); + check(trtmc::platform::nccl_library() == trtmc::platform::default_nccl_library(), + "empty override keeps the default"); + set_env(trtmc::platform::kNcclLibraryEnv, "/custom/nccl-build/nccl.dll"); + check(trtmc::platform::nccl_library() == "/custom/nccl-build/nccl.dll", + "TRTMC_NCCL_LIBRARY overrides the NCCL library"); + set_env(trtmc::platform::kNcclLibraryEnv, nullptr); +} + +void test_missing_library(const fs::path& directory) { + const auto missing = directory / trtmc::platform::shared_library_filename("no_such_nccl"); + const auto message = + error_of([&] { DynamicLibrary library(missing.string(), "Unit test: NCCL"); }); + check(contains(message, "Unit test: NCCL"), "load error names the purpose: " + message); + check(contains(message, "unable to load"), "load error says it cannot load: " + message); + check(contains(message, missing.string()), "load error names the library: " + message); + + const auto bare = error_of([] { DynamicLibrary library("trtmc_no_such_library_xyz", "x"); }); + check(contains(bare, "trtmc_no_such_library_xyz"), "bare-name load error: " + bare); + + const auto empty = error_of([] { DynamicLibrary library("", "empty"); }); + check(contains(empty, "empty shared library name"), "empty name error: " + empty); +} + +void test_partial_library(const fs::path& partial) { + DynamicLibrary library(partial.string(), "Unit test: NCCL"); + check(library.name() == partial.string(), "name is the requested path"); + check(fs::equivalent(fs::path(library.loaded_path()), partial), + "loaded_path is the mapped file: " + library.loaded_path()); + + const auto get_version = library.require("ncclGetVersion"); + int version = 0; + check(get_version(&version) == 0 && version == 23007, "resolved symbol is callable"); + check(library.find_symbol("ncclGetUniqueId") != nullptr, "find_symbol finds exports"); + check(library.find_symbol("ncclCommInitRank") == nullptr, "find_symbol returns null"); + check(library.find_symbol(nullptr) == nullptr, "find_symbol(nullptr) is null"); + + const auto message = error_of([&] { (void)library.require_symbol("ncclCommInitRank"); }); + check(contains(message, "Unit test: NCCL"), "symbol error names the purpose: " + message); + check(contains(message, "missing required symbol 'ncclCommInitRank'"), + "symbol error names the symbol: " + message); + check(contains(message, partial.filename().string()), + "symbol error names the library: " + message); +} + +void test_module_paths(const char* argv0) { + const auto executable = trtmc::platform::current_executable_path(); + check(executable.is_absolute(), "executable path is absolute: " + executable.string()); + check(executable.stem() == fs::path(argv0).stem(), + "current_executable_path is this test: " + executable.string()); + const auto containing = trtmc::platform::module_path_containing( + reinterpret_cast(&trtmc::platform::nccl_library)); + check(containing.filename() == trtmc::platform::shared_library_filename("trtmc_core"), + "module_path_containing finds trtmc_core: " + containing.string()); + check(containing.is_absolute(), "module path is absolute"); +} + +} // namespace + +int main(int argc, char** argv) { + if (argc != 2) { + std::cerr << "usage: test_dynamic_library \n"; + return 2; + } + const fs::path partial = fs::absolute(argv[1]); + try { + test_file_names(); + test_nccl_library_override(); + test_missing_library(partial.parent_path()); + test_partial_library(partial); + test_module_paths(argv[0]); + } catch (const std::exception& error) { + std::cerr << "FAIL: unexpected exception: " << error.what() << '\n'; + return 1; + } + if (failures != 0) + return 1; + std::cout << "dynamic_library: all checks passed\n"; + return 0; +} diff --git a/tools/launch_ranks.py b/tools/launch_ranks.py new file mode 100644 index 0000000000..2a5037fff2 --- /dev/null +++ b/tools/launch_ranks.py @@ -0,0 +1,305 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Start N local ranks of a TRTMC command without MPI (Linux and Windows). + +The distributed family runtimes discover their rank from the OpenMPI +environment contract (``OMPI_COMM_WORLD_SIZE``/``_RANK``/``_LOCAL_RANK``) and +exchange the NCCL unique id through the file named by +``TRTMC_NCCL_RENDEZVOUS``. On Linux, ``mpirun`` provides the rank variables. +This launcher provides the same contract on one machine where OpenMPI is not +available, such as native Windows: + + python tools/launch_ranks.py -n 2 --gpus 0,1 -- trtmc generate-video BUNDLE ... + +Every rank sees the same ``CUDA_VISIBLE_DEVICES`` list and selects its device +by local rank, exactly as under ``mpirun``. Each launch uses a fresh +rendezvous file, so a stale unique id from an earlier run cannot be read. +Rank output is prefixed like ``mpirun --tag-output`` (``[1,]:``) +so existing rank-0 output parsers work unchanged. When one rank fails, the +remaining ranks are terminated and the launcher returns the first non-zero +exit status. +""" + +from __future__ import annotations + +import argparse +import os +import signal +import subprocess +import sys +import tempfile +import threading +import time +import uuid +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from pathlib import Path +from typing import TextIO + +RANK_ENV_NAMES = ( + "OMPI_COMM_WORLD_SIZE", + "OMPI_COMM_WORLD_RANK", + "OMPI_COMM_WORLD_LOCAL_RANK", + "OMPI_COMM_WORLD_LOCAL_SIZE", +) +RENDEZVOUS_ENV = "TRTMC_NCCL_RENDEZVOUS" +NCCL_LIBRARY_ENV = "TRTMC_NCCL_LIBRARY" + + +def tag(rank: int, stream: str) -> str: + """Return the ``mpirun --tag-output`` prefix for one rank and stream.""" + return f"[1,{rank}]<{stream}>:" + + +def library_path_variable(platform: str = sys.platform) -> str: + """Environment variable the platform loader searches for shared libraries.""" + return "PATH" if platform.startswith("win") else "LD_LIBRARY_PATH" + + +def parse_gpus(value: str | None, world_size: int) -> list[str] | None: + """Validate a comma-separated GPU list; None keeps the caller's visibility.""" + if value is None or value == "": + return None + gpus = [item.strip() for item in value.split(",")] + if any(not item for item in gpus): + raise ValueError(f"invalid --gpus list: {value!r}") + if len(set(gpus)) != len(gpus): + raise ValueError(f"--gpus lists a device twice: {value!r}") + if len(gpus) < world_size: + raise ValueError(f"--gpus lists {len(gpus)} device(s) for {world_size} ranks") + return gpus + + +def new_rendezvous_path(directory: str | os.PathLike[str] | None = None) -> Path: + """Unique, not-yet-existing rendezvous file for one launch.""" + root = Path(directory) if directory is not None else Path(tempfile.gettempdir()) + return root / f"trtmc_nccl_{os.getpid()}_{uuid.uuid4().hex}.bin" + + +def rank_environments( + base: Mapping[str, str], + world_size: int, + rendezvous: Path, + gpus: Sequence[str] | None = None, + nccl_library: str | None = None, + library_dirs: Sequence[str] = (), + platform: str = sys.platform, +) -> list[dict[str, str]]: + """Build the environment of every rank from ``base`` (not modified).""" + if world_size < 1: + raise ValueError("world size must be at least 1") + shared = dict(base) + if gpus is not None: + shared["CUDA_VISIBLE_DEVICES"] = ",".join(gpus) + shared[RENDEZVOUS_ENV] = str(rendezvous) + if nccl_library: + shared[NCCL_LIBRARY_ENV] = nccl_library + if library_dirs: + variable = library_path_variable(platform) + # Windows environment names are case-insensitive; reuse the existing key. + key = next((name for name in shared if name.upper() == variable.upper()), variable) + current = shared.get(key, "") + separator = ";" if platform.startswith("win") else ":" + shared[key] = separator.join([*library_dirs, *([current] if current else [])]) + environments = [] + for rank in range(world_size): + env = dict(shared) + env["OMPI_COMM_WORLD_SIZE"] = str(world_size) + env["OMPI_COMM_WORLD_RANK"] = str(rank) + env["OMPI_COMM_WORLD_LOCAL_RANK"] = str(rank) + env["OMPI_COMM_WORLD_LOCAL_SIZE"] = str(world_size) + environments.append(env) + return environments + + +@dataclass +class _Rank: + rank: int + process: subprocess.Popen + threads: list[threading.Thread] = field(default_factory=list) + + +def _pump(source, rank: int, stream: str, sink: TextIO, lock: threading.Lock, tagged: bool): + prefix = tag(rank, stream) if tagged else "" + for raw in iter(source.readline, b""): + line = raw.decode("utf-8", errors="replace").rstrip("\r\n") + with lock: + sink.write(f"{prefix}{line}\n") + sink.flush() + source.close() + + +def _terminate(ranks: Sequence[_Rank], grace_s: float) -> None: + for item in ranks: + if item.process.poll() is None: + item.process.terminate() + deadline = time.monotonic() + grace_s + for item in ranks: + remaining = max(0.0, deadline - time.monotonic()) + try: + item.process.wait(timeout=remaining) + except subprocess.TimeoutExpired: + item.process.kill() + item.process.wait() + + +def launch( + command: Sequence[str], + world_size: int, + *, + gpus: Sequence[str] | None = None, + rendezvous_dir: str | os.PathLike[str] | None = None, + nccl_library: str | None = None, + library_dirs: Sequence[str] = (), + tagged: bool = True, + timeout_s: float | None = None, + grace_s: float = 10.0, + base_env: Mapping[str, str] | None = None, + stdout: TextIO | None = None, + stderr: TextIO | None = None, +) -> int: + """Run ``command`` as ``world_size`` ranks; return the launch exit status.""" + if not command: + raise ValueError("no command to launch") + stdout = stdout if stdout is not None else sys.stdout + stderr = stderr if stderr is not None else sys.stderr + rendezvous = new_rendezvous_path(rendezvous_dir) + rendezvous.parent.mkdir(parents=True, exist_ok=True) + environments = rank_environments( + os.environ if base_env is None else base_env, + world_size, + rendezvous, + gpus=gpus, + nccl_library=nccl_library, + library_dirs=library_dirs, + ) + lock = threading.Lock() + ranks: list[_Rank] = [] + status = 0 + try: + for rank, env in enumerate(environments): + process = subprocess.Popen( + list(command), + env=env, + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + item = _Rank(rank, process) + for source, name, sink in ( + (process.stdout, "stdout", stdout), + (process.stderr, "stderr", stderr), + ): + thread = threading.Thread( + target=_pump, args=(source, rank, name, sink, lock, tagged), daemon=True + ) + thread.start() + item.threads.append(thread) + ranks.append(item) + + deadline = None if timeout_s is None else time.monotonic() + timeout_s + pending = list(ranks) + while pending: + for item in list(pending): + code = item.process.poll() + if code is None: + continue + pending.remove(item) + if code != 0 and status == 0: + status = code + with lock: + stderr.write( + f"[launch_ranks] rank {item.rank} exited with status {code}; " + "terminating the remaining ranks\n" + ) + stderr.flush() + _terminate(pending, grace_s) + if deadline is not None and pending and time.monotonic() > deadline: + with lock: + stderr.write(f"[launch_ranks] timed out after {timeout_s:g} s\n") + stderr.flush() + _terminate(pending, grace_s) + status = status or 124 + break + if pending: + time.sleep(0.05) + except KeyboardInterrupt: + _terminate(ranks, grace_s) + status = status or 130 + finally: + _terminate(ranks, grace_s) + for item in ranks: + for thread in item.threads: + thread.join() + for leftover in (rendezvous, rendezvous.with_name(rendezvous.name + ".tmp")): + try: + leftover.unlink() + except FileNotFoundError: + pass + except OSError: + pass + return status + + +def _interrupt(_signum, _frame) -> None: + raise KeyboardInterrupt + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description=__doc__.split("\n\n", 1)[0], + epilog="Everything after '--' is the command each rank runs.", + ) + parser.add_argument("-n", "--np", dest="world_size", type=int, required=True) + parser.add_argument( + "--gpus", help="comma-separated CUDA devices; sets CUDA_VISIBLE_DEVICES for all ranks" + ) + parser.add_argument( + "--rendezvous-dir", help="directory for the per-launch NCCL rendezvous file (default: temp)" + ) + parser.add_argument( + "--nccl-library", help=f"NCCL shared library for the ranks (sets {NCCL_LIBRARY_ENV})" + ) + parser.add_argument( + "--library-dir", + action="append", + default=[], + help="prepend a directory to the shared-library search path (PATH on Windows, " + "LD_LIBRARY_PATH elsewhere); repeatable", + ) + parser.add_argument("--no-tag-output", action="store_true", help="do not prefix rank output") + parser.add_argument("--timeout", type=float, help="terminate all ranks after this many seconds") + parser.add_argument("command", nargs=argparse.REMAINDER) + return parser + + +def main(argv: Sequence[str] | None = None) -> int: + args = _parser().parse_args(argv) + command = list(args.command) + if command and command[0] == "--": + command = command[1:] + if not command: + _parser().error("missing command after '--'") + if args.world_size < 1: + _parser().error("-n must be at least 1") + try: + gpus = parse_gpus(args.gpus, args.world_size) + except ValueError as error: + _parser().error(str(error)) + if hasattr(signal, "SIGTERM"): + signal.signal(signal.SIGTERM, _interrupt) + return launch( + command, + args.world_size, + gpus=gpus, + rendezvous_dir=args.rendezvous_dir, + nccl_library=args.nccl_library, + library_dirs=args.library_dir, + tagged=not args.no_tag_output, + timeout_s=args.timeout, + ) + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/tests/test_architecture.py b/tools/tests/test_architecture.py index c95a86f7af..2f87c80d95 100644 --- a/tools/tests/test_architecture.py +++ b/tools/tests/test_architecture.py @@ -393,6 +393,7 @@ def test_shared_python_and_native_trees_are_closed_minimal_sets() -> None: "core/runtime/primitives/cuda_common.cpp", "core/runtime/primitives/cuda_common.h", "core/runtime/primitives/device_tensor.cpp", + "core/runtime/primitives/dynamic_library.cpp", "core/runtime/primitives/trt_common.cpp", "core/runtime/primitives/trt_common.h", "core/runtime/loader/family_loader.cpp", @@ -420,6 +421,7 @@ def test_shared_python_and_native_trees_are_closed_minimal_sets() -> None: "core/runtime/include/trtmc/internal/tracking.h", "core/runtime/include/trtmc/internal/video.h", "core/runtime/include/trtmc/runtime/device_tensor.h", + "core/runtime/include/trtmc/runtime/dynamic_library.h", "core/runtime/include/trtmc/runtime/family_factory.h", "core/runtime/include/trtmc/runtime/family_loader.h", "core/runtime/include/trtmc/runtime/span.h", @@ -428,7 +430,9 @@ def test_shared_python_and_native_trees_are_closed_minimal_sets() -> None: "core/runtime/include/trtmc/runtime/trt_module.h", "core/runtime/tests/fake_backend.cpp", "core/runtime/tests/fake_family.cpp", + "core/runtime/tests/fake_partial_nccl.cpp", "core/runtime/tests/test_bundle_format_v1.cpp", + "core/runtime/tests/test_dynamic_library.cpp", "core/runtime/tests/test_byok_shape_spec.cpp", "core/runtime/tests/test_family_loader.cpp", "core/runtime/tests/test_internal_config.cpp", @@ -505,6 +509,7 @@ def test_shared_python_and_native_trees_are_closed_minimal_sets() -> None: "tools/e2e_evidence.py", "tools/e2e_report.py", "tools/legal_header_exceptions.toml", + "tools/launch_ranks.py", "tools/legal_headers.py", "tools/model_ci.py", "tools/model_benchmark.py", @@ -592,6 +597,7 @@ def test_shared_python_and_native_trees_are_closed_minimal_sets() -> None: "tools/tests/test_devtoolkit_capabilities.py", "tools/tests/test_e2e_evidence.py", "tools/tests/test_family_impact.py", + "tools/tests/test_launch_ranks.py", "tools/tests/test_merge_ready_slack_alert.py", "tools/tests/test_new_ci.py", "tools/tests/test_model_benchmark.py", @@ -1361,15 +1367,21 @@ def test_runtime_has_no_retired_shared_implementation_surface() -> None: assert violations == [] +# NCCL is loaded at run time, either directly as libnccl.so.2 (ELF-only +# families) or through the portable loader, which honors TRTMC_NCCL_LIBRARY +# and defaults to libnccl.so.2 (ELF) or nccl.dll (Windows). +NCCL_LOADERS = ('dlopen("libnccl.so.2"', "platform::nccl_library()") + + def test_distributed_runtimes_use_one_explicit_launcher_contract() -> None: required = ( '"OMPI_COMM_WORLD_SIZE"', '"OMPI_COMM_WORLD_RANK"', '"OMPI_COMM_WORLD_LOCAL_RANK"', '"TRTMC_NCCL_RENDEZVOUS"', - 'dlopen("libnccl.so.2"', "ncclCommInitRank", ) + nccl_loaders = NCCL_LOADERS forbidden = ( '"PMI_SIZE"', '"PMI_RANK"', @@ -1389,6 +1401,8 @@ def test_distributed_runtimes_use_one_explicit_launcher_contract() -> None: for token in required: if token not in source: violations.append(f"{family.name}:missing:{token}") + if sum(loader in source for loader in nccl_loaders) != 1: + violations.append(f"{family.name}:nccl_loader") for token in forbidden: if token in source: violations.append(f"{family.name}:forbidden:{token}") @@ -1594,7 +1608,7 @@ def test_family_tp_runtimes_load_nccl_only_for_collective_communicators() -> Non sources.append(source) if 'getenv("RANK")' in source: violations.append(f"{path.relative_to(REPO)}:RANK") - loads_nccl = 'dlopen("libnccl.so.2"' in source + loads_nccl = any(loader in source for loader in NCCL_LOADERS) initializes_nccl = "ncclCommInitRank" in source if loads_nccl and not initializes_nccl: violations.append(f"{path.relative_to(REPO)}:NCCL-without-communicator") @@ -1604,10 +1618,11 @@ def test_family_tp_runtimes_load_nccl_only_for_collective_communicators() -> Non '"OMPI_COMM_WORLD_SIZE"', '"OMPI_COMM_WORLD_RANK"', '"OMPI_COMM_WORLD_LOCAL_RANK"', - 'dlopen("libnccl.so.2"', ): if token not in source: violations.append(f"{path.relative_to(REPO)}:missing:{token}") + if not loads_nccl: + violations.append(f"{path.relative_to(REPO)}:missing:NCCL loader") runtime = "\n".join(sources) build_source = "\n".join( @@ -1622,7 +1637,7 @@ def test_family_tp_runtimes_load_nccl_only_for_collective_communicators() -> Non violations.append(f"{family.name}:collective-runtime-mismatch") if 'getenv("OMPI_COMM_WORLD_RANK")' in runtime and not initializes_nccl: rank_only_consumers += 1 - if 'dlopen("libnccl.so.2"' in runtime: + if any(loader in runtime for loader in NCCL_LOADERS): violations.append(f"{family.name}:rank-only-NCCL-loader") if family.name == "patchtsmixer": for token in ( diff --git a/tools/tests/test_launch_ranks.py b/tools/tests/test_launch_ranks.py new file mode 100644 index 0000000000..091b5109e2 --- /dev/null +++ b/tools/tests/test_launch_ranks.py @@ -0,0 +1,232 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import io +import json +import re +import sys +import time +from pathlib import Path + +import pytest + +from tools import launch_ranks + +TAGGED = re.compile(r"^\[1,(\d+)\]<(stdout|stderr)>:(.*)$") + + +def _child(code: str) -> list[str]: + return [sys.executable, "-c", code] + + +def _run(command, world_size, **kwargs): + out, err = io.StringIO(), io.StringIO() + status = launch_ranks.launch(command, world_size, stdout=out, stderr=err, **kwargs) + return status, out.getvalue(), err.getvalue() + + +def _rank_lines(text: str, stream: str) -> dict[int, list[str]]: + lines: dict[int, list[str]] = {} + for line in text.splitlines(): + match = TAGGED.fullmatch(line) + if match and match.group(2) == stream: + lines.setdefault(int(match.group(1)), []).append(match.group(3)) + return lines + + +def test_rank_environments_follow_the_openmpi_contract(tmp_path: Path) -> None: + base = {"KEEP": "1", "OMPI_COMM_WORLD_RANK": "stale"} + rendezvous = tmp_path / "id.bin" + envs = launch_ranks.rank_environments( + base, 2, rendezvous, gpus=["3", "5"], nccl_library="/opt/nccl/libnccl.so.2" + ) + assert base == {"KEEP": "1", "OMPI_COMM_WORLD_RANK": "stale"} + assert [env["OMPI_COMM_WORLD_RANK"] for env in envs] == ["0", "1"] + assert [env["OMPI_COMM_WORLD_LOCAL_RANK"] for env in envs] == ["0", "1"] + for env in envs: + assert env["OMPI_COMM_WORLD_SIZE"] == "2" + assert env["OMPI_COMM_WORLD_LOCAL_SIZE"] == "2" + # Every rank sees every selected GPU and picks one by local rank, as under mpirun. + assert env["CUDA_VISIBLE_DEVICES"] == "3,5" + assert env["TRTMC_NCCL_RENDEZVOUS"] == str(rendezvous) + assert env["TRTMC_NCCL_LIBRARY"] == "/opt/nccl/libnccl.so.2" + assert env["KEEP"] == "1" + assert set(launch_ranks.RANK_ENV_NAMES) <= set(envs[0]) + + +def test_rank_environments_keep_visibility_and_nccl_when_not_requested(tmp_path: Path) -> None: + base = {"CUDA_VISIBLE_DEVICES": "7", "TRTMC_NCCL_LIBRARY": "custom"} + (env,) = launch_ranks.rank_environments(base, 1, tmp_path / "id.bin") + assert env["CUDA_VISIBLE_DEVICES"] == "7" + assert env["TRTMC_NCCL_LIBRARY"] == "custom" + + +def test_library_dirs_prepend_to_the_platform_search_path(tmp_path: Path) -> None: + (linux,) = launch_ranks.rank_environments( + {"LD_LIBRARY_PATH": "/usr/lib"}, + 1, + tmp_path / "id", + library_dirs=["/a", "/b"], + platform="linux", + ) + assert linux["LD_LIBRARY_PATH"] == "/a:/b:/usr/lib" + # Windows environment names are case-insensitive: extend the existing "Path" entry. + (windows,) = launch_ranks.rank_environments( + {"Path": r"C:\Windows"}, + 1, + tmp_path / "id", + library_dirs=[r"C:\nccl\bin"], + platform="win32", + ) + assert windows == {**windows, "Path": r"C:\nccl\bin;C:\Windows"} + assert "PATH" not in windows + (empty,) = launch_ranks.rank_environments( + {}, 1, tmp_path / "id", library_dirs=[r"C:\x"], platform="win32" + ) + assert empty["PATH"] == r"C:\x" + + +@pytest.mark.parametrize( + ("value", "world_size", "expected"), + [(None, 2, None), ("", 2, None), ("0,1", 2, ["0", "1"]), (" 1 , 0 ,2", 2, ["1", "0", "2"])], +) +def test_parse_gpus_accepts_valid_lists(value, world_size, expected) -> None: + assert launch_ranks.parse_gpus(value, world_size) == expected + + +@pytest.mark.parametrize("value", ["0", "0,,1", "0,0", ","]) +def test_parse_gpus_rejects_invalid_lists(value) -> None: + with pytest.raises(ValueError): + launch_ranks.parse_gpus(value, 2) + + +def test_each_launch_gets_a_fresh_rendezvous_path(tmp_path: Path) -> None: + first = launch_ranks.new_rendezvous_path(tmp_path) + second = launch_ranks.new_rendezvous_path(tmp_path) + assert first != second + assert first.parent == tmp_path and not first.exists() + assert first.name.startswith("trtmc_nccl_") and first.suffix == ".bin" + + +def test_library_path_variable() -> None: + assert launch_ranks.library_path_variable("win32") == "PATH" + assert launch_ranks.library_path_variable("linux") == "LD_LIBRARY_PATH" + + +def test_launch_tags_output_and_passes_rank_environment(tmp_path: Path) -> None: + code = ( + "import json, os, sys\n" + "keys = ['OMPI_COMM_WORLD_SIZE', 'OMPI_COMM_WORLD_RANK', 'OMPI_COMM_WORLD_LOCAL_RANK'," + " 'CUDA_VISIBLE_DEVICES', 'TRTMC_NCCL_RENDEZVOUS']\n" + "print(json.dumps({k: os.environ.get(k) for k in keys}))\n" + "print('rank-stderr', os.environ['OMPI_COMM_WORLD_RANK'], file=sys.stderr)\n" + ) + status, out, err = _run(_child(code), 3, gpus=["0", "1", "2"], rendezvous_dir=tmp_path) + assert status == 0, err + stdout_lines = _rank_lines(out, "stdout") + assert sorted(stdout_lines) == [0, 1, 2] + payloads = {rank: json.loads(lines[0]) for rank, lines in stdout_lines.items()} + for rank, payload in payloads.items(): + assert payload["OMPI_COMM_WORLD_RANK"] == str(rank) + assert payload["OMPI_COMM_WORLD_LOCAL_RANK"] == str(rank) + assert payload["OMPI_COMM_WORLD_SIZE"] == "3" + assert payload["CUDA_VISIBLE_DEVICES"] == "0,1,2" + rendezvous = {payload["TRTMC_NCCL_RENDEZVOUS"] for payload in payloads.values()} + assert len(rendezvous) == 1 + (path,) = rendezvous + assert Path(path).parent == tmp_path + assert not Path(path).exists() + assert _rank_lines(err, "stderr") == { + 0: ["rank-stderr 0"], + 1: ["rank-stderr 1"], + 2: ["rank-stderr 2"], + } + + +def test_launch_supports_the_file_rendezvous_contract(tmp_path: Path) -> None: + # Mirrors the family runtimes: rank 0 writes the 128-byte id to ".tmp" + # and renames it; other ranks poll for the file and read it. + code = ( + "import os, sys, time\n" + "path = os.environ['TRTMC_NCCL_RENDEZVOUS']\n" + "rank = int(os.environ['OMPI_COMM_WORLD_RANK'])\n" + "if rank == 0:\n" + " time.sleep(0.3)\n" + " data = os.urandom(128)\n" + " open(path + '.tmp', 'wb').write(data)\n" + " os.replace(path + '.tmp', path)\n" + "else:\n" + " deadline = time.time() + 30\n" + " while not os.path.exists(path):\n" + " assert time.time() < deadline\n" + " time.sleep(0.02)\n" + " data = open(path, 'rb').read()\n" + "print(data.hex())\n" + ) + status, out, err = _run(_child(code), 2, rendezvous_dir=tmp_path) + assert status == 0, err + lines = _rank_lines(out, "stdout") + assert len(lines[0][0]) == 256 + assert lines[0] == lines[1] + assert list(tmp_path.iterdir()) == [] + + +def test_a_failing_rank_terminates_the_others_and_sets_the_status(tmp_path: Path) -> None: + code = ( + "import os, sys, time\n" + "if os.environ['OMPI_COMM_WORLD_RANK'] == '1':\n" + " print('boom', file=sys.stderr)\n" + " sys.exit(3)\n" + "time.sleep(120)\n" + ) + start = time.monotonic() + status, _out, err = _run(_child(code), 2, rendezvous_dir=tmp_path, grace_s=5) + assert status == 3 + assert time.monotonic() - start < 60 + assert "[1,1]:boom" in err + assert "rank 1 exited with status 3" in err + + +def test_timeout_terminates_every_rank(tmp_path: Path) -> None: + status, _out, err = _run( + _child("import time; time.sleep(120)"), + 2, + rendezvous_dir=tmp_path, + timeout_s=1, + grace_s=5, + ) + assert status == 124 + assert "timed out" in err + + +def test_untagged_output(tmp_path: Path) -> None: + status, out, _err = _run(_child("print('plain')"), 1, rendezvous_dir=tmp_path, tagged=False) + assert status == 0 + assert out == "plain\n" + + +def test_cli_runs_the_command_after_double_dash(tmp_path: Path, capsys) -> None: + status = launch_ranks.main( + [ + "-n", + "2", + "--rendezvous-dir", + str(tmp_path), + "--", + sys.executable, + "-c", + "import os; print(os.environ['OMPI_COMM_WORLD_RANK'])", + ] + ) + assert status == 0 + out = capsys.readouterr().out + assert sorted(out.splitlines()) == ["[1,0]:0", "[1,1]:1"] + + +def test_cli_rejects_too_few_gpus(capsys) -> None: + with pytest.raises(SystemExit) as error: + launch_ranks.main(["-n", "2", "--gpus", "0", "--", sys.executable, "-c", "pass"]) + assert error.value.code == 2 + assert "1 device(s) for 2 ranks" in capsys.readouterr().err From 5ae7cab137433e08ed67a6df7c6eb078f39e35fb Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 14:25:39 -0700 Subject: [PATCH 02/15] feat(rtx): find the Windows TensorRT-RTX import library and keep per-rank runtime caches - CMake finds the versioned Windows import library (tensorrt_rtx__.lib) when building the TensorRT-RTX backend. - --runtime-cache expands {rank} to OMPI_COMM_WORLD_RANK (0 when unset), so distributed ranks keep separate TensorRT-RTX runtime caches. Documented and covered by the CLI unit test. (cherry picked from commit 0e5d3d727c76c323eb3a7c42b45bec7792971934, without the Cosmos3 change) Signed-off-by: Peter Kisfaludi --- CMakeLists.txt | 14 +++++++++++++- apps/cli/cli.cpp | 17 ++++++++++++++++- apps/cli/tests/test_cli.cpp | 25 +++++++++++++++++++++++++ website/docs/api/cli-reference.md | 2 +- 4 files changed, 55 insertions(+), 3 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 3173657c8e..8fe31d8cb6 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -251,8 +251,20 @@ if(TRTMC_BUILD_BACKEND_RTX) "TRTMC_BUILD_BACKEND_RTX=ON requires TRTMC_RTX_LIBRARY_DIR" ) endif() + # Windows TensorRT-RTX packages name the import library by version + # (tensorrt_rtx__.lib); prefer the newest one in the directory. + set(_trtmc_rtx_names tensorrt_rtx) + if(WIN32) + file(GLOB _trtmc_rtx_versioned RELATIVE "${TRTMC_RTX_LIBRARY_DIR}" + "${TRTMC_RTX_LIBRARY_DIR}/tensorrt_rtx_*_*.lib") + list(SORT _trtmc_rtx_versioned COMPARE NATURAL ORDER DESCENDING) + foreach(_trtmc_rtx_lib IN LISTS _trtmc_rtx_versioned) + get_filename_component(_trtmc_rtx_name "${_trtmc_rtx_lib}" NAME_WE) + list(APPEND _trtmc_rtx_names "${_trtmc_rtx_name}") + endforeach() + endif() find_library(TRTMC_RTX_LIBRARY - NAMES tensorrt_rtx + NAMES ${_trtmc_rtx_names} PATHS "${TRTMC_RTX_LIBRARY_DIR}" NO_DEFAULT_PATH ) diff --git a/apps/cli/cli.cpp b/apps/cli/cli.cpp index 2510bbd54e..10c7dfe422 100644 --- a/apps/cli/cli.cpp +++ b/apps/cli/cli.cpp @@ -183,6 +183,20 @@ std::string take_value(int argc, char** argv, int& index, const std::string& opt return value; } +// Distributed ranks share one command line (mpirun or tools/launch_ranks.py), but +// each rank needs its own TensorRT-RTX runtime cache file. "{rank}" in the path +// becomes the OpenMPI world rank, or 0 for a single-process run. +std::string expand_rank_placeholder(std::string path) { + static const std::string placeholder = "{rank}"; + const char* rank = std::getenv("OMPI_COMM_WORLD_RANK"); + const std::string value = rank != nullptr && *rank != '\0' ? rank : "0"; + for (auto at = path.find(placeholder); at != std::string::npos; + at = path.find(placeholder, at + value.size())) { + path.replace(at, placeholder.size(), value); + } + return path; +} + std::uint64_t parse_byte_size(const std::string& text) { std::uint64_t multiplier = 1; std::string number = text; @@ -789,7 +803,8 @@ Command parse_args(int argc, char** argv) { if (option == "--runtime-cache") { if (!command.runtime_cache_path.empty()) throw std::invalid_argument("--runtime-cache may be specified only once"); - command.runtime_cache_path = take_value(argc, argv, index, option); + command.runtime_cache_path = + expand_rank_placeholder(take_value(argc, argv, index, option)); continue; } if (option == "--cuda-graphs") { diff --git a/apps/cli/tests/test_cli.cpp b/apps/cli/tests/test_cli.cpp index 65bb4f38d2..766b9072a7 100644 --- a/apps/cli/tests/test_cli.cpp +++ b/apps/cli/tests/test_cli.cpp @@ -8,6 +8,7 @@ #include "cli/sdk_dispatch.h" #include +#include #include #include #include @@ -39,6 +40,15 @@ trtmc::cli::Command parse(std::vector arguments) { return trtmc::cli::parse_args(static_cast(argv.size()), argv.data()); } +// An empty value counts as unset for the runtime-cache {rank} expansion. +void set_env(const char* name, const std::string& value) { +#ifdef _WIN32 + _putenv_s(name, value.c_str()); +#else + setenv(name, value.c_str(), 1); +#endif +} + bool parse_throws(std::vector arguments) { try { (void)parse(std::move(arguments)); @@ -394,6 +404,21 @@ int main() { "--runtime-cache", "kernels.cache", "--cuda-graphs"}); check(rtx.runtime_cache_path == "kernels.cache" && rtx.cuda_graphs, "TensorRT-RTX runtime options are retained directly"); + { + const char* saved = std::getenv("OMPI_COMM_WORLD_RANK"); + const std::string previous = saved != nullptr ? saved : ""; + set_env("OMPI_COMM_WORLD_RANK", ""); + check(parse({"trtmc", "run", "model.bundle", "--runtime-root", "lib", "--runtime-cache", + "k.rank{rank}.cache"}) + .runtime_cache_path == "k.rank0.cache", + "runtime cache {rank} defaults to rank 0 outside a distributed launch"); + set_env("OMPI_COMM_WORLD_RANK", "3"); + check(parse({"trtmc", "run", "model.bundle", "--runtime-root", "lib", "--runtime-cache", + "{rank}/k.{rank}.cache"}) + .runtime_cache_path == "3/k.3.cache", + "runtime cache {rank} expands to the OpenMPI world rank"); + set_env("OMPI_COMM_WORLD_RANK", previous); + } check(parse_throws({"trtmc", "run", "model.bundle", "--runtime-root", "lib", "--cuda-graphs", "--cuda-graphs"}), "duplicate TensorRT-RTX graph option is rejected"); diff --git a/website/docs/api/cli-reference.md b/website/docs/api/cli-reference.md index 78503f8021..48b4774c00 100644 --- a/website/docs/api/cli-reference.md +++ b/website/docs/api/cli-reference.md @@ -90,7 +90,7 @@ installed fallback. Common load options are: | --- | --- | | `--runtime-root DIR` | Required exact DSO root. | | `--kv-cache-size BYTES\|GB\|GiB` | Runtime-sized KV capacity for a compatible bundle. | -| `--runtime-cache PATH` | TensorRT-RTX cache path; rejected by the standard TensorRT backend. | +| `--runtime-cache PATH` | TensorRT-RTX cache path; rejected by the standard TensorRT backend. `{rank}` in `PATH` becomes `OMPI_COMM_WORLD_RANK` (0 when unset), so distributed ranks keep separate caches. | | `--cuda-graphs` | Enable TensorRT-RTX CUDA graphs; rejected by the standard backend. | | `--byok-library DSO`, `--byok-function NAME`, `--byok-name NAME` | Load one TVM-FFI BYOK binding. All three are required together. | From 6c91a392472780de8aa4e5d648c869b2877b6abd Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Tue, 22 Sep 2026 22:10:32 -0700 Subject: [PATCH 03/15] fix(build): keep conanfile free of requirements; test module paths portably The package build uses a preinstalled offline toolchain, and CI rejects self.requires() in conanfile.py. Drop the Windows nlohmann_json requirement; Windows builds install it separately and pass its CMake package directory through CMAKE_PREFIX_PATH, which generate() now forwards like TRT_ROOT. test_dynamic_library took the address of an imported function to find trtmc_core. On Windows that address is the import thunk inside the executable, so the test now checks module_path_containing() with data in the executable and a symbol inside a loaded library. Signed-off-by: Peter Kisfaludi (cherry picked from commit c415c56320fadd91f0ab5e08c02bee543efaa3e1) Signed-off-by: Peter Kisfaludi --- conanfile.py | 11 ++++------- core/runtime/tests/test_dynamic_library.cpp | 22 ++++++++++++++------- 2 files changed, 19 insertions(+), 14 deletions(-) diff --git a/conanfile.py b/conanfile.py index 37069316d9..0e8f0c21dc 100644 --- a/conanfile.py +++ b/conanfile.py @@ -11,7 +11,7 @@ from conan import ConanFile from conan.errors import ConanException -from conan.tools.cmake import CMake, CMakeDeps, CMakeToolchain, cmake_layout +from conan.tools.cmake import CMake, CMakeToolchain, cmake_layout from conan.tools.files import copy @@ -59,11 +59,6 @@ def layout(self) -> None: # CMakeToolchain derives install directories from the package layout. self.cpp.package.libdirs = ["bin"] - def requirements(self) -> None: - # Linux images provide nlohmann-json3-dev; MSVC builds take it from Conan. - if self._windows(): - self.requires("nlohmann_json/3.11.3") - def generate(self) -> None: toolchain = CMakeToolchain(self) toolchain.cache_variables["TRTMC_BUILD_TESTS"] = False @@ -71,6 +66,9 @@ def generate(self) -> None: "TRT_ROOT", "CMAKE_CUDA_ARCHITECTURES", "TRTMC_FAMILIES", + # Windows hosts have no system nlohmann_json; point CMake at an + # installed package (for example a conan install --requires output). + "CMAKE_PREFIX_PATH", ): value = os.environ.get(name) if value: @@ -80,7 +78,6 @@ def generate(self) -> None: # families; the server, BYOK bridge, and examples stay ELF-only. for option in ("TRTMC_BUILD_SERVER", "TRTMC_ENABLE_BYOK", "TRTMC_BUILD_EXAMPLES"): toolchain.cache_variables[option] = False - CMakeDeps(self).generate() toolchain.generate() def build(self) -> None: diff --git a/core/runtime/tests/test_dynamic_library.cpp b/core/runtime/tests/test_dynamic_library.cpp index 94314a044c..fa6d04c3a8 100644 --- a/core/runtime/tests/test_dynamic_library.cpp +++ b/core/runtime/tests/test_dynamic_library.cpp @@ -113,16 +113,24 @@ void test_partial_library(const fs::path& partial) { "symbol error names the library: " + message); } -void test_module_paths(const char* argv0) { +void test_module_paths(const char* argv0, const fs::path& partial) { const auto executable = trtmc::platform::current_executable_path(); check(executable.is_absolute(), "executable path is absolute: " + executable.string()); check(executable.stem() == fs::path(argv0).stem(), "current_executable_path is this test: " + executable.string()); - const auto containing = trtmc::platform::module_path_containing( - reinterpret_cast(&trtmc::platform::nccl_library)); - check(containing.filename() == trtmc::platform::shared_library_filename("trtmc_core"), - "module_path_containing finds trtmc_core: " + containing.string()); - check(containing.is_absolute(), "module path is absolute"); + + static const char executable_anchor = 0; + const auto self = trtmc::platform::module_path_containing(&executable_anchor); + check(fs::equivalent(self, executable), + "module_path_containing(executable data) is the executable: " + self.string()); + + // An address inside a loaded library maps back to that library file. + DynamicLibrary library(partial.string(), "Unit test: NCCL"); + const auto containing = + trtmc::platform::module_path_containing(library.find_symbol("ncclGetVersion")); + check(containing.is_absolute(), "module path is absolute: " + containing.string()); + check(fs::equivalent(containing, partial), + "module_path_containing(library symbol) is the library: " + containing.string()); } } // namespace @@ -138,7 +146,7 @@ int main(int argc, char** argv) { test_nccl_library_override(); test_missing_library(partial.parent_path()); test_partial_library(partial); - test_module_paths(argv[0]); + test_module_paths(argv[0], partial); } catch (const std::exception& error) { std::cerr << "FAIL: unexpected exception: " << error.what() << '\n'; return 1; From 6476420cb85dc8c56707c95900c32c318c14954a Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Wed, 30 Sep 2026 15:10:19 -0700 Subject: [PATCH 04/15] docs(multi-device): build NCCL from source on Windows NVIDIA publishes no Windows NCCL binaries, so native Windows multi-GPU runs need an nccl.dll built from the upstream source. Document the requirements (TCC mode, TensorRT 11.4+ or TensorRT-RTX 1.7.1+), the CMake build of NVIDIA/nccl v2.32.3-1 with CUDA 13.x, the two extra settings CUDA 12.9 needs until the upstream fixes land, and how TRTMC_NCCL_LIBRARY and tools/launch_ranks.py pick the library up. Validated on 2x RTX PRO 6000 Blackwell (TCC) with TensorRT 11.5.0.30: both builds, done exactly as documented, produce an nccl.dll that imports only Windows system DLLs. With either one, the LTX-Video CP=2 run gives 161/161 frames bit-identical to the reference run. Signed-off-by: Peter Kisfaludi (cherry picked from commit 181768081ad44e69bfe1f54fdb4486aff34c7291) --- website/docs/features/multi-device.md | 59 +++++++++++++++++++++++++++ 1 file changed, 59 insertions(+) diff --git a/website/docs/features/multi-device.md b/website/docs/features/multi-device.md index 35a3c996f7..aa6d5144e9 100644 --- a/website/docs/features/multi-device.md +++ b/website/docs/features/multi-device.md @@ -53,6 +53,65 @@ The family maps launcher rank to a visible device and must keep its communicator alive for the TensorRT engines that use it. Runtime process count, visible devices, and bundle topology must agree. +## Native Windows + +NVIDIA does not publish NCCL binaries for Windows, so you build `nccl.dll` +from the upstream source. Requirements: + +- Two or more GPUs in TCC mode (`nvidia-smi -g -dm 1` from an elevated + prompt). Under MCDM or WDDM, the GPUs report no peer access and NCCL has no + working transport. +- TensorRT 11.4 or newer, or TensorRT-RTX 1.7.1 or newer. Older Windows + releases do not load `nccl.dll`. +- Visual Studio 2022 (MSVC v143), CMake 3.25 or newer, Ninja, Git, Python 3, + and a CUDA 13.x toolkit. CUDA 12.9 also works; see the note below. + +Build it in PowerShell. Set `CMAKE_CUDA_ARCHITECTURES` to your GPUs' compute +capabilities, for example `120` for RTX PRO 6000 Blackwell or +`86-real;89-real;120-real` for a mix. If you leave it out, you get nvcc's default +architecture instead of your GPUs'. + +```powershell +git clone --branch v2.32.3-1 https://github.com/NVIDIA/nccl.git C:\nccl\src +& "C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools\Common7\Tools\Launch-VsDevShell.ps1" -Arch amd64 -HostArch amd64 -SkipAutomaticLocation +$cuda = $env:CUDA_PATH_V13_3 -replace '\\', '/' +cmake -S C:\nccl\src -B C:\nccl\build -G Ninja -DCMAKE_BUILD_TYPE=Release ` + "-DCMAKE_CUDA_COMPILER=$cuda/bin/nvcc.exe" "-DCUDAToolkit_ROOT=$cuda" ` + "-DCMAKE_CUDA_ARCHITECTURES=120" -DCMAKE_INSTALL_PREFIX=C:\nccl\install +cmake --build C:\nccl\build --parallel +cmake --install C:\nccl\build +``` + +With CUDA 12.9, NCCL v2.32.3-1 needs two extra configure settings. Its CUDA 12 +build is C++14, which MSVC rejects in `sym_kernels.h`. And it links +`cudart64_12.dll` dynamically, which no TensorRT package ships. Write a CMake +include that switches the build to C++17, and link the static runtime: + +```powershell +@' +set(CMAKE_CUDA14_STANDARD_COMPILE_OPTION "-std=c++17") +set(CMAKE_CUDA14_EXTENSION_COMPILE_OPTION "-std=c++17") +set(CMAKE_CXX14_STANDARD_COMPILE_OPTION "-std:c++17") +set(CMAKE_CXX14_EXTENSION_COMPILE_OPTION "-std:c++17") +'@ | Set-Content -Encoding ascii C:\nccl\cuda12_windows.cmake +# add to the cmake configure line above: +# "-DCUDA_cudart_LIBRARY=$cuda/lib/x64/cudart_static.lib" +# -DCMAKE_PROJECT_NCCL_INCLUDE=C:/nccl/cuda12_windows.cmake +``` + +Point the runtime at the result. Either set `TRTMC_NCCL_LIBRARY` to +`C:\nccl\install\bin\nccl.dll` or put `C:\nccl\install\bin` first on `PATH`. +TensorRT loads `nccl.dll` by name and gets the module that TRTMC already +loaded. Native Windows has no `mpirun`, so `tools/launch_ranks.py` starts the +ranks and sets the same rank and rendezvous variables: + +```powershell +python tools\launch_ranks.py -n 2 --gpus 0,1 ` + --nccl-library C:\nccl\install\bin\nccl.dll ` + --library-dir ` + -- trtmc generate-video model-cp2.bundle --runtime-root ... +``` + ## Find exact support Do not infer support from the generic flags. Search family manifests for the From d9a4541cbe51b6a03cef8d24e2bfdc383dda1c76 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 21:27:37 -0700 Subject: [PATCH 05/15] fix(cli): quote Python family arguments on Windows _spawnvp joins its argument array with spaces and does not quote it, so a family command value with spaces, quotes, or trailing backslashes reached Python as different arguments. Quote every forwarded argument with the rules the C runtime and CommandLineToArgvW use to split a command line. The quoting function is portable and unit tested on every platform; on Windows the test also round-trips the arguments through CommandLineToArgvW. Also load family CLI adapters with critical-error dialogs suppressed, so a missing dependent DLL is reported as an error instead of blocking an unattended run on a modal loader dialog. Signed-off-by: Peter Kisfaludi --- apps/cli/family_cli.cpp | 37 ++++++++++++++++++++++++++--- apps/cli/family_cli.h | 7 ++++++ apps/cli/tests/test_cli.cpp | 46 +++++++++++++++++++++++++++++++++++++ 3 files changed, 87 insertions(+), 3 deletions(-) diff --git a/apps/cli/family_cli.cpp b/apps/cli/family_cli.cpp index e25c24b3a4..5a4d8fd572 100644 --- a/apps/cli/family_cli.cpp +++ b/apps/cli/family_cli.cpp @@ -391,10 +391,15 @@ class CliLibrary { public: explicit CliLibrary(const fs::path& path) { #if defined(_WIN32) + // Report a missing dependent DLL as an error instead of a modal loader + // dialog that would block unattended runs. + const UINT previous_mode = SetErrorMode(SEM_FAILCRITICALERRORS | SEM_NOOPENFILEERRORBOX); handle_ = LoadLibraryExW(path.wstring().c_str(), nullptr, LOAD_WITH_ALTERED_SEARCH_PATH); + const DWORD error = handle_ == nullptr ? GetLastError() : ERROR_SUCCESS; + SetErrorMode(previous_mode); if (handle_ == nullptr) { throw std::runtime_error("cannot load family CLI: " + path.string() + - " (Windows error " + std::to_string(GetLastError()) + ")"); + " (Windows error " + std::to_string(error) + ")"); } #else handle_ = dlopen(path.c_str(), RTLD_NOW | RTLD_LOCAL); @@ -450,9 +455,14 @@ int invoke(const fs::path& executable, const std::string& family, const Json& co #if defined(_WIN32) // Windows has no exec: run the Python command as a child process and // return its exit status. The launcher is "python" on Windows. - std::vector arguments{"python", "-m", "tensorrt_model_connect"}; + // _spawnvp joins the arguments with spaces without quoting them, so + // quote each one to keep values with spaces or quotes intact. + std::vector quoted; for (int i = 1; i < argc; ++i) - arguments.push_back(argv[i]); + quoted.push_back(quote_windows_argument(argv[i])); + std::vector arguments{"python", "-m", "tensorrt_model_connect"}; + for (const auto& argument : quoted) + arguments.push_back(argument.c_str()); arguments.push_back(nullptr); output.flush(); error.flush(); @@ -499,6 +509,27 @@ int invoke(const fs::path& executable, const std::string& family, const Json& co } } // namespace +std::string quote_windows_argument(const std::string& argument) { + if (!argument.empty() && argument.find_first_of(" \t\n\v\"") == std::string::npos) + return argument; + // Backslashes are literal unless they precede a quote: double those before + // an embedded quote (which is then escaped) and before the closing quote. + std::string quoted = "\""; + std::size_t backslashes = 0; + for (const char c : argument) { + if (c == '\\') { + ++backslashes; + continue; + } + quoted.append(c == '"' ? backslashes * 2 + 1 : backslashes, '\\'); + quoted.push_back(c); + backslashes = 0; + } + quoted.append(backslashes * 2, '\\'); + quoted.push_back('"'); + return quoted; +} + std::optional run_family_cli(int argc, char** argv, std::ostream& output, std::ostream& error, const fs::path& executable_override) { try { diff --git a/apps/cli/family_cli.h b/apps/cli/family_cli.h index 39d7d0cd39..b815b2ff3a 100644 --- a/apps/cli/family_cli.h +++ b/apps/cli/family_cli.h @@ -8,6 +8,7 @@ #include #include #include +#include namespace trtmc::cli { @@ -16,4 +17,10 @@ namespace trtmc::cli { std::optional run_family_cli(int argc, char** argv, std::ostream& output, std::ostream& error, const std::filesystem::path& executable = {}); +// Quotes one argument for a Windows command line so that the C runtime and +// CommandLineToArgvW parse it back as exactly that argument. Arguments without +// whitespace or quotes are returned unchanged. Portable, so it is unit tested +// on every platform. +std::string quote_windows_argument(const std::string& argument); + } // namespace trtmc::cli diff --git a/apps/cli/tests/test_cli.cpp b/apps/cli/tests/test_cli.cpp index 766b9072a7..199ca06536 100644 --- a/apps/cli/tests/test_cli.cpp +++ b/apps/cli/tests/test_cli.cpp @@ -4,6 +4,7 @@ */ #include "cli/cli.h" +#include "cli/family_cli.h" #include "cli/io.h" #include "cli/sdk_dispatch.h" @@ -21,6 +22,15 @@ #include #include +#ifdef _WIN32 +#ifndef NOMINMAX +#define NOMINMAX +#endif +#include +// shellapi.h depends on windows.h. +#include +#endif + namespace { int failures = 0; @@ -49,6 +59,41 @@ void set_env(const char* name, const std::string& value) { #endif } +// Family Python commands reach the interpreter through a Windows command line. +void check_windows_argument_quoting() { + using trtmc::cli::quote_windows_argument; + const std::vector> cases{ + {"plain", "plain"}, + {"C:\\models\\ltx2", "C:\\models\\ltx2"}, + {"", "\"\""}, + {"a red fox", "\"a red fox\""}, + {"tab\there", "\"tab\there\""}, + {"C:\\Program Files\\model\\", "\"C:\\Program Files\\model\\\\\""}, + {"say \"hi\"", "\"say \\\"hi\\\"\""}, + {"a\\\"b", "\"a\\\\\\\"b\""}, + {"\"", "\"\\\"\""}, + }; + for (const auto& [argument, expected] : cases) + check(quote_windows_argument(argument) == expected, + "Windows argument quoting follows the CommandLineToArgvW rules"); +#ifdef _WIN32 + std::wstring line = L"python"; + for (const auto& entry : cases) { + const auto quoted = quote_windows_argument(entry.first); + line += L' '; + line += std::wstring(quoted.begin(), quoted.end()); + } + int count = 0; + LPWSTR* parsed = CommandLineToArgvW(line.c_str(), &count); + bool round_trip = parsed != nullptr && count == static_cast(cases.size()) + 1; + for (std::size_t i = 0; round_trip && i < cases.size(); ++i) + round_trip = std::wstring(parsed[i + 1]) == + std::wstring(cases[i].first.begin(), cases[i].first.end()); + LocalFree(parsed); + check(round_trip, "CommandLineToArgvW parses quoted arguments back unchanged"); +#endif +} + bool parse_throws(std::vector arguments) { try { (void)parse(std::move(arguments)); @@ -332,6 +377,7 @@ bool dispatch_throws(const trtmc::cli::Command& command, trtmc::ITask& task) { } // namespace int main() { + check_windows_argument_quoting(); const std::vector execution_commands{ "run", "encode", From a33f80bbbdf93551cd37d30ca18237c9be4d6447 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 21:27:37 -0700 Subject: [PATCH 06/15] fix(build): translate GCC compile options per generator expression for MSVC The MSVC option translation stripped -W flags textually, which turned -Xcompiler=-Wall,-Wextra into a malformed -Xcompiler=, and left -O3 in C++ options, where cl.exe ignores it with D9002. It could also drop the closing '>' of a split generator expression. Rebuild each $<$:...> expression from its translated options instead: drop -W flags, drop GCC entries from nvcc host-compiler pass-through options (and the option when nothing remains), keep -O for nvcc only, and omit expressions that end up empty. Signed-off-by: Peter Kisfaludi --- CMakeLists.txt | 81 +++++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 71 insertions(+), 10 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 8fe31d8cb6..adcd0b3749 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -1191,10 +1191,39 @@ install(FILES COMPONENT sdk ) -# GCC/Clang warning flags are spelled inline on each target. MSVC does not -# understand them (and nvcc forwards them to cl.exe), so translate every -# target's options once, after all targets exist. +# GCC/Clang warning and optimization flags are spelled inline on each target. +# MSVC does not understand them (and nvcc forwards them to cl.exe), so +# translate every target's options once, after all targets exist. if(MSVC) + # Translates one option that applies to `language`: CUDA, HOST (C/C++ only), + # or ANY. Sets `result` to the MSVC option, or to "" to drop it. + function(_trtmc_msvc_translate_option option language result) + set(_value "${option}") + if(option MATCHES "^-W") + set(_value "") + elseif(option MATCHES "^(-Xcompiler=|--compiler-options=)(.*)$") + # nvcc host-compiler pass-through: drop the GCC entries and keep the + # option only when cl.exe entries remain. + set(_prefix "${CMAKE_MATCH_1}") + string(REPLACE "," ";" _host "${CMAKE_MATCH_2}") + list(FILTER _host EXCLUDE REGEX "^-[WO]") + list(JOIN _host "," _host) + set(_value "") + if(NOT _host STREQUAL "") + set(_value "${_prefix}${_host}") + endif() + elseif(option MATCHES "^-O[0-9s]?$") + # nvcc accepts -O; cl.exe does not, and the MSVC build type already + # selects /O2 or /Od for C and C++. + if(language STREQUAL "HOST") + set(_value "") + elseif(language STREQUAL "ANY") + set(_value "$<$:${option}>") + endif() + endif() + set(${result} "${_value}" PARENT_SCOPE) + endfunction() + function(_trtmc_msvc_translate_warning_flags directory) get_property(_targets DIRECTORY "${directory}" PROPERTY BUILDSYSTEM_TARGETS) foreach(_target IN LISTS _targets) @@ -1206,18 +1235,50 @@ if(MSVC) if(NOT _options) continue() endif() + # Generator expressions such as $<$:-O3;-Wall> + # arrive split on ';'. Collect the options inside one, translate them + # for its language, and rebuild it only when options remain. set(_translated) + set(_open "") + set(_language ANY) foreach(_option IN LISTS _options) - if(_option MATCHES "^-W" OR _option MATCHES "^-Werror") - continue() + if(_open STREQUAL "" AND _option MATCHES "^(\\$<\\$:)(.*)$") + set(_open "${CMAKE_MATCH_1}") + set(_languages "${CMAKE_MATCH_2}") + set(_option "${CMAKE_MATCH_3}") + set(_inner) + if(_languages STREQUAL "CUDA") + set(_language CUDA) + elseif(_languages MATCHES "CUDA") + set(_language ANY) + else() + set(_language HOST) + endif() + endif() + set(_close FALSE) + if(NOT _open STREQUAL "" AND _option MATCHES "^(.*)>$") + set(_option "${CMAKE_MATCH_1}") + set(_close TRUE) endif() - # Generator expressions such as $<$:-Wall;-Wextra> - # arrive split on ';'; drop GCC flags inside them and keep the rest. - string(REGEX REPLACE "-W[A-Za-z0-9=_-]+" "" _option "${_option}") - if(_option MATCHES "^\\$<\\$:>?$" OR _option STREQUAL "" OR _option STREQUAL ">") + _trtmc_msvc_translate_option("${_option}" ${_language} _option) + if(_open STREQUAL "") + if(NOT _option STREQUAL "") + list(APPEND _translated "${_option}") + endif() continue() endif() - list(APPEND _translated "${_option}") + if(NOT _option STREQUAL "") + list(APPEND _inner "${_option}") + endif() + if(_close) + list(LENGTH _inner _count) + if(_count GREATER 0) + list(JOIN _inner ";" _inner) + list(APPEND _translated "${_open}${_inner}>") + endif() + set(_open "") + set(_language ANY) + endif() endforeach() list(APPEND _translated "$<$:/W3>") set_property(TARGET ${_target} PROPERTY COMPILE_OPTIONS "${_translated}") From f9ccf6d88a120ab8b6a8896ed05f0a98eb5db443 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 21:27:37 -0700 Subject: [PATCH 07/15] fix(build): package family CLI declarations in the Windows package The Windows package did not contain families//cli.json, so an installed trtmc.exe could not resolve family commands. Copy the declarations of the packaged families beside the executable and check that the trtmc_cli_.dll set matches the families that declare native commands, as the Linux package already does. Signed-off-by: Peter Kisfaludi --- conanfile.py | 26 ++++++++++++++++++++++++++ 1 file changed, 26 insertions(+) diff --git a/conanfile.py b/conanfile.py index 0e8f0c21dc..d47724a7ea 100644 --- a/conanfile.py +++ b/conanfile.py @@ -114,6 +114,32 @@ def _package_windows(self) -> None: f"family DLL set does not match: missing={sorted(expected - packaged)}, " f"extra={sorted(packaged - expected)}" ) + self._package_windows_cli(source, module_bin, expected) + + def _package_windows_cli(self, source: Path, module_bin: Path, families: set[str]) -> None: + # trtmc.exe resolves family commands from families//cli.json + # beside the executable and loads native adapters from that directory. + expected = set() + for family in sorted(families): + declaration = source / "families" / family / "cli.json" + if not declaration.is_file(): + continue + copy( + self, + declaration.name, + src=str(declaration.parent), + dst=str(module_bin / "families" / family), + keep_path=False, + ) + commands = json.loads(declaration.read_text(encoding="utf-8"))["commands"] + if any(command["executor"] == "native" for command in commands): + expected.add(f"trtmc_cli_{family}.dll") + packaged = {path.name for path in module_bin.glob("trtmc_cli_*.dll")} + if packaged != expected: + raise ConanException( + f"family CLI adapter set does not match CLI declarations: " + f"missing={sorted(expected - packaged)}, extra={sorted(packaged - expected)}" + ) def package(self) -> None: if self._windows(): From 715bd534f39112c4f1d7d905f3bb0beb9d348977 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 15:17:43 -0700 Subject: [PATCH 08/15] feat(api): accept worker completion for synchronized audio-video results Distributed families return results only from the output rank; the other ranks return the worker-completion sentinel. Video results already accept it, but an audio-video result required decoded frames, audio and an audio clock origin on every rank, so a context-parallel text-to-audio-video family could not report a worker rank's completion. An audio-video result whose video is the worker-completion sentinel is now accepted when it also carries no audio samples and no audio clock origin; its view has no frames and no samples. The video fixture covers the worker case. Signed-off-by: Peter Kisfaludi --- core/api/runtime/video.cpp | 9 +++++++++ core/api/tests/video_family.cpp | 8 ++++++++ core/api/tests/video_test.cpp | 4 ++++ 3 files changed, 21 insertions(+) diff --git a/core/api/runtime/video.cpp b/core/api/runtime/video.cpp index 60e1f007fb..ad701ed02e 100644 --- a/core/api/runtime/video.cpp +++ b/core/api/runtime/video.cpp @@ -409,6 +409,15 @@ struct ActionVideoStorage final : ResultStorage { struct AudioVideoStorage final : ResultStorage { explicit AudioVideoStorage(internal::AudioVideoResult result) : video(std::move(result.video)), audio(std::move(result.audio)) { + if (internal::is_worker_completion(video.result.frames)) { + // A distributed worker rank returns no media; rank 0 owns the synchronized result. + output_check(audio.result.samples.empty() && !result.audio_start_seconds, + "worker completion must not contain audio or an audio clock origin"); + view.video = video.view; + fill_audio_result_view(audio, &view.audio); + view.audio_start_seconds = 0.0; + return; + } output_check( result.audio_start_seconds && std::isfinite(*result.audio_start_seconds) && !video.result.timestamps_seconds.empty() && !audio.result.samples.empty(), diff --git a/core/api/tests/video_family.cpp b/core/api/tests/video_family.cpp index 1ff9c6f541..2337d15cce 100644 --- a/core/api/tests/video_family.cpp +++ b/core/api/tests/video_family.cpp @@ -322,6 +322,14 @@ class VideoFixture final : public IModel, } AudioVideoResult run(const TextToAudioVideoRequest& input, ConfigView config) override { ++single_calls; + if (mode_ == "worker") { + (void)gain(config); + AudioVideoResult result; + result.video.frames.num_frames = 0; + result.audio.sample_rate = 48000; + result.audio.channels = 2; + return result; + } auto result = audio_video(16, float(input.prompt.size()) / 100, gain(config)); if (mode_ == "bad_av") result.audio_start_seconds.reset(); diff --git a/core/api/tests/video_test.cpp b/core/api/tests/video_test.cpp index 74f4375cbe..140aa7bbd4 100644 --- a/core/api/tests/video_test.cpp +++ b/core/api/tests/video_test.cpp @@ -355,6 +355,10 @@ int main(int argc, char** argv) { participant.timestamps_seconds().empty() && participant.conditioned_prefix_frames() == 0, "non-output distributed completion is not decoded media or an inference failure"); + auto av_participant = load("worker").task().run({"participate"}); + check(av_participant.frames().empty() && av_participant.timestamps_seconds().empty() && + av_participant.audio().samples.empty(), + "a distributed worker's audio-video completion carries neither frames nor audio"); auto none = load("none"); check(none.tasks().empty() && !none.supports(), "interface inheritance alone does not advertise video support"); From 73bef4495f9667d63c0290d10ff2bffda2d7dfaf Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 15:17:43 -0700 Subject: [PATCH 09/15] feat(cli): write text_to_audio_video results as frames and audio.wav generate-video selects TextToAudioVideo for bundles whose task is text_to_audio_video (without --image). It writes the frames like other video Tasks and the soundtrack as OUTPUT/audio.wav (interleaved PCM at the result's sample rate and channel count), and reports the audio path, rate, channels and audio start time in the JSON output. Worker ranks report {"worker": true} without writing files. Signed-off-by: Peter Kisfaludi --- apps/cli/sdk_video.cpp | 66 +++++++++++++++++++++++++++++++++--------- 1 file changed, 53 insertions(+), 13 deletions(-) diff --git a/apps/cli/sdk_video.cpp b/apps/cli/sdk_video.cpp index 08696caf6e..5b0e831a2f 100644 --- a/apps/cli/sdk_video.cpp +++ b/apps/cli/sdk_video.cpp @@ -39,12 +39,8 @@ std::vector intrinsics(const Command& command) { return values; } -nlohmann::json write_video(const VideoGenerationResult& video, const std::string& directory) { - const auto frames = video.frames(); - // The C result has already validated the exact worker-completion sentinel. - // It carries no image/frame payload; never create an empty output directory. - if (frames.empty()) - return {{"worker", true}}; +nlohmann::json write_frames(Span frames, Span times, + const std::string& directory) { std::filesystem::create_directories(directory); nlohmann::json files = nlohmann::json::array(); for (std::size_t frame = 0; frame < frames.size(); ++frame) { @@ -59,18 +55,50 @@ nlohmann::json write_video(const VideoGenerationResult& video, const std::string path); files.push_back(path); } - nlohmann::json times = nlohmann::json::array(); - for (const auto time : video.timestamps_seconds()) - times.push_back(time); + nlohmann::json timestamps = nlohmann::json::array(); + for (const auto time : times) + timestamps.push_back(time); return {{"output", directory}, {"frames", std::move(files)}, {"height", frames[0].height}, {"width", frames[0].width}, {"channels", frames[0].channels}, - {"timestamps_seconds", std::move(times)}, - {"conditioned_prefix_frames", video.conditioned_prefix_frames()}, - {"setup_ms", video.setup_ms()}, - {"inference_ms", video.inference_ms()}}; + {"timestamps_seconds", std::move(timestamps)}}; +} + +nlohmann::json write_video(const VideoGenerationResult& video, const std::string& directory) { + const auto frames = video.frames(); + // The C result has already validated the exact worker-completion sentinel. + // It carries no image/frame payload; never create an empty output directory. + if (frames.empty()) + return {{"worker", true}}; + auto json = write_frames(frames, video.timestamps_seconds(), directory); + json["conditioned_prefix_frames"] = video.conditioned_prefix_frames(); + json["setup_ms"] = video.setup_ms(); + json["inference_ms"] = video.inference_ms(); + return json; +} + +// Synchronized video + audio: the frames as in generate-video plus OUTPUT/audio.wav +// (interleaved PCM at the result's own rate and channel count). +nlohmann::json write_audio_video(const AudioVideoGenerationResult& result, + const std::string& directory) { + const auto frames = result.frames(); + if (frames.empty()) + return {{"worker", true}}; + auto json = write_frames(frames, result.timestamps_seconds(), directory); + const auto audio = result.audio(); + const auto audio_path = (std::filesystem::path(directory) / "audio.wav").string(); + io::write_wav_interleaved(audio.samples, + static_cast(audio.sample_rate.value_or(0)), + static_cast(audio.channels), audio_path); + json["audio"] = audio_path; + json["audio_sample_rate"] = audio.sample_rate.value_or(0); + json["audio_channels"] = audio.channels; + json["audio_start_seconds"] = result.audio_start_seconds(); + json["setup_ms"] = result.video_view().setup_ms; + json["inference_ms"] = result.video_view().inference_ms; + return json; } void generate_world(const Command& command, const Model& model, std::string_view id, @@ -127,6 +155,8 @@ std::string_view video_task_for_command(const Command& command, const Model& mod if (!command.selected_task.empty()) return {}; if (command.kind == CommandKind::kGenerateVideo) { + if (!has_option(command, "--image") && model.info().bundle_task == TextToAudioVideo::kTask) + return TextToAudioVideo::kTask; if (has_option(command, "--image") || model.info().bundle_task == InitialImageTextToVideo::kTask) return InitialImageTextToVideo::kTask; @@ -165,6 +195,16 @@ bool dispatch_sdk_video(const Command& command, const Model& model, std::string_ detail::write_json(output, write_video(result, require_option(command, "--output"))); return true; } + if (command.kind == CommandKind::kGenerateVideo && id == TextToAudioVideo::kTask) { + if (has_option(command, "--image") || has_option(command, "--initial-latents-raw")) + throw std::invalid_argument("TextToAudioVideo accepts only --prompt and Task config"); + const auto task = model.task(); + const auto config = + detail::task_config(command, task.config_fields(), {"--prompt", "--output"}); + const auto result = task.run({require_option(command, "--prompt")}, config); + detail::write_json(output, write_audio_video(result, require_option(command, "--output"))); + return true; + } if (command.kind != CommandKind::kGenerateVideo || id != TextToVideo::kTask) return false; if (has_option(command, "--image")) From 0f919dd2b5b5408b127bb12fd34aa13f3f85e89c Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 15:17:54 -0700 Subject: [PATCH 10/15] feat(ltx2): add the LTX-2.5 text-to-audio-video family Onboard Lightricks LTX-2.5 (diffusers LTX2Pipeline) as the ltx2 family. A bundle generates a video and its 48 kHz stereo soundtrack with the distilled transformer's 8-step schedule, on one GPU or with the DiT context parallel over two GPUs. Builders (TensorRT network API, bf16 strongly typed with fp32 norm, RoPE and activation islands): - text encoder: the Gemma 4 text tower plus the LTX-2 text connectors that produce the video and audio contexts; - denoiser: the joint audio/video DiT, including the audio-video cross attention and gated attention. With context_parallel_size=2, one plan serves both ranks. Each rank owns half of the video tokens, video self-attention all-gathers the normed and rotated keys and values, and video-to-audio attention merges per-rank softmax statistics through one fp32 all-gather. Audio and text stay replicated. The graph uses no all-to-all collective; - video VAE decoder and the audio VAE decoder with the bandwidth-extension vocoder. model.build validates the request before loading any builder. It accepts only task text_to_audio_video, bf16, batch 1 and CP 1 or 2. With backend=trt_rtx and CP > 1 it requires TensorRT-RTX >= 1.7.1, because 1.6.x has no multi-device support. The C++ runtime implements ITextToAudioVideo. It tokenizes the prompt, draws the seeded noise, runs the Euler loop and decodes the video and audio on rank 0; the other ranks return the worker completion. TRTMC_LTX2_PROGRESS=1 prints per-step progress. Tests: tiny-random parity against diffusers for each engine, a 2-rank context-parallel DiT check with torch-free ranks, build-request and version-gate contract tests, the support identity test, and a manifest-driven E2E that compares the native CLI with LTX2Pipeline from the same noise. Signed-off-by: Peter Kisfaludi --- families/ltx2/README.md | 107 ++ families/ltx2/__init__.py | 4 + families/ltx2/audio_builder.py | 337 ++++ families/ltx2/checkpoint.py | 75 + families/ltx2/dit_builder.py | 491 ++++++ families/ltx2/graph.py | 271 +++ families/ltx2/layers.py | 188 +++ families/ltx2/model.py | 208 +++ families/ltx2/parallel.py | 85 + families/ltx2/requirements.txt | 10 + families/ltx2/runtime/CMakeLists.txt | 51 + families/ltx2/runtime/bpe_tokenizer.cpp | 1489 +++++++++++++++++ families/ltx2/runtime/distributed_runtime.cpp | 215 +++ families/ltx2/runtime/distributed_runtime.h | 24 + families/ltx2/runtime/pipeline.cpp | 403 +++++ families/ltx2/runtime/pipeline.h | 98 ++ families/ltx2/runtime/plugin.cpp | 83 + families/ltx2/runtime/portable_normal.h | 56 + families/ltx2/runtime/progress_log.h | 77 + families/ltx2/runtime/runtime_config.cpp | 48 + families/ltx2/runtime/runtime_config.h | 29 + families/ltx2/runtime/runtime_math.h | 57 + families/ltx2/runtime/tokenizer.h | 30 + families/ltx2/support.py | 15 + families/ltx2/tests/__init__.py | 2 + families/ltx2/tests/conftest.py | 30 + families/ltx2/tests/cp_tiny_prep.py | 63 + .../ltx2/tests/cpp/test_runtime_contract.cpp | 80 + families/ltx2/tests/dist_dit_cp_check.py | 74 + families/ltx2/tests/dist_helpers.py | 73 + families/ltx2/tests/engine_runner.py | 75 + .../tests/manifests/ltx25-distilled-cp2.json | 21 + .../tests/manifests/ltx25-distilled-l0.json | 22 + families/ltx2/tests/np_engine.py | 82 + families/ltx2/tests/test_audio_parity.py | 121 ++ families/ltx2/tests/test_context_parallel.py | 69 + families/ltx2/tests/test_dit_parity.py | 160 ++ families/ltx2/tests/test_e2e.py | 469 ++++++ families/ltx2/tests/test_model_contract.py | 70 + families/ltx2/tests/test_support.py | 12 + .../ltx2/tests/test_text_encoder_parity.py | 185 ++ families/ltx2/tests/test_vae_parity.py | 91 + .../tests/thresholds/ltx25-distilled-cp2.json | 8 + .../tests/thresholds/ltx25-distilled-l0.json | 8 + families/ltx2/text_encoder_builder.py | 440 +++++ families/ltx2/vae_builder.py | 188 +++ 46 files changed, 6794 insertions(+) create mode 100644 families/ltx2/README.md create mode 100644 families/ltx2/__init__.py create mode 100644 families/ltx2/audio_builder.py create mode 100644 families/ltx2/checkpoint.py create mode 100644 families/ltx2/dit_builder.py create mode 100644 families/ltx2/graph.py create mode 100644 families/ltx2/layers.py create mode 100644 families/ltx2/model.py create mode 100644 families/ltx2/parallel.py create mode 100644 families/ltx2/requirements.txt create mode 100644 families/ltx2/runtime/CMakeLists.txt create mode 100644 families/ltx2/runtime/bpe_tokenizer.cpp create mode 100644 families/ltx2/runtime/distributed_runtime.cpp create mode 100644 families/ltx2/runtime/distributed_runtime.h create mode 100644 families/ltx2/runtime/pipeline.cpp create mode 100644 families/ltx2/runtime/pipeline.h create mode 100644 families/ltx2/runtime/plugin.cpp create mode 100644 families/ltx2/runtime/portable_normal.h create mode 100644 families/ltx2/runtime/progress_log.h create mode 100644 families/ltx2/runtime/runtime_config.cpp create mode 100644 families/ltx2/runtime/runtime_config.h create mode 100644 families/ltx2/runtime/runtime_math.h create mode 100644 families/ltx2/runtime/tokenizer.h create mode 100644 families/ltx2/support.py create mode 100644 families/ltx2/tests/__init__.py create mode 100644 families/ltx2/tests/conftest.py create mode 100644 families/ltx2/tests/cp_tiny_prep.py create mode 100644 families/ltx2/tests/cpp/test_runtime_contract.cpp create mode 100644 families/ltx2/tests/dist_dit_cp_check.py create mode 100644 families/ltx2/tests/dist_helpers.py create mode 100644 families/ltx2/tests/engine_runner.py create mode 100644 families/ltx2/tests/manifests/ltx25-distilled-cp2.json create mode 100644 families/ltx2/tests/manifests/ltx25-distilled-l0.json create mode 100644 families/ltx2/tests/np_engine.py create mode 100644 families/ltx2/tests/test_audio_parity.py create mode 100644 families/ltx2/tests/test_context_parallel.py create mode 100644 families/ltx2/tests/test_dit_parity.py create mode 100644 families/ltx2/tests/test_e2e.py create mode 100644 families/ltx2/tests/test_model_contract.py create mode 100644 families/ltx2/tests/test_support.py create mode 100644 families/ltx2/tests/test_text_encoder_parity.py create mode 100644 families/ltx2/tests/test_vae_parity.py create mode 100644 families/ltx2/tests/thresholds/ltx25-distilled-cp2.json create mode 100644 families/ltx2/tests/thresholds/ltx25-distilled-l0.json create mode 100644 families/ltx2/text_encoder_builder.py create mode 100644 families/ltx2/vae_builder.py diff --git a/families/ltx2/README.md b/families/ltx2/README.md new file mode 100644 index 0000000000..05fc8e6cba --- /dev/null +++ b/families/ltx2/README.md @@ -0,0 +1,107 @@ +# LTX-2.5 (`ltx2`) + +Text-to-audio-video for Lightricks LTX-2.5 diffusers checkpoints (`LTX2Pipeline`, for example +`Lightricks/LTX-2.5-Diffusers`). One bundle generates a video and its 48 kHz stereo soundtrack. + +## Scope + +- Task `text_to_audio_video`, precision `bf16`, batch 1. +- The distilled transformer with its 8-step schedule (guidance 1, no STG). The bundle bakes + the schedule into `runtime.json`. The full transformer (30 steps, CFG/STG/modality + guidance) is not supported yet. +- The video size and frame count are fixed at build time (`--image-height`, `--image-width`, + `--video-num-frames`). Height and width must be multiples of 32, and the frame count must + be `8n+1`. The default is 960x544, 121 frames at 24 fps. +- `--context-parallel-size 1` runs on one GPU. `--context-parallel-size 2` splits the + video tokens of the DiT across two GPUs. The text encoder, video VAE and audio decoder + run on rank 0 only. + +## Bundle + +| Section | Contents | +|---|---| +| `text_encoder.plan` | Gemma 4 text tower and the LTX-2 text connectors (video and audio context) | +| `denoiser.plan` | Joint audio/video DiT. With CP=2, one plan serves both ranks. | +| `vae.plan` | Video VAE decoder | +| `audio.plan` | Audio VAE decoder and vocoder with bandwidth extension | +| `tokenizer.json`, `runtime.json` | Tokenizer, shapes and schedule | + +Context parallelism keeps the audio stream and text replicated and shards the video +tokens. Video self-attention all-gathers each rank's normed and rotated keys and values. +Video-to-audio attention merges per-rank softmax statistics through one small all-gather. +The network uses no all-to-all collective. + +## Build and run + +```bash +python -m tensorrt_model_connect build /models/LTX-2.5-Diffusers --family ltx2 \ + --backend trt_rtx --precision bf16 \ + --image-width 960 --image-height 544 --video-num-frames 121 \ + --context-parallel-size 2 -o ltx25-cp2.bundle + +trtmc generate-video ltx25-cp2.bundle --runtime-root \ + --prompt "A red fox walking through a snowy forest at dawn" \ + --output out --set seed=42 +``` + +`generate-video` writes `out/frame-NNNNNN.png` and `out/audio.wav`. Set `TRTMC_LTX2_PROGRESS=1` +to print one progress line per stage and denoising step. + +## TensorRT-RTX + +- Context parallelism with `--backend trt_rtx` requires TensorRT-RTX 1.7.1 or newer. + TensorRT-RTX 1.6.x has no multi-device support, so the build stops with an error + for `--context-parallel-size 2` when an older version is installed. +- Build and run with the same TensorRT-RTX version. A bundle's plans are specific to the + version that built them. +- TRTMC does not pin `tensorrt-rtx` in `pyproject.toml`. Install the Python package that + matches your TensorRT-RTX runtime. + +## Native Windows with two GPUs + +- Put both GPUs in TCC mode. Under WDDM or MCDM they report no peer access and NCCL has + no transport. +- Build `nccl.dll` from source and point the runtime at it with `TRTMC_NCCL_LIBRARY`. + See [Multi-Device Execution](../../website/docs/features/multi-device.md). +- Start the ranks with `tools/launch_ranks.py`, because native Windows has no `mpirun`: + +```powershell +python tools\launch_ranks.py -n 2 --gpus 0,1 --nccl-library C:\nccl\install\bin\nccl.dll ` + --library-dir ` + -- trtmc generate-video ltx25-cp2.bundle --runtime-root ` + --prompt "..." --output out --set seed=42 +``` + +To give each rank its own TensorRT-RTX runtime cache, put `{rank}` in the path, for example +`--runtime-cache cache-rank{rank}.bin`. + +## Tests + +- `tests/test_*_parity.py` build each engine from tiny random weights and compare it + with diffusers. +- `tests/test_context_parallel.py` runs the CP=2 DiT on two GPUs (torch-free ranks) against + the single-device plan and diffusers. It needs `TRTMC_NCCL_LIBRARY`. +- `tests/test_e2e.py` builds the real checkpoint and runs the native CLI. It compares the + output with `LTX2Pipeline` started from the same noise. Select it with `--e2e-model ltx2`. + +### E2E thresholds + +Both sides run the 48-layer DiT in bf16. A single forward of the native engine differs from +an fp32 diffusers forward by about 3% relative L2, as does diffusers in bf16. The 8 distilled +steps amplify that difference into detail-level changes of the same scene. Exact pixel +agreement is therefore not expected. The thresholds separate "same trajectory" from +"different or broken output". They were set from these measurements on 2x RTX PRO 6000 with +TensorRT-RTX 1.7.1: + +| Comparison (seed 42) | Frame PSNR mean / min | Log-spectrogram corr | RMS ratio | +|---|---|---|---| +| Native vs diffusers bf16, 384x640x49 (CP=1) | 19.1 / 18.4 dB | 0.908 | 1.00 | +| Native vs diffusers bf16, 960x544x121 (CP=2) | 18.2 / 16.2 dB | 0.892 | 1.13 | +| Diffusers bf16 vs diffusers fp32 (both sizes) | 22.7-23.8 / 21.9-22.9 dB | 0.92 | 1.00 | +| Diffusers bf16, seed 42 vs seed 7 (unrelated output) | 11.5-12.9 / 10.1-12.2 dB | 0.75-0.81 | 0.25-5.9 | + +The thresholds are a mean of at least 15 dB, a minimum of at least 14 dB, a correlation of +at least 0.85 and an RMS ratio in [0.8, 1.25]. Each one sits between the measured native +values and the unrelated-output values. The audio correlation is lower than the video +agreement because the released pipeline runs its vocoder in bf16; against an fp32 vocoder, +the native audio reaches 0.95-0.97. diff --git a/families/ltx2/__init__.py b/families/ltx2/__init__.py new file mode 100644 index 0000000000..c2c518715d --- /dev/null +++ b/families/ltx2/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 (Lightricks LTX2Pipeline) audio + video family.""" diff --git a/families/ltx2/audio_builder.py b/families/ltx2/audio_builder.py new file mode 100644 index 0000000000..cd1e85aa79 --- /dev/null +++ b/families/ltx2/audio_builder.py @@ -0,0 +1,337 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 audio decoder: ``AutoencoderKLLTX2Audio.decoder`` + ``LTX2VocoderWithBWE`` as one TensorRT plan. + +Engine I/O: + Inputs: + audio_latents [1, Sa, 128] fp32 packed, normalized audio latents (DiT layout) + Outputs: + waveform [1, 2, N] fp32 48 kHz stereo in [-1, 1], N = Sa*4-3 mel frames * 160 * 3 + mel [1, 2, T, 64] fp32 (debug builds only) the audio VAE log-mel output + +Everything runs in fp32 (the vocoder's SnakeBeta / log-mel path is precision sensitive and cheap): +- latent denormalization with the audio VAE ``latents_mean`` / ``latents_std``, unpack to ``[1, 8, Sa, 16]``; +- the causal (time axis) Conv2d decoder with pixel norms, nearest x2 upsampling that drops the first + time row, and the crop to ``Sa*4-3`` frames x 64 mel bins; +- stage-1 vocoder (BigVGAN-style: transposed-conv upsamplers, three parallel dilated ResBlocks per + stage averaged, anti-aliased SnakeBeta activations, Kaiser up/down filters with replicate padding); +- causal STFT (conv1d with the checkpoint's windowed DFT basis), magnitude, mel matmul, log clamp; +- BWE generator on that mel plus the Hann sinc x3 resampler of the stage-1 waveform, clamp, crop. +""" + +from __future__ import annotations + +import math +import sys +from pathlib import Path + +import numpy as np +import tensorrt as trt + +from .checkpoint import Checkpoint +from .graph import Graph, build_plan, make_logger, new_network + +F32 = np.float32 + + +# ---------------------------------------------------------------------- 1D / 2D primitives (fp32) + + +def _pad_replicate_last(g: Graph, x, left: int, right: int): + """Replicate padding on the last axis of ``[B, C, L]``.""" + b, c, n = (int(s) for s in x.shape) + parts = [] + if left: + parts += [g.slice(x, (0, 0, 0), (b, c, 1))] * left + parts.append(x) + if right: + parts += [g.slice(x, (0, 0, n - 1), (b, c, 1))] * right + return g.concat(parts, axis=2) if len(parts) > 1 else x + + +def _conv1d(g: Graph, x, weight: np.ndarray, bias: np.ndarray | None, *, stride: int = 1, dilation: int = 1, + padding: int = 0, groups: int = 1): + """``F.conv1d`` on ``[B, C, L]`` (as a 1xK Conv2d).""" + b, c, n = (int(s) for s in x.shape) + out_c, _, k = (int(s) for s in weight.shape) + x4 = g.reshape(x, (b, c, 1, n)) + layer = g.net.add_convolution_nd(x4, out_c, (1, k), g.weights(weight.reshape(out_c, -1, 1, k), trt.float32), + g.weights(bias, trt.float32) if bias is not None else trt.Weights()) + layer.stride_nd = (1, stride) + layer.dilation_nd = (1, dilation) + layer.padding_nd = (0, padding) + layer.num_groups = groups + y = layer.get_output(0) + return g.reshape(y, (b, out_c, int(y.shape[3]))) + + +def _conv_transpose1d(g: Graph, x, weight: np.ndarray, bias: np.ndarray | None, *, stride: int, padding: int = 0, + groups: int = 1): + """``F.conv_transpose1d`` on ``[B, C, L]``; ``weight`` is ``[C_in, C_out / groups, K]``.""" + b, c, n = (int(s) for s in x.shape) + cin, cout_g, k = (int(s) for s in weight.shape) + out_c = cout_g * groups + x4 = g.reshape(x, (b, c, 1, n)) + layer = g.net.add_deconvolution_nd(x4, out_c, (1, k), g.weights(weight.reshape(cin, cout_g, 1, k), trt.float32), + g.weights(bias, trt.float32) if bias is not None else trt.Weights()) + layer.stride_nd = (1, stride) + layer.padding_nd = (0, padding) + layer.num_groups = groups + y = layer.get_output(0) + return g.reshape(y, (b, out_c, int(y.shape[3]))) + + +def _snake_beta(g: Graph, x, alpha: np.ndarray, beta: np.ndarray, eps: float = 1e-9): + """``x + 1 / (exp(beta) + eps) * sin(x * exp(alpha))^2`` (logscale SnakeBeta).""" + c = int(x.shape[1]) + a = g.const(np.exp(alpha.astype(np.float64)).astype(F32).reshape(1, c, 1), trt.float32) + inv_b = g.const((1.0 / (np.exp(beta.astype(np.float64)) + eps)).astype(F32).reshape(1, c, 1), trt.float32) + s = g.unary(g.mul(x, a), trt.UnaryOperation.SIN) + return g.add(x, g.mul(inv_b, g.mul(s, s))) + + +def _kaiser_up(g: Graph, x, filt: np.ndarray, ratio: int): + k = filt.size + pad = k // ratio - 1 + pad_left = pad * ratio + (k - ratio) // 2 + pad_right = pad * ratio + (k - ratio + 1) // 2 + c = int(x.shape[1]) + xp = _pad_replicate_last(g, x, pad, pad) + w = np.tile(filt.reshape(1, 1, k), (c, 1, 1)).astype(F32) * F32(ratio) + y = _conv_transpose1d(g, xp, w, None, stride=ratio, groups=c) + n = int(y.shape[2]) + return g.slice(y, (0, 0, pad_left), (int(y.shape[0]), c, n - pad_left - pad_right)) + + +def _kaiser_down(g: Graph, x, filt: np.ndarray, ratio: int): + k = filt.size + pad_left = k // 2 + (k % 2) - 1 + pad_right = k // 2 + c = int(x.shape[1]) + xp = _pad_replicate_last(g, x, pad_left, pad_right) + w = np.tile(filt.reshape(1, 1, k), (c, 1, 1)).astype(F32) + return _conv1d(g, xp, w, None, stride=ratio, groups=c) + + +def _aa_act(g: Graph, ck: Checkpoint, p: str, x, ratio: int): + """``AntiAliasAct1d(SnakeBeta)``: Kaiser x2 up, SnakeBeta, Kaiser x2 down.""" + x = _kaiser_up(g, x, ck.get(f"{p}.upsample.filter", F32).reshape(-1), ratio) + x = _snake_beta(g, x, ck.get(f"{p}.act.alpha", F32), ck.get(f"{p}.act.beta", F32)) + return _kaiser_down(g, x, ck.get(f"{p}.downsample.filter", F32).reshape(-1), ratio) + + +def _resblock(g: Graph, ck: Checkpoint, p: str, x, kernel: int, dilations, ratio: int): + for j, d in enumerate(dilations): + xt = _aa_act(g, ck, f"{p}.acts1.{j}", x, ratio) + xt = _conv1d(g, xt, ck.get(f"{p}.convs1.{j}.weight", F32), ck.maybe(f"{p}.convs1.{j}.bias", F32), + dilation=d, padding=d * (kernel - 1) // 2) + xt = _aa_act(g, ck, f"{p}.acts2.{j}", xt, ratio) + xt = _conv1d(g, xt, ck.get(f"{p}.convs2.{j}.weight", F32), ck.maybe(f"{p}.convs2.{j}.bias", F32), + padding=(kernel - 1) // 2) + x = g.add(x, xt) + return x + + +def _vocoder_stage(g: Graph, ck: Checkpoint, p: str, mel, cfg: dict, *, prefix_cfg: str = ""): + """``LTX2Vocoder.forward`` (snakebeta + antialias variant) on ``mel`` ``[B, C, T, M]``.""" + def c(key): + return cfg[prefix_cfg + key] + + if c("act_fn") not in ("snakebeta", "snake") or not c("antialias"): + raise NotImplementedError("only the anti-aliased SnakeBeta vocoder is implemented") + b, ch, t, m = (int(s) for s in mel.shape) + x = g.reshape(g.transpose(mel, (0, 1, 3, 2)), (b, ch * m, t)) + x = _conv1d(g, x, ck.get(f"{p}.conv_in.weight", F32), ck.maybe(f"{p}.conv_in.bias", F32), padding=3) + ratio = int(c("antialias_ratio")) + kernels = c("resnet_kernel_sizes") + dils = c("resnet_dilations") + n_res = len(kernels) + for i, (stride, k) in enumerate(zip(c("upsample_factors"), c("upsample_kernel_sizes"))): + x = _conv_transpose1d(g, x, ck.get(f"{p}.upsamplers.{i}.weight", F32), ck.maybe(f"{p}.upsamplers.{i}.bias", F32), + stride=int(stride), padding=(int(k) - int(stride)) // 2) + outs = [_resblock(g, ck, f"{p}.resnets.{i * n_res + j}", x, int(kernels[j]), dils[j], ratio) + for j in range(n_res)] + acc = outs[0] + for o in outs[1:]: + acc = g.add(acc, o) + x = g.mul(acc, g.scalar(1.0 / n_res, trt.float32, 3)) + x = _aa_act(g, ck, f"{p}.act_out", x, ratio) + x = _conv1d(g, x, ck.get(f"{p}.conv_out.weight", F32), ck.maybe(f"{p}.conv_out.bias", F32), padding=3) + final = c("final_act_fn") + if final == "tanh": + x = g.net.add_activation(x, trt.ActivationType.TANH).get_output(0) + elif final == "clamp": + x = g.maximum(g.minimum(x, g.scalar(1.0, trt.float32, 3)), g.scalar(-1.0, trt.float32, 3)) + return x + + +def hann_resampler_filter(ratio: int) -> tuple[np.ndarray, int, int, int]: + """diffusers ``UpSample1d(window_type="hann")``: (filter, pad, pad_left, pad_right).""" + rolloff = 0.99 + lowpass_filter_width = 6 + width = math.ceil(lowpass_filter_width / rolloff) + k = 2 * width * ratio + 1 + t = (np.arange(k, dtype=F32) / F32(ratio) - F32(width)) * F32(rolloff) + tc = np.clip(t, -lowpass_filter_width, lowpass_filter_width) + window = np.cos(tc * F32(math.pi) / F32(lowpass_filter_width) / F32(2)) ** 2 + filt = (np.sinc(t) * window * F32(rolloff) / F32(ratio)).astype(F32) + return filt, width, 2 * width * ratio, k - ratio + + +def _bwe_vocoder(g: Graph, ck: Checkpoint, cfg: dict, mel, *, debug: bool = False): + x = _vocoder_stage(g, ck, "vocoder", mel, cfg) # [1, 2, n] at 16 kHz + if debug: + g.mark_output(g.cast(x, trt.float32), "stage1", trt.float32) + b, ch, n = (int(s) for s in x.shape) + hop = int(cfg["hop_length"]) + if n % hop: + x = g.concat([x, g.const(np.zeros((b, ch, hop - n % hop), F32), trt.float32)], axis=2) + n_pad = int(x.shape[2]) + # MelSTFT on every channel: causal left zero pad, windowed DFT conv, magnitude, mel, log clamp. + win = int(cfg["window_length"]) + left = max(0, win - hop) + w = g.reshape(x, (b * ch, 1, n_pad)) + w = g.concat([g.const(np.zeros((b * ch, 1, left), F32), trt.float32), w], axis=2) + basis = ck.get("mel_stft.stft_fn.forward_basis", F32) + spec = _conv1d(g, w, basis, None, stride=hop) # [B*C, 2*nf, frames] + nf = basis.shape[0] // 2 + frames = int(spec.shape[2]) + re = g.slice(spec, (0, 0, 0), (b * ch, nf, frames)) + im = g.slice(spec, (0, nf, 0), (b * ch, nf, frames)) + mag = g.unary(g.add(g.mul(re, re), g.mul(im, im)), trt.UnaryOperation.SQRT) + mel_basis = ck.get("mel_stft.mel_basis", F32) # [n_mels, nf] + mel2 = g.net.add_matrix_multiply(g.const(mel_basis.reshape(1, *mel_basis.shape), trt.float32), + trt.MatrixOperation.NONE, mag, trt.MatrixOperation.NONE).get_output(0) + log_mel = g.unary(g.maximum(mel2, g.scalar(1e-5, trt.float32, 3)), trt.UnaryOperation.LOG) + n_mels = int(mel_basis.shape[0]) + mel_bwe = g.transpose(g.reshape(log_mel, (b, ch, n_mels, frames)), (0, 1, 3, 2)) # [B, C, frames, mels] + residual = _vocoder_stage(g, ck, "bwe_generator", mel_bwe, cfg, prefix_cfg="bwe_") + ratio = int(cfg["output_sampling_rate"]) // int(cfg["input_sampling_rate"]) + filt, pad, pad_left, pad_right = hann_resampler_filter(ratio) + xp = _pad_replicate_last(g, x, pad, pad) + k = filt.size + skip = _conv_transpose1d(g, xp, np.tile(filt.reshape(1, 1, k), (ch, 1, 1)) * F32(ratio), None, stride=ratio, + groups=ch) + skip = g.slice(skip, (0, 0, pad_left), (b, ch, int(skip.shape[2]) - pad_left - pad_right)) + out = g.add(residual, skip) + out = g.maximum(g.minimum(out, g.scalar(1.0, trt.float32, 3)), g.scalar(-1.0, trt.float32, 3)) + n_out = n * int(cfg["output_sampling_rate"]) // int(cfg["input_sampling_rate"]) + return g.slice(out, (0, 0, 0), (b, ch, n_out)) + + +# ---------------------------------------------------------------------- audio VAE decoder + + +def _pixel_norm(g: Graph, x, eps: float): + ms = g.reduce(g.mul(x, x), trt.ReduceOperation.AVG, 1) + inv = g.unary(g.unary(g.add(ms, g.scalar(eps, trt.float32, 4)), trt.UnaryOperation.SQRT), + trt.UnaryOperation.RECIP) + return g.mul(x, inv) + + +def _causal_conv2d(g: Graph, x, weight: np.ndarray, bias: np.ndarray | None, axis: str): + out_c, _, kh, kw = (int(s) for s in weight.shape) + ph, pw = kh - 1, kw - 1 + if axis == "height": + pre, post = (ph, pw // 2), (0, pw - pw // 2) + elif axis == "none": + pre, post = (ph // 2, pw // 2), (ph - ph // 2, pw - pw // 2) + else: + raise NotImplementedError(f"audio VAE causality_axis={axis!r}") + layer = g.net.add_convolution_nd(x, out_c, (kh, kw), g.weights(weight, trt.float32), + g.weights(bias, trt.float32) if bias is not None else trt.Weights()) + layer.pre_padding = pre + layer.post_padding = post + return layer.get_output(0) + + +def _audio_resnet(g: Graph, ck: Checkpoint, p: str, x, axis: str, eps: float): + h = g.silu(_pixel_norm(g, x, eps)) + h = _causal_conv2d(g, h, ck.get(f"{p}.conv1.conv.weight", F32), ck.maybe(f"{p}.conv1.conv.bias", F32), axis) + h = g.silu(_pixel_norm(g, h, eps)) + h = _causal_conv2d(g, h, ck.get(f"{p}.conv2.conv.weight", F32), ck.maybe(f"{p}.conv2.conv.bias", F32), axis) + if ck.has(f"{p}.nin_shortcut.conv.weight"): + x = _causal_conv2d(g, x, ck.get(f"{p}.nin_shortcut.conv.weight", F32), + ck.maybe(f"{p}.nin_shortcut.conv.bias", F32), axis) + elif ck.has(f"{p}.conv_shortcut.conv.weight"): + x = _causal_conv2d(g, x, ck.get(f"{p}.conv_shortcut.conv.weight", F32), + ck.maybe(f"{p}.conv_shortcut.conv.bias", F32), axis) + return g.add(x, h) + + +def _nearest2x(g: Graph, x): + b, c, h, w = (int(s) for s in x.shape) + x = g.gather(x, g.const(np.arange(2 * h, dtype=np.int32) // 2, trt.int32), 2) + return g.gather(x, g.const(np.arange(2 * w, dtype=np.int32) // 2, trt.int32), 3) + + +def add_audio_vae_decoder(g: Graph, ck: Checkpoint, cfg: dict, z): + """``LTX2AudioDecoder.forward`` on ``z`` ``[1, 8, frames, mel/4]`` -> mel ``[1, out_ch, T, mel_bins]``.""" + if cfg.get("norm_type", "pixel") != "pixel" or cfg.get("mid_block_add_attention") or cfg.get("attn_resolutions"): + raise NotImplementedError("only the pixel-norm, attention-free LTX-2 audio decoder is implemented") + axis = cfg.get("causality_axis", "height") + eps = 1e-6 + ch_mult = list(cfg["ch_mult"]) + frames = int(z.shape[2]) + target_t = frames * 4 + if axis is not None: + target_t = max(target_t - 3, 1) + target_m = int(cfg.get("mel_bins") or z.shape[3]) + out_ch = int(cfg["output_channels"]) + x = _causal_conv2d(g, z, ck.get("decoder.conv_in.conv.weight", F32), ck.maybe("decoder.conv_in.conv.bias", F32), + axis) + x = _audio_resnet(g, ck, "decoder.mid.block_1", x, axis, eps) + x = _audio_resnet(g, ck, "decoder.mid.block_2", x, axis, eps) + for level in reversed(range(len(ch_mult))): + for bi in range(int(cfg["num_res_blocks"]) + 1): + x = _audio_resnet(g, ck, f"decoder.up.{level}.block.{bi}", x, axis, eps) + if level != 0: + x = _nearest2x(g, x) + x = _causal_conv2d(g, x, ck.get(f"decoder.up.{level}.upsample.conv.conv.weight", F32), + ck.maybe(f"decoder.up.{level}.upsample.conv.conv.bias", F32), axis) + if axis == "height": + b, c, h, w = (int(s) for s in x.shape) + x = g.slice(x, (0, 0, 1, 0), (b, c, h - 1, w)) + x = g.silu(_pixel_norm(g, x, eps)) + x = _causal_conv2d(g, x, ck.get("decoder.conv_out.conv.weight", F32), ck.maybe("decoder.conv_out.conv.bias", F32), + axis) + b, c, t, m = (int(s) for s in x.shape) + x = g.slice(x, (0, 0, 0, 0), (b, min(c, out_ch), min(t, target_t), min(m, target_m))) + if t < target_t or m < target_m: + x = g.net.add_padding_nd(x, (0, 0), (max(target_t - t, 0), max(target_m - m, 0))).get_output(0) + return x + + +# ---------------------------------------------------------------------- engine + + +def build_audio_decoder_engine(model_dir: str | Path, *, audio_frames: int, debug_mel: bool = False, + verbose: bool = False, tf32: bool = True): + model_dir = Path(model_dir) + vae = Checkpoint(model_dir / "audio_vae") + voc = Checkpoint(model_dir / "vocoder") + vcfg, ocfg = vae.config(), voc.config() + latent_c = int(vcfg["latent_channels"]) + mel_bins = int(vcfg["mel_bins"]) + latent_m = mel_bins // 4 + packed = latent_c * latent_m + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + z = network.add_input("audio_latents", trt.float32, (1, audio_frames, packed)) + mean = vae.get("latents_mean", F32).reshape(1, 1, -1) + std = vae.get("latents_std", F32).reshape(1, 1, -1) + if mean.shape[-1] != packed: + raise ValueError(f"audio latent statistics have {mean.shape[-1]} channels, packed latents {packed}") + x = g.add(g.mul(z, g.const(std, trt.float32)), g.const(mean, trt.float32)) + # [1, L, C*M] -> unflatten(2, (C, M)) -> transpose(1, 2) -> [1, C, L, M] + x = g.reshape(x, (1, audio_frames, latent_c, latent_m)) + x = g.transpose(x, (0, 2, 1, 3)) + mel = add_audio_vae_decoder(g, vae, vcfg, x) + if debug_mel: + g.mark_output(mel, "mel", trt.float32) + wave = _bwe_vocoder(g, voc, ocfg, mel, debug=debug_mel) + g.mark_output(wave, "waveform", trt.float32) + print(f"[ltx2] Building audio decoder engine (audio latents {audio_frames} -> mel {int(mel.shape[2])} frames -> " + f"{int(wave.shape[2])} samples @ {ocfg['output_sampling_rate']} Hz) ...", file=sys.stderr) + return build_plan(builder, network, label="audio decoder", tf32=tf32) diff --git a/families/ltx2/checkpoint.py b/families/ltx2/checkpoint.py new file mode 100644 index 0000000000..22ec3e5457 --- /dev/null +++ b/families/ltx2/checkpoint.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Safetensors access for the LTX-2.5 diffusers checkpoint folders.""" + +from __future__ import annotations + +import json +from pathlib import Path + +import ml_dtypes # noqa: F401 - registers the bf16 NumPy dtype used by safetensors +import numpy as np +from safetensors import safe_open + +BF16 = ml_dtypes.bfloat16 + +_INDEX_NAMES = ("model.safetensors.index.json", "diffusion_pytorch_model.safetensors.index.json") +_SINGLE_NAMES = ("model.safetensors", "diffusion_pytorch_model.safetensors") + + +class Checkpoint: + """Name -> tensor access over one (possibly sharded) safetensors folder. + + ``get`` converts to the requested NumPy dtype with round-to-nearest-even, the + same rounding torch applies when a pipeline loads a fp32 tensor as bf16. + """ + + def __init__(self, folder: str | Path): + self.folder = Path(folder) + self._readers: dict[str, object] = {} + self._map: dict[str, str] = {} + for name in _INDEX_NAMES: + index = self.folder / name + if index.is_file(): + weight_map = json.loads(index.read_text(encoding="utf-8"))["weight_map"] + self._map = {str(k): str(v) for k, v in weight_map.items()} + break + else: + for name in _SINGLE_NAMES: + single = self.folder / name + if single.is_file(): + reader = self._open(name) + self._map = {k: name for k in reader.keys()} + break + else: + raise FileNotFoundError(f"no safetensors checkpoint in {self.folder}") + + def _open(self, shard: str): + reader = self._readers.get(shard) + if reader is None: + reader = safe_open(str(self.folder / shard), framework="numpy") + self._readers[shard] = reader + return reader + + def keys(self) -> list[str]: + return list(self._map) + + def has(self, name: str) -> bool: + return name in self._map + + def get(self, name: str, dtype=BF16) -> np.ndarray: + shard = self._map.get(name) + if shard is None: + raise KeyError(f"{self.folder.name}: tensor not found: {name}") + arr = self._open(shard).get_tensor(name) + if dtype is not None and arr.dtype != np.dtype(dtype): + arr = arr.astype(dtype) + return np.ascontiguousarray(arr) + + def maybe(self, name: str, dtype=BF16) -> np.ndarray | None: + return self.get(name, dtype) if self.has(name) else None + + def config(self) -> dict: + path = self.folder / "config.json" + return json.loads(path.read_text(encoding="utf-8")) if path.is_file() else {} diff --git a/families/ltx2/dit_builder.py b/families/ltx2/dit_builder.py new file mode 100644 index 0000000000..c7ca945c56 --- /dev/null +++ b/families/ltx2/dit_builder.py @@ -0,0 +1,491 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 joint audio + video DiT (``LTX2VideoTransformer3DModel``) as a TensorRT plan. + +One builder serves the single-device plan (``cp_size=1``) and the rank-dynamic +context-parallel plan (``cp_size>1``, one serialized plan shared by every rank). + +Engine I/O (``B`` = batch of guidance branches; 1 for the distilled model): + Inputs: + video_latent [B, S, 128] fp32 packed video latents (full sequence on every rank) + audio_latent [B, Sa, 128] fp32 packed audio latents + video_context [B, L, 4096] bf16 video connector output + audio_context [B, L, 2048] bf16 audio connector output + timestep [B] fp32 sigma * 1000 (also used as ``sigma``, cross timestep) + stg_keep [B] fp32 1 = normal; 0 = spatio-temporal guidance branch + (self-attention context -> V in the STG blocks) + av_keep [B] fp32 1 = normal; 0 = modality-isolated branch + (a2v / v2a residual adds are multiplied by 0) + Outputs: + video_velocity [B, S, 128] fp32 (full sequence on every rank) + audio_velocity [B, Sa, 128] fp32 + +Precision: bf16 strongly typed (diffusers runs this model in bf16), with fp32 RMSNorm / +LayerNorm statistics, fp32 split RoPE, fp32 timestep sinusoids and fp32 GELU/SiLU islands. + +Context parallelism (``cp_size`` ranks, contiguous token shards of ``S / cp`` rows): + - every rank derives its index with one tiny REDUCE_SCATTER and gathers its own latent + rows and RoPE rows; + - video self-attention: local queries attend all keys/values. Q/K RMSNorm spans all heads + of a token and RoPE is per token, so K and V are normed and rotated on their own rows and + then ALL_GATHERed. For two ranks this moves the same bytes as a Ulysses head exchange and + needs no ALL_TO_ALL. bf16 payloads cross the wire as fp16 (exact inside the fp16 range, + which covers every activation of the reference run; see ``_ag_tokens``); + - video -> text and audio -> video (a2v) use local queries against replicated keys; + - video -> audio (v2a) needs every video key: each rank computes its partial softmax + statistics (row max, row sum, unnormalized context) over its key shard in fp32, one + ALL_GATHER (~1 MB) shares them and every rank merges them identically (exact + log-sum-exp merge, as in flash-decoding); + - the audio stream (126 tokens) is replicated and computed redundantly on every rank; + - the video output rows are restored with one fp32 ALL_GATHER. +""" + +from __future__ import annotations + +import json +import sys +from dataclasses import dataclass +from pathlib import Path + +import numpy as np +import tensorrt as trt + +from .checkpoint import Checkpoint +from .graph import Graph, build_plan, make_logger, new_network +from .layers import ( + AttnWeights, + RopeTables, + feed_forward, + from_heads, + gate_and_project_out, + gather_rope_rows, + ltx_attention, + modulate, + project_qkv, + rope_constants, + split_rope_freqs, + stg_lerp, + to_heads, +) +from .parallel import add_collective, local_row_indices + +EPS = 1e-6 +_FP16_SAFE_BF16_MAX = 65280.0 # largest bf16 value that is finite in fp16 + + +@dataclass(frozen=True) +class DiTConfig: + in_channels: int + out_channels: int + heads: int + head_dim: int + audio_in_channels: int + audio_out_channels: int + audio_heads: int + audio_head_dim: int + cross_attention_dim: int + audio_cross_attention_dim: int + layers: int + cross_attn_mod: bool + audio_cross_attn_mod: bool + prompt_adaln: bool + gated: bool + audio_gated: bool + rope_theta: float + rope_double_precision: bool + pos_embed_max_pos: int + audio_pos_embed_max_pos: int + base_height: int + base_width: int + vae_scale_factors: tuple[int, int, int] + causal_offset: int + audio_sampling_rate: int + audio_hop_length: int + audio_scale_factor: int + timestep_scale_multiplier: float + cross_attn_timestep_scale_multiplier: float + + @property + def dim(self) -> int: + return self.heads * self.head_dim + + @property + def audio_dim(self) -> int: + return self.audio_heads * self.audio_head_dim + + @staticmethod + def from_dict(c: dict) -> "DiTConfig": + for key, expected in (("rope_type", "split"), ("qk_norm", "rms_norm_across_heads"), + ("patch_size", 1), ("patch_size_t", 1), ("audio_patch_size", 1), + ("audio_patch_size_t", 1), ("norm_elementwise_affine", False), + ("use_prompt_embeddings", False), ("activation_fn", "gelu-approximate")): + if c.get(key, expected) != expected: + raise NotImplementedError(f"LTX-2.5 DiT builder expects {key}={expected!r}, got {c.get(key)!r}") + return DiTConfig( + in_channels=int(c["in_channels"]), out_channels=int(c.get("out_channels") or c["in_channels"]), + heads=int(c["num_attention_heads"]), head_dim=int(c["attention_head_dim"]), + audio_in_channels=int(c["audio_in_channels"]), + audio_out_channels=int(c.get("audio_out_channels") or c["audio_in_channels"]), + audio_heads=int(c["audio_num_attention_heads"]), audio_head_dim=int(c["audio_attention_head_dim"]), + cross_attention_dim=int(c["cross_attention_dim"]), + audio_cross_attention_dim=int(c["audio_cross_attention_dim"]), + layers=int(c["num_layers"]), + cross_attn_mod=bool(c.get("cross_attn_mod", False)), + audio_cross_attn_mod=bool(c.get("audio_cross_attn_mod", False)), + prompt_adaln=bool(c.get("use_prompt_adaln_single", True)), + gated=bool(c.get("gated_attn", False)), audio_gated=bool(c.get("audio_gated_attn", False)), + rope_theta=float(c.get("rope_theta", 10000.0)), + rope_double_precision=bool(c.get("rope_double_precision", True)), + pos_embed_max_pos=int(c.get("pos_embed_max_pos", 20)), + audio_pos_embed_max_pos=int(c.get("audio_pos_embed_max_pos", 20)), + base_height=int(c.get("base_height", 2048)), base_width=int(c.get("base_width", 2048)), + vae_scale_factors=tuple(int(v) for v in c.get("vae_scale_factors", (8, 32, 32))), + causal_offset=int(c.get("causal_offset", 1)), + audio_sampling_rate=int(c.get("audio_sampling_rate", 16000)), + audio_hop_length=int(c.get("audio_hop_length", 160)), + audio_scale_factor=int(c.get("audio_scale_factor", 4)), + timestep_scale_multiplier=float(c.get("timestep_scale_multiplier", 1000)), + cross_attn_timestep_scale_multiplier=float(c.get("cross_attn_timestep_scale_multiplier", 1000)), + ) + + +@dataclass(frozen=True) +class DiTShape: + batch: int + latent_frames: int + latent_height: int + latent_width: int + audio_frames: int + text_len: int + fps: float + + @property + def video_tokens(self) -> int: + return self.latent_frames * self.latent_height * self.latent_width + + +def audio_latent_frames(num_frames: int, fps: float, *, sampling_rate: int = 16000, hop_length: int = 160, + temporal_compression: int = 4) -> int: + """``LTX2Pipeline``: ``round(num_frames / fps * sr / hop / compression)``.""" + return int(round(num_frames / float(fps) * (float(sampling_rate) / float(hop_length) / + float(temporal_compression)))) + + +# ---------------------------------------------------------------------- RoPE grids (diffusers, fp32) + + +def video_midpoints(cfg: DiTConfig, shape: DiTShape) -> np.ndarray: + """``prepare_video_coords`` patch midpoints: ``[3, S]`` fp32 (time in seconds, space in pixels).""" + f32 = np.float32 + gf, gh, gw = np.meshgrid(np.arange(shape.latent_frames, dtype=f32), np.arange(shape.latent_height, dtype=f32), + np.arange(shape.latent_width, dtype=f32), indexing="ij") + start = np.stack([gf, gh, gw], 0).reshape(3, -1) + end = start + f32(1.0) + scale = np.asarray(cfg.vae_scale_factors, dtype=f32).reshape(3, 1) + ps, pe = start * scale, end * scale + t = f32(cfg.vae_scale_factors[0]) + ps[0] = np.maximum(ps[0] + f32(cfg.causal_offset) - t, f32(0.0)) / f32(shape.fps) + pe[0] = np.maximum(pe[0] + f32(cfg.causal_offset) - t, f32(0.0)) / f32(shape.fps) + return ((ps + pe) / f32(2.0)).astype(f32) + + +def video_grid(cfg: DiTConfig, shape: DiTShape) -> np.ndarray: + """Video self-attention RoPE positions: midpoints / (max frames, base height, base width): ``[S, 3]``.""" + maxpos = np.asarray([cfg.pos_embed_max_pos, cfg.base_height, cfg.base_width], dtype=np.float32).reshape(3, 1) + return (video_midpoints(cfg, shape) / maxpos).T.astype(np.float32) + + +def audio_grid(cfg: DiTConfig, shape: DiTShape, max_pos: int) -> np.ndarray: + """``prepare_audio_coords`` + midpoint + division: ``[Sa, 1]`` fp32 (seconds / max_pos).""" + f32 = np.float32 + gf = np.arange(shape.audio_frames, dtype=f32) + sf = f32(cfg.audio_scale_factor) + start = np.maximum(gf * sf + f32(cfg.causal_offset) - sf, f32(0.0)) + start = start * f32(cfg.audio_hop_length) / f32(cfg.audio_sampling_rate) + end = np.maximum((gf + f32(1.0)) * sf + f32(cfg.causal_offset) - sf, f32(0.0)) + end = end * f32(cfg.audio_hop_length) / f32(cfg.audio_sampling_rate) + mid = (start + end) / f32(2.0) + return (mid / f32(max_pos)).reshape(-1, 1).astype(f32) + + +def rope_tables(cfg: DiTConfig, shape: DiTShape) -> dict[str, tuple[np.ndarray, np.ndarray]]: + vg = video_grid(cfg, shape) + cross_max = max(cfg.pos_embed_max_pos, cfg.audio_pos_embed_max_pos) + # cross_attn_rope sees video_coords[:, 0:1] (time only) with max position max(video, audio). + v_time = (video_midpoints(cfg, shape)[0:1] / np.float32(cross_max)).T.astype(np.float32) + th, dp = cfg.rope_theta, cfg.rope_double_precision + return { + "video": split_rope_freqs(vg, cfg.dim, cfg.heads, th, dp), + "audio": split_rope_freqs(audio_grid(cfg, shape, cfg.audio_pos_embed_max_pos), cfg.audio_dim, + cfg.audio_heads, th, dp), + "ca_video": split_rope_freqs(v_time, cfg.audio_cross_attention_dim, cfg.heads, th, dp), + "ca_audio": split_rope_freqs(audio_grid(cfg, shape, cross_max), cfg.audio_cross_attention_dim, + cfg.audio_heads, th, dp), + } + + +# ---------------------------------------------------------------------- AdaLN helpers + + +def _adaln(g: Graph, ckpt: Checkpoint, prefix: str, t_sin_bf16): + """``LTX2AdaLayerNormSingle``: returns (mod ``[B, n*D]``, embedded_timestep ``[B, D]``) in bf16.""" + p = f"{prefix}.emb.timestep_embedder" + e = g.linear(t_sin_bf16, ckpt.get(f"{p}.linear_1.weight"), ckpt.get(f"{p}.linear_1.bias")) + e = g.linear(g.silu(e), ckpt.get(f"{p}.linear_2.weight"), ckpt.get(f"{p}.linear_2.bias")) + mod = g.linear(g.silu(e), ckpt.get(f"{prefix}.linear.weight"), ckpt.get(f"{prefix}.linear.bias")) + return mod, e + + +def _mod_params(g: Graph, table: np.ndarray, temb, batch: int): + """``table[None, None] + temb.reshape(B, 1, n, D)`` unbound into n tensors ``[B, 1, D]`` (bf16).""" + n, d = (int(s) for s in table.shape) + vals = g.add(g.reshape(temb, (batch, 1, n, d)), g.const(table.reshape(1, 1, n, d), trt.bfloat16)) + return [g.reshape(g.slice(vals, (0, 0, i, 0), (batch, 1, 1, d)), (batch, 1, d)) for i in range(n)] + + +# ---------------------------------------------------------------------- context-parallel attention + + +def _ag_tokens(g: Graph, x, cp: int): + """ALL_GATHER of ``[B, S/cp, D]`` token shards into ``[B, S, D]`` (bf16 carried as fp16).""" + b, s_loc, d = (int(v) for v in x.shape) + t = g.transpose(x, (1, 0, 2)) # [S/cp, B, D]: ALL_GATHER concatenates dim 0 + if t.dtype == trt.bfloat16: + t = g.maximum(g.minimum(t, g.scalar(_FP16_SAFE_BF16_MAX, trt.bfloat16, 3)), + g.scalar(-_FP16_SAFE_BF16_MAX, trt.bfloat16, 3)) + out = g.cast(add_collective(g.net, g.cast(t, trt.float16), trt.CollectiveOperation.ALL_GATHER, cp), + trt.bfloat16) + else: + out = add_collective(g.net, t, trt.CollectiveOperation.ALL_GATHER, cp) + return g.transpose(out, (1, 0, 2)) + + +def _video_self_attention(g: Graph, aw: AttnWeights, x_n, *, heads: int, rope: RopeTables, cp: int, + gated: bool, stg_keep): + if cp == 1: + return ltx_attention(g, aw, x_n, x_n, heads=heads, eps=EPS, q_rope=rope, k_rope=rope, gated=gated, + stg_keep=stg_keep) + # Local queries attend the gathered keys/values; Q/K RMSNorm and RoPE are row-local. + q, k, v = project_qkv(g, aw, x_n, x_n, eps=EPS, q_rope=rope, k_rope=rope) + kf, vf = _ag_tokens(g, k, cp), _ag_tokens(g, v, cp) + ctx = from_heads(g, g.attention(to_heads(g, q, heads), to_heads(g, kf, heads), to_heads(g, vf, heads))) + if stg_keep is not None: + ctx = stg_lerp(g, v, ctx, stg_keep) + return gate_and_project_out(g, aw, ctx, x_n, heads, gated=gated) + + +def _v2a_attention(g: Graph, aw: AttnWeights, a_q, v_kv, *, heads: int, q_rope: RopeTables, + k_rope: RopeTables, cp: int, gated: bool): + """video -> audio cross attention (Q audio, K/V video).""" + if cp == 1: + return ltx_attention(g, aw, a_q, v_kv, heads=heads, eps=EPS, q_rope=q_rope, k_rope=k_rope, gated=gated) + q, k, v = project_qkv(g, aw, a_q, v_kv, eps=EPS, q_rope=q_rope, k_rope=k_rope) + q4 = g.cast(to_heads(g, q, heads), trt.float32) # [B, H, Sa, d] + k4 = g.cast(to_heads(g, k, heads), trt.float32) # [B, H, S/cp, d] + v4 = g.cast(to_heads(g, v, heads), trt.float32) + b, h, sa, d = (int(s) for s in q4.shape) + q4 = g.mul(q4, g.scalar(1.0 / np.sqrt(d), trt.float32, 4)) + scores = g.net.add_matrix_multiply(q4, trt.MatrixOperation.NONE, k4, + trt.MatrixOperation.TRANSPOSE).get_output(0) # [B,H,Sa,S/cp] + m = g.reduce(scores, trt.ReduceOperation.MAX, 3) # [B,H,Sa,1] + p = g.unary(g.sub(scores, m), trt.UnaryOperation.EXP) + l_sum = g.reduce(p, trt.ReduceOperation.SUM, 3) + num = g.net.add_matrix_multiply(p, trt.MatrixOperation.NONE, v4, trt.MatrixOperation.NONE).get_output(0) + stats = g.reshape(g.concat([m, l_sum, num], axis=3), (1, b, h, sa, d + 2)) + allst = add_collective(g.net, stats, trt.CollectiveOperation.ALL_GATHER, cp) # [cp, B, H, Sa, d+2] + # Merge over a trailing rank axis: slicing / reducing the gathered leading axis directly is + # mis-compiled by TensorRT-RTX 1.7.1 (wrong values), the transposed form is exact. + t = g.transpose(allst, (1, 2, 3, 4, 0)) # [B, H, Sa, d+2, cp] + m_all = g.slice(t, (0, 0, 0, 0, 0), (b, h, sa, 1, cp)) + l_all = g.slice(t, (0, 0, 0, 1, 0), (b, h, sa, 1, cp)) + n_all = g.slice(t, (0, 0, 0, 2, 0), (b, h, sa, d, cp)) + m_max = g.reduce(m_all, trt.ReduceOperation.MAX, 4) # [B, H, Sa, 1, 1] + w = g.unary(g.sub(m_all, m_max), trt.UnaryOperation.EXP) + denom = g.reduce(g.mul(l_all, w), trt.ReduceOperation.SUM, 4) + numer = g.reduce(g.mul(n_all, w), trt.ReduceOperation.SUM, 4) + ctx4 = g.reshape(g.div(numer, denom), (b, h, sa, d)) + ctx = from_heads(g, g.cast(ctx4, trt.bfloat16)) + return gate_and_project_out(g, aw, ctx, a_q, heads, gated=gated) + + +# ---------------------------------------------------------------------- network + + +def _scalar_like(g: Graph, x, value: float): + return g.scalar(value, x.dtype, len(x.shape)) + + +def add_dit(g: Graph, ckpt: Checkpoint, cfg: DiTConfig, shape: DiTShape, inputs: dict, *, cp: int = 1, + stg_blocks: tuple[int, ...] = (), num_layers: int | None = None): + """Adds the DiT; returns (video_velocity ``[B, S_local, C]`` bf16, audio_velocity ``[B, Sa, C]`` bf16).""" + B = shape.batch + S = shape.video_tokens + if S % cp: + raise ValueError(f"video tokens {S} are not divisible by context_parallel_size {cp}") + for name, h in (("video", cfg.heads), ("audio", cfg.audio_heads)): + if h % cp: + raise ValueError(f"{name} heads {h} are not divisible by context_parallel_size {cp}") + s_loc = S // cp + D, Da = cfg.dim, cfg.audio_dim + n_layers = cfg.layers if num_layers is None else num_layers + + tables = rope_tables(cfg, shape) + rope_v = rope_constants(g, *tables["video"]) + rope_ca_v = rope_constants(g, *tables["ca_video"]) + rope_a = rope_constants(g, *tables["audio"]) + rope_ca_a = rope_constants(g, *tables["ca_audio"]) + + video_latent = g.cast(inputs["video_latent"], trt.bfloat16) + if cp > 1: + rows = local_row_indices(g, cp=cp, local_rows=s_loc) + video_latent = g.gather(video_latent, rows, 1) + rope_v = gather_rope_rows(g, rope_v, rows) + rope_ca_v = gather_rope_rows(g, rope_ca_v, rows) + audio_latent = g.cast(inputs["audio_latent"], trt.bfloat16) + + x = g.linear(video_latent, ckpt.get("proj_in.weight"), ckpt.get("proj_in.bias")) # [B, S_loc, D] + a = g.linear(audio_latent, ckpt.get("audio_proj_in.weight"), ckpt.get("audio_proj_in.bias")) + + t = inputs["timestep"] + t_sin = g.cast(g.timestep_sinusoid(t), trt.bfloat16) + gate_factor = cfg.cross_attn_timestep_scale_multiplier / cfg.timestep_scale_multiplier + t_gate_sin = t_sin if gate_factor == 1.0 else g.cast( + g.timestep_sinusoid(g.mul(t, g.scalar(gate_factor, trt.float32, 1))), trt.bfloat16) + temb_v, emb_v = _adaln(g, ckpt, "time_embed", t_sin) + temb_a, emb_a = _adaln(g, ckpt, "audio_time_embed", t_sin) + temb_pv = temb_pa = None + if (cfg.cross_attn_mod or cfg.audio_cross_attn_mod) and cfg.prompt_adaln: + temb_pv, _ = _adaln(g, ckpt, "prompt_adaln", t_sin) + temb_pa, _ = _adaln(g, ckpt, "audio_prompt_adaln", t_sin) + # use_cross_timestep=True with one shared sigma: every cross-modal modulation sees t. + ca_v, _ = _adaln(g, ckpt, "av_cross_attn_video_scale_shift", t_sin) + gate_v, _ = _adaln(g, ckpt, "av_cross_attn_video_a2v_gate", t_gate_sin) + ca_a, _ = _adaln(g, ckpt, "av_cross_attn_audio_scale_shift", t_sin) + gate_a, _ = _adaln(g, ckpt, "av_cross_attn_audio_v2a_gate", t_gate_sin) + + ctx_v_in = inputs["video_context"] + ctx_a_in = inputs["audio_context"] + stg_keep = inputs.get("stg_keep") + av_keep = g.reshape(g.cast(inputs["av_keep"], trt.bfloat16), (B, 1, 1)) + + for i in range(n_layers): + p = f"transformer_blocks.{i}" + stg = stg_keep if i in stg_blocks else None + vm = _mod_params(g, ckpt.get(f"{p}.scale_shift_table"), temb_v, B) + am = _mod_params(g, ckpt.get(f"{p}.audio_scale_shift_table"), temb_a, B) + # 1. self-attention + xn = modulate(g, g.rms_norm(x, None, EPS), vm[1], vm[0]) + x = g.add(x, g.mul(_video_self_attention(g, AttnWeights(ckpt, f"{p}.attn1"), xn, heads=cfg.heads, + rope=rope_v, cp=cp, gated=cfg.gated, stg_keep=stg), vm[2])) + an = modulate(g, g.rms_norm(a, None, EPS), am[1], am[0]) + a = g.add(a, g.mul(ltx_attention(g, AttnWeights(ckpt, f"{p}.audio_attn1"), an, an, heads=cfg.audio_heads, + eps=EPS, q_rope=rope_a, k_rope=rope_a, gated=cfg.audio_gated, + stg_keep=stg), am[2])) + # 2. text cross-attention + ctx_v, ctx_a = ctx_v_in, ctx_a_in + if cfg.cross_attn_mod or cfg.audio_cross_attn_mod: + if temb_pv is not None: + pv = _mod_params(g, ckpt.get(f"{p}.prompt_scale_shift_table"), temb_pv, B) + pa = _mod_params(g, ckpt.get(f"{p}.audio_prompt_scale_shift_table"), temb_pa, B) + else: + tv = ckpt.get(f"{p}.prompt_scale_shift_table") + ta = ckpt.get(f"{p}.audio_prompt_scale_shift_table") + pv = [g.const(tv[j].reshape(1, 1, -1), trt.bfloat16) for j in range(2)] + pa = [g.const(ta[j].reshape(1, 1, -1), trt.bfloat16) for j in range(2)] + ctx_v = modulate(g, ctx_v_in, pv[1], pv[0]) + ctx_a = modulate(g, ctx_a_in, pa[1], pa[0]) + xn = g.rms_norm(x, None, EPS) + if cfg.cross_attn_mod: + xn = modulate(g, xn, vm[7], vm[6]) + out = ltx_attention(g, AttnWeights(ckpt, f"{p}.attn2"), xn, ctx_v, heads=cfg.heads, eps=EPS, gated=cfg.gated) + if cfg.cross_attn_mod: + out = g.mul(out, vm[8]) + x = g.add(x, out) + an = g.rms_norm(a, None, EPS) + if cfg.audio_cross_attn_mod: + an = modulate(g, an, am[7], am[6]) + out = ltx_attention(g, AttnWeights(ckpt, f"{p}.audio_attn2"), an, ctx_a, heads=cfg.audio_heads, eps=EPS, + gated=cfg.audio_gated) + if cfg.audio_cross_attn_mod: + out = g.mul(out, am[8]) + a = g.add(a, out) + # 3. audio <-> video cross-attention + xn = g.rms_norm(x, None, EPS) + an = g.rms_norm(a, None, EPS) + vt = ckpt.get(f"{p}.video_a2v_cross_attn_scale_shift_table") + at = ckpt.get(f"{p}.audio_a2v_cross_attn_scale_shift_table") + v_a2v_scale, v_a2v_shift, v_v2a_scale, v_v2a_shift = _mod_params(g, vt[:4], ca_v, B) + a_a2v_scale, a_a2v_shift, a_v2a_scale, a_v2a_shift = _mod_params(g, at[:4], ca_a, B) + a2v_gate = _mod_params(g, vt[4:], gate_v, B)[0] + v2a_gate = _mod_params(g, at[4:], gate_a, B)[0] + a2v = ltx_attention(g, AttnWeights(ckpt, f"{p}.audio_to_video_attn"), modulate(g, xn, v_a2v_scale, v_a2v_shift), + modulate(g, an, a_a2v_scale, a_a2v_shift), heads=cfg.audio_heads, eps=EPS, + q_rope=rope_ca_v, k_rope=rope_ca_a, gated=cfg.gated) + v2a = _v2a_attention(g, AttnWeights(ckpt, f"{p}.video_to_audio_attn"), + modulate(g, an, a_v2a_scale, a_v2a_shift), modulate(g, xn, v_v2a_scale, v_v2a_shift), + heads=cfg.audio_heads, q_rope=rope_ca_a, k_rope=rope_ca_v, cp=cp, gated=cfg.audio_gated) + x = g.add(x, g.mul(g.mul(a2v_gate, a2v), av_keep)) + a = g.add(a, g.mul(g.mul(v2a_gate, v2a), av_keep)) + # 4. feed-forward + xn = modulate(g, g.rms_norm(x, None, EPS), vm[4], vm[3]) + x = g.add(x, g.mul(feed_forward(g, ckpt, f"{p}.ff", xn), vm[5])) + an = modulate(g, g.rms_norm(a, None, EPS), am[4], am[3]) + a = g.add(a, g.mul(feed_forward(g, ckpt, f"{p}.audio_ff", an), am[5])) + + def head(h, emb, table_name, proj, dim): + sst = ckpt.get(table_name) # [2, dim] + vals = g.add(g.reshape(emb, (B, 1, 1, dim)), g.const(sst.reshape(1, 1, 2, dim), trt.bfloat16)) + shift = g.reshape(g.slice(vals, (0, 0, 0, 0), (B, 1, 1, dim)), (B, 1, dim)) + scale = g.reshape(g.slice(vals, (0, 0, 1, 0), (B, 1, 1, dim)), (B, 1, dim)) + y = modulate(g, g.layer_norm(h, EPS), scale, shift) + return g.linear(y, ckpt.get(f"{proj}.weight"), ckpt.get(f"{proj}.bias")) + + return head(x, emb_v, "scale_shift_table", "proj_out", D), head(a, emb_a, "audio_scale_shift_table", + "audio_proj_out", Da) + + +def _gather_video_rows(g: Graph, y, cp: int, batch: int): + """fp32 ALL_GATHER of ``[B, S/cp, C]`` token shards into ``[B, S, C]`` (rank-ordered rows).""" + yf = g.cast(y, trt.float32) + b, s_loc, c = (int(v) for v in yf.shape) + if b == 1: + out = add_collective(g.net, g.reshape(yf, (s_loc, c)), trt.CollectiveOperation.ALL_GATHER, cp) + return g.reshape(out, (1, s_loc * cp, c)) + t = g.transpose(yf, (1, 0, 2)) # [S/cp, B, C] + out = add_collective(g.net, t, trt.CollectiveOperation.ALL_GATHER, cp) # [S, B, C] + return g.transpose(out, (1, 0, 2)) + + +def build_dit_engine(transformer_dir: str | Path, shape: DiTShape, *, cp_size: int = 1, + stg_blocks: tuple[int, ...] = (28,), num_layers: int | None = None, + verbose: bool = False) -> bytes: + ckpt = Checkpoint(transformer_dir) + cfg = DiTConfig.from_dict(ckpt.config()) + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + B, S, Sa, L = shape.batch, shape.video_tokens, shape.audio_frames, shape.text_len + inputs = { + "video_latent": network.add_input("video_latent", trt.float32, (B, S, cfg.in_channels)), + "audio_latent": network.add_input("audio_latent", trt.float32, (B, Sa, cfg.audio_in_channels)), + "video_context": network.add_input("video_context", trt.bfloat16, (B, L, cfg.cross_attention_dim)), + "audio_context": network.add_input("audio_context", trt.bfloat16, (B, L, cfg.audio_cross_attention_dim)), + "timestep": network.add_input("timestep", trt.float32, (B,)), + "stg_keep": network.add_input("stg_keep", trt.float32, (B,)), + "av_keep": network.add_input("av_keep", trt.float32, (B,)), + } + if cfg.cross_attention_dim != cfg.dim or cfg.audio_cross_attention_dim != cfg.audio_dim: + raise NotImplementedError("LTX-2.5 DiT builder expects the connector widths to match the streams") + video, audio = add_dit(g, ckpt, cfg, shape, inputs, cp=cp_size, stg_blocks=stg_blocks, num_layers=num_layers) + if cp_size > 1: + video = _gather_video_rows(g, video, cp_size, B) + g.mark_output(video, "video_velocity", trt.float32) + g.mark_output(audio, "audio_velocity", trt.float32) + print(f"[ltx2] Building DiT engine (batch={B}, video_tokens={S}, audio_tokens={Sa}, cp={cp_size}, " + f"layers={num_layers or cfg.layers}) ...", file=sys.stderr) + return build_plan(builder, network, label="DiT") + + +def load_dit_config(transformer_dir: str | Path) -> DiTConfig: + return DiTConfig.from_dict(json.loads((Path(transformer_dir) / "config.json").read_text(encoding="utf-8"))) + diff --git a/families/ltx2/graph.py b/families/ltx2/graph.py new file mode 100644 index 0000000000..465c472753 --- /dev/null +++ b/families/ltx2/graph.py @@ -0,0 +1,271 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned TensorRT network helpers for LTX-2.5 engine builds. + +Every network is strongly typed. Activations are bf16 (the precision LTX-2.5 and +its Gemma 4 text encoder are trained and run in); normalization statistics, +RoPE, timestep sinusoids and softmax-free reductions run in fp32 islands. + +Linear weights stay in the checkpoint's ``[out, in]`` layout and are consumed by +a transposed matrix multiply, so multi-GB checkpoints are neither transposed nor +copied on the host. bf16 constants are passed to TensorRT without a copy and kept +alive by the :class:`Graph` that created them. +""" + +from __future__ import annotations + +import math +from typing import Sequence + +import ml_dtypes +import numpy as np +import tensorrt as trt + +BF16 = ml_dtypes.bfloat16 + + +def np_dtype_for(dtype: "trt.DataType"): + if dtype == trt.bfloat16: + return BF16 + if dtype == trt.float16: + return np.float16 + if dtype == trt.float32: + return np.float32 + if dtype == trt.int32: + return np.int32 + raise ValueError(f"unsupported constant dtype {dtype}") + + +class Graph: + """Thin stateful wrapper around one ``INetworkDefinition``.""" + + def __init__(self, network: "trt.INetworkDefinition"): + self.net = network + self._keepalive: list[np.ndarray] = [] + + # ------------------------------------------------------------------ constants + + def const(self, values, dtype: "trt.DataType" = trt.float32, shape: Sequence[int] | None = None): + """Constant tensor of ``dtype``; ``values`` is converted (round-to-nearest-even) if needed.""" + arr = np.asarray(values) + target = np_dtype_for(dtype) + if arr.dtype != target: + arr = arr.astype(target) + arr = np.ascontiguousarray(arr) + if shape is None: + shape = arr.shape if arr.ndim else (1,) + # The explicit (type, pointer, count) form: implicit NumPy dtype detection differs per + # platform (e.g. int32 on Windows) and does not know bf16. TensorRT does not copy. + weights = trt.Weights(dtype, arr.ctypes.data, arr.size) + self._keepalive.append(arr) + return self.net.add_constant(tuple(int(s) for s in shape), weights).get_output(0) + + def weights(self, values, dtype: "trt.DataType") -> "trt.Weights": + """``trt.Weights`` in ``dtype`` for layer parameters (kept alive with the graph).""" + arr = np.ascontiguousarray(np.asarray(values).astype(np_dtype_for(dtype))) + self._keepalive.append(arr) + return trt.Weights(dtype, arr.ctypes.data, arr.size) + + def scalar(self, value: float, dtype: "trt.DataType", rank: int): + return self.const(np.full((1,) * rank, value, dtype=np.float32), dtype) + + # ------------------------------------------------------------------ basic ops + + def cast(self, x, dtype: "trt.DataType"): + if x.dtype == dtype: + return x + return self.net.add_cast(x, dtype).get_output(0) + + def ew(self, a, b, op): + return self.net.add_elementwise(a, b, op).get_output(0) + + def add(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.SUM) + + def sub(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.SUB) + + def mul(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.PROD) + + def div(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.DIV) + + def maximum(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.MAX) + + def minimum(self, a, b): + return self.ew(a, b, trt.ElementWiseOperation.MIN) + + def unary(self, x, op): + return self.net.add_unary(x, op).get_output(0) + + def reduce(self, x, op, axis: int, keep_dims: bool = True): + rank = len(x.shape) + axis = axis % rank + return self.net.add_reduce(x, op, 1 << axis, keep_dims).get_output(0) + + def reshape(self, x, shape: Sequence[int], *, first: Sequence[int] | None = None, + second: Sequence[int] | None = None): + layer = self.net.add_shuffle(x) + if first is not None: + layer.first_transpose = trt.Permutation(list(first)) + layer.reshape_dims = tuple(int(s) for s in shape) + if second is not None: + layer.second_transpose = trt.Permutation(list(second)) + return layer.get_output(0) + + def transpose(self, x, perm: Sequence[int]): + layer = self.net.add_shuffle(x) + layer.first_transpose = trt.Permutation(list(perm)) + return layer.get_output(0) + + def slice(self, x, start: Sequence[int], size: Sequence[int], stride: Sequence[int] | None = None): + if stride is None: + stride = (1,) * len(start) + return self.net.add_slice(x, tuple(start), tuple(size), tuple(stride)).get_output(0) + + def concat(self, xs, axis: int): + layer = self.net.add_concatenation(list(xs)) + layer.axis = axis + return layer.get_output(0) + + def gather(self, x, indices, axis: int): + return self.net.add_gather(x, indices, axis).get_output(0) + + def select(self, cond, a, b): + return self.net.add_select(cond, a, b).get_output(0) + + def mark_output(self, x, name: str, dtype: "trt.DataType | None" = None): + if dtype is not None: + x = self.cast(x, dtype) + x.name = name + self.net.mark_output(x) + return x + + # ------------------------------------------------------------------ layers + + def linear(self, x, weight: np.ndarray, bias: np.ndarray | None = None): + """``x @ weight.T + bias`` with ``weight`` in checkpoint ``[out, in]`` layout.""" + rank = len(x.shape) + out_f, in_f = (int(s) for s in weight.shape) + w = self.const(weight, x.dtype, shape=(1,) * (rank - 2) + (out_f, in_f)) + y = self.net.add_matrix_multiply( + x, trt.MatrixOperation.NONE, w, trt.MatrixOperation.TRANSPOSE + ).get_output(0) + if bias is not None: + b = self.const(bias, x.dtype, shape=(1,) * (rank - 1) + (out_f,)) + y = self.add(y, b) + return y + + def rms_norm(self, x, weight: np.ndarray | None, eps: float, *, out_dtype=None): + """RMSNorm over the last axis with fp32 statistics (``x * rsqrt(mean(x^2) + eps) * w``).""" + out_dtype = out_dtype or x.dtype + rank = len(x.shape) + xf = self.cast(x, trt.float32) + ms = self.reduce(self.mul(xf, xf), trt.ReduceOperation.AVG, -1) + inv = self.unary(self.unary(self.add(ms, self.scalar(eps, trt.float32, rank)), + trt.UnaryOperation.SQRT), trt.UnaryOperation.RECIP) + y = self.mul(xf, inv) + if weight is not None: + y = self.mul(y, self.const(np.asarray(weight, dtype=np.float32), trt.float32, + shape=(1,) * (rank - 1) + (int(weight.shape[-1]),))) + return self.cast(y, out_dtype) + + def layer_norm(self, x, eps: float, *, out_dtype=None): + """Non-affine LayerNorm over the last axis in fp32.""" + out_dtype = out_dtype or x.dtype + rank = len(x.shape) + xf = self.cast(x, trt.float32) + mean = self.reduce(xf, trt.ReduceOperation.AVG, -1) + centered = self.sub(xf, mean) + var = self.reduce(self.mul(centered, centered), trt.ReduceOperation.AVG, -1) + inv = self.unary(self.unary(self.add(var, self.scalar(eps, trt.float32, rank)), + trt.UnaryOperation.SQRT), trt.UnaryOperation.RECIP) + return self.cast(self.mul(centered, inv), out_dtype) + + def gelu_tanh(self, x): + """GELU, tanh approximation (``gelu_pytorch_tanh`` / ``gelu-approximate``). + + Evaluated in fp32 and rounded once to x.dtype, like torch's bf16 kernel. + """ + rank = len(x.shape) + out_dtype = x.dtype + x = self.cast(x, trt.float32) + c = lambda v: self.scalar(v, trt.float32, rank) # noqa: E731 + inner = self.mul(c(math.sqrt(2.0 / math.pi)), + self.add(x, self.mul(c(0.044715), self.mul(self.mul(x, x), x)))) + t = self.net.add_activation(inner, trt.ActivationType.TANH).get_output(0) + return self.cast(self.mul(self.mul(c(0.5), x), self.add(c(1.0), t)), out_dtype) + + def silu(self, x): + """SiLU evaluated in fp32 and rounded once to x.dtype.""" + out_dtype = x.dtype + x = self.cast(x, trt.float32) + s = self.net.add_activation(x, trt.ActivationType.SIGMOID).get_output(0) + return self.cast(self.mul(x, s), out_dtype) + + def sigmoid(self, x): + return self.net.add_activation(x, trt.ActivationType.SIGMOID).get_output(0) + + def attention(self, q, k, v, *, scale: float | None = None, mask=None): + """Softmax attention over ``[B, H, S, D]`` tensors (TensorRT IAttention). + + IAttention computes raw ``Q @ K^T``; Q is pre-scaled by ``scale`` + (default ``1/sqrt(D)``; ``scale=1.0`` skips the multiply). + """ + if scale is None: + scale = 1.0 / math.sqrt(int(q.shape[-1])) + if scale != 1.0: + q = self.mul(q, self.scalar(scale, q.dtype, 4)) + layer = self.net.add_attention(q, k, v, trt.AttentionNormalizationOp.SOFTMAX, False) + if layer is None: + raise RuntimeError("TensorRT rejected the attention layer") + layer.decomposable = True + if mask is not None: + layer.mask = mask + return layer.get_output(0) + + def timestep_sinusoid(self, t, dim: int = 256, max_period: float = 10000.0): + """diffusers ``Timesteps(dim, flip_sin_to_cos=True, downscale_freq_shift=0)``: [B] -> [B, dim] fp32.""" + half = dim // 2 + freqs = np.exp(-math.log(max_period) * np.arange(half, dtype=np.float32) / half).astype(np.float32) + b = int(t.shape[0]) + t2 = self.reshape(self.cast(t, trt.float32), (b, 1)) + args = self.mul(t2, self.const(freqs.reshape(1, half), trt.float32)) + return self.concat([self.unary(args, trt.UnaryOperation.COS), + self.unary(args, trt.UnaryOperation.SIN)], axis=1) + + +def new_network(logger: "trt.ILogger"): + builder = trt.Builder(logger) + flags = 1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED) + network = builder.create_network(flags) + return builder, network + + +def build_plan(builder, network, *, label: str = "engine", tf32: bool = True): + """Serialized plan (a bytes-like ``IHostMemory``; multi-GB plans are not copied again). + + ``tf32=False`` keeps fp32 convolutions / matrix multiplies in full fp32 (TensorRT allows + TF32 for fp32 layers by default). + """ + config = builder.create_builder_config() + config.builder_optimization_level = 3 + if not tf32: + config.clear_flag(trt.BuilderFlag.TF32) + plan = builder.build_serialized_network(network, config) + if plan is None: + raise RuntimeError(f"TensorRT failed to build the LTX-2.5 {label}") + return plan + + +_LOGGERS: dict[bool, "trt.ILogger"] = {} + + +def make_logger(verbose: bool = False): + """One TensorRT logger per process and verbosity.""" + if verbose not in _LOGGERS: + _LOGGERS[verbose] = trt.Logger(trt.Logger.VERBOSE if verbose else trt.Logger.WARNING) + return _LOGGERS[verbose] diff --git a/families/ltx2/layers.py b/families/ltx2/layers.py new file mode 100644 index 0000000000..8a1e6caf38 --- /dev/null +++ b/families/ltx2/layers.py @@ -0,0 +1,188 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2 attention, RoPE and feed-forward lowering shared by the connectors and the DiT. + +Mirrors diffusers ``LTX2Attention`` + ``LTX2AudioVideoAttnProcessor`` / +``LTX2PerturbedAttnProcessor`` and ``apply_split_rotary_emb``: + +- Q/K RMSNorm spans all heads of a token (``rms_norm_across_heads``). +- Split RoPE rotates the two halves of every head: + ``out1 = x1*cos - x2*sin``, ``out2 = x2*cos + x1*sin`` in fp32, applied + elementwise on a ``[B, T, H, 2, r]`` view (no dense rotate-half matrix). +- Gated attention multiplies each head's context by ``2*sigmoid(to_gate_logits(x_q))``. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import tensorrt as trt + +from .graph import Graph + + +@dataclass +class RopeTables: + """Split-RoPE cos/sin tensors of shape ``[1, T, H, 1, r]`` (fp32, already in the network).""" + + cos: object + sin: object + heads: int + half: int + + +def split_rope_freqs(grid: np.ndarray, dim: int, num_heads: int, theta: float, + double_precision: bool = True) -> tuple[np.ndarray, np.ndarray]: + """diffusers ``LTX2AudioVideoRotaryPosEmbed.forward`` / ``LTX2RotaryPosEmbed1d`` (split type). + + ``grid``: ``[T, P]`` fractional positions (already divided by the max positions). + Returns cos, sin of shape ``[T, num_heads, dim // num_heads // 2]`` in fp32. + """ + grid = np.asarray(grid, dtype=np.float32) + t, p = grid.shape + steps = dim // (2 * p) + fdt = np.float64 if double_precision else np.float32 + pow_indices = np.power(fdt(theta), np.linspace(0.0, 1.0, steps, dtype=fdt)) + freqs = (pow_indices * np.pi / 2.0).astype(np.float32) # [F] + angles = (grid[:, :, None] * np.float32(2.0) - np.float32(1.0)) * freqs[None, None, :] # [T, P, F] + angles = np.transpose(angles, (0, 2, 1)).reshape(t, steps * p).astype(np.float32) # freq-major + cos = np.cos(angles).astype(np.float32) + sin = np.sin(angles).astype(np.float32) + pad = dim // 2 - angles.shape[1] + if pad: + cos = np.concatenate([np.ones((t, pad), np.float32), cos], axis=1) + sin = np.concatenate([np.zeros((t, pad), np.float32), sin], axis=1) + half = dim // num_heads // 2 + return cos.reshape(t, num_heads, half), sin.reshape(t, num_heads, half) + + +def rope_constants(g: Graph, cos: np.ndarray, sin: np.ndarray) -> RopeTables: + t, h, r = cos.shape + return RopeTables(g.const(cos.reshape(1, t, h, 1, r), trt.float32), + g.const(sin.reshape(1, t, h, 1, r), trt.float32), h, r) + + +def gather_rope_rows(g: Graph, rope: RopeTables, rows) -> RopeTables: + """Rows ``rows`` (int32 [T_local]) of a RoPE table, e.g. one context-parallel shard.""" + return RopeTables(g.gather(rope.cos, rows, 1), g.gather(rope.sin, rows, 1), rope.heads, rope.half) + + +def apply_split_rope(g: Graph, x, rope: RopeTables): + """``x``: ``[B, T, H*2r]`` -> same shape/dtype, rotated in fp32.""" + b, t, d = (int(s) for s in x.shape) + h, r = rope.heads, rope.half + out_dtype = x.dtype + x5 = g.reshape(g.cast(x, trt.float32), (b, t, h, 2, r)) + x1 = g.slice(x5, (0, 0, 0, 0, 0), (b, t, h, 1, r)) + x2 = g.slice(x5, (0, 0, 0, 1, 0), (b, t, h, 1, r)) + o1 = g.sub(g.mul(x1, rope.cos), g.mul(x2, rope.sin)) + o2 = g.add(g.mul(x2, rope.cos), g.mul(x1, rope.sin)) + return g.cast(g.reshape(g.concat([o1, o2], axis=3), (b, t, d)), out_dtype) + + +def to_heads(g: Graph, x, heads: int): + """``[B, T, H*D]`` -> ``[B, H, T, D]``.""" + b, t, d = (int(s) for s in x.shape) + return g.reshape(x, (b, t, heads, d // heads), second=(0, 2, 1, 3)) + + +def from_heads(g: Graph, x): + """``[B, H, T, D]`` -> ``[B, T, H*D]``.""" + b, h, t, d = (int(s) for s in x.shape) + return g.reshape(x, (b, t, h * d), first=(0, 2, 1, 3)) + + +class AttnWeights: + """Weights of one ``LTX2Attention`` (checkpoint layout) fetched lazily from a checkpoint.""" + + def __init__(self, ckpt, prefix: str): + self.prefix = prefix + self.ckpt = ckpt + + def w(self, name: str): + return self.ckpt.get(f"{self.prefix}.{name}") + + def maybe(self, name: str): + return self.ckpt.maybe(f"{self.prefix}.{name}") + + def f32(self, name: str): + return self.ckpt.get(f"{self.prefix}.{name}", np.float32) + + +def project_qkv(g: Graph, aw: AttnWeights, x_q, x_kv, *, eps: float, + q_rope: RopeTables | None, k_rope: RopeTables | None, need_q: bool = True): + """to_q/to_k/to_v + across-heads RMSNorm + split RoPE; returns (q, k, v) as ``[B, T, inner]``.""" + q = None + if need_q: + q = g.linear(x_q, aw.w("to_q.weight"), aw.maybe("to_q.bias")) + q = g.rms_norm(q, aw.f32("norm_q.weight"), eps) + if q_rope is not None: + q = apply_split_rope(g, q, q_rope) + k = g.linear(x_kv, aw.w("to_k.weight"), aw.maybe("to_k.bias")) + k = g.rms_norm(k, aw.f32("norm_k.weight"), eps) + if k_rope is not None: + k = apply_split_rope(g, k, k_rope) + v = g.linear(x_kv, aw.w("to_v.weight"), aw.maybe("to_v.bias")) + return q, k, v + + +def gate_and_project_out(g: Graph, aw: AttnWeights, ctx, x_q, heads: int, *, gated: bool): + """Per-head ``2*sigmoid`` gates (from the query input) and ``to_out.0``; ``ctx`` is ``[B, T, inner]``.""" + if gated: + b, t, inner = (int(s) for s in ctx.shape) + logits = g.linear(x_q, aw.w("to_gate_logits.weight"), aw.maybe("to_gate_logits.bias")) # [B,T,H] + gates = g.mul(g.sigmoid(g.cast(logits, trt.float32)), g.scalar(2.0, trt.float32, 3)) + gates = g.cast(g.reshape(gates, (b, t, heads, 1)), ctx.dtype) + ctx = g.reshape(g.mul(g.reshape(ctx, (b, t, heads, inner // heads)), gates), (b, t, inner)) + return g.linear(ctx, aw.w("to_out.0.weight"), aw.maybe("to_out.0.bias")) + + +def ltx_attention(g: Graph, aw: AttnWeights, x_q, x_kv, *, heads: int, eps: float, + q_rope: RopeTables | None = None, k_rope: RopeTables | None = None, + gated: bool = True, mask=None, stg_keep=None): + """Full ``LTX2Attention`` forward (single device). + + ``stg_keep``: optional ``[B, 1, 1]`` mask for spatio-temporal guidance; batch rows + with 0 replace the attention context by the value projection + (``torch.lerp(value, attn, mask)`` in ``LTX2PerturbedAttnProcessor``). + """ + q, k, v = project_qkv(g, aw, x_q, x_kv, eps=eps, q_rope=q_rope, k_rope=k_rope) + ctx = from_heads(g, g.attention(to_heads(g, q, heads), to_heads(g, k, heads), to_heads(g, v, heads), + mask=mask)) + if stg_keep is not None: + ctx = stg_lerp(g, v, ctx, stg_keep) + return gate_and_project_out(g, aw, ctx, x_q, heads, gated=gated) + + +def stg_lerp(g: Graph, value, ctx, keep): + """Per batch row: ``ctx`` where ``keep`` > 0.5, else ``value`` (STG masks are exactly 0 or 1). + + A select keeps both branches bit-exact, like ``torch.lerp`` at weights 0 and 1. + """ + rank = len(ctx.shape) + keep = g.reshape(g.cast(keep, trt.float32), (int(keep.shape[0]),) + (1,) * (rank - 1)) + cond = g.ew(keep, g.scalar(0.5, trt.float32, rank), trt.ElementWiseOperation.GREATER) + return g.select(cond, ctx, value) + + +def feed_forward(g: Graph, ckpt, prefix: str, x): + """diffusers ``FeedForward(activation_fn="gelu-approximate")``: proj -> GELU(tanh) -> proj.""" + h = g.linear(x, ckpt.get(f"{prefix}.net.0.proj.weight"), ckpt.maybe(f"{prefix}.net.0.proj.bias")) + h = g.gelu_tanh(h) + return g.linear(h, ckpt.get(f"{prefix}.net.2.weight"), ckpt.maybe(f"{prefix}.net.2.bias")) + + +def modulate(g: Graph, x, scale, shift): + """``x * (1 + scale) + shift`` in x.dtype (scale/shift broadcast ``[B, 1, D]``).""" + one = g.scalar(1.0, x.dtype, len(x.shape)) + return g.add(g.mul(x, g.add(one, g.cast(scale, x.dtype))), g.cast(shift, x.dtype)) + + +def connector_rope_tables(seq_len: int, dim: int, heads: int, *, base_seq_len: int, theta: float, + double_precision: bool = True): + """``LTX2RotaryPosEmbed1d`` (split): positions ``arange(L) / base_seq_len``.""" + grid = (np.arange(seq_len, dtype=np.float32) / np.float32(base_seq_len)).reshape(seq_len, 1) + return split_rope_freqs(grid, dim, heads, theta, double_precision) + diff --git a/families/ltx2/model.py b/families/ltx2/model.py new file mode 100644 index 0000000000..ce2f57ed2a --- /dev/null +++ b/families/ltx2/model.py @@ -0,0 +1,208 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 family plugin: text-to-audio-video bundles for Lightricks ``LTX2Pipeline``. + +Builds a native TRTMC bundle from a diffusers LTX-2.5 checkpoint (distilled ``transformer/``): + + - ``text_encoder.plan``: Gemma 4 text tower + ``LTX2TextConnectors`` + - ``denoiser.plan``: the joint audio/video DiT, single device (``context_parallel_size=1``) + or context parallel over the video tokens (``context_parallel_size=2``, one rank-dynamic plan) + - ``vae.plan``: the video VAE decoder + - ``audio.plan``: the audio VAE decoder + vocoder with bandwidth extension (48 kHz stereo) + - ``tokenizer.json`` and ``runtime.json`` + +Every engine is built directly with the TensorRT network API in bf16 (the precision LTX-2.5 is +trained and released in). The runtime path is C++ + TensorRT (or TensorRT-RTX) only. +""" + +from __future__ import annotations + +import json +import sys +import time +from pathlib import Path +from typing import TYPE_CHECKING + +from .parallel import ParallelConfig, validate_context_parallel_layout + +if TYPE_CHECKING: + from tensorrt_model_connect.build import BuildRequest + from tensorrt_model_connect.bundle_writer import BundleWriter + +TASK = "text_to_audio_video" +PIPELINE_CLASS = "LTX2Pipeline" + +# diffusers ``pipelines/ltx2/utils.py`` DISTILLED_SIGMA_VALUES: the distilled checkpoint's +# 8-step schedule. LTX-2.5's shipped scheduler config disables dynamic shifting, so the +# pipeline uses these values unshifted (timesteps = sigma * 1000), then the terminal 0. +DISTILLED_SIGMAS = (1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875) + +DEFAULT_HEIGHT = 544 # the LTX-2.5 model card's 960x544, 121 frames at 24 fps +DEFAULT_WIDTH = 960 +DEFAULT_FRAMES = 121 +FRAME_RATE = 24.0 +TEXT_SEQ_LEN = 1024 +SPATIAL_COMPRESSION = 32 +TEMPORAL_COMPRESSION = 8 + +# TensorRT-RTX builds before the rel-11.4 line (1.6.x and older) ship without multi-device +# (IDistCollectiveLayer) support; 1.7.1 is the first with it and with the Myelin communicator +# handoff that context parallelism needs. +MIN_RTX_FOR_CONTEXT_PARALLEL = (1, 7, 1) + + +def _version_tuple(text: str) -> tuple[int, ...]: + parts = [] + for piece in text.split("."): + digits = "".join(ch for ch in piece if ch.isdigit()) + if not digits: + break + parts.append(int(digits)) + return tuple(parts) + + +def require_rtx_context_parallel(rtx_version: str | None, cp_size: int) -> None: + """Fail early when a TensorRT-RTX build cannot run multi-device engines.""" + if cp_size <= 1: + return + found = rtx_version or "not installed" + if rtx_version is None or _version_tuple(rtx_version)[:3] < MIN_RTX_FOR_CONTEXT_PARALLEL: + raise RuntimeError( + "LTX-2 context parallelism on TensorRT-RTX requires TensorRT-RTX >= 1.7.1 " + f"(found {found}); 1.6.x ships without multi-device support" + ) + + +def installed_rtx_version() -> str | None: + from importlib import metadata + + for name in ("tensorrt-rtx", "tensorrt_rtx"): + try: + return metadata.version(name) + except metadata.PackageNotFoundError: + continue + return None + + +def _read_json(path: Path) -> dict: + if not path.is_file(): + raise FileNotFoundError(f"LTX-2.5 checkpoint file is missing: {path}") + return json.loads(path.read_text(encoding="utf-8")) + + +def _log(message: str) -> None: + print(f"[ltx2] {message}", file=sys.stderr, flush=True) + + +def _write_plan(writer: "BundleWriter", name: str, plan) -> None: + with writer.open_section(name) as section: + section.write(memoryview(plan)) + + +def build(request: "BuildRequest", writer: "BundleWriter") -> None: + """Build one LTX-2.5 text-to-audio-video bundle.""" + if request.task != TASK: + raise ValueError(f"ltx2 supports only task={TASK}") + if request.dynamic_kv_cache: + raise NotImplementedError("ltx2 does not support dynamic_kv_cache") + if request.tensor_parallel_size != 1: + raise NotImplementedError("ltx2 requires tensor_parallel_size=1 (it shards the video tokens)") + if request.max_batch_size != 1: + raise NotImplementedError("ltx2 requires max_batch_size=1") + if request.quantization not in (None, "none"): + raise NotImplementedError("ltx2 does not support quantization") + if request.fp32_layers: + raise NotImplementedError("ltx2 does not support fp32_layers (its fp32 islands are fixed in the graph)") + if request.precision != "bf16": + raise ValueError("ltx2 builds bf16 engines (precision=bf16), the precision LTX-2.5 runs in") + parallel = ParallelConfig(cp_size=int(request.context_parallel_size)) + if parallel.cp_size not in (1, 2): + raise ValueError("ltx2 supports context_parallel_size 1 or 2") + if request.backend == "trt_rtx": + require_rtx_context_parallel(installed_rtx_version(), parallel.cp_size) + + # The builders import the bound TensorRT module; validate the request before loading them. + from .audio_builder import build_audio_decoder_engine + from .dit_builder import DiTConfig, DiTShape, audio_latent_frames, build_dit_engine + from .text_encoder_builder import build_text_encoder_engine + from .vae_builder import build_vae_decoder_engine + + model_dir = Path(request.model_dir) + index = _read_json(model_dir / "model_index.json") + if index.get("_class_name") != PIPELINE_CLASS: + raise ValueError(f"expected a diffusers {PIPELINE_CLASS} checkpoint, got {index.get('_class_name')!r}") + scheduler = _read_json(model_dir / "scheduler" / "scheduler_config.json") + if scheduler.get("use_dynamic_shifting", False) or scheduler.get("shift_terminal"): + raise NotImplementedError("ltx2 builds the distilled schedule; this scheduler config shifts sigmas") + if float(scheduler.get("shift", 1.0)) != 1.0: + raise NotImplementedError("ltx2 builds the distilled schedule with shift 1.0") + vae_cfg = _read_json(model_dir / "vae" / "config.json") + dit_cfg = DiTConfig.from_dict(_read_json(model_dir / "transformer" / "config.json")) + audio_cfg = _read_json(model_dir / "audio_vae" / "config.json") + vocoder_cfg = _read_json(model_dir / "vocoder" / "config.json") + + height = int(request.image_height or DEFAULT_HEIGHT) + width = int(request.image_width or DEFAULT_WIDTH) + frames = int(request.video_num_frames or DEFAULT_FRAMES) + spatial = int(vae_cfg.get("spatial_compression_ratio", SPATIAL_COMPRESSION)) + temporal = int(vae_cfg.get("temporal_compression_ratio", TEMPORAL_COMPRESSION)) + if height % spatial or width % spatial: + raise ValueError(f"ltx2 image_height/image_width must be multiples of {spatial}") + if (frames - 1) % temporal: + raise ValueError(f"ltx2 video_num_frames must equal {temporal}*n+1") + text_len = int(request.max_sequence_length or TEXT_SEQ_LEN) + shape = DiTShape( + batch=1, + latent_frames=(frames - 1) // temporal + 1, + latent_height=height // spatial, + latent_width=width // spatial, + audio_frames=audio_latent_frames(frames, FRAME_RATE, sampling_rate=int(audio_cfg.get("sample_rate", 16000)), + hop_length=int(audio_cfg.get("mel_hop_length", 160))), + text_len=text_len, + fps=FRAME_RATE, + ) + validate_context_parallel_layout(parallel, video_tokens=shape.video_tokens, video_heads=dit_cfg.heads, + audio_heads=dit_cfg.audio_heads) + + writer.set_header(family="ltx2", task=request.task, backend=request.backend) + started = time.perf_counter() + _write_plan(writer, "text_encoder.plan", build_text_encoder_engine(model_dir, seq_len=text_len, + verbose=request.verbose)) + _log(f"text encoder engine built in {time.perf_counter() - started:.1f} s") + started = time.perf_counter() + _write_plan(writer, "denoiser.plan", build_dit_engine(model_dir / "transformer", shape, cp_size=parallel.cp_size, + verbose=request.verbose)) + _log(f"DiT engine built in {time.perf_counter() - started:.1f} s (cp={parallel.cp_size}, " + f"{shape.video_tokens} video + {shape.audio_frames} audio tokens)") + started = time.perf_counter() + _write_plan(writer, "vae.plan", build_vae_decoder_engine(model_dir / "vae", latent_frames=shape.latent_frames, + latent_height=shape.latent_height, + latent_width=shape.latent_width, + verbose=request.verbose)) + _log(f"video VAE engine built in {time.perf_counter() - started:.1f} s") + started = time.perf_counter() + _write_plan(writer, "audio.plan", build_audio_decoder_engine(model_dir, audio_frames=shape.audio_frames, + verbose=request.verbose)) + _log(f"audio decoder engine built in {time.perf_counter() - started:.1f} s") + writer.add_bytes("tokenizer.json", (model_dir / "tokenizer" / "tokenizer.json").read_bytes()) + writer.add_json("runtime.json", { + "video_frames": frames, + "video_height": height, + "video_width": width, + "latent_frames": shape.latent_frames, + "latent_height": shape.latent_height, + "latent_width": shape.latent_width, + "latent_channels": dit_cfg.in_channels, + "audio_frames": shape.audio_frames, + "audio_latent_channels": dit_cfg.audio_in_channels, + "text_seq_len": text_len, + "frame_rate": FRAME_RATE, + "pad_token_id": 0, + "sigmas": [*DISTILLED_SIGMAS, 0.0], + "dit_batch": 1, + "audio_sample_rate": int(vocoder_cfg.get("output_sampling_rate", 48000)), + "audio_channels": int(vocoder_cfg.get("out_channels", 2)), + "parallel_mode": parallel.mode, + "parallel_size": parallel.world_size, + }) diff --git a/families/ltx2/parallel.py b/families/ltx2/parallel.py new file mode 100644 index 0000000000..5b75a072f3 --- /dev/null +++ b/families/ltx2/parallel.py @@ -0,0 +1,85 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 model-owned context-parallel build primitives.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np + +SUPPORTED_CONTEXT_PARALLEL_SIZES = (1, 2, 4, 8) + + +@dataclass(frozen=True) +class ParallelConfig: + """LTX-2.5 runs single-device or context parallel (video token shards); TP is not supported.""" + + cp_size: int = 1 + + @property + def cp_enabled(self) -> bool: + return self.cp_size > 1 + + @property + def mode(self) -> str: + return "context_parallel" if self.cp_enabled else "single" + + @property + def world_size(self) -> int: + return self.cp_size + + def validate(self) -> None: + if self.cp_size not in SUPPORTED_CONTEXT_PARALLEL_SIZES: + raise ValueError("LTX-2.5 context_parallel_size must be one of 1, 2, 4, 8") + + +def validate_context_parallel_layout(parallel: ParallelConfig, *, video_tokens: int, video_heads: int, + audio_heads: int) -> None: + """Reject layouts whose video tokens or attention heads cannot be split evenly.""" + parallel.validate() + if not parallel.cp_enabled: + return + cp = parallel.cp_size + if video_tokens % cp: + raise ValueError(f"LTX-2.5 context parallel needs the video token count ({video_tokens}) " + f"divisible by context_parallel_size ({cp})") + for name, heads in (("video", video_heads), ("audio", audio_heads)): + if heads % cp: + raise ValueError(f"LTX-2.5 context parallel needs {name} heads ({heads}) divisible by " + f"context_parallel_size ({cp})") + + +def add_collective(network, tensor, operation, cp_size: int, *, reduce_operation=None): + """One TensorRT distributed collective spanning the whole CP world.""" + import tensorrt as trt + + if reduce_operation is None: + reduce_operation = trt.ReduceOperation.NONE + layer = network.add_dist_collective(tensor, operation, reduce_operation, -1, []) + if layer is None: + raise RuntimeError(f"TensorRT failed to add the LTX-2.5 {operation} collective") + layer.num_ranks = int(cp_size) + return layer.get_output(0) + + +def rank_selector_values(cp_size: int) -> np.ndarray: + """Replicated ``[CP, 1]`` values whose SUM reduce-scatter yields each rank's index. + + Every rank contributes ``r / CP`` at row ``r``; summing CP identical copies hands + rank ``r`` exactly ``r`` (CP is a power of two, so this is exact in fp32). + """ + return (np.arange(cp_size, dtype=np.float32) / np.float32(cp_size)).reshape(cp_size, 1) + + +def local_row_indices(g, *, cp: int, local_rows: int): + """int32 ``[local_rows]`` indices of the contiguous token shard owned by this rank.""" + import tensorrt as trt + + selector = g.const(rank_selector_values(cp), trt.float32) + rank_f = add_collective(g.net, selector, trt.CollectiveOperation.REDUCE_SCATTER, cp, + reduce_operation=trt.ReduceOperation.SUM) + rank_i = g.reshape(g.cast(rank_f, trt.int32), (1,)) + start = g.mul(rank_i, g.const(np.array([local_rows], np.int32), trt.int32)) + return g.add(g.const(np.arange(local_rows, dtype=np.int32), trt.int32), start) diff --git a/families/ltx2/requirements.txt b/families/ltx2/requirements.txt new file mode 100644 index 0000000000..4f165f0f3d --- /dev/null +++ b/families/ltx2/requirements.txt @@ -0,0 +1,10 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +numpy +ml_dtypes +safetensors +torch +transformers==5.18.0 +diffusers==0.40.0 +cuda-python diff --git a/families/ltx2/runtime/CMakeLists.txt b/families/ltx2/runtime/CMakeLists.txt new file mode 100644 index 0000000000..56d53912b9 --- /dev/null +++ b/families/ltx2/runtime/CMakeLists.txt @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +add_library(trtmc_model_ltx2 SHARED + bpe_tokenizer.cpp + distributed_runtime.cpp + pipeline.cpp + plugin.cpp + runtime_config.cpp +) + +target_include_directories(trtmc_model_ltx2 + PRIVATE + ${PROJECT_SOURCE_DIR} + ${PROJECT_SOURCE_DIR}/core/runtime/include +) +target_include_directories(trtmc_model_ltx2 SYSTEM PRIVATE + ${TRTMC_CUDA_INCLUDE_DIR} +) +target_link_libraries(trtmc_model_ltx2 PRIVATE + trtmc_core + nlohmann_json::nlohmann_json + ${TRTMC_CUDART_LIBRARY} + ${CMAKE_DL_LIBS} +) +target_compile_options(trtmc_model_ltx2 PRIVATE + "$<$:-Wall;-Wextra;-Wpedantic>" +) +set_target_properties(trtmc_model_ltx2 PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" + BUILD_RPATH "\$ORIGIN" + INSTALL_RPATH "\$ORIGIN" +) +install(TARGETS trtmc_model_ltx2 + LIBRARY DESTINATION ${CMAKE_INSTALL_LIBDIR} +) + +if(TRTMC_BUILD_TESTS) + add_executable(test_ltx2_runtime_contract + ../tests/cpp/test_runtime_contract.cpp + ) + target_include_directories(test_ltx2_runtime_contract PRIVATE + ${PROJECT_SOURCE_DIR} + ) + target_compile_options(test_ltx2_runtime_contract PRIVATE + -Wall -Wextra -Wpedantic + ) + add_test(NAME ltx2_runtime_contract + COMMAND test_ltx2_runtime_contract + ) +endif() diff --git a/families/ltx2/runtime/bpe_tokenizer.cpp b/families/ltx2/runtime/bpe_tokenizer.cpp new file mode 100644 index 0000000000..1d79344d7a --- /dev/null +++ b/families/ltx2/runtime/bpe_tokenizer.cpp @@ -0,0 +1,1489 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "families/ltx2/runtime/tokenizer.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc { +namespace { + +// ─── UTF-8 helpers ─── + +// Decode one UTF-8 codepoint from s starting at pos, advance pos. +inline char32_t utf8_to_char32(const std::string& s, size_t& pos) { + unsigned char c = static_cast(s[pos]); + if (c < 0x80) { + ++pos; + return static_cast(c); + } + if ((c & 0xE0) == 0xC0 && pos + 1 < s.size()) { + char32_t cp = (static_cast(c & 0x1F) << 6) | + static_cast(static_cast(s[pos + 1]) & 0x3F); + pos += 2; + return cp; + } + if ((c & 0xF0) == 0xE0 && pos + 2 < s.size()) { + char32_t cp = (static_cast(c & 0x0F) << 12) | + (static_cast(static_cast(s[pos + 1]) & 0x3F) << 6) | + static_cast(static_cast(s[pos + 2]) & 0x3F); + pos += 3; + return cp; + } + if ((c & 0xF8) == 0xF0 && pos + 3 < s.size()) { + char32_t cp = (static_cast(c & 0x07) << 18) | + (static_cast(static_cast(s[pos + 1]) & 0x3F) << 12) | + (static_cast(static_cast(s[pos + 2]) & 0x3F) << 6) | + static_cast(static_cast(s[pos + 3]) & 0x3F); + pos += 4; + return cp; + } + ++pos; + return 0xFFFD; +} + +inline std::string utf32_to_utf8(char32_t cp) { + std::string r; + if (cp <= 0x7F) { + r.push_back(static_cast(cp)); + } else if (cp <= 0x7FF) { + r.push_back(static_cast(0xC0 | ((cp >> 6) & 0x1F))); + r.push_back(static_cast(0x80 | (cp & 0x3F))); + } else if (cp <= 0xFFFF) { + r.push_back(static_cast(0xE0 | ((cp >> 12) & 0x0F))); + r.push_back(static_cast(0x80 | ((cp >> 6) & 0x3F))); + r.push_back(static_cast(0x80 | (cp & 0x3F))); + } else if (cp <= 0x10FFFF) { + r.push_back(static_cast(0xF0 | ((cp >> 18) & 0x07))); + r.push_back(static_cast(0x80 | ((cp >> 12) & 0x3F))); + r.push_back(static_cast(0x80 | ((cp >> 6) & 0x3F))); + r.push_back(static_cast(0x80 | (cp & 0x3F))); + } + return r; +} + +// Read one UTF-8 codepoint from raw bytes, advance ptr. Returns 0xFFFD on error. +inline char32_t read_utf8(const char*& p, const char* end) { + if (p >= end) + return 0xFFFD; + unsigned char c = static_cast(*p); + if (c < 0x80) { + ++p; + return c; + } + if ((c & 0xE0) == 0xC0 && p + 1 < end) { + char32_t cp = + (static_cast(c & 0x1F) << 6) | (static_cast(p[1]) & 0x3F); + p += 2; + return cp; + } + if ((c & 0xF0) == 0xE0 && p + 2 < end) { + char32_t cp = (static_cast(c & 0x0F) << 12) | + (static_cast(static_cast(p[1]) & 0x3F) << 6) | + (static_cast(p[2]) & 0x3F); + p += 3; + return cp; + } + if ((c & 0xF8) == 0xF0 && p + 3 < end) { + char32_t cp = (static_cast(c & 0x07) << 18) | + (static_cast(static_cast(p[1]) & 0x3F) << 12) | + (static_cast(static_cast(p[2]) & 0x3F) << 6) | + (static_cast(p[3]) & 0x3F); + p += 4; + return cp; + } + ++p; + return 0xFFFD; +} + +// ─── GPT-2 byte encoder: byte value <-> Unicode codepoint ─── + +struct ByteEncoderTables { + // byte -> UTF-8 encoded string (precomputed for speed) + std::string byte_to_utf8[256]; + // Unicode codepoint -> byte value + std::unordered_map cp_to_byte; + + ByteEncoderTables() { + // GPT-2 byte encoder: printable bytes map to themselves, + // others map to 256+ to avoid control chars. + bool direct[256] = {}; + for (int b = 33; b <= 126; ++b) + direct[b] = true; // !"#$...~ + for (int b = 161; b <= 172; ++b) + direct[b] = true; // non-breaking space area + for (int b = 174; b <= 255; ++b) + direct[b] = true; // extended latin + + int n = 0; + for (int b = 0; b < 256; ++b) { + char32_t cp = direct[b] ? static_cast(b) : static_cast(256 + n++); + byte_to_utf8[b] = utf32_to_utf8(cp); + cp_to_byte[cp] = static_cast(b); + } + } +}; + +static const ByteEncoderTables& byte_tables() { + static ByteEncoderTables tables; + return tables; +} + +// ─── Hand-written BPE pre-tokenizer ─── +// +// Supports two regex variants: +// +// GPT-2 (used by GPT-2, Falcon, OPT): +// 'contractions| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+ +// +// Qwen3 (used by Qwen3, LLaMA-3, Mistral, Phi): +// (?i:contractions)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| +// ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+ +// +// Key differences: Qwen3 allows any non-CR/LF/letter/digit as optional prefix before +// letters, matches single digits, and has explicit newline handling. + +namespace pretok { + +enum class Variant { kGpt2, kQwen3, kBloom, kDeepSeek }; + +// ── Character classification ── + +struct UnicodeRange { + char32_t lo, hi; +}; + +constexpr UnicodeRange kLetterRanges[] = { + {'A', 'Z'}, {'a', 'z'}, {0xB5, 0xB5}, // µ (micro sign, treated as letter) + {0xC0, 0xD6}, {0xD8, 0xF6}, {0xF8, 0x1BA}, // Latin-1 + Extended A/B + {0x1BC, 0x1BF}, {0x1C4, 0x293}, {0x295, 0x2AF}, // Latin Extended B cont. + {0x370, 0x373}, {0x376, 0x377}, {0x37B, 0x37D}, {0x37F, 0x37F}, // Greek + {0x386, 0x386}, {0x388, 0x38A}, {0x38C, 0x38C}, {0x38E, 0x3A1}, + {0x3A3, 0x3F5}, {0x3F7, 0x481}, {0x48A, 0x52F}, // Greek + Cyrillic + {0x531, 0x556}, {0x559, 0x559}, // Armenian + {0x560, 0x588}, // Armenian lowercase + {0x600, 0x6FF}, // Arabic + {0x900, 0x97F}, // Devanagari + {0xE00, 0xE7F}, // Thai + {0x10A0, 0x10C5}, // Georgian + {0x13A0, 0x13F5}, // Cherokee + {0x1C90, 0x1CBA}, {0x1CBD, 0x1CBF}, // Georgian Extended + {0x1D00, 0x1D2B}, {0x1D6B, 0x1D77}, {0x1D79, 0x1D9A}, // Phonetic Extensions + {0x1E00, 0x1F15}, {0x1F18, 0x1F1D}, {0x1F20, 0x1F45}, // Latin Ext. Additional + Greek Ext. + {0x2C00, 0x2C5F}, // Glagolitic + {0x3040, 0x309F}, // Hiragana + {0x30A0, 0x30FF}, // Katakana + {0x3400, 0x4DBF}, // CJK Extension A + {0x4E00, 0x9FFF}, // CJK Unified Ideographs + {0xAC00, 0xD7AF}, // Hangul Syllables + {0xFB00, 0xFDFF}, // Alphabetic Presentation Forms + Arabic Forms + {0x10000, 0x1007F}, // Linear B Syllabary +}; + +inline bool is_letter(char32_t cp) { + for (const auto& r : kLetterRanges) { + if (cp >= r.lo && cp <= r.hi) + return true; + } + return false; +} + +inline bool is_digit(char32_t cp) { + return (cp >= '0' && cp <= '9'); +} + +inline bool is_whitespace(char32_t cp) { + return cp == ' ' || cp == '\t' || cp == '\n' || cp == '\r' || cp == 0x0B || cp == 0x0C // VT, FF + || cp == 0xA0 // non-breaking space + || cp == 0x2000 || cp == 0x200A // en space through hair space + || cp == 0x3000; // ideographic space +} + +// BLOOM punctuation set: .,!?... +constexpr char32_t kBloomPunct[] = { + '.', ',', '!', '?', + 0x2026, // ellipsis + 0x3002, // ideographic full stop + 0xFF0C, // fullwidth comma + 0x3001, // ideographic comma + 0x0964, // Devanagari danda + 0x06D4, // Arabic full stop + 0x060C, // Arabic comma +}; + +inline bool is_bloom_punct(char32_t cp) { + for (auto c : kBloomPunct) { + if (cp == c) + return true; + } + return false; +} + +inline bool is_newline(char32_t cp) { + return cp == '\n' || cp == '\r'; +} + +// Is this char a valid optional prefix before a word? +// GPT-2/BLOOM: only space (0x20) +// Qwen3: any char except CR, LF, letter, digit +// DeepSeek: any whitespace (\s?) +inline bool is_prefix(char32_t cp, Variant v) { + if (v == Variant::kQwen3) { + return !is_newline(cp) && !is_letter(cp) && !is_digit(cp); + } + if (v == Variant::kDeepSeek) { + return is_whitespace(cp); + } + return cp == ' '; +} + +// BLOOM word char: not whitespace and not BLOOM punctuation +inline bool is_bloom_word_char(char32_t cp) { + return !is_whitespace(cp) && !is_bloom_punct(cp); +} + +// ── Scanning helpers (advance pointer past matching chars) ── + +inline void scan_letters(const char*& p, const char* end) { + while (p < end) { + const char* peek = p; + if (!is_letter(read_utf8(peek, end))) + break; + p = peek; + } +} + +inline void scan_digits(const char*& p, const char* end, int max_group = 0) { + int count = 0; + while (p < end) { + if (max_group > 0 && count >= max_group) + break; + const char* peek = p; + if (!is_digit(read_utf8(peek, end))) + break; + p = peek; + ++count; + } +} + +inline void scan_others(const char*& p, const char* end) { + while (p < end) { + const char* peek = p; + char32_t nc = read_utf8(peek, end); + if (is_whitespace(nc) || is_letter(nc) || is_digit(nc)) + break; + p = peek; + } +} + +inline void scan_newlines(const char*& p, const char* end) { + while (p < end) { + const char* peek = p; + if (!is_newline(read_utf8(peek, end))) + break; + p = peek; + } +} + +inline void scan_all_whitespace(const char*& p, const char* end) { + while (p < end) { + const char* peek = p; + if (!is_whitespace(read_utf8(peek, end))) + break; + p = peek; + } +} + +// Qwen3 regex branch: \s*[\r\n]+ +// Consume any leading whitespace up to the first newline, then consume the +// contiguous newline run itself, but stop before whitespace that follows the +// last newline. +inline void scan_qwen3_newline_chunk(char32_t first_cp, const char*& p, const char* end) { + bool saw_newline = is_newline(first_cp); + while (p < end) { + const char* peek = p; + char32_t nc = read_utf8(peek, end); + if (!saw_newline) { + if (is_newline(nc)) { + saw_newline = true; + p = peek; + continue; + } + if (is_whitespace(nc)) { + p = peek; + continue; + } + break; + } + if (!is_newline(nc)) + break; + p = peek; + } +} + +inline void scan_bloom_words(const char*& p, const char* end) { + while (p < end) { + const char* peek = p; + if (!is_bloom_word_char(read_utf8(peek, end))) + break; + p = peek; + } +} + +// Emit whitespace run, leaving last ws char for next token's optional prefix +inline void emit_whitespace_leave_last(const char*& p, const char* end, const char* start, + std::vector& result) { + const char* last_ws_start = start; + while (p < end) { + const char* peek = p; + char32_t nc = read_utf8(peek, end); + if (!is_whitespace(nc)) + break; + last_ws_start = p; + p = peek; + } + if (p < end && last_ws_start > start) { + p = last_ws_start; + } + result.emplace_back(start, p); +} + +// ── Contraction matching ── + +inline bool is_two_char_contraction(char c1, char c2) { + return (c1 == 'r' && c2 == 'e') || (c1 == 'v' && c2 == 'e') || (c1 == 'l' && c2 == 'l'); +} + +// Returns length of contraction suffix after apostrophe ('s 't 'm 'd 're 've 'll) +inline int match_contraction_suffix(const char* p, const char* end) { + if (p >= end) + return 0; + char c = *p; + if (c == 's' || c == 't' || c == 'm' || c == 'd') + return 1; + if (p + 1 < end && is_two_char_contraction(c, p[1])) + return 2; + return 0; +} + +// ── GPT-2/Qwen3 pre-tokenize dispatch helpers ── + +// Handle apostrophe: either contraction or "other" chars run +inline bool try_contraction(char32_t cp, const char*& p, const char* end, const char* start, + std::vector& result) { + if (cp != '\'') + return false; + int suffix = match_contraction_suffix(p, end); + if (suffix > 0) { + p += suffix; + } else { + scan_others(p, end); + } + result.emplace_back(start, p); + return true; +} + +// Handle optional prefix + letter/digit/other run +inline bool try_prefix_run(char32_t cp, const char*& p, const char* end, const char* start, + Variant variant, std::vector& result) { + if (!is_prefix(cp, variant) || p >= end) + return false; + const char* after_prefix = p; + char32_t next_cp = read_utf8(p, end); + + if (is_letter(next_cp)) { + scan_letters(p, end); + result.emplace_back(start, p); + return true; + } + if (is_digit(next_cp)) { + // The regex prefix [^\r\n\p{L}\p{N}]? only applies before \p{L}+ (letters). + // Digits are matched by standalone \p{N} — no prefix. Back up so the + // digit is handled by try_simple_run instead. + if (variant == Variant::kQwen3) { + p = after_prefix; + return false; + } + scan_digits(p, end); + result.emplace_back(start, p); + return true; + } + if (!is_whitespace(next_cp)) { + scan_others(p, end); + if (variant == Variant::kQwen3) + scan_newlines(p, end); + result.emplace_back(start, p); + return true; + } + // Prefix followed by whitespace — back up, let whitespace handler deal with it + p = after_prefix; + return false; +} + +// Check if a whitespace run contains any newline character +inline bool has_newline_in_ws(char32_t first_cp, const char* p, const char* end) { + if (is_newline(first_cp)) + return true; + const char* scan = p; + while (scan < end) { + const char* peek = scan; + char32_t nc = read_utf8(peek, end); + if (!is_whitespace(nc)) + break; + if (is_newline(nc)) + return true; + scan = peek; + } + return false; +} + +// Handle whitespace (Qwen3 newline sequences + general whitespace) +inline bool try_whitespace_run(char32_t cp, const char*& p, const char* end, const char* start, + Variant variant, std::vector& result) { + if (!is_whitespace(cp)) + return false; + + // Qwen3: \s*[\r\n]+ — newline sequences take priority + if (variant == Variant::kQwen3 && has_newline_in_ws(cp, p, end)) { + scan_qwen3_newline_chunk(cp, p, end); + result.emplace_back(start, p); + return true; + } + + // General whitespace: leave last ws char for next token's prefix + emit_whitespace_leave_last(p, end, start, result); + return true; +} + +// Handle letter or digit run (no prefix) +inline bool try_simple_run(char32_t cp, const char*& p, const char* end, const char* start, + Variant variant, std::vector& result, int digit_group = 0) { + if (is_letter(cp)) { + scan_letters(p, end); + result.emplace_back(start, p); + return true; + } + if (is_digit(cp)) { + if (variant == Variant::kQwen3) { + if (digit_group > 1) + scan_digits(p, end, digit_group - 1); + } else { + scan_digits(p, end); + } + result.emplace_back(start, p); + return true; + } + return false; +} + +// ── Main pre-tokenize functions ── + +std::vector pre_tokenize(const std::string& text, Variant variant, + int digit_group = 0) { + std::vector result; + if (text.empty()) + return result; + + const char* p = text.data(); + const char* end = p + text.size(); + + while (p < end) { + const char* start = p; + char32_t cp = read_utf8(p, end); + + if (try_contraction(cp, p, end, start, result)) + continue; + if (try_prefix_run(cp, p, end, start, variant, result)) + continue; + if (try_whitespace_run(cp, p, end, start, variant, result)) + continue; + if (try_simple_run(cp, p, end, start, variant, result, digit_group)) + continue; + + // Other chars (punctuation/symbols) + scan_others(p, end); + if (variant == pretok::Variant::kQwen3) { + // Qwen3's ` ?[^\s\p{L}\p{N}]+[\r\n]*` keeps trailing newlines + // attached to punctuation/symbol runs even without an optional prefix. + scan_newlines(p, end); + } + result.emplace_back(start, p); + } + + return result; +} + +// ── BLOOM pre-tokenize helpers ── + +// Handle optional leading space + word chars for BLOOM +inline bool try_bloom_space_word(char32_t cp, const char*& p, const char* end, const char* start, + std::vector& result) { + if (cp != ' ') + return false; + if (p >= end) { + result.emplace_back(start, p); + return true; + } + const char* after_space = p; + char32_t next_cp = read_utf8(p, end); + if (is_bloom_word_char(next_cp)) { + scan_bloom_words(p, end); + result.emplace_back(start, p); + return true; + } + // Space followed by whitespace or punctuation — back up + p = after_space; + return false; +} + +// BLOOM pre-tokenizer: " ?[^(\s|[.,!?...])]+". +// Simpler than GPT-2: optional space + non-whitespace-non-punct chars. +// Punctuation chars become individual tokens. +std::vector bloom_pre_tokenize(const std::string& text) { + std::vector result; + if (text.empty()) + return result; + + const char* p = text.data(); + const char* end = p + text.size(); + + while (p < end) { + const char* start = p; + char32_t cp = read_utf8(p, end); + + if (try_bloom_space_word(cp, p, end, start, result)) + continue; + + if (is_whitespace(cp)) { + emit_whitespace_leave_last(p, end, start, result); + continue; + } + + if (is_bloom_punct(cp)) { + // Keep consecutive non-word chars together (Split "Isolated" behavior: + // unmatched text stays as one segment for BPE to merge) + while (p < end) { + const char* next_start = p; + char32_t next_cp = read_utf8(p, end); + if (!is_bloom_punct(next_cp) || is_whitespace(next_cp)) { + p = next_start; + break; + } + } + result.emplace_back(start, p); + continue; + } + + // Word chars (no leading space) + scan_bloom_words(p, end); + result.emplace_back(start, p); + } + + return result; +} + +} // namespace pretok + +// ─── BpeTokenizer implementation ─── + +class BpeTokenizer final : public ITokenizer { + public: + static std::unique_ptr Create(const char* tokenizer_json_data, + std::size_t tokenizer_json_size, + bool add_special_tokens = false) { + auto tokenizer = std::unique_ptr(new BpeTokenizer()); + tokenizer->mAddSpecialTokens = add_special_tokens; + tokenizer->parse_tokenizer_json(tokenizer_json_data, tokenizer_json_size); + return tokenizer; + } + + void encode_segment(const std::string& text, std::vector& result) const { + if (mIsSentencePiece) { + encode_sentencepiece(text, result); + } else if (mIsMetaspace) { + encode_metaspace(text, result); + } else { + encode_bytelevel(text, result); + } + } + + std::vector encode(const std::string& text) const override { + std::vector result; + + if (mAddSpecialTokens) { + for (int32_t bos_id : mPostBosIds) + result.push_back(bos_id); + } + + if (!text.empty()) { + auto segments = split_added_tokens(text); + for (const auto& seg : segments) { + if (seg.added_id >= 0) + result.push_back(seg.added_id); + else + encode_segment(seg.text, result); + } + } + + if (mAddSpecialTokens) { + for (int32_t eos_id : mPostEosIds) + result.push_back(eos_id); + } + + return result; + } + + std::string decode(const std::vector& ids) const override { + std::string joined = join_vocab_tokens(ids); + switch (mDecoderType) { + case DecoderType::kByteLevel: + return byte_decode(joined); + case DecoderType::kMetaspace: + return decode_metaspace(joined); + case DecoderType::kSequence: + return decode_sequence(joined); + } + return byte_decode(joined); + } + + int32_t id_for_token(std::string_view token) const override { + auto it = mTokenToId.find(std::string(token)); + return it != mTokenToId.end() ? it->second : -1; + } + + std::string token_for_id(int32_t id) const override { + if (id >= 0 && static_cast(id) < mVocab.size()) { + return mVocab[id]; + } + return ""; + } + + private: + BpeTokenizer() = default; + + enum class DecoderType { kByteLevel, kMetaspace, kSequence }; + + // ─── Added token segmentation ─── + + struct Segment { + std::string text; + int32_t added_id; /* -1 = normal */ + }; + + std::pair find_longest_added_token(const std::string& text, size_t pos) const { + int32_t best_id = -1; + size_t best_len = 0; + for (const auto& [content, id] : mAddedTokenPatterns) { + if (pos + content.size() <= text.size() && content.size() > best_len && + text.compare(pos, content.size(), content) == 0) { + best_id = id; + best_len = content.size(); + } + } + return {best_id, best_len}; + } + + std::vector split_added_tokens(const std::string& text) const { + std::vector segments; + if (mAddedTokenPatterns.empty()) { + segments.push_back({text, -1}); + return segments; + } + size_t pos = 0; + while (pos < text.size()) { + auto [best_id, best_len] = find_longest_added_token(text, pos); + if (best_id >= 0) { + segments.push_back({text.substr(pos, best_len), best_id}); + pos += best_len; + } else { + if (segments.empty() || segments.back().added_id >= 0) { + segments.push_back({"", -1}); + } + segments.back().text.push_back(text[pos]); + ++pos; + } + } + return segments; + } + + // ─── Encoding helpers ─── + + // Metaspace encode (DeepSeek style): raw UTF-8 char split, BPE merge. + // Spaces are dropped (not in vocab), Ġ (U+0120) used as separator. + void encode_metaspace(const std::string& text, std::vector& result) const { + std::vector chars; + const char* cp = text.data(); + const char* ce = cp + text.size(); + while (cp < ce) { + const char* cs = cp; + read_utf8(cp, ce); + std::string ch(cs, cp); + if (mTokenToId.count(ch)) { + chars.push_back(std::move(ch)); + } + } + auto tokens = apply_merges(std::move(chars)); + for (const auto& token : tokens) { + auto it = mTokenToId.find(token); + if (it != mTokenToId.end()) { + result.push_back(it->second); + } + } + } + + // SentencePiece-style encode: replace spaces with ▁, split to chars, BPE merge. + // Used for Metaspace pre-tokenizer and Sequence decoder models (LLaMA, Mistral, Phi-3). + // Normalize text for SentencePiece: replace spaces with ▁, handle prepend + std::string normalize_sentencepiece(const std::string& text) const { + static const std::string sp = "\xe2\x96\x81"; // U+2581 + std::string out; + if (mSentencePiecePrependAlways) { + out = sp; + } + for (char c : text) { + out += (c == ' ') ? sp : std::string(1, c); + } + // Metaspace prepend_scheme=first: prepend if not already starting with ▁ + if (mSentencePiecePrefixIfMissing && !mSentencePiecePrependAlways && + (out.empty() || out.compare(0, sp.size(), sp) != 0)) { + out = sp + out; + } + return out; + } + + void encode_sentencepiece(const std::string& text, std::vector& result) const { + std::string normalized = normalize_sentencepiece(text); + + // Split into UTF-8 characters + std::vector chars; + const char* cp = normalized.data(); + const char* ce = cp + normalized.size(); + while (cp < ce) { + const char* cs = cp; + read_utf8(cp, ce); + chars.emplace_back(cs, cp); + } + + // byte_fallback: replace unknown chars with <0xXX> + if (mByteFallback) { + chars = apply_byte_fallback_encode(std::move(chars)); + } + + // BPE merge and lookup + auto tokens = apply_merges(std::move(chars)); + for (const auto& token : tokens) { + auto it = mTokenToId.find(token); + if (it != mTokenToId.end()) { + result.push_back(it->second); + } + } + } + + // For byte_fallback: replace chars not in vocab with <0xXX> byte tokens + std::vector apply_byte_fallback_encode(std::vector chars) const { + std::vector result; + for (auto& ch : chars) { + if (mTokenToId.count(ch)) { + result.push_back(std::move(ch)); + } else { + // Split into individual bytes as <0xXX> + for (unsigned char byte : ch) { + char buf[8]; + std::snprintf(buf, sizeof(buf), "<0x%02X>", byte); + result.push_back(std::string(buf)); + } + } + } + return result; + } + + void encode_bytelevel(const std::string& text, std::vector& result) const { + std::vector words; + if (!mUsePreTokenizer) { + words = fallback_pre_tokenize(text); + } else if (mPreTokenizerVariant == pretok::Variant::kBloom) { + words = pretok::bloom_pre_tokenize(text); + } else { + words = pretok::pre_tokenize(text, mPreTokenizerVariant, mPreTokenizerDigitGroup); + } + + for (const auto& word : words) { + auto chars = byte_encode(word); + if (!chars.empty() && !mEndOfWordSuffix.empty()) { + chars.back() += mEndOfWordSuffix; + } + auto tokens = apply_merges(std::move(chars)); + for (const auto& token : tokens) { + auto it = mTokenToId.find(token); + if (it != mTokenToId.end()) { + result.push_back(it->second); + } + } + } + } + + // ─── Decoding helpers ─── + + std::string join_vocab_tokens(const std::vector& ids) const { + std::string joined; + for (int32_t id : ids) { + if (mSpecialIds.count(id)) + continue; + if (id >= 0 && static_cast(id) < mVocab.size()) { + joined += mVocab[id]; + } + } + return joined; + } + + std::string decode_metaspace(const std::string& joined) const { + std::string result; + const std::string g_char = utf32_to_utf8(0x0120); + const std::string spiece_char = "\xe2\x96\x81"; // U+2581 + size_t pos = 0; + while (pos < joined.size()) { + if (pos + g_char.size() <= joined.size() && + joined.compare(pos, g_char.size(), g_char) == 0) { + result.push_back(' '); + pos += g_char.size(); + } else if (pos + spiece_char.size() <= joined.size() && + joined.compare(pos, spiece_char.size(), spiece_char) == 0) { + result.push_back(' '); + pos += spiece_char.size(); + } else { + result.push_back(joined[pos]); + ++pos; + } + } + if (!result.empty() && result[0] == ' ') { + result.erase(0, 1); + } + return result; + } + + // ─── Sequence decoder (SentencePiece BPE models: LLaMA, Mistral, Phi-3) ─── + + std::string decode_sequence(const std::string& joined) const { + std::string text = joined; + // Step 1: Apply Replace operations (e.g. ▁ → space) + for (const auto& rep : mSeqDecoderReplaces) { + text = string_replace_all(text, rep.pattern, rep.content); + } + // Step 2: ByteFallback — convert <0xXX> tokens to raw bytes + if (mSeqDecoderByteFallback) { + text = apply_byte_fallback(text); + } + // Step 3: Fuse is implicit (tokens already joined) + // Step 4: Strip leading space + if (mSeqDecoderStripLeft && !text.empty() && text[0] == ' ') { + text.erase(0, 1); + } + return text; + } + + static std::string string_replace_all(const std::string& input, const std::string& from, + const std::string& to) { + if (from.empty()) + return input; + std::string result; + result.reserve(input.size()); + size_t pos = 0; + while (pos < input.size()) { + if (pos + from.size() <= input.size() && input.compare(pos, from.size(), from) == 0) { + result += to; + pos += from.size(); + } else { + result.push_back(input[pos]); + ++pos; + } + } + return result; + } + + static int hex_char_value(char c) { + if (c >= '0' && c <= '9') + return c - '0'; + if (c >= 'a' && c <= 'f') + return c - 'a' + 10; + if (c >= 'A' && c <= 'F') + return c - 'A' + 10; + return -1; + } + + // Try to parse <0xXX> at position pos. Returns parsed byte or -1. + static int try_parse_byte_token(const std::string& text, size_t pos) { + if (pos + 6 > text.size()) + return -1; + if (text[pos] != '<' || text[pos + 1] != '0' || text[pos + 2] != 'x' || + text[pos + 5] != '>') + return -1; + int h = hex_char_value(text[pos + 3]); + int l = hex_char_value(text[pos + 4]); + if (h < 0 || l < 0) + return -1; + return (h << 4) | l; + } + + static std::string apply_byte_fallback(const std::string& text) { + std::string result; + result.reserve(text.size()); + size_t pos = 0; + while (pos < text.size()) { + int byte_val = try_parse_byte_token(text, pos); + if (byte_val >= 0) { + result.push_back(static_cast(byte_val)); + pos += 6; + } else { + result.push_back(text[pos]); + ++pos; + } + } + return result; + } + + // ─── Byte-level encoding ─── + + static std::vector byte_encode(const std::string& text) { + const auto& tables = byte_tables(); + std::vector result; + result.reserve(text.size()); + for (unsigned char byte : text) { + result.push_back(tables.byte_to_utf8[byte]); + } + return result; + } + + static std::string byte_decode(const std::string& text) { + const auto& tables = byte_tables(); + std::string result; + result.reserve(text.size()); + size_t pos = 0; + while (pos < text.size()) { + size_t start = pos; + char32_t cp = utf8_to_char32(text, pos); + auto it = tables.cp_to_byte.find(cp); + if (it != tables.cp_to_byte.end()) { + result.push_back(static_cast(it->second)); + } else { + // Not a GPT-2 byte-encoded codepoint — pass through raw UTF-8 + result.append(text, start, pos - start); + } + } + return result; + } + + // Generic pre-tokenizer used when no GPT-2 pattern is declared. + + static std::vector fallback_pre_tokenize(const std::string& text) { + std::vector result; + if (text.empty()) + return result; + result.push_back(text); + return result; + } + + // ─── BPE merge algorithm ─── + + struct MergeCandidate { + int rank; + std::string first; + std::string second; + }; + + MergeCandidate find_best_merge(const std::vector& tokens) const { + MergeCandidate best{INT_MAX, "", ""}; + for (size_t i = 0; i + 1 < tokens.size(); ++i) { + auto it = + mMergeRank.find(std::make_pair(std::cref(tokens[i]), std::cref(tokens[i + 1]))); + if (it != mMergeRank.end() && it->second < best.rank) { + best = {it->second, tokens[i], tokens[i + 1]}; + } + } + return best; + } + + static std::vector merge_all_pairs(std::vector tokens, + const std::string& first, + const std::string& second) { + std::string merged = first + second; + std::vector result; + result.reserve(tokens.size()); + for (size_t i = 0; i < tokens.size(); ++i) { + if (i + 1 < tokens.size() && tokens[i] == first && tokens[i + 1] == second) { + result.push_back(merged); + ++i; // skip next + } else { + result.push_back(std::move(tokens[i])); + } + } + return result; + } + + // Optimized: merge ALL occurrences of the best pair per pass. + std::vector apply_merges(std::vector tokens) const { + if (tokens.size() <= 1) + return tokens; + + while (true) { + auto best = find_best_merge(tokens); + if (best.rank == INT_MAX) + break; + tokens = merge_all_pairs(std::move(tokens), best.first, best.second); + if (tokens.size() <= 1) + break; + } + + return tokens; + } + + // ─── JSON parsing helpers ─── + + static bool is_eos_content(const std::string& s) { + return s == "<|endoftext|>" || s == "" || s == "<|end_of_text|>"; + } + + void parse_vocab(const nlohmann::json& j) { + auto& vocab_obj = j["model"]["vocab"]; + size_t vocab_size = vocab_obj.size(); + mVocab.resize(vocab_size); + + for (auto& [token, id] : vocab_obj.items()) { + int32_t token_id = id.get(); + if (token_id >= 0 && token_id < static_cast(vocab_size)) { + mVocab[token_id] = token; + mTokenToId[token] = token_id; + } + } + } + + void parse_merges(const nlohmann::json& j) { + if (!j["model"].contains("merges")) + throw std::runtime_error("Invalid tokenizer.json: missing model.merges"); + + auto& merges_arr = j["model"]["merges"]; + mMergeRank.reserve(merges_arr.size()); + + for (size_t i = 0; i < merges_arr.size(); ++i) { + std::string first, second; + + if (merges_arr[i].is_array()) { + auto arr = merges_arr[i].get>(); + if (arr.size() != 2) + continue; + first = std::move(arr[0]); + second = std::move(arr[1]); + } else if (merges_arr[i].is_string()) { + std::string merge_str = merges_arr[i].get(); + auto space_pos = merge_str.find(' '); + if (space_pos == std::string::npos) + continue; + first = merge_str.substr(0, space_pos); + second = merge_str.substr(space_pos + 1); + } else { + continue; + } + + auto pair = std::make_pair(std::move(first), std::move(second)); + mMergeRank[pair] = static_cast(i); + } + } + + void parse_added_tokens(const nlohmann::json& j) { + if (!j.contains("added_tokens")) + return; + for (auto& tok : j["added_tokens"]) { + std::string content = tok["content"].get(); + int32_t id = tok["id"].get(); + + if (id >= 0 && static_cast(id) >= mVocab.size()) { + mVocab.resize(static_cast(id) + 1); + } + if (id >= 0) { + mVocab[id] = content; + mTokenToId[content] = id; + } + + if (tok.value("special", false)) { + mSpecialTokens[content] = id; + mSpecialIds.insert(id); + if (is_eos_content(content)) + mEosId = id; + } + // All added tokens (special and non-special) participate in pre-split + // matching during encode, matching HuggingFace's AddedToken behavior. + // The special flag only controls decode filtering (mSpecialIds) and + // post_processor BOS/EOS insertion. + mAddedTokenPatterns.push_back({content, id}); + } + // Sort by length descending for longest-match-first + std::sort(mAddedTokenPatterns.begin(), mAddedTokenPatterns.end(), + [](const auto& a, const auto& b) { return a.first.size() > b.first.size(); }); + } + + // Parse digit group size from regex pattern like \p{N}{1,3} → 3. + // Returns 0 if no digit grouping found (single digit or unlimited). + static int parse_digit_group(const std::string& regex) { + // Search for "\p{N}{" — the standalone grouped variant, not [^\p{N}] + auto pos = regex.find("\\p{N}{"); + if (pos == std::string::npos) + return 0; + auto open = pos + 6; // after "\\p{N}{" + auto close = regex.find('}', open); + if (close == std::string::npos) + return 0; + auto comma = regex.find(',', open); + if (comma == std::string::npos || comma >= close) + return 0; + return std::stoi(regex.substr(comma + 1, close - comma - 1)); + } + + // Classify a single Split regex string into a pre-tokenizer variant. + static pretok::Variant classify_split_regex(const std::string& regex, int& digit_group_out) { + if (regex.find("[^\r\n") != std::string::npos || + regex.find("[^\\r\\n") != std::string::npos) { + digit_group_out = parse_digit_group(regex); + return pretok::Variant::kQwen3; + } + if (regex.find("[^(\\s") != std::string::npos || + regex.find("[^(\\\\s") != std::string::npos) { + return pretok::Variant::kBloom; + } + if (regex.find("\\s?[A-Za-z") != std::string::npos) { + return pretok::Variant::kDeepSeek; + } + return pretok::Variant::kGpt2; + } + + // Detect variant from the Split regex inside a Sequence pre_tokenizer. + static pretok::Variant detect_split_variant(const nlohmann::json& pt, int& digit_group_out) { + digit_group_out = 0; + if (!pt.contains("pretokenizers")) + return pretok::Variant::kGpt2; + for (auto& sub : pt["pretokenizers"]) { + if (sub.value("type", "") != "Split") + continue; + if (!sub.contains("pattern") || !sub["pattern"].contains("Regex")) + continue; + return classify_split_regex(sub["pattern"]["Regex"].get(), + digit_group_out); + } + return pretok::Variant::kGpt2; + } + + static bool is_space_split_pre_tokenizer(const nlohmann::json& pt, const std::string& pt_type) { + if (pt_type != "Split") + return false; + if (!pt.contains("pattern") || !pt["pattern"].contains("String")) + return false; + return pt["pattern"]["String"].get() == " "; + } + + void detect_standalone_split_pre_tokenizer(const nlohmann::json& pt) { + if (pt.contains("pattern") && pt["pattern"].contains("Regex")) { + int digit_group = 0; + mPreTokenizerVariant = + classify_split_regex(pt["pattern"]["Regex"].get(), digit_group); + mPreTokenizerDigitGroup = digit_group; + return; + } + if (pt.contains("pattern") && pt["pattern"].contains("String") && + pt["pattern"]["String"] == " ") { + mUsePreTokenizer = false; + } + } + + void detect_pre_tokenizer(const nlohmann::json& j) { + mUsePreTokenizer = true; + mPreTokenizerVariant = pretok::Variant::kGpt2; + + if (!j.contains("pre_tokenizer") || j["pre_tokenizer"].is_null()) + return; + auto& pt = j["pre_tokenizer"]; + std::string pt_type = pt.value("type", ""); + + if (pt_type == "ByteLevel") { + mPreTokenizerVariant = pretok::Variant::kGpt2; + } else if (pt_type == "Sequence") { + int digit_group = 0; + mPreTokenizerVariant = detect_split_variant(pt, digit_group); + mPreTokenizerDigitGroup = digit_group; + } else if (is_space_split_pre_tokenizer(pt, pt_type)) { + // Gemma SentencePiece-BPE tokenizers use a direct Split(" ") + // pre-tokenizer with spaces already normalized to U+2581. + mUsePreTokenizer = false; + } else if (pt_type == "Metaspace") { + mIsMetaspace = true; + mUsePreTokenizer = false; + } else if (pt_type == "Split") { + detect_standalone_split_pre_tokenizer(pt); + } else if (pt_type.empty()) { + mUsePreTokenizer = false; + } else { + throw std::runtime_error("Unsupported pre_tokenizer type: " + pt_type); + } + } + + DecoderType classify_decoder_type(const nlohmann::json& j) const { + if (!j.contains("decoder") || j["decoder"].is_null()) { + return mIsMetaspace ? DecoderType::kMetaspace : DecoderType::kByteLevel; + } + std::string dt = j["decoder"].value("type", ""); + if (dt == "ByteLevel") + return DecoderType::kByteLevel; + if (dt == "Metaspace") + return DecoderType::kMetaspace; + if (dt == "Sequence") + return DecoderType::kSequence; + return mIsMetaspace ? DecoderType::kMetaspace : DecoderType::kByteLevel; + } + + void detect_decoder(const nlohmann::json& j) { + mDecoderType = classify_decoder_type(j); + if (mDecoderType == DecoderType::kSequence) { + parse_sequence_decoder(j["decoder"]); + } + // SentencePiece encode: vocab uses ▁ (U+2581) for spaces. + // Detect by: (1) ▁ in vocabulary, or (2) normalizer prepends ▁. + static const std::string spiece_marker = "\xe2\x96\x81"; // U+2581 + mIsSentencePiece = mTokenToId.count(spiece_marker) > 0 || mSentencePiecePrependAlways; + } + + void parse_seq_decoder_replace(const nlohmann::json& sub) { + SeqDecoderReplace rep; + if (sub.contains("pattern") && sub["pattern"].contains("String")) { + rep.pattern = sub["pattern"]["String"].get(); + } + rep.content = sub.value("content", ""); + if (!rep.pattern.empty()) { + mSeqDecoderReplaces.push_back(std::move(rep)); + } + } + + void parse_sequence_decoder(const nlohmann::json& dec) { + if (!dec.contains("decoders")) + return; + for (auto& sub : dec["decoders"]) { + std::string sub_type = sub.value("type", ""); + if (sub_type == "Replace") { + parse_seq_decoder_replace(sub); + } else if (sub_type == "ByteFallback") { + mSeqDecoderByteFallback = true; + } else if (sub_type == "Strip") { + if (sub.value("content", " ") == " " && sub.value("start", 0) > 0) { + mSeqDecoderStripLeft = true; + } + } + // Fuse is implicit (tokens already joined) + } + } + + static std::string optional_model_string(const nlohmann::json& model, const char* key) { + auto it = model.find(key); + if (it != model.end() && it->is_string()) + return it->get(); + return {}; + } + + void parse_tokenizer_json(const char* json_data, std::size_t json_size) { + nlohmann::json j; + try { + j = nlohmann::json::parse(json_data, json_data + json_size); + } catch (const std::exception& e) { + throw std::runtime_error(std::string("Failed to parse tokenizer.json: ") + e.what()); + } + + if (!j.contains("model")) + throw std::runtime_error("Invalid tokenizer.json: missing model"); + auto& model = j["model"]; + if (!model.contains("vocab") || !model["vocab"].is_object()) + throw std::runtime_error("Invalid tokenizer.json: model.vocab must be an object"); + if (!model.contains("merges") || !model["merges"].is_array()) + throw std::runtime_error("Invalid tokenizer.json: model.merges must be an array"); + + mByteFallback = model.value("byte_fallback", false); + mEndOfWordSuffix = optional_model_string(model, "end_of_word_suffix"); + parse_vocab(j); + parse_merges(j); + parse_added_tokens(j); + parse_post_processor(j); + detect_normalizer(j); + detect_pre_tokenizer(j); + detect_decoder(j); + } + + // Detect normalizer: check for Prepend (always prepend ▁) vs none + static bool normalizer_replaces_space(const nlohmann::json& norm) { + return norm.value("type", "") == "Replace" && norm.contains("pattern") && + norm["pattern"].contains("String") && norm["pattern"]["String"] == " " && + norm.value("content", "") == "\xe2\x96\x81"; + } + + void apply_sentencepiece_normalizer(const nlohmann::json& norm) { + if (norm.value("type", "") == "Prepend") { + mSentencePiecePrependAlways = true; + } else if (normalizer_replaces_space(norm)) { + mSentencePiecePrefixIfMissing = false; + } + } + + void detect_normalizer(const nlohmann::json& j) { + if (!j.contains("normalizer") || j["normalizer"].is_null()) + return; + auto& norm = j["normalizer"]; + std::string norm_type = norm.value("type", ""); + if (normalizer_replaces_space(norm)) { + mSentencePiecePrefixIfMissing = false; + return; + } + if (norm_type == "Sequence" && norm.contains("normalizers")) { + for (auto& sub : norm["normalizers"]) + apply_sentencepiece_normalizer(sub); + } + } + + // Parse post_processor to extract BOS/EOS tokens for add_special_tokens + void parse_post_processor(const nlohmann::json& j) { + if (!j.contains("post_processor") || j["post_processor"].is_null()) + return; + auto& pp = j["post_processor"]; + std::string pp_type = pp.value("type", ""); + + if (pp_type == "TemplateProcessing") { + parse_template_post_processor(pp); + } else if (pp_type == "Sequence") { + parse_sequence_post_processor(pp); + } else if (pp_type == "RobertaProcessing") { + parse_roberta_post_processor(pp); + } + // ByteLevel post_processor doesn't add tokens, nothing to do + } + + // Iterate Sequence post_processor's processors array to find TemplateProcessing + void parse_sequence_post_processor(const nlohmann::json& pp) { + if (!pp.contains("processors")) + return; + for (auto& sub : pp["processors"]) { + if (sub.value("type", "") == "TemplateProcessing") { + parse_template_post_processor(sub); + return; // first TemplateProcessing wins + } + } + } + + // RobertaProcessing: {"type":"RobertaProcessing","cls":["",0],"sep":["",2],...} + void parse_roberta_post_processor(const nlohmann::json& pp) { + if (pp.contains("cls") && pp["cls"].is_array() && pp["cls"].size() >= 2) { + int32_t cls_id = pp["cls"][1].get(); + mPostBosIds.push_back(cls_id); + } + if (pp.contains("sep") && pp["sep"].is_array() && pp["sep"].size() >= 2) { + int32_t sep_id = pp["sep"][1].get(); + mPostEosIds.push_back(sep_id); + } + } + + int32_t resolve_special_token_id(const nlohmann::json& entry) const { + if (!entry.contains("SpecialToken")) + return -1; + std::string token_id = entry["SpecialToken"].value("id", ""); + if (token_id.empty()) + return -1; + auto it = mSpecialTokens.find(token_id); + return it != mSpecialTokens.end() ? it->second : -1; + } + + void parse_template_post_processor(const nlohmann::json& pp) { + if (!pp.contains("single")) + return; + bool seen_sequence = false; + for (auto& entry : pp["single"]) { + if (entry.contains("Sequence")) { + seen_sequence = true; + continue; + } + int32_t id = resolve_special_token_id(entry); + if (id < 0) + continue; + if (seen_sequence) { + mPostEosIds.push_back(id); + } else { + mPostBosIds.push_back(id); + } + } + } + + // ─── Data members ─── + + std::vector mVocab; + std::unordered_map mTokenToId; + + struct PairHash { + size_t operator()(const std::pair& p) const { + size_t h1 = std::hash{}(p.first); + size_t h2 = std::hash{}(p.second); + return h1 ^ (h2 * 0x9e3779b97f4a7c15ULL + 0x9e3779b9 + (h1 << 6) + (h1 >> 2)); + } + }; + + std::unordered_map, int, PairHash> mMergeRank; + + std::unordered_map mSpecialTokens; + std::unordered_set mSpecialIds; // O(1) lookup in decode() + int32_t mEosId = -1; + + // Non-special added tokens: matched before pre-tokenization (longest first) + std::vector> mAddedTokenPatterns; + + bool mAddSpecialTokens = false; + bool mUsePreTokenizer = true; + bool mIsMetaspace = false; // Set by detect_pre_tokenizer. + bool mIsSentencePiece = false; + bool mSentencePiecePrependAlways = + false; // true for Normalizer Prepend, false for Metaspace first + bool mSentencePiecePrefixIfMissing = true; + bool mByteFallback = false; + std::string mEndOfWordSuffix; + + // Post-processor: BOS/EOS token IDs to add when add_special_tokens=true + // Vectors to support multiple BOS/EOS tokens (e.g. GLM-4: [gMASK] + ) + std::vector mPostBosIds; + std::vector mPostEosIds; + pretok::Variant mPreTokenizerVariant = pretok::Variant::kGpt2; + int mPreTokenizerDigitGroup = 0; // 0=unlimited, 3=\p{N}{1,3} (LLaMA 3.1/GLM-4) + + DecoderType mDecoderType = DecoderType::kByteLevel; + + // Sequence decoder config (parsed from tokenizer.json decoder field) + struct SeqDecoderReplace { + std::string pattern; + std::string content; + }; + std::vector mSeqDecoderReplaces; + bool mSeqDecoderByteFallback = false; + bool mSeqDecoderStripLeft = false; +}; + +} // namespace + +std::unique_ptr CreateLtx2BpeTokenizer(const char* tokenizer_json_data, + std::size_t tokenizer_json_size, + bool add_special_tokens) { + return BpeTokenizer::Create(tokenizer_json_data, tokenizer_json_size, add_special_tokens); +} + +} // namespace trtmc diff --git a/families/ltx2/runtime/distributed_runtime.cpp b/families/ltx2/runtime/distributed_runtime.cpp new file mode 100644 index 0000000000..8b8fb76ac0 --- /dev/null +++ b/families/ltx2/runtime/distributed_runtime.cpp @@ -0,0 +1,215 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "families/ltx2/runtime/distributed_runtime.h" + +#include "trtmc/runtime/dynamic_library.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc::ltx2 { +namespace { + +struct NcclUniqueId { + char internal[128]; +}; + +using NcclComm = void*; +using NcclResult = int; +using NcclGetUniqueIdFn = NcclResult (*)(NcclUniqueId*); +using NcclCommInitRankFn = NcclResult (*)(NcclComm*, int, NcclUniqueId, int); +using NcclCommDestroyFn = NcclResult (*)(NcclComm); +using NcclGetErrorStringFn = const char* (*)(NcclResult); +using NcclGetVersionFn = NcclResult (*)(int*); + +int require_env_int(const char* name) { + const char* raw = std::getenv(name); + if (raw == nullptr || *raw == '\0') + throw std::runtime_error(std::string("LTX-2.5 distributed runtime requires ") + name); + char* end = nullptr; + const long value = std::strtol(raw, &end, 10); + if (end == raw || *end != '\0') + throw std::runtime_error(std::string("LTX-2.5 distributed runtime has invalid ") + name); + return static_cast(value); +} + +int detect_world_size() { + return require_env_int("OMPI_COMM_WORLD_SIZE"); +} + +int detect_rank() { + return require_env_int("OMPI_COMM_WORLD_RANK"); +} + +int detect_local_rank() { + return require_env_int("OMPI_COMM_WORLD_LOCAL_RANK"); +} + +std::filesystem::path rendezvous_path() { + const char* path = std::getenv("TRTMC_NCCL_RENDEZVOUS"); + if (path == nullptr || *path == '\0') + throw std::runtime_error("LTX-2.5 distributed runtime requires TRTMC_NCCL_RENDEZVOUS"); + return path; +} + +class NcclRuntime { + public: + NcclRuntime() { + // NCCL is resolved at run time: TRTMC_NCCL_LIBRARY, else libnccl.so.2 + // (ELF) or nccl.dll (Windows) through the platform library search path. + const std::string library = platform::nccl_library(); + try { + library_ = std::make_unique(library, "LTX-2.5 runtime: NCCL"); + } catch (const std::exception& error) { + throw std::runtime_error(std::string(error.what()) + ". Set " + + platform::kNcclLibraryEnv + + " to the NCCL shared library to use."); + } + get_unique_id_ = load("ncclGetUniqueId"); + comm_init_rank_ = load("ncclCommInitRank"); + comm_destroy_ = load("ncclCommDestroy"); + get_error_string_ = load("ncclGetErrorString"); + int version = 0; + const auto get_version = + reinterpret_cast(library_->find_symbol("ncclGetVersion")); + if (get_version != nullptr && get_version(&version) != 0) + version = 0; + std::cerr << "[ltx2] NCCL " << version << " loaded from " << library_->loaded_path() + << std::endl; + } + + ~NcclRuntime() { + if (comm_ != nullptr) { + comm_destroy_(comm_); + comm_ = nullptr; + } + } + + void init(int size, int rank, const NcclUniqueId& id) { + check(comm_init_rank_(&comm_, size, id, rank), "ncclCommInitRank"); + } + + NcclUniqueId unique_id() { + NcclUniqueId id{}; + check(get_unique_id_(&id), "ncclGetUniqueId"); + return id; + } + + void* communicator() const { return comm_; } + + private: + template + T load(const char* symbol) { + return library_->require(symbol); + } + + void check(NcclResult result, const char* operation) const { + if (result == 0) + return; + const char* message = get_error_string_(result); + throw std::runtime_error(std::string(operation) + " failed: " + message); + } + + std::unique_ptr library_; + NcclComm comm_{nullptr}; + NcclGetUniqueIdFn get_unique_id_{nullptr}; + NcclCommInitRankFn comm_init_rank_{nullptr}; + NcclCommDestroyFn comm_destroy_{nullptr}; + NcclGetErrorStringFn get_error_string_{nullptr}; +}; + +void write_unique_id(const std::filesystem::path& path, const NcclUniqueId& id) { + if (!path.parent_path().empty()) + std::filesystem::create_directories(path.parent_path()); + const auto temporary = path.string() + ".tmp"; + { + std::ofstream output(temporary, std::ios::binary | std::ios::trunc); + if (!output) + throw std::runtime_error("Failed to write NCCL rendezvous file: " + temporary); + output.write(id.internal, sizeof(id.internal)); + } + std::filesystem::rename(temporary, path); +} + +NcclUniqueId read_unique_id(const std::filesystem::path& path) { + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(60); + while (!std::filesystem::exists(path)) { + if (std::chrono::steady_clock::now() > deadline) + throw std::runtime_error("Timed out waiting for NCCL rendezvous file: " + + path.string()); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + NcclUniqueId id{}; + std::ifstream input(path, std::ios::binary); + if (!input) + throw std::runtime_error("Failed to read NCCL rendezvous file: " + path.string()); + input.read(id.internal, sizeof(id.internal)); + if (input.gcount() != static_cast(sizeof(id.internal))) + throw std::runtime_error("Short NCCL rendezvous file: " + path.string()); + return id; +} + +void bind_cuda_device_for_local_rank(int local_rank) { + int count = 0; + const auto count_status = cudaGetDeviceCount(&count); + if (count_status != cudaSuccess) { + throw std::runtime_error(std::string("cudaGetDeviceCount failed for LTX-2.5 runtime: ") + + cudaGetErrorString(count_status)); + } + if (local_rank < 0 || local_rank >= count) { + throw std::runtime_error( + "LTX-2.5 distributed local rank is outside the visible CUDA device range"); + } + const auto status = cudaSetDevice(local_rank); + if (status != cudaSuccess) { + throw std::runtime_error( + std::string("cudaSetDevice failed for LTX-2.5 distributed rank: ") + + cudaGetErrorString(status)); + } +} + +} // namespace + +DistributedRuntimeGroup initialize_parallel_group(int parallel_size) { + DistributedRuntimeGroup group; + group.parallel_size = parallel_size; + if (parallel_size <= 1) + return group; + + group.world_size = detect_world_size(); + group.rank = detect_rank(); + if (group.world_size != parallel_size) { + throw std::runtime_error( + "LTX-2.5 distributed runtime requires launcher world size to equal parallel_size"); + } + if (group.rank < 0 || group.rank >= parallel_size) + throw std::runtime_error("LTX-2.5 distributed rank is outside parallel_size"); + + bind_cuda_device_for_local_rank(detect_local_rank()); + auto runtime = std::make_shared(); + const auto path = rendezvous_path(); + NcclUniqueId id{}; + if (group.rank == 0) { + id = runtime->unique_id(); + write_unique_id(path, id); + } else { + id = read_unique_id(path); + } + runtime->init(parallel_size, group.rank, id); + group.communicator = runtime->communicator(); + group.owner = std::move(runtime); + return group; +} + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/runtime/distributed_runtime.h b/families/ltx2/runtime/distributed_runtime.h new file mode 100644 index 0000000000..991d435129 --- /dev/null +++ b/families/ltx2/runtime/distributed_runtime.h @@ -0,0 +1,24 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include + +namespace trtmc::ltx2 { + +struct DistributedRuntimeGroup { + int world_size{1}; + int rank{0}; + int parallel_size{1}; + void* communicator{nullptr}; + std::shared_ptr owner; +}; + +// Initialize the NCCL communicator consumed by TensorRT distributed layers. +// Launcher discovery and communicator ownership remain local to LTX-2.5. +DistributedRuntimeGroup initialize_parallel_group(int parallel_size); + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/runtime/pipeline.cpp b/families/ltx2/runtime/pipeline.cpp new file mode 100644 index 0000000000..e2f1631e21 --- /dev/null +++ b/families/ltx2/runtime/pipeline.cpp @@ -0,0 +1,403 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "families/ltx2/runtime/pipeline.h" + +#include "families/ltx2/runtime/portable_normal.h" +#include "families/ltx2/runtime/runtime_math.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc { +namespace { + +using Clock = std::chrono::steady_clock; + +double elapsed_ms(Clock::time_point start, Clock::time_point end) { + return std::chrono::duration(end - start).count(); +} + +float half_to_float(uint16_t h) { + const uint32_t sign = (static_cast(h) & 0x8000U) << 16U; + const uint32_t exp = (h >> 10U) & 0x1FU; + uint32_t mant = h & 0x3FFU; + uint32_t bits = sign; + if (exp == 31U) { + bits |= 0x7F800000U | (mant << 13U); + } else if (exp != 0U) { + bits |= ((exp - 15U + 127U) << 23U) | (mant << 13U); + } else if (mant != 0U) { + int32_t e = -1; + do { + ++e; + mant <<= 1U; + } while ((mant & 0x400U) == 0U); + bits |= (static_cast(127 - 15 - e) << 23U) | ((mant & 0x3FFU) << 13U); + } + float out; + std::memcpy(&out, &bits, sizeof(out)); + return out; +} + +float bf16_to_float(uint16_t h) { + const uint32_t bits = static_cast(h) << 16U; + float out; + std::memcpy(&out, &bits, sizeof(out)); + return out; +} + +const Tensor& require_output(const TensorMap& outputs, const std::string& name, DType dtype, + std::size_t count) { + const auto it = outputs.find(name); + if (it == outputs.end()) + throw std::runtime_error("LTX-2.5 engine output is missing: " + name); + if (it->second.dtype != dtype || it->second.numel() != count || it->second.data == nullptr) + throw std::runtime_error("LTX-2.5 engine output does not match its contract: " + name); + return it->second; +} + +std::vector float_output(const TensorMap& outputs, const std::string& name, + std::size_t count) { + const auto it = outputs.find(name); + if (it == outputs.end()) + throw std::runtime_error("LTX-2.5 engine output is missing: " + name); + const auto& tensor = it->second; + if (tensor.data == nullptr || tensor.numel() != count) + throw std::runtime_error("LTX-2.5 engine output does not match its contract: " + name); + std::vector out(count); + if (tensor.dtype == DType::kFloat32) { + std::memcpy(out.data(), tensor.data, count * sizeof(float)); + } else if (tensor.dtype == DType::kFloat16) { + const auto* src = static_cast(tensor.data); + for (std::size_t i = 0; i < count; ++i) + out[i] = half_to_float(src[i]); + } else if (tensor.dtype == DType::kBFloat16) { + const auto* src = static_cast(tensor.data); + for (std::size_t i = 0; i < count; ++i) + out[i] = bf16_to_float(src[i]); + } else { + throw std::runtime_error("LTX-2.5 engine output has an unsupported dtype: " + name); + } + return out; +} + +std::vector bf16_output(const TensorMap& outputs, const std::string& name, + std::size_t count) { + const auto& tensor = require_output(outputs, name, DType::kBFloat16, count); + const auto* src = static_cast(tensor.data); + return {src, src + count}; +} + +std::string trim(const std::string& text) { + const auto begin = text.find_first_not_of(" \t\r\n"); + if (begin == std::string::npos) + return {}; + const auto end = text.find_last_not_of(" \t\r\n"); + return text.substr(begin, end - begin + 1); +} + +// Optional family diagnostics (all off by default): +// TRTMC_LTX2_INITIAL_LATENTS raw fp32 file: packed video [S, C] then audio [Sa, Ca] noise +// (replaces the seeded noise, e.g. a reference pipeline's draw) +// TRTMC_LTX2_DUMP_LATENTS raw fp32 file written with the final video then audio latents +std::vector read_f32_file(const char* path) { + std::ifstream input(path, std::ios::binary | std::ios::ate); + if (!input) + throw std::runtime_error(std::string("cannot read LTX-2.5 initial latents: ") + path); + const auto size = static_cast(input.tellg()); + if (size % sizeof(float) != 0) + throw std::runtime_error("LTX-2.5 initial latents file is not fp32"); + std::vector values(size / sizeof(float)); + input.seekg(0); + input.read(reinterpret_cast(values.data()), static_cast(size)); + return values; +} + +void maybe_dump(const std::vector& video, const std::vector& audio) { + const char* path = std::getenv("TRTMC_LTX2_DUMP_LATENTS"); + if (path == nullptr || *path == '\0') + return; + std::ofstream output(path, std::ios::binary | std::ios::trunc); + output.write(reinterpret_cast(video.data()), + static_cast(video.size() * 4)); + output.write(reinterpret_cast(audio.data()), + static_cast(audio.size() * 4)); + std::cerr << "[ltx2] wrote final latents (" << video.size() << " + " << audio.size() + << " fp32) to " << path << "\n"; +} + +const std::array& config_fields() { + static const std::array fields{{ + {"seed", internal::ConfigKind::I64, internal::ConfigValue{std::int64_t{0}}, + "Seed of the initial video and audio noise (portable std::mt19937 + normal draws)."}, + }}; + return fields; +} + +internal::AudioVideoResult worker_completion(const LTX2Options& options) { + internal::AudioVideoResult result; + result.video.frames.num_frames = 0; + result.video.frames.height = 0; + result.video.frames.width = 0; + result.video.frames.channels = 3; + result.audio.sample_rate = static_cast(options.audio_sample_rate); + result.audio.channels = static_cast(options.audio_channels); + return result; +} + +} // namespace + +LTX2Options parse_ltx2_options(const std::string& runtime_json) { + const auto doc = nlohmann::json::parse(runtime_json); + LTX2Options o; + o.video_frames = doc.at("video_frames").get(); + o.video_height = doc.at("video_height").get(); + o.video_width = doc.at("video_width").get(); + o.latent_frames = doc.at("latent_frames").get(); + o.latent_height = doc.at("latent_height").get(); + o.latent_width = doc.at("latent_width").get(); + o.latent_channels = doc.at("latent_channels").get(); + o.audio_frames = doc.at("audio_frames").get(); + o.audio_latent_channels = doc.at("audio_latent_channels").get(); + o.text_seq_len = doc.at("text_seq_len").get(); + o.frame_rate = doc.at("frame_rate").get(); + o.pad_token_id = doc.at("pad_token_id").get(); + o.sigmas = doc.at("sigmas").get>(); + o.audio_sample_rate = doc.at("audio_sample_rate").get(); + o.audio_channels = doc.at("audio_channels").get(); + if (o.sigmas.size() < 2 || o.sigmas.back() != 0.0F) + throw std::runtime_error("LTX-2.5 runtime.json sigmas must end with the terminal 0"); + if (doc.at("dit_batch").get() != 1) + throw std::runtime_error("LTX-2.5 runtime runs the distilled (batch 1) denoiser"); + if (o.video_tokens() <= 0 || o.audio_frames <= 0 || o.text_seq_len <= 0) + throw std::runtime_error("LTX-2.5 runtime.json has invalid shapes"); + return o; +} + +LTX2Pipeline::LTX2Pipeline(std::unique_ptr text_encoder, + std::unique_ptr denoiser, std::unique_ptr vae, + std::unique_ptr audio, LTX2Options options, + std::shared_ptr tokenizer, + LTX2DistributedContext distributed) + : distributed_(std::move(distributed)), text_encoder_(std::move(text_encoder)), + denoiser_(std::move(denoiser)), vae_(std::move(vae)), audio_(std::move(audio)), + options_(std::move(options)), tokenizer_(std::move(tokenizer)), progress_(distributed_.rank) { +} + +LTX2Pipeline::~LTX2Pipeline() = default; + +std::vector LTX2Pipeline::task_bindings() { + const auto& fields = config_fields(); + return {internal::bind(*this, {fields.data(), fields.size()})}; +} + +LTX2Pipeline::TextContext LTX2Pipeline::encode(const std::string& text) { + std::vector ids; + std::vector mask; + ltx2_prompt_ids(tokenizer_->encode(trim(text)), options_.text_seq_len, options_.pad_token_id, + ids, mask); + const int64_t L = options_.text_seq_len; + TensorMap inputs; + inputs["input_ids"] = Tensor{ids.data(), {1, L}, DType::kInt32}; + inputs["attention_mask"] = Tensor{mask.data(), {1, L}, DType::kInt32}; + const auto outputs = text_encoder_->forward(inputs); + const auto video_dim = + static_cast(text_encoder_->tensor_shape("video_context").back()); + const auto audio_dim = + static_cast(text_encoder_->tensor_shape("audio_context").back()); + TextContext context; + context.video = bf16_output(outputs, "video_context", static_cast(L) * video_dim); + context.audio = bf16_output(outputs, "audio_context", static_cast(L) * audio_dim); + return context; +} + +void LTX2Pipeline::run_dit(const std::vector& video, const std::vector& audio, + const TextContext& text, float timestep, std::vector& video_out, + std::vector& audio_out) { + const int64_t S = options_.video_tokens(); + const int64_t Sa = options_.audio_frames; + const int64_t L = options_.text_seq_len; + std::vector t{timestep}; + std::vector keep{1.0F}; + TensorMap inputs; + inputs["video_latent"] = + Tensor{const_cast(video.data()), {1, S, options_.latent_channels}, DType::kFloat32}; + inputs["audio_latent"] = Tensor{ + const_cast(audio.data()), {1, Sa, options_.audio_latent_channels}, DType::kFloat32}; + inputs["video_context"] = Tensor{const_cast(text.video.data()), + {1, L, static_cast(text.video.size()) / L}, + DType::kBFloat16}; + inputs["audio_context"] = Tensor{const_cast(text.audio.data()), + {1, L, static_cast(text.audio.size()) / L}, + DType::kBFloat16}; + inputs["timestep"] = Tensor{t.data(), {1}, DType::kFloat32}; + inputs["stg_keep"] = Tensor{keep.data(), {1}, DType::kFloat32}; + inputs["av_keep"] = Tensor{keep.data(), {1}, DType::kFloat32}; + const auto outputs = denoiser_->forward(inputs); + video_out = float_output(outputs, "video_velocity", video.size()); + audio_out = float_output(outputs, "audio_velocity", audio.size()); +} + +std::vector LTX2Pipeline::decode_video(const std::vector& video_latents) { + TensorMap inputs; + inputs["latents"] = Tensor{const_cast(video_latents.data()), + {1, options_.video_tokens(), options_.latent_channels}, + DType::kFloat32}; + const auto outputs = vae_->forward(inputs); + const auto count = static_cast(options_.video_frames) * options_.video_height * + options_.video_width * 3U; + return float_output(outputs, "frames", count); +} + +std::vector LTX2Pipeline::decode_audio(const std::vector& audio_latents) { + TensorMap inputs; + inputs["audio_latents"] = Tensor{const_cast(audio_latents.data()), + {1, options_.audio_frames, options_.audio_latent_channels}, + DType::kFloat32}; + const auto outputs = audio_->forward(inputs); + const auto shape = audio_->tensor_shape("waveform"); + std::size_t count = 1; + for (const auto dim : shape) + count *= static_cast(dim); + return float_output(outputs, "waveform", count); +} + +internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoRequest& request, + internal::ConfigView config) { + const auto& fields = config_fields(); + internal::validate_config({fields.data(), fields.size()}, config); + const auto seed = + internal::config_get(config, {fields.data(), fields.size()}, "seed").value(); + const std::string prompt(request.prompt); + const int32_t steps = static_cast(options_.sigmas.size()) - 1; + const auto video_count = + static_cast(options_.video_tokens()) * options_.latent_channels; + const auto audio_count = + static_cast(options_.audio_frames) * options_.audio_latent_channels; + + const auto t_start = Clock::now(); + if (progress_.enabled()) { + std::ostringstream detail; + detail << "world_size=" << distributed_.world_size << " frames=" << options_.video_frames + << " width=" << options_.video_width << " height=" << options_.video_height + << " steps=" << steps; + progress_.start(detail.str()); + progress_.emit("encode_begin"); + } + const auto text = encode(prompt); + const auto t_text = Clock::now(); + progress_.emit("encode_end"); + + std::vector video(video_count); + std::vector audio(audio_count); + if (const char* path = std::getenv("TRTMC_LTX2_INITIAL_LATENTS"); + path != nullptr && *path != '\0') { + const auto values = read_f32_file(path); + if (values.size() != video_count + audio_count) + throw std::runtime_error( + "TRTMC_LTX2_INITIAL_LATENTS must hold the packed video then audio noise"); + std::copy_n(values.begin(), video_count, video.begin()); + std::copy_n(values.begin() + static_cast(video_count), audio_count, + audio.begin()); + } else { + std::mt19937 generator(static_cast(seed)); + ltx2::LibstdcxxNormalFloat normal; + for (auto& v : video) + v = normal(generator); + for (auto& v : audio) + v = normal(generator); + } + + progress_.emit("denoise_begin"); + std::vector step_ms; + std::vector video_v; + std::vector audio_v; + for (int32_t step = 0; step < steps; ++step) { + const auto step_start = Clock::now(); + const float sigma = options_.sigmas[static_cast(step)]; + const float sigma_next = options_.sigmas[static_cast(step) + 1]; + run_dit(video, audio, text, sigma * 1000.0F, video_v, audio_v); + ltx2_euler_step(video, video_v, sigma, sigma_next); + ltx2_euler_step(audio, audio_v, sigma, sigma_next); + step_ms.push_back(elapsed_ms(step_start, Clock::now())); + if (progress_.enabled()) { + std::ostringstream detail; + detail << "step=" << (step + 1) << "/" << steps << " step_ms=" << std::fixed + << std::setprecision(3) << step_ms.back(); + progress_.emit("step", detail.str()); + } + } + const auto t_denoise = Clock::now(); + progress_.emit("denoise_end"); + + if (distributed_.world_size > 1 && distributed_.rank != 0) { + std::cerr << "[ltx2] context-parallel rank " << distributed_.rank + << " finished denoising in " << elapsed_ms(t_text, t_denoise) + << " ms; rank 0 decodes the video and audio\n"; + progress_.emit("worker_done"); + return worker_completion(options_); + } + maybe_dump(video, audio); + + progress_.emit("vae_begin"); + auto frames = decode_video(video); + const auto t_vae = Clock::now(); + progress_.emit("vae_end"); + progress_.emit("audio_begin"); + const auto wave = decode_audio(audio); + const auto t_audio = Clock::now(); + progress_.emit("audio_end"); + + internal::AudioVideoResult result; + result.video.frames.pixels = std::move(frames); + result.video.frames.height = options_.video_height; + result.video.frames.width = options_.video_width; + result.video.frames.channels = 3; + result.video.frames.num_frames = options_.video_frames; + result.video.timestamps_seconds.reserve(static_cast(options_.video_frames)); + for (int32_t f = 0; f < options_.video_frames; ++f) + result.video.timestamps_seconds.push_back(static_cast(f) / options_.frame_rate); + result.audio.samples = ltx2_interleave(wave, options_.audio_channels); + result.audio.sample_rate = static_cast(options_.audio_sample_rate); + result.audio.channels = static_cast(options_.audio_channels); + result.audio_start_seconds = 0.0; + result.video.inference_ms = elapsed_ms(t_start, t_audio); + + std::vector sorted = step_ms; + std::sort(sorted.begin(), sorted.end()); + const double median = sorted.empty() ? 0.0 : sorted[sorted.size() / 2]; + std::cerr << std::fixed << std::setprecision(3) + << "[ltx2-perf-json] {\"world_size\":" << distributed_.world_size + << ",\"text_encode_ms\":" << elapsed_ms(t_start, t_text) + << ",\"denoise_ms\":" << elapsed_ms(t_text, t_denoise) + << ",\"median_step_ms\":" << median + << ",\"vae_decode_ms\":" << elapsed_ms(t_denoise, t_vae) + << ",\"audio_decode_ms\":" << elapsed_ms(t_vae, t_audio) + << ",\"generate_ms\":" << elapsed_ms(t_start, t_audio) << ",\"num_steps\":" << steps + << "}\n"; + if (progress_.enabled()) { + std::ostringstream detail; + detail << "generate_ms=" << std::fixed << std::setprecision(3) + << elapsed_ms(t_start, t_audio); + progress_.emit("done", detail.str()); + } + return result; +} + +} // namespace trtmc diff --git a/families/ltx2/runtime/pipeline.h b/families/ltx2/runtime/pipeline.h new file mode 100644 index 0000000000..9d05a54dab --- /dev/null +++ b/families/ltx2/runtime/pipeline.h @@ -0,0 +1,98 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +// LTX2Pipeline: native C++ runtime for Lightricks LTX-2.5 (text -> synchronized video + audio). +// All model execution goes through TensorRT component engines: +// text_encoder.plan Gemma 4 + LTX2TextConnectors -> video/audio text contexts +// denoiser.plan joint audio/video DiT (single device or context parallel) +// vae.plan video VAE decoder -> RGB frames +// audio.plan audio VAE decoder + vocoder with BWE -> 48 kHz stereo + +#include "families/ltx2/runtime/progress_log.h" +#include "families/ltx2/runtime/tokenizer.h" +#include "trtmc/internal/model.h" +#include "trtmc/internal/video.h" +#include "trtmc/runtime/trt_module.h" + +#include +#include +#include +#include +#include + +namespace trtmc { + +struct LTX2Options { + int32_t video_frames{121}; + int32_t video_height{544}; + int32_t video_width{960}; + int32_t latent_frames{16}; + int32_t latent_height{17}; + int32_t latent_width{30}; + int32_t latent_channels{128}; + int32_t audio_frames{126}; + int32_t audio_latent_channels{128}; + int32_t text_seq_len{1024}; + float frame_rate{24.0F}; + int32_t pad_token_id{0}; + // Full schedule including the terminal 0: the model runs at sigmas[0..n-1]. + std::vector sigmas; + int32_t audio_sample_rate{48000}; + int32_t audio_channels{2}; + + int64_t video_tokens() const { return int64_t(latent_frames) * latent_height * latent_width; } +}; + +LTX2Options parse_ltx2_options(const std::string& runtime_json); + +// Context-parallel participation. The owner keeps the NCCL communicator used by the +// denoiser engine alive for the pipeline lifetime. Rank 0 decodes and returns media; +// other ranks return the worker completion. +struct LTX2DistributedContext { + std::shared_ptr owner; + int32_t rank{0}; + int32_t world_size{1}; +}; + +class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudioVideo { + public: + LTX2Pipeline(std::unique_ptr text_encoder, std::unique_ptr denoiser, + std::unique_ptr vae, std::unique_ptr audio, + LTX2Options options, std::shared_ptr tokenizer, + LTX2DistributedContext distributed = {}); + ~LTX2Pipeline() override; + + const char* task() const noexcept override { return ITextToAudioVideo::kTask.data(); } + std::vector task_bindings() override; + internal::AudioVideoResult run(const internal::TextToAudioVideoRequest& request, + internal::ConfigView config) override; + + struct TextContext { + std::vector video; // [1, L, 4096] bf16 bits + std::vector audio; // [1, L, 2048] bf16 bits + }; + + private: + TextContext encode(const std::string& text); + void run_dit(const std::vector& video, const std::vector& audio, + const TextContext& text, float timestep, std::vector& video_out, + std::vector& audio_out); + std::vector decode_video(const std::vector& video_latents); + std::vector decode_audio(const std::vector& audio_latents); + + // Declared first so the communicator outlives every engine that uses it. + LTX2DistributedContext distributed_; + std::unique_ptr text_encoder_; + std::unique_ptr denoiser_; + std::unique_ptr vae_; + std::unique_ptr audio_; + LTX2Options options_; + std::shared_ptr tokenizer_; + LTX2ProgressLog progress_; +}; + +} // namespace trtmc diff --git a/families/ltx2/runtime/plugin.cpp b/families/ltx2/runtime/plugin.cpp new file mode 100644 index 0000000000..190195bb1e --- /dev/null +++ b/families/ltx2/runtime/plugin.cpp @@ -0,0 +1,83 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "families/ltx2/runtime/distributed_runtime.h" +#include "families/ltx2/runtime/pipeline.h" +#include "families/ltx2/runtime/runtime_config.h" +#include "families/ltx2/runtime/tokenizer.h" +#include "trtmc/runtime/family_factory.h" +#include "trtmc/runtime/trt_backend.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc::ltx2 { +namespace { + +std::vector require_section(const BundleReader& bundle, const char* name) { + const auto* section = bundle.find_section(name); + if (section == nullptr || section->length == 0) + throw std::runtime_error("LTX-2.5 bundle section is missing or empty: " + + std::string(name)); + return bundle.read_section(name); +} + +std::unique_ptr load(IBackend& backend, const BundleReader& bundle, const char* name, + const ModuleCreateOptions& options) { + const auto start = std::chrono::steady_clock::now(); + auto plan = require_section(bundle, name); + auto module = backend.create_module(plan.data(), plan.size(), options); + if (!module || !module->ok()) + throw std::runtime_error(std::string("failed to load LTX-2.5 engine ") + name); + module->set_timing_label(name); + std::cerr << "[trtmc.load_timing] label=\"" << name << "\" load_deserialize_ms=" + << std::chrono::duration(std::chrono::steady_clock::now() - start) + .count() + << " plan_bytes=" << plan.size() << '\n'; + return module; +} + +} // namespace +} // namespace trtmc::ltx2 + +extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext& context) { + using namespace trtmc; + if (context.kv_cache_size_bytes != 0) + throw std::invalid_argument("ltx2 does not support --kv-cache-size"); + const auto runtime_data = ltx2::require_section(context.reader, "runtime.json"); + const std::string runtime(runtime_data.begin(), runtime_data.end()); + const auto parallel = ltx2::parse_parallel_runtime_config(runtime); + // Binds this rank's CUDA device before any engine is deserialized. + const auto group = ltx2::initialize_parallel_group(parallel.size); + auto options = parse_ltx2_options(runtime); + ModuleCreateOptions plain{}; + ModuleCreateOptions denoiser_options{}; + if (parallel.distributed()) { + denoiser_options.distributed_communicator = group.communicator; + denoiser_options.distributed_owner = group.owner; + } + auto text = ltx2::load(context.backend, context.reader, "text_encoder.plan", plain); + auto denoiser = ltx2::load(context.backend, context.reader, "denoiser.plan", denoiser_options); + // Only rank 0 decodes and returns media; worker ranks never load the decoders. + std::unique_ptr vae; + std::unique_ptr audio; + if (group.rank == 0) { + vae = ltx2::load(context.backend, context.reader, "vae.plan", plain); + audio = ltx2::load(context.backend, context.reader, "audio.plan", plain); + } + const auto tokenizer_data = ltx2::require_section(context.reader, "tokenizer.json"); + std::shared_ptr tokenizer = CreateLtx2BpeTokenizer( + tokenizer_data.data(), tokenizer_data.size(), /*add_special_tokens=*/false); + if (!tokenizer) + throw std::runtime_error("LTX-2.5 bundle tokenizer.json is not a supported BPE tokenizer"); + return new LTX2Pipeline(std::move(text), std::move(denoiser), std::move(vae), std::move(audio), + std::move(options), std::move(tokenizer), + LTX2DistributedContext{group.owner, group.rank, group.world_size}); +} diff --git a/families/ltx2/runtime/portable_normal.h b/families/ltx2/runtime/portable_normal.h new file mode 100644 index 0000000000..f4e389a3e7 --- /dev/null +++ b/families/ltx2/runtime/portable_normal.h @@ -0,0 +1,56 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include +#include +#include + +namespace trtmc::ltx2 { + +// std::normal_distribution is implementation-defined: libstdc++ and the MSVC +// STL draw different values from the same std::mt19937 seed, so one seed +// would give a different video per platform. This class reproduces the +// libstdc++ algorithm (generate_canonical plus the Marsaglia polar +// method with one cached value) step by step in float/double arithmetic, so +// other standard libraries draw the same initial latents as Linux builds, up +// to last-bit logf differences between C runtimes. +class LibstdcxxNormalFloat { + public: + float operator()(std::mt19937& generator) { + if (saved_available_) { + saved_available_ = false; + return saved_; + } + float x = 0.0F; + float y = 0.0F; + float r2 = 0.0F; + do { + x = static_cast(static_cast(2.0F * canonical(generator)) - 1.0); + y = static_cast(static_cast(2.0F * canonical(generator)) - 1.0); + r2 = x * x + y * y; + } while (r2 > 1.0 || r2 == 0.0); + // Same float log/sqrt calls as libstdc++. On other C runtimes, logf + // may differ from glibc in the last bit for a few inputs. + const float multiplier = std::sqrt(-2.0F * std::log(r2) / r2); + saved_ = x * multiplier; + saved_available_ = true; + return y * multiplier; + } + + private: + // generate_canonical over a 32-bit engine: one draw, rounded + // to float, scaled by 2^-32, clamped below 1. + static float canonical(std::mt19937& generator) { + const float value = static_cast(generator() - std::mt19937::min()) / 4294967296.0F; + return value >= 1.0F ? std::nextafter(1.0F, 0.0F) : value; + } + + float saved_{0.0F}; + bool saved_available_{false}; +}; + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/runtime/progress_log.h b/families/ltx2/runtime/progress_log.h new file mode 100644 index 0000000000..ef5fd713ef --- /dev/null +++ b/families/ltx2/runtime/progress_log.h @@ -0,0 +1,77 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +// Opt-in LTX-2.5 progress log (same line format as LTX-Video) with monotonic timestamps. +// +// Set TRTMC_LTX2_PROGRESS=1 to print one flushed stdout line per generation +// phase boundary and per denoising step, e.g. +// +// [ltx-progress] rank=0 t_ms=1234.567 event=step step=3/50 step_ms=305.123 +// +// t_ms is std::chrono::steady_clock time since the start of generate_image +// (text encoder -> denoise -> video and audio decode), after all engines are loaded. Unset, +// empty or "0" disables the log; nothing else about the run changes. + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc { + +constexpr const char* kLTX2ProgressEnv = "TRTMC_LTX2_PROGRESS"; + +inline bool ltx2_progress_enabled(const char* value) { + return value != nullptr && *value != '\0' && std::strcmp(value, "0") != 0; +} + +inline std::string format_ltx2_progress(int32_t rank, double t_ms, const std::string& event, + const std::string& detail) { + char stamp[32]; + std::snprintf(stamp, sizeof(stamp), "%.3f", t_ms); + std::ostringstream line; + line << "[ltx-progress] rank=" << rank << " t_ms=" << stamp << " event=" << event; + if (!detail.empty()) + line << " " << detail; + return line.str(); +} + +class LTX2ProgressLog { + public: + using Clock = std::chrono::steady_clock; + + explicit LTX2ProgressLog(int32_t rank = 0) + : enabled_(ltx2_progress_enabled(std::getenv(kLTX2ProgressEnv))), rank_(rank), + origin_(Clock::now()) {} + + bool enabled() const { return enabled_; } + + // Resets t=0 and emits event=start. + void start(const std::string& detail = "") { + origin_ = Clock::now(); + emit("start", detail); + } + + void emit(const std::string& event, const std::string& detail = "") const { + if (!enabled_) + return; + const double t_ms = + std::chrono::duration(Clock::now() - origin_).count(); + std::cout << format_ltx2_progress(rank_, t_ms, event, detail) << std::endl; + } + + private: + bool enabled_; + int32_t rank_; + Clock::time_point origin_; +}; + +} // namespace trtmc diff --git a/families/ltx2/runtime/runtime_config.cpp b/families/ltx2/runtime/runtime_config.cpp new file mode 100644 index 0000000000..bd92cd3cd9 --- /dev/null +++ b/families/ltx2/runtime/runtime_config.cpp @@ -0,0 +1,48 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "families/ltx2/runtime/runtime_config.h" + +#include +#include + +namespace trtmc::ltx2 { +namespace { + +bool supported_context_parallel_size(std::int32_t size) { + return size == 2 || size == 4 || size == 8; +} + +} // namespace + +ParallelRuntimeConfig parse_parallel_runtime_config(const std::string& json) { + const auto document = nlohmann::json::parse(json); + const bool has_mode = document.contains("parallel_mode"); + const bool has_size = document.contains("parallel_size"); + if (!has_mode && !has_size) + return {}; + if (!has_mode || !document.at("parallel_mode").is_string()) + throw std::runtime_error("LTX-2.5 runtime.json requires string parallel_mode"); + if (!has_size || !document.at("parallel_size").is_number_integer()) + throw std::runtime_error("LTX-2.5 runtime.json requires integer parallel_size"); + + ParallelRuntimeConfig config; + const auto mode = document.at("parallel_mode").get(); + if (mode == "single") + config.mode = ParallelMode::Single; + else if (mode == "context_parallel") + config.mode = ParallelMode::Context; + else + throw std::runtime_error("LTX-2.5 runtime.json has unsupported parallel_mode"); + config.size = document.at("parallel_size").get(); + + if ((config.mode == ParallelMode::Single && config.size != 1) || + (config.mode == ParallelMode::Context && !supported_context_parallel_size(config.size))) { + throw std::runtime_error("LTX-2.5 runtime.json has invalid parallel settings"); + } + return config; +} + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/runtime/runtime_config.h b/families/ltx2/runtime/runtime_config.h new file mode 100644 index 0000000000..4af3dbe8f7 --- /dev/null +++ b/families/ltx2/runtime/runtime_config.h @@ -0,0 +1,29 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include +#include + +namespace trtmc::ltx2 { + +enum class ParallelMode { + Single, + Context, +}; + +struct ParallelRuntimeConfig { + ParallelMode mode{ParallelMode::Single}; + std::int32_t size{1}; + + bool distributed() const { return mode != ParallelMode::Single; } +}; + +// Bundles built before context parallelism carry no parallel keys and run on +// one device. Context-parallel bundles share one rank-dynamic denoiser plan. +ParallelRuntimeConfig parse_parallel_runtime_config(const std::string& json); + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/runtime/runtime_math.h b/families/ltx2/runtime/runtime_math.h new file mode 100644 index 0000000000..87c8b8f729 --- /dev/null +++ b/families/ltx2/runtime/runtime_math.h @@ -0,0 +1,57 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +// Pure host-side helpers of the LTX-2.5 runtime (header-only so the contract tests compile +// them without engines). + +#include +#include +#include +#include +#include + +namespace trtmc { + +// Euler step of diffusers FlowMatchEulerDiscreteScheduler: x + (sigma_next - sigma) * v. +inline void ltx2_euler_step(std::vector& x, const std::vector& v, float sigma, + float sigma_next) { + if (x.size() != v.size()) + throw std::runtime_error("LTX-2.5 scheduler step: latent and velocity sizes differ"); + const float dt = sigma_next - sigma; + for (std::size_t i = 0; i < x.size(); ++i) + x[i] = x[i] + dt * v[i]; +} + +// Gemma prompt ids for one prompt: the tokenizer ids (no special tokens) right-truncated to +// seq_len and left-padded with pad_id, as LTX2Pipeline tokenizes with padding_side="left". +// mask is 1 on tokens and 0 on padding. +inline void ltx2_prompt_ids(const std::vector& ids, int32_t seq_len, int32_t pad_id, + std::vector& padded, std::vector& mask) { + const auto n = std::min(ids.size(), static_cast(seq_len)); + padded.assign(static_cast(seq_len), pad_id); + mask.assign(static_cast(seq_len), 0); + const auto offset = static_cast(seq_len) - n; + for (std::size_t i = 0; i < n; ++i) { + padded[offset + i] = ids[i]; + mask[offset + i] = 1; + } +} + +// Planar [channels][samples] -> interleaved [samples][channels]. +inline std::vector ltx2_interleave(const std::vector& planar, int32_t channels) { + if (channels <= 0 || planar.size() % static_cast(channels) != 0) + throw std::runtime_error("LTX-2.5 audio has an invalid channel layout"); + const auto samples = planar.size() / static_cast(channels); + std::vector out(planar.size()); + for (std::size_t s = 0; s < samples; ++s) + for (int32_t c = 0; c < channels; ++c) + out[s * static_cast(channels) + static_cast(c)] = + planar[static_cast(c) * samples + s]; + return out; +} + +} // namespace trtmc diff --git a/families/ltx2/runtime/tokenizer.h b/families/ltx2/runtime/tokenizer.h new file mode 100644 index 0000000000..3091ef8e47 --- /dev/null +++ b/families/ltx2/runtime/tokenizer.h @@ -0,0 +1,30 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include +#include +#include +#include +#include +#include + +namespace trtmc { + +class ITokenizer { + public: + virtual ~ITokenizer() = default; + virtual std::vector encode(const std::string& text) const = 0; + virtual std::string decode(const std::vector& ids) const = 0; + virtual std::int32_t id_for_token(std::string_view token) const = 0; + virtual std::string token_for_id(std::int32_t id) const = 0; +}; + +std::unique_ptr CreateLtx2BpeTokenizer(const char* tokenizer_json_data, + std::size_t tokenizer_json_size, + bool add_special_tokens); + +} // namespace trtmc diff --git a/families/ltx2/support.py b/families/ltx2/support.py new file mode 100644 index 0000000000..73a4fd7d2f --- /dev/null +++ b/families/ltx2/support.py @@ -0,0 +1,15 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned model and task support for LTX-2.5 (``LTX2Pipeline``).""" + +from tensorrt_model_connect.model_support import family_support + + +describe = family_support( + model_types=("ltx2", "ltx-2", "ltx-2.5"), + pipeline_classes=("LTX2Pipeline",), + tasks=("text_to_audio_video",), + default_task="text_to_audio_video", + default_precision="bf16", +) diff --git a/families/ltx2/tests/__init__.py b/families/ltx2/tests/__init__.py new file mode 100644 index 0000000000..52a7a9daf0 --- /dev/null +++ b/families/ltx2/tests/__init__.py @@ -0,0 +1,2 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 diff --git a/families/ltx2/tests/conftest.py b/families/ltx2/tests/conftest.py new file mode 100644 index 0000000000..6e4b95cf9d --- /dev/null +++ b/families/ltx2/tests/conftest.py @@ -0,0 +1,30 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 test backend selection. + +The builders import ``tensorrt``. Hosts that only have TensorRT-RTX installed (the +Windows RTX demo host) bind it through the same alias the build CLI uses +(``select_backend("trt_rtx")``). ``TRTMC_LTX2_TEST_BACKEND=trt|trt_rtx`` forces a choice. +""" + +from __future__ import annotations + +import importlib.util +import os + + +def _bind_backend() -> None: + choice = os.environ.get("TRTMC_LTX2_TEST_BACKEND", "").strip() + if not choice: + if importlib.util.find_spec("tensorrt") is not None: + return + if importlib.util.find_spec("tensorrt_rtx") is None: + return + choice = "trt_rtx" + from tensorrt_model_connect.build import select_backend + + select_backend(choice) + + +_bind_backend() diff --git a/families/ltx2/tests/cp_tiny_prep.py b/families/ltx2/tests/cp_tiny_prep.py new file mode 100644 index 0000000000..2e085d0b08 --- /dev/null +++ b/families/ltx2/tests/cp_tiny_prep.py @@ -0,0 +1,63 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Single-rank preparation for the multi-rank CP check (uses torch + diffusers). + +``python -m families.ltx2.tests.cp_tiny_prep OUT_DIR CP``: builds the tiny random DiT, its +single-device and CP plans, and the diffusers reference outputs for the normal, mixed-STG and +modality-isolated batches. The multi-rank step (``dist_dit_cp_check``) is torch-free. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +from families.ltx2.tests import conftest # noqa: F401 - binds the TensorRT backend + +CASES = (("plain", [1.0, 1.0], [1.0, 1.0]), ("stg_mixed", [1.0, 0.0], [1.0, 1.0]), + ("isolated", [1.0, 1.0], [0.0, 0.0])) + + +def main() -> int: + import numpy as np + import safetensors.torch as st + import torch + from diffusers import LTX2VideoTransformer3DModel + + from families.ltx2.dit_builder import DiTShape, audio_latent_frames, build_dit_engine + from families.ltx2.tests import test_dit_parity as tp + + out = Path(sys.argv[1]) + cp = int(sys.argv[2]) + out.mkdir(parents=True, exist_ok=True) + model = LTX2VideoTransformer3DModel(**tp.TINY_DIT).eval() + tp._randomize(model, 5) + folder = out / "transformer" + folder.mkdir(exist_ok=True) + st.save_file({k: v.to(torch.bfloat16).contiguous() for k, v in model.state_dict().items()}, + str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(tp.TINY_DIT), encoding="utf-8") + sa = audio_latent_frames((tp.FRAMES - 1) * 8 + 1, tp.FPS) + shape = DiTShape(batch=2, latent_frames=tp.FRAMES, latent_height=tp.LH, latent_width=tp.LW, audio_frames=sa, + text_len=tp.TEXT, fps=tp.FPS) + (out / "single.plan").write_bytes(build_dit_engine(folder, shape, cp_size=1, stg_blocks=(tp.STG_BLOCK,))) + (out / f"cp{cp}.plan").write_bytes(build_dit_engine(folder, shape, cp_size=cp, stg_blocks=(tp.STG_BLOCK,))) + inp = tp._inputs(shape) + arrays = {k: v.float().numpy() for k, v in inp.items()} + for case, stg, av in CASES: + rv, ra = tp._reference(model, shape, inp, torch.float32, + stg_mask=torch.tensor(stg) if case == "stg_mixed" else None, + isolate=(case == "isolated")) + arrays[f"{case}_ref_video"] = rv.numpy() + arrays[f"{case}_ref_audio"] = ra.numpy() + arrays[f"{case}_stg_keep"] = np.asarray(stg, np.float32) + arrays[f"{case}_av_keep"] = np.asarray(av, np.float32) + np.savez(out / "reference.npz", **arrays) + print(f"prepared {out} (video tokens {shape.video_tokens}, audio {sa}, cp {cp})", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/families/ltx2/tests/cpp/test_runtime_contract.cpp b/families/ltx2/tests/cpp/test_runtime_contract.cpp new file mode 100644 index 0000000000..9cb47556b3 --- /dev/null +++ b/families/ltx2/tests/cpp/test_runtime_contract.cpp @@ -0,0 +1,80 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +// LTX-2.5 host-side runtime contract: scheduler step, prompt padding, audio interleave. + +#include "families/ltx2/runtime/progress_log.h" +#include "families/ltx2/runtime/runtime_math.h" + +#include +#include +#include +#include +#include + +namespace { + +int failures = 0; + +void check(bool ok, const char* what) { + if (!ok) { + std::fprintf(stderr, "FAIL: %s\n", what); + ++failures; + } +} + +void test_euler_step_matches_flow_match_euler() { + // diffusers: prev = x + (sigma_next - sigma) * v; the last distilled step lands on x0. + std::vector x{1.0F, -2.0F, 0.5F}; + const std::vector v{0.5F, 1.0F, -4.0F}; + trtmc::ltx2_euler_step(x, v, 0.421875F, 0.0F); + check(std::fabs(x[0] - (1.0F - 0.421875F * 0.5F)) < 1e-7F, "euler step value 0"); + check(std::fabs(x[1] - (-2.0F - 0.421875F)) < 1e-7F, "euler step value 1"); + check(std::fabs(x[2] - (0.5F + 0.421875F * 4.0F)) < 1e-7F, "euler step value 2"); + bool threw = false; + try { + std::vector short_v{1.0F}; + trtmc::ltx2_euler_step(x, short_v, 1.0F, 0.5F); + } catch (const std::runtime_error&) { + threw = true; + } + check(threw, "euler step rejects mismatched sizes"); +} + +void test_prompt_ids_left_pad_and_truncate() { + std::vector ids; + std::vector mask; + trtmc::ltx2_prompt_ids({11, 12, 13}, 6, 0, ids, mask); + check((ids == std::vector{0, 0, 0, 11, 12, 13}), "left padding"); + check((mask == std::vector{0, 0, 0, 1, 1, 1}), "mask marks tokens"); + trtmc::ltx2_prompt_ids({1, 2, 3, 4, 5}, 3, 0, ids, mask); + check((ids == std::vector{1, 2, 3}), "right truncation keeps the first tokens"); + check((mask == std::vector{1, 1, 1}), "truncated mask is full"); +} + +void test_interleave_stereo() { + const auto out = trtmc::ltx2_interleave({1.0F, 2.0F, 3.0F, -1.0F, -2.0F, -3.0F}, 2); + check((out == std::vector{1.0F, -1.0F, 2.0F, -2.0F, 3.0F, -3.0F}), + "planar to interleaved"); +} + +void test_progress_line_format() { + const auto line = trtmc::format_ltx2_progress(1, 12.5, "step", "step=2/8 step_ms=3.000"); + check(line == "[ltx-progress] rank=1 t_ms=12.500 event=step step=2/8 step_ms=3.000", + "progress line keeps the LTX format"); +} + +} // namespace + +int main() { + test_euler_step_matches_flow_match_euler(); + test_prompt_ids_left_pad_and_truncate(); + test_interleave_stereo(); + test_progress_line_format(); + if (failures != 0) + return EXIT_FAILURE; + std::puts("ltx2 runtime contract: OK"); + return EXIT_SUCCESS; +} diff --git a/families/ltx2/tests/dist_dit_cp_check.py b/families/ltx2/tests/dist_dit_cp_check.py new file mode 100644 index 0000000000..ac9ab74160 --- /dev/null +++ b/families/ltx2/tests/dist_dit_cp_check.py @@ -0,0 +1,74 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Multi-rank check of the LTX-2.5 context-parallel DiT (tiny random weights), torch-free. + +``tools/launch_ranks.py -n CP --gpus ... -- python -m families.ltx2.tests.dist_dit_cp_check PREP_DIR`` +after ``cp_tiny_prep`` wrote the plans and references into ``PREP_DIR``. Every rank runs the CP +plan with its NCCL communicator (a hung collective aborts the communicator instead of blocking), +runs the single-device plan, and compares the full outputs with the single-device plan and with +diffusers, shard by shard. Writes ``cp_rank.json``; exits non-zero on failure. +""" + +from __future__ import annotations + +import json +import os +import sys +from pathlib import Path + +from families.ltx2.tests import conftest # noqa: F401 - binds the TensorRT backend + + +def main() -> int: + import numpy as np + from cuda.bindings import runtime as rt + + from families.ltx2.tests.dist_helpers import NcclComm + from families.ltx2.tests.np_engine import NpEngine, ck, cosine + + ck(rt.cudaSetDevice(int(os.environ.get("OMPI_COMM_WORLD_LOCAL_RANK", "0")))) + ck(rt.cudaFree(0)) + prep = Path(sys.argv[1]) + comm = NcclComm() + ref = np.load(prep / "reference.npz") + cp_engine = NpEngine((prep / f"cp{comm.world}.plan").read_bytes(), comm.capsule(), on_timeout=comm.abort) + single = NpEngine((prep / "single.plan").read_bytes()) + report = {"rank": comm.rank, "world": comm.world, "cases": {}} + ok = True + base = {k: ref[k] for k in ("video_latent", "audio_latent", "video_context", "audio_context", "timestep")} + for case in ("plain", "stg_mixed", "isolated"): + feed = dict(base, stg_keep=ref[f"{case}_stg_keep"], av_keep=ref[f"{case}_av_keep"]) + try: + got = cp_engine(feed, timeout_s=120) + except TimeoutError as exc: + print(f"[rank {comm.rank}] {case}: {exc}; communicator aborted", flush=True) + report["cases"][case] = {"ok": False, "error": str(exc)} + ok = False + break + one = single(feed) + gv, ga = got["video_velocity"], got["audio_velocity"] + s = gv.shape[1] // comm.world + rec = { + "video_cos_per_shard_vs_single": [cosine(gv[:, r * s:(r + 1) * s], one["video_velocity"][:, r * s:(r + 1) * s]) + for r in range(comm.world)], + "audio_cos_vs_single": cosine(ga, one["audio_velocity"]), + "video_cos_vs_diffusers": cosine(gv, ref[f"{case}_ref_video"]), + "audio_cos_vs_diffusers": cosine(ga, ref[f"{case}_ref_audio"]), + "max_abs_vs_single": float(np.abs(gv - one["video_velocity"]).max()), + "finite": bool(np.isfinite(gv).all() and np.isfinite(ga).all()), + } + rec["ok"] = bool(rec["finite"] and min(rec["video_cos_per_shard_vs_single"]) > 0.9999 + and rec["audio_cos_vs_single"] > 0.9999 and rec["video_cos_vs_diffusers"] > 0.999 + and rec["audio_cos_vs_diffusers"] > 0.999) + ok &= rec["ok"] + report["cases"][case] = rec + print(f"[rank {comm.rank}] {case}: {json.dumps(rec)}", flush=True) + (prep / f"cp_rank{comm.rank}.json").write_text(json.dumps(report, indent=1), encoding="utf-8") + if comm.comm: + comm.destroy() + return 0 if ok else 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/families/ltx2/tests/dist_helpers.py b/families/ltx2/tests/dist_helpers.py new file mode 100644 index 0000000000..8b5680c8d0 --- /dev/null +++ b/families/ltx2/tests/dist_helpers.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""NCCL communicator for multi-rank engine tests (launch with ``tools/launch_ranks.py``). + +Select the rank's CUDA device before constructing it (``cudaSetDevice(local_rank)``). + +Uses the launcher's contract: ``OMPI_COMM_WORLD_{SIZE,RANK,LOCAL_RANK}``, the unique-id +rendezvous file ``TRTMC_NCCL_RENDEZVOUS`` and the library ``TRTMC_NCCL_LIBRARY``. +""" + +from __future__ import annotations + +import ctypes +import os +import sys +import time +from pathlib import Path + + +class _UniqueId(ctypes.Structure): + _fields_ = [("internal", ctypes.c_char * 128)] + + +class NcclComm: + def __init__(self): + self.rank = int(os.environ["OMPI_COMM_WORLD_RANK"]) + self.world = int(os.environ["OMPI_COMM_WORLD_SIZE"]) + self.local_rank = int(os.environ.get("OMPI_COMM_WORLD_LOCAL_RANK", self.rank)) + default = "nccl.dll" if sys.platform.startswith("win") else "libnccl.so.2" + self.lib = ctypes.CDLL(os.environ.get("TRTMC_NCCL_LIBRARY") or default) + self.lib.ncclCommInitRank.argtypes = [ctypes.POINTER(ctypes.c_void_p), ctypes.c_int, _UniqueId, ctypes.c_int] + self.lib.ncclCommAbort.argtypes = [ctypes.c_void_p] + self.lib.ncclCommDestroy.argtypes = [ctypes.c_void_p] + uid = _UniqueId() + path = Path(os.environ["TRTMC_NCCL_RENDEZVOUS"]) + if self.rank == 0: + self._check(self.lib.ncclGetUniqueId(ctypes.byref(uid)), "ncclGetUniqueId") + tmp = path.with_name(path.name + ".tmp") + tmp.write_bytes(ctypes.string_at(ctypes.addressof(uid), 128)) + os.replace(tmp, path) + else: + deadline = time.monotonic() + 120 + while not path.exists(): + if time.monotonic() > deadline: + raise TimeoutError(f"no NCCL unique id at {path}") + time.sleep(0.05) + ctypes.memmove(ctypes.addressof(uid), path.read_bytes(), 128) + self.comm = ctypes.c_void_p() + self._check(self.lib.ncclCommInitRank(ctypes.byref(self.comm), self.world, uid, self.rank), + "ncclCommInitRank") + + @staticmethod + def _check(status: int, what: str) -> None: + if status != 0: + raise RuntimeError(f"{what} failed with NCCL status {status}") + + def capsule(self): + new = ctypes.pythonapi.PyCapsule_New + new.restype = ctypes.py_object + new.argtypes = [ctypes.c_void_p, ctypes.c_char_p, ctypes.c_void_p] + return new(self.comm.value, None, None) + + def abort(self) -> None: + """``ncclCommAbort``: makes in-flight NCCL kernels exit (use instead of killing a hung rank).""" + if self.comm: + self.lib.ncclCommAbort(self.comm) + self.comm = ctypes.c_void_p() + + def destroy(self) -> None: + if self.comm: + self.lib.ncclCommDestroy(self.comm) + self.comm = ctypes.c_void_p() diff --git a/families/ltx2/tests/engine_runner.py b/families/ltx2/tests/engine_runner.py new file mode 100644 index 0000000000..b2c6319440 --- /dev/null +++ b/families/ltx2/tests/engine_runner.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Run a serialized TensorRT plan on torch CUDA tensors (test helper).""" + +from __future__ import annotations + +import tensorrt as trt +import torch + +from families.ltx2.graph import make_logger + +_TORCH = { + trt.float32: torch.float32, + trt.float16: torch.float16, + trt.bfloat16: torch.bfloat16, + trt.int32: torch.int32, + trt.bool: torch.bool, +} + + +class Engine: + """Deserialized plan + execution context, reusable across calls.""" + + def __init__(self, plan: bytes, communicator=None): + self.runtime = trt.Runtime(make_logger()) + self.engine = self.runtime.deserialize_cuda_engine(plan) + if self.engine is None: + raise RuntimeError("failed to deserialize plan") + self.context = self.engine.create_execution_context() + if communicator is not None and not self.context.set_communicator(communicator): + raise RuntimeError("set_communicator failed") + self.stream = torch.cuda.Stream() + + def __call__(self, inputs: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + return _execute(self.engine, self.context, inputs, self.stream) + + +def run_plan(plan: bytes, inputs: dict[str, torch.Tensor], communicator=None) -> dict[str, torch.Tensor]: + return Engine(plan, communicator)(inputs) + + +def _execute(engine, context, inputs, stream) -> dict[str, torch.Tensor]: + keep = [] + outputs: dict[str, torch.Tensor] = {} + for i in range(engine.num_io_tensors): + name = engine.get_tensor_name(i) + dtype = _TORCH[engine.get_tensor_dtype(name)] + shape = tuple(engine.get_tensor_shape(name)) + if engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: + t = inputs[name].to(device="cuda", dtype=dtype).contiguous() + if tuple(t.shape) != shape: + raise ValueError(f"{name}: shape {tuple(t.shape)} != engine {shape}") + keep.append(t) + else: + t = torch.empty(shape, dtype=dtype, device="cuda") + outputs[name] = t + context.set_tensor_address(name, t.data_ptr()) + torch.cuda.current_stream().synchronize() + if not context.execute_async_v3(stream.cuda_stream): + raise RuntimeError("execute_async_v3 failed") + stream.synchronize() + return outputs + + +def cosine(a: torch.Tensor, b: torch.Tensor) -> float: + a = a.detach().double().flatten() + b = b.detach().double().flatten() + return float((a @ b) / (a.norm() * b.norm() + 1e-30)) + + +def rel_l2(a: torch.Tensor, b: torch.Tensor) -> float: + a = a.detach().double() + b = b.detach().double() + return float((a - b).norm() / (b.norm() + 1e-30)) diff --git a/families/ltx2/tests/manifests/ltx25-distilled-cp2.json b/families/ltx2/tests/manifests/ltx25-distilled-cp2.json new file mode 100644 index 0000000000..5e12e482f7 --- /dev/null +++ b/families/ltx2/tests/manifests/ltx25-distilled-cp2.json @@ -0,0 +1,21 @@ +{ + "name": "ltx25-distilled-cp2", + "hf_id": "Lightricks/LTX-2.5-Diffusers", + "hf_revision": "426936f8b22dc28e4def61e515478b0b7e4a53cc", + "bundle": "ltx25-distilled-cp2.bundle", + "family": "ltx2", + "task": "text_to_audio_video", + "precision": "bf16", + "testcases": [ + { + "name": "ltx25-distilled-cp2", + "test_prompt": "A cinematic shot of a red fox walking through a snowy forest at dawn, the camera tracking alongside, snow crunching underfoot.", + "seed": 42, + "runtime_timeout_s": 3600 + } + ], + "context_parallel_size": 2, + "image_height": 544, + "image_width": 960, + "video_num_frames": 121 +} diff --git a/families/ltx2/tests/manifests/ltx25-distilled-l0.json b/families/ltx2/tests/manifests/ltx25-distilled-l0.json new file mode 100644 index 0000000000..d97178544f --- /dev/null +++ b/families/ltx2/tests/manifests/ltx25-distilled-l0.json @@ -0,0 +1,22 @@ +{ + "name": "ltx25-distilled-l0", + "hf_id": "Lightricks/LTX-2.5-Diffusers", + "hf_revision": "426936f8b22dc28e4def61e515478b0b7e4a53cc", + "bundle": "ltx25-distilled-l0.bundle", + "family": "ltx2", + "task": "text_to_audio_video", + "precision": "bf16", + "testcases": [ + { + "name": "ltx25-distilled-l0", + "premerge": true, + "test_prompt": "A cinematic shot of a red fox walking through a snowy forest at dawn, the camera tracking alongside, snow crunching underfoot.", + "seed": 42, + "runtime_timeout_s": 3600 + } + ], + "context_parallel_size": 1, + "image_height": 384, + "image_width": 640, + "video_num_frames": 49 +} diff --git a/families/ltx2/tests/np_engine.py b/families/ltx2/tests/np_engine.py new file mode 100644 index 0000000000..b384240cca --- /dev/null +++ b/families/ltx2/tests/np_engine.py @@ -0,0 +1,82 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Torch-free TensorRT plan runner (cuda-python + NumPy) for multi-rank tests. + +Multi-rank tests keep each rank light: no torch, only the plan and the NCCL communicator. +Waits poll ``cudaStreamQuery`` with a deadline and abort the NCCL communicator instead of +blocking forever, so a hung collective never leaves kernels running on the GPUs (killing +a rank with NCCL kernels in flight can leave the devices busy until a reset). +""" + +from __future__ import annotations + +import time + +import ml_dtypes +import numpy as np +import tensorrt as trt +from cuda.bindings import runtime as rt + +from families.ltx2.graph import make_logger + +_NP = {trt.float32: np.float32, trt.float16: np.float16, trt.bfloat16: ml_dtypes.bfloat16, trt.int32: np.int32, + trt.bool: np.bool_} + + +def ck(result): + if int(result[0]) != 0: + raise RuntimeError(f"CUDA error {result[0]}") + return result[1] if len(result) == 2 else result[1:] + + +class NpEngine: + def __init__(self, plan: bytes, communicator=None, on_timeout=None): + self.runtime = trt.Runtime(make_logger()) + self.engine = self.runtime.deserialize_cuda_engine(plan) + if self.engine is None: + raise RuntimeError("failed to deserialize plan") + self.context = self.engine.create_execution_context() + if communicator is not None and not self.context.set_communicator(communicator): + raise RuntimeError("set_communicator failed") + self.stream = int(ck(rt.cudaStreamCreate())) + self.on_timeout = on_timeout + self.buffers: dict[str, tuple[int, tuple, object]] = {} + for i in range(self.engine.num_io_tensors): + name = self.engine.get_tensor_name(i) + shape = tuple(self.engine.get_tensor_shape(name)) + dtype = _NP[self.engine.get_tensor_dtype(name)] + nbytes = int(np.prod(shape)) * np.dtype(dtype).itemsize + ptr = int(ck(rt.cudaMalloc(max(nbytes, 1)))) + self.buffers[name] = (ptr, shape, dtype) + self.context.set_tensor_address(name, ptr) + + def __call__(self, inputs: dict[str, np.ndarray], *, timeout_s: float = 120.0) -> dict[str, np.ndarray]: + outs = {} + for name, (ptr, shape, dtype) in self.buffers.items(): + if self.engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: + a = np.ascontiguousarray(np.asarray(inputs[name]).astype(dtype)) + if a.shape != shape: + raise ValueError(f"{name}: shape {a.shape} != engine {shape}") + ck(rt.cudaMemcpy(ptr, a.ctypes.data, a.nbytes, rt.cudaMemcpyKind.cudaMemcpyHostToDevice)) + if not self.context.execute_async_v3(self.stream): + raise RuntimeError("execute_async_v3 failed") + deadline = time.monotonic() + timeout_s + while int(rt.cudaStreamQuery(self.stream)[0]) != 0: + if time.monotonic() > deadline: + if self.on_timeout is not None: + self.on_timeout() + raise TimeoutError(f"engine did not finish within {timeout_s} s") + time.sleep(0.002) + for name, (ptr, shape, dtype) in self.buffers.items(): + if self.engine.get_tensor_mode(name) == trt.TensorIOMode.OUTPUT: + a = np.empty(shape, dtype=dtype) + ck(rt.cudaMemcpy(a.ctypes.data, ptr, a.nbytes, rt.cudaMemcpyKind.cudaMemcpyDeviceToHost)) + outs[name] = a.astype(np.float32) if dtype is not np.int32 else a + return outs + + +def cosine(a, b) -> float: + a = np.asarray(a, dtype=np.float64).ravel() + b = np.asarray(b, dtype=np.float64).ravel() + return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-30)) diff --git a/families/ltx2/tests/test_audio_parity.py b/families/ltx2/tests/test_audio_parity.py new file mode 100644 index 0000000000..6d9cfa7972 --- /dev/null +++ b/families/ltx2/tests/test_audio_parity.py @@ -0,0 +1,121 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tiny-random parity: LTX-2.5 audio decoder engine (audio VAE decoder + vocoder with BWE) vs diffusers. + +Real structure (causal pixel-norm audio decoder; BigVGAN-style anti-aliased SnakeBeta vocoder with +the real 160x stage-1 and 240x BWE upsampling, causal STFT + mel, Hann x3 resampler), narrow channels. +""" + +from __future__ import annotations + +import json + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("tensorrt") +if not torch.cuda.is_available(): + pytest.skip("CUDA is required for LTX-2.5 engine parity tests", allow_module_level=True) +safetensors_torch = pytest.importorskip("safetensors.torch") + +from families.ltx2.tests.engine_runner import cosine, rel_l2, run_plan # noqa: E402 + +TINY_AUDIO_VAE = { + "attn_resolutions": None, "base_channels": 32, "causality_axis": "height", "ch_mult": [1, 2, 4], + "double_z": True, "dropout": 0.0, "in_channels": 2, "is_causal": True, "latent_channels": 4, "mel_bins": 32, + "mel_hop_length": 160, "mid_block_add_attention": False, "norm_type": "pixel", "num_res_blocks": 1, + "output_channels": 2, "resolution": 256, "sample_rate": 16000, +} +TINY_VOCODER = { + "act_fn": "snakebeta", "antialias": True, "antialias_kernel_size": 12, "antialias_ratio": 2, + "bwe_act_fn": "snakebeta", "bwe_antialias": True, "bwe_antialias_kernel_size": 12, "bwe_antialias_ratio": 2, + "bwe_final_act_fn": None, "bwe_final_bias": False, "bwe_hidden_channels": 32, "bwe_in_channels": 64, + "bwe_leaky_relu_negative_slope": 0.1, "bwe_out_channels": 2, + "bwe_resnet_dilations": [[1, 3, 5], [1, 3, 5], [1, 3, 5]], "bwe_resnet_kernel_sizes": [3, 7, 11], + "bwe_upsample_factors": [6, 5, 2, 2, 2], "bwe_upsample_kernel_sizes": [12, 11, 4, 4, 4], + "filter_length": 512, "final_act_fn": None, "final_bias": False, "hidden_channels": 64, "hop_length": 80, + "in_channels": 64, "input_sampling_rate": 16000, "leaky_relu_negative_slope": 0.1, "num_mel_channels": 32, + "out_channels": 2, "output_sampling_rate": 48000, "resnet_dilations": [[1, 3, 5], [1, 3, 5], [1, 3, 5]], + "resnet_kernel_sizes": [3, 7, 11], "upsample_factors": [5, 2, 2, 2, 2, 2], + "upsample_kernel_sizes": [11, 4, 4, 4, 4, 4], "window_length": 512, +} +SA = 6 + + +def _save(module, folder, cfg): + folder.mkdir() + safetensors_torch.save_file({k: v.contiguous() for k, v in module.state_dict().items()}, + str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(cfg), encoding="utf-8") + + +@pytest.fixture(scope="module") +def tiny(tmp_path_factory): + return make_tiny(tmp_path_factory.mktemp("tiny_audio")) + + +def make_tiny(root): + """Random tiny audio VAE + vocoder saved as diffusers folders under ``root``.""" + from diffusers import AutoencoderKLLTX2Audio + from diffusers.pipelines.ltx2.vocoder import LTX2VocoderWithBWE + + gen = torch.Generator().manual_seed(4) + vae = AutoencoderKLLTX2Audio(**TINY_AUDIO_VAE).eval() + voc = LTX2VocoderWithBWE(**TINY_VOCODER).eval() + with torch.no_grad(): + for name, p in list(vae.named_parameters()) + list(voc.named_parameters()): + if name.endswith(("alpha", "beta")): + p.copy_(0.2 * torch.randn(p.shape, generator=gen)) + elif p.ndim >= 2: + p.copy_(torch.randn(p.shape, generator=gen) * (1.0 / p[0].numel()) ** 0.5) + else: + p.copy_(0.05 * torch.randn(p.shape, generator=gen)) + vae.latents_mean.copy_(0.2 * torch.randn(vae.latents_mean.shape, generator=gen)) + vae.latents_std.copy_(0.5 + torch.rand(vae.latents_std.shape, generator=gen)) + # A real windowed DFT basis and a positive mel filterbank keep the log-mel well conditioned. + n_fft = 512 + nf = n_fft // 2 + 1 + t = torch.arange(n_fft, dtype=torch.float64) + win = torch.hann_window(n_fft, periodic=True, dtype=torch.float64) + f = torch.arange(nf, dtype=torch.float64)[:, None] + cos = torch.cos(2 * torch.pi * f * t / n_fft) * win + sin = -torch.sin(2 * torch.pi * f * t / n_fft) * win + voc.mel_stft.stft_fn.forward_basis.copy_(torch.cat([cos, sin], 0).unsqueeze(1).float()) + voc.mel_stft.mel_basis.copy_(torch.rand(voc.mel_stft.mel_basis.shape, generator=gen) / 20) + _save(vae, root / "audio_vae", TINY_AUDIO_VAE) + _save(voc, root / "vocoder", TINY_VOCODER) + return vae, voc, root + + +def test_audio_decoder_tiny_parity(tiny) -> None: + from families.ltx2.audio_builder import build_audio_decoder_engine + + vae, voc, root = tiny + lat_m = TINY_AUDIO_VAE["mel_bins"] // 4 + packed = torch.randn(1, SA, TINY_AUDIO_VAE["latent_channels"] * lat_m, generator=torch.Generator().manual_seed(1)) + plan = build_audio_decoder_engine(root, audio_frames=SA, debug_mel=True) + got = run_plan(plan, {"audio_latents": packed}) + # fp64 CPU reference: the engine computes in fp32, and TensorRT(-RTX) may run fp32 + # convolutions at TF32-class internal precision, which a random-weight vocoder (~120 + # stacked convolutions without normalization) amplifies into a percent-level waveform + # difference. The mel path (audio VAE) has no such stack and must match tightly; the + # waveform must stay within that convolution-precision envelope. + vae = vae.to("cpu", torch.float64) + voc = voc.to("cpu", torch.float64) + with torch.no_grad(): + z = packed.double() * vae.latents_std + vae.latents_mean + z = z.unflatten(2, (-1, lat_m)).transpose(1, 2) + mel = vae.decode(z, return_dict=False)[0] + wave = voc(mel) + got_mel = got["mel"].double().cpu() + got_wave = got["waveform"].double().cpu() + cm = cosine(got_mel, mel) + cw = cosine(got_wave, wave) + print(f"audio tiny: mel {tuple(got_mel.shape)} cos {cm:.6f} relL2 {rel_l2(got_mel, mel):.2e} | " + f"wave {tuple(got_wave.shape)} cos {cw:.6f} relL2 {rel_l2(got_wave, wave):.2e}") + assert tuple(got_mel.shape) == tuple(mel.shape) + assert tuple(got_wave.shape) == tuple(wave.shape) + assert torch.isfinite(got_wave).all() + assert cm > 0.9999 + assert cw > 0.995 diff --git a/families/ltx2/tests/test_context_parallel.py b/families/ltx2/tests/test_context_parallel.py new file mode 100644 index 0000000000..f452467de1 --- /dev/null +++ b/families/ltx2/tests/test_context_parallel.py @@ -0,0 +1,69 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Context-parallel DiT: layout checks and a 2-rank tiny-random parity run.""" + +from __future__ import annotations + +import json +import os +import subprocess +import sys +from pathlib import Path + +import pytest + +from families.ltx2.parallel import ParallelConfig, rank_selector_values, validate_context_parallel_layout + +REPO = Path(__file__).resolve().parents[3] + + +def test_rank_selector_sums_to_rank_index() -> None: + for cp in (2, 4, 8): + values = rank_selector_values(cp) + assert [float(v) * cp for v in values[:, 0]] == list(range(cp)) + + +@pytest.mark.parametrize("tokens,ok", [(8160, True), (8161, False)]) +def test_layout_validation(tokens: int, ok: bool) -> None: + parallel = ParallelConfig(cp_size=2) + if ok: + validate_context_parallel_layout(parallel, video_tokens=tokens, video_heads=32, audio_heads=32) + else: + with pytest.raises(ValueError): + validate_context_parallel_layout(parallel, video_tokens=tokens, video_heads=32, audio_heads=32) + + +def test_cp2_tiny_parity_two_ranks(tmp_path: Path) -> None: + """2-rank CP DiT vs single-device plan and diffusers (tiny random weights). + + The multi-rank step is torch-free (cuda-python + NumPy) so each rank stays light and + every engine wait can abort its NCCL communicator on a timeout. The torch/diffusers + reference runs in a separate single-rank preparation process. + """ + torch = pytest.importorskip("torch") + pytest.importorskip("tensorrt") + pytest.importorskip("diffusers") + pytest.importorskip("cuda.bindings") + if not torch.cuda.is_available() or torch.cuda.device_count() < 2: + pytest.skip("two CUDA devices are required") + nccl = os.environ.get("TRTMC_NCCL_LIBRARY") + if not nccl: + pytest.skip("TRTMC_NCCL_LIBRARY must point at the NCCL library") + env = dict(os.environ) + env.pop("CUDA_VISIBLE_DEVICES", None) + env["PYTHONPATH"] = os.pathsep.join([str(REPO), str(REPO / "core" / "builder"), env.get("PYTHONPATH", "")]) + prep = subprocess.run([sys.executable, "-m", "families.ltx2.tests.cp_tiny_prep", str(tmp_path), "2"], + cwd=REPO, env=env, capture_output=True, text=True, timeout=900) + print(prep.stdout[-2000:], prep.stderr[-2000:]) + assert prep.returncode == 0 + cmd = [sys.executable, str(REPO / "tools" / "launch_ranks.py"), "-n", "2", "--gpus", "0,1", + "--nccl-library", nccl, "--timeout", "600", "--", + sys.executable, "-m", "families.ltx2.tests.dist_dit_cp_check", str(tmp_path)] + proc = subprocess.run(cmd, cwd=REPO, env=env, capture_output=True, text=True, timeout=900) + print(proc.stdout[-6000:]) + print(proc.stderr[-3000:]) + assert proc.returncode == 0 + for rank in range(2): + report = json.loads((tmp_path / f"cp_rank{rank}.json").read_text(encoding="utf-8")) + assert all(case["ok"] for case in report["cases"].values()) diff --git a/families/ltx2/tests/test_dit_parity.py b/families/ltx2/tests/test_dit_parity.py new file mode 100644 index 0000000000..42b984d3fa --- /dev/null +++ b/families/ltx2/tests/test_dit_parity.py @@ -0,0 +1,160 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tiny-random parity: LTX-2.5 joint audio/video DiT engine vs diffusers ``LTX2VideoTransformer3DModel``. + +The shrunk config keeps every LTX-2.5 feature (9-parameter AdaLN, prompt AdaLN, gated attention, +split RoPE with fps-scaled video coordinates, time-only cross-modal RoPE, a2v / v2a, STG on one +block, bias-free video FFN, cross timestep) and a non-square latent grid. +""" + +from __future__ import annotations + +import json + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("tensorrt") +if not torch.cuda.is_available(): + pytest.skip("CUDA is required for LTX-2.5 engine parity tests", allow_module_level=True) +safetensors_torch = pytest.importorskip("safetensors.torch") + +from families.ltx2.tests.engine_runner import cosine, rel_l2, run_plan # noqa: E402 + +TINY_DIT = { + "activation_fn": "gelu-approximate", "attention_bias": True, "attention_head_dim": 16, + "attention_out_bias": True, "audio_attention_head_dim": 8, "audio_cross_attention_dim": 32, + "audio_cross_attn_mod": True, "audio_ff_bias": True, "audio_gated_attn": True, "audio_hop_length": 160, + "audio_in_channels": 16, "audio_num_attention_heads": 4, "audio_out_channels": 16, "audio_patch_size": 1, + "audio_patch_size_t": 1, "audio_pos_embed_max_pos": 20, "audio_sampling_rate": 16000, + "audio_scale_factor": 4, "base_height": 2048, "base_width": 2048, "caption_channels": 48, + "causal_offset": 1, "cross_attention_dim": 64, "cross_attn_mod": True, + "cross_attn_timestep_scale_multiplier": 1000, "ff_bias": False, "gated_attn": True, "in_channels": 16, + "norm_elementwise_affine": False, "norm_eps": 1e-6, "num_attention_heads": 4, "num_layers": 3, + "out_channels": 16, "patch_size": 1, "patch_size_t": 1, "perturbed_attn": True, "pos_embed_max_pos": 20, + "qk_norm": "rms_norm_across_heads", "rope_double_precision": True, "rope_theta": 10000.0, + "rope_type": "split", "timestep_scale_multiplier": 1000, "use_keyframes_abs_pos_embedding": True, + "use_prompt_adaln_single": True, "use_prompt_embeddings": False, "vae_scale_factors": [8, 32, 32], +} +FRAMES, LH, LW = 3, 4, 6 # latent grid -> 17 pixel frames +FPS = 24.0 +TEXT = 16 +STG_BLOCK = 1 + + +def _randomize(module, seed: int) -> None: + gen = torch.Generator().manual_seed(seed) + with torch.no_grad(): + for name, p in module.named_parameters(): + if "norm" in name and name.endswith("weight"): + p.copy_(1.0 + 0.2 * torch.randn(p.shape, generator=gen)) + elif "scale_shift_table" in name: + p.copy_(0.3 * torch.randn(p.shape, generator=gen)) + elif p.ndim >= 2: + p.copy_(torch.randn(p.shape, generator=gen) / p.shape[-1] ** 0.5) + else: + p.copy_(0.1 * torch.randn(p.shape, generator=gen)) + + +@pytest.fixture(scope="module") +def tiny(tmp_path_factory): + from diffusers import LTX2VideoTransformer3DModel + + from families.ltx2.dit_builder import DiTShape, audio_latent_frames, build_dit_engine + + model = LTX2VideoTransformer3DModel(**TINY_DIT).eval() + _randomize(model, 5) + folder = tmp_path_factory.mktemp("tiny_dit") / "transformer" + folder.mkdir() + safetensors_torch.save_file({k: v.to(torch.bfloat16).contiguous() for k, v in model.state_dict().items()}, + str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(TINY_DIT), encoding="utf-8") + num_frames = (FRAMES - 1) * 8 + 1 + sa = audio_latent_frames(num_frames, FPS) + shape = DiTShape(batch=2, latent_frames=FRAMES, latent_height=LH, latent_width=LW, audio_frames=sa, + text_len=TEXT, fps=FPS) + plan = build_dit_engine(folder, shape, stg_blocks=(STG_BLOCK,)) + return model, plan, shape + + +def _inputs(shape, seed=0): + gen = torch.Generator().manual_seed(seed) + b, s = shape.batch, shape.video_tokens + return { + "video_latent": torch.randn(b, s, 16, generator=gen), + "audio_latent": torch.randn(b, shape.audio_frames, 16, generator=gen), + "video_context": torch.randn(b, TEXT, 64, generator=gen).to(torch.bfloat16), + "audio_context": torch.randn(b, TEXT, 32, generator=gen).to(torch.bfloat16), + "timestep": torch.full((b,), 909.375), + } + + +def _reference(model, shape, inp, dtype, *, stg_mask=None, isolate=False): + model = model.to("cuda", dtype) + b = shape.batch + with torch.no_grad(): + v, a = model( + hidden_states=inp["video_latent"].cuda().to(dtype), + audio_hidden_states=inp["audio_latent"].cuda().to(dtype), + encoder_hidden_states=inp["video_context"].cuda().to(dtype), + audio_encoder_hidden_states=inp["audio_context"].cuda().to(dtype), + timestep=inp["timestep"].cuda(), sigma=inp["timestep"].cuda(), + encoder_attention_mask=torch.ones(b, TEXT, device="cuda"), + audio_encoder_attention_mask=torch.ones(b, TEXT, device="cuda"), + num_frames=shape.latent_frames, height=shape.latent_height, width=shape.latent_width, fps=FPS, + audio_num_frames=shape.audio_frames, use_cross_timestep=True, isolate_modalities=isolate, + spatio_temporal_guidance_blocks=[STG_BLOCK] if stg_mask is not None else None, + perturbation_mask=None if stg_mask is None else stg_mask.cuda().to(dtype), return_dict=False) + return v.float().cpu(), a.float().cpu() + + +@pytest.mark.parametrize("case", ["plain", "stg_mixed", "isolated"]) +def test_dit_tiny_parity(tiny, case) -> None: + model, plan, shape = tiny + inp = _inputs(shape) + stg = torch.ones(shape.batch) + av = torch.ones(shape.batch) + stg_ref = None + if case == "stg_mixed": + stg = torch.tensor([1.0, 0.0]) + stg_ref = stg + if case == "isolated": + av = torch.zeros(shape.batch) + got = run_plan(plan, {**inp, "stg_keep": stg, "av_keep": av}) + for dtype, floor in ((torch.float32, 0.999), (torch.bfloat16, 0.999)): + rv, ra = _reference(model, shape, inp, dtype, stg_mask=stg_ref, isolate=(case == "isolated")) + cv = cosine(got["video_velocity"].cpu(), rv) + ca = cosine(got["audio_velocity"].cpu(), ra) + print(f"dit tiny {case} vs {dtype}: video cos {cv:.6f} relL2 {rel_l2(got['video_velocity'].cpu(), rv):.3e}" + f" | audio cos {ca:.6f} relL2 {rel_l2(got['audio_velocity'].cpu(), ra):.3e}") + assert cv > floor + assert ca > floor + + +def test_rope_grids_match_diffusers() -> None: + from diffusers.models.transformers.transformer_ltx2 import LTX2AudioVideoRotaryPosEmbed + + from families.ltx2.dit_builder import DiTConfig, DiTShape, audio_latent_frames, rope_tables + + cfg = DiTConfig.from_dict(dict(TINY_DIT, num_layers=1)) + sa = audio_latent_frames((FRAMES - 1) * 8 + 1, FPS) + shape = DiTShape(1, FRAMES, LH, LW, sa, TEXT, FPS) + ours = rope_tables(cfg, shape) + kw = dict(theta=10000.0, causal_offset=1, double_precision=True, rope_type="split") + video = LTX2AudioVideoRotaryPosEmbed(dim=64, base_num_frames=20, scale_factors=(8, 32, 32), modality="video", + num_attention_heads=4, **kw) + audio = LTX2AudioVideoRotaryPosEmbed(dim=32, base_num_frames=20, scale_factors=[4], modality="audio", + num_attention_heads=4, **kw) + ca_v = LTX2AudioVideoRotaryPosEmbed(dim=32, base_num_frames=20, scale_factors=(8, 32, 32), modality="video", + num_attention_heads=4, **kw) + ca_a = LTX2AudioVideoRotaryPosEmbed(dim=32, base_num_frames=20, scale_factors=(8, 32, 32), modality="audio", + num_attention_heads=4, **kw) + vc = video.prepare_video_coords(1, FRAMES, LH, LW, "cpu", fps=FPS) + ac = audio.prepare_audio_coords(1, sa, "cpu") + refs = {"video": video(vc), "audio": audio(ac), "ca_video": ca_v(vc[:, 0:1]), "ca_audio": ca_a(ac[:, 0:1])} + for key, (cos_ref, sin_ref) in refs.items(): + cos, sin = ours[key] + # diffusers: [B, H, T, r]; ours: [T, H, r] + assert torch.allclose(torch.from_numpy(cos).permute(1, 0, 2), cos_ref[0], atol=2e-6), key + assert torch.allclose(torch.from_numpy(sin).permute(1, 0, 2), sin_ref[0], atol=2e-6), key diff --git a/families/ltx2/tests/test_e2e.py b/families/ltx2/tests/test_e2e.py new file mode 100644 index 0000000000..79e9f4dd41 --- /dev/null +++ b/families/ltx2/tests/test_e2e.py @@ -0,0 +1,469 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Direct build, native-runtime and diffusers-reference E2E for ltx2 (LTX-2.5 distilled). + +The native run and the diffusers ``LTX2Pipeline`` start from the same seeded noise: the +native runtime reads it through ``TRTMC_LTX2_INITIAL_LATENTS`` (packed, normalized video +then audio noise) and the reference receives the equivalent unpacked, denormalized +``latents``/``audio_latents``, which its ``prepare_latents`` maps back to the same noise. +Both run the checkpoint's 8-step distilled schedule without guidance. + +Video frames are compared by PSNR. The released pipeline runs its vocoder in bf16, which +alone moves the waveform by a log-spectrum L1 of about 1.2 relative to an fp32 vocoder, +so the soundtrack is compared by the correlation of the log spectrograms instead of +sample-wise. +""" + +from __future__ import annotations + +import gc +import json +import os +import shutil +import subprocess +import sys +from pathlib import Path + +import numpy as np +import pytest + +from tensorrt_model_connect import BuildRequest, build +from tools.e2e_evidence import evidence_stage, record_evidence + +FAMILY = "ltx2" +TASKS = frozenset({"text_to_audio_video"}) +TEST_ROOT = Path(__file__).resolve().parent +REPO = TEST_ROOT.parents[2] +MANIFEST_ROOT = TEST_ROOT / "manifests" +THRESHOLD_ROOT = TEST_ROOT / "thresholds" + +# diffusers ``pipelines/ltx2/utils.py`` DISTILLED_SIGMA_VALUES (also baked into the bundle). +DISTILLED_SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875] +FRAME_RATE = 24.0 +LATENT_CHANNELS = 128 +AUDIO_LATENT_CHANNELS = 8 +AUDIO_LATENT_MEL_BINS = 16 + + +def _case_index() -> dict[str, tuple[Path, dict, dict]]: + result = {} + for path in sorted(MANIFEST_ROOT.glob("*.json")): + manifest = json.loads(path.read_text(encoding="utf-8")) + assert manifest["family"] == FAMILY + assert manifest["task"] in TASKS + for case in manifest["testcases"]: + name = str(case["name"]) + assert name not in result + result[name] = (path, manifest, case) + return result + + +CASES = _case_index() + + +def _selected_cases(config) -> tuple[list[str], bool]: + model_filters = set() + for raw in config.getoption("--e2e-model") or []: + model_filters.update(item.strip() for item in str(raw).split(",") if item.strip()) + models_file = config.getoption("--e2e-models-file") + if models_file: + model_filters.update( + line.strip() + for line in Path(models_file).read_text(encoding="utf-8").splitlines() + if line.strip() and not line.lstrip().startswith("#") + ) + testcase_filters = set() + for raw in config.getoption("--e2e-testcase") or []: + testcase_filters.update(item.strip() for item in str(raw).split(",") if item.strip()) + if not model_filters and not testcase_filters: + return sorted(CASES), False + selected = [] + for name, (_, manifest, _) in CASES.items(): + model_match = ( + not model_filters + or FAMILY in model_filters + or name in model_filters + or manifest["name"] in model_filters + ) + testcase_match = not testcase_filters or name in testcase_filters + if model_match and testcase_match: + selected.append(name) + return sorted(selected), True + + +def pytest_generate_tests(metafunc) -> None: + if "case_name" in metafunc.fixturenames: + names, enabled = _selected_cases(metafunc.config) + parameters = names + if not enabled: + parameters = [ + pytest.param( + name, + marks=pytest.mark.skip( + reason="direct E2E requires one of the three explicit E2E selectors" + ), + ) + for name in names + ] + metafunc.parametrize("case_name", parameters, ids=names) + + +def _required_path(value: str | None, label: str) -> Path: + assert value, f"selected {FAMILY} E2E requires {label}" + path = Path(value) + assert path.exists(), f"selected {FAMILY} E2E {label} does not exist: {path}" + return path + + +def _model_dir(manifest: dict) -> Path: + explicit = os.environ.get(f"TRTMC_{FAMILY.upper()}_MODEL_DIR") + if explicit: + return _required_path(explicit, f"TRTMC_{FAMILY.upper()}_MODEL_DIR") + from huggingface_hub import snapshot_download + + try: + snapshot = snapshot_download( + repo_id=manifest["hf_id"], + revision=manifest.get("hf_revision"), + local_files_only=True, + ) + except Exception as error: + raise AssertionError( + f"selected {FAMILY} E2E requires the exact cached checkpoint {manifest['hf_id']}" + ) from error + return Path(snapshot) + + +def _parallel_size(manifest: dict) -> int: + return int(manifest.get("context_parallel_size", 1)) + + +def _backend() -> str: + """The build backend: ``trt`` when TensorRT is installed, else ``trt_rtx`` (as conftest binds).""" + import importlib.util + + choice = os.environ.get("TRTMC_LTX2_TEST_BACKEND", "").strip() + if choice: + return choice + if sys.modules.get("tensorrt_rtx") is not None: # conftest bound TensorRT-RTX + return "trt_rtx" + return "trt" if importlib.util.find_spec("tensorrt") is not None else "trt_rtx" + + +def _library(runtime_root: Path, name: str) -> Path: + return runtime_root / (f"{name}.dll" if sys.platform == "win32" else f"lib{name}.so") + + +def _runtime(manifest: dict) -> tuple[Path, Path]: + binary = _required_path(os.environ.get("TRTMC_BINARY"), "TRTMC_BINARY") + runtime_root = _required_path(os.environ.get("TRTMC_RUNTIME_ROOT"), "TRTMC_RUNTIME_ROOT") + assert _library(runtime_root, f"trtmc_backend_{_backend()}").is_file() + assert _library(runtime_root, f"trtmc_model_{FAMILY}").is_file() + import torch + + required_gpus = _parallel_size(manifest) + assert torch.cuda.is_available(), f"selected {FAMILY} E2E requires CUDA" + assert torch.cuda.device_count() >= required_gpus, ( + f"selected {FAMILY} E2E requires {required_gpus} GPUs, found {torch.cuda.device_count()}" + ) + return binary, runtime_root + + +def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: + build( + BuildRequest( + model_dir=model_dir, + output_path=bundle, + family=FAMILY, + task=manifest["task"], + precision=manifest["precision"], + backend=_backend(), + image_height=manifest.get("image_height"), + image_width=manifest.get("image_width"), + video_num_frames=manifest.get("video_num_frames"), + context_parallel_size=_parallel_size(manifest), + ) + ) + + +def _select_json_payload(stdout: str, parallel_size: int) -> dict: + """The output rank's JSON line (rank 0 under ``mpirun --tag-output`` or ``launch_ranks``).""" + payloads = [] + for line in stdout.splitlines(): + if parallel_size > 1 and not line.startswith("[1,0]:"): + continue + start = line.find("{") + if start >= 0: + try: + payloads.append(json.loads(line[start:])) + except json.JSONDecodeError: + pass + payloads = [payload for payload in payloads if not payload.get("worker")] + assert len(payloads) == 1, f"expected one output-rank JSON payload: {stdout[-2000:]}" + return payloads[0] + + +def test_json_selection_uses_only_the_output_rank() -> None: + stdout = "\n".join( + ( + '[1,1]:{"worker": true}', + '[1,0]:{"output": "frames", "audio": "frames/audio.wav"}', + ) + ) + + assert _select_json_payload(stdout, 2) == {"output": "frames", "audio": "frames/audio.wav"} + + +def _run_json( + binary: Path, + runtime_root: Path, + bundle: Path, + manifest: dict, + case: dict, + noise_path: Path, + *arguments: str, +) -> dict: + invocation = [ + str(binary), + "generate-video", + str(bundle), + "--runtime-root", + str(runtime_root), + *arguments, + ] + parallel_size = _parallel_size(manifest) + env = os.environ.copy() + env["TRTMC_LTX2_INITIAL_LATENTS"] = str(noise_path) + if parallel_size > 1 and shutil.which("mpirun"): + env["TRTMC_NCCL_RENDEZVOUS"] = str(bundle.with_suffix(".nccl-rendezvous")) + invocation = [ + shutil.which("mpirun"), + "--tag-output", + "-x", + "LD_LIBRARY_PATH", + "-x", + "TRTMC_NCCL_RENDEZVOUS", + "-x", + "TRTMC_LTX2_INITIAL_LATENTS", + "-np", + str(parallel_size), + *invocation, + ] + elif parallel_size > 1: + # Hosts without OpenMPI (native Windows): the repository's local rank launcher + # provides the same rank environment and a fresh rendezvous file per launch. + invocation = [ + sys.executable, + str(REPO / "tools" / "launch_ranks.py"), + "-n", + str(parallel_size), + "--", + *invocation, + ] + if sys.platform == "win32": + env["PATH"] = os.pathsep.join(value for value in (str(runtime_root), env.get("PATH", "")) if value) + else: + env["LD_LIBRARY_PATH"] = ":".join( + value for value in (str(runtime_root), env.get("LD_LIBRARY_PATH", "")) if value + ) + completed = subprocess.run( + invocation, + capture_output=True, + text=True, + env=env, + timeout=int(case.get("runtime_timeout_s", 3600)), + ) + record_evidence("commands", {"argv": getattr(completed, "args", None)}) + record_evidence("native", {"stdout": completed.stdout[-20000:], "stderr": completed.stderr[-20000:]}) + assert completed.returncode == 0, ( + f"native generate-video failed ({completed.returncode}): {completed.stderr[-4000:]}" + ) + return _select_json_payload(completed.stdout, parallel_size) + + +def _thresholds(case_name: str) -> dict: + path = THRESHOLD_ROOT / f"{case_name}.json" + assert path.is_file(), f"selected {FAMILY} E2E requires exact thresholds: {path}" + return json.loads(path.read_text(encoding="utf-8"))["threshold_overrides"] + + +def _case_text(case: dict) -> str: + value = str(case.get("test_prompt") or "") + assert value, f"selected {FAMILY} E2E requires a direct prompt" + return value + + +def _latent_layout(manifest: dict) -> tuple[int, int, int, int]: + frames = int(manifest["video_num_frames"]) + height = int(manifest["image_height"]) + width = int(manifest["image_width"]) + assert (frames - 1) % 8 == 0 and height % 32 == 0 and width % 32 == 0 + # LTX2Pipeline: round(duration * sampling_rate / hop_length / temporal_compression). + audio_frames = round(frames / FRAME_RATE * 16000 / 160 / 4) + return (frames - 1) // 8 + 1, height // 32, width // 32, audio_frames + + +def _initial_noise(manifest: dict, case: dict) -> tuple[np.ndarray, np.ndarray]: + """Packed, normalized noise: video ``[S, 128]`` (tokens f, h, w) and audio ``[Sa, 128]``.""" + latent_frames, latent_height, latent_width, audio_frames = _latent_layout(manifest) + rng = np.random.default_rng(int(case["seed"])) + video = rng.standard_normal((latent_frames * latent_height * latent_width, LATENT_CHANNELS), dtype=np.float32) + audio = rng.standard_normal((audio_frames, AUDIO_LATENT_CHANNELS * AUDIO_LATENT_MEL_BINS), dtype=np.float32) + return video, audio + + +def _native(binary, runtime_root, bundle, manifest, case, tmp_path, noise) -> dict: + noise_path = tmp_path / "initial-noise.f32" + np.concatenate([noise[0].ravel(), noise[1].ravel()]).astype(np.float32).tofile(noise_path) + record_evidence("inputs", {"raw_file": noise_path}) + output = tmp_path / "native-frames" + payload = _run_json( + binary, runtime_root, bundle, manifest, case, noise_path, + "--prompt", _case_text(case), "--output", str(output), "--set", f"seed={int(case['seed'])}", + ) + payload["artifact"] = str(output) + return payload + + +def _official_reference(model_dir: Path, manifest: dict, case: dict, noise) -> dict: + import torch + from diffusers import LTX2Pipeline + + # The prompt enhancer (with its processor) and the duration head are optional components the + # distilled text-to-audio-video call does not use; the bundle does not contain them either. + pipeline = LTX2Pipeline.from_pretrained( + model_dir, torch_dtype=torch.bfloat16, local_files_only=True, + processor=None, prompt_enhancer=None, duration_head=None, + ).to("cuda") + latent_frames, latent_height, latent_width, audio_frames = _latent_layout(manifest) + # Denormalize so that the pipeline's prepare_latents normalizes back to the same noise. + vae, audio_vae = pipeline.vae, pipeline.audio_vae + video = torch.from_numpy(noise[0]).reshape(1, latent_frames, latent_height, latent_width, LATENT_CHANNELS) + video = video.permute(0, 4, 1, 2, 3).to("cuda", torch.float32) + mean = vae.latents_mean.view(1, -1, 1, 1, 1).to(video) + std = vae.latents_std.view(1, -1, 1, 1, 1).to(video) + video = video * std / vae.config.scaling_factor + mean + audio = torch.from_numpy(noise[1]).reshape(1, audio_frames, -1).to("cuda", torch.float32) + audio = audio * audio_vae.latents_std.to(audio) + audio_vae.latents_mean.to(audio) + audio = audio.unflatten(2, (AUDIO_LATENT_CHANNELS, AUDIO_LATENT_MEL_BINS)).transpose(1, 2) + frames, waveform = pipeline( + prompt=_case_text(case), + width=int(manifest["image_width"]), + height=int(manifest["image_height"]), + num_frames=int(manifest["video_num_frames"]), + frame_rate=FRAME_RATE, + sigmas=DISTILLED_SIGMAS, + guidance_scale=1.0, + audio_guidance_scale=1.0, + stg_scale=0.0, + audio_stg_scale=0.0, + modality_scale=1.0, + audio_modality_scale=1.0, + guidance_rescale=0.0, + audio_guidance_rescale=0.0, + spatio_temporal_guidance_blocks=None, + use_cross_timestep=True, + enable_prompt_enhancement=False, + latents=video, + audio_latents=audio, + generator=torch.Generator("cuda").manual_seed(int(case["seed"])), + output_type="np", + return_dict=False, + ) + result = { + "frames": (np.clip(frames[0], 0, 1) * 255).round().astype(np.uint8), + "audio": waveform[0].float().cpu().numpy(), + "audio_sample_rate": int(pipeline.vocoder.config.output_sampling_rate), + } + # Release the reference before the next case's native run needs the GPU memory. + del pipeline, frames, waveform + gc.collect() + torch.cuda.empty_cache() + return result + + +def _read_wav(path: Path) -> tuple[np.ndarray, int]: + import struct + + data = path.read_bytes() + position, layout, samples = 12, None, b"" + while position + 8 <= len(data): + chunk, size = data[position:position + 4], struct.unpack(" np.ndarray: + size, hop = 2048, 512 + windows = np.lib.stride_tricks.sliding_window_view(wave, size, axis=-1)[..., ::hop, :] + return np.log(np.abs(np.fft.rfft(windows * np.hanning(size), axis=-1)) + 1e-5) + + +def _metrics(actual: dict, expected: dict) -> dict: + from PIL import Image + + paths = sorted(Path(actual["artifact"]).glob("frame-*.png")) + frames = np.stack([np.asarray(Image.open(path).convert("RGB")) for path in paths]).astype(np.float64) + reference = expected["frames"].astype(np.float64) + assert frames.shape == reference.shape, (frames.shape, reference.shape) + mse = ((frames - reference) ** 2).mean(axis=(1, 2, 3)) + psnr = 10.0 * np.log10(255.0**2 / np.maximum(mse, 1e-12)) + wave, rate = _read_wav(Path(actual["audio"])) + reference_wave = expected["audio"] + assert rate == expected["audio_sample_rate"] + assert wave.shape == reference_wave.shape, (wave.shape, reference_wave.shape) + spectrogram = _log_spectrogram(wave).ravel() + reference_spectrogram = _log_spectrogram(reference_wave).ravel() + return { + "frames": int(frames.shape[0]), + "frame_psnr_mean_db": float(psnr.mean()), + "frame_psnr_min_db": float(psnr.min()), + "audio_log_spectrogram_corr": float(np.corrcoef(spectrogram, reference_spectrogram)[0, 1]), + "audio_rms_ratio": float(np.sqrt((wave**2).mean()) / max(np.sqrt((reference_wave**2).mean()), 1e-12)), + } + + +def _assert_contract(metrics: dict, manifest: dict, thresholds: dict) -> None: + assert metrics["frames"] == int(manifest["video_num_frames"]) + assert metrics["frame_psnr_mean_db"] >= float(thresholds["min_frame_psnr_mean_db"]) + assert metrics["frame_psnr_min_db"] >= float(thresholds["min_frame_psnr_min_db"]) + assert metrics["audio_log_spectrogram_corr"] >= float(thresholds["min_audio_log_spectrogram_corr"]) + low, high = (float(value) for value in thresholds["audio_rms_ratio_range"]) + assert low <= metrics["audio_rms_ratio"] <= high + + +def test_official_checkpoint_e2e(case_name: str, tmp_path: Path) -> None: + _, manifest, case = CASES[case_name] + record_evidence("inputs", {"manifest": manifest, "case": case}) + 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")}, + ) + binary, runtime_root = _runtime(manifest) + bundle = tmp_path / manifest["bundle"] + with evidence_stage("build"): + _build(model_dir, bundle, manifest) + noise = _initial_noise(manifest, case) + with evidence_stage("native"): + actual = _native(binary, runtime_root, bundle, manifest, case, tmp_path, noise) + record_evidence("native", {"payload": {key: actual[key] for key in actual if key != "frames"}}) + with evidence_stage("reference"): + expected = _official_reference(model_dir, manifest, case, noise) + with evidence_stage("compare"): + metrics = record_evidence("metrics", _metrics(actual, expected)) + print(f"[ltx2-e2e] {case_name} {json.dumps(metrics)}") + _assert_contract(metrics, manifest, record_evidence("thresholds", _thresholds(case_name))) diff --git a/families/ltx2/tests/test_model_contract.py b/families/ltx2/tests/test_model_contract.py new file mode 100644 index 0000000000..fba44587c5 --- /dev/null +++ b/families/ltx2/tests/test_model_contract.py @@ -0,0 +1,70 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Build-request contract of the LTX-2.5 family (no TensorRT, torch or checkpoint needed).""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from families.ltx2 import model + + +def _request(**overrides) -> SimpleNamespace: + fields = dict(task="text_to_audio_video", dynamic_kv_cache=False, tensor_parallel_size=1, + max_batch_size=1, quantization=None, fp32_layers=(), precision="bf16", context_parallel_size=1, + backend="trt", model_dir="/nonexistent/ltx2") + fields.update(overrides) + return SimpleNamespace(**fields) + + +@pytest.mark.parametrize("version", ["1.7.1", "1.7.1.107", "1.8.0", "2.0"]) +def test_rtx_context_parallel_gate_accepts_multi_device_releases(version: str) -> None: + model.require_rtx_context_parallel(version, 2) + + +@pytest.mark.parametrize("version", [None, "1.6.0.0", "1.6.1.4", "1.7.0.32"]) +def test_rtx_context_parallel_gate_rejects_older_releases(version: str | None) -> None: + with pytest.raises(RuntimeError) as error: + model.require_rtx_context_parallel(version, 2) + message = str(error.value) + assert "requires TensorRT-RTX >= 1.7.1" in message + assert f"(found {version or 'not installed'})" in message + assert "1.6.x ships without multi-device support" in message + + +@pytest.mark.parametrize("version", [None, "1.6.0.0"]) +def test_rtx_gate_does_not_apply_to_a_single_device(version: str | None) -> None: + model.require_rtx_context_parallel(version, 1) + + +def test_rtx_backend_build_checks_the_installed_version(monkeypatch) -> None: + monkeypatch.setattr(model, "installed_rtx_version", lambda: "1.6.1.4") + with pytest.raises(RuntimeError, match="TensorRT-RTX >= 1.7.1"): + model.build(_request(backend="trt_rtx", context_parallel_size=2), writer=None) + + +@pytest.mark.parametrize( + "overrides,error", + [ + ({"task": "text_to_video"}, ValueError), + ({"precision": "fp16"}, ValueError), + ({"tensor_parallel_size": 2}, NotImplementedError), + ({"max_batch_size": 2}, NotImplementedError), + ({"quantization": "fp8"}, NotImplementedError), + ({"fp32_layers": (3,)}, NotImplementedError), + ({"dynamic_kv_cache": True}, NotImplementedError), + ({"context_parallel_size": 4}, ValueError), + ], +) +def test_build_rejects_unsupported_requests(overrides: dict, error: type) -> None: + with pytest.raises(error): + model.build(_request(**overrides), writer=None) + + +def test_distilled_schedule_is_the_eight_step_checkpoint_schedule() -> None: + assert len(model.DISTILLED_SIGMAS) == 8 + assert model.DISTILLED_SIGMAS[0] == 1.0 + assert list(model.DISTILLED_SIGMAS) == sorted(model.DISTILLED_SIGMAS, reverse=True) diff --git a/families/ltx2/tests/test_support.py b/families/ltx2/tests/test_support.py new file mode 100644 index 0000000000..a5f312c0a1 --- /dev/null +++ b/families/ltx2/tests/test_support.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from tensorrt_model_connect.model_support import ModelMetadata, resolve_family + + +def test_ltx2_pipeline_resolves_to_text_to_audio_video() -> None: + family, support = resolve_family(ModelMetadata({}, {"_class_name": "LTX2Pipeline"})) + + assert family == "ltx2" + assert support.default_task == "text_to_audio_video" + assert support.default_precision == "bf16" diff --git a/families/ltx2/tests/test_text_encoder_parity.py b/families/ltx2/tests/test_text_encoder_parity.py new file mode 100644 index 0000000000..ca8ef4375e --- /dev/null +++ b/families/ltx2/tests/test_text_encoder_parity.py @@ -0,0 +1,185 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tiny-random parity: Gemma 4 text tower and LTX2TextConnectors vs transformers / diffusers. + +No checkpoint is needed: random modules are built from shrunk configs that keep every +feature of the real LTX-2.5 configs (sliding + full layers, head_dim 2x on full layers, +one global KV head with k_eq_v, proportional partial RoPE, a sliding window shorter than +the sequence, left padding, per-modality projections, registers, gated split-RoPE +connectors), saved as safetensors and built with the family builders. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("tensorrt") +if not torch.cuda.is_available(): + pytest.skip("CUDA is required for LTX-2.5 engine parity tests", allow_module_level=True) +safetensors_torch = pytest.importorskip("safetensors.torch") + +from families.ltx2.tests.engine_runner import cosine, rel_l2, run_plan # noqa: E402 + +SEQ = 16 +PAD = 5 + +TINY_GEMMA = { + "vocab_size": 512, + "hidden_size": 64, + "intermediate_size": 160, + "num_hidden_layers": 6, + "num_attention_heads": 4, + "num_key_value_heads": 2, + "head_dim": 16, + "global_head_dim": 32, + "num_global_key_value_heads": 1, + "attention_k_eq_v": True, + "sliding_window": 8, + "layer_types": ["sliding_attention", "sliding_attention", "full_attention", + "sliding_attention", "sliding_attention", "full_attention"], + "rope_parameters": { + "full_attention": {"partial_rotary_factor": 0.25, "rope_theta": 1000000.0, "rope_type": "proportional"}, + "sliding_attention": {"rope_theta": 10000.0, "rope_type": "default"}, + }, + "rms_norm_eps": 1e-6, + "hidden_activation": "gelu_pytorch_tanh", + "hidden_size_per_layer_input": 0, + "num_kv_shared_layers": 0, + "enable_moe_block": False, + "use_double_wide_mlp": False, + "attention_bias": False, + "pad_token_id": 0, +} + +TINY_CONNECTORS = { + "caption_channels": 64, + "text_proj_in_factor": 7, + "video_connector_num_attention_heads": 4, + "video_connector_attention_head_dim": 16, + "video_connector_num_layers": 2, + "video_connector_num_learnable_registers": 8, + "video_gated_attn": True, + "audio_connector_num_attention_heads": 4, + "audio_connector_attention_head_dim": 8, + "audio_connector_num_layers": 2, + "audio_connector_num_learnable_registers": 8, + "audio_gated_attn": True, + "connector_rope_base_seq_len": 4096, + "rope_theta": 10000.0, + "rope_double_precision": True, + "rope_type": "split", + "per_modality_projections": True, + "video_hidden_dim": 64, + "audio_hidden_dim": 32, + "proj_bias": True, + "causal_temporal_positioning": False, +} + + +def _randomize(module: "torch.nn.Module", seed: int) -> None: + gen = torch.Generator().manual_seed(seed) + with torch.no_grad(): + for name, p in module.named_parameters(): + if name.endswith("norm.weight") or "layernorm" in name or name.endswith("_norm.weight"): + p.copy_(1.0 + 0.2 * torch.randn(p.shape, generator=gen)) + elif p.ndim >= 2: + p.copy_(torch.randn(p.shape, generator=gen) / p.shape[-1] ** 0.5) + else: + p.copy_(0.1 * torch.randn(p.shape, generator=gen)) + for name, b in module.named_buffers(): + if name.endswith("layer_scalar"): + b.copy_(0.5 + torch.rand(b.shape, generator=gen)) + + +def _ids_and_mask(vocab: int): + gen = torch.Generator().manual_seed(7) + ids = torch.randint(3, vocab, (1, SEQ), generator=gen, dtype=torch.int64) + ids[:, :PAD] = 0 + mask = torch.ones(1, SEQ, dtype=torch.int64) + mask[:, :PAD] = 0 + return ids, mask + + +def _tiny_gemma(tmp_path: Path): + from transformers.models.gemma4_unified.configuration_gemma4_unified import Gemma4UnifiedTextConfig + from transformers.models.gemma4_unified.modeling_gemma4_unified import Gemma4UnifiedTextModel + + extra = {"global_head_dim", "num_global_key_value_heads"} + cfg = Gemma4UnifiedTextConfig(**{k: v for k, v in TINY_GEMMA.items() + if k in extra or hasattr(Gemma4UnifiedTextConfig, k)}) + cfg._attn_implementation = "eager" + model = Gemma4UnifiedTextModel(cfg).eval() + _randomize(model, 11) + folder = tmp_path / "text_encoder" + folder.mkdir() + state = {f"model.language_model.{k}": v.to(torch.bfloat16).contiguous() + for k, v in model.state_dict().items() if "rotary_emb" not in k} + safetensors_torch.save_file(state, str(folder / "model.safetensors")) + (folder / "config.json").write_text(json.dumps({"text_config": TINY_GEMMA}), encoding="utf-8") + return model, folder + + +def _gemma_reference(model, ids, mask, dtype): + model = model.to("cuda", dtype) + with torch.no_grad(): + out = model(input_ids=ids.cuda(), attention_mask=mask.cuda(), output_hidden_states=True) + return torch.stack(out.hidden_states, dim=-1).flatten(2, 3).float().cpu() + + +def test_gemma4_tiny_parity(tmp_path: Path) -> None: + from families.ltx2.text_encoder_builder import build_gemma_engine + + model, folder = _tiny_gemma(tmp_path) + ids, mask = _ids_and_mask(TINY_GEMMA["vocab_size"]) + plan = build_gemma_engine(folder, seq_len=SEQ) + got = run_plan(plan, {"input_ids": ids.int(), "attention_mask": mask.int()})["packed"].float().cpu() + ref32 = _gemma_reference(model, ids, mask, torch.float32) + ref16 = _gemma_reference(model, ids, mask, torch.bfloat16) + valid = slice(PAD, SEQ) # padding rows are zeroed by the connectors and never compared + assert torch.isfinite(got).all(), "padding rows must stay finite (they feed the connectors' select)" + c32 = cosine(got[:, valid], ref32[:, valid]) + c16 = cosine(got[:, valid], ref16[:, valid]) + base = cosine(ref16[:, valid], ref32[:, valid]) + print(f"gemma4 tiny: cos vs fp32 {c32:.6f}, vs bf16 {c16:.6f}, relL2 fp32 " + f"{rel_l2(got[:, valid], ref32[:, valid]):.4e}; torch bf16 vs fp32 cos {base:.6f} " + f"relL2 {rel_l2(ref16[:, valid], ref32[:, valid]):.4e}") + n = TINY_GEMMA["num_hidden_layers"] + 1 + per_layer = [cosine(got[:, valid].view(-1, TINY_GEMMA["hidden_size"], n)[..., i], + ref32[:, valid].view(-1, TINY_GEMMA["hidden_size"], n)[..., i]) for i in range(n)] + print("gemma4 tiny per-state cos vs fp32:", " ".join(f"{c:.5f}" for c in per_layer)) + assert c32 > 0.999 + assert c16 > 0.999 + + +def test_connectors_tiny_parity(tmp_path: Path) -> None: + from diffusers.pipelines.ltx2.connectors import LTX2TextConnectors + + from families.ltx2.text_encoder_builder import build_connectors_engine + + conn = LTX2TextConnectors(**TINY_CONNECTORS).eval() + _randomize(conn, 23) + folder = tmp_path / "connectors" + folder.mkdir() + safetensors_torch.save_file({k: v.contiguous() for k, v in conn.state_dict().items()}, + str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(TINY_CONNECTORS), encoding="utf-8") + width = TINY_CONNECTORS["caption_channels"] * TINY_CONNECTORS["text_proj_in_factor"] + packed = (3.0 * torch.randn(1, SEQ, width, generator=torch.Generator().manual_seed(3))).to(torch.bfloat16) + _, mask = _ids_and_mask(16) + plan = build_connectors_engine(folder, seq_len=SEQ) + got = run_plan(plan, {"packed": packed, "attention_mask": mask.int()}) + for dtype in (torch.float32, torch.bfloat16): + ref = conn.to("cuda", dtype) + with torch.no_grad(): + v, a, m = ref(packed.cuda().to(dtype), mask.cuda()) + assert bool((m == 1).all()) + cv = cosine(got["video_context"].float().cpu(), v.float().cpu()) + ca = cosine(got["audio_context"].float().cpu(), a.float().cpu()) + print(f"connectors tiny ({dtype}): video cos {cv:.6f}, audio cos {ca:.6f}") + assert cv > 0.999 + assert ca > 0.999 diff --git a/families/ltx2/tests/test_vae_parity.py b/families/ltx2/tests/test_vae_parity.py new file mode 100644 index 0000000000..701850279d --- /dev/null +++ b/families/ltx2/tests/test_vae_parity.py @@ -0,0 +1,91 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tiny-random parity: LTX-2.5 video VAE decoder engine vs diffusers ``AutoencoderKLLTX2Video``. + +Same block structure as the real LTX-2.5 VAE (four up blocks: spatiotemporal, spatiotemporal, +temporal, spatial; upsample factors 2/1/2/2; non-causal zero-padded decoder; patch 4), narrow +channels, non-trivial latent statistics. +""" + +from __future__ import annotations + +import json + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("tensorrt") +if not torch.cuda.is_available(): + pytest.skip("CUDA is required for LTX-2.5 engine parity tests", allow_module_level=True) +safetensors_torch = pytest.importorskip("safetensors.torch") + +from families.ltx2.tests.engine_runner import cosine, run_plan # noqa: E402 + +TINY_VAE = { + "block_out_channels": [8, 16, 32, 32], + "decoder_block_out_channels": [16, 32, 32, 64], + "decoder_causal": False, + "decoder_inject_noise": [False, False, False, False, False], + "decoder_layers_per_block": [1, 2, 1, 1, 1], + "decoder_spatial_padding_mode": "zeros", + "decoder_spatio_temporal_scaling": [True, True, True, True], + "down_block_types": ["LTX2VideoDownBlock3D"] * 4, + "downsample_type": ["spatial", "temporal", "spatiotemporal", "spatiotemporal"], + "encoder_causal": True, + "encoder_spatial_padding_mode": "zeros", + "in_channels": 3, + "latent_channels": 16, + "layers_per_block": [1, 1, 1, 1, 1], + "out_channels": 3, + "patch_size": 4, + "patch_size_t": 1, + "resnet_norm_eps": 1e-6, + "scaling_factor": 1.0, + "spatio_temporal_scaling": [True, True, True, True], + "timestep_conditioning": False, + "upsample_factor": [2, 2, 1, 2], + "upsample_residual": [False, False, False, False], + "upsample_type": ["spatiotemporal", "spatiotemporal", "temporal", "spatial"], +} +F, H, W = 3, 2, 3 + + +def test_vae_decoder_tiny_parity(tmp_path) -> None: + from diffusers import AutoencoderKLLTX2Video + + from families.ltx2.vae_builder import build_vae_decoder_engine + + vae = AutoencoderKLLTX2Video(**TINY_VAE).eval() + gen = torch.Generator().manual_seed(9) + with torch.no_grad(): + for name, p in vae.named_parameters(): + if p.ndim >= 2: + p.copy_(torch.randn(p.shape, generator=gen) * (1.0 / max(1, p[0].numel())) ** 0.5) + else: + p.copy_(1.0 + 0.1 * torch.randn(p.shape, generator=gen) if "norm" in name + else 0.05 * torch.randn(p.shape, generator=gen)) + vae.latents_mean.copy_(0.3 * torch.randn(16, generator=gen)) + vae.latents_std.copy_(0.5 + torch.rand(16, generator=gen)) + folder = tmp_path / "vae" + folder.mkdir() + state = {k: (v.to(torch.bfloat16) if k not in ("latents_mean", "latents_std") else v).contiguous() + for k, v in vae.state_dict().items()} + safetensors_torch.save_file(state, str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(TINY_VAE), encoding="utf-8") + + packed = torch.randn(1, F * H * W, 16, generator=gen) + plan = build_vae_decoder_engine(folder, latent_frames=F, latent_height=H, latent_width=W) + got = run_plan(plan, {"latents": packed})["frames"].float().cpu() # [T, H, W, 3] + for dtype in (torch.float32, torch.bfloat16): + ref_vae = vae.to("cuda", dtype) + z = packed.cuda().reshape(1, F, H, W, 16).permute(0, 4, 1, 2, 3) + z = z * ref_vae.latents_std.view(1, -1, 1, 1, 1).float() + ref_vae.latents_mean.view(1, -1, 1, 1, 1).float() + with torch.no_grad(): + video = ref_vae.decode(z.to(dtype), return_dict=False)[0].float() + ref = (video / 2 + 0.5).clamp(0, 1)[0].permute(1, 2, 3, 0).cpu() + assert tuple(got.shape) == tuple(ref.shape), (got.shape, ref.shape) + c = cosine(got - 0.5, ref - 0.5) + mae = float((got - ref).abs().mean()) + print(f"vae tiny vs {dtype}: cos(centered) {c:.6f} mean|diff| {mae:.2e} out {tuple(got.shape)}") + assert c > 0.999 diff --git a/families/ltx2/tests/thresholds/ltx25-distilled-cp2.json b/families/ltx2/tests/thresholds/ltx25-distilled-cp2.json new file mode 100644 index 0000000000..893de8645e --- /dev/null +++ b/families/ltx2/tests/thresholds/ltx25-distilled-cp2.json @@ -0,0 +1,8 @@ +{ + "threshold_overrides": { + "min_frame_psnr_mean_db": 15.0, + "min_frame_psnr_min_db": 14.0, + "min_audio_log_spectrogram_corr": 0.85, + "audio_rms_ratio_range": [0.8, 1.25] + } +} diff --git a/families/ltx2/tests/thresholds/ltx25-distilled-l0.json b/families/ltx2/tests/thresholds/ltx25-distilled-l0.json new file mode 100644 index 0000000000..893de8645e --- /dev/null +++ b/families/ltx2/tests/thresholds/ltx25-distilled-l0.json @@ -0,0 +1,8 @@ +{ + "threshold_overrides": { + "min_frame_psnr_mean_db": 15.0, + "min_frame_psnr_min_db": 14.0, + "min_audio_log_spectrogram_corr": 0.85, + "audio_rms_ratio_range": [0.8, 1.25] + } +} diff --git a/families/ltx2/text_encoder_builder.py b/families/ltx2/text_encoder_builder.py new file mode 100644 index 0000000000..e9f4e13fd6 --- /dev/null +++ b/families/ltx2/text_encoder_builder.py @@ -0,0 +1,440 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 prompt encoder: Gemma 4 text tower + LTX2TextConnectors as one TensorRT plan. + +Engine I/O (one prompt per call, left padded to ``seq_len``): + Inputs: + input_ids [1, L] int32 Gemma token ids (pad id 0 on the left) + attention_mask [1, L] int32 1 = token, 0 = padding + Outputs: + video_context [1, L, 4096] bf16 video connector output (``connector_prompt_embeds``) + audio_context [1, L, 2048] bf16 audio connector output + packed [1, L, 3840*49] bf16 (debug builds only) the stacked Gemma hidden states + +Gemma 4 (``Gemma4UnifiedForConditionalGeneration``, text tower only) is computed exactly as +transformers 5.18 does for a text-only, left-padded batch: positions ``arange(L)``, causal + +padding mask (a finite large negative, so padding rows stay finite), sliding layers with +head_dim 256 / GQA and default RoPE (theta 1e4), full layers with head_dim 512, one KV head, +``attention_k_eq_v`` (V = the raw K projection) and proportional partial RoPE (theta 1e6), +q/k RMSNorm with scale, v RMSNorm without scale, attention scale 1.0, sandwich norms and the +per-layer ``layer_scalar``. ``hidden_states`` are the embeddings, every layer output and the +final norm applied to the last layer; the connectors stack them hidden-major / layer-minor. + +The connector path is diffusers ``LTX2TextConnectors`` with ``per_modality_projections``: +per-token RMSNorm over the hidden axis of every layer, padding rows zeroed, per-modality +rescale and projection, valid tokens moved to the front with learnable registers in the tail +(an exact index computation replaces the argsort), then the 1D transformer stacks. +""" + +from __future__ import annotations + +import json +import math +import sys +from dataclasses import dataclass +from pathlib import Path + +import numpy as np +import tensorrt as trt + +from .checkpoint import Checkpoint +from .graph import Graph, build_plan, make_logger, new_network +from .layers import ( + AttnWeights, + connector_rope_tables, + feed_forward, + ltx_attention, + rope_constants, +) + +_NEG_MASK = -1.0e9 + + +@dataclass(frozen=True) +class Gemma4TextConfig: + hidden: int + intermediate: int + layers: int + heads: int + kv_heads: int + head_dim: int + global_head_dim: int + global_kv_heads: int + k_eq_v: bool + layer_types: tuple[str, ...] + sliding_window: int + local_theta: float + global_theta: float + global_partial: float + global_rope_type: str + eps: float + activation: str + + @staticmethod + def from_dict(cfg: dict) -> "Gemma4TextConfig": + tc = cfg.get("text_config", cfg) + unsupported = [] + if tc.get("enable_moe_block"): + unsupported.append("MoE blocks") + if int(tc.get("hidden_size_per_layer_input") or 0): + unsupported.append("per-layer input embeddings") + if int(tc.get("num_kv_shared_layers") or 0): + unsupported.append("KV-shared layers") + if tc.get("use_double_wide_mlp"): + unsupported.append("double-wide MLP") + if tc.get("attention_bias"): + unsupported.append("attention bias") + if unsupported: + raise NotImplementedError("LTX-2.5 Gemma 4 builder does not implement: " + ", ".join(unsupported)) + rope = tc.get("rope_parameters", {}) + local = rope.get("sliding_attention", {}) + glob = rope.get("full_attention", {}) + if local.get("rope_type", "default") != "default": + raise NotImplementedError(f"sliding RoPE type {local.get('rope_type')}") + if glob.get("rope_type", "default") not in ("default", "proportional"): + raise NotImplementedError(f"full-attention RoPE type {glob.get('rope_type')}") + k_eq_v = bool(tc.get("attention_k_eq_v", False)) + kv_heads = int(tc["num_key_value_heads"]) + global_kv = tc.get("num_global_key_value_heads") + return Gemma4TextConfig( + hidden=int(tc["hidden_size"]), + intermediate=int(tc["intermediate_size"]), + layers=int(tc["num_hidden_layers"]), + heads=int(tc["num_attention_heads"]), + kv_heads=kv_heads, + head_dim=int(tc["head_dim"]), + global_head_dim=int(tc.get("global_head_dim") or tc["head_dim"]), + global_kv_heads=int(global_kv) if (k_eq_v and global_kv) else kv_heads, + k_eq_v=k_eq_v, + layer_types=tuple(tc["layer_types"]), + sliding_window=int(tc.get("sliding_window") or 0), + local_theta=float(local.get("rope_theta", 10000.0)), + global_theta=float(glob.get("rope_theta", 10000.0)), + global_partial=float(glob.get("partial_rotary_factor", 1.0)), + global_rope_type=str(glob.get("rope_type", "default")), + eps=float(tc.get("rms_norm_eps", 1e-6)), + activation=str(tc.get("hidden_activation", "gelu_pytorch_tanh")), + ) + + +@dataclass(frozen=True) +class ConnectorConfig: + caption_channels: int + proj_in_factor: int + video_heads: int + video_head_dim: int + video_layers: int + video_registers: int + video_gated: bool + audio_heads: int + audio_head_dim: int + audio_layers: int + audio_registers: int + audio_gated: bool + video_hidden: int + audio_hidden: int + rope_base_seq_len: int + rope_theta: float + rope_double_precision: bool + + @staticmethod + def from_dict(cfg: dict) -> "ConnectorConfig": + if not cfg.get("per_modality_projections", False): + raise NotImplementedError("LTX-2.5 connectors builder expects per_modality_projections") + if cfg.get("rope_type", "split") != "split": + raise NotImplementedError("LTX-2.5 connectors builder expects split RoPE") + if cfg.get("causal_temporal_positioning", False): + raise NotImplementedError("causal temporal positioning is not used by LTX-2.5") + return ConnectorConfig( + caption_channels=int(cfg["caption_channels"]), + proj_in_factor=int(cfg["text_proj_in_factor"]), + video_heads=int(cfg["video_connector_num_attention_heads"]), + video_head_dim=int(cfg["video_connector_attention_head_dim"]), + video_layers=int(cfg["video_connector_num_layers"]), + video_registers=int(cfg.get("video_connector_num_learnable_registers") or 0), + video_gated=bool(cfg.get("video_gated_attn", False)), + audio_heads=int(cfg["audio_connector_num_attention_heads"]), + audio_head_dim=int(cfg["audio_connector_attention_head_dim"]), + audio_layers=int(cfg["audio_connector_num_layers"]), + audio_registers=int(cfg.get("audio_connector_num_learnable_registers") or 0), + audio_gated=bool(cfg.get("audio_gated_attn", False)), + video_hidden=int(cfg["video_hidden_dim"]), + audio_hidden=int(cfg["audio_hidden_dim"]), + rope_base_seq_len=int(cfg.get("connector_rope_base_seq_len", 4096)), + rope_theta=float(cfg.get("rope_theta", 10000.0)), + rope_double_precision=bool(cfg.get("rope_double_precision", True)), + ) + + +# ---------------------------------------------------------------------- Gemma 4 + + +def gemma_rope_tables(cfg: Gemma4TextConfig, seq_len: int, layer_type: str): + """cos/sin ``[L, head_dim]`` (rotate-half layout) for one layer type, rounded to bf16 like HF.""" + import ml_dtypes + + if layer_type == "full_attention": + head_dim = cfg.global_head_dim + base = cfg.global_theta + if cfg.global_rope_type == "proportional": + rope_angles = int(cfg.global_partial * head_dim // 2) + inv = 1.0 / (base ** (np.arange(0, 2 * rope_angles, 2, dtype=np.float32) / np.float32(head_dim))) + inv = np.concatenate([inv.astype(np.float32), + np.zeros(head_dim // 2 - rope_angles, dtype=np.float32)]) + else: + inv = 1.0 / (base ** (np.arange(0, head_dim, 2, dtype=np.float32) / np.float32(head_dim))) + else: + head_dim = cfg.head_dim + base = cfg.local_theta + inv = 1.0 / (base ** (np.arange(0, head_dim, 2, dtype=np.float32) / np.float32(head_dim))) + inv = inv.astype(np.float32) + pos = np.arange(seq_len, dtype=np.float32) + freqs = pos[:, None] * inv[None, :] + emb = np.concatenate([freqs, freqs], axis=1).astype(np.float32) + cos = np.cos(emb).astype(ml_dtypes.bfloat16).astype(np.float32) + sin = np.sin(emb).astype(ml_dtypes.bfloat16).astype(np.float32) + return cos, sin + + +def _gemma_rope(g: Graph, x4, cos_t, sin_t): + """Rotate-half RoPE on ``[1, L, H, D]`` (fp32 math, bf16 tables), result in x.dtype.""" + b, l, h, d = (int(s) for s in x4.shape) + out_dtype = x4.dtype + xf = g.cast(x4, trt.float32) + x1 = g.slice(xf, (0, 0, 0, 0), (b, l, h, d // 2)) + x2 = g.slice(xf, (0, 0, 0, d // 2), (b, l, h, d // 2)) + neg = g.mul(x2, g.scalar(-1.0, trt.float32, 4)) + rot = g.concat([neg, x1], axis=3) + return g.cast(g.add(g.mul(xf, cos_t), g.mul(rot, sin_t)), out_dtype) + + +def _repeat_heads(g: Graph, x, kv_heads: int, heads: int): + """``repeat_kv`` on ``[1, kvH, L, D]`` -> ``[1, H, L, D]``.""" + if kv_heads == heads: + return x + rep = heads // kv_heads + idx = g.const(np.repeat(np.arange(kv_heads, dtype=np.int32), rep), trt.int32) + return g.gather(x, idx, 1) + + +def _gemma_mask(g: Graph, mask_i, seq_len: int, window: int | None): + """Additive ``[1, 1, L, L]`` bf16 mask: causal (and sliding window) x key padding.""" + i = np.arange(seq_len)[:, None] + j = np.arange(seq_len)[None, :] + allowed = j <= i + if window: + allowed &= (i - j) < window + allowed_t = g.const(allowed.astype(np.float32).reshape(1, 1, seq_len, seq_len), trt.float32) + keys = g.reshape(g.cast(mask_i, trt.float32), (1, 1, 1, seq_len)) + keep = g.mul(allowed_t, keys) + add = g.mul(g.sub(g.scalar(1.0, trt.float32, 4), keep), g.scalar(_NEG_MASK, trt.float32, 4)) + return g.cast(add, trt.bfloat16) + + +def add_gemma4_text(g: Graph, ckpt: Checkpoint, cfg: Gemma4TextConfig, input_ids, mask_i, *, + seq_len: int, prefix: str = "model.language_model", num_layers: int | None = None): + """Gemma 4 text tower; returns the ``hidden_states`` list (embeddings, layers, final norm).""" + n_layers = cfg.layers if num_layers is None else num_layers + hidden = cfg.hidden + emb = ckpt.get(f"{prefix}.embed_tokens.weight") + table = g.const(emb, trt.bfloat16) + ids = g.reshape(input_ids, (seq_len,)) + x = g.reshape(g.gather(table, ids, 0), (1, seq_len, hidden)) + # Gemma4UnifiedTextScaledWordEmbedding: x * embed_scale, the scale rounded to the weight dtype. + x = g.mul(x, g.scalar(math.sqrt(hidden), trt.bfloat16, 3)) + + rope = {} + masks = {} + for lt in set(cfg.layer_types[:n_layers]): + cos, sin = gemma_rope_tables(cfg, seq_len, lt) + d = cos.shape[1] + rope[lt] = (g.const(cos.reshape(1, seq_len, 1, d), trt.float32), + g.const(sin.reshape(1, seq_len, 1, d), trt.float32)) + window = cfg.sliding_window if lt == "sliding_attention" and cfg.sliding_window < seq_len else None + key = window or 0 + if key not in masks: + masks[key] = _gemma_mask(g, mask_i, seq_len, window) + masks[lt] = masks[key] + + states = [x] + for li in range(n_layers): + p = f"{prefix}.layers.{li}" + lt = cfg.layer_types[li] + full = lt == "full_attention" + hd = cfg.global_head_dim if full else cfg.head_dim + kvh = cfg.global_kv_heads if full else cfg.kv_heads + alt = cfg.k_eq_v and full + + residual = x + h = g.rms_norm(x, ckpt.get(f"{p}.input_layernorm.weight", np.float32), cfg.eps) + q = g.reshape(g.linear(h, ckpt.get(f"{p}.self_attn.q_proj.weight")), (1, seq_len, cfg.heads, hd)) + q = g.rms_norm(q, ckpt.get(f"{p}.self_attn.q_norm.weight", np.float32), cfg.eps) + q = _gemma_rope(g, q, *rope[lt]) + k_raw = g.reshape(g.linear(h, ckpt.get(f"{p}.self_attn.k_proj.weight")), (1, seq_len, kvh, hd)) + v_raw = k_raw if alt else g.reshape(g.linear(h, ckpt.get(f"{p}.self_attn.v_proj.weight")), + (1, seq_len, kvh, hd)) + k = g.rms_norm(k_raw, ckpt.get(f"{p}.self_attn.k_norm.weight", np.float32), cfg.eps) + k = _gemma_rope(g, k, *rope[lt]) + v = g.rms_norm(v_raw, None, cfg.eps) + q4 = g.transpose(q, (0, 2, 1, 3)) + k4 = _repeat_heads(g, g.transpose(k, (0, 2, 1, 3)), kvh, cfg.heads) + v4 = _repeat_heads(g, g.transpose(v, (0, 2, 1, 3)), kvh, cfg.heads) + ctx = g.attention(q4, k4, v4, scale=1.0, mask=masks[lt]) + ctx = g.reshape(ctx, (1, seq_len, cfg.heads * hd), first=(0, 2, 1, 3)) + attn = g.linear(ctx, ckpt.get(f"{p}.self_attn.o_proj.weight")) + attn = g.rms_norm(attn, ckpt.get(f"{p}.post_attention_layernorm.weight", np.float32), cfg.eps) + x = g.add(residual, attn) + + residual = x + h = g.rms_norm(x, ckpt.get(f"{p}.pre_feedforward_layernorm.weight", np.float32), cfg.eps) + gate = g.gelu_tanh(g.linear(h, ckpt.get(f"{p}.mlp.gate_proj.weight"))) + up = g.linear(h, ckpt.get(f"{p}.mlp.up_proj.weight")) + mlp = g.linear(g.mul(gate, up), ckpt.get(f"{p}.mlp.down_proj.weight")) + mlp = g.rms_norm(mlp, ckpt.get(f"{p}.post_feedforward_layernorm.weight", np.float32), cfg.eps) + x = g.add(residual, mlp) + if ckpt.has(f"{p}.layer_scalar"): + scalar = ckpt.get(f"{p}.layer_scalar").reshape(1, 1, 1) + x = g.mul(x, g.const(scalar, trt.bfloat16)) + states.append(x) + if n_layers == cfg.layers: + states[-1] = g.rms_norm(x, ckpt.get(f"{prefix}.norm.weight", np.float32), cfg.eps) + return states + + +def stack_hidden_states(g: Graph, states): + """``torch.stack(states, dim=-1)``: ``[1, L, H]`` x N -> ``[1, L, H, N]``.""" + b, l, h = (int(s) for s in states[0].shape) + return g.concat([g.reshape(s, (b, l, h, 1)) for s in states], axis=3) + + +# ---------------------------------------------------------------------- connectors + + +def _front_aligned_rows(g: Graph, x, mask_i, registers: np.ndarray | None, seq_len: int): + """``LTX2ConnectorTransformer1d`` register replacement for a left-padded prompt. + + Valid tokens (the last n rows) move to rows 0..n-1 in order; rows n.. take the learnable + registers tiled to the sequence length. Exact replacement of the stable argsort. + """ + if registers is None: + return x + _, _, d = (int(s) for s in x.shape) + n_f = g.reduce(g.cast(mask_i, trt.float32), trt.ReduceOperation.SUM, 1, keep_dims=False) # [1] + n = g.cast(n_f, trt.int32) + pos = g.const(np.arange(seq_len, dtype=np.int32), trt.int32) # [L] + offset = g.sub(g.const(np.array([seq_len], np.int32), trt.int32), n) # L - n + src = g.minimum(g.add(pos, offset), g.const(np.array([seq_len - 1], np.int32), trt.int32)) + front = g.gather(x, src, 1) # [1, L, D] + is_valid = g.ew(pos, n, trt.ElementWiseOperation.LESS) # [L] + reps = seq_len // registers.shape[0] + tiled = np.tile(registers, (reps, 1)).reshape(1, seq_len, d) + reg_t = g.const(tiled, x.dtype) + cond = g.reshape(is_valid, (1, seq_len, 1)) + return g.select(cond, front, reg_t) + + +def _connector_stack(g: Graph, ckpt: Checkpoint, name: str, x, *, heads: int, head_dim: int, + layers: int, gated: bool, rope): + for li in range(layers): + p = f"{name}.transformer_blocks.{li}" + h = g.rms_norm(x, None, 1e-6) + x = g.add(x, ltx_attention(g, AttnWeights(ckpt, f"{p}.attn1"), h, h, heads=heads, eps=1e-6, + q_rope=rope, k_rope=rope, gated=gated)) + h = g.rms_norm(x, None, 1e-6) + x = g.add(x, feed_forward(g, ckpt, f"{p}.ff", h)) + return g.rms_norm(x, None, 1e-6) + + +def add_connectors(g: Graph, ckpt: Checkpoint, cfg: ConnectorConfig, stacked, mask_i, *, seq_len: int): + """diffusers ``LTX2TextConnectors.forward`` (per-modality projections); returns (video, audio).""" + b, l, h, n = (int(s) for s in stacked.shape) + if h != cfg.caption_channels or n != cfg.proj_in_factor: + raise ValueError(f"packed text states {h}x{n} do not match the connectors " + f"({cfg.caption_channels}x{cfg.proj_in_factor})") + # per_token_rms_norm over the hidden axis of each layer, eps 1e-6, fp32 statistics. + xf = g.cast(stacked, trt.float32) + ms = g.reduce(g.mul(xf, xf), trt.ReduceOperation.AVG, 2) + inv = g.unary(g.unary(g.add(ms, g.scalar(1e-6, trt.float32, 4)), trt.UnaryOperation.SQRT), + trt.UnaryOperation.RECIP) + valid = g.reshape(g.cast(mask_i, trt.float32), (1, seq_len, 1, 1)) + normed = g.mul(g.mul(xf, inv), valid) # padding rows -> 0 + flat = g.reshape(normed, (1, seq_len, h * n)) + + outs = [] + for mod, hidden, heads, head_dim, layers, regs, gated in ( + ("video", cfg.video_hidden, cfg.video_heads, cfg.video_head_dim, cfg.video_layers, + cfg.video_registers, cfg.video_gated), + ("audio", cfg.audio_hidden, cfg.audio_heads, cfg.audio_head_dim, cfg.audio_layers, + cfg.audio_registers, cfg.audio_gated), + ): + scale = math.sqrt(hidden / cfg.caption_channels) + x = g.cast(g.mul(flat, g.scalar(scale, trt.float32, 3)), trt.bfloat16) + x = g.linear(x, ckpt.get(f"{mod}_text_proj_in.weight"), ckpt.maybe(f"{mod}_text_proj_in.bias")) + registers = ckpt.maybe(f"{mod}_connector.learnable_registers") if regs else None + x = _front_aligned_rows(g, x, mask_i, registers, seq_len) + cos, sin = connector_rope_tables(seq_len, heads * head_dim, heads, base_seq_len=cfg.rope_base_seq_len, + theta=cfg.rope_theta, double_precision=cfg.rope_double_precision) + rope = rope_constants(g, cos, sin) + outs.append(_connector_stack(g, ckpt, f"{mod}_connector", x, heads=heads, head_dim=head_dim, + layers=layers, gated=gated, rope=rope)) + return outs[0], outs[1] + + +# ---------------------------------------------------------------------- engines + + +def build_text_encoder_engine(model_dir: str | Path, *, seq_len: int = 1024, debug_packed: bool = False, + verbose: bool = False, gemma_layers: int | None = None) -> bytes: + """Gemma 4 + connectors plan for one LTX-2.5 diffusers folder (``text_encoder/``, ``connectors/``).""" + model_dir = Path(model_dir) + te_cfg = json.loads((model_dir / "text_encoder" / "config.json").read_text(encoding="utf-8")) + gcfg = Gemma4TextConfig.from_dict(te_cfg) + ccfg = ConnectorConfig.from_dict(json.loads((model_dir / "connectors" / "config.json").read_text("utf-8"))) + if ccfg.caption_channels != gcfg.hidden or ccfg.proj_in_factor != gcfg.layers + 1: + raise ValueError("connectors do not match the Gemma text encoder") + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + ids = network.add_input("input_ids", trt.int32, (1, seq_len)) + mask = network.add_input("attention_mask", trt.int32, (1, seq_len)) + te = Checkpoint(model_dir / "text_encoder") + states = add_gemma4_text(g, te, gcfg, ids, mask, seq_len=seq_len, num_layers=gemma_layers) + stacked = stack_hidden_states(g, states) + if debug_packed: + g.mark_output(g.reshape(stacked, (1, seq_len, gcfg.hidden * len(states))), "packed") + video, audio = add_connectors(g, Checkpoint(model_dir / "connectors"), ccfg, stacked, mask, seq_len=seq_len) + g.mark_output(video, "video_context") + g.mark_output(audio, "audio_context") + print(f"[ltx2] Building text encoder engine (Gemma 4 {gemma_layers or gcfg.layers} layers + connectors, " + f"seq_len={seq_len}) ...", file=sys.stderr) + return build_plan(builder, network, label="text encoder") + + +def build_connectors_engine(connectors_dir: str | Path, *, seq_len: int, verbose: bool = False) -> bytes: + """Connectors only, from a packed ``[1, L, H*N]`` bf16 input (parity tests).""" + ckpt = Checkpoint(connectors_dir) + ccfg = ConnectorConfig.from_dict(ckpt.config()) + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + h, n = ccfg.caption_channels, ccfg.proj_in_factor + packed = network.add_input("packed", trt.bfloat16, (1, seq_len, h * n)) + mask = network.add_input("attention_mask", trt.int32, (1, seq_len)) + stacked = g.reshape(packed, (1, seq_len, h, n)) + video, audio = add_connectors(g, ckpt, ccfg, stacked, mask, seq_len=seq_len) + g.mark_output(video, "video_context") + g.mark_output(audio, "audio_context") + return build_plan(builder, network, label="connectors") + + +def build_gemma_engine(text_encoder_dir: str | Path, *, seq_len: int, verbose: bool = False) -> bytes: + """Gemma 4 text tower only, output ``packed`` (parity tests).""" + ckpt = Checkpoint(text_encoder_dir) + gcfg = Gemma4TextConfig.from_dict(ckpt.config()) + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + ids = network.add_input("input_ids", trt.int32, (1, seq_len)) + mask = network.add_input("attention_mask", trt.int32, (1, seq_len)) + states = add_gemma4_text(g, ckpt, gcfg, ids, mask, seq_len=seq_len) + stacked = stack_hidden_states(g, states) + g.mark_output(g.reshape(stacked, (1, seq_len, gcfg.hidden * len(states))), "packed") + return build_plan(builder, network, label="gemma") diff --git a/families/ltx2/vae_builder.py b/families/ltx2/vae_builder.py new file mode 100644 index 0000000000..eb36566a6b --- /dev/null +++ b/families/ltx2/vae_builder.py @@ -0,0 +1,188 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 video VAE decoder (``AutoencoderKLLTX2Video.decoder``) as a TensorRT plan. + +Engine I/O: + Inputs: + latents [1, S, 128] fp32 packed, normalized video latents (DiT layout, S = F*H*W tokens) + Outputs: + frames [T, H*32, W*32, 3] fp16 RGB in [0, 1] (``(x + 1) / 2`` clamped, the pipeline's + ``postprocess_video``) + +The engine denormalizes with the VAE's ``latents_mean`` / ``latents_std`` (``scaling_factor``), +unpacks the tokens to ``[1, C, F, H, W]`` and runs the non-causal ``LTX2VideoDecoder3d``: +temporal replicate padding + spatial zero padding for every 3x3x3 convolution, per-pixel RMS +norms (eps 1e-8) and SiLU in fp32 islands, bf16 convolutions, per-block upsamplers with their own +stride (spatiotemporal / temporal / spatial) and depth-to-space shuffles, then the 4x4 spatial +unpatchify. ``timestep_conditioning`` and decoder noise injection are not used by LTX-2.5 and +are rejected. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import numpy as np +import tensorrt as trt + +from .checkpoint import Checkpoint +from .graph import Graph, build_plan, make_logger, new_network + +_UPSAMPLE_STRIDES = {"spatiotemporal": (2, 2, 2), "temporal": (2, 1, 1), "spatial": (1, 2, 2)} + + +def _conv3d(g: Graph, x, weight: np.ndarray, bias: np.ndarray | None, *, causal: bool = False): + """``LTX2VideoCausalConv3d`` (non-causal mode) or a plain 1x1x1 ``nn.Conv3d``.""" + b, c, t, h, w = (int(s) for s in x.shape) + out_c, _, kt, kh, kw = (int(s) for s in weight.shape) + if kt > 1: + pad = kt - 1 if causal else (kt - 1) // 2 + first = g.slice(x, (0, 0, 0, 0, 0), (b, c, 1, h, w)) + parts = [first] * pad + [x] + if not causal: + last = g.slice(x, (0, 0, t - 1, 0, 0), (b, c, 1, h, w)) + parts += [last] * pad + x = g.concat(parts, axis=2) + layer = g.net.add_convolution_nd(x, out_c, (kt, kh, kw), g.weights(weight, x.dtype), + g.weights(bias, x.dtype) if bias is not None else trt.Weights()) + layer.stride_nd = (1, 1, 1) + layer.padding_nd = (0, kh // 2, kw // 2) + return layer.get_output(0) + + +def _norm_silu(g: Graph, x, eps: float = 1e-8): + """``PerChannelRMSNorm`` (``x / sqrt(mean_c(x^2) + eps)``) then SiLU, one fp32 island.""" + out = x.dtype + xf = g.cast(x, trt.float32) + ms = g.reduce(g.mul(xf, xf), trt.ReduceOperation.AVG, 1) + inv = g.unary(g.unary(g.add(ms, g.scalar(eps, trt.float32, 5)), trt.UnaryOperation.SQRT), + trt.UnaryOperation.RECIP) + return g.cast(g.silu(g.mul(xf, inv)), out) + + +def _channel_layer_norm(g: Graph, x, gamma: np.ndarray, beta: np.ndarray, eps: float): + out = x.dtype + c = int(x.shape[1]) + xf = g.cast(x, trt.float32) + mean = g.reduce(xf, trt.ReduceOperation.AVG, 1) + cen = g.sub(xf, mean) + var = g.reduce(g.mul(cen, cen), trt.ReduceOperation.AVG, 1) + inv = g.unary(g.unary(g.add(var, g.scalar(eps, trt.float32, 5)), trt.UnaryOperation.SQRT), + trt.UnaryOperation.RECIP) + y = g.mul(cen, inv) + y = g.add(g.mul(y, g.const(gamma.reshape(1, c, 1, 1, 1), trt.float32)), + g.const(beta.reshape(1, c, 1, 1, 1), trt.float32)) + return g.cast(y, out) + + +def _resnet(g: Graph, ck: Checkpoint, p: str, x, eps: float): + h = _norm_silu(g, x) + h = _conv3d(g, h, ck.get(f"{p}.conv1.conv.weight"), ck.get(f"{p}.conv1.conv.bias")) + h = _norm_silu(g, h) + h = _conv3d(g, h, ck.get(f"{p}.conv2.conv.weight"), ck.get(f"{p}.conv2.conv.bias")) + shortcut = x + if ck.has(f"{p}.norm3.weight"): + shortcut = _channel_layer_norm(g, shortcut, ck.get(f"{p}.norm3.weight", np.float32), + ck.get(f"{p}.norm3.bias", np.float32), eps) + for key in (f"{p}.conv_shortcut", f"{p}.conv_shortcut.conv"): + if ck.has(f"{key}.weight"): + shortcut = _conv3d(g, shortcut, ck.get(f"{key}.weight"), ck.maybe(f"{key}.bias")) + break + return g.add(h, shortcut) + + +def _depth_to_space(g: Graph, x, stride: tuple[int, int, int]): + """``LTX2VideoUpsampler3d`` shuffle: ``[B, C*s0*s1*s2, F, H, W]`` -> ``[B, C, F*s0-(s0-1), H*s1, W*s2]``.""" + b, cc, f, h, w = (int(s) for s in x.shape) + s0, s1, s2 = stride + c = cc // (s0 * s1 * s2) + y = g.reshape(x, (b, c, s0, s1, s2, f, h, w)) + y = g.reshape(y, (b, c, f * s0, h * s1, w * s2), first=(0, 1, 5, 2, 6, 3, 7, 4)) + if s0 > 1: + y = g.slice(y, (0, 0, s0 - 1, 0, 0), (b, c, f * s0 - (s0 - 1), h * s1, w * s2)) + return y + + +def _count(ck: Checkpoint, fmt: str) -> int: + n = 0 + while ck.has(fmt.format(n)): + n += 1 + return n + + +def build_vae_decoder_engine(vae_dir: str | Path, *, latent_frames: int, latent_height: int, latent_width: int, + verbose: bool = False, precision: str = "bf16") -> bytes: + ck = Checkpoint(vae_dir) + cfg = ck.config() + if cfg.get("timestep_conditioning"): + raise NotImplementedError("timestep-conditioned LTX-2 VAE decoders are not supported") + if any(bool(v) for v in (cfg.get("decoder_inject_noise") or ())): + raise NotImplementedError("LTX-2 VAE decoder noise injection is not supported") + if cfg.get("decoder_causal", False): + raise NotImplementedError("the LTX-2.5 decoder is non-causal; causal decoding is not implemented") + if cfg.get("decoder_spatial_padding_mode", "zeros") != "zeros": + raise NotImplementedError("only zero spatial padding is implemented") + if any(bool(v) for v in (cfg.get("upsample_residual") or ())): + raise NotImplementedError("residual upsamplers are not used by LTX-2.5 and are not implemented") + patch, patch_t = int(cfg.get("patch_size", 4)), int(cfg.get("patch_size_t", 1)) + if patch_t != 1: + raise NotImplementedError("temporal patching is not implemented") + eps = float(cfg.get("resnet_norm_eps", 1e-6)) + channels = list(reversed(cfg["decoder_block_out_channels"])) + factors = list(reversed(cfg["upsample_factor"])) + scaling = list(reversed(cfg["decoder_spatio_temporal_scaling"])) + up_types = list(cfg["upsample_type"]) + latent_c = int(cfg.get("latent_channels", 128)) + dt = {"bf16": trt.bfloat16, "fp16": trt.float16, "fp32": trt.float32}[precision] + + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + f, h, w = latent_frames, latent_height, latent_width + s = f * h * w + z = network.add_input("latents", trt.float32, (1, s, latent_c)) + mean = ck.get("latents_mean", np.float32).reshape(1, 1, latent_c) + std = ck.get("latents_std", np.float32).reshape(1, 1, latent_c) + sf = float(cfg.get("scaling_factor", 1.0)) + x = g.add(g.mul(z, g.const(std / sf, trt.float32)), g.const(mean, trt.float32)) + x = g.reshape(g.transpose(x, (0, 2, 1)), (1, latent_c, f, h, w)) + x = g.cast(x, dt) + + x = _conv3d(g, x, ck.get("decoder.conv_in.conv.weight"), ck.get("decoder.conv_in.conv.bias")) + for i in range(_count(ck, "decoder.mid_block.resnets.{}.conv1.conv.weight")): + x = _resnet(g, ck, f"decoder.mid_block.resnets.{i}", x, eps) + for bi in range(len(channels)): + p = f"decoder.up_blocks.{bi}" + if ck.has(f"{p}.conv_in.conv1.conv.weight"): + x = _resnet(g, ck, f"{p}.conv_in", x, eps) + if scaling[bi]: + stride = _UPSAMPLE_STRIDES[up_types[bi]] + x = _conv3d(g, x, ck.get(f"{p}.upsamplers.0.conv.conv.weight"), + ck.get(f"{p}.upsamplers.0.conv.conv.bias")) + x = _depth_to_space(g, x, stride) + for ri in range(_count(ck, p + ".resnets.{}.conv1.conv.weight")): + x = _resnet(g, ck, f"{p}.resnets.{ri}", x, eps) + expected = channels[bi] // factors[bi] + if int(x.shape[1]) != expected: + raise ValueError(f"VAE up block {bi}: {int(x.shape[1])} channels, config expects {expected}") + x = _norm_silu(g, x) + x = _conv3d(g, x, ck.get("decoder.conv_out.conv.weight"), ck.get("decoder.conv_out.conv.bias")) + b, cc, t, hh, ww = (int(v) for v in x.shape) + c = cc // (patch * patch) + # reshape(B, C, p_t, p, p, F, H, W).permute(0, 1, 5, 2, 6, 4, 7, 3) -> [B, C, F, H*p, W*p] + x = g.reshape(x, (b, c, 1, patch, patch, t, hh, ww)) + x = g.reshape(x, (c, t, hh * patch, ww * patch), first=(0, 1, 5, 2, 6, 4, 7, 3)) + x = g.cast(x, trt.float32) + x = g.mul(g.add(x, g.scalar(1.0, trt.float32, 4)), g.scalar(0.5, trt.float32, 4)) + x = g.maximum(g.minimum(x, g.scalar(1.0, trt.float32, 4)), g.scalar(0.0, trt.float32, 4)) + frames = g.transpose(x, (1, 2, 3, 0)) # [T, H, W, 3] + g.mark_output(frames, "frames", trt.float16) + print(f"[ltx2] Building video VAE decoder engine (latent {f}x{h}x{w} -> {t}x{hh * patch}x{ww * patch}, " + f"{precision}) ...", file=sys.stderr) + return build_plan(builder, network, label="video VAE decoder") + + +def vae_config(vae_dir: str | Path) -> dict: + return json.loads((Path(vae_dir) / "config.json").read_text(encoding="utf-8")) From cea37d589bf009fe69503906e0d1f1788f8fb551 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 21:00:10 -0700 Subject: [PATCH 11/15] fix(ltx2): declare tensor_parallel_size in the E2E manifests The docs model-support inventory requires every manifest to declare precision and tensor_parallel_size. Declare tensor_parallel_size=1 in both ltx2 manifests, assert it when indexing the cases, and pass it into the build request so the manifest key is a used family test input. Signed-off-by: Peter Kisfaludi --- families/ltx2/tests/manifests/ltx25-distilled-cp2.json | 1 + families/ltx2/tests/manifests/ltx25-distilled-l0.json | 1 + families/ltx2/tests/test_e2e.py | 2 ++ 3 files changed, 4 insertions(+) diff --git a/families/ltx2/tests/manifests/ltx25-distilled-cp2.json b/families/ltx2/tests/manifests/ltx25-distilled-cp2.json index 5e12e482f7..4aa3b31ad1 100644 --- a/families/ltx2/tests/manifests/ltx25-distilled-cp2.json +++ b/families/ltx2/tests/manifests/ltx25-distilled-cp2.json @@ -14,6 +14,7 @@ "runtime_timeout_s": 3600 } ], + "tensor_parallel_size": 1, "context_parallel_size": 2, "image_height": 544, "image_width": 960, diff --git a/families/ltx2/tests/manifests/ltx25-distilled-l0.json b/families/ltx2/tests/manifests/ltx25-distilled-l0.json index d97178544f..f26fbfcc81 100644 --- a/families/ltx2/tests/manifests/ltx25-distilled-l0.json +++ b/families/ltx2/tests/manifests/ltx25-distilled-l0.json @@ -15,6 +15,7 @@ "runtime_timeout_s": 3600 } ], + "tensor_parallel_size": 1, "context_parallel_size": 1, "image_height": 384, "image_width": 640, diff --git a/families/ltx2/tests/test_e2e.py b/families/ltx2/tests/test_e2e.py index 79e9f4dd41..e09e5ea18a 100644 --- a/families/ltx2/tests/test_e2e.py +++ b/families/ltx2/tests/test_e2e.py @@ -52,6 +52,7 @@ def _case_index() -> dict[str, tuple[Path, dict, dict]]: manifest = json.loads(path.read_text(encoding="utf-8")) assert manifest["family"] == FAMILY assert manifest["task"] in TASKS + assert manifest["tensor_parallel_size"] == 1 for case in manifest["testcases"]: name = str(case["name"]) assert name not in result @@ -182,6 +183,7 @@ def _build(model_dir: Path, bundle: Path, manifest: dict) -> None: image_height=manifest.get("image_height"), image_width=manifest.get("image_width"), video_num_frames=manifest.get("video_num_frames"), + tensor_parallel_size=int(manifest["tensor_parallel_size"]), context_parallel_size=_parallel_size(manifest), ) ) From a2282585e38bfbe59e18cace5f76bcadb5626953 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Thu, 1 Oct 2026 21:28:06 -0700 Subject: [PATCH 12/15] fix(ltx2): remove the NCCL rendezvous file after communicator init Rank 0 wrote the unique-id file and never removed it. When a launch reused the path, a non-zero rank could read the previous run's id before rank 0 replaced it, and ncclCommInitRank, which has no timeout, then blocked every rank. ncclCommInitRank returns on rank 0 only after all ranks joined, so rank 0 now removes the file at that point, in the runtime and in the test communicator helper. The mpirun E2E lane also clears a file left by an interrupted launch before it starts the ranks. Signed-off-by: Peter Kisfaludi --- families/ltx2/runtime/distributed_runtime.cpp | 8 ++++++++ families/ltx2/tests/dist_helpers.py | 4 ++++ families/ltx2/tests/test_e2e.py | 5 ++++- 3 files changed, 16 insertions(+), 1 deletion(-) diff --git a/families/ltx2/runtime/distributed_runtime.cpp b/families/ltx2/runtime/distributed_runtime.cpp index 8b8fb76ac0..e0a8bf07f3 100644 --- a/families/ltx2/runtime/distributed_runtime.cpp +++ b/families/ltx2/runtime/distributed_runtime.cpp @@ -16,6 +16,7 @@ #include #include #include +#include #include namespace trtmc::ltx2 { @@ -207,6 +208,13 @@ DistributedRuntimeGroup initialize_parallel_group(int parallel_size) { id = read_unique_id(path); } runtime->init(parallel_size, group.rank, id); + if (group.rank == 0) { + // ncclCommInitRank returns only after every rank joined, so every rank + // has read the ID. Remove it so a reused path cannot hand a stale ID to + // the next launch. + std::error_code ignored; + std::filesystem::remove(path, ignored); + } group.communicator = runtime->communicator(); group.owner = std::move(runtime); return group; diff --git a/families/ltx2/tests/dist_helpers.py b/families/ltx2/tests/dist_helpers.py index 8b5680c8d0..9320e6c9d5 100644 --- a/families/ltx2/tests/dist_helpers.py +++ b/families/ltx2/tests/dist_helpers.py @@ -49,6 +49,10 @@ def __init__(self): self.comm = ctypes.c_void_p() self._check(self.lib.ncclCommInitRank(ctypes.byref(self.comm), self.world, uid, self.rank), "ncclCommInitRank") + if self.rank == 0: + # Every rank joined, so every rank read the id; a reused path must not + # hand this stale id to the next launch. + path.unlink(missing_ok=True) @staticmethod def _check(status: int, what: str) -> None: diff --git a/families/ltx2/tests/test_e2e.py b/families/ltx2/tests/test_e2e.py index e09e5ea18a..7c5acfd3ea 100644 --- a/families/ltx2/tests/test_e2e.py +++ b/families/ltx2/tests/test_e2e.py @@ -238,7 +238,10 @@ def _run_json( env = os.environ.copy() env["TRTMC_LTX2_INITIAL_LATENTS"] = str(noise_path) if parallel_size > 1 and shutil.which("mpirun"): - env["TRTMC_NCCL_RENDEZVOUS"] = str(bundle.with_suffix(".nccl-rendezvous")) + rendezvous = bundle.with_suffix(".nccl-rendezvous") + # A file left by an interrupted launch would hand its stale id to the ranks. + rendezvous.unlink(missing_ok=True) + env["TRTMC_NCCL_RENDEZVOUS"] = str(rendezvous) invocation = [ shutil.which("mpirun"), "--tag-output", From 50bb0654061ad70ed94b68276e64d7344abc8837 Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Mon, 5 Oct 2026 15:21:25 -0700 Subject: [PATCH 13/15] fix(cli): reject an audio-video result without a sample rate write_audio_video wrote audio.wav with sample_rate.value_or(0) when a text_to_audio_video family returned no sample rate, producing a WAV header with rate 0 and reporting audio_sample_rate 0 while the command succeeded. Fail explicitly instead. Signed-off-by: Peter Kisfaludi --- apps/cli/sdk_video.cpp | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/apps/cli/sdk_video.cpp b/apps/cli/sdk_video.cpp index 5b0e831a2f..c3b5abc182 100644 --- a/apps/cli/sdk_video.cpp +++ b/apps/cli/sdk_video.cpp @@ -88,12 +88,13 @@ nlohmann::json write_audio_video(const AudioVideoGenerationResult& result, return {{"worker", true}}; auto json = write_frames(frames, result.timestamps_seconds(), directory); const auto audio = result.audio(); + if (!audio.sample_rate || *audio.sample_rate == 0) + throw std::runtime_error("text_to_audio_video result has no audio sample rate"); const auto audio_path = (std::filesystem::path(directory) / "audio.wav").string(); - io::write_wav_interleaved(audio.samples, - static_cast(audio.sample_rate.value_or(0)), + io::write_wav_interleaved(audio.samples, static_cast(*audio.sample_rate), static_cast(audio.channels), audio_path); json["audio"] = audio_path; - json["audio_sample_rate"] = audio.sample_rate.value_or(0); + json["audio_sample_rate"] = *audio.sample_rate; json["audio_channels"] = audio.channels; json["audio_start_seconds"] = result.audio_start_seconds(); json["setup_ms"] = result.video_view().setup_ms; From 0c126dba5a8d67e8e48956600252b8d5b565794f Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Fri, 2 Oct 2026 10:44:38 -0700 Subject: [PATCH 14/15] feat(ltx2): decode the video VAE in tiles shared across CP ranks The video VAE decoded the whole clip on rank 0, so the second GPU idled for the whole decode. The VAE now decodes overlapping tiles of one shape with one static tile plan. The tiles are blended with linear ramps that are normalized by the summed weights, as in the Lightricks and TRT-LLM tiled_decode. The build computes the tile plan (512 px tiles, >= 64 px overlap; time splits into 256-frame tiles only for longer clips) and writes it into runtime.json. With context parallelism the ranks decode disjoint tiles. Worker ranks send their fp16 tiles to rank 0 with NCCL point-to-point on the engines' communicator, and the last rank also decodes the audio and sends the waveform. Rank 0 blends every tile in tile order, so the output is the same bit for bit as the single-GPU tiled decode. A transfer that misses its deadline aborts the communicator. On one GPU, the host blend overlaps the audio decode. trtmc ltx2 build takes the tile options; sizes of 0 build the untiled decoder. Bundles without a tile plan keep the rank-0 decode. Signed-off-by: Peter Kisfaludi --- families/ltx2/README.md | 33 +- families/ltx2/audio_builder.py | 5 +- families/ltx2/cli.json | 28 ++ families/ltx2/cli.py | 89 +++++ families/ltx2/model.py | 48 +-- families/ltx2/runtime/distributed_runtime.cpp | 95 +++++- families/ltx2/runtime/distributed_runtime.h | 25 ++ families/ltx2/runtime/pipeline.cpp | 304 +++++++++++++++++- families/ltx2/runtime/pipeline.h | 43 ++- families/ltx2/runtime/plugin.cpp | 37 ++- families/ltx2/runtime/vae_tiling.h | 231 +++++++++++++ .../ltx2/tests/cpp/test_runtime_contract.cpp | 120 ++++++- families/ltx2/tests/dist_helpers.py | 24 ++ families/ltx2/tests/dist_vae_tile_check.py | 109 +++++++ families/ltx2/tests/test_context_parallel.py | 35 ++ families/ltx2/tests/test_vae_parity.py | 75 ++++- families/ltx2/tests/test_vae_tiling.py | 112 +++++++ families/ltx2/tests/vae_tile_prep.py | 49 +++ families/ltx2/vae_builder.py | 11 +- families/ltx2/vae_tiling.py | 192 +++++++++++ 20 files changed, 1600 insertions(+), 65 deletions(-) create mode 100644 families/ltx2/cli.json create mode 100644 families/ltx2/cli.py create mode 100644 families/ltx2/runtime/vae_tiling.h create mode 100644 families/ltx2/tests/dist_vae_tile_check.py create mode 100644 families/ltx2/tests/test_vae_tiling.py create mode 100644 families/ltx2/tests/vae_tile_prep.py create mode 100644 families/ltx2/vae_tiling.py diff --git a/families/ltx2/README.md b/families/ltx2/README.md index 05fc8e6cba..0e1723590c 100644 --- a/families/ltx2/README.md +++ b/families/ltx2/README.md @@ -13,8 +13,9 @@ Text-to-audio-video for Lightricks LTX-2.5 diffusers checkpoints (`LTX2Pipeline` `--video-num-frames`). Height and width must be multiples of 32, and the frame count must be `8n+1`. The default is 960x544, 121 frames at 24 fps. - `--context-parallel-size 1` runs on one GPU. `--context-parallel-size 2` splits the - video tokens of the DiT across two GPUs. The text encoder, video VAE and audio decoder - run on rank 0 only. + video tokens of the DiT across two GPUs. The text encoder runs on every rank. The video + VAE decodes in tiles that the ranks share, and the last rank decodes the audio + (see [Tiled video decode](#tiled-video-decode)). ## Bundle @@ -22,7 +23,7 @@ Text-to-audio-video for Lightricks LTX-2.5 diffusers checkpoints (`LTX2Pipeline` |---|---| | `text_encoder.plan` | Gemma 4 text tower and the LTX-2 text connectors (video and audio context) | | `denoiser.plan` | Joint audio/video DiT. With CP=2, one plan serves both ranks. | -| `vae.plan` | Video VAE decoder | +| `vae.plan` | Video VAE decoder for one tile shape (whole video with `--vae-tile-pixels 0 --vae-tile-frames 0`) | | `audio.plan` | Audio VAE decoder and vocoder with bandwidth extension | | `tokenizer.json`, `runtime.json` | Tokenizer, shapes and schedule | @@ -31,6 +32,28 @@ tokens. Video self-attention all-gathers each rank's normed and rotated keys and Video-to-audio attention merges per-rank softmax statistics through one small all-gather. The network uses no all-to-all collective. +## Tiled video decode + +The video VAE decodes the latent video as overlapping tiles of one shape, so one static plan serves +every tile. The tiles are blended with linear ramps over their overlaps and normalized by the summed +weights, as in the Lightricks and TensorRT-LLM `tiled_decode`. Tiles are at most 512 pixels with at +least 64 pixels of overlap; clips longer than 257 frames also split in time, into tiles of up to +256 frames that overlap by at least 24 frames. `families/ltx2/vae_tiling.py` computes the tile plan +and writes it into `runtime.json`. + +- With context parallelism, the ranks decode disjoint tiles. The worker ranks send their decoded + tiles to rank 0 over NCCL point-to-point on the engines' communicator, and rank 0 blends every + tile in tile order. The blended video is the same bit for bit as the single-GPU tiled decode. +- The last rank decodes the audio while rank 0 decodes its tiles. On one GPU, the host blend runs + while the GPU decodes the audio. +- A transfer that does not finish within 10 minutes aborts the NCCL communicator instead of + leaving NCCL kernels running. + +`trtmc ltx2 build` (`python -m tensorrt_model_connect ltx2 build`) takes the tile options +`--vae-tile-pixels`, `--vae-tile-overlap-pixels`, `--vae-tile-frames` and +`--vae-tile-overlap-frames`. A size of 0 leaves that axis untiled, and setting both sizes to 0 +builds the untiled decoder. The shared `trtmc build --family ltx2` uses the defaults. + ## Build and run ```bash @@ -80,7 +103,9 @@ To give each rank its own TensorRT-RTX runtime cache, put `{rank}` in the path, - `tests/test_*_parity.py` build each engine from tiny random weights and compare it with diffusers. - `tests/test_context_parallel.py` runs the CP=2 DiT on two GPUs (torch-free ranks) against - the single-device plan and diffusers. It needs `TRTMC_NCCL_LIBRARY`. + the single-device plan and diffusers. It also runs the tile-parallel VAE decode and checks that + it matches the single-GPU tiled decode bit for bit. It needs `TRTMC_NCCL_LIBRARY`. +- `tests/test_vae_tiling.py` checks the tile plan (coverage, ramps, rank assignment) without a GPU. - `tests/test_e2e.py` builds the real checkpoint and runs the native CLI. It compares the output with `LTX2Pipeline` started from the same noise. Select it with `--e2e-model ltx2`. diff --git a/families/ltx2/audio_builder.py b/families/ltx2/audio_builder.py index cd1e85aa79..01960c9743 100644 --- a/families/ltx2/audio_builder.py +++ b/families/ltx2/audio_builder.py @@ -307,7 +307,8 @@ def add_audio_vae_decoder(g: Graph, ck: Checkpoint, cfg: dict, z): def build_audio_decoder_engine(model_dir: str | Path, *, audio_frames: int, debug_mel: bool = False, - verbose: bool = False, tf32: bool = True): + verbose: bool = False, tf32: bool = True, shapes: dict | None = None): + """Serialized plan; ``shapes`` (optional) receives the ``waveform`` output shape.""" model_dir = Path(model_dir) vae = Checkpoint(model_dir / "audio_vae") voc = Checkpoint(model_dir / "vocoder") @@ -332,6 +333,8 @@ def build_audio_decoder_engine(model_dir: str | Path, *, audio_frames: int, debu g.mark_output(mel, "mel", trt.float32) wave = _bwe_vocoder(g, voc, ocfg, mel, debug=debug_mel) g.mark_output(wave, "waveform", trt.float32) + if shapes is not None: + shapes["waveform"] = [int(v) for v in wave.shape] print(f"[ltx2] Building audio decoder engine (audio latents {audio_frames} -> mel {int(mel.shape[2])} frames -> " f"{int(wave.shape[2])} samples @ {ocfg['output_sampling_rate']} Hz) ...", file=sys.stderr) return build_plan(builder, network, label="audio decoder", tf32=tf32) diff --git a/families/ltx2/cli.json b/families/ltx2/cli.json new file mode 100644 index 0000000000..8c084036c4 --- /dev/null +++ b/families/ltx2/cli.json @@ -0,0 +1,28 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one LTX-2.5 text-to-audio-video 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": "precision", "flags": ["--precision"], "type": "string", "choices": ["bf16"], "default": "bf16"}, + {"name": "backend", "flags": ["--backend"], "type": "string", "choices": ["trt", "trt_rtx"], "default": "trt"}, + {"name": "image_height", "flags": ["--image-height"], "type": "int"}, + {"name": "image_width", "flags": ["--image-width"], "type": "int"}, + {"name": "video_num_frames", "flags": ["--video-num-frames"], "type": "int"}, + {"name": "max_sequence_length", "flags": ["--max-sequence-length"], "type": "int"}, + {"name": "context_parallel_size", "flags": ["--context-parallel-size"], "type": "int", "choices": [1, 2], "default": 1}, + {"name": "vae_tile_pixels", "flags": ["--vae-tile-pixels"], "type": "int", "default": 512, "help": "Video VAE tile size in pixels (0: untiled spatial axes)"}, + {"name": "vae_tile_overlap_pixels", "flags": ["--vae-tile-overlap-pixels"], "type": "int", "default": 64, "help": "Minimum spatial tile overlap in pixels"}, + {"name": "vae_tile_frames", "flags": ["--vae-tile-frames"], "type": "int", "default": 256, "help": "Video VAE tile length in frames (0: untiled time axis)"}, + {"name": "vae_tile_overlap_frames", "flags": ["--vae-tile-overlap-frames"], "type": "int", "default": 24, "help": "Minimum temporal tile overlap in frames"}, + {"name": "verbose", "flags": ["--verbose"], "type": "bool", "action": "store_true", "default": false} + ] + } + ] +} diff --git a/families/ltx2/cli.py b/families/ltx2/cli.py new file mode 100644 index 0000000000..0c6fa21600 --- /dev/null +++ b/families/ltx2/cli.py @@ -0,0 +1,89 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 build command (``trtmc ltx2 build``) and typed build inputs; importing this module is CPU-only.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.model_support import resolve_model + +from .vae_tiling import TileConfig + +TASK = "text_to_audio_video" + + +@dataclass(frozen=True) +class BuildRequest: + model_dir: Path + task: str = TASK + precision: str = "bf16" + backend: str = "trt" + max_sequence_length: int | None = None + image_height: int | None = None + image_width: int | None = None + video_num_frames: int | None = None + context_parallel_size: int = 1 + vae_tiles: TileConfig = field(default_factory=TileConfig) + verbose: bool = False + + def __post_init__(self) -> None: + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + self.vae_tiles.validate() + + +def coerce_request(request: object) -> BuildRequest: + """Accept the shared ``trtmc build`` request; reject the options LTX-2.5 cannot honour.""" + if isinstance(request, BuildRequest): + return request + if getattr(request, "dynamic_kv_cache", False): + raise NotImplementedError("ltx2 does not support dynamic_kv_cache") + if getattr(request, "tensor_parallel_size", 1) != 1: + raise NotImplementedError("ltx2 requires tensor_parallel_size=1 (it shards the video tokens)") + if getattr(request, "max_batch_size", 1) != 1: + raise NotImplementedError("ltx2 requires max_batch_size=1") + if getattr(request, "quantization", None) not in (None, "none"): + raise NotImplementedError("ltx2 does not support quantization") + if getattr(request, "fp32_layers", ()): + raise NotImplementedError("ltx2 does not support fp32_layers (its fp32 islands are fixed in the graph)") + return BuildRequest( + model_dir=Path(request.model_dir), task=request.task, precision=request.precision, + backend=getattr(request, "backend", "trt"), + max_sequence_length=getattr(request, "max_sequence_length", None), + image_height=getattr(request, "image_height", None), image_width=getattr(request, "image_width", None), + video_num_frames=getattr(request, "video_num_frames", None), + context_parallel_size=int(getattr(request, "context_parallel_size", 1)), + verbose=bool(getattr(request, "verbose", False))) + + +def build_bundle(request: BuildRequest, output: Path) -> None: + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + + +def build(*, model: str, output: Path, revision: str | None = None, precision: str = "bf16", + backend: str = "trt", image_height: int | None = None, image_width: int | None = None, + video_num_frames: int | None = None, max_sequence_length: int | None = None, + context_parallel_size: int = 1, vae_tile_pixels: int = 512, vae_tile_overlap_pixels: int = 64, + vae_tile_frames: int = 256, vae_tile_overlap_frames: int = 24, verbose: bool = False) -> int: + tiles = TileConfig(tile_pixels=vae_tile_pixels, overlap_pixels=vae_tile_overlap_pixels, + tile_frames=vae_tile_frames, overlap_frames=vae_tile_overlap_frames) + request = BuildRequest(model_dir=resolve_model(model, revision), precision=precision, backend=backend, + max_sequence_length=max_sequence_length, image_height=image_height, + image_width=image_width, video_num_frames=video_num_frames, + context_parallel_size=context_parallel_size, vae_tiles=tiles, verbose=verbose) + build_bundle(request, output) + return 0 diff --git a/families/ltx2/model.py b/families/ltx2/model.py index ce2f57ed2a..a5908a0541 100644 --- a/families/ltx2/model.py +++ b/families/ltx2/model.py @@ -8,7 +8,8 @@ - ``text_encoder.plan``: Gemma 4 text tower + ``LTX2TextConnectors`` - ``denoiser.plan``: the joint audio/video DiT, single device (``context_parallel_size=1``) or context parallel over the video tokens (``context_parallel_size=2``, one rank-dynamic plan) - - ``vae.plan``: the video VAE decoder + - ``vae.plan``: the video VAE decoder, by default one tile-shaped plan for the tiled decode + (``vae_tiling.py``; context-parallel ranks decode disjoint tiles, rank 0 blends them) - ``audio.plan``: the audio VAE decoder + vocoder with bandwidth extension (48 kHz stereo) - ``tokenizer.json`` and ``runtime.json`` @@ -24,13 +25,13 @@ from pathlib import Path from typing import TYPE_CHECKING +from .cli import TASK, BuildRequest, coerce_request from .parallel import ParallelConfig, validate_context_parallel_layout +from .vae_tiling import plan_tiles if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter -TASK = "text_to_audio_video" PIPELINE_CLASS = "LTX2Pipeline" # diffusers ``pipelines/ltx2/utils.py`` DISTILLED_SIGMA_VALUES: the distilled checkpoint's @@ -100,20 +101,11 @@ def _write_plan(writer: "BundleWriter", name: str, plan) -> None: section.write(memoryview(plan)) -def build(request: "BuildRequest", writer: "BundleWriter") -> None: - """Build one LTX-2.5 text-to-audio-video bundle.""" +def build(request: BuildRequest, writer: "BundleWriter") -> None: + """Build one LTX-2.5 text-to-audio-video bundle (``trtmc ltx2 build`` or the shared ``trtmc build``).""" + request = coerce_request(request) if request.task != TASK: raise ValueError(f"ltx2 supports only task={TASK}") - if request.dynamic_kv_cache: - raise NotImplementedError("ltx2 does not support dynamic_kv_cache") - if request.tensor_parallel_size != 1: - raise NotImplementedError("ltx2 requires tensor_parallel_size=1 (it shards the video tokens)") - if request.max_batch_size != 1: - raise NotImplementedError("ltx2 requires max_batch_size=1") - if request.quantization not in (None, "none"): - raise NotImplementedError("ltx2 does not support quantization") - if request.fp32_layers: - raise NotImplementedError("ltx2 does not support fp32_layers (its fp32 islands are fixed in the graph)") if request.precision != "bf16": raise ValueError("ltx2 builds bf16 engines (precision=bf16), the precision LTX-2.5 runs in") parallel = ParallelConfig(cp_size=int(request.context_parallel_size)) @@ -176,17 +168,23 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: _log(f"DiT engine built in {time.perf_counter() - started:.1f} s (cp={parallel.cp_size}, " f"{shape.video_tokens} video + {shape.audio_frames} audio tokens)") started = time.perf_counter() - _write_plan(writer, "vae.plan", build_vae_decoder_engine(model_dir / "vae", latent_frames=shape.latent_frames, - latent_height=shape.latent_height, - latent_width=shape.latent_width, - verbose=request.verbose)) - _log(f"video VAE engine built in {time.perf_counter() - started:.1f} s") + tiling = None + vae_grid = (shape.latent_frames, shape.latent_height, shape.latent_width) + if request.vae_tiles.enabled: + tiling = plan_tiles(*vae_grid, request.vae_tiles, world=parallel.world_size) + vae_grid = tuple(tiling["tile_latent"]) + _write_plan(writer, "vae.plan", build_vae_decoder_engine(model_dir / "vae", latent_frames=vae_grid[0], + latent_height=vae_grid[1], latent_width=vae_grid[2], + clamp_output=tiling is None, verbose=request.verbose)) + _log(f"video VAE engine built in {time.perf_counter() - started:.1f} s" + + (f" ({len(tiling['tiles'])} tiles of {vae_grid} latents)" if tiling else " (untiled)")) started = time.perf_counter() + audio_shapes: dict = {} _write_plan(writer, "audio.plan", build_audio_decoder_engine(model_dir, audio_frames=shape.audio_frames, - verbose=request.verbose)) + verbose=request.verbose, shapes=audio_shapes)) _log(f"audio decoder engine built in {time.perf_counter() - started:.1f} s") writer.add_bytes("tokenizer.json", (model_dir / "tokenizer" / "tokenizer.json").read_bytes()) - writer.add_json("runtime.json", { + runtime = { "video_frames": frames, "video_height": height, "video_width": width, @@ -205,4 +203,8 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: "audio_channels": int(vocoder_cfg.get("out_channels", 2)), "parallel_mode": parallel.mode, "parallel_size": parallel.world_size, - }) + "audio_waveform_shape": audio_shapes["waveform"], + } + if tiling is not None: + runtime["vae_tiling"] = tiling + writer.add_json("runtime.json", runtime) diff --git a/families/ltx2/runtime/distributed_runtime.cpp b/families/ltx2/runtime/distributed_runtime.cpp index e0a8bf07f3..47a8bee5c3 100644 --- a/families/ltx2/runtime/distributed_runtime.cpp +++ b/families/ltx2/runtime/distributed_runtime.cpp @@ -33,6 +33,11 @@ using NcclCommInitRankFn = NcclResult (*)(NcclComm*, int, NcclUniqueId, int); using NcclCommDestroyFn = NcclResult (*)(NcclComm); using NcclGetErrorStringFn = const char* (*)(NcclResult); using NcclGetVersionFn = NcclResult (*)(int*); +using NcclSendFn = NcclResult (*)(const void*, std::size_t, int, int, NcclComm, cudaStream_t); +using NcclRecvFn = NcclResult (*)(void*, std::size_t, int, int, NcclComm, cudaStream_t); +using NcclGroupFn = NcclResult (*)(); +using NcclCommAbortFn = NcclResult (*)(NcclComm); +constexpr int kNcclUint8 = 1; // ncclUint8: transfers are byte copies int require_env_int(const char* name) { const char* raw = std::getenv(name); @@ -64,7 +69,7 @@ std::filesystem::path rendezvous_path() { return path; } -class NcclRuntime { +class NcclRuntime final : public PeerChannel { public: NcclRuntime() { // NCCL is resolved at run time: TRTMC_NCCL_LIBRARY, else libnccl.so.2 @@ -81,6 +86,11 @@ class NcclRuntime { comm_init_rank_ = load("ncclCommInitRank"); comm_destroy_ = load("ncclCommDestroy"); get_error_string_ = load("ncclGetErrorString"); + send_ = load("ncclSend"); + recv_ = load("ncclRecv"); + group_start_ = load("ncclGroupStart"); + group_end_ = load("ncclGroupEnd"); + comm_abort_ = load("ncclCommAbort"); int version = 0; const auto get_version = reinterpret_cast(library_->find_symbol("ncclGetVersion")); @@ -90,15 +100,72 @@ class NcclRuntime { << std::endl; } - ~NcclRuntime() { + ~NcclRuntime() override { + if (token_ != nullptr) + cudaFree(token_); + if (stream_ != nullptr) + cudaStreamDestroy(stream_); if (comm_ != nullptr) { comm_destroy_(comm_); comm_ = nullptr; } } + void run(const std::vector& transfers, + std::chrono::milliseconds timeout) override { + if (comm_ == nullptr) + throw std::runtime_error("LTX-2.5 peer transfer: the NCCL communicator was aborted"); + if (stream_ == nullptr && + cudaStreamCreateWithFlags(&stream_, cudaStreamNonBlocking) != cudaSuccess) + throw std::runtime_error("LTX-2.5 peer transfer: cudaStreamCreate failed"); + enqueue(transfers); + // Poll instead of blocking, so a missing peer aborts the communicator rather than + // leaving NCCL kernels spinning on the GPU. + const auto deadline = std::chrono::steady_clock::now() + timeout; + for (;;) { + const auto status = cudaStreamQuery(stream_); + if (status == cudaSuccess) + return; + if (status != cudaErrorNotReady) + throw std::runtime_error(std::string("LTX-2.5 peer transfer failed: ") + + cudaGetErrorString(status)); + if (std::chrono::steady_clock::now() > deadline) { + comm_abort_(comm_); + comm_ = nullptr; + throw std::runtime_error( + "LTX-2.5 peer transfer timed out; the NCCL communicator was aborted"); + } + std::this_thread::sleep_for(std::chrono::microseconds(200)); + } + } + + // Rank 0 hears from every rank, then answers each: nobody leaves before all arrived. The + // token is 64 KiB: with the Windows NCCL build used for the RTX PRO 6000 host, point-to-point + // transfers below 32 KiB never complete (a 1-byte barrier hangs), larger ones do. + void barrier(std::chrono::milliseconds timeout) override { + constexpr std::size_t kToken = 64 * 1024; + if (token_ == nullptr && + cudaMalloc(&token_, kToken * static_cast(size_)) != cudaSuccess) + throw std::runtime_error("LTX-2.5 rank barrier: cudaMalloc failed"); + auto* token = static_cast(token_); + std::vector gather; + std::vector release; + for (int peer = 1; rank_ == 0 && peer < size_; ++peer) { + gather.push_back({peer, token + kToken * peer, kToken, false}); + release.push_back({peer, token + kToken * peer, kToken, true}); + } + if (rank_ != 0) { + gather.push_back({0, token, kToken, true}); + release.push_back({0, token, kToken, false}); + } + run(gather, timeout); + run(release, timeout); + } + void init(int size, int rank, const NcclUniqueId& id) { check(comm_init_rank_(&comm_, size, id, rank), "ncclCommInitRank"); + size_ = size; + rank_ = rank; } NcclUniqueId unique_id() { @@ -115,6 +182,20 @@ class NcclRuntime { return library_->require(symbol); } + void enqueue(const std::vector& transfers) { + check(group_start_(), "ncclGroupStart"); + for (const auto& t : transfers) { + const auto status = t.send + ? send_(t.device, t.bytes, kNcclUint8, t.peer, comm_, stream_) + : recv_(t.device, t.bytes, kNcclUint8, t.peer, comm_, stream_); + if (status != 0) { + (void)group_end_(); + check(status, t.send ? "ncclSend" : "ncclRecv"); + } + } + check(group_end_(), "ncclGroupEnd"); + } + void check(NcclResult result, const char* operation) const { if (result == 0) return; @@ -128,6 +209,15 @@ class NcclRuntime { NcclCommInitRankFn comm_init_rank_{nullptr}; NcclCommDestroyFn comm_destroy_{nullptr}; NcclGetErrorStringFn get_error_string_{nullptr}; + NcclSendFn send_{nullptr}; + NcclRecvFn recv_{nullptr}; + NcclGroupFn group_start_{nullptr}; + NcclGroupFn group_end_{nullptr}; + NcclCommAbortFn comm_abort_{nullptr}; + cudaStream_t stream_{nullptr}; + void* token_{nullptr}; + int size_{1}; + int rank_{0}; }; void write_unique_id(const std::filesystem::path& path, const NcclUniqueId& id) { @@ -216,6 +306,7 @@ DistributedRuntimeGroup initialize_parallel_group(int parallel_size) { std::filesystem::remove(path, ignored); } group.communicator = runtime->communicator(); + group.channel = runtime; group.owner = std::move(runtime); return group; } diff --git a/families/ltx2/runtime/distributed_runtime.h b/families/ltx2/runtime/distributed_runtime.h index 991d435129..212d5c510a 100644 --- a/families/ltx2/runtime/distributed_runtime.h +++ b/families/ltx2/runtime/distributed_runtime.h @@ -5,16 +5,41 @@ #pragma once +#include +#include #include +#include namespace trtmc::ltx2 { +// One point-to-point copy of a device buffer between this rank and `peer`. +struct PeerTransfer { + int peer{0}; + void* device{nullptr}; + std::size_t bytes{0}; + bool send{false}; +}; + +// Point-to-point transfers on the communicator the TensorRT engines use. Runs only between +// engine executions, so it never interleaves with an engine's collectives. +class PeerChannel { + public: + virtual ~PeerChannel() = default; + // Runs the transfers as one group and waits for them. When they do not finish within + // `timeout`, aborts the communicator (in-flight transfers exit) and throws. + virtual void run(const std::vector& transfers, + std::chrono::milliseconds timeout) = 0; + // Returns once every rank has called it (same timeout and abort behavior as run). + virtual void barrier(std::chrono::milliseconds timeout) = 0; +}; + struct DistributedRuntimeGroup { int world_size{1}; int rank{0}; int parallel_size{1}; void* communicator{nullptr}; std::shared_ptr owner; + std::shared_ptr channel; // null on a single device }; // Initialize the NCCL communicator consumed by TensorRT distributed layers. diff --git a/families/ltx2/runtime/pipeline.cpp b/families/ltx2/runtime/pipeline.cpp index e2f1631e21..b3955898cb 100644 --- a/families/ltx2/runtime/pipeline.cpp +++ b/families/ltx2/runtime/pipeline.cpp @@ -13,6 +13,8 @@ #include #include #include +#include +#include #include #include #include @@ -21,6 +23,7 @@ #include #include #include +#include #include #include @@ -116,6 +119,9 @@ std::string trim(const std::string& text) { // TRTMC_LTX2_INITIAL_LATENTS raw fp32 file: packed video [S, C] then audio [Sa, Ca] noise // (replaces the seeded noise, e.g. a reference pipeline's draw) // TRTMC_LTX2_DUMP_LATENTS raw fp32 file written with the final video then audio latents +// TRTMC_LTX2_DECODE_LATENTS raw fp32 file in the TRTMC_LTX2_DUMP_LATENTS layout; replaces the +// denoised latents before the decode (decoder checks, e.g. the +// single-GPU vs tile-parallel decode of identical latents) std::vector read_f32_file(const char* path) { std::ifstream input(path, std::ios::binary | std::ios::ate); if (!input) @@ -142,6 +148,20 @@ void maybe_dump(const std::vector& video, const std::vector& audio << " fp32) to " << path << "\n"; } +void replace_final_latents(std::vector& video, std::vector& audio) { + const char* path = std::getenv("TRTMC_LTX2_DECODE_LATENTS"); + if (path == nullptr || *path == '\0') + return; + const auto values = read_f32_file(path); + if (values.size() != video.size() + audio.size()) + throw std::runtime_error( + "TRTMC_LTX2_DECODE_LATENTS must hold the packed final video then audio latents"); + std::copy_n(values.begin(), video.size(), video.begin()); + std::copy_n(values.begin() + static_cast(video.size()), audio.size(), + audio.begin()); + std::cerr << "[ltx2] decoding the latents of " << path << "\n"; +} + const std::array& config_fields() { static const std::array fields{{ {"seed", internal::ConfigKind::I64, internal::ConfigValue{std::int64_t{0}}, @@ -150,6 +170,56 @@ const std::array& config_fields() { return fields; } +ltx2::VaeTilePlan parse_tile_plan(const nlohmann::json& doc) { + ltx2::VaeTilePlan plan; + plan.tile_latent = doc.at("tile_latent").get>(); + plan.tile_pixels = doc.at("tile_pixels").get>(); + for (const auto& item : doc.at("tiles")) { + ltx2::VaeTile tile; + tile.latent_start = item.at("latent_start").get>(); + tile.pixel_start = item.at("pixel_start").get>(); + tile.ramps = item.at("ramps").get, 3>>(); + tile.rank = item.at("rank").get(); + plan.tiles.push_back(tile); + } + return plan; +} + +// Device scratch for one decode (tiles a worker rank sends or rank 0 receives). +class DeviceBuffer { + public: + explicit DeviceBuffer(std::size_t bytes) { + if (bytes != 0 && cudaMalloc(&ptr_, bytes) != cudaSuccess) + throw std::runtime_error("LTX-2.5 tiled decode: cudaMalloc failed"); + } + ~DeviceBuffer() { + if (ptr_ != nullptr) + cudaFree(ptr_); + } + DeviceBuffer(const DeviceBuffer&) = delete; + DeviceBuffer& operator=(const DeviceBuffer&) = delete; + uint8_t* get() const { return static_cast(ptr_); } + + private: + void* ptr_{nullptr}; +}; + +void cuda_copy(void* dst, const void* src, std::size_t bytes, cudaMemcpyKind kind) { + const auto status = cudaMemcpy(dst, src, bytes, kind); + if (status != cudaSuccess) + throw std::runtime_error(std::string("LTX-2.5 tiled decode: cudaMemcpy failed: ") + + cudaGetErrorString(status)); +} + +std::size_t numel(const std::vector& shape) { + std::size_t count = 1; + for (const auto dim : shape) + count *= static_cast(dim); + return count; +} + +constexpr std::chrono::minutes kPeerTimeout{10}; + internal::AudioVideoResult worker_completion(const LTX2Options& options) { internal::AudioVideoResult result; result.video.frames.num_frames = 0; @@ -163,7 +233,7 @@ internal::AudioVideoResult worker_completion(const LTX2Options& options) { } // namespace -LTX2Options parse_ltx2_options(const std::string& runtime_json) { +LTX2Options parse_ltx2_options(const std::string& runtime_json, int32_t world_size) { const auto doc = nlohmann::json::parse(runtime_json); LTX2Options o; o.video_frames = doc.at("video_frames").get(); @@ -187,6 +257,16 @@ LTX2Options parse_ltx2_options(const std::string& runtime_json) { throw std::runtime_error("LTX-2.5 runtime runs the distilled (batch 1) denoiser"); if (o.video_tokens() <= 0 || o.audio_frames <= 0 || o.text_seq_len <= 0) throw std::runtime_error("LTX-2.5 runtime.json has invalid shapes"); + if (doc.contains("audio_waveform_shape")) + o.audio_waveform_shape = doc.at("audio_waveform_shape").get>(); + if (doc.contains("vae_tiling")) { + const auto& tiling = doc.at("vae_tiling"); + if (tiling.at("world_size").get() != world_size) + throw std::runtime_error("LTX-2.5 VAE tile plan was built for another world size"); + o.vae_tiling = parse_tile_plan(tiling); + ltx2::vae_validate_plan(o.vae_tiling, {o.latent_frames, o.latent_height, o.latent_width}, + {o.video_frames, o.video_height, o.video_width}, world_size); + } return o; } @@ -278,6 +358,179 @@ std::vector LTX2Pipeline::decode_audio(const std::vector& audio_la return float_output(outputs, "waveform", count); } +uint8_t* LTX2Pipeline::tile_host_buffer(std::size_t bytes) { + if (tile_host_bytes_ < bytes) { + tile_host_.reset(); + tile_host_bytes_ = 0; + void* ptr = nullptr; + if (cudaMallocHost(&ptr, bytes) != cudaSuccess) + throw std::runtime_error("LTX-2.5 tiled decode: cudaMallocHost failed"); + tile_host_ = std::shared_ptr(static_cast(ptr), + [](uint8_t* p) { cudaFreeHost(p); }); + tile_host_bytes_ = bytes; + } + return tile_host_.get(); +} + +LTX2Pipeline::Decoded LTX2Pipeline::decode_untiled(const std::vector& video, + const std::vector& audio) { + Decoded out; + auto start = Clock::now(); + out.frames = decode_video(video); + out.tiles_ms = elapsed_ms(start, Clock::now()); + start = Clock::now(); + out.wave = decode_audio(audio); + out.audio_ms = elapsed_ms(start, Clock::now()); + return out; +} + +// Decodes this rank's tiles in tile order: into host slot k on rank 0, else packed into the +// device send buffer. +void LTX2Pipeline::decode_own_tiles(const std::vector& video, uint8_t* host, void* device, + Decoded& out) { + const auto& plan = options_.vae_tiling; + const auto bytes = plan.tile_values() * sizeof(uint16_t); + const std::array latent{options_.latent_frames, options_.latent_height, + options_.latent_width}; + std::vector tile_latents; + std::size_t packed = 0; + const auto start = Clock::now(); + for (std::size_t k = 0; k < plan.tiles.size(); ++k) { + if (plan.tiles[k].rank != distributed_.rank) + continue; + const auto tile_start = Clock::now(); + ltx2::vae_gather_tile_latents(video, latent, options_.latent_channels, plan, plan.tiles[k], + tile_latents); + TensorMap inputs; + inputs["latents"] = + Tensor{tile_latents.data(), + {1, static_cast(tile_latents.size()) / options_.latent_channels, + options_.latent_channels}, + DType::kFloat32}; + vae_->forward_async(inputs); + vae_->sync(); + const void* frames = vae_->device_ptr("frames"); + if (host != nullptr) + cuda_copy(host + k * bytes, frames, bytes, cudaMemcpyDeviceToHost); + else + cuda_copy(static_cast(device) + (packed++) * bytes, frames, bytes, + cudaMemcpyDeviceToDevice); + ++out.tiles; + if (progress_.enabled()) { + std::ostringstream detail; + detail << "tile=" << k << " tile_ms=" << std::fixed << std::setprecision(3) + << elapsed_ms(tile_start, Clock::now()); + progress_.emit("vae_tile", detail.str()); + } + } + out.tiles_ms = elapsed_ms(start, Clock::now()); +} + +// Rank 0: receives every worker's tiles (in their tile order) and the audio rank's waveform. +void LTX2Pipeline::receive_peer_tiles(uint8_t* host, Decoded& out) { + const auto& plan = options_.vae_tiling; + const auto bytes = plan.tile_values() * sizeof(uint16_t); + const auto world = static_cast(distributed_.world_size); + const auto audio_rank = static_cast(options_.audio_rank(distributed_.world_size)); + const auto wave_count = numel(options_.audio_waveform_shape); + std::vector peer_bytes(world, 0); + for (const auto& tile : plan.tiles) + peer_bytes[static_cast(tile.rank)] += bytes; + if (audio_rank != 0) + peer_bytes[audio_rank] += wave_count * sizeof(float); + std::vector cursor(world, 0); + std::size_t total = 0; + for (std::size_t p = 1; p < world; ++p) { + cursor[p] = total; + total += peer_bytes[p]; + } + DeviceBuffer recv(total); + std::vector transfers; + for (std::size_t p = 1; p < world; ++p) { + if (peer_bytes[p] != 0) + transfers.push_back( + {static_cast(p), recv.get() + cursor[p], peer_bytes[p], false}); + } + const auto start = Clock::now(); + distributed_.channel->run(transfers, kPeerTimeout); + for (std::size_t k = 0; k < plan.tiles.size(); ++k) { + const auto rank = static_cast(plan.tiles[k].rank); + if (rank == 0) + continue; + cuda_copy(host + k * bytes, recv.get() + cursor[rank], bytes, cudaMemcpyDeviceToHost); + cursor[rank] += bytes; + } + if (audio_rank != 0) { + out.wave.resize(wave_count); + cuda_copy(out.wave.data(), recv.get() + cursor[audio_rank], wave_count * sizeof(float), + cudaMemcpyDeviceToHost); + } + out.exchange_ms = elapsed_ms(start, Clock::now()); +} + +LTX2Pipeline::Decoded LTX2Pipeline::decode_tiled(const std::vector& video, + const std::vector& audio) { + const auto& plan = options_.vae_tiling; + const auto bytes = plan.tile_values() * sizeof(uint16_t); + const int32_t audio_rank = options_.audio_rank(distributed_.world_size); + Decoded out; + if (distributed_.rank != 0) { + std::size_t mine = 0; + for (const auto& tile : plan.tiles) + mine += tile.rank == distributed_.rank ? 1U : 0U; + const auto wave_bytes = numel(options_.audio_waveform_shape) * sizeof(float); + const bool sends_audio = distributed_.rank == audio_rank; + DeviceBuffer send(mine * bytes + (sends_audio ? wave_bytes : 0)); + decode_own_tiles(video, nullptr, send.get(), out); + if (sends_audio) { + const auto start = Clock::now(); + (void)decode_audio(audio); + cuda_copy(send.get() + mine * bytes, audio_->device_ptr("waveform"), wave_bytes, + cudaMemcpyDeviceToDevice); + out.audio_ms = elapsed_ms(start, Clock::now()); + } + const auto start = Clock::now(); + const auto total = mine * bytes + (sends_audio ? wave_bytes : 0); + if (total != 0) // rank 0 posts a receive only for peers with data + distributed_.channel->run({{0, send.get(), total, true}}, kPeerTimeout); + out.exchange_ms = elapsed_ms(start, Clock::now()); + return out; + } + uint8_t* host = tile_host_buffer(plan.tiles.size() * bytes); + decode_own_tiles(video, host, nullptr, out); + if (distributed_.world_size > 1) + receive_peer_tiles(host, out); + std::vector tiles(plan.tiles.size()); + for (std::size_t k = 0; k < tiles.size(); ++k) + tiles[k] = reinterpret_cast(host + k * bytes); + // The host blend overlaps the audio decode when rank 0 decodes the audio. + std::exception_ptr blend_error; + std::thread blend([&] { + try { + const auto start = Clock::now(); + ltx2::vae_blend_tiles(plan, tiles, options_.video_frames, options_.video_height, + options_.video_width, out.frames); + out.blend_ms = elapsed_ms(start, Clock::now()); + } catch (...) { + blend_error = std::current_exception(); + } + }); + if (audio_rank == 0) { + try { + const auto start = Clock::now(); + out.wave = decode_audio(audio); + out.audio_ms = elapsed_ms(start, Clock::now()); + } catch (...) { + blend.join(); + throw; + } + } + blend.join(); + if (blend_error) + std::rethrow_exception(blend_error); + return out; +} + internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoRequest& request, internal::ConfigView config) { const auto& fields = config_fields(); @@ -291,6 +544,10 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq const auto audio_count = static_cast(options_.audio_frames) * options_.audio_latent_channels; + // Ranks finish loading their engines at different times (the ranks load different decoders); + // start together so the first collective does not charge one rank's load to the generation. + if (distributed_.channel) + distributed_.channel->barrier(kPeerTimeout); const auto t_start = Clock::now(); if (progress_.enabled()) { std::ostringstream detail; @@ -345,24 +602,35 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq } const auto t_denoise = Clock::now(); progress_.emit("denoise_end"); + replace_final_latents(video, audio); - if (distributed_.world_size > 1 && distributed_.rank != 0) { + const bool tiled = options_.vae_tiling.enabled(); + if (distributed_.world_size > 1 && distributed_.rank != 0 && !tiled) { std::cerr << "[ltx2] context-parallel rank " << distributed_.rank << " finished denoising in " << elapsed_ms(t_text, t_denoise) << " ms; rank 0 decodes the video and audio\n"; progress_.emit("worker_done"); return worker_completion(options_); } - maybe_dump(video, audio); - - progress_.emit("vae_begin"); - auto frames = decode_video(video); - const auto t_vae = Clock::now(); - progress_.emit("vae_end"); - progress_.emit("audio_begin"); - const auto wave = decode_audio(audio); + if (distributed_.rank == 0) + maybe_dump(video, audio); + + progress_.emit("decode_begin"); + auto decoded = tiled ? decode_tiled(video, audio) : decode_untiled(video, audio); const auto t_audio = Clock::now(); - progress_.emit("audio_end"); + progress_.emit("decode_end"); + if (distributed_.rank != 0) { + std::cerr << std::fixed << std::setprecision(3) + << "[ltx2-worker-perf-json] {\"rank\":" << distributed_.rank + << ",\"denoise_ms\":" << elapsed_ms(t_text, t_denoise) + << ",\"vae_tiles\":" << decoded.tiles << ",\"vae_tiles_ms\":" << decoded.tiles_ms + << ",\"audio_decode_ms\":" << decoded.audio_ms + << ",\"vae_send_ms\":" << decoded.exchange_ms << "}\n"; + progress_.emit("worker_done"); + return worker_completion(options_); + } + auto frames = std::move(decoded.frames); + const auto& wave = decoded.wave; internal::AudioVideoResult result; result.video.frames.pixels = std::move(frames); @@ -386,9 +654,17 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq << "[ltx2-perf-json] {\"world_size\":" << distributed_.world_size << ",\"text_encode_ms\":" << elapsed_ms(t_start, t_text) << ",\"denoise_ms\":" << elapsed_ms(t_text, t_denoise) - << ",\"median_step_ms\":" << median - << ",\"vae_decode_ms\":" << elapsed_ms(t_denoise, t_vae) - << ",\"audio_decode_ms\":" << elapsed_ms(t_vae, t_audio) + << ",\"median_step_ms\":" << median << ",\"decode_ms\":" + << elapsed_ms(t_denoise, t_audio) + // Untiled: the video then the audio decode. Tiled: the audio overlaps the blend + // (one device) or runs on the audio rank, so the video path spans the phase. + << ",\"vae_decode_ms\":" + << (tiled ? elapsed_ms(t_denoise, t_audio) : decoded.tiles_ms) + << ",\"audio_decode_ms\":" << decoded.audio_ms << ",\"vae_tiles\":" << decoded.tiles + << ",\"vae_tiles_ms\":" << decoded.tiles_ms + << ",\"vae_exchange_ms\":" << decoded.exchange_ms + << ",\"vae_blend_ms\":" << decoded.blend_ms + << ",\"audio_rank\":" << options_.audio_rank(distributed_.world_size) << ",\"generate_ms\":" << elapsed_ms(t_start, t_audio) << ",\"num_steps\":" << steps << "}\n"; if (progress_.enabled()) { diff --git a/families/ltx2/runtime/pipeline.h b/families/ltx2/runtime/pipeline.h index 9d05a54dab..24f24f33b5 100644 --- a/families/ltx2/runtime/pipeline.h +++ b/families/ltx2/runtime/pipeline.h @@ -9,11 +9,14 @@ // All model execution goes through TensorRT component engines: // text_encoder.plan Gemma 4 + LTX2TextConnectors -> video/audio text contexts // denoiser.plan joint audio/video DiT (single device or context parallel) -// vae.plan video VAE decoder -> RGB frames +// vae.plan video VAE decoder -> RGB frames (whole video, or one tile shape when the +// bundle carries a tile plan; context-parallel ranks decode disjoint tiles) // audio.plan audio VAE decoder + vocoder with BWE -> 48 kHz stereo +#include "families/ltx2/runtime/distributed_runtime.h" #include "families/ltx2/runtime/progress_log.h" #include "families/ltx2/runtime/tokenizer.h" +#include "families/ltx2/runtime/vae_tiling.h" #include "trtmc/internal/model.h" #include "trtmc/internal/video.h" #include "trtmc/runtime/trt_module.h" @@ -43,19 +46,32 @@ struct LTX2Options { std::vector sigmas; int32_t audio_sample_rate{48000}; int32_t audio_channels{2}; + // Tiled video decode (empty: vae.plan decodes the whole video on rank 0). + ltx2::VaeTilePlan vae_tiling; + // Waveform shape of audio.plan; lets another rank decode the audio for rank 0. + std::vector audio_waveform_shape; int64_t video_tokens() const { return int64_t(latent_frames) * latent_height * latent_width; } + // Rank that decodes the audio: the last context-parallel rank when the tiles spread the + // video decode over every rank, else rank 0. + int32_t audio_rank(int32_t world_size) const { + return world_size > 1 && vae_tiling.enabled() && !audio_waveform_shape.empty() + ? world_size - 1 + : 0; + } }; -LTX2Options parse_ltx2_options(const std::string& runtime_json); +LTX2Options parse_ltx2_options(const std::string& runtime_json, int32_t world_size = 1); // Context-parallel participation. The owner keeps the NCCL communicator used by the -// denoiser engine alive for the pipeline lifetime. Rank 0 decodes and returns media; -// other ranks return the worker completion. +// denoiser engine alive for the pipeline lifetime. Rank 0 returns the media; other ranks +// return the worker completion. With a tile plan every rank decodes its video tiles and sends +// them (and, on the audio rank, the waveform) to rank 0 over `channel`. struct LTX2DistributedContext { std::shared_ptr owner; int32_t rank{0}; int32_t world_size{1}; + std::shared_ptr channel; }; class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudioVideo { @@ -76,6 +92,17 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi std::vector audio; // [1, L, 2048] bf16 bits }; + // Decoded media (rank 0) and the decode phase timings of this rank. + struct Decoded { + std::vector frames; + std::vector wave; + double tiles_ms{0.0}; + double audio_ms{0.0}; + double exchange_ms{0.0}; + double blend_ms{0.0}; + int32_t tiles{0}; + }; + private: TextContext encode(const std::string& text); void run_dit(const std::vector& video, const std::vector& audio, @@ -83,6 +110,12 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi std::vector& audio_out); std::vector decode_video(const std::vector& video_latents); std::vector decode_audio(const std::vector& audio_latents); + Decoded decode_untiled(const std::vector& video, const std::vector& audio); + Decoded decode_tiled(const std::vector& video, const std::vector& audio); + void decode_own_tiles(const std::vector& video, uint8_t* host, void* device, + Decoded& out); + void receive_peer_tiles(uint8_t* host, Decoded& out); + uint8_t* tile_host_buffer(std::size_t bytes); // Declared first so the communicator outlives every engine that uses it. LTX2DistributedContext distributed_; @@ -93,6 +126,8 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi LTX2Options options_; std::shared_ptr tokenizer_; LTX2ProgressLog progress_; + std::shared_ptr tile_host_; // pinned decoded tiles on rank 0 + std::size_t tile_host_bytes_{0}; }; } // namespace trtmc diff --git a/families/ltx2/runtime/plugin.cpp b/families/ltx2/runtime/plugin.cpp index 190195bb1e..997847efbf 100644 --- a/families/ltx2/runtime/plugin.cpp +++ b/families/ltx2/runtime/plugin.cpp @@ -44,6 +44,19 @@ std::unique_ptr load(IBackend& backend, const BundleReader& bundle, return module; } +// The tile plan and vae.plan come from one build; reject a bundle whose plan shapes disagree. +void require_tile_engine(const ITrtModule& vae, const LTX2Options& options) { + const auto& plan = options.vae_tiling; + const std::vector latents{ + 1, int64_t(plan.tile_latent[0]) * plan.tile_latent[1] * plan.tile_latent[2], + options.latent_channels}; + const std::vector frames{plan.tile_pixels[0], plan.tile_pixels[1], plan.tile_pixels[2], + 3}; + if (vae.tensor_shape("latents") != latents || vae.tensor_shape("frames") != frames || + vae.tensor_dtype("frames") != DType::kFloat16) + throw std::runtime_error("LTX-2.5 vae.plan does not match the bundle's VAE tile plan"); +} + } // namespace } // namespace trtmc::ltx2 @@ -56,7 +69,7 @@ extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext& context const auto parallel = ltx2::parse_parallel_runtime_config(runtime); // Binds this rank's CUDA device before any engine is deserialized. const auto group = ltx2::initialize_parallel_group(parallel.size); - auto options = parse_ltx2_options(runtime); + auto options = parse_ltx2_options(runtime, group.world_size); ModuleCreateOptions plain{}; ModuleCreateOptions denoiser_options{}; if (parallel.distributed()) { @@ -65,19 +78,29 @@ extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext& context } auto text = ltx2::load(context.backend, context.reader, "text_encoder.plan", plain); auto denoiser = ltx2::load(context.backend, context.reader, "denoiser.plan", denoiser_options); - // Only rank 0 decodes and returns media; worker ranks never load the decoders. + // Rank 0 returns the media. With a tile plan every rank decodes video tiles and the audio + // rank decodes the audio; otherwise worker ranks never load the decoders. + const bool tiled = options.vae_tiling.enabled(); std::unique_ptr vae; std::unique_ptr audio; - if (group.rank == 0) { + if (group.rank == 0 || tiled) vae = ltx2::load(context.backend, context.reader, "vae.plan", plain); + if (group.rank == options.audio_rank(group.world_size)) audio = ltx2::load(context.backend, context.reader, "audio.plan", plain); - } + if (tiled) + ltx2::require_tile_engine(*vae, options); + // Rank 0 sizes the received waveform from runtime.json; it must match audio.plan. + if (audio && !options.audio_waveform_shape.empty() && + audio->tensor_shape("waveform") != options.audio_waveform_shape) + throw std::runtime_error( + "LTX-2.5 audio.plan does not match runtime.json audio_waveform_shape"); const auto tokenizer_data = ltx2::require_section(context.reader, "tokenizer.json"); std::shared_ptr tokenizer = CreateLtx2BpeTokenizer( tokenizer_data.data(), tokenizer_data.size(), /*add_special_tokens=*/false); if (!tokenizer) throw std::runtime_error("LTX-2.5 bundle tokenizer.json is not a supported BPE tokenizer"); - return new LTX2Pipeline(std::move(text), std::move(denoiser), std::move(vae), std::move(audio), - std::move(options), std::move(tokenizer), - LTX2DistributedContext{group.owner, group.rank, group.world_size}); + return new LTX2Pipeline( + std::move(text), std::move(denoiser), std::move(vae), std::move(audio), std::move(options), + std::move(tokenizer), + LTX2DistributedContext{group.owner, group.rank, group.world_size, group.channel}); } diff --git a/families/ltx2/runtime/vae_tiling.h b/families/ltx2/runtime/vae_tiling.h new file mode 100644 index 0000000000..7bbed2c5dc --- /dev/null +++ b/families/ltx2/runtime/vae_tiling.h @@ -0,0 +1,231 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +// Host side of the LTX-2.5 tiled video VAE decode (header-only so the contract tests compile it +// without engines). The build writes the tile plan (families/ltx2/vae_tiling.py) into +// runtime.json; every tile is decoded by one tile-shaped plan, and the tiles are blended here with +// linear ramps normalized by the summed weights. Each output value accumulates its tiles in tile +// order, so the result is the same bit for bit whichever rank decoded a tile. + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc::ltx2 { + +struct VaeTile { + std::array latent_start{}; // latent frame, row, column + std::array pixel_start{}; // output frame, row, column + // (left, right) ramp lengths in output frames / pixels for time, height, width. + std::array, 3> ramps{}; + int32_t rank{0}; +}; + +struct VaeTilePlan { + std::array tile_latent{}; // latent frames, rows, columns of every tile + std::array tile_pixels{}; // decoded frames, rows, columns of every tile + std::vector tiles; + + bool enabled() const { return !tiles.empty(); } + std::size_t tile_values() const { + return static_cast(tile_pixels[0]) * tile_pixels[1] * tile_pixels[2] * 3U; + } +}; + +// Ramp weights of one tile axis in fp32 (vae_tiling.py axis_weights). Spatial ramps fade in as +// k / (r + 1), k = 1..r; temporal left ramps fade in from 0 as k / r, k = 0..r-1; right ramps fade +// out as 1 - k / (r + 1), k = 1..r. +inline std::vector vae_axis_weights(int32_t length, int32_t left, int32_t right, + bool temporal) { + if (length <= 0 || left < 0 || right < 0 || left > length || right > length) + throw std::runtime_error("LTX-2.5 VAE tile ramp does not fit its tile"); + std::vector w(static_cast(length), 1.0F); + for (int32_t k = 0; k < left; ++k) { + w[static_cast(k)] = + temporal ? static_cast(k) / static_cast(left) + : static_cast(k + 1) / static_cast(left + 1); + } + for (int32_t k = 1; k <= right; ++k) { + w[static_cast(length - right + k - 1)] = + 1.0F - static_cast(k) / static_cast(right + 1); + } + return w; +} + +// Checks that every tile lies inside the latent grid and the decoded video. +inline void vae_validate_plan(const VaeTilePlan& plan, const std::array& latent, + const std::array& video, int32_t world_size) { + if (plan.tiles.empty()) + throw std::runtime_error("LTX-2.5 VAE tile plan has no tiles"); + for (const auto& tile : plan.tiles) { + if (tile.rank < 0 || tile.rank >= world_size) + throw std::runtime_error("LTX-2.5 VAE tile is assigned to a rank outside the world"); + for (std::size_t axis = 0; axis < 3; ++axis) { + if (tile.latent_start[axis] < 0 || + tile.latent_start[axis] + plan.tile_latent[axis] > latent[axis] || + tile.pixel_start[axis] < 0 || + tile.pixel_start[axis] + plan.tile_pixels[axis] > video[axis]) + throw std::runtime_error("LTX-2.5 VAE tile lies outside the video"); + (void)vae_axis_weights(plan.tile_pixels[axis], tile.ramps[axis][0], tile.ramps[axis][1], + axis == 0); + } + } +} + +// Copies one tile's packed latent tokens [tf * th * tw, C] out of the packed video +// [F * H * W, C] (token order frame, row, column). +inline void vae_gather_tile_latents(const std::vector& packed, + const std::array& latent, int32_t channels, + const VaeTilePlan& plan, const VaeTile& tile, + std::vector& out) { + const auto [tf, th, tw] = plan.tile_latent; + const auto row = static_cast(tw) * static_cast(channels); + out.resize(static_cast(tf) * th * row); + if (packed.size() != static_cast(latent[0]) * latent[1] * latent[2] * channels) + throw std::runtime_error("LTX-2.5 VAE tile gather: packed latents have the wrong size"); + std::size_t dst = 0; + for (int32_t f = 0; f < tf; ++f) { + for (int32_t h = 0; h < th; ++h) { + const auto token = (static_cast(tile.latent_start[0] + f) * latent[1] + + static_cast(tile.latent_start[1] + h)) * + latent[2] + + static_cast(tile.latent_start[2]); + std::memcpy(out.data() + dst, packed.data() + token * channels, row * sizeof(float)); + dst += row; + } + } +} + +inline float vae_half_to_float(uint16_t h) { + const uint32_t sign = (static_cast(h) & 0x8000U) << 16U; + const uint32_t exp = (h >> 10U) & 0x1FU; + uint32_t mant = h & 0x3FFU; + uint32_t bits = sign; + if (exp == 31U) { + bits |= 0x7F800000U | (mant << 13U); + } else if (exp != 0U) { + bits |= ((exp - 15U + 127U) << 23U) | (mant << 13U); + } else if (mant != 0U) { + int32_t e = -1; + do { + ++e; + mant <<= 1U; + } while ((mant & 0x400U) == 0U); + bits |= (static_cast(127 - 15 - e) << 23U) | ((mant & 0x3FFU) << 13U); + } + float out; + std::memcpy(&out, &bits, sizeof(out)); + return out; +} + +namespace detail { + +inline const std::vector& half_table() { + static const std::vector table = [] { + std::vector values(65536); + for (uint32_t i = 0; i < 65536U; ++i) + values[i] = vae_half_to_float(static_cast(i)); + return values; + }(); + return table; +} + +struct AxisWeights { + std::vector t, y, x; +}; + +// Accumulates one tile into the frame scratch and returns false when the tile misses the frame. +inline bool accumulate_tile(const VaeTilePlan& plan, const VaeTile& tile, const AxisWeights& w, + const uint16_t* values, int32_t frame, int32_t width, + std::vector& num, std::vector& den) { + const auto [tt, th, tw] = plan.tile_pixels; + const int32_t local_t = frame - tile.pixel_start[0]; + if (local_t < 0 || local_t >= tt) + return false; + const float wt = w.t[static_cast(local_t)]; + if (wt == 0.0F) + return true; + const auto& lut = half_table(); + const uint16_t* src = values + static_cast(local_t) * th * tw * 3U; + for (int32_t y = 0; y < th; ++y) { + const float wty = wt * w.y[static_cast(y)]; + const auto out_row = static_cast(tile.pixel_start[1] + y) * width + + static_cast(tile.pixel_start[2]); + float* n = num.data() + out_row * 3U; + float* d = den.data() + out_row; + for (int32_t x = 0; x < tw; ++x) { + const float wxy = wty * w.x[static_cast(x)]; + n[3 * x + 0] += wxy * lut[src[0]]; + n[3 * x + 1] += wxy * lut[src[1]]; + n[3 * x + 2] += wxy * lut[src[2]]; + d[x] += wxy; + src += 3; + } + } + return true; +} + +} // namespace detail + +// Blends decoded tiles (tiles[k]: fp16 [T, H, W, 3] of plan.tiles[k]) into clamped fp32 +// [frames, height, width, 3]. Output frames are independent, so threads split the frames; every +// value accumulates its tiles in tile order whatever the thread count. +inline void vae_blend_tiles(const VaeTilePlan& plan, const std::vector& tiles, + int32_t frames, int32_t height, int32_t width, std::vector& out, + unsigned threads = 0) { + if (tiles.size() != plan.tiles.size()) + throw std::runtime_error("LTX-2.5 VAE blend: one decoded buffer per tile is required"); + std::vector weights(plan.tiles.size()); + for (std::size_t k = 0; k < plan.tiles.size(); ++k) { + const auto& r = plan.tiles[k].ramps; + weights[k].t = vae_axis_weights(plan.tile_pixels[0], r[0][0], r[0][1], true); + weights[k].y = vae_axis_weights(plan.tile_pixels[1], r[1][0], r[1][1], false); + weights[k].x = vae_axis_weights(plan.tile_pixels[2], r[2][0], r[2][1], false); + } + const auto plane = static_cast(height) * width; + out.resize(static_cast(frames) * plane * 3U); + if (threads == 0) + threads = std::max(1U, std::min(32U, std::thread::hardware_concurrency())); + threads = std::min(threads, static_cast(std::max(frames, 1))); + auto work = [&](int32_t first, int32_t last) { + std::vector num(plane * 3U); + std::vector den(plane); + for (int32_t f = first; f < last; ++f) { + std::fill(num.begin(), num.end(), 0.0F); + std::fill(den.begin(), den.end(), 0.0F); + for (std::size_t k = 0; k < plan.tiles.size(); ++k) + detail::accumulate_tile(plan, plan.tiles[k], weights[k], tiles[k], f, width, num, + den); + float* dst = out.data() + static_cast(f) * plane * 3U; + for (std::size_t p = 0; p < plane; ++p) { + for (std::size_t c = 0; c < 3U; ++c) { + const float v = num[p * 3U + c] / den[p]; + dst[p * 3U + c] = std::min(1.0F, std::max(0.0F, v)); + } + } + } + }; + std::vector pool; + const int32_t chunk = + (frames + static_cast(threads) - 1) / static_cast(threads); + for (unsigned i = 1; i < threads; ++i) { + const int32_t first = static_cast(i) * chunk; + if (first >= frames) + break; + pool.emplace_back(work, first, std::min(frames, first + chunk)); + } + work(0, std::min(frames, chunk)); + for (auto& thread : pool) + thread.join(); +} + +} // namespace trtmc::ltx2 diff --git a/families/ltx2/tests/cpp/test_runtime_contract.cpp b/families/ltx2/tests/cpp/test_runtime_contract.cpp index 9cb47556b3..27ff311aa8 100644 --- a/families/ltx2/tests/cpp/test_runtime_contract.cpp +++ b/families/ltx2/tests/cpp/test_runtime_contract.cpp @@ -3,14 +3,18 @@ * SPDX-License-Identifier: Apache-2.0 */ -// LTX-2.5 host-side runtime contract: scheduler step, prompt padding, audio interleave. +// LTX-2.5 host-side runtime contract: scheduler step, prompt padding, audio interleave, tiled +// VAE ramps, tile latent gather and blend. #include "families/ltx2/runtime/progress_log.h" #include "families/ltx2/runtime/runtime_math.h" +#include "families/ltx2/runtime/vae_tiling.h" #include +#include #include #include +#include #include #include @@ -66,6 +70,116 @@ void test_progress_line_format() { "progress line keeps the LTX format"); } +uint16_t to_half(float value) { + // Exact for the small dyadic test values used below. + const auto bits = [&] { + uint32_t b; + std::memcpy(&b, &value, sizeof(b)); + return b; + }(); + const uint32_t sign = (bits >> 16U) & 0x8000U; + const int32_t exp = static_cast((bits >> 23U) & 0xFFU) - 127 + 15; + if ((bits & 0x7FFFFFFFU) == 0U) + return static_cast(sign); + return static_cast(sign | (static_cast(exp) << 10U) | + ((bits >> 13U) & 0x3FFU)); +} + +void test_vae_axis_weights() { + using trtmc::ltx2::vae_axis_weights; + const auto spatial = vae_axis_weights(6, 2, 2, false); + check(spatial[0] == 1.0F / 3.0F && spatial[1] == 2.0F / 3.0F && spatial[2] == 1.0F, + "spatial ramp fades in as k / (r + 1)"); + check(spatial[4] == 1.0F - 1.0F / 3.0F && spatial[5] == 1.0F - 2.0F / 3.0F, + "spatial ramp fades out as 1 - k / (r + 1)"); + const auto temporal = vae_axis_weights(5, 3, 0, true); + check(temporal[0] == 0.0F && temporal[1] == 1.0F / 3.0F && temporal[3] == 1.0F, + "temporal ramp fades in from zero"); + bool threw = false; + try { + (void)vae_axis_weights(2, 3, 0, false); + } catch (const std::runtime_error&) { + threw = true; + } + check(threw, "ramp longer than its tile is rejected"); +} + +// Two tiles along the width of a 1-frame, 1x6 video, overlapping by 2 pixels. +trtmc::ltx2::VaeTilePlan two_tile_plan() { + trtmc::ltx2::VaeTilePlan plan; + plan.tile_latent = {1, 1, 4}; + plan.tile_pixels = {1, 1, 4}; + plan.tiles.push_back({{0, 0, 0}, {0, 0, 0}, {{{0, 0}, {0, 0}, {0, 2}}}, 0}); + plan.tiles.push_back({{0, 0, 2}, {0, 0, 2}, {{{0, 0}, {0, 0}, {2, 0}}}, 1}); + return plan; +} + +void test_vae_blend_normalizes_overlaps() { + const auto plan = two_tile_plan(); + std::vector left(12, to_half(0.25F)); + std::vector right(12, to_half(0.75F)); + std::vector out; + trtmc::ltx2::vae_blend_tiles(plan, {left.data(), right.data()}, 1, 1, 6, out, 1); + check(out.size() == 18U, "blend output covers the video"); + check(out[0] == 0.25F && out[3 * 5] == 0.75F, "unshared pixels keep their tile"); + const float w = 1.0F / 3.0F; + const float expected = ((1.0F - w) * 0.25F + w * 0.75F) / ((1.0F - w) + w); + check(out[3 * 2] == expected, "overlap is the weight-normalized blend"); + std::vector big(12, to_half(2.0F)); + trtmc::ltx2::vae_blend_tiles(plan, {big.data(), big.data()}, 1, 1, 6, out, 1); + check(out[3 * 3] == 1.0F, "blended values are clamped to [0, 1]"); +} + +void test_vae_blend_is_thread_count_invariant() { + trtmc::ltx2::VaeTilePlan plan; + plan.tile_latent = {2, 1, 2}; + plan.tile_pixels = {9, 2, 3}; + plan.tiles.push_back({{0, 0, 0}, {0, 0, 0}, {{{0, 0}, {0, 0}, {0, 2}}}, 0}); + plan.tiles.push_back({{0, 0, 1}, {0, 0, 1}, {{{0, 0}, {0, 0}, {2, 0}}}, 1}); + std::vector a(9 * 2 * 3 * 3); + std::vector b(a.size()); + for (std::size_t i = 0; i < a.size(); ++i) { + a[i] = to_half(static_cast(i % 7) / 8.0F); + b[i] = to_half(static_cast(i % 5) / 8.0F); + } + std::vector one; + std::vector many; + trtmc::ltx2::vae_blend_tiles(plan, {a.data(), b.data()}, 9, 2, 4, one, 1); + trtmc::ltx2::vae_blend_tiles(plan, {a.data(), b.data()}, 9, 2, 4, many, 4); + check(one == many, "blend is bit-identical for any thread count"); +} + +void test_vae_tile_latents_and_validation() { + // Packed [F=2, H=2, W=3, C=2] latents; tile of 2x1x2 latents at (0, 1, 1). + std::vector packed(2 * 2 * 3 * 2); + for (std::size_t i = 0; i < packed.size(); ++i) + packed[i] = static_cast(i); + trtmc::ltx2::VaeTilePlan plan; + plan.tile_latent = {2, 1, 2}; + plan.tile_pixels = {9, 32, 64}; + plan.tiles.push_back({{0, 1, 1}, {0, 32, 32}, {{{0, 0}, {0, 0}, {0, 0}}}, 0}); + std::vector tile; + trtmc::ltx2::vae_gather_tile_latents(packed, {2, 2, 3}, 2, plan, plan.tiles[0], tile); + // Tokens (f, h, w) = (0,1,1), (0,1,2), (1,1,1), (1,1,2) -> token ids 4, 5, 10, 11. + check((tile == std::vector{8, 9, 10, 11, 20, 21, 22, 23}), "tile latent gather"); + trtmc::ltx2::vae_validate_plan(plan, {2, 2, 3}, {9, 64, 96}, 1); + bool threw = false; + try { + trtmc::ltx2::vae_validate_plan(plan, {2, 2, 3}, {9, 64, 64}, 1); + } catch (const std::runtime_error&) { + threw = true; + } + check(threw, "tile outside the video is rejected"); + threw = false; + try { + plan.tiles[0].rank = 1; + trtmc::ltx2::vae_validate_plan(plan, {2, 2, 3}, {9, 64, 96}, 1); + } catch (const std::runtime_error&) { + threw = true; + } + check(threw, "tile rank outside the world is rejected"); +} + } // namespace int main() { @@ -73,6 +187,10 @@ int main() { test_prompt_ids_left_pad_and_truncate(); test_interleave_stereo(); test_progress_line_format(); + test_vae_axis_weights(); + test_vae_blend_normalizes_overlaps(); + test_vae_blend_is_thread_count_invariant(); + test_vae_tile_latents_and_validation(); if (failures != 0) return EXIT_FAILURE; std::puts("ltx2 runtime contract: OK"); diff --git a/families/ltx2/tests/dist_helpers.py b/families/ltx2/tests/dist_helpers.py index 9320e6c9d5..65f99928dc 100644 --- a/families/ltx2/tests/dist_helpers.py +++ b/families/ltx2/tests/dist_helpers.py @@ -32,6 +32,9 @@ def __init__(self): self.lib.ncclCommInitRank.argtypes = [ctypes.POINTER(ctypes.c_void_p), ctypes.c_int, _UniqueId, ctypes.c_int] self.lib.ncclCommAbort.argtypes = [ctypes.c_void_p] self.lib.ncclCommDestroy.argtypes = [ctypes.c_void_p] + p2p_args = [ctypes.c_void_p, ctypes.c_size_t, ctypes.c_int, ctypes.c_int, ctypes.c_void_p, ctypes.c_void_p] + self.lib.ncclSend.argtypes = p2p_args + self.lib.ncclRecv.argtypes = p2p_args uid = _UniqueId() path = Path(os.environ["TRTMC_NCCL_RENDEZVOUS"]) if self.rank == 0: @@ -65,6 +68,27 @@ def capsule(self): new.argtypes = [ctypes.c_void_p, ctypes.c_char_p, ctypes.c_void_p] return new(self.comm.value, None, None) + def p2p(self, transfers, stream: int, *, timeout_s: float = 120.0) -> None: + """One NCCL group of ``(send, device_ptr, nbytes, peer)`` byte transfers on ``stream``. + + Polls the stream with a deadline and aborts the communicator on a timeout, like the + runtime's ``PeerChannel``. + """ + from cuda.bindings import runtime as rt + + self._check(self.lib.ncclGroupStart(), "ncclGroupStart") + for send, ptr, nbytes, peer in transfers: + op = self.lib.ncclSend if send else self.lib.ncclRecv + self._check(op(ctypes.c_void_p(ptr), nbytes, 1, peer, self.comm, ctypes.c_void_p(stream)), + "ncclSend" if send else "ncclRecv") + self._check(self.lib.ncclGroupEnd(), "ncclGroupEnd") + deadline = time.monotonic() + timeout_s + while int(rt.cudaStreamQuery(stream)[0]) != 0: + if time.monotonic() > deadline: + self.abort() + raise TimeoutError(f"NCCL transfers did not finish within {timeout_s} s; communicator aborted") + time.sleep(0.002) + def abort(self) -> None: """``ncclCommAbort``: makes in-flight NCCL kernels exit (use instead of killing a hung rank).""" if self.comm: diff --git a/families/ltx2/tests/dist_vae_tile_check.py b/families/ltx2/tests/dist_vae_tile_check.py new file mode 100644 index 0000000000..aaa2e47391 --- /dev/null +++ b/families/ltx2/tests/dist_vae_tile_check.py @@ -0,0 +1,109 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Multi-rank check of the tile-parallel LTX-2.5 video VAE decode (tiny random weights), torch-free. + +``tools/launch_ranks.py -n WORLD --gpus ... -- python -m families.ltx2.tests.dist_vae_tile_check PREP_DIR`` +after ``vae_tile_prep`` wrote the tile plan, the tile-shaped plan and the reference into ``PREP_DIR``. + +Mirrors the runtime protocol: every rank decodes its tiles of the plan; the worker ranks pack their +fp16 tiles into one device buffer and send it to rank 0 (NCCL point-to-point on the communicator, +polled with a deadline that aborts the communicator); rank 0 blends every tile in tile order. Rank 0 +also decodes every tile itself and checks that the tile-parallel blend equals that single-rank tiled +decode bit for bit, and that it matches the blend of diffusers' tiles. Writes ``vae_rank.json``. +""" + +from __future__ import annotations + +import json +import os +import sys +from pathlib import Path + +from families.ltx2.tests import conftest # noqa: F401 - binds the TensorRT backend + + +def _tile_latents(latents, grid, plan, tile): + f, h, w = grid + tf, th, tw = plan["tile_latent"] + f0, h0, w0 = tile["latent_start"] + part = latents.reshape(1, f, h, w, -1)[:, f0:f0 + tf, h0:h0 + th, w0:w0 + tw] + return part.reshape(1, tf * th * tw, -1) + + +def main() -> int: + import numpy as np + from cuda.bindings import runtime as rt + + from families.ltx2.tests.dist_helpers import NcclComm + from families.ltx2.tests.np_engine import NpEngine, ck, cosine + from families.ltx2.vae_tiling import blend_tiles + + ck(rt.cudaSetDevice(int(os.environ.get("OMPI_COMM_WORLD_LOCAL_RANK", "0")))) + ck(rt.cudaFree(0)) + prep = Path(sys.argv[1]) + comm = NcclComm() + spec = json.loads((prep / "plan.json").read_text(encoding="utf-8")) + grid, plan = spec["grid"], spec["plan"] + latents = np.load(prep / "latents.npy") + engine = NpEngine((prep / "tile.plan").read_bytes()) + out_ptr, out_shape, _ = engine.buffers["frames"] + tile_bytes = int(np.prod(out_shape)) * 2 + stream = int(ck(rt.cudaStreamCreate())) + tiles = plan["tiles"] + mine = [k for k, t in enumerate(tiles) if t["rank"] == comm.rank] + report = {"rank": comm.rank, "world": comm.world, "tiles": mine} + ok = True + try: + if comm.rank != 0: + send = int(ck(rt.cudaMalloc(max(len(mine) * tile_bytes, 1)))) + for i, k in enumerate(mine): + engine({"latents": _tile_latents(latents, grid, plan, tiles[k])}) + ck(rt.cudaMemcpy(send + i * tile_bytes, out_ptr, tile_bytes, + rt.cudaMemcpyKind.cudaMemcpyDeviceToDevice)) + if mine: + comm.p2p([(True, send, len(mine) * tile_bytes, 0)], stream) + else: + decoded: list = [None] * len(tiles) + for k in mine: + decoded[k] = engine({"latents": _tile_latents(latents, grid, plan, tiles[k])})["frames"] + peer_tiles = {p: [k for k, t in enumerate(tiles) if t["rank"] == p] for p in range(1, comm.world)} + recv = {p: int(ck(rt.cudaMalloc(max(len(ks) * tile_bytes, 1)))) for p, ks in peer_tiles.items()} + comm.p2p([(False, recv[p], len(ks) * tile_bytes, p) for p, ks in peer_tiles.items() if ks], stream) + peer_identical = True + for p, ks in peer_tiles.items(): + for i, k in enumerate(ks): + host = np.empty(out_shape, np.float16) + ck(rt.cudaMemcpy(host.ctypes.data, recv[p] + i * tile_bytes, tile_bytes, + rt.cudaMemcpyKind.cudaMemcpyDeviceToHost)) + decoded[k] = host + own = engine({"latents": _tile_latents(latents, grid, plan, tiles[k])})["frames"] + peer_identical &= bool(np.array_equal(own.astype(np.float16), host)) + single = [engine({"latents": _tile_latents(latents, grid, plan, t)})["frames"].astype(np.float16) + for t in tiles] + f, h, w = grid + shape = ((f - 1) * 8 + 1, h * 32, w * 32) + parallel = blend_tiles(plan, [np.asarray(d, np.float16) for d in decoded], *shape) + serial = blend_tiles(plan, single, *shape) + reference = np.load(prep / "reference.npy") + report.update( + peer_tiles_identical_to_rank0_decode=peer_identical, + blend_bit_identical_to_single_rank=bool(np.array_equal(parallel, serial)), + cos_vs_diffusers_tiles=cosine(parallel - 0.5, reference - 0.5), + finite=bool(np.isfinite(parallel).all()), + ) + ok = (report["blend_bit_identical_to_single_rank"] and report["finite"] + and report["cos_vs_diffusers_tiles"] > 0.999) + except TimeoutError as exc: + report["error"] = str(exc) + ok = False + report["ok"] = ok + print(f"[rank {comm.rank}] {json.dumps(report)}", flush=True) + (prep / f"vae_rank{comm.rank}.json").write_text(json.dumps(report, indent=1), encoding="utf-8") + if comm.comm: + comm.destroy() + return 0 if ok else 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/families/ltx2/tests/test_context_parallel.py b/families/ltx2/tests/test_context_parallel.py index f452467de1..042f1c048c 100644 --- a/families/ltx2/tests/test_context_parallel.py +++ b/families/ltx2/tests/test_context_parallel.py @@ -67,3 +67,38 @@ def test_cp2_tiny_parity_two_ranks(tmp_path: Path) -> None: for rank in range(2): report = json.loads((tmp_path / f"cp_rank{rank}.json").read_text(encoding="utf-8")) assert all(case["ok"] for case in report["cases"].values()) + + +def test_tile_parallel_vae_two_ranks(tmp_path: Path) -> None: + """2-rank tile-parallel VAE decode vs the single-rank tiled decode (bit-exact) and diffusers. + + Rank 1 sends its decoded tiles to rank 0 over NCCL point-to-point, as the runtime does; the + ranks are torch-free and abort the communicator instead of hanging on a missing peer. + """ + torch = pytest.importorskip("torch") + pytest.importorskip("tensorrt") + pytest.importorskip("diffusers") + pytest.importorskip("cuda.bindings") + if not torch.cuda.is_available() or torch.cuda.device_count() < 2: + pytest.skip("two CUDA devices are required") + nccl = os.environ.get("TRTMC_NCCL_LIBRARY") + if not nccl: + pytest.skip("TRTMC_NCCL_LIBRARY must point at the NCCL library") + env = dict(os.environ) + env.pop("CUDA_VISIBLE_DEVICES", None) + env["PYTHONPATH"] = os.pathsep.join([str(REPO), str(REPO / "core" / "builder"), env.get("PYTHONPATH", "")]) + prep = subprocess.run([sys.executable, "-m", "families.ltx2.tests.vae_tile_prep", str(tmp_path), "2"], + cwd=REPO, env=env, capture_output=True, text=True, timeout=900) + print(prep.stdout[-2000:], prep.stderr[-2000:]) + assert prep.returncode == 0 + cmd = [sys.executable, str(REPO / "tools" / "launch_ranks.py"), "-n", "2", "--gpus", "0,1", + "--nccl-library", nccl, "--timeout", "600", "--", + sys.executable, "-m", "families.ltx2.tests.dist_vae_tile_check", str(tmp_path)] + proc = subprocess.run(cmd, cwd=REPO, env=env, capture_output=True, text=True, timeout=900) + print(proc.stdout[-6000:]) + print(proc.stderr[-3000:]) + assert proc.returncode == 0 + reports = [json.loads((tmp_path / f"vae_rank{rank}.json").read_text(encoding="utf-8")) for rank in range(2)] + assert all(report["ok"] for report in reports) + assert reports[0]["blend_bit_identical_to_single_rank"] + assert reports[0]["tiles"] and reports[1]["tiles"] diff --git a/families/ltx2/tests/test_vae_parity.py b/families/ltx2/tests/test_vae_parity.py index 701850279d..af5d9256e2 100644 --- a/families/ltx2/tests/test_vae_parity.py +++ b/families/ltx2/tests/test_vae_parity.py @@ -51,10 +51,13 @@ F, H, W = 3, 2, 3 -def test_vae_decoder_tiny_parity(tmp_path) -> None: - from diffusers import AutoencoderKLLTX2Video +TILE_GRID = (7, 6, 7) # latent frames, rows, columns: 2 x 1 x 2 tiles of 5 x 6 x 5 latents +TILE_CONFIG = dict(tile_pixels=192, overlap_pixels=64, tile_frames=40, overlap_frames=8) - from families.ltx2.vae_builder import build_vae_decoder_engine + +def tiny_vae(folder): + """Tiny random VAE saved under ``folder``; returns it and its (advanced) generator.""" + from diffusers import AutoencoderKLLTX2Video vae = AutoencoderKLLTX2Video(**TINY_VAE).eval() gen = torch.Generator().manual_seed(9) @@ -67,12 +70,74 @@ def test_vae_decoder_tiny_parity(tmp_path) -> None: else 0.05 * torch.randn(p.shape, generator=gen)) vae.latents_mean.copy_(0.3 * torch.randn(16, generator=gen)) vae.latents_std.copy_(0.5 + torch.rand(16, generator=gen)) - folder = tmp_path / "vae" - folder.mkdir() + folder.mkdir(parents=True, exist_ok=True) state = {k: (v.to(torch.bfloat16) if k not in ("latents_mean", "latents_std") else v).contiguous() for k, v in vae.state_dict().items()} safetensors_torch.save_file(state, str(folder / "diffusion_pytorch_model.safetensors")) (folder / "config.json").write_text(json.dumps(TINY_VAE), encoding="utf-8") + return vae, gen + + +def diffusers_tile(vae, packed, grid, plan, tile, dtype=torch.float32): + """diffusers decode of one tile's latents, ``(x + 1) / 2`` unclamped: ``[T, H, W, 3]``.""" + f, h, w = grid + tf, th, tw = plan["tile_latent"] + f0, h0, w0 = tile["latent_start"] + z = packed.reshape(1, f, h, w, 16)[:, f0:f0 + tf, h0:h0 + th, w0:w0 + tw].permute(0, 4, 1, 2, 3).cuda() + z = z * vae.latents_std.view(1, -1, 1, 1, 1).float() + vae.latents_mean.view(1, -1, 1, 1, 1).float() + with torch.no_grad(): + video = vae.decode(z.to(dtype), return_dict=False)[0].float() + return (video / 2 + 0.5)[0].permute(1, 2, 3, 0).cpu() + + +def tile_latents(packed, grid, plan, tile): + """Packed ``[1, tf*th*tw, C]`` latents of one tile.""" + f, h, w = grid + tf, th, tw = plan["tile_latent"] + f0, h0, w0 = tile["latent_start"] + part = packed.reshape(1, f, h, w, -1)[:, f0:f0 + tf, h0:h0 + th, w0:w0 + tw] + return part.reshape(1, tf * th * tw, -1) + + +def test_vae_tile_engine_tiny_parity(tmp_path) -> None: + """Every tile of a tile plan vs diffusers on the same latents, and the blended video.""" + import numpy as np + + from families.ltx2.tests.engine_runner import Engine + from families.ltx2.vae_builder import build_vae_decoder_engine + from families.ltx2.vae_tiling import TileConfig, blend_tiles, plan_tiles + + vae, _ = tiny_vae(tmp_path / "vae") + ref_vae = vae.to("cuda", torch.float32) + f, h, w = TILE_GRID + plan = plan_tiles(f, h, w, TileConfig(**TILE_CONFIG)) + tf, th, tw = plan["tile_latent"] + assert len(plan["tiles"]) == 4 and plan["tile_latent"] == [5, 6, 5] + engine = Engine(build_vae_decoder_engine(tmp_path / "vae", latent_frames=tf, latent_height=th, latent_width=tw, + clamp_output=False)) + packed = torch.randn(1, f * h * w, 16, generator=torch.Generator().manual_seed(4)) + got, ref = [], [] + for tile in plan["tiles"]: + out = engine({"latents": tile_latents(packed, TILE_GRID, plan, tile)})["frames"].float().cpu() + expected = diffusers_tile(ref_vae, packed, TILE_GRID, plan, tile) + c = cosine(out - 0.5, expected - 0.5) + print(f"tile {tile['latent_start']}: cos(centered) {c:.6f}") + assert c > 0.999 + got.append(out.numpy().astype(np.float16)) + ref.append(expected.numpy()) + frames, height, width = (f - 1) * 8 + 1, h * 32, w * 32 + blended = blend_tiles(plan, got, frames, height, width) + expected = blend_tiles(plan, ref, frames, height, width) + c = cosine(torch.from_numpy(blended) - 0.5, torch.from_numpy(expected) - 0.5) + print(f"blended tiles vs blended diffusers tiles: cos(centered) {c:.6f}") + assert c > 0.999 + + +def test_vae_decoder_tiny_parity(tmp_path) -> None: + from families.ltx2.vae_builder import build_vae_decoder_engine + + vae, gen = tiny_vae(tmp_path / "vae") + folder = tmp_path / "vae" packed = torch.randn(1, F * H * W, 16, generator=gen) plan = build_vae_decoder_engine(folder, latent_frames=F, latent_height=H, latent_width=W) diff --git a/families/ltx2/tests/test_vae_tiling.py b/families/ltx2/tests/test_vae_tiling.py new file mode 100644 index 0000000000..de6575f3a6 --- /dev/null +++ b/families/ltx2/tests/test_vae_tiling.py @@ -0,0 +1,112 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tile plan of the tiled / tile-parallel video VAE decode (no TensorRT, torch or GPU needed).""" + +from __future__ import annotations + +import numpy as np +import pytest + +from families.ltx2.vae_tiling import TileConfig, assign_lpt, axis_weights, blend_tiles, plan_tiles, split_axis + +GRIDS = [(31, 22, 40), (16, 17, 30), (31, 11, 20), (5, 6, 7), (3, 2, 3), (40, 9, 9)] +CONFIGS = [TileConfig(), TileConfig(tile_frames=0), TileConfig(tile_frames=136, overlap_frames=16), + TileConfig(tile_pixels=192, overlap_pixels=64, tile_frames=40, overlap_frames=8), + TileConfig(tile_frames=96, overlap_frames=24)] + + +def _weight_sum(plan: dict, frames: int, height: int, width: int) -> np.ndarray: + tt, th, tw = plan["tile_pixels"] + den = np.zeros((frames, height, width), np.float64) + for tile in plan["tiles"]: + t0, y0, x0 = tile["pixel_start"] + (tl, tr), (yl, yr), (xl, xr) = tile["ramps"] + w = (axis_weights(tt, tl, tr, temporal=True)[:, None, None] + * axis_weights(th, yl, yr, temporal=False)[None, :, None] + * axis_weights(tw, xl, xr, temporal=False)[None, None, :]) + den[t0:t0 + tt, y0:y0 + th, x0:x0 + tw] += w + return den + + +@pytest.mark.parametrize("config", CONFIGS) +@pytest.mark.parametrize("grid", GRIDS) +def test_tiles_cover_the_video_with_unit_weight_sums(grid, config) -> None: + f, h, w = grid + plan = plan_tiles(f, h, w, config, world=2) + frames, height, width = (f - 1) * 8 + 1, h * 32, w * 32 + tf, th, tw = plan["tile_latent"] + assert plan["tile_pixels"] == [(tf - 1) * 8 + 1, th * 32, tw * 32] + for tile in plan["tiles"]: + f0, h0, w0 = tile["latent_start"] + assert 0 <= f0 and f0 + tf <= f and 0 <= h0 and h0 + th <= h and 0 <= w0 and w0 + tw <= w + assert tile["pixel_start"] == [f0 * 8, h0 * 32, w0 * 32] + np.testing.assert_allclose(_weight_sum(plan, frames, height, width), 1.0, atol=1e-6) + + +def test_default_plan_for_the_large_config() -> None: + plan = plan_tiles(31, 22, 40, TileConfig(), world=2) + # Spatial 512 px tiles with >= 64 px overlap (the diffusers enable_tiling geometry, equalized); + # 241 frames fit one temporal tile. + assert plan["tile_latent"] == [31, 12, 15] + assert len(plan["tiles"]) == 6 + assert [t["rank"] for t in plan["tiles"]] == [0, 1, 0, 1, 0, 1] + assert {t["latent_start"][2] for t in plan["tiles"]} == {0, 13, 25} + assert {t["latent_start"][1] for t in plan["tiles"]} == {0, 10} + + +def test_split_overlaps_are_at_least_the_minimum() -> None: + for length in range(1, 80): + for tile, overlap in ((16, 2), (8, 2), (12, 3), (32, 4)): + size, starts = split_axis(length, tile, overlap) + assert size <= max(tile, length if length <= tile else 0) or len(starts) == 1 + assert starts[0] == 0 and starts[-1] + size == length + overlaps = [starts[i] + size - starts[i + 1] for i in range(len(starts) - 1)] + assert all(o >= overlap for o in overlaps) + assert all(a + b <= size for a, b in zip([0, *overlaps], [*overlaps, 0])) + + +def test_temporal_ramps_follow_the_causal_frame_mapping() -> None: + plan = plan_tiles(31, 4, 4, TileConfig(tile_pixels=0, tile_frames=136, overlap_frames=16)) + first, second = plan["tiles"] + assert first["latent_start"][0] == 0 and second["latent_start"][0] == 14 + # Latent overlap 3 -> (3 - 1) * 8 + 1 = 17 shared frames: the later tile fades in from 0 over + # all of them, the earlier one fades out over the last 16. + assert second["ramps"][0] == [17, 0] and first["ramps"][0] == [0, 16] + w = axis_weights(plan["tile_pixels"][0], 17, 0, temporal=True) + assert w[0] == 0.0 and w[16] == np.float32(16) / np.float32(17) and w[17] == 1.0 + + +def test_lpt_balances_unequal_volumes() -> None: + ranks = assign_lpt([5, 4, 3, 3, 1], 2) + load = [sum(v for v, r in zip([5, 4, 3, 3, 1], ranks) if r == k) for k in range(2)] + assert sorted(load) == [8, 8] + + +def test_blend_of_constant_tiles_is_the_constant() -> None: + plan = plan_tiles(5, 6, 7, TileConfig(tile_pixels=192, overlap_pixels=64, tile_frames=40, overlap_frames=8)) + tiles = [np.full(plan["tile_pixels"] + [3], 0.375, np.float16) for _ in plan["tiles"]] + out = blend_tiles(plan, tiles, 33, 192, 224) + np.testing.assert_allclose(out, 0.375, rtol=1e-6) + + +@pytest.mark.parametrize("config", [TileConfig(tile_pixels=48), TileConfig(overlap_pixels=512), + TileConfig(overlap_pixels=256), TileConfig(tile_frames=20), + TileConfig(tile_frames=64, overlap_frames=24), TileConfig(tile_pixels=-32)]) +def test_invalid_tile_configs_are_rejected(config) -> None: + with pytest.raises(ValueError): + config.validate() + + +def test_family_build_command_declares_the_tile_options() -> None: + from tensorrt_model_connect import family_cli + + declaration = family_cli.load_family_cli("ltx2") + build = next(c for c in declaration["commands"] if c["name"] == "build") + names = {a["name"] for a in build["arguments"]} + assert {"vae_tile_pixels", "vae_tile_overlap_pixels", "vae_tile_frames", "vae_tile_overlap_frames"} <= names + defaults = {a["name"]: a.get("default") for a in build["arguments"]} + config = TileConfig() + assert (defaults["vae_tile_pixels"], defaults["vae_tile_overlap_pixels"], defaults["vae_tile_frames"], + defaults["vae_tile_overlap_frames"]) == (config.tile_pixels, config.overlap_pixels, config.tile_frames, + config.overlap_frames) diff --git a/families/ltx2/tests/vae_tile_prep.py b/families/ltx2/tests/vae_tile_prep.py new file mode 100644 index 0000000000..8574549bf6 --- /dev/null +++ b/families/ltx2/tests/vae_tile_prep.py @@ -0,0 +1,49 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Single-rank preparation for the tile-parallel VAE check (uses torch + diffusers). + +``python -m families.ltx2.tests.vae_tile_prep OUT_DIR WORLD``: builds the tiny random VAE, its tile +plan for WORLD ranks and the tile-shaped plan, random latents, and the blend of diffusers' decode of +every tile. The multi-rank step (``dist_vae_tile_check``) is torch-free. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +from families.ltx2.tests import conftest # noqa: F401 - binds the TensorRT backend + + +def main() -> int: + import numpy as np + import torch + + from families.ltx2.tests import test_vae_parity as tv + from families.ltx2.vae_builder import build_vae_decoder_engine + from families.ltx2.vae_tiling import TileConfig, blend_tiles, plan_tiles + + out = Path(sys.argv[1]) + world = int(sys.argv[2]) + out.mkdir(parents=True, exist_ok=True) + vae, _ = tv.tiny_vae(out / "vae") + vae = vae.to("cuda", torch.float32) + f, h, w = tv.TILE_GRID + plan = plan_tiles(f, h, w, TileConfig(**tv.TILE_CONFIG), world=world) + tf, th, tw = plan["tile_latent"] + (out / "tile.plan").write_bytes(build_vae_decoder_engine(out / "vae", latent_frames=tf, latent_height=th, + latent_width=tw, clamp_output=False)) + packed = torch.randn(1, f * h * w, 16, generator=torch.Generator().manual_seed(4)) + tiles = [tv.diffusers_tile(vae, packed, tv.TILE_GRID, plan, tile).numpy() for tile in plan["tiles"]] + reference = blend_tiles(plan, tiles, (f - 1) * 8 + 1, h * 32, w * 32) + np.save(out / "latents.npy", packed.numpy()) + np.save(out / "reference.npy", reference) + (out / "plan.json").write_text(json.dumps({"grid": [f, h, w], "plan": plan}), encoding="utf-8") + print(f"prepared {out}: {len(plan['tiles'])} tiles of {plan['tile_latent']} latents for {world} ranks", flush=True) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/families/ltx2/vae_builder.py b/families/ltx2/vae_builder.py index eb36566a6b..2a30e3b41d 100644 --- a/families/ltx2/vae_builder.py +++ b/families/ltx2/vae_builder.py @@ -8,7 +8,9 @@ latents [1, S, 128] fp32 packed, normalized video latents (DiT layout, S = F*H*W tokens) Outputs: frames [T, H*32, W*32, 3] fp16 RGB in [0, 1] (``(x + 1) / 2`` clamped, the pipeline's - ``postprocess_video``) + ``postprocess_video``). Tile plans (``clamp_output=False``) + leave ``(x + 1) / 2`` unclamped: the runtime clamps after + blending the tiles (``vae_tiling.py``). The engine denormalizes with the VAE's ``latents_mean`` / ``latents_std`` (``scaling_factor``), unpacks the tokens to ``[1, C, F, H, W]`` and runs the non-causal ``LTX2VideoDecoder3d``: @@ -114,7 +116,7 @@ def _count(ck: Checkpoint, fmt: str) -> int: def build_vae_decoder_engine(vae_dir: str | Path, *, latent_frames: int, latent_height: int, latent_width: int, - verbose: bool = False, precision: str = "bf16") -> bytes: + verbose: bool = False, precision: str = "bf16", clamp_output: bool = True) -> bytes: ck = Checkpoint(vae_dir) cfg = ck.config() if cfg.get("timestep_conditioning"): @@ -176,11 +178,12 @@ def build_vae_decoder_engine(vae_dir: str | Path, *, latent_frames: int, latent_ x = g.reshape(x, (c, t, hh * patch, ww * patch), first=(0, 1, 5, 2, 6, 4, 7, 3)) x = g.cast(x, trt.float32) x = g.mul(g.add(x, g.scalar(1.0, trt.float32, 4)), g.scalar(0.5, trt.float32, 4)) - x = g.maximum(g.minimum(x, g.scalar(1.0, trt.float32, 4)), g.scalar(0.0, trt.float32, 4)) + if clamp_output: + x = g.maximum(g.minimum(x, g.scalar(1.0, trt.float32, 4)), g.scalar(0.0, trt.float32, 4)) frames = g.transpose(x, (1, 2, 3, 0)) # [T, H, W, 3] g.mark_output(frames, "frames", trt.float16) print(f"[ltx2] Building video VAE decoder engine (latent {f}x{h}x{w} -> {t}x{hh * patch}x{ww * patch}, " - f"{precision}) ...", file=sys.stderr) + f"{precision}{'' if clamp_output else ', tile'}) ...", file=sys.stderr) return build_plan(builder, network, label="video VAE decoder") diff --git a/families/ltx2/vae_tiling.py b/families/ltx2/vae_tiling.py new file mode 100644 index 0000000000..529c755511 --- /dev/null +++ b/families/ltx2/vae_tiling.py @@ -0,0 +1,192 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tile plan of the LTX-2.5 tiled (and tile-parallel) video VAE decode. + +The decode splits the latent video into overlapping tiles, decodes every tile independently with +one static tile-shaped plan and blends the decoded tiles with linear ramps over their overlaps, +normalized by the summed weights (``out = sum_k w_k * tile_k / sum_k w_k``, the Lightricks / +TensorRT-LLM ``tiled_decode`` blend). Each tile is an independent forward, so context-parallel +ranks decode disjoint tile subsets. Rank 0 blends every tile in tile order; the result is the same +bit for bit whichever rank decoded a tile. + +Geometry (all tiles share one shape, so one static plan serves every tile): + +- spatial axes: the fewest (then smallest) equal tiles of at most ``T`` latents with evenly spread + starts, every overlap at least ``O`` latents and both overlaps of a tile inside the tile (usually + ``n = ceil((L - O) / (T - O))`` tiles of ``ceil((L + (n - 1) O) / n)`` latents). ``T`` / ``O`` come + from the pixel tile size and overlap (512 / 64 px by default, as in diffusers ``enable_tiling``); +- time: the same split with a minimum overlap of ``O_t + 1`` latent frames (by default, clips of up + to 257 frames decode as one temporal tile and longer clips split into 256-frame tiles overlapping + by at least 24 frames). A tile covering latent frames ``[a, b)`` decodes ``(b - a - 1) * 8 + 1`` + frames placed at frame ``a * 8`` (its first latent frame decodes to a single frame, ``ltx-core`` + ``map_temporal_interval_to_frame``); +- ramps: a spatial overlap of ``r`` pixels fades in as ``k / (r + 1)`` (``k = 1..r``) and out as + ``1 - k / (r + 1)``; a temporal overlap fades the later tile in from 0 (``k / r``, ``k = 0..r-1``) + because its first frame stands in for a whole latent group, and the earlier tile out over the + remaining ``r - 1`` frames. + +The runtime (``runtime/vae_tiling.h``) consumes the plan written into ``runtime.json`` and owns +the ramp evaluation and blend; :func:`blend_tiles` is the NumPy reference of the same arithmetic. +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass + +import numpy as np + +SPATIAL_SCALE = 32 +TEMPORAL_SCALE = 8 + + +@dataclass(frozen=True) +class TileConfig: + """Tile size and minimum overlap in output pixels / frames (0 disables tiling on that axis).""" + + tile_pixels: int = 512 + overlap_pixels: int = 64 + tile_frames: int = 256 + overlap_frames: int = 24 + + def validate(self) -> None: + if self.tile_pixels < 0 or self.tile_frames < 0: + raise ValueError("VAE tile sizes must be non-negative") + # Overlaps below half a tile keep every ramp inside its tile and let only neighbouring + # tiles overlap. + if self.tile_pixels and (self.tile_pixels % SPATIAL_SCALE or self.overlap_pixels % SPATIAL_SCALE + or not 0 < 2 * self.overlap_pixels < self.tile_pixels): + raise ValueError(f"VAE spatial tiles need multiples of {SPATIAL_SCALE} px with " + "0 < overlap < tile / 2") + if self.tile_frames and (self.tile_frames % TEMPORAL_SCALE or self.overlap_frames % TEMPORAL_SCALE + or not 0 < self.overlap_frames + or 2 * (self.overlap_frames + TEMPORAL_SCALE) >= self.tile_frames): + raise ValueError(f"VAE temporal tiles need multiples of {TEMPORAL_SCALE} frames with " + f"0 < overlap < tile / 2 - {TEMPORAL_SCALE}") + + @property + def enabled(self) -> bool: + return self.tile_pixels > 0 or self.tile_frames > 0 + + +def split_axis(length: int, tile: int, overlap: int) -> tuple[int, list[int]]: + """``(size, starts)`` of equal tiles covering ``[0, length)`` with overlaps of at least ``overlap``.""" + if tile <= 0 or length <= tile: + return length, [0] + if not 0 < overlap < tile: + raise ValueError("tile overlap must be positive and smaller than the tile") + # Fewest tiles first, then the smallest size: every overlap >= `overlap`, and each tile's two + # overlaps fit in the tile (ramps never collide and only neighbours overlap). + for count in range(math.ceil((length - overlap) / (tile - overlap)), length + 1): + for size in range(math.ceil((length + (count - 1) * overlap) / count), tile + 1): + span = length - size + starts = [(2 * i * span + count - 1) // (2 * (count - 1)) for i in range(count)] + overlaps = [0] + [starts[i] + size - starts[i + 1] for i in range(count - 1)] + [0] + if min(overlaps[1:-1]) >= overlap and all(overlaps[i] + overlaps[i + 1] <= size + for i in range(count)): + return size, starts + raise ValueError(f"tiles of at most {tile} cannot cover {length} with overlaps of {overlap}") + + +def _spatial_ramps(starts: list[int], size: int) -> list[tuple[int, int]]: + """Per tile (left, right) ramps in pixels: both sides of an overlap ramp over all of it.""" + ramps = [[0, 0] for _ in starts] + for i in range(len(starts) - 1): + overlap = (starts[i] + size - starts[i + 1]) * SPATIAL_SCALE + ramps[i][1] = overlap + ramps[i + 1][0] = overlap + return [tuple(r) for r in ramps] + + +def _temporal_ramps(starts: list[int], size: int) -> list[tuple[int, int]]: + """Per tile (left, right) ramps in frames. + + Latent overlap ``ov`` shares ``(ov - 1) * 8 + 1`` frames: the later tile ramps in from 0 over all + of them, the earlier tile ramps out over the last ``(ov - 1) * 8`` (weights sum to 1 per frame). + """ + ramps = [[0, 0] for _ in starts] + for i in range(len(starts) - 1): + shared = (starts[i] + size - starts[i + 1] - 1) * TEMPORAL_SCALE + 1 + ramps[i][1] = shared - 1 + ramps[i + 1][0] = shared + return [tuple(r) for r in ramps] + + +def assign_lpt(volumes: list[int], world: int) -> list[int]: + """Longest-processing-time rank per tile (stable: ties keep tile order and the lowest rank).""" + load = [0] * world + ranks = [0] * len(volumes) + for index in sorted(range(len(volumes)), key=lambda i: -volumes[i]): + rank = min(range(world), key=lambda r: load[r]) + load[rank] += volumes[index] + ranks[index] = rank + return ranks + + +def plan_tiles(latent_frames: int, latent_height: int, latent_width: int, config: TileConfig, + world: int = 1) -> dict: + """Tile plan for ``runtime.json`` (``vae_tiling``).""" + config.validate() + if not config.enabled: + raise ValueError("VAE tiling is disabled") + s_tile = config.tile_pixels // SPATIAL_SCALE + s_overlap = config.overlap_pixels // SPATIAL_SCALE + t_tile = config.tile_frames // TEMPORAL_SCALE + t_overlap = config.overlap_frames // TEMPORAL_SCALE + 1 + tf, f_starts = split_axis(latent_frames, t_tile, t_overlap) + th, h_starts = split_axis(latent_height, s_tile, s_overlap) + tw, w_starts = split_axis(latent_width, s_tile, s_overlap) + if tf < 2 and latent_frames >= 2: + raise ValueError("VAE temporal tiles must span at least two latent frames") + f_ramps = _temporal_ramps(f_starts, tf) + h_ramps = _spatial_ramps(h_starts, th) + w_ramps = _spatial_ramps(w_starts, tw) + tiles = [] + for fi, f0 in enumerate(f_starts): + for hi, h0 in enumerate(h_starts): + for wi, w0 in enumerate(w_starts): + tiles.append({"latent_start": [f0, h0, w0], + "pixel_start": [f0 * TEMPORAL_SCALE, h0 * SPATIAL_SCALE, w0 * SPATIAL_SCALE], + "ramps": [list(f_ramps[fi]), list(h_ramps[hi]), list(w_ramps[wi])]}) + for tile, rank in zip(tiles, assign_lpt([1] * len(tiles), world)): + tile["rank"] = rank + return { + "tile_latent": [tf, th, tw], + "tile_pixels": [(tf - 1) * TEMPORAL_SCALE + 1, th * SPATIAL_SCALE, tw * SPATIAL_SCALE], + "config": {"tile_pixels": config.tile_pixels, "overlap_pixels": config.overlap_pixels, + "tile_frames": config.tile_frames, "overlap_frames": config.overlap_frames}, + "world_size": world, + "tiles": tiles, + } + + +def axis_weights(length: int, left: int, right: int, *, temporal: bool) -> np.ndarray: + """fp32 ramp weights of one tile axis (the runtime evaluates the same expressions in fp32).""" + w = np.ones(length, dtype=np.float32) + if left: + if temporal: + w[:left] = np.arange(left, dtype=np.float32) / np.float32(left) + else: + w[:left] = np.arange(1, left + 1, dtype=np.float32) / np.float32(left + 1) + if right: + denom = np.float32(right + 1) + w[length - right:] = np.float32(1.0) - np.arange(1, right + 1, dtype=np.float32) / denom + return w + + +def blend_tiles(plan: dict, decoded: list[np.ndarray], frames: int, height: int, width: int) -> np.ndarray: + """NumPy reference blend: ``decoded[k]`` is tile k ``[T, H, W, 3]``; returns clamped ``[frames, H, W, 3]``.""" + num = np.zeros((frames, height, width, 3), dtype=np.float32) + den = np.zeros((frames, height, width, 1), dtype=np.float32) + tt, th, tw = plan["tile_pixels"] + for tile, values in zip(plan["tiles"], decoded): + t0, y0, x0 = tile["pixel_start"] + (tl, tr), (yl, yr), (xl, xr) = tile["ramps"] + wt = axis_weights(tt, tl, tr, temporal=True) + wy = axis_weights(th, yl, yr, temporal=False) + wx = axis_weights(tw, xl, xr, temporal=False) + w = (wt[:, None, None] * wy[None, :, None]) * wx[None, None, :] + num[t0:t0 + tt, y0:y0 + th, x0:x0 + tw] += w[..., None] * values.astype(np.float32) + den[t0:t0 + tt, y0:y0 + th, x0:x0 + tw] += w[..., None] + return np.clip(num / den, np.float32(0.0), np.float32(1.0)) From 19f1b5789141842eb989dbf0b54437896d49984e Mon Sep 17 00:00:00 2001 From: Peter Kisfaludi Date: Fri, 2 Oct 2026 10:57:04 -0700 Subject: [PATCH 15/15] feat(ltx2): add the two-stage pipeline The diffusers LTX-2.5 two-stage recipe denoises at half resolution, doubles the latent grid with the learned latent upsampler and refines at full resolution with three distilled sigmas. Stage 1 runs about a quarter of the full-resolution tokens, so the run replaces 8 full DiT steps with 8 small steps and 3 full ones. trtmc ltx2 build --two-stage adds latent_upsampler.plan (bf16, with fp32 GroupNorm statistics) and builds the DiT for both grids. The video token count becomes a run-time dimension of one plan. The RoPE tables of both grids are constants that the plan selects by token count, and the CP row shards follow the run-time count. The engine I/O is unchanged. --set two_stage=true runs stage 1 with the distilled sigmas, upsamples on every rank, re-noises the video and audio latents to 0.909375 with draws that continue the seeded stream, and runs stage 2 with 0.909375/0.725/0.421875. Both stages use the distilled transformer; no LoRA is involved. Single-stage generation stays the default, also on two-stage bundles. Signed-off-by: Peter Kisfaludi --- families/ltx2/README.md | 25 +- families/ltx2/cli.json | 1 + families/ltx2/cli.py | 7 +- families/ltx2/dit_builder.py | 89 +++++-- families/ltx2/graph.py | 38 ++- families/ltx2/layers.py | 6 +- families/ltx2/model.py | 42 +++- families/ltx2/parallel.py | 19 +- families/ltx2/runtime/pipeline.cpp | 234 +++++++++++++----- families/ltx2/runtime/pipeline.h | 46 +++- families/ltx2/runtime/plugin.cpp | 7 +- families/ltx2/runtime/runtime_math.h | 11 + families/ltx2/tests/cp_tiny_prep.py | 10 + .../ltx2/tests/cpp/test_runtime_contract.cpp | 16 ++ families/ltx2/tests/dist_dit_cp_check.py | 69 ++++-- families/ltx2/tests/engine_runner.py | 24 +- families/ltx2/tests/np_engine.py | 46 ++-- families/ltx2/tests/test_dit_parity.py | 37 ++- families/ltx2/tests/test_model_contract.py | 14 ++ families/ltx2/tests/test_upsampler_parity.py | 77 ++++++ families/ltx2/upsampler_builder.py | 128 ++++++++++ 21 files changed, 802 insertions(+), 144 deletions(-) create mode 100644 families/ltx2/tests/test_upsampler_parity.py create mode 100644 families/ltx2/upsampler_builder.py diff --git a/families/ltx2/README.md b/families/ltx2/README.md index 0e1723590c..8cda96bcc3 100644 --- a/families/ltx2/README.md +++ b/families/ltx2/README.md @@ -22,9 +22,10 @@ Text-to-audio-video for Lightricks LTX-2.5 diffusers checkpoints (`LTX2Pipeline` | Section | Contents | |---|---| | `text_encoder.plan` | Gemma 4 text tower and the LTX-2 text connectors (video and audio context) | -| `denoiser.plan` | Joint audio/video DiT. With CP=2, one plan serves both ranks. | +| `denoiser.plan` | Joint audio/video DiT. With CP=2, one plan serves both ranks. Two-stage bundles serve the full and the half-resolution grid. | | `vae.plan` | Video VAE decoder for one tile shape (whole video with `--vae-tile-pixels 0 --vae-tile-frames 0`) | | `audio.plan` | Audio VAE decoder and vocoder with bandwidth extension | +| `latent_upsampler.plan` | Two-stage bundles only: the 2x spatial latent upsampler | | `tokenizer.json`, `runtime.json` | Tokenizer, shapes and schedule | Context parallelism keeps the audio stream and text replicated and shards the video @@ -54,6 +55,26 @@ and writes it into `runtime.json`. `--vae-tile-overlap-frames`. A size of 0 leaves that axis untiled, and setting both sizes to 0 builds the untiled decoder. The shared `trtmc build --family ltx2` uses the defaults. +## Two-stage pipeline + +`trtmc ltx2 build --two-stage` builds a bundle that also runs the two-stage recipe of the diffusers +LTX-2.5 documentation (distilled checkpoint in both stages, no LoRA). Run it with +`--set two_stage=true`; single-stage generation stays the default, also on two-stage bundles. + +1. Stage 1 denoises at half the width and height with the 8 distilled sigmas. +2. The latent upsampler (`latent_upsampler/`) doubles the latent grid. +3. The video and audio latents are re-noised to the first stage 2 sigma + (`noise_scale * noise + (1 - noise_scale) * latents`, noise_scale 0.909375). The noise draws + continue the seeded stream after the stage 1 noise. +4. Stage 2 refines at full resolution with `STAGE_2_DISTILLED_SIGMA_VALUES` + (0.909375, 0.725, 0.421875), then the video and audio are decoded. + +Width and height must be multiples of 64, and the checkpoint must include `latent_upsampler/`. One +DiT plan serves both grids: the video token count is a run-time dimension, and the RoPE tables of +both grids are constants that the plan selects by token count. Stage 1 runs about a quarter of the +full-resolution tokens, so the two-stage run does 8 small steps and 3 full steps instead of 8 full +steps. Context parallelism applies to both stages, and every rank upsamples the latents itself. + ## Build and run ```bash @@ -101,7 +122,7 @@ To give each rank its own TensorRT-RTX runtime cache, put `{rank}` in the path, ## Tests - `tests/test_*_parity.py` build each engine from tiny random weights and compare it - with diffusers. + with diffusers, including the latent upsampler and the two-grid DiT plan. - `tests/test_context_parallel.py` runs the CP=2 DiT on two GPUs (torch-free ranks) against the single-device plan and diffusers. It also runs the tile-parallel VAE decode and checks that it matches the single-GPU tiled decode bit for bit. It needs `TRTMC_NCCL_LIBRARY`. diff --git a/families/ltx2/cli.json b/families/ltx2/cli.json index 8c084036c4..438e9fc1c4 100644 --- a/families/ltx2/cli.json +++ b/families/ltx2/cli.json @@ -21,6 +21,7 @@ {"name": "vae_tile_overlap_pixels", "flags": ["--vae-tile-overlap-pixels"], "type": "int", "default": 64, "help": "Minimum spatial tile overlap in pixels"}, {"name": "vae_tile_frames", "flags": ["--vae-tile-frames"], "type": "int", "default": 256, "help": "Video VAE tile length in frames (0: untiled time axis)"}, {"name": "vae_tile_overlap_frames", "flags": ["--vae-tile-overlap-frames"], "type": "int", "default": 24, "help": "Minimum temporal tile overlap in frames"}, + {"name": "two_stage", "flags": ["--two-stage"], "type": "bool", "action": "store_true", "default": false, "help": "Also build the two-stage pipeline (half-resolution stage 1, latent upsampler, full-resolution refinement); run it with --set two_stage=true"}, {"name": "verbose", "flags": ["--verbose"], "type": "bool", "action": "store_true", "default": false} ] } diff --git a/families/ltx2/cli.py b/families/ltx2/cli.py index 0c6fa21600..c250655415 100644 --- a/families/ltx2/cli.py +++ b/families/ltx2/cli.py @@ -29,6 +29,7 @@ class BuildRequest: video_num_frames: int | None = None context_parallel_size: int = 1 vae_tiles: TileConfig = field(default_factory=TileConfig) + two_stage: bool = False verbose: bool = False def __post_init__(self) -> None: @@ -78,12 +79,14 @@ def build(*, model: str, output: Path, revision: str | None = None, precision: s backend: str = "trt", image_height: int | None = None, image_width: int | None = None, video_num_frames: int | None = None, max_sequence_length: int | None = None, context_parallel_size: int = 1, vae_tile_pixels: int = 512, vae_tile_overlap_pixels: int = 64, - vae_tile_frames: int = 256, vae_tile_overlap_frames: int = 24, verbose: bool = False) -> int: + vae_tile_frames: int = 256, vae_tile_overlap_frames: int = 24, two_stage: bool = False, + verbose: bool = False) -> int: tiles = TileConfig(tile_pixels=vae_tile_pixels, overlap_pixels=vae_tile_overlap_pixels, tile_frames=vae_tile_frames, overlap_frames=vae_tile_overlap_frames) request = BuildRequest(model_dir=resolve_model(model, revision), precision=precision, backend=backend, max_sequence_length=max_sequence_length, image_height=image_height, image_width=image_width, video_num_frames=video_num_frames, - context_parallel_size=context_parallel_size, vae_tiles=tiles, verbose=verbose) + context_parallel_size=context_parallel_size, vae_tiles=tiles, + two_stage=two_stage, verbose=verbose) build_bundle(request, output) return 0 diff --git a/families/ltx2/dit_builder.py b/families/ltx2/dit_builder.py index c7ca945c56..e4e9600607 100644 --- a/families/ltx2/dit_builder.py +++ b/families/ltx2/dit_builder.py @@ -21,6 +21,10 @@ video_velocity [B, S, 128] fp32 (full sequence on every rank) audio_velocity [B, Sa, 128] fp32 +Two-stage plans (``extra_shapes``) also serve the half-resolution stage 1 grid: ``S`` is a run-time +dimension (one optimization profile spanning both token counts). The RoPE tables of both grids are +baked in one after the other, and the run-time token count selects the rows of its grid. + Precision: bf16 strongly typed (diffusers runs this model in bf16), with fp32 RMSNorm / LayerNorm statistics, fp32 split RoPE, fp32 timestep sinusoids and fp32 GELU/SiLU islands. @@ -68,7 +72,7 @@ stg_lerp, to_heads, ) -from .parallel import add_collective, local_row_indices +from .parallel import add_collective, local_row_indices, local_row_start EPS = 1e-6 _FP16_SAFE_BF16_MAX = 65280.0 # largest bf16 value that is finite in fp16 @@ -316,28 +320,74 @@ def _scalar_like(g: Graph, x, value: float): return g.scalar(value, x.dtype, len(x.shape)) -def add_dit(g: Graph, ckpt: Checkpoint, cfg: DiTConfig, shape: DiTShape, inputs: dict, *, cp: int = 1, - stg_blocks: tuple[int, ...] = (), num_layers: int | None = None): - """Adds the DiT; returns (video_velocity ``[B, S_local, C]`` bf16, audio_velocity ``[B, Sa, C]`` bf16).""" - B = shape.batch - S = shape.video_tokens - if S % cp: - raise ValueError(f"video tokens {S} are not divisible by context_parallel_size {cp}") +def _check_grids(shape: DiTShape, extra_shapes: tuple[DiTShape, ...], cp: int, cfg: DiTConfig) -> None: + for grid in (shape, *extra_shapes): + if grid.video_tokens % cp: + raise ValueError(f"video tokens {grid.video_tokens} are not divisible by context_parallel_size {cp}") + if (grid.batch, grid.audio_frames, grid.text_len, grid.fps) != (shape.batch, shape.audio_frames, + shape.text_len, shape.fps): + raise ValueError("every video grid of one DiT plan needs the same batch, audio, text and fps") + if len(extra_shapes) > 1 or (extra_shapes and extra_shapes[0].video_tokens == shape.video_tokens): + raise ValueError("a DiT plan serves its grid plus at most one grid with a different token count") for name, h in (("video", cfg.heads), ("audio", cfg.audio_heads)): if h % cp: raise ValueError(f"{name} heads {h} are not divisible by context_parallel_size {cp}") + + +def _grid_rows(g: Graph, video_latent, shape: DiTShape, extra: DiTShape, cp: int): + """Run-time grid selection of a two-grid plan. + + Returns (latent rows of this rank or None, RoPE rows): the RoPE tables hold the rows of + ``shape`` followed by the rows of ``extra``, and the run-time token count picks the block. + """ + s_main, s_extra = shape.video_tokens, extra.video_tokens + tokens = g.dim(video_latent, 1) + # 0 for the main grid, 1 for the extra grid (exact integer arithmetic on the two counts). + main = g.const(np.array([s_main], np.int32), trt.int32) + distance = g.sub(main, tokens) if s_extra < s_main else g.sub(tokens, main) + which = g.ew(distance, g.const(np.array([abs(s_extra - s_main)], np.int32), trt.int32), + trt.ElementWiseOperation.FLOOR_DIV) + offset = g.mul(which, main) + if cp == 1: + return None, g.arange(tokens, offset) + local = g.ew(tokens, g.const(np.array([cp], np.int32), trt.int32), trt.ElementWiseOperation.FLOOR_DIV) + rows = g.arange(local, local_row_start(g, cp=cp, local_rows=local)) + return rows, g.add(rows, offset) + + +def add_dit(g: Graph, ckpt: Checkpoint, cfg: DiTConfig, shape: DiTShape, inputs: dict, *, cp: int = 1, + stg_blocks: tuple[int, ...] = (), num_layers: int | None = None, + extra_shapes: tuple[DiTShape, ...] = ()): + """Adds the DiT; returns (video_velocity ``[B, S_local, C]`` bf16, audio_velocity ``[B, Sa, C]`` bf16). + + ``extra_shapes`` (at most one): a second video grid served by the same plan. The video token + count is then a run-time dimension and picks the RoPE rows of its grid. + """ + B = shape.batch + S = shape.video_tokens + _check_grids(shape, extra_shapes, cp, cfg) s_loc = S // cp D, Da = cfg.dim, cfg.audio_dim n_layers = cfg.layers if num_layers is None else num_layers tables = rope_tables(cfg, shape) + for extra in extra_shapes: + more = rope_tables(cfg, extra) + for key in ("video", "ca_video"): + tables[key] = tuple(np.concatenate([a, b], axis=0) for a, b in zip(tables[key], more[key])) rope_v = rope_constants(g, *tables["video"]) rope_ca_v = rope_constants(g, *tables["ca_video"]) rope_a = rope_constants(g, *tables["audio"]) rope_ca_a = rope_constants(g, *tables["ca_audio"]) video_latent = g.cast(inputs["video_latent"], trt.bfloat16) - if cp > 1: + if extra_shapes: + rows, rope_rows = _grid_rows(g, inputs["video_latent"], shape, extra_shapes[0], cp) + if rows is not None: + video_latent = g.gather(video_latent, rows, 1) + rope_v = gather_rope_rows(g, rope_v, rope_rows) + rope_ca_v = gather_rope_rows(g, rope_ca_v, rope_rows) + elif cp > 1: rows = local_row_indices(g, cp=cp, local_rows=s_loc) video_latent = g.gather(video_latent, rows, 1) rope_v = gather_rope_rows(g, rope_v, rows) @@ -451,7 +501,7 @@ def _gather_video_rows(g: Graph, y, cp: int, batch: int): b, s_loc, c = (int(v) for v in yf.shape) if b == 1: out = add_collective(g.net, g.reshape(yf, (s_loc, c)), trt.CollectiveOperation.ALL_GATHER, cp) - return g.reshape(out, (1, s_loc * cp, c)) + return g.reshape(out, (1, s_loc * cp if s_loc >= 0 else -1, c)) t = g.transpose(yf, (1, 0, 2)) # [S/cp, B, C] out = add_collective(g.net, t, trt.CollectiveOperation.ALL_GATHER, cp) # [S, B, C] return g.transpose(out, (1, 0, 2)) @@ -459,14 +509,17 @@ def _gather_video_rows(g: Graph, y, cp: int, batch: int): def build_dit_engine(transformer_dir: str | Path, shape: DiTShape, *, cp_size: int = 1, stg_blocks: tuple[int, ...] = (28,), num_layers: int | None = None, - verbose: bool = False) -> bytes: + verbose: bool = False, extra_shapes: tuple[DiTShape, ...] = ()) -> bytes: + """Serialized DiT plan for ``shape``; ``extra_shapes`` adds one more video grid (run-time token count).""" ckpt = Checkpoint(transformer_dir) cfg = DiTConfig.from_dict(ckpt.config()) builder, network = new_network(make_logger(verbose)) g = Graph(network) B, S, Sa, L = shape.batch, shape.video_tokens, shape.audio_frames, shape.text_len + tokens = sorted({S, *(extra.video_tokens for extra in extra_shapes)}) inputs = { - "video_latent": network.add_input("video_latent", trt.float32, (B, S, cfg.in_channels)), + "video_latent": network.add_input("video_latent", trt.float32, + (B, S if len(tokens) == 1 else -1, cfg.in_channels)), "audio_latent": network.add_input("audio_latent", trt.float32, (B, Sa, cfg.audio_in_channels)), "video_context": network.add_input("video_context", trt.bfloat16, (B, L, cfg.cross_attention_dim)), "audio_context": network.add_input("audio_context", trt.bfloat16, (B, L, cfg.audio_cross_attention_dim)), @@ -476,14 +529,18 @@ def build_dit_engine(transformer_dir: str | Path, shape: DiTShape, *, cp_size: i } if cfg.cross_attention_dim != cfg.dim or cfg.audio_cross_attention_dim != cfg.audio_dim: raise NotImplementedError("LTX-2.5 DiT builder expects the connector widths to match the streams") - video, audio = add_dit(g, ckpt, cfg, shape, inputs, cp=cp_size, stg_blocks=stg_blocks, num_layers=num_layers) + video, audio = add_dit(g, ckpt, cfg, shape, inputs, cp=cp_size, stg_blocks=stg_blocks, num_layers=num_layers, + extra_shapes=tuple(extra_shapes)) if cp_size > 1: video = _gather_video_rows(g, video, cp_size, B) g.mark_output(video, "video_velocity", trt.float32) g.mark_output(audio, "audio_velocity", trt.float32) - print(f"[ltx2] Building DiT engine (batch={B}, video_tokens={S}, audio_tokens={Sa}, cp={cp_size}, " - f"layers={num_layers or cfg.layers}) ...", file=sys.stderr) - return build_plan(builder, network, label="DiT") + profile = None + if len(tokens) > 1: + profile = {"video_latent": tuple((B, s, cfg.in_channels) for s in (tokens[0], S, tokens[-1]))} + print(f"[ltx2] Building DiT engine (batch={B}, video_tokens={'/'.join(map(str, tokens))}, audio_tokens={Sa}, " + f"cp={cp_size}, layers={num_layers or cfg.layers}) ...", file=sys.stderr) + return build_plan(builder, network, label="DiT", profile=profile) def load_dit_config(transformer_dir: str | Path) -> DiTConfig: diff --git a/families/ltx2/graph.py b/families/ltx2/graph.py index 465c472753..a385c01dbe 100644 --- a/families/ltx2/graph.py +++ b/families/ltx2/graph.py @@ -137,6 +137,34 @@ def gather(self, x, indices, axis: int): def select(self, cond, a, b): return self.net.add_select(cond, a, b).get_output(0) + # ------------------------------------------------------------------ run-time shapes + + def dim(self, x, axis: int): + """int32 ``[1]`` run-time size of ``x`` along ``axis``.""" + shape = self.cast(self.net.add_shape(x).get_output(0), trt.int32) + return self.slice(shape, (axis,), (1,)) + + def arange(self, length, start): + """int32 ``[length]`` values ``start, start + 1, ...`` for int32 ``[1]`` tensors ``length`` / ``start``. + + The fill only depends on the (shape) length; ``start`` may be a device value, e.g. derived from + a collective, and is added afterwards. + """ + layer = self.net.add_fill((1,), trt.FillOperation.LINSPACE, trt.int32) + layer.set_input(0, length) + layer.set_input(1, self.const(np.zeros((), np.int32), trt.int32, shape=())) + layer.set_input(2, self.const(np.ones(1, np.int32), trt.int32)) + return self.add(layer.get_output(0), start) + + def take(self, x, axis: int, index: int): + """``x[..., index:index + 1, ...]`` along ``axis``; also when other axes are only known at run time.""" + if all(int(s) >= 0 for s in x.shape): + start = [0] * len(x.shape) + size = [int(s) for s in x.shape] + start[axis], size[axis] = index, 1 + return self.slice(x, start, size) + return self.gather(x, self.const(np.array([index], np.int32), trt.int32), axis) + def mark_output(self, x, name: str, dtype: "trt.DataType | None" = None): if dtype is not None: x = self.cast(x, dtype) @@ -245,16 +273,22 @@ def new_network(logger: "trt.ILogger"): return builder, network -def build_plan(builder, network, *, label: str = "engine", tf32: bool = True): +def build_plan(builder, network, *, label: str = "engine", tf32: bool = True, profile: dict | None = None): """Serialized plan (a bytes-like ``IHostMemory``; multi-GB plans are not copied again). ``tf32=False`` keeps fp32 convolutions / matrix multiplies in full fp32 (TensorRT allows - TF32 for fp32 layers by default). + TF32 for fp32 layers by default). ``profile`` maps each run-time shaped input to its + ``(min, opt, max)`` shapes (one optimization profile). """ config = builder.create_builder_config() config.builder_optimization_level = 3 if not tf32: config.clear_flag(trt.BuilderFlag.TF32) + if profile: + shapes = builder.create_optimization_profile() + for name, (low, opt, high) in profile.items(): + shapes.set_shape(name, low, opt, high) + config.add_optimization_profile(shapes) plan = builder.build_serialized_network(network, config) if plan is None: raise RuntimeError(f"TensorRT failed to build the LTX-2.5 {label}") diff --git a/families/ltx2/layers.py b/families/ltx2/layers.py index 8a1e6caf38..1c50a94262 100644 --- a/families/ltx2/layers.py +++ b/families/ltx2/layers.py @@ -70,13 +70,13 @@ def gather_rope_rows(g: Graph, rope: RopeTables, rows) -> RopeTables: def apply_split_rope(g: Graph, x, rope: RopeTables): - """``x``: ``[B, T, H*2r]`` -> same shape/dtype, rotated in fp32.""" + """``x``: ``[B, T, H*2r]`` -> same shape/dtype, rotated in fp32 (``T`` may be a run-time size).""" b, t, d = (int(s) for s in x.shape) h, r = rope.heads, rope.half out_dtype = x.dtype x5 = g.reshape(g.cast(x, trt.float32), (b, t, h, 2, r)) - x1 = g.slice(x5, (0, 0, 0, 0, 0), (b, t, h, 1, r)) - x2 = g.slice(x5, (0, 0, 0, 1, 0), (b, t, h, 1, r)) + x1 = g.take(x5, 3, 0) + x2 = g.take(x5, 3, 1) o1 = g.sub(g.mul(x1, rope.cos), g.mul(x2, rope.sin)) o2 = g.add(g.mul(x2, rope.cos), g.mul(x1, rope.sin)) return g.cast(g.reshape(g.concat([o1, o2], axis=3), (b, t, d)), out_dtype) diff --git a/families/ltx2/model.py b/families/ltx2/model.py index a5908a0541..d87d3b2516 100644 --- a/families/ltx2/model.py +++ b/families/ltx2/model.py @@ -11,6 +11,8 @@ - ``vae.plan``: the video VAE decoder, by default one tile-shaped plan for the tiled decode (``vae_tiling.py``; context-parallel ranks decode disjoint tiles, rank 0 blends them) - ``audio.plan``: the audio VAE decoder + vocoder with bandwidth extension (48 kHz stereo) + - ``latent_upsampler.plan`` (two-stage bundles): the 2x spatial latent upsampler; the DiT plan + then also serves the half-resolution stage 1 grid - ``tokenizer.json`` and ``runtime.json`` Every engine is built directly with the TensorRT network API in bf16 (the precision LTX-2.5 is @@ -22,6 +24,7 @@ import json import sys import time +from dataclasses import replace from pathlib import Path from typing import TYPE_CHECKING @@ -38,6 +41,10 @@ # 8-step schedule. LTX-2.5's shipped scheduler config disables dynamic shifting, so the # pipeline uses these values unshifted (timesteps = sigma * 1000), then the terminal 0. DISTILLED_SIGMAS = (1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875) +# diffusers ``pipelines/ltx2/utils.py`` STAGE_2_DISTILLED_SIGMA_VALUES: the two-stage refinement +# schedule at full resolution. Stage 2 re-noises the upsampled stage 1 latents (video and audio) +# to its first sigma and runs the same distilled transformer (no LoRA). +STAGE_2_DISTILLED_SIGMAS = (0.909375, 0.725, 0.421875) DEFAULT_HEIGHT = 544 # the LTX-2.5 model card's 960x544, 121 frames at 24 fps DEFAULT_WIDTH = 960 @@ -101,6 +108,15 @@ def _write_plan(writer: "BundleWriter", name: str, plan) -> None: section.write(memoryview(plan)) +def _stage1_shape(model_dir: Path, shape, height: int, width: int, spatial: int): + """Half-resolution stage 1 grid of the two-stage pipeline (diffusers: ``width // 2``, ``height // 2``).""" + if height % (2 * spatial) or width % (2 * spatial): + raise ValueError(f"ltx2 two-stage bundles need image_height/image_width divisible by {2 * spatial}") + if not (model_dir / "latent_upsampler" / "config.json").is_file(): + raise FileNotFoundError(f"ltx2 two-stage bundles need the checkpoint's latent_upsampler/ ({model_dir})") + return replace(shape, latent_height=shape.latent_height // 2, latent_width=shape.latent_width // 2) + + def build(request: BuildRequest, writer: "BundleWriter") -> None: """Build one LTX-2.5 text-to-audio-video bundle (``trtmc ltx2 build`` or the shared ``trtmc build``).""" request = coerce_request(request) @@ -156,6 +172,10 @@ def build(request: BuildRequest, writer: "BundleWriter") -> None: ) validate_context_parallel_layout(parallel, video_tokens=shape.video_tokens, video_heads=dit_cfg.heads, audio_heads=dit_cfg.audio_heads) + stage1 = _stage1_shape(model_dir, shape, height, width, spatial) if request.two_stage else None + if stage1 is not None: + validate_context_parallel_layout(parallel, video_tokens=stage1.video_tokens, video_heads=dit_cfg.heads, + audio_heads=dit_cfg.audio_heads) writer.set_header(family="ltx2", task=request.task, backend=request.backend) started = time.perf_counter() @@ -164,9 +184,19 @@ def build(request: BuildRequest, writer: "BundleWriter") -> None: _log(f"text encoder engine built in {time.perf_counter() - started:.1f} s") started = time.perf_counter() _write_plan(writer, "denoiser.plan", build_dit_engine(model_dir / "transformer", shape, cp_size=parallel.cp_size, - verbose=request.verbose)) + verbose=request.verbose, + extra_shapes=(stage1,) if stage1 else ())) _log(f"DiT engine built in {time.perf_counter() - started:.1f} s (cp={parallel.cp_size}, " - f"{shape.video_tokens} video + {shape.audio_frames} audio tokens)") + f"{shape.video_tokens}{f' / {stage1.video_tokens}' if stage1 else ''} video + " + f"{shape.audio_frames} audio tokens)") + if stage1 is not None: + from .upsampler_builder import build_latent_upsampler_engine + + started = time.perf_counter() + _write_plan(writer, "latent_upsampler.plan", build_latent_upsampler_engine( + model_dir / "latent_upsampler", model_dir / "vae", latent_frames=stage1.latent_frames, + latent_height=stage1.latent_height, latent_width=stage1.latent_width, verbose=request.verbose)) + _log(f"latent upsampler engine built in {time.perf_counter() - started:.1f} s") started = time.perf_counter() tiling = None vae_grid = (shape.latent_frames, shape.latent_height, shape.latent_width) @@ -207,4 +237,12 @@ def build(request: BuildRequest, writer: "BundleWriter") -> None: } if tiling is not None: runtime["vae_tiling"] = tiling + if stage1 is not None: + runtime["two_stage"] = { + "latent_height": stage1.latent_height, + "latent_width": stage1.latent_width, + "sigmas": [*DISTILLED_SIGMAS, 0.0], + "stage2_sigmas": [*STAGE_2_DISTILLED_SIGMAS, 0.0], + "noise_scale": STAGE_2_DISTILLED_SIGMAS[0], + } writer.add_json("runtime.json", runtime) diff --git a/families/ltx2/parallel.py b/families/ltx2/parallel.py index 5b75a072f3..93fb5313d3 100644 --- a/families/ltx2/parallel.py +++ b/families/ltx2/parallel.py @@ -73,13 +73,24 @@ def rank_selector_values(cp_size: int) -> np.ndarray: return (np.arange(cp_size, dtype=np.float32) / np.float32(cp_size)).reshape(cp_size, 1) -def local_row_indices(g, *, cp: int, local_rows: int): - """int32 ``[local_rows]`` indices of the contiguous token shard owned by this rank.""" +def rank_index(g, *, cp: int): + """int32 ``[1]`` index of this rank (one tiny REDUCE_SCATTER of replicated values).""" import tensorrt as trt selector = g.const(rank_selector_values(cp), trt.float32) rank_f = add_collective(g.net, selector, trt.CollectiveOperation.REDUCE_SCATTER, cp, reduce_operation=trt.ReduceOperation.SUM) - rank_i = g.reshape(g.cast(rank_f, trt.int32), (1,)) - start = g.mul(rank_i, g.const(np.array([local_rows], np.int32), trt.int32)) + return g.reshape(g.cast(rank_f, trt.int32), (1,)) + + +def local_row_indices(g, *, cp: int, local_rows: int): + """int32 ``[local_rows]`` indices of the contiguous token shard owned by this rank.""" + import tensorrt as trt + + start = g.mul(rank_index(g, cp=cp), g.const(np.array([local_rows], np.int32), trt.int32)) return g.add(g.const(np.arange(local_rows, dtype=np.int32), trt.int32), start) + + +def local_row_start(g, *, cp: int, local_rows): + """int32 ``[1]`` first row of this rank's shard for a run-time shard length ``local_rows`` (int32 ``[1]``).""" + return g.mul(rank_index(g, cp=cp), local_rows) diff --git a/families/ltx2/runtime/pipeline.cpp b/families/ltx2/runtime/pipeline.cpp index b3955898cb..1a833228cc 100644 --- a/families/ltx2/runtime/pipeline.cpp +++ b/families/ltx2/runtime/pipeline.cpp @@ -117,8 +117,12 @@ std::string trim(const std::string& text) { // Optional family diagnostics (all off by default): // TRTMC_LTX2_INITIAL_LATENTS raw fp32 file: packed video [S, C] then audio [Sa, Ca] noise -// (replaces the seeded noise, e.g. a reference pipeline's draw) -// TRTMC_LTX2_DUMP_LATENTS raw fp32 file written with the final video then audio latents +// (replaces the seeded noise, e.g. a reference pipeline's draw). +// Two-stage runs append the stage 2 re-noise draws: packed +// full-resolution video [S2, C] then audio [Sa, Ca]. +// TRTMC_LTX2_DUMP_LATENTS raw fp32 file written with the final video then audio latents; +// two-stage runs also write .stage1 (stage 1 video, audio) and +// .upsampled (upsampled video, audio) // TRTMC_LTX2_DECODE_LATENTS raw fp32 file in the TRTMC_LTX2_DUMP_LATENTS layout; replaces the // denoised latents before the decode (decoder checks, e.g. the // single-GPU vs tile-parallel decode of identical latents) @@ -135,10 +139,12 @@ std::vector read_f32_file(const char* path) { return values; } -void maybe_dump(const std::vector& video, const std::vector& audio) { - const char* path = std::getenv("TRTMC_LTX2_DUMP_LATENTS"); - if (path == nullptr || *path == '\0') +void maybe_dump(const std::vector& video, const std::vector& audio, + const char* suffix = "") { + const char* base = std::getenv("TRTMC_LTX2_DUMP_LATENTS"); + if (base == nullptr || *base == '\0') return; + const std::string path = std::string(base) + suffix; std::ofstream output(path, std::ios::binary | std::ios::trunc); output.write(reinterpret_cast(video.data()), static_cast(video.size() * 4)); @@ -162,14 +168,33 @@ void replace_final_latents(std::vector& video, std::vector& audio) std::cerr << "[ltx2] decoding the latents of " << path << "\n"; } -const std::array& config_fields() { - static const std::array fields{{ +const std::array& config_fields() { + static const std::array fields{{ {"seed", internal::ConfigKind::I64, internal::ConfigValue{std::int64_t{0}}, "Seed of the initial video and audio noise (portable std::mt19937 + normal draws)."}, + {"two_stage", internal::ConfigKind::Bool, internal::ConfigValue{false}, + "Two-stage pipeline: half-resolution stage 1, latent upsampler, full-resolution " + "refinement (bundles built with trtmc ltx2 build --two-stage)."}, }}; return fields; } +LTX2TwoStage parse_two_stage(const nlohmann::json& doc) { + LTX2TwoStage two; + two.latent_height = doc.at("latent_height").get(); + two.latent_width = doc.at("latent_width").get(); + two.sigmas = doc.at("sigmas").get>(); + two.stage2_sigmas = doc.at("stage2_sigmas").get>(); + two.noise_scale = doc.at("noise_scale").get(); + for (const auto* schedule : {&two.sigmas, &two.stage2_sigmas}) { + if (schedule->size() < 2 || schedule->back() != 0.0F) + throw std::runtime_error("LTX-2.5 two-stage sigmas must end with the terminal 0"); + } + if (two.latent_height <= 0 || two.latent_width <= 0) + throw std::runtime_error("LTX-2.5 runtime.json has an invalid two-stage grid"); + return two; +} + ltx2::VaeTilePlan parse_tile_plan(const nlohmann::json& doc) { ltx2::VaeTilePlan plan; plan.tile_latent = doc.at("tile_latent").get>(); @@ -259,6 +284,12 @@ LTX2Options parse_ltx2_options(const std::string& runtime_json, int32_t world_si throw std::runtime_error("LTX-2.5 runtime.json has invalid shapes"); if (doc.contains("audio_waveform_shape")) o.audio_waveform_shape = doc.at("audio_waveform_shape").get>(); + if (doc.contains("two_stage")) { + o.two_stage = parse_two_stage(doc.at("two_stage")); + if (o.two_stage.latent_height * 2 != o.latent_height || + o.two_stage.latent_width * 2 != o.latent_width) + throw std::runtime_error("LTX-2.5 two-stage grid must be half the video latent grid"); + } if (doc.contains("vae_tiling")) { const auto& tiling = doc.at("vae_tiling"); if (tiling.at("world_size").get() != world_size) @@ -274,11 +305,12 @@ LTX2Pipeline::LTX2Pipeline(std::unique_ptr text_encoder, std::unique_ptr denoiser, std::unique_ptr vae, std::unique_ptr audio, LTX2Options options, std::shared_ptr tokenizer, - LTX2DistributedContext distributed) + LTX2DistributedContext distributed, + std::unique_ptr upsampler) : distributed_(std::move(distributed)), text_encoder_(std::move(text_encoder)), denoiser_(std::move(denoiser)), vae_(std::move(vae)), audio_(std::move(audio)), - options_(std::move(options)), tokenizer_(std::move(tokenizer)), progress_(distributed_.rank) { -} + upsampler_(std::move(upsampler)), options_(std::move(options)), + tokenizer_(std::move(tokenizer)), progress_(distributed_.rank) {} LTX2Pipeline::~LTX2Pipeline() = default; @@ -308,9 +340,9 @@ LTX2Pipeline::TextContext LTX2Pipeline::encode(const std::string& text) { } void LTX2Pipeline::run_dit(const std::vector& video, const std::vector& audio, - const TextContext& text, float timestep, std::vector& video_out, - std::vector& audio_out) { - const int64_t S = options_.video_tokens(); + const TextContext& text, float timestep, int64_t video_tokens, + std::vector& video_out, std::vector& audio_out) { + const int64_t S = video_tokens; const int64_t Sa = options_.audio_frames; const int64_t L = options_.text_seq_len; std::vector t{timestep}; @@ -334,6 +366,95 @@ void LTX2Pipeline::run_dit(const std::vector& video, const std::vector sorted = step_ms; + std::sort(sorted.begin(), sorted.end()); + return sorted[sorted.size() / 2]; +} + +LTX2Pipeline::Noise LTX2Pipeline::initial_noise(int64_t seed, bool two_stage) const { + const auto channels = static_cast(options_.latent_channels); + const auto full = static_cast(options_.video_tokens()) * channels; + const auto stage1 = + two_stage + ? static_cast(options_.two_stage.video_tokens(options_.latent_frames)) * + channels + : full; + const auto audio = + static_cast(options_.audio_frames) * options_.audio_latent_channels; + Noise noise; + noise.video.resize(stage1); + noise.audio.resize(audio); + if (two_stage) { + noise.video_stage2.resize(full); + noise.audio_stage2.resize(audio); + } + const std::vector*> draws{&noise.video, &noise.audio, &noise.video_stage2, + &noise.audio_stage2}; + if (const char* path = std::getenv("TRTMC_LTX2_INITIAL_LATENTS"); + path != nullptr && *path != '\0') { + const auto values = read_f32_file(path); + std::size_t total = 0; + for (const auto* d : draws) + total += d->size(); + if (values.size() != total) + throw std::runtime_error("TRTMC_LTX2_INITIAL_LATENTS must hold the packed video then " + "audio noise (then the stage 2 video and audio draws)"); + auto it = values.begin(); + for (auto* d : draws) { + std::copy_n(it, d->size(), d->begin()); + it += static_cast(d->size()); + } + return noise; + } + std::mt19937 generator(static_cast(seed)); + ltx2::LibstdcxxNormalFloat normal; + for (auto* d : draws) + for (auto& v : *d) + v = normal(generator); + return noise; +} + +void LTX2Pipeline::denoise(std::vector& video, std::vector& audio, + const TextContext& text, const std::vector& sigmas, + int64_t video_tokens, const char* stage, StageTimes& times) { + const auto start = Clock::now(); + const int32_t steps = static_cast(sigmas.size()) - 1; + std::vector video_v; + std::vector audio_v; + for (int32_t step = 0; step < steps; ++step) { + const auto step_start = Clock::now(); + const float sigma = sigmas[static_cast(step)]; + const float sigma_next = sigmas[static_cast(step) + 1]; + run_dit(video, audio, text, sigma * 1000.0F, video_tokens, video_v, audio_v); + ltx2_euler_step(video, video_v, sigma, sigma_next); + ltx2_euler_step(audio, audio_v, sigma, sigma_next); + times.step_ms.push_back(elapsed_ms(step_start, Clock::now())); + if (progress_.enabled()) { + std::ostringstream detail; + detail << "stage=" << stage << " step=" << (step + 1) << "/" << steps + << " step_ms=" << std::fixed << std::setprecision(3) << times.step_ms.back(); + progress_.emit("step", detail.str()); + } + } + times.total_ms = elapsed_ms(start, Clock::now()); +} + +std::vector LTX2Pipeline::upsample(const std::vector& video, int64_t video_tokens) { + if (!upsampler_) + throw std::runtime_error("LTX-2.5 two-stage run without latent_upsampler.plan"); + TensorMap inputs; + inputs["latents"] = Tensor{const_cast(video.data()), + {1, video_tokens, options_.latent_channels}, + DType::kFloat32}; + const auto outputs = upsampler_->forward(inputs); + return float_output(outputs, "upsampled", + static_cast(options_.video_tokens()) * + options_.latent_channels); +} + std::vector LTX2Pipeline::decode_video(const std::vector& video_latents) { TensorMap inputs; inputs["latents"] = Tensor{const_cast(video_latents.data()), @@ -537,12 +658,16 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq internal::validate_config({fields.data(), fields.size()}, config); const auto seed = internal::config_get(config, {fields.data(), fields.size()}, "seed").value(); + const bool two_stage = + internal::config_get(config, {fields.data(), fields.size()}, "two_stage").value(); + if (two_stage && !options_.two_stage.enabled()) + throw internal::ConfigError( + "two_stage=true needs a two-stage bundle (trtmc ltx2 build --two-stage)"); const std::string prompt(request.prompt); - const int32_t steps = static_cast(options_.sigmas.size()) - 1; - const auto video_count = - static_cast(options_.video_tokens()) * options_.latent_channels; - const auto audio_count = - static_cast(options_.audio_frames) * options_.audio_latent_channels; + const auto& stage1_sigmas = two_stage ? options_.two_stage.sigmas : options_.sigmas; + const int32_t steps = + static_cast(stage1_sigmas.size()) - 1 + + (two_stage ? static_cast(options_.two_stage.stage2_sigmas.size()) - 1 : 0); // Ranks finish loading their engines at different times (the ranks load different decoders); // start together so the first collective does not charge one rank's load to the generation. @@ -553,7 +678,7 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq std::ostringstream detail; detail << "world_size=" << distributed_.world_size << " frames=" << options_.video_frames << " width=" << options_.video_width << " height=" << options_.video_height - << " steps=" << steps; + << " steps=" << steps << " two_stage=" << (two_stage ? 1 : 0); progress_.start(detail.str()); progress_.emit("encode_begin"); } @@ -561,48 +686,39 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq const auto t_text = Clock::now(); progress_.emit("encode_end"); - std::vector video(video_count); - std::vector audio(audio_count); - if (const char* path = std::getenv("TRTMC_LTX2_INITIAL_LATENTS"); - path != nullptr && *path != '\0') { - const auto values = read_f32_file(path); - if (values.size() != video_count + audio_count) - throw std::runtime_error( - "TRTMC_LTX2_INITIAL_LATENTS must hold the packed video then audio noise"); - std::copy_n(values.begin(), video_count, video.begin()); - std::copy_n(values.begin() + static_cast(video_count), audio_count, - audio.begin()); - } else { - std::mt19937 generator(static_cast(seed)); - ltx2::LibstdcxxNormalFloat normal; - for (auto& v : video) - v = normal(generator); - for (auto& v : audio) - v = normal(generator); - } - + auto noise = initial_noise(seed, two_stage); + auto& video = noise.video; + auto& audio = noise.audio; + const int64_t stage1_tokens = two_stage + ? options_.two_stage.video_tokens(options_.latent_frames) + : options_.video_tokens(); progress_.emit("denoise_begin"); - std::vector step_ms; - std::vector video_v; - std::vector audio_v; - for (int32_t step = 0; step < steps; ++step) { - const auto step_start = Clock::now(); - const float sigma = options_.sigmas[static_cast(step)]; - const float sigma_next = options_.sigmas[static_cast(step) + 1]; - run_dit(video, audio, text, sigma * 1000.0F, video_v, audio_v); - ltx2_euler_step(video, video_v, sigma, sigma_next); - ltx2_euler_step(audio, audio_v, sigma, sigma_next); - step_ms.push_back(elapsed_ms(step_start, Clock::now())); - if (progress_.enabled()) { - std::ostringstream detail; - detail << "step=" << (step + 1) << "/" << steps << " step_ms=" << std::fixed - << std::setprecision(3) << step_ms.back(); - progress_.emit("step", detail.str()); - } + StageTimes stage1; + StageTimes stage2; + denoise(video, audio, text, stage1_sigmas, stage1_tokens, "stage1", stage1); + double upsample_ms = 0.0; + if (two_stage) { + const auto up_start = Clock::now(); + if (distributed_.rank == 0) + maybe_dump(video, audio, ".stage1"); + video = upsample(video, stage1_tokens); + if (distributed_.rank == 0) + maybe_dump(video, audio, ".upsampled"); + // diffusers prepare_latents / prepare_audio_latents: noise_scale * noise + (1 - + // noise_scale) * x. + const float s = options_.two_stage.noise_scale; + ltx2_renoise(video, noise.video_stage2, s); + ltx2_renoise(audio, noise.audio_stage2, s); + upsample_ms = elapsed_ms(up_start, Clock::now()); + progress_.emit("upsample_end"); + denoise(video, audio, text, options_.two_stage.stage2_sigmas, options_.video_tokens(), + "stage2", stage2); } const auto t_denoise = Clock::now(); progress_.emit("denoise_end"); replace_final_latents(video, audio); + std::vector step_ms = stage1.step_ms; + step_ms.insert(step_ms.end(), stage2.step_ms.begin(), stage2.step_ms.end()); const bool tiled = options_.vae_tiling.enabled(); if (distributed_.world_size > 1 && distributed_.rank != 0 && !tiled) { @@ -654,7 +770,11 @@ internal::AudioVideoResult LTX2Pipeline::run(const internal::TextToAudioVideoReq << "[ltx2-perf-json] {\"world_size\":" << distributed_.world_size << ",\"text_encode_ms\":" << elapsed_ms(t_start, t_text) << ",\"denoise_ms\":" << elapsed_ms(t_text, t_denoise) - << ",\"median_step_ms\":" << median << ",\"decode_ms\":" + << ",\"median_step_ms\":" << median << ",\"two_stage\":" << (two_stage ? 1 : 0) + << ",\"stage1_ms\":" << stage1.total_ms + << ",\"stage1_median_step_ms\":" << stage1.median_ms() + << ",\"upsample_ms\":" << upsample_ms << ",\"stage2_ms\":" << stage2.total_ms + << ",\"stage2_median_step_ms\":" << stage2.median_ms() << ",\"decode_ms\":" << elapsed_ms(t_denoise, t_audio) // Untiled: the video then the audio decode. Tiled: the audio overlaps the blend // (one device) or runs on the audio rank, so the video path spans the phase. diff --git a/families/ltx2/runtime/pipeline.h b/families/ltx2/runtime/pipeline.h index 24f24f33b5..c3a5a6e4c9 100644 --- a/families/ltx2/runtime/pipeline.h +++ b/families/ltx2/runtime/pipeline.h @@ -12,6 +12,7 @@ // vae.plan video VAE decoder -> RGB frames (whole video, or one tile shape when the // bundle carries a tile plan; context-parallel ranks decode disjoint tiles) // audio.plan audio VAE decoder + vocoder with BWE -> 48 kHz stereo +// latent_upsampler.plan (two-stage bundles) 2x spatial latent upsampler between the stages #include "families/ltx2/runtime/distributed_runtime.h" #include "families/ltx2/runtime/progress_log.h" @@ -29,6 +30,22 @@ namespace trtmc { +// Two-stage pipeline of a bundle built with `trtmc ltx2 build --two-stage`: stage 1 denoises the +// half-resolution grid, latent_upsampler.plan doubles it, stage 2 re-noises the video and audio +// latents to noise_scale and refines at full resolution. +struct LTX2TwoStage { + int32_t latent_height{0}; + int32_t latent_width{0}; + std::vector sigmas; // stage 1 schedule including the terminal 0 + std::vector stage2_sigmas; // stage 2 schedule including the terminal 0 + float noise_scale{0.0F}; + + bool enabled() const { return latent_height > 0; } + int64_t video_tokens(int32_t latent_frames) const { + return int64_t(latent_frames) * latent_height * latent_width; + } +}; + struct LTX2Options { int32_t video_frames{121}; int32_t video_height{544}; @@ -50,6 +67,7 @@ struct LTX2Options { ltx2::VaeTilePlan vae_tiling; // Waveform shape of audio.plan; lets another rank decode the audio for rank 0. std::vector audio_waveform_shape; + LTX2TwoStage two_stage; int64_t video_tokens() const { return int64_t(latent_frames) * latent_height * latent_width; } // Rank that decodes the audio: the last context-parallel rank when the tiles spread the @@ -79,7 +97,8 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi LTX2Pipeline(std::unique_ptr text_encoder, std::unique_ptr denoiser, std::unique_ptr vae, std::unique_ptr audio, LTX2Options options, std::shared_ptr tokenizer, - LTX2DistributedContext distributed = {}); + LTX2DistributedContext distributed = {}, + std::unique_ptr upsampler = nullptr); ~LTX2Pipeline() override; const char* task() const noexcept override { return ITextToAudioVideo::kTask.data(); } @@ -103,11 +122,31 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi int32_t tiles{0}; }; + // Initial noise: stage 1 video then audio; for two-stage runs the stage 2 re-noise draws + // (video at full resolution, then audio) continue the same stream. + struct Noise { + std::vector video; + std::vector audio; + std::vector video_stage2; + std::vector audio_stage2; + }; + + struct StageTimes { + std::vector step_ms; + double total_ms{0.0}; + double median_ms() const; + }; + private: TextContext encode(const std::string& text); void run_dit(const std::vector& video, const std::vector& audio, - const TextContext& text, float timestep, std::vector& video_out, - std::vector& audio_out); + const TextContext& text, float timestep, int64_t video_tokens, + std::vector& video_out, std::vector& audio_out); + Noise initial_noise(int64_t seed, bool two_stage) const; + void denoise(std::vector& video, std::vector& audio, const TextContext& text, + const std::vector& sigmas, int64_t video_tokens, const char* stage, + StageTimes& times); + std::vector upsample(const std::vector& video, int64_t video_tokens); std::vector decode_video(const std::vector& video_latents); std::vector decode_audio(const std::vector& audio_latents); Decoded decode_untiled(const std::vector& video, const std::vector& audio); @@ -123,6 +162,7 @@ class LTX2Pipeline final : public internal::IModel, public internal::ITextToAudi std::unique_ptr denoiser_; std::unique_ptr vae_; std::unique_ptr audio_; + std::unique_ptr upsampler_; LTX2Options options_; std::shared_ptr tokenizer_; LTX2ProgressLog progress_; diff --git a/families/ltx2/runtime/plugin.cpp b/families/ltx2/runtime/plugin.cpp index 997847efbf..7cc2bad429 100644 --- a/families/ltx2/runtime/plugin.cpp +++ b/families/ltx2/runtime/plugin.cpp @@ -94,6 +94,10 @@ extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext& context audio->tensor_shape("waveform") != options.audio_waveform_shape) throw std::runtime_error( "LTX-2.5 audio.plan does not match runtime.json audio_waveform_shape"); + // Every rank upsamples the stage 1 latents itself (a small plan; no transfer needed). + std::unique_ptr upsampler; + if (options.two_stage.enabled()) + upsampler = ltx2::load(context.backend, context.reader, "latent_upsampler.plan", plain); const auto tokenizer_data = ltx2::require_section(context.reader, "tokenizer.json"); std::shared_ptr tokenizer = CreateLtx2BpeTokenizer( tokenizer_data.data(), tokenizer_data.size(), /*add_special_tokens=*/false); @@ -102,5 +106,6 @@ extern "C" trtmc::ITask* trtmc_create_family(const trtmc::FamilyContext& context return new LTX2Pipeline( std::move(text), std::move(denoiser), std::move(vae), std::move(audio), std::move(options), std::move(tokenizer), - LTX2DistributedContext{group.owner, group.rank, group.world_size, group.channel}); + LTX2DistributedContext{group.owner, group.rank, group.world_size, group.channel}, + std::move(upsampler)); } diff --git a/families/ltx2/runtime/runtime_math.h b/families/ltx2/runtime/runtime_math.h index 87c8b8f729..eef6eb88ac 100644 --- a/families/ltx2/runtime/runtime_math.h +++ b/families/ltx2/runtime/runtime_math.h @@ -26,6 +26,17 @@ inline void ltx2_euler_step(std::vector& x, const std::vector& v, x[i] = x[i] + dt * v[i]; } +// Two-stage re-noise (diffusers LTX2Pipeline._create_noised_state): +// x = noise_scale * noise + (1 - noise_scale) * x. +inline void ltx2_renoise(std::vector& x, const std::vector& noise, + float noise_scale) { + if (x.size() != noise.size()) + throw std::runtime_error("LTX-2.5 re-noise: latent and noise sizes differ"); + const float keep = 1.0F - noise_scale; + for (std::size_t i = 0; i < x.size(); ++i) + x[i] = noise_scale * noise[i] + keep * x[i]; +} + // Gemma prompt ids for one prompt: the tokenizer ids (no special tokens) right-truncated to // seq_len and left-padded with pad_id, as LTX2Pipeline tokenizes with padding_side="left". // mask is 1 on tokens and 0 on padding. diff --git a/families/ltx2/tests/cp_tiny_prep.py b/families/ltx2/tests/cp_tiny_prep.py index 2e085d0b08..ea31252e16 100644 --- a/families/ltx2/tests/cp_tiny_prep.py +++ b/families/ltx2/tests/cp_tiny_prep.py @@ -54,6 +54,16 @@ def main() -> int: arrays[f"{case}_ref_audio"] = ra.numpy() arrays[f"{case}_stg_keep"] = np.asarray(stg, np.float32) arrays[f"{case}_av_keep"] = np.asarray(av, np.float32) + # Two-stage plans: one CP plan serving the full grid and the half-resolution stage 1 grid. + small = tp.half_grid(shape) + (out / "single_small.plan").write_bytes(build_dit_engine(folder, small, cp_size=1, stg_blocks=(tp.STG_BLOCK,))) + (out / f"cp{cp}_two_grid.plan").write_bytes(build_dit_engine(folder, shape, cp_size=cp, stg_blocks=(tp.STG_BLOCK,), + extra_shapes=(small,))) + small_inp = tp._inputs(small, seed=7) + rv, ra = tp._reference(model, small, small_inp, torch.float32) + arrays.update({f"small_{k}": v.float().numpy() for k, v in small_inp.items()}) + arrays["small_ref_video"] = rv.numpy() + arrays["small_ref_audio"] = ra.numpy() np.savez(out / "reference.npz", **arrays) print(f"prepared {out} (video tokens {shape.video_tokens}, audio {sa}, cp {cp})", flush=True) return 0 diff --git a/families/ltx2/tests/cpp/test_runtime_contract.cpp b/families/ltx2/tests/cpp/test_runtime_contract.cpp index 27ff311aa8..1f712c870a 100644 --- a/families/ltx2/tests/cpp/test_runtime_contract.cpp +++ b/families/ltx2/tests/cpp/test_runtime_contract.cpp @@ -47,6 +47,21 @@ void test_euler_step_matches_flow_match_euler() { check(threw, "euler step rejects mismatched sizes"); } +void test_two_stage_renoise() { + // diffusers _create_noised_state: noise_scale * noise + (1 - noise_scale) * latents. + std::vector x{1.0F, -2.0F}; + trtmc::ltx2_renoise(x, {0.5F, 4.0F}, 0.25F); + check(x[0] == 0.25F * 0.5F + 0.75F * 1.0F && x[1] == 0.25F * 4.0F + 0.75F * -2.0F, + "re-noise mixes noise and latents"); + bool threw = false; + try { + trtmc::ltx2_renoise(x, {1.0F}, 0.5F); + } catch (const std::runtime_error&) { + threw = true; + } + check(threw, "re-noise rejects mismatched sizes"); +} + void test_prompt_ids_left_pad_and_truncate() { std::vector ids; std::vector mask; @@ -184,6 +199,7 @@ void test_vae_tile_latents_and_validation() { int main() { test_euler_step_matches_flow_match_euler(); + test_two_stage_renoise(); test_prompt_ids_left_pad_and_truncate(); test_interleave_stereo(); test_progress_line_format(); diff --git a/families/ltx2/tests/dist_dit_cp_check.py b/families/ltx2/tests/dist_dit_cp_check.py index ac9ab74160..9eb2ea35e7 100644 --- a/families/ltx2/tests/dist_dit_cp_check.py +++ b/families/ltx2/tests/dist_dit_cp_check.py @@ -7,7 +7,9 @@ after ``cp_tiny_prep`` wrote the plans and references into ``PREP_DIR``. Every rank runs the CP plan with its NCCL communicator (a hung collective aborts the communicator instead of blocking), runs the single-device plan, and compares the full outputs with the single-device plan and with -diffusers, shard by shard. Writes ``cp_rank.json``; exits non-zero on failure. +diffusers, shard by shard. The two-stage CP plan (full grid plus the half-resolution grid, run-time +token count) is checked the same way at both grids. Writes ``cp_rank.json``; exits non-zero on +failure. """ from __future__ import annotations @@ -19,6 +21,27 @@ from families.ltx2.tests import conftest # noqa: F401 - binds the TensorRT backend +INPUTS = ("video_latent", "audio_latent", "video_context", "audio_context", "timestep") + + +def _compare(np, cosine, world: int, got: dict, one: dict, ref_video, ref_audio) -> dict: + gv, ga = got["video_velocity"], got["audio_velocity"] + s = gv.shape[1] // world + rec = { + "video_tokens": int(gv.shape[1]), + "video_cos_per_shard_vs_single": [cosine(gv[:, r * s:(r + 1) * s], one["video_velocity"][:, r * s:(r + 1) * s]) + for r in range(world)], + "audio_cos_vs_single": cosine(ga, one["audio_velocity"]), + "video_cos_vs_diffusers": cosine(gv, ref_video), + "audio_cos_vs_diffusers": cosine(ga, ref_audio), + "max_abs_vs_single": float(np.abs(gv - one["video_velocity"]).max()), + "finite": bool(np.isfinite(gv).all() and np.isfinite(ga).all()), + } + rec["ok"] = bool(rec["finite"] and min(rec["video_cos_per_shard_vs_single"]) > 0.9999 + and rec["audio_cos_vs_single"] > 0.9999 and rec["video_cos_vs_diffusers"] > 0.999 + and rec["audio_cos_vs_diffusers"] > 0.999) + return rec + def main() -> int: import numpy as np @@ -33,37 +56,35 @@ def main() -> int: comm = NcclComm() ref = np.load(prep / "reference.npz") cp_engine = NpEngine((prep / f"cp{comm.world}.plan").read_bytes(), comm.capsule(), on_timeout=comm.abort) + two_grid = NpEngine((prep / f"cp{comm.world}_two_grid.plan").read_bytes(), comm.capsule(), + on_timeout=comm.abort) single = NpEngine((prep / "single.plan").read_bytes()) + single_small = NpEngine((prep / "single_small.plan").read_bytes()) report = {"rank": comm.rank, "world": comm.world, "cases": {}} + base = {k: ref[k] for k in INPUTS} + plain = dict(base, stg_keep=ref["plain_stg_keep"], av_keep=ref["plain_av_keep"]) + small = dict({k: ref[f"small_{k}"] for k in INPUTS}, stg_keep=ref["plain_stg_keep"], + av_keep=ref["plain_av_keep"]) + cases = [(case, cp_engine, single, dict(base, stg_keep=ref[f"{case}_stg_keep"], av_keep=ref[f"{case}_av_keep"]), + case) for case in ("plain", "stg_mixed", "isolated")] + # The two-stage plan switches grids on one context: full, half resolution, full again. + cases += [("two_grid_full", two_grid, single, plain, "plain"), + ("two_grid_small", two_grid, single_small, small, "small"), + ("two_grid_full_again", two_grid, single, plain, "plain")] ok = True - base = {k: ref[k] for k in ("video_latent", "audio_latent", "video_context", "audio_context", "timestep")} - for case in ("plain", "stg_mixed", "isolated"): - feed = dict(base, stg_keep=ref[f"{case}_stg_keep"], av_keep=ref[f"{case}_av_keep"]) + for name, engine, reference_plan, feed, ref_key in cases: try: - got = cp_engine(feed, timeout_s=120) + got = engine(feed, timeout_s=120) except TimeoutError as exc: - print(f"[rank {comm.rank}] {case}: {exc}; communicator aborted", flush=True) - report["cases"][case] = {"ok": False, "error": str(exc)} + print(f"[rank {comm.rank}] {name}: {exc}; communicator aborted", flush=True) + report["cases"][name] = {"ok": False, "error": str(exc)} ok = False break - one = single(feed) - gv, ga = got["video_velocity"], got["audio_velocity"] - s = gv.shape[1] // comm.world - rec = { - "video_cos_per_shard_vs_single": [cosine(gv[:, r * s:(r + 1) * s], one["video_velocity"][:, r * s:(r + 1) * s]) - for r in range(comm.world)], - "audio_cos_vs_single": cosine(ga, one["audio_velocity"]), - "video_cos_vs_diffusers": cosine(gv, ref[f"{case}_ref_video"]), - "audio_cos_vs_diffusers": cosine(ga, ref[f"{case}_ref_audio"]), - "max_abs_vs_single": float(np.abs(gv - one["video_velocity"]).max()), - "finite": bool(np.isfinite(gv).all() and np.isfinite(ga).all()), - } - rec["ok"] = bool(rec["finite"] and min(rec["video_cos_per_shard_vs_single"]) > 0.9999 - and rec["audio_cos_vs_single"] > 0.9999 and rec["video_cos_vs_diffusers"] > 0.999 - and rec["audio_cos_vs_diffusers"] > 0.999) + rec = _compare(np, cosine, comm.world, got, reference_plan(feed), ref[f"{ref_key}_ref_video"], + ref[f"{ref_key}_ref_audio"]) ok &= rec["ok"] - report["cases"][case] = rec - print(f"[rank {comm.rank}] {case}: {json.dumps(rec)}", flush=True) + report["cases"][name] = rec + print(f"[rank {comm.rank}] {name}: {json.dumps(rec)}", flush=True) (prep / f"cp_rank{comm.rank}.json").write_text(json.dumps(report, indent=1), encoding="utf-8") if comm.comm: comm.destroy() diff --git a/families/ltx2/tests/engine_runner.py b/families/ltx2/tests/engine_runner.py index b2c6319440..0be7937424 100644 --- a/families/ltx2/tests/engine_runner.py +++ b/families/ltx2/tests/engine_runner.py @@ -43,18 +43,22 @@ def run_plan(plan: bytes, inputs: dict[str, torch.Tensor], communicator=None) -> def _execute(engine, context, inputs, stream) -> dict[str, torch.Tensor]: keep = [] outputs: dict[str, torch.Tensor] = {} - for i in range(engine.num_io_tensors): - name = engine.get_tensor_name(i) + names = [engine.get_tensor_name(i) for i in range(engine.num_io_tensors)] + is_input = {name: engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT for name in names} + for name in (n for n in names if is_input[n]): # inputs first: they fix run-time dimensions dtype = _TORCH[engine.get_tensor_dtype(name)] shape = tuple(engine.get_tensor_shape(name)) - if engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: - t = inputs[name].to(device="cuda", dtype=dtype).contiguous() - if tuple(t.shape) != shape: - raise ValueError(f"{name}: shape {tuple(t.shape)} != engine {shape}") - keep.append(t) - else: - t = torch.empty(shape, dtype=dtype, device="cuda") - outputs[name] = t + t = inputs[name].to(device="cuda", dtype=dtype).contiguous() + if -1 in shape: + context.set_input_shape(name, tuple(t.shape)) + elif tuple(t.shape) != shape: + raise ValueError(f"{name}: shape {tuple(t.shape)} != engine {shape}") + keep.append(t) + context.set_tensor_address(name, t.data_ptr()) + for name in (n for n in names if not is_input[n]): + t = torch.empty(tuple(context.get_tensor_shape(name)), dtype=_TORCH[engine.get_tensor_dtype(name)], + device="cuda") + outputs[name] = t context.set_tensor_address(name, t.data_ptr()) torch.cuda.current_stream().synchronize() if not context.execute_async_v3(stream.cuda_stream): diff --git a/families/ltx2/tests/np_engine.py b/families/ltx2/tests/np_engine.py index b384240cca..d2867d75b7 100644 --- a/families/ltx2/tests/np_engine.py +++ b/families/ltx2/tests/np_engine.py @@ -42,23 +42,37 @@ def __init__(self, plan: bytes, communicator=None, on_timeout=None): self.stream = int(ck(rt.cudaStreamCreate())) self.on_timeout = on_timeout self.buffers: dict[str, tuple[int, tuple, object]] = {} - for i in range(self.engine.num_io_tensors): - name = self.engine.get_tensor_name(i) + names = [self.engine.get_tensor_name(i) for i in range(self.engine.num_io_tensors)] + self.inputs = [n for n in names if self.engine.get_tensor_mode(n) == trt.TensorIOMode.INPUT] + for name in self.inputs: # run-time dimensions: allocate the profile maximum shape = tuple(self.engine.get_tensor_shape(name)) - dtype = _NP[self.engine.get_tensor_dtype(name)] - nbytes = int(np.prod(shape)) * np.dtype(dtype).itemsize - ptr = int(ck(rt.cudaMalloc(max(nbytes, 1)))) - self.buffers[name] = (ptr, shape, dtype) - self.context.set_tensor_address(name, ptr) + if -1 in shape: + shape = tuple(self.engine.get_tensor_profile_shape(name, 0)[2]) + self.context.set_input_shape(name, shape) + self._allocate(name, shape) + for name in names: + if name not in self.inputs: + self._allocate(name, tuple(self.context.get_tensor_shape(name))) + + def _allocate(self, name: str, shape: tuple) -> None: + dtype = _NP[self.engine.get_tensor_dtype(name)] + nbytes = int(np.prod(shape)) * np.dtype(dtype).itemsize + ptr = int(ck(rt.cudaMalloc(max(nbytes, 1)))) + self.buffers[name] = (ptr, shape, dtype) + self.context.set_tensor_address(name, ptr) def __call__(self, inputs: dict[str, np.ndarray], *, timeout_s: float = 120.0) -> dict[str, np.ndarray]: outs = {} - for name, (ptr, shape, dtype) in self.buffers.items(): - if self.engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: - a = np.ascontiguousarray(np.asarray(inputs[name]).astype(dtype)) - if a.shape != shape: - raise ValueError(f"{name}: shape {a.shape} != engine {shape}") - ck(rt.cudaMemcpy(ptr, a.ctypes.data, a.nbytes, rt.cudaMemcpyKind.cudaMemcpyHostToDevice)) + for name in self.inputs: + ptr, shape, dtype = self.buffers[name] + a = np.ascontiguousarray(np.asarray(inputs[name]).astype(dtype)) + if -1 in tuple(self.engine.get_tensor_shape(name)): + if a.size > int(np.prod(shape)): + raise ValueError(f"{name}: shape {a.shape} exceeds the profile maximum {shape}") + self.context.set_input_shape(name, a.shape) + elif a.shape != shape: + raise ValueError(f"{name}: shape {a.shape} != engine {shape}") + ck(rt.cudaMemcpy(ptr, a.ctypes.data, a.nbytes, rt.cudaMemcpyKind.cudaMemcpyHostToDevice)) if not self.context.execute_async_v3(self.stream): raise RuntimeError("execute_async_v3 failed") deadline = time.monotonic() + timeout_s @@ -68,9 +82,9 @@ def __call__(self, inputs: dict[str, np.ndarray], *, timeout_s: float = 120.0) - self.on_timeout() raise TimeoutError(f"engine did not finish within {timeout_s} s") time.sleep(0.002) - for name, (ptr, shape, dtype) in self.buffers.items(): - if self.engine.get_tensor_mode(name) == trt.TensorIOMode.OUTPUT: - a = np.empty(shape, dtype=dtype) + for name, (ptr, _, dtype) in self.buffers.items(): + if name not in self.inputs: + a = np.empty(tuple(self.context.get_tensor_shape(name)), dtype=dtype) ck(rt.cudaMemcpy(a.ctypes.data, ptr, a.nbytes, rt.cudaMemcpyKind.cudaMemcpyDeviceToHost)) outs[name] = a.astype(np.float32) if dtype is not np.int32 else a return outs diff --git a/families/ltx2/tests/test_dit_parity.py b/families/ltx2/tests/test_dit_parity.py index 42b984d3fa..6362807c31 100644 --- a/families/ltx2/tests/test_dit_parity.py +++ b/families/ltx2/tests/test_dit_parity.py @@ -75,7 +75,7 @@ def tiny(tmp_path_factory): shape = DiTShape(batch=2, latent_frames=FRAMES, latent_height=LH, latent_width=LW, audio_frames=sa, text_len=TEXT, fps=FPS) plan = build_dit_engine(folder, shape, stg_blocks=(STG_BLOCK,)) - return model, plan, shape + return model, plan, shape, folder def _inputs(shape, seed=0): @@ -111,7 +111,7 @@ def _reference(model, shape, inp, dtype, *, stg_mask=None, isolate=False): @pytest.mark.parametrize("case", ["plain", "stg_mixed", "isolated"]) def test_dit_tiny_parity(tiny, case) -> None: - model, plan, shape = tiny + model, plan, shape, _ = tiny inp = _inputs(shape) stg = torch.ones(shape.batch) av = torch.ones(shape.batch) @@ -158,3 +158,36 @@ def test_rope_grids_match_diffusers() -> None: # diffusers: [B, H, T, r]; ours: [T, H, r] assert torch.allclose(torch.from_numpy(cos).permute(1, 0, 2), cos_ref[0], atol=2e-6), key assert torch.allclose(torch.from_numpy(sin).permute(1, 0, 2), sin_ref[0], atol=2e-6), key + + +def half_grid(shape): + """The two-stage stage 1 grid: half the latent rows and columns.""" + from dataclasses import replace + + return replace(shape, latent_height=shape.latent_height // 2, latent_width=shape.latent_width // 2) + + +def test_dit_two_grid_plan_tiny_parity(tiny) -> None: + """One plan serving the full grid and the half-resolution stage 1 grid (run-time token count).""" + from families.ltx2.dit_builder import build_dit_engine + from families.ltx2.tests.engine_runner import Engine + + model, static_plan, shape, folder = tiny + small = half_grid(shape) + assert small.video_tokens != shape.video_tokens + engine = Engine(build_dit_engine(folder, shape, stg_blocks=(STG_BLOCK,), extra_shapes=(small,))) + keep = {"stg_keep": torch.ones(shape.batch), "av_keep": torch.ones(shape.batch)} + static = run_plan(static_plan, {**_inputs(shape), **keep}) + for grid in (shape, small, shape): # switch grids back and forth on one context + inp = _inputs(grid, seed=grid.video_tokens) + got = engine({**inp, **keep}) + assert tuple(got["video_velocity"].shape) == (grid.batch, grid.video_tokens, 16) + rv, ra = _reference(model, grid, inp, torch.float32) + cv = cosine(got["video_velocity"].cpu(), rv) + ca = cosine(got["audio_velocity"].cpu(), ra) + print(f"two-grid plan at {grid.video_tokens} tokens vs fp32: video cos {cv:.6f} | audio cos {ca:.6f}") + assert cv > 0.999 and ca > 0.999 + got = engine({**_inputs(shape), **keep}) + c = cosine(got["video_velocity"].cpu(), static["video_velocity"].cpu()) + print(f"two-grid plan vs the static plan on the full grid: video cos {c:.6f}") + assert c > 0.9999 diff --git a/families/ltx2/tests/test_model_contract.py b/families/ltx2/tests/test_model_contract.py index fba44587c5..7ed87e0789 100644 --- a/families/ltx2/tests/test_model_contract.py +++ b/families/ltx2/tests/test_model_contract.py @@ -64,6 +64,20 @@ def test_build_rejects_unsupported_requests(overrides: dict, error: type) -> Non model.build(_request(**overrides), writer=None) +def test_stage_two_schedule_is_the_distilled_refinement_tail() -> None: + # diffusers STAGE_2_DISTILLED_SIGMA_VALUES: the last three distilled sigmas. + assert model.STAGE_2_DISTILLED_SIGMAS == (0.909375, 0.725, 0.421875) + assert model.STAGE_2_DISTILLED_SIGMAS == model.DISTILLED_SIGMAS[-3:] + + +def test_two_stage_grid_needs_64_aligned_sizes_and_the_upsampler(tmp_path) -> None: + shape = SimpleNamespace(latent_height=22, latent_width=40) + with pytest.raises(ValueError, match="divisible by 64"): + model._stage1_shape(tmp_path, shape, 672, 1280, 32) + with pytest.raises(FileNotFoundError, match="latent_upsampler"): + model._stage1_shape(tmp_path, shape, 704, 1280, 32) + + def test_distilled_schedule_is_the_eight_step_checkpoint_schedule() -> None: assert len(model.DISTILLED_SIGMAS) == 8 assert model.DISTILLED_SIGMAS[0] == 1.0 diff --git a/families/ltx2/tests/test_upsampler_parity.py b/families/ltx2/tests/test_upsampler_parity.py new file mode 100644 index 0000000000..a4c997a94c --- /dev/null +++ b/families/ltx2/tests/test_upsampler_parity.py @@ -0,0 +1,77 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tiny-random parity: LTX-2.5 latent upsampler engine vs diffusers ``LTX2LatentUpsamplerModel``. + +The engine takes the packed, normalized stage 1 latents and returns packed, normalized latents on +the 2x grid; the reference denormalizes, runs the diffusers model (``LTX2LatentUpsamplePipeline`` +with ``latents_normalized=False`` semantics) and normalizes again, as stage 2's ``prepare_latents`` +does. Both upsampler variants are covered (``upsampler.0`` and the 2x rational resampler). +""" + +from __future__ import annotations + +import json + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("tensorrt") +if not torch.cuda.is_available(): + pytest.skip("CUDA is required for LTX-2.5 engine parity tests", allow_module_level=True) +safetensors_torch = pytest.importorskip("safetensors.torch") + +from families.ltx2.tests.engine_runner import cosine, rel_l2, run_plan # noqa: E402 + +C = 16 +F, H, W = 3, 2, 3 + + +def _tiny(tmp_path, rational: bool): + from diffusers.pipelines.ltx2.latent_upsampler import LTX2LatentUpsamplerModel + + config = {"in_channels": C, "mid_channels": 64, "num_blocks_per_stage": 2, "dims": 3, + "spatial_upsample": True, "temporal_upsample": False, "rational_spatial_scale": 2.0, + "use_rational_resampler": rational} + model = LTX2LatentUpsamplerModel(**config).eval() + gen = torch.Generator().manual_seed(3) + with torch.no_grad(): + for name, p in model.named_parameters(): + if "norm" in name: + p.copy_((1.0 if name.endswith("weight") else 0.0) + 0.1 * torch.randn(p.shape, generator=gen)) + elif p.ndim >= 2: + p.copy_(torch.randn(p.shape, generator=gen) * (1.0 / p[0].numel()) ** 0.5) + else: + p.copy_(0.05 * torch.randn(p.shape, generator=gen)) + folder = tmp_path / "latent_upsampler" + folder.mkdir() + safetensors_torch.save_file({k: v.contiguous() for k, v in model.state_dict().items()}, + str(folder / "diffusion_pytorch_model.safetensors")) + (folder / "config.json").write_text(json.dumps(config), encoding="utf-8") + vae = tmp_path / "vae" + vae.mkdir() + stats = {"latents_mean": 0.3 * torch.randn(C, generator=gen), "latents_std": 0.5 + torch.rand(C, generator=gen)} + safetensors_torch.save_file(stats, str(vae / "diffusion_pytorch_model.safetensors")) + (vae / "config.json").write_text(json.dumps({"scaling_factor": 1.0}), encoding="utf-8") + return model, folder, vae, stats + + +@pytest.mark.parametrize("rational", [False, True]) +def test_latent_upsampler_tiny_parity(tmp_path, rational: bool) -> None: + from families.ltx2.upsampler_builder import build_latent_upsampler_engine + + model, folder, vae, stats = _tiny(tmp_path, rational) + plan = build_latent_upsampler_engine(folder, vae, latent_frames=F, latent_height=H, latent_width=W) + packed = torch.randn(1, F * H * W, C, generator=torch.Generator().manual_seed(11)) + got = run_plan(plan, {"latents": packed})["upsampled"].float().cpu() + assert tuple(got.shape) == (1, F * 2 * H * 2 * W, C) + mean, std = stats["latents_mean"].view(1, -1, 1, 1, 1), stats["latents_std"].view(1, -1, 1, 1, 1) + z = packed.reshape(1, F, H, W, C).permute(0, 4, 1, 2, 3) * std + mean + for dtype in (torch.float32, torch.bfloat16): + ref_model = model.to("cuda", dtype) + with torch.no_grad(): + up = ref_model(z.cuda().to(dtype)).float().cpu() + ref = ((up - mean) / std).permute(0, 2, 3, 4, 1).reshape(1, -1, C) + c = cosine(got, ref) + print(f"upsampler tiny (rational={rational}) vs {dtype}: cos {c:.6f} relL2 {rel_l2(got, ref):.3e}") + assert c > 0.999 diff --git a/families/ltx2/upsampler_builder.py b/families/ltx2/upsampler_builder.py new file mode 100644 index 0000000000..a454dccf58 --- /dev/null +++ b/families/ltx2/upsampler_builder.py @@ -0,0 +1,128 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""LTX-2.5 latent spatial upsampler (``LTX2LatentUpsamplerModel``) as a TensorRT plan. + +The two-stage pipeline denoises at half resolution, doubles the latent grid with this model and +refines at full resolution (diffusers ``LTX2LatentUpsamplePipeline`` with ``latents_normalized=False`` +on the stage 1 ``output_type="latent"`` latents, then stage 2's ``prepare_latents``). + +Engine I/O: + Inputs: + latents [1, F*H*W, 128] fp32 packed, normalized stage 1 video latents (DiT layout) + Outputs: + upsampled [1, F*2H*2W, 128] fp32 packed, normalized latents on the 2x spatial grid + +The engine denormalizes with the VAE's ``latents_mean`` / ``latents_std`` (``scaling_factor``) in +fp32, unpacks to ``[1, C, F, H, W]`` and runs the upsampler in bf16 (the precision diffusers runs it +in): zero-padded 3x3x3 convolutions, GroupNorm(32) with fp32 statistics, SiLU, residual blocks, +the per-frame 3x3 convolution + 2x pixel shuffle, residual blocks and the final convolution. It +then normalizes again (in fp32; diffusers renormalizes the bf16 tensor in bf16) and packs the +tokens. Temporal upsampling and rational scales other than 2 are not used by LTX-2.5 and are +rejected. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import tensorrt as trt + +from .checkpoint import Checkpoint +from .graph import Graph, build_plan, make_logger, new_network + +GROUP_NORM_EPS = 1e-5 + + +def _conv(g: Graph, x, weight: np.ndarray, bias: np.ndarray): + """Zero-padded ``nn.Conv3d`` (or a per-frame ``nn.Conv2d`` given a ``[O, I, 1, k, k]`` weight).""" + out_c, _, kt, kh, kw = (int(s) for s in weight.shape) + layer = g.net.add_convolution_nd(x, out_c, (kt, kh, kw), g.weights(weight, x.dtype), g.weights(bias, x.dtype)) + layer.stride_nd = (1, 1, 1) + layer.padding_nd = (kt // 2, kh // 2, kw // 2) + return layer.get_output(0) + + +def _group_norm(g: Graph, x, gamma: np.ndarray, beta: np.ndarray, groups: int = 32): + """``nn.GroupNorm`` with fp32 statistics and affine, rounded once to the input dtype.""" + b, c, f, h, w = (int(s) for s in x.shape) + xf = g.reshape(g.cast(x, trt.float32), (b, groups, (c // groups) * f * h * w)) + mean = g.reduce(xf, trt.ReduceOperation.AVG, 2) + centered = g.sub(xf, mean) + var = g.reduce(g.mul(centered, centered), trt.ReduceOperation.AVG, 2) + inv = g.unary(g.unary(g.add(var, g.scalar(GROUP_NORM_EPS, trt.float32, 3)), trt.UnaryOperation.SQRT), + trt.UnaryOperation.RECIP) + y = g.reshape(g.mul(centered, inv), (b, c, f, h, w)) + y = g.add(g.mul(y, g.const(gamma.reshape(1, c, 1, 1, 1), trt.float32)), + g.const(beta.reshape(1, c, 1, 1, 1), trt.float32)) + return g.cast(y, x.dtype) + + +def _res_block(g: Graph, ck: Checkpoint, p: str, x): + h = _conv(g, x, ck.get(f"{p}.conv1.weight"), ck.get(f"{p}.conv1.bias")) + h = g.silu(_group_norm(g, h, ck.get(f"{p}.norm1.weight", np.float32), ck.get(f"{p}.norm1.bias", np.float32))) + h = _conv(g, h, ck.get(f"{p}.conv2.weight"), ck.get(f"{p}.conv2.bias")) + h = _group_norm(g, h, ck.get(f"{p}.norm2.weight", np.float32), ck.get(f"{p}.norm2.bias", np.float32)) + return g.silu(g.add(h, x)) + + +def _pixel_shuffle_2d(g: Graph, x): + """``PixelShuffleND(2)`` per frame: ``[B, C*4, F, H, W]`` -> ``[B, C, F, 2H, 2W]``.""" + b, cc, f, h, w = (int(s) for s in x.shape) + c = cc // 4 + y = g.reshape(x, (b, c, 2, 2, f, h, w)) + return g.reshape(y, (b, c, f, 2 * h, 2 * w), first=(0, 1, 4, 5, 2, 6, 3)) + + +def _upsampler_weights(ck: Checkpoint, cfg: dict) -> tuple[np.ndarray, np.ndarray]: + """The 2x spatial upsampler convolution as a per-frame ``[O, I, 1, 3, 3]`` weight.""" + if cfg.get("use_rational_resampler", True): + if float(cfg.get("rational_spatial_scale", 2.0)) != 2.0: + raise NotImplementedError("only the 2x rational resampler (no blur downsample) is implemented") + prefix = "upsampler.conv" + else: + prefix = "upsampler.0" + weight = ck.get(f"{prefix}.weight") + return weight.reshape(weight.shape[0], weight.shape[1], 1, *weight.shape[2:]), ck.get(f"{prefix}.bias") + + +def build_latent_upsampler_engine(upsampler_dir: str | Path, vae_dir: str | Path, *, latent_frames: int, + latent_height: int, latent_width: int, verbose: bool = False) -> bytes: + ck = Checkpoint(upsampler_dir) + cfg = ck.config() + if int(cfg.get("dims", 3)) != 3: + raise NotImplementedError("the LTX-2.5 latent upsampler uses 3D convolutions (dims=3)") + if not cfg.get("spatial_upsample", True) or cfg.get("temporal_upsample", False): + raise NotImplementedError("only the spatial 2x latent upsampler is implemented") + vae = Checkpoint(vae_dir) + latent_c = int(cfg.get("in_channels", 128)) + blocks = int(cfg.get("num_blocks_per_stage", 4)) + scaling = float(vae.config().get("scaling_factor", 1.0)) + mean = vae.get("latents_mean", np.float32).reshape(1, 1, latent_c) + std = vae.get("latents_std", np.float32).reshape(1, 1, latent_c) + + builder, network = new_network(make_logger(verbose)) + g = Graph(network) + f, h, w = latent_frames, latent_height, latent_width + z = network.add_input("latents", trt.float32, (1, f * h * w, latent_c)) + x = g.add(g.mul(z, g.const(std / scaling, trt.float32)), g.const(mean, trt.float32)) + x = g.cast(g.reshape(g.transpose(x, (0, 2, 1)), (1, latent_c, f, h, w)), trt.bfloat16) + + x = _conv(g, x, ck.get("initial_conv.weight"), ck.get("initial_conv.bias")) + x = g.silu(_group_norm(g, x, ck.get("initial_norm.weight", np.float32), ck.get("initial_norm.bias", np.float32))) + for i in range(blocks): + x = _res_block(g, ck, f"res_blocks.{i}", x) + x = _pixel_shuffle_2d(g, _conv(g, x, *_upsampler_weights(ck, cfg))) + for i in range(blocks): + x = _res_block(g, ck, f"post_upsample_res_blocks.{i}", x) + x = _conv(g, x, ck.get("final_conv.weight"), ck.get("final_conv.bias")) + + tokens = f * 2 * h * 2 * w + y = g.reshape(g.cast(x, trt.float32), (1, latent_c, tokens), second=(0, 2, 1)) + y = g.mul(g.sub(y, g.const(mean, trt.float32)), g.const(scaling / std, trt.float32)) + g.mark_output(y, "upsampled", trt.float32) + print(f"[ltx2] Building latent upsampler engine (latent {f}x{h}x{w} -> {f}x{2 * h}x{2 * w}) ...", + file=sys.stderr) + return build_plan(builder, network, label="latent upsampler")