diff --git a/CMakeLists.txt b/CMakeLists.txt index 92e11c3ee0..adcd0b3749 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 @@ -221,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 ) @@ -304,7 +346,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 +357,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 +418,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 +582,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 +1138,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 +1190,103 @@ install(FILES DESTINATION ${CMAKE_INSTALL_DATADIR}/cmake/trtmc COMPONENT sdk ) + +# 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) + 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() + # 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(_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() + _trtmc_msvc_translate_option("${_option}" ${_language} _option) + if(_open STREQUAL "") + if(NOT _option STREQUAL "") + list(APPEND _translated "${_option}") + endif() + continue() + endif() + 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}") + 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/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/family_cli.cpp b/apps/cli/family_cli.cpp index 1d24435862..5a4d8fd572 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,102 @@ 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) + // 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(error) + ")"); + } +#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. + // _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) + 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(); + 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 +481,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,32 +495,46 @@ 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; } } // 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 { - 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..b815b2ff3a 100644 --- a/apps/cli/family_cli.h +++ b/apps/cli/family_cli.h @@ -8,12 +8,19 @@ #include #include #include +#include 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 = {}); +// 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/sdk_video.cpp b/apps/cli/sdk_video.cpp index 08696caf6e..c3b5abc182 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,51 @@ 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(); + 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), + static_cast(audio.channels), audio_path); + json["audio"] = audio_path; + 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; + 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 +156,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 +196,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")) diff --git a/apps/cli/tests/test_cli.cpp b/apps/cli/tests/test_cli.cpp index 65bb4f38d2..199ca06536 100644 --- a/apps/cli/tests/test_cli.cpp +++ b/apps/cli/tests/test_cli.cpp @@ -4,10 +4,12 @@ */ #include "cli/cli.h" +#include "cli/family_cli.h" #include "cli/io.h" #include "cli/sdk_dispatch.h" #include +#include #include #include #include @@ -20,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; @@ -39,6 +50,50 @@ 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 +} + +// 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)); @@ -322,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", @@ -394,6 +450,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/conanfile.py b/conanfile.py index 24c715a709..d47724a7ea 100644 --- a/conanfile.py +++ b/conanfile.py @@ -51,6 +51,9 @@ 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. @@ -62,10 +65,19 @@ def generate(self) -> None: for name in ( "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: 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 toolchain.generate() def build(self) -> None: @@ -73,7 +85,66 @@ 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)}" + ) + 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(): + 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/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"); 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..fa6d04c3a8 --- /dev/null +++ b/core/runtime/tests/test_dynamic_library.cpp @@ -0,0 +1,158 @@ +/* + * 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 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()); + + 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 + +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], partial); + } 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/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..e0a8bf07f3 --- /dev/null +++ b/families/ltx2/runtime/distributed_runtime.cpp @@ -0,0 +1,223 @@ +/* + * 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 +#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); + 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; +} + +} // 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..9320e6c9d5 --- /dev/null +++ b/families/ltx2/tests/dist_helpers.py @@ -0,0 +1,77 @@ +# 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") + 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: + 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..4aa3b31ad1 --- /dev/null +++ b/families/ltx2/tests/manifests/ltx25-distilled-cp2.json @@ -0,0 +1,22 @@ +{ + "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 + } + ], + "tensor_parallel_size": 1, + "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..f26fbfcc81 --- /dev/null +++ b/families/ltx2/tests/manifests/ltx25-distilled-l0.json @@ -0,0 +1,23 @@ +{ + "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 + } + ], + "tensor_parallel_size": 1, + "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..7c5acfd3ea --- /dev/null +++ b/families/ltx2/tests/test_e2e.py @@ -0,0 +1,474 @@ +# 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 + assert manifest["tensor_parallel_size"] == 1 + 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"), + tensor_parallel_size=int(manifest["tensor_parallel_size"]), + 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"): + 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", + "-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")) 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 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. | 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