diff --git a/CHANGELOG.md b/CHANGELOG.md index 2c2e416c7c..cb7bb85107 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,16 @@ # TensorRT OSS Release Changelog +## 11.4 GA - 2026-10-09 +- Parsers + - Added `IRefitterObserver` class and `IParser::setRefitObserver` to better handle refittable weights when parsing. + +- Plugins + - Added various C++20 updates to plugin source code. + +- Samples + - Moved samples/common files only relevant to trtexec to samples/trtexecCommon. + + ## 11.3 GA - 2026-09-22 - General - Updated default CUDA version to 13.4 diff --git a/CMakeLists.txt b/CMakeLists.txt index 790d8b9b2d..243614361d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -53,10 +53,6 @@ option(BUILD_PYTHON "Build TensorRT python bindings" OFF) option(TRT_SAFETY_INFERENCE_ONLY "Build only the safety inference components (no safety builders)" OFF) option(TRT_BUILD_TESTING "Build gtests for TensorRT components" OFF) option(TRT_BUILD_ALLOW_CCACHE "Allow ccache to be used for builds" ON) -set(TRT_BUILD_PRODUCT "" CACHE STRING - "TensorRT product package to import: enterprise, automotive, or safe_inference. Empty preserves legacy selection.") -set_property(CACHE TRT_BUILD_PRODUCT PROPERTY STRINGS - "" enterprise automotive safe_inference) if(TRT_SAFETY_INFERENCE_ONLY AND NOT BUILD_SAFE_SAMPLES) message(STATUS "TRT_SAFETY_INFERENCE_ONLY is ON, enabling BUILD_SAFE_SAMPLES") @@ -106,45 +102,12 @@ message(STATUS "Generating CUDA code for SMs: ${CMAKE_CUDA_ARCHITECTURES}") # OSS vendors third-party headers under third_party/. set(TRT_EXTERNALS_DIR ${CMAKE_CURRENT_SOURCE_DIR}/third_party) -set(_TRT_OSS_SUPPORTED_PRODUCTS enterprise automotive safe_inference) -if(NOT TRT_BUILD_PRODUCT STREQUAL "") - if(NOT "${TRT_BUILD_PRODUCT}" IN_LIST _TRT_OSS_SUPPORTED_PRODUCTS) - message(FATAL_ERROR - "Invalid TRT_BUILD_PRODUCT='${TRT_BUILD_PRODUCT}'. Expected one of: " - "enterprise, automotive, safe_inference.") +if(BUILD_SAFE_SAMPLES) + if(TRT_SAFETY_INFERENCE_ONLY) + find_package(TensorRT-SafeInference CONFIG REQUIRED) + else() + find_package(TensorRT-Automotive CONFIG REQUIRED) endif() - set(_TRT_OSS_SELECTED_PRODUCT "${TRT_BUILD_PRODUCT}") -elseif(TRT_SAFETY_INFERENCE_ONLY) - set(_TRT_OSS_SELECTED_PRODUCT safe_inference) -elseif(BUILD_SAFE_SAMPLES) - set(_TRT_OSS_SELECTED_PRODUCT automotive) -else() - set(_TRT_OSS_SELECTED_PRODUCT enterprise) -endif() - -if(TRT_SAFETY_INFERENCE_ONLY - AND NOT _TRT_OSS_SELECTED_PRODUCT STREQUAL "safe_inference") - message(FATAL_ERROR - "TRT_SAFETY_INFERENCE_ONLY=ON requires " - "TRT_BUILD_PRODUCT=safe_inference.") -endif() -if(_TRT_OSS_SELECTED_PRODUCT STREQUAL "safe_inference" - AND NOT TRT_SAFETY_INFERENCE_ONLY) - message(FATAL_ERROR - "TRT_BUILD_PRODUCT=safe_inference requires " - "TRT_SAFETY_INFERENCE_ONLY=ON.") -endif() -if(BUILD_SAFE_SAMPLES AND _TRT_OSS_SELECTED_PRODUCT STREQUAL "enterprise") - message(FATAL_ERROR - "BUILD_SAFE_SAMPLES=ON is not supported with " - "TRT_BUILD_PRODUCT=enterprise.") -endif() - -message(STATUS "Importing TensorRT product: ${_TRT_OSS_SELECTED_PRODUCT}") -if(_TRT_OSS_SELECTED_PRODUCT STREQUAL "safe_inference") - find_package(TensorRT-SafeInference CONFIG REQUIRED) -elseif(_TRT_OSS_SELECTED_PRODUCT STREQUAL "automotive") - find_package(TensorRT-Automotive CONFIG REQUIRED) else() find_package(TensorRT-Enterprise CONFIG REQUIRED) endif() @@ -207,18 +170,6 @@ if(TRT_SAFETY_INFERENCE_ONLY) Threads::Threads ) - # CUDA compiler identification receives these PDK dependencies through the - # QNX Safety toolchain. Add them here only for CXX-linked sample targets; - # CUDA-linked targets already inherit the toolchain flags. - if(CMAKE_SYSTEM_NAME STREQUAL "QNX" AND TRT_QNX_SAFE_PDK_LIBRARIES) - foreach(TRT_QNX_SAFE_PDK_LIBRARY IN LISTS TRT_QNX_SAFE_PDK_LIBRARIES) - target_link_libraries(trt_global_definitions INTERFACE - "$<$:${TRT_QNX_SAFE_PDK_LIBRARY}>") - endforeach() - target_link_options(trt_global_definitions INTERFACE - "$<$:-lslog2>") - endif() - if(NOT WIN32 AND NOT CMAKE_SYSTEM_NAME STREQUAL "QNX") target_link_libraries(trt_global_definitions INTERFACE dl rt) endif() @@ -314,7 +265,6 @@ if(BUILD_SAMPLES OR BUILD_SAFE_SAMPLES) # Map OSS option names to the internal names used by samples/CMakeLists.txt. set(TRT_BUILD_SAMPLES ${BUILD_SAMPLES}) - set(TRT_BUILD_TRTEXEC ${BUILD_SAMPLES}) # nvonnxparser is always available, either built from source or imported # from the TensorRT package. set(TRT_BUILD_ONNX_PARSER ON) @@ -325,19 +275,10 @@ if(BUILD_SAMPLES OR BUILD_SAFE_SAMPLES) option(TRT_BUILD_ENABLE_DLA "Build TensorRT with DLA features enabled." OFF) option(TRT_BUILD_ENABLE_UNIFIED_BUILDER "Build TensorRT with unified builder (safety) features enabled." ${BUILD_SAFE_SAMPLES}) - # The companion library a safety engine pairs with only exists in a TensorRT built with the unified - # builder, so the sample code that saves and loads one compiles only there. TensorRT forwards the same - # flag under the same name when it builds these samples itself. - target_compile_definitions(trt_global_definitions INTERFACE - ENABLE_UNIFIED_BUILDER=$ - ) - option(TRT_BUILD_ENABLE_MULTIDEVICE "Build TensorRT with multi-device support." OFF) set(TRT_BUILD_SAMPLES_LINK_STATIC_TRT OFF CACHE INTERNAL "") - if(NOT (TRT_SAFETY_INFERENCE_ONLY AND CMAKE_SYSTEM_NAME STREQUAL "QNX")) - add_subdirectory(shared) - endif() + add_subdirectory(shared) include(InstallUtils) diff --git a/README.md b/README.md index aa1b66a4a8..4cbb950adb 100644 --- a/README.md +++ b/README.md @@ -48,7 +48,7 @@ To build the TensorRT-OSS components, you will first need the following software **TensorRT GA build** -- TensorRT v11.3.0.99 +- TensorRT v11.4.0.106 - Available from direct download links listed below **System Packages** @@ -103,24 +103,24 @@ To build the TensorRT-OSS components, you will first need the following software Else download and extract the TensorRT GA build from [NVIDIA Developer Zone](https://developer.nvidia.com) with the direct links below: - - [TensorRT 11.3.0.99 for CUDA 13.4, Linux x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.3.0/tars/TensorRT-Enterprise-11.3.0.99-Linux-x86_64-cuda-13.4-Release-external.tar.zst) - - [TensorRT 11.3.0.99 for CUDA 12.9, Linux x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.3.0/tars/TensorRT-Enterprise-11.3.0.99-Linux-x86_64-cuda-12.9-Release-external.tar.zst) - - [TensorRT 11.3.0.99 for CUDA 13.4, Windows x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.3.0/zip/TensorRT-Enterprise-11.3.0.99-Windows-amd64-cuda-13.4-Release-external.zip) - - [TensorRT 11.3.0.99 for CUDA 12.9, Windows x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.3.0/zip/TensorRT-Enterprise-11.3.0.99-Windows-amd64-cuda-12.9-Release-external.zip) + - [TensorRT 11.4.0.106 for CUDA 13.4, Linux x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.4.0/tars/TensorRT-Enterprise-11.4.0.106-Linux-x86_64-cuda-13.4-Release-external.tar.zst) + - [TensorRT 11.4.0.106 for CUDA 12.9, Linux x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.4.0/tars/TensorRT-Enterprise-11.4.0.106-Linux-x86_64-cuda-12.9-Release-external.tar.zst) + - [TensorRT 11.4.0.106 for CUDA 13.4, Windows x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.4.0/zip/TensorRT-Enterprise-11.4.0.106-Windows-amd64-cuda-13.4-Release-external.zip) + - [TensorRT 11.4.0.106 for CUDA 12.9, Windows x86_64](https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.4.0/zip/TensorRT-Enterprise-11.4.0.106-Windows-amd64-cuda-12.9-Release-external.zip) **Example: Ubuntu 22.04 on x86-64 with cuda-13.4** ```bash cd ~/Downloads - tar --zstd -xvf TensorRT-Enterprise-11.3.0.99-Linux-x86_64-cuda-13.4-Release-external.tar.zst - export TRT_LIBPATH=`pwd`/TensorRT-11.3.0.99/lib + tar --zstd -xvf TensorRT-Enterprise-11.4.0.106-Linux-x86_64-cuda-13.4-Release-external.tar.zst + export TRT_LIBPATH=`pwd`/TensorRT-11.4.0.106/lib ``` **Example: Windows on x86-64 with cuda-12.9** ```powershell - Expand-Archive -Path TensorRT-Enterprise-11.3.0.99-Windows-amd64-cuda-12.9-Release-external.zip - $env:TRT_LIBPATH="$pwd\TensorRT-11.3.0.99\lib" + Expand-Archive -Path TensorRT-Enterprise-11.4.0.106-Windows-amd64-cuda-12.9-Release-external.zip + $env:TRT_LIBPATH="$pwd\TensorRT-11.4.0.106\lib" ``` ## Setting Up The Build Environment @@ -222,8 +222,7 @@ For Linux platforms, we recommend that you generate a docker container for build ```bash cd $TRT_OSSPATH mkdir -p build && cd build - cmake .. -DTRT_BUILD_PRODUCT=automotive -DCMAKE_PREFIX_PATH=$TRT_ROOT \ - -DCMAKE_TOOLCHAIN_FILE=$TRT_OSSPATH/cmake/toolchains/cmake_aarch64_cross.toolchain + cmake .. -DCMAKE_PREFIX_PATH=$TRT_ROOT -DCMAKE_TOOLCHAIN_FILE=$TRT_OSSPATH/cmake/toolchains/cmake_aarch64_cross.toolchain make -j$(nproc) ``` @@ -256,7 +255,6 @@ For Linux platforms, we recommend that you generate a docker container for build - `BUILD_SAMPLES`: Specify if the samples should be built, for example [`ON`] | `OFF`. - `BUILD_SAFE_SAMPLES`: Specify if safety samples should be built, for example [`ON`] | `OFF`. - `TRT_SAFETY_INFERENCE_ONLY`: Specify if only build the safety inference components, for example [`ON`] | `OFF`. If turned ON, all other components will be turned OFF except `BUILD_SAFE_SAMPLES`. - - `TRT_BUILD_PRODUCT`: Select the TensorRT product package to import: `enterprise`, `automotive`, or `safe_inference`. If omitted, the build will infer the product based on the values of `BUILD_SAFE_SAMPLES` and `TRT_SAFETY_INFERENCE_ONLY`. - `TRT_BUILD_ENABLE_MULTIDEVICE`: Enable the multi-device sample (`sampleDistCollective`). Use `-DTRT_BUILD_ENABLE_MULTIDEVICE=ON` to build it; requires [NCCL](https://developer.nvidia.com/nccl/nccl-download) >= v2.19, < v3.0. - `TRT_BUILD_TESTING` : Build gTests for samples. Requires [gtest](https://github.com/google/googletest) if available; otherwise fetches googletest at configure time. @@ -270,7 +268,6 @@ For Linux platforms, we recommend that you generate a docker container for build cd $TRT_OSSPATH mkdir -p build && cd build cmake .. -DBUILD_SAMPLES=ON -DBUILD_PLUGINS=OFF -DBUILD_PARSERS=OFF \ - -DTRT_BUILD_PRODUCT=automotive \ -DCMAKE_RUNTIME_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ -DCMAKE_ARCHIVE_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ @@ -304,8 +301,7 @@ For Linux platforms, we recommend that you generate a docker container for build -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=`pwd`/out \ -DCMAKE_ARCHIVE_OUTPUT_DIRECTORY=`pwd`/out \ -DCMAKE_TOOLCHAIN_FILE=$TRT_OSSPATH/cmake/toolchains/cmake_aarch64-native.toolchain \ - -DBUILD_SAMPLES=ON -DBUILD_PLUGINS=OFF -DBUILD_PARSERS=OFF \ - -DTRT_BUILD_PRODUCT=automotive + -DBUILD_SAMPLES=ON -DBUILD_PLUGINS=OFF -DBUILD_PARSERS=OFF make -j$(nproc) ``` @@ -369,13 +365,12 @@ For Linux platforms, we recommend that you generate a docker container for build mkdir -p build && cd build export CUDA_VERSION=13.4 export CUDA=cuda-$CUDA_VERSION - export CUDA_ROOT=/usr/local/cuda-$CUDA_VERSION + export CUDA_ROOT=/usr/local/cuda-safe-$CUDA_VERSION export QNX_BASE=/drive/toolchains/qnx_toolchain # Set to your QNX toolchain installation path export QNX_HOST=$QNX_BASE/host/linux/x86_64/ export QNX_TARGET=$QNX_BASE/target/qnx/ export PATH=$PATH:$QNX_HOST/usr/bin cmake .. -DBUILD_SAMPLES=ON -DBUILD_PLUGINS=OFF -DBUILD_PARSERS=OFF -DBUILD_SAFE_SAMPLES=OFF \ - -DTRT_BUILD_PRODUCT=automotive \ -DCMAKE_CUDA_COMPILER=$CUDA_ROOT/bin/nvcc \ -DCMAKE_RUNTIME_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ @@ -389,41 +384,6 @@ For Linux platforms, we recommend that you generate a docker container for build > NOTE: Set `QNX_BASE` to your QNX toolchain installation path. > If your CUDA version is not the same as in the example, set `CUDA_VERSION` (for examples that use it in multiple places) or add `-DCUDA_VERSION=` to the cmake command. - **Example: Cross-Compile for DOS7 QNX Safety (aarch64)** - - DOS7 QNX and QNX Safety use the same QNX 8 SDK through `QNX_HOST` and - `QNX_TARGET`. The QNX Safety build additionally uses SafeCUDA, the matching - TensorRT SafeInference package, and `cmake_qnx_safe.toolchain`. The standard - DriveOS QNX Safety build environment provides `PDK_TOP`; the toolchain uses - the libraries under `$PDK_TOP/drive-qnx-safety/lib-target`. - - ```bash - cd $TRT_OSSPATH - mkdir -p build && cd build - export CUDA_VERSION=13.4 - export CUDA=cuda-$CUDA_VERSION - export CUDA_ROOT=/usr/local/cuda-$CUDA_VERSION-safe - export QNX_BASE=/drive/toolchains/qnx_toolchain # Set to your QNX 8 toolchain installation path - export QNX_HOST=$QNX_BASE/host/linux/x86_64/ - export QNX_TARGET=$QNX_BASE/target/qnx/ - export PATH=$PATH:$QNX_HOST/usr/bin - cmake .. -DBUILD_SAMPLES=OFF -DBUILD_SAFE_SAMPLES=ON -DBUILD_PLUGINS=OFF -DBUILD_PARSERS=OFF \ - -DTRT_BUILD_PRODUCT=safe_inference \ - -DTRT_SAFETY_INFERENCE_ONLY=ON -DCMAKE_BUILD_TYPE=Release \ - -DCMAKE_RUNTIME_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ - -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ - -DCMAKE_ARCHIVE_OUTPUT_DIRECTORY=`pwd`/bin_dynamic_cross \ - -DCMAKE_PREFIX_PATH=$TRT_ROOT \ - -DTensorRT-SafeInference_DIR=$TRT_ROOT/cmake/TensorRT-SafeInference \ - -DCMAKE_TOOLCHAIN_FILE=$TRT_OSSPATH/cmake/toolchains/cmake_qnx_safe.toolchain \ - -DCUDA_VERSION=$CUDA_VERSION -DCMAKE_CUDA_COMPILER=$CUDA_ROOT/bin/nvcc \ - -DCMAKE_CUDA_ARCHITECTURES=110 - make -j$(nproc) - ``` - - > NOTE: Set `QNX_BASE` to the same QNX 8 SDK used for DOS7 QNX builds. The - > generated QNX Safety binaries are placed in `build/bin_dynamic_cross`. - # References ## TensorRT Resources diff --git a/VERSION b/VERSION index 0877031594..d3864c623b 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -11.3.0.99 +11.4.0.106 diff --git a/cmake/toolchains/cmake_qnx_safe.toolchain b/cmake/toolchains/cmake_qnx_safe.toolchain index 6985c2ce24..70d714858c 100644 --- a/cmake/toolchains/cmake_qnx_safe.toolchain +++ b/cmake/toolchains/cmake_qnx_safe.toolchain @@ -45,32 +45,12 @@ else() set(QNX_GCC_VERSION "12.2.0" CACHE STRING "") endif() -if(${QNX_VERSION} VERSION_GREATER_EQUAL "8.0.0") - # DOS7 QNX Safety samples are consumed directly from the documented build - # output. Do not embed build-host CUDA, PDK, or TensorRT paths in them. - set(CMAKE_SKIP_BUILD_RPATH TRUE CACHE BOOL - "Do not embed build-tree library paths in DOS7 QNX Safety outputs") -endif() - set(QNX_TRIPLE aarch64-unknown-nto-qnx${QNX_VERSION}) set(CMAKE_C_COMPILER ${QNX_HOST}/usr/bin/${QNX_TRIPLE}-gcc) set(CMAKE_CXX_COMPILER ${QNX_HOST}/usr/bin/${QNX_TRIPLE}-g++) set(CMAKE_CUDA_HOST_COMPILER ${QNX_HOST}/usr/bin/${QNX_TRIPLE}-g++ CACHE FILEPATH "") -# The DOS7 QNX Safety compiler supports optimization levels through O2. Keep -# Release builds within that supported range for every compiled language. -if(${QNX_VERSION} VERSION_GREATER_EQUAL "8.0.0") - if(NOT DEFINED CMAKE_CXX_FLAGS_RELEASE) - set(CMAKE_CXX_FLAGS_RELEASE "-O2 -DNDEBUG" CACHE STRING - "Flags used by the CXX compiler for Release builds") - endif() - if(NOT DEFINED CMAKE_CUDA_FLAGS_RELEASE) - set(CMAKE_CUDA_FLAGS_RELEASE "-O2 -DNDEBUG" CACHE STRING - "Flags used by the CUDA compiler for Release builds") - endif() -endif() - # These linker behaviors aren't setup by default for QNX # They work the same as they would on Linux set(CMAKE_LINK_LIBRARY_USING_WHOLE_ARCHIVE @@ -97,12 +77,8 @@ include_directories(BEFORE SYSTEM ${QNX_HOST}/usr/lib/gcc/${QNX_TRIPLE}/${QNX_GCC_VERSION}/include ${QNX_TARGET}/usr/include/c++/${QNX_GCC_VERSION}/${QNX_TRIPLE} ${QNX_TARGET}/usr/include + /usr/include/aarch64-unknown-nto-qnx-safety # TensorRT Safety Headers are installed here. ) -if(${QNX_VERSION} VERSION_LESS "8.0.0") - # Legacy QNX Safety packages install TensorRT headers in this location. - # DOS7 builds obtain matching headers from the selected CMake package. - include_directories(BEFORE SYSTEM /usr/include/aarch64-unknown-nto-qnx-safety) -endif() set(CMAKE_C_FLAGS_INIT ${TRT_MAGIC_QNX_FLAGS}) set(CMAKE_CXX_FLAGS_INIT ${TRT_MAGIC_QNX_FLAGS}) @@ -130,109 +106,12 @@ if(NOT DEFINED CMAKE_CUDA_COMPILER) set(CMAKE_CUDA_COMPILER "${CUDA_ROOT}/bin/nvcc" CACHE FILEPATH "Path to nvcc compiler") endif() -# SafeCUDA packages have used both aarch64-qnx-safe and aarch64-qnx target -# layouts. Select the target from the installed CUDART payload because the QNX -# SDK version does not determine the SafeCUDA package layout. Resolve target -# directories before counting them so a compatibility symlink does not make -# one installed payload appear ambiguous. -set(_TRT_SAFE_CUDA_TARGET_CANDIDATES aarch64-qnx-safe aarch64-qnx) -set(_TRT_SAFE_CUDA_TARGETS_FOUND "") -set(_TRT_SAFE_CUDA_TARGET_REAL_PATHS_FOUND "") -foreach(_TRT_SAFE_CUDA_TARGET IN LISTS _TRT_SAFE_CUDA_TARGET_CANDIDATES) - set(_TRT_SAFE_CUDA_TARGET_PATH - "${CUDA_ROOT}/targets/${_TRT_SAFE_CUDA_TARGET}") - if(EXISTS "${_TRT_SAFE_CUDA_TARGET_PATH}/lib/libcudart.so") - file(REAL_PATH "${_TRT_SAFE_CUDA_TARGET_PATH}" - _TRT_SAFE_CUDA_TARGET_REAL_PATH) - if(NOT _TRT_SAFE_CUDA_TARGET_REAL_PATH IN_LIST - _TRT_SAFE_CUDA_TARGET_REAL_PATHS_FOUND) - list(APPEND _TRT_SAFE_CUDA_TARGET_REAL_PATHS_FOUND - "${_TRT_SAFE_CUDA_TARGET_REAL_PATH}") - get_filename_component(_TRT_SAFE_CUDA_TARGET_CANONICAL - "${_TRT_SAFE_CUDA_TARGET_REAL_PATH}" NAME) - if(_TRT_SAFE_CUDA_TARGET_CANONICAL IN_LIST - _TRT_SAFE_CUDA_TARGET_CANDIDATES) - list(APPEND _TRT_SAFE_CUDA_TARGETS_FOUND - "${_TRT_SAFE_CUDA_TARGET_CANONICAL}") - else() - # Preserve a single package-provided logical target whose - # symlink backing directory uses an implementation name. - list(APPEND _TRT_SAFE_CUDA_TARGETS_FOUND - "${_TRT_SAFE_CUDA_TARGET}") - endif() - endif() - endif() -endforeach() - -list(LENGTH _TRT_SAFE_CUDA_TARGETS_FOUND _TRT_SAFE_CUDA_TARGET_COUNT) -if(_TRT_SAFE_CUDA_TARGET_COUNT EQUAL 0) - message(FATAL_ERROR - "SafeCUDA CUDART target not found. Checked " - "${CUDA_ROOT}/targets/aarch64-qnx-safe/lib/libcudart.so and " - "${CUDA_ROOT}/targets/aarch64-qnx/lib/libcudart.so") -elseif(_TRT_SAFE_CUDA_TARGET_COUNT GREATER 1) - message(FATAL_ERROR - "SafeCUDA CUDART target is ambiguous; found " - "'${_TRT_SAFE_CUDA_TARGETS_FOUND}' under ${CUDA_ROOT}/targets") -endif() -list(GET _TRT_SAFE_CUDA_TARGETS_FOUND 0 TRT_SAFE_CUDA_TARGET) -message(STATUS "TensorRT SafeCUDA target: ${TRT_SAFE_CUDA_TARGET}") - -# The aarch64-qnx SafeCUDA payload provides shared CUDART. Keep the legacy -# aarch64-qnx-safe nvcc behavior unchanged. -set(_TRT_SAFE_CUDART_FLAG "") -if(TRT_SAFE_CUDA_TARGET STREQUAL "aarch64-qnx") - set(_TRT_SAFE_CUDART_FLAG "--cudart=shared") -endif() - -# The aarch64-qnx SafeCUDA runtime depends on libraries supplied by the -# DriveOS QNX Safety PDK. Prefer an explicit CMake setting, then the matching -# environment variable, then the standard PDK_TOP layout. -set(_TRT_SAFE_PDK_FLAGS "") -set(TRT_QNX_SAFE_PDK_LIBRARIES "" CACHE INTERNAL - "DriveOS QNX Safety PDK libraries" FORCE) -if(TRT_SAFE_CUDA_TARGET STREQUAL "aarch64-qnx") - list(APPEND CMAKE_TRY_COMPILE_PLATFORM_VARIABLES TRT_PDK_LIB_TARGET) - if((NOT DEFINED TRT_PDK_LIB_TARGET OR "${TRT_PDK_LIB_TARGET}" STREQUAL "") - AND DEFINED ENV{TRT_PDK_LIB_TARGET} - AND NOT "$ENV{TRT_PDK_LIB_TARGET}" STREQUAL "") - set(TRT_PDK_LIB_TARGET "$ENV{TRT_PDK_LIB_TARGET}") - elseif((NOT DEFINED TRT_PDK_LIB_TARGET OR "${TRT_PDK_LIB_TARGET}" STREQUAL "") - AND DEFINED ENV{PDK_TOP} - AND NOT "$ENV{PDK_TOP}" STREQUAL "") - set(TRT_PDK_LIB_TARGET "$ENV{PDK_TOP}/drive-qnx-safety/lib-target") - endif() - - if(NOT DEFINED TRT_PDK_LIB_TARGET OR "${TRT_PDK_LIB_TARGET}" STREQUAL "") - message(FATAL_ERROR - "The aarch64-qnx SafeCUDA target requires the QNX Safety PDK libraries. " - "Set TRT_PDK_LIB_TARGET or PDK_TOP.") - endif() - set(TRT_PDK_LIB_TARGET "${TRT_PDK_LIB_TARGET}" CACHE PATH - "DriveOS QNX Safety PDK lib-target directory") - - set(_TRT_SAFE_PDK_LIBRARIES "") - foreach(_TRT_SAFE_PDK_LIBRARY IN ITEMS nvdvms_client nvos_s3_safety) - set(_TRT_SAFE_PDK_LIBRARY_PATH - "${TRT_PDK_LIB_TARGET}/lib${_TRT_SAFE_PDK_LIBRARY}.so") - if(NOT EXISTS "${_TRT_SAFE_PDK_LIBRARY_PATH}") - message(FATAL_ERROR "Missing QNX Safety PDK library: ${_TRT_SAFE_PDK_LIBRARY_PATH}") - endif() - list(APPEND _TRT_SAFE_PDK_LIBRARIES "${_TRT_SAFE_PDK_LIBRARY_PATH}") - endforeach() - set(TRT_QNX_SAFE_PDK_LIBRARIES "${_TRT_SAFE_PDK_LIBRARIES}" - CACHE INTERNAL "DriveOS QNX Safety PDK libraries" FORCE) - set(_TRT_SAFE_PDK_FLAGS - "-L${TRT_PDK_LIB_TARGET} -lnvdvms_client -lnvos_s3_safety") -endif() - # We need to set a couple additional flags to compile with SafeCUDA. # 1. SafeCUDA does not contain cudadevrt, so we need to disable it by setting --cudadevrt=none # 2. SafeCUDA depends on the QNX Slogger2 library, which must be linked via -lslog2. It is a system library for QNX-Safe. -# 3. Use --target-directory to keep nvcc's implicit paths aligned with the -# installed CUDART target selected above. -set(CMAKE_CUDA_FLAGS_INIT - "${CMAKE_CUDA_FLAGS_INIT} --cudadevrt=none ${_TRT_SAFE_CUDART_FLAG} -lslog2 --target-directory ${TRT_SAFE_CUDA_TARGET} -legacy-launch-seq ${_TRT_SAFE_PDK_FLAGS}") +# 3. Use --target-directory to tell nvcc to use aarch64-qnx-safe instead of aarch64-qnx +# This is critical to override nvcc's internal _TARGET_DIR_ variable +set(CMAKE_CUDA_FLAGS_INIT "${CMAKE_CUDA_FLAGS_INIT} --cudadevrt=none -lslog2 --target-directory aarch64-qnx-safe -legacy-launch-seq") # We need to explicitly populate `CMAKE_CUDA_FLAGS` with the initial flags so they propagate to the CMake Compiler Identification phase. # Otherwise, they will only be initialized afterwards, which will cause the identification to fail. @@ -240,17 +119,17 @@ set(CMAKE_CUDA_FLAGS ${CMAKE_CUDA_FLAGS_INIT}) # The CUDA setup for TRT-OSS is a bit wonky, so we enable the includes globally. include_directories(BEFORE SYSTEM - ${CUDA_ROOT}/targets/${TRT_SAFE_CUDA_TARGET}/include + ${CUDA_ROOT}/targets/aarch64-qnx-safe/include ) link_directories( - ${CUDA_ROOT}/targets/${TRT_SAFE_CUDA_TARGET}/lib - ${CUDA_ROOT}/targets/${TRT_SAFE_CUDA_TARGET}/lib/stubs + ${CUDA_ROOT}/targets/aarch64-qnx-safe/lib + ${CUDA_ROOT}/targets/aarch64-qnx-safe/lib/stubs ) add_link_options( - "LINKER:-rpath-link=${CUDA_ROOT}/targets/${TRT_SAFE_CUDA_TARGET}/lib" - "LINKER:-rpath-link=${CUDA_ROOT}/targets/${TRT_SAFE_CUDA_TARGET}/lib/stubs" + "LINKER:-rpath-link=${CUDA_ROOT}/targets/aarch64-qnx-safe/lib" + "LINKER:-rpath-link=${CUDA_ROOT}/targets/aarch64-qnx-safe/lib/stubs" ) # Disable CMake-based cudart support as we handle the linkage manually. diff --git a/docker/downloadTRT.sh b/docker/downloadTRT.sh index 7eaee31e3e..0574eb3b35 100755 --- a/docker/downloadTRT.sh +++ b/docker/downloadTRT.sh @@ -2,7 +2,7 @@ set -e -TRT_VERSION="11.3.0.99" +TRT_VERSION="11.4.0.106" usage() { echo "Usage: $0 [--x86 | --aarch64] [--cuda 13.4 | --cuda 12.9]" @@ -44,7 +44,7 @@ case "$CUDA_VERSION" in ;; esac -URL="https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.3.0/tars/TensorRT-Enterprise-11.3.0.99-Linux-${ARCH}-cuda-${CUDA_VERSION}-Release-external.tar.zst" +URL="https://developer.nvidia.com/downloads/compute/machine-learning/tensorrt/11.4.0/tars/TensorRT-Enterprise-11.4.0.106-Linux-${ARCH}-cuda-${CUDA_VERSION}-Release-external.tar.zst" echo "Downloading TensorRT package from: $URL" cd /opt diff --git a/docker/rockylinux8.Dockerfile b/docker/rockylinux8.Dockerfile index eec9d3f9b6..41072aa6a9 100644 --- a/docker/rockylinux8.Dockerfile +++ b/docker/rockylinux8.Dockerfile @@ -20,7 +20,7 @@ ARG CUDA_VERSION=13.4.1 FROM nvidia/cuda:${CUDA_VERSION}-devel-rockylinux8 LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account diff --git a/docker/rockylinux9.Dockerfile b/docker/rockylinux9.Dockerfile index 4fc9cbf0c7..b09285438a 100644 --- a/docker/rockylinux9.Dockerfile +++ b/docker/rockylinux9.Dockerfile @@ -20,7 +20,7 @@ ARG CUDA_VERSION=13.4.1 FROM nvidia/cuda:${CUDA_VERSION}-devel-rockylinux9 LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account diff --git a/docker/ubuntu-22.04.Dockerfile b/docker/ubuntu-22.04.Dockerfile index be840d7ab0..483f5b5cd4 100644 --- a/docker/ubuntu-22.04.Dockerfile +++ b/docker/ubuntu-22.04.Dockerfile @@ -20,7 +20,7 @@ ARG CUDA_VERSION=13.4.1 FROM nvidia/cuda:${CUDA_VERSION}-devel-ubuntu22.04 LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account diff --git a/docker/ubuntu-24.04-aarch64.Dockerfile b/docker/ubuntu-24.04-aarch64.Dockerfile index 4f433ddaad..3f5234b631 100644 --- a/docker/ubuntu-24.04-aarch64.Dockerfile +++ b/docker/ubuntu-24.04-aarch64.Dockerfile @@ -20,7 +20,7 @@ ARG CUDA_VERSION=13.4.1 # Multi-arch container support available in non-cudnn containers. FROM nvidia/cuda:${CUDA_VERSION}-devel-ubuntu24.04 -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account and edit default account diff --git a/docker/ubuntu-24.04.Dockerfile b/docker/ubuntu-24.04.Dockerfile index 81c3350ec2..7e1832ac44 100644 --- a/docker/ubuntu-24.04.Dockerfile +++ b/docker/ubuntu-24.04.Dockerfile @@ -21,7 +21,7 @@ ARG CUDA_VERSION=13.4.1 FROM nvidia/cuda:${CUDA_VERSION}-devel-ubuntu24.04 LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account and edit default account @@ -43,7 +43,7 @@ RUN apt-key adv --fetch-keys https://developer.download.nvidia.com/compute/cuda/ # Install requried libraries RUN apt-get update && apt-get install -y software-properties-common -RUN apt-get update && apt-get install -y --no-install-recommends \ +RUN apt-get update && apt-get install -y --no-install-recommends --allow-change-held-packages \ libcurl4-openssl-dev \ wget \ git \ diff --git a/docker/ubuntu-26.04.Dockerfile b/docker/ubuntu-26.04.Dockerfile index 8a76fed2fd..b720e5ea90 100644 --- a/docker/ubuntu-26.04.Dockerfile +++ b/docker/ubuntu-26.04.Dockerfile @@ -21,7 +21,7 @@ ARG CUDA_VERSION=13.4.1 FROM nvidia/cuda:${CUDA_VERSION}-devel-ubuntu26.04 LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 SHELL ["/bin/bash", "-c"] # Setup user account and edit default account diff --git a/docker/ubuntu-cross-aarch64.Dockerfile b/docker/ubuntu-cross-aarch64.Dockerfile index d0ecd39247..8809a7865b 100644 --- a/docker/ubuntu-cross-aarch64.Dockerfile +++ b/docker/ubuntu-cross-aarch64.Dockerfile @@ -21,7 +21,7 @@ ARG OS_VERSION=24.04 FROM nvidia/cuda:${CUDA_VERSION}-devel-ubuntu${OS_VERSION} LABEL maintainer="NVIDIA CORPORATION" -ENV TRT_VERSION=11.3.0.99 +ENV TRT_VERSION=11.4.0.106 ENV DEBIAN_FRONTEND=noninteractive # Setup user account and edit default account diff --git a/include/NvInfer.h b/include/NvInfer.h index e50377f804..6f2bb8fa88 100644 --- a/include/NvInfer.h +++ b/include/NvInfer.h @@ -6524,6 +6524,9 @@ inline INormalizationLayer::~INormalizationLayer() noexcept = default; //! \brief Layer that represents a squeeze operation, removing unit dimensions of the first input tensor //! on a set of axes specified by the second input tensor. //! +//! When the axes input is a nullptr, all dimensions of the first input whose size is statically 1 in +//! the network definition are removed. +//! //! \warning Do not inherit from this class, as doing so will break forward-compatibility of the API and ABI. //! class ISqueezeLayer : public ILayer @@ -6541,6 +6544,9 @@ class ISqueezeLayer : public ILayer //! - 0: Input data tensor. //! - 1: The axes to remove. Must resolve to a constant Int32 or Int64 1D shape tensor. //! + //! When the axes input is set to a nullptr, all dimensions of the first input whose size is statically 1 in + //! the network definition are removed. + //! using ILayer::setInput; protected: @@ -9849,6 +9855,35 @@ class INetworkDefinition : public INoCopy return mImpl->addSqueeze(input, axes); } + //! + //! \brief Add a squeeze layer to the network, with axes given by a tensor or a nullptr. + //! + //! \param input The input tensor to the layer. + //! \param axes The axes to remove unit dimensions on, or nullptr to remove every dimension that is + //! statically 1 in the network definition. + //! + //! \see ISqueezeLayer + //! + //! When axes is non-null, the behavior is identical to addSqueeze(ITensor&, ITensor&): axes must + //! be resolvable to a constant Int32 or Int64 1D shape tensor, values in axes must be unique and + //! in the range of [-r, r-1] where r is the rank of the input tensor, and for each axis value the + //! corresponding dimension in the input tensor must be one. + //! + //! When axes is nullptr, removes every dimension of the input tensor whose size is statically 1 in + //! the network definition. The set of removed dimensions, and hence the output rank, is decided + //! when the layer is added or its input is replaced; it does not depend on the optimization + //! profile. Dynamic dimensions are retained and must not be 1, because removing one would change + //! the output rank. That requirement is validated at build time when the optimization profile + //! decides it, and otherwise at runtime. To squeeze a dynamic dimension, provide an explicit axes + //! tensor. Zero-sized dimensions are retained, so empty tensors are supported. + //! + //! \return The new Squeeze layer, or nullptr if it could not be created. + //! + ISqueezeLayer* addSqueeze(ITensor& input, ITensor* axes) noexcept + { + return mImpl->addSqueezeV2(input, axes); + } + //! //! \brief Add an unsqueeze layer to the network. //! @@ -10309,6 +10344,7 @@ enum class MemoryPoolType : int32_t //! The size of this pool must be at least 4 KiB and must be a power of 2. //! This defaults to 1 MiB. //! Orin has capacity of 1 MiB per core. + //! Each loadable is given the whole pool. //! kDLA_MANAGED_SRAM = 1, @@ -10316,6 +10352,8 @@ enum class MemoryPoolType : int32_t //! kDLA_LOCAL_DRAM is host RAM used by DLA to share intermediate tensor data across operations. //! The size of this pool must be at least 4 KiB and must be a power of 2. //! This defaults to 1 GiB. + //! \note the compiled loadable may require less than this amount; at runtime, TensorRT will + //! allocate only as much as is required. //! kDLA_LOCAL_DRAM = 2, @@ -10323,6 +10361,8 @@ enum class MemoryPoolType : int32_t //! kDLA_GLOBAL_DRAM is host RAM used by DLA to store weights and metadata for execution. //! The size of this pool must be at least 4 KiB and must be a power of 2. //! This defaults to 512 MiB. + //! \note the compiled loadable may require less than this amount; at runtime, TensorRT will + //! allocate only as much as is required. //! kDLA_GLOBAL_DRAM = 3, @@ -10430,6 +10470,9 @@ enum class HardwareCompatibilityLevel : int32_t //! //! This option is only supported for engines built on NVIDIA Turing and later GPUs. //! + //! \warning CUDA green context profile streams are not supported with this option. The build will continue, but + //! the resulting engine may cause crashes or other undefined behavior during execution. + //! kSAME_COMPUTE_CAPABILITY = 2, }; @@ -10815,21 +10858,48 @@ class IBuilderConfig : public INoCopy //! //! \brief Set the CUDA stream that is used to profile this network. //! + //! This is the default engine-level profile stream. A stream set through + //! IOptimizationProfile::setProfileStream() overrides it for that optimization profile. + //! + //! If \p stream belongs to a CUDA green context, TensorRT automatically queries the context's SM count and + //! co-scheduled-SM alignment, profiles each optimization profile that inherits this stream on it, and stores the + //! resulting execution-resource contract for each such profile in the engine. The contract conservatively limits + //! CGA size to the number of SMs guaranteed to be co-scheduled. If an optimization profile's effective stream is + //! not associated with a CUDA green context, that profile is built without a CUDA green context execution + //! contract. + //! + //! Using a CUDA green context requires both a TensorRT build based on CUDA Toolkit 13.0 or newer and a CUDA driver + //! compatible with CUDA 13.0 or newer. Ordinary profile streams remain supported with older CUDA Toolkit and + //! driver versions supported by TensorRT. TensorRT uses the provided stream for tactic profiling and creates + //! auxiliary profiling streams in the same CUDA green context. + //! + //! At runtime, TensorRT compares a CUDA green context stream's resources with the active optimization profile's + //! stored execution contract. If the stream has fewer SMs or a smaller CGA size, TensorRT emits a warning and + //! continues the enqueue. Enqueuing a profile built without CUDA green context constraints on a CUDA green context + //! stream also emits a warning. In either case, the configuration may cause enqueue crashes or other undefined + //! behavior. + //! + //! TensorRT does not inspect, redirect, or constrain CUDA streams, library handles, or kernel tactics created by + //! user plugins; each plugin is responsible for keeping its own CUDA work within the intended CUDA green context + //! and its resource limits. + //! + //! The application must keep the stream and its CUDA green context alive until the build completes. + //! //! \param stream The CUDA stream used for profiling by the builder. //! - //! \see getProfileStream() + //! \see getProfileStream(), IOptimizationProfile::setProfileStream() //! - void setProfileStream(const cudaStream_t stream) noexcept + void setProfileStream(cudaStream_t const stream) noexcept { return mImpl->setProfileStream(stream); } //! - //! \brief Get the CUDA stream that is used to profile this network. + //! \brief Get the default engine-level CUDA stream used to profile this network. //! - //! \return The CUDA stream set by setProfileStream, nullptr if setProfileStream has not been called. + //! \return The CUDA stream set by setProfileStream, or nullptr if setProfileStream has not been called. //! - //! \see setProfileStream() + //! \see setProfileStream(), IOptimizationProfile::getProfileStream() //! cudaStream_t getProfileStream() const noexcept { @@ -11659,11 +11729,12 @@ class IBuilder : public INoCopy //! \param network Network definition. //! \param config Builder configuration. //! - //! \return A pointer to a IHostMemory object that contains a serialized network. + //! \return A pointer to a IHostMemory object that contains a serialized engine. //! - //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() before returning. + //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() and every + //! profile-specific stream before returning. //! - //! \see INetworkDefinition, IBuilderConfig, IHostMemory + //! \see INetworkDefinition, IBuilderConfig, IHostMemory, ICudaEngine //! nvinfer1::IHostMemory* buildSerializedNetwork(INetworkDefinition& network, IBuilderConfig& config) noexcept { @@ -11671,10 +11742,10 @@ class IBuilder : public INoCopy } //! - //! \brief Builds and serializes a network into stream for the given INetworkDefinition and IBuilderConfig. + //! \brief Builds \p network using the \p config configuration, and serializes an engine into \p writer. //! - //! This function allows building and serialization of a network without creating an engine. The engine is - //! finally serialized into the writer stream. + //! This function allows building and serialization of a network without returning an engine. The engine is + //! finally serialized into the \p writer stream. //! //! \param network Network definition. //! \param config Builder configuration. @@ -11682,9 +11753,10 @@ class IBuilder : public INoCopy //! //! \return true if build succeed, otherwise false. //! - //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() before returning. + //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() and every + //! profile-specific stream before returning. //! - //! \see INetworkDefinition, IBuilderConfig, IStreamWriter + //! \see INetworkDefinition, IBuilderConfig, IStreamWriter, ICudaEngine //! bool buildSerializedNetworkToStream( INetworkDefinition& network, IBuilderConfig& config, IStreamWriter& writer) noexcept @@ -11693,22 +11765,23 @@ class IBuilder : public INoCopy } //! - //! \brief Extended form of buildSerializedNetwork that optionally permits getting the kernelText. + //! \brief Extended form of buildSerializedNetwork that optionally permits getting the kernel text. //! //! Similar to two-argument form, except that if an engine with safe capability is successfully built - //! and there are kernels, sets kernelText to ..... Otherwise sets kernelText=nullptr. + //! and there are kernels, sets \p kernelText to the kernel source code. Otherwise leaves \p kernelText unmodified. //! - //! This function allows building and serialization of a network without creating an engine. + //! This function allows building and serialization of a network without returning an engine. //! //! \param network Network definition. //! \param config Builder configuration. //! \param kernelText A reference to a pointer to a IHostMemory object that will be set to the kernel CPP code text //! - //! \return A pointer to a IHostMemory object that contains a serialized network. + //! \return A pointer to a IHostMemory object that contains a serialized engine. //! - //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() before returning. + //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() and every + //! profile-specific stream before returning. //! - //! \see INetworkDefinition, IBuilderConfig, IHostMemory + //! \see INetworkDefinition, IBuilderConfig, IHostMemory, ICudaEngine //! nvinfer1::IHostMemory* buildSerializedNetwork( INetworkDefinition& network, IBuilderConfig& config, IHostMemory*& kernelText) noexcept @@ -11725,7 +11798,8 @@ class IBuilder : public INoCopy //! //! \return A pointer to a ICudaEngine object that contains an engine. //! - //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() before returning. + //! \note This function will synchronize the CUDA stream returned by \p config.getProfileStream() and every + //! profile-specific stream before returning. //! //! \note This function does not support \p BuilderFlag::kVERSION_COMPATIBLE. //! Please use \p buildSerializedNetwork to get a version compatible engine. diff --git a/include/NvInferImpl.h b/include/NvInferImpl.h index 989da0c747..32844714b6 100644 --- a/include/NvInferImpl.h +++ b/include/NvInferImpl.h @@ -40,6 +40,7 @@ class IProfiler; } // namespace v_1_0 using IProfiler = v_1_0::IProfiler; + namespace v_1_0 { class IOutputAllocator; @@ -303,6 +304,7 @@ class VRuntime : public VRoot TRT_NODISCARD virtual DLAWorkspaceAllocationStrategy getDLAWorkspaceAllocationStrategy() const noexcept = 0; }; + class VRefitter : public VRoot { public: @@ -327,6 +329,8 @@ class VRefitter : public VRoot virtual bool getWeightsValidation() const noexcept = 0; virtual bool refitCudaEngineAsync(cudaStream_t stream) noexcept = 0; virtual Weights getWeightsPrototype(char const* weightsName) const noexcept = 0; + // Added in TensorRT 11.4 + virtual bool releaseRefitResources() noexcept = 0; }; class VOptimizationProfile : public VRoot @@ -340,9 +344,14 @@ class VOptimizationProfile : public VRoot virtual bool isValid() const noexcept = 0; // Added in TensorRT 10.11 TRT_NODISCARD virtual bool setShapeValuesV2( - char const* inputName, OptProfileSelector select, int64_t const* values, int32_t nbValues) noexcept = 0; + char const* inputName, OptProfileSelector select, int64_t const* values, int32_t nbValues) noexcept + = 0; TRT_NODISCARD virtual int64_t const* getShapeValuesV2( - char const* inputName, OptProfileSelector select) const noexcept = 0; + char const* inputName, OptProfileSelector select) const noexcept + = 0; + // Added in TensorRT 11.4 + virtual void setProfileStream(cudaStream_t stream) noexcept = 0; + TRT_NODISCARD virtual cudaStream_t getProfileStream() const noexcept = 0; }; class VCudaEngine : public VRoot @@ -1267,6 +1276,7 @@ class VNetworkDefinition : public VRoot ReduceOperation reduceOp, int64_t root, int64_t* groups, int64_t groupSize) noexcept = 0; virtual IAttention* addAttentionV2(ITensor& query, ITensor& key, ITensor& value, AttentionNormalizationOp normOp, CausalMaskKind causalKind) noexcept = 0; + virtual ISqueezeLayer* addSqueezeV2(ITensor& input, ITensor* axes) noexcept = 0; }; class VTimingCache : public VRoot diff --git a/include/NvInferRuntime.h b/include/NvInferRuntime.h index add449fc3c..340f35eb0a 100644 --- a/include/NvInferRuntime.h +++ b/include/NvInferRuntime.h @@ -927,7 +927,7 @@ class IPluginV3OneBuild : public IPluginCapability //! \warning DataType:kBOOL and DataType::kUINT8 are not supported. //! virtual int32_t getOutputDataTypes( - DataType* outputTypes, int32_t nbOutputs, const DataType* inputTypes, int32_t nbInputs) const noexcept = 0; + DataType* outputTypes, int32_t nbOutputs, DataType const* inputTypes, int32_t nbInputs) const noexcept = 0; //! //! \brief Provide expressions for computing dimensions of the output tensors from dimensions of the input tensors. @@ -2571,6 +2571,9 @@ class IRefitter : public INoCopy //! weights repeatedly for multiple refit calls as the weights memory can be updated directly instead. The weights //! updating task should use the same stream as the one used for the refit call. //! + //! \warning CUDA green context streams are not supported. Passing one may cause crashes or other undefined + //! behavior. Use an ordinary CUDA stream instead. + //! bool refitCudaEngineAsync(cudaStream_t stream) noexcept { return mImpl->refitCudaEngineAsync(stream); @@ -2594,6 +2597,20 @@ class IRefitter : public INoCopy return mImpl->getWeightsPrototype(weightsName); } + //! + //! \brief Release resources cached for the engine associated with this refitter. + //! + //! \return True on success, false otherwise. + //! + //! A later refit recreates the resources as needed. The application must ensure that all previously enqueued + //! asynchronous refit work from every refitter sharing the same engine has completed and that no such refitter + //! executes a refit operation concurrently with this call. + //! + bool releaseRefitResources() noexcept + { + return mImpl->releaseRefitResources(); + } + protected: apiv::VRefitter* mImpl; }; @@ -2813,6 +2830,44 @@ class IOptimizationProfile : public INoCopy return mImpl->getShapeValuesV2(inputName, select); } + //! + //! \brief Set the CUDA stream that is used to profile this optimization profile. + //! + //! This stream overrides IBuilderConfig::setProfileStream() for this optimization profile. Passing nullptr clears + //! the override, and the optimization profile uses the stream from IBuilderConfig instead. + //! + //! If \p stream belongs to a CUDA green context, TensorRT automatically queries the context's SM count and + //! co-scheduled-SM alignment, profiles this optimization profile on that stream, and stores the resulting + //! execution-resource contract for this optimization profile in the engine. + //! + //! CUDA green context use requires both a TensorRT build based on CUDA Toolkit 13.0 or newer and a CUDA driver + //! compatible with CUDA 13.0 or newer. Ordinary profile streams remain supported with older CUDA Toolkit and + //! driver versions supported by TensorRT. + //! + //! The application must keep the stream and its CUDA green context alive until the build completes. + //! + //! \param stream The CUDA stream used to profile this optimization profile, or nullptr to use the stream from + //! IBuilderConfig. + //! + //! \see getProfileStream(), IBuilderConfig::setProfileStream() + //! + void setProfileStream(cudaStream_t stream) noexcept + { + mImpl->setProfileStream(stream); + } + + //! + //! \brief Get the CUDA stream set for this optimization profile. + //! + //! \return The stream set by setProfileStream(), or nullptr if no profile-specific stream is set. + //! + //! \see setProfileStream(), IBuilderConfig::getProfileStream() + //! + cudaStream_t getProfileStream() const noexcept + { + return mImpl->getProfileStream(); + } + protected: apiv::VOptimizationProfile* mImpl; virtual ~IOptimizationProfile() noexcept = 0; @@ -3053,6 +3108,7 @@ class IRuntimeConfig : public INoCopy inline IRuntimeConfig::~IRuntimeConfig() noexcept = default; + //! //! \enum EngineStat //! @@ -3884,6 +3940,7 @@ class ICudaEngine : public INoCopy return mImpl->getEngineStat(stat); } + protected: apiv::VCudaEngine* mImpl; }; @@ -4718,6 +4775,19 @@ class IExecutionContext : public INoCopy //! \warning If the Engine is streaming weights, enqueueV3 will become synchronous, and //! the graph will not be capturable. //! + //! \note When \p stream belongs to a CUDA green context, TensorRT compares its resources with the active + //! optimization profile's stored execution contract. TensorRT emits a warning if the stream has fewer SMs or + //! a smaller CGA size, or if the profile was built without CUDA green context constraints, and continues the + //! enqueue. This configuration may cause enqueue crashes or other undefined behavior. CUDA green context + //! execution requires both a TensorRT runtime built with CUDA Toolkit 13.0 or newer and a CUDA driver + //! compatible with CUDA 13.0 or newer. + //! + //! \warning When capturing this call in a CUDA graph, capture on a stream in the target CUDA green context and + //! execute the graph in that same context. The application must keep the green context alive until each + //! captured graph and every graph executable instantiated from it has been destroyed; another context with + //! identical resources is not interchangeable. A graph captured on an ordinary stream does not acquire + //! resource isolation when launched on a CUDA green context stream. + //! bool enqueueV3(cudaStream_t stream) noexcept { return mImpl->enqueueV3(stream); @@ -4810,6 +4880,23 @@ class IExecutionContext : public INoCopy //! \note The provided auxiliary streams must not be the default stream and must all be different to avoid //! deadlocks. //! + //! \note If any provided auxiliary stream used by TensorRT or the main stream passed to enqueueV3() is associated + //! with a CUDA green context, all provided auxiliary streams used by TensorRT and the main stream must belong + //! to the same CUDA green context. Otherwise, enqueueV3() reports an invalid-argument error and returns + //! false. When the main stream belongs to a CUDA green context, TensorRT creates any remaining auxiliary + //! streams in that same context. + //! + //! CUDA green context use requires both a TensorRT runtime built with CUDA Toolkit 13.0 or newer and a CUDA driver + //! compatible with CUDA 13.0 or newer. Ordinary auxiliary streams remain supported with older CUDA Toolkit and + //! driver versions supported by TensorRT. + //! + //! \warning TensorRT does not take ownership of the provided auxiliary streams. The application must keep each + //! provided auxiliary stream used by TensorRT, and any CUDA green context to which it belongs, alive while + //! the stream is configured for use by enqueueV3() and until all work enqueued on it has completed. If + //! TensorRT creates auxiliary streams in a CUDA green context, the application must keep that CUDA green + //! context alive while the execution context retains those streams and until all work enqueued on them + //! has completed. + //! //! \see enqueueV3(), IBuilderConfig::setMaxAuxStreams(), ICudaEngine::getNbAuxStreams() //! void setAuxStreams(cudaStream_t* auxStreams, int32_t nbStreams) noexcept diff --git a/include/NvInferVersion.h b/include/NvInferVersion.h index 94a852454d..0e3985c68c 100644 --- a/include/NvInferVersion.h +++ b/include/NvInferVersion.h @@ -24,9 +24,9 @@ #define NV_INFER_VERSION_H #define TRT_MAJOR_ENTERPRISE 11 -#define TRT_MINOR_ENTERPRISE 3 +#define TRT_MINOR_ENTERPRISE 4 #define TRT_PATCH_ENTERPRISE 0 -#define TRT_BUILD_ENTERPRISE 99 +#define TRT_BUILD_ENTERPRISE 106 #define NV_TENSORRT_MAJOR TRT_MAJOR_ENTERPRISE //!< TensorRT major version. #define NV_TENSORRT_MINOR TRT_MINOR_ENTERPRISE //!< TensorRT minor version. #define NV_TENSORRT_PATCH TRT_PATCH_ENTERPRISE //!< TensorRT patch version. diff --git a/include/NvOnnxParser.h b/include/NvOnnxParser.h index 37ebe83a1e..20e68ecab2 100644 --- a/include/NvOnnxParser.h +++ b/include/NvOnnxParser.h @@ -176,6 +176,8 @@ class IParserError virtual ~IParserError() {} }; +class IRefitterObserver; + //! //! \class IParser //! @@ -437,6 +439,48 @@ class IParser //! \return true if the IBuilderConfig was set successfully, false otherwise. //! virtual bool setBuilderConfig(const nvinfer1::IBuilderConfig* const builderConfig) noexcept = 0; + + //! + //! \brief Set or clear an optional observer notified once per refittable weight during parsing. + //! + //! When attached, the parser emits one RefitRecord per network weight it names via + //! INetworkDefinition::setWeightsName, at the moment the weight is created during + //! parse/parseModelProto. This is the same record schema emitted by + //! IParserRefitter::setRefitObserver during refit, produced without building or + //! deserializing an engine. Because the network has not been built yet, the records are a + //! candidate superset of the built engine's refittable weights: the builder may fold or + //! absorb some of them. Consumers replaying the records against an engine must skip + //! records whose trtName the engine's nvinfer1::IRefitter does not report in + //! getAllWeights, and rely on refitCudaEngine's missing-weights check for coverage. + //! + //! TensorRT-RTX only: When an observer is attached to the parser associated with a BuilderConfig that sets + //! kREFIT_INDIVIDUAL and kSTRIP_PLAN, the parser will provide placeholder weights for the listed initializers + //! of the following ONNX nodes: + //! * Conv kernels and biases + //! * Gemm operands and bias + //! * MatMul operands + //! * DequantizeLinear scales + //! * DequantizeLinear FP8 or FP4 data + //! * TRT_FP8DequantizeLinear/TRT_MXFP8DequantizeLinear FP8 data + //! A placeholder weight is represented by nvinfer1::Weights with the same type and count as the original + //! initializer, but with nullptr values. All placeholder weights are expected to be refit prior to inference. + //! + //! A placeholder weight will not be produced under the following circumstances: + //! * If the initializer is also a graph input or output + //! * If the initializer is referenced by a nested graph + //! + //! Placeholders will not be created when importing a model whose producer name is "TensorRT". Networks + //! containing placeholders must be built with IBuilder::buildSerializedNetwork(). + //! + //! Records are only valid if the subsequent parse call returns true. + //! + //! May be called any time before parse / parseModelProto. Pass nullptr to detach. The + //! observer must outlive the parse call, or be detached before destruction. Ownership + //! remains with the caller. + //! + //! \see IRefitterObserver IParserRefitter::setRefitObserver + //! + virtual void setRefitObserver(IRefitterObserver* observer) noexcept = 0; }; //! diff --git a/parsers/onnx b/parsers/onnx index 37fbf784e6..8f60063873 160000 --- a/parsers/onnx +++ b/parsers/onnx @@ -1 +1 @@ -Subproject commit 37fbf784e64076cab1d385682f30ddf362910fe6 +Subproject commit 8f60063873ddc9cc69125c60425fb93189851f18 diff --git a/plugin/common/plugin.h b/plugin/common/plugin.h index fbb6e088fa..e8b4f7990c 100644 --- a/plugin/common/plugin.h +++ b/plugin/common/plugin.h @@ -121,6 +121,21 @@ OutType read(BufferType const*& buffer) return val; } +//! Read a value of type OutType from the front of \p buffer and advance \p buffer past it. Unlike the +//! pointer overload, this validates that \p buffer holds enough bytes and throws PluginError if it does +//! not, so callers need no separate remaining-length check before each read. +//! \return The value read from the front of \p buffer. +template +OutType read(std::span& buffer) +{ + static_assert(std::is_trivially_copyable_v, "read<> requires a trivially copyable type."); + PLUGIN_VALIDATE(buffer.size() >= sizeof(OutType)); + OutType val{}; + std::memcpy(&val, buffer.data(), sizeof(OutType)); + buffer = buffer.subspan(sizeof(OutType)); + return val; +} + inline int32_t getTrtSmVersionDec(int32_t majorVersion, int32_t minorVersion) { return majorVersion * 10 + minorVersion; @@ -151,7 +166,9 @@ inline int32_t getSmVersion() int32_t device{-1}; PLUGIN_CHECK_CUDA(cudaGetDevice(&device)); auto const cc = DeviceComputeCapability::forDevice(device); - return getTrtSmVersionDec(cc.major, cc.minor); + auto const smVersion = getTrtSmVersionDec(cc.major, cc.minor); + // SM88 reuses SM87 plugin cubins on T238. + return smVersion == 88 ? 87 : smVersion; } // Check that all required field names are present in the PluginFieldCollection. diff --git a/plugin/decodeBbox3DPlugin/decodeBbox3D.cpp b/plugin/decodeBbox3DPlugin/decodeBbox3D.cpp index 5b16e2dac1..ea31038b63 100644 --- a/plugin/decodeBbox3DPlugin/decodeBbox3D.cpp +++ b/plugin/decodeBbox3DPlugin/decodeBbox3D.cpp @@ -53,7 +53,7 @@ DecodeBbox3DPlugin::DecodeBbox3DPlugin(float xMin, float xMax, float yMin, float mAnchorBottomHeight = anchorBottomHeight; mAnchors = anchors; mNumClasses = static_cast(mAnchorBottomHeight.size()); - PLUGIN_VALIDATE(static_cast(mNumClasses) * 2 * 4 == mAnchors.size()); + PLUGIN_VALIDATE(mNumClasses * 2 * 4 == std::ssize(mAnchors)); } DecodeBbox3DPlugin::DecodeBbox3DPlugin(float xMin, float xMax, float yMin, float yMax, float zMin, float zMax, @@ -75,7 +75,7 @@ DecodeBbox3DPlugin::DecodeBbox3DPlugin(float xMin, float xMax, float yMin, float mAnchorBottomHeight = anchorBottomHeight; mAnchors = anchors; mNumClasses = static_cast(mAnchorBottomHeight.size()); - PLUGIN_VALIDATE(static_cast(mNumClasses) * 2 * 4 == mAnchors.size()); + PLUGIN_VALIDATE(mNumClasses * 2 * 4 == std::ssize(mAnchors)); } DecodeBbox3DPlugin::DecodeBbox3DPlugin(void const* data, size_t length) diff --git a/plugin/flattenConcat/flattenConcat.cpp b/plugin/flattenConcat/flattenConcat.cpp index 3286770bfa..e6a3ab6561 100644 --- a/plugin/flattenConcat/flattenConcat.cpp +++ b/plugin/flattenConcat/flattenConcat.cpp @@ -78,14 +78,14 @@ FlattenConcat::FlattenConcat(void const* data, size_t length) ensureAvailable(static_cast(mNumInputs) * sizeof(int32_t)); mInputConcatAxis.resize(mNumInputs); - std::for_each(mInputConcatAxis.begin(), mInputConcatAxis.end(), [&](int32_t& inp) { inp = read(d); }); + std::ranges::for_each(mInputConcatAxis, [&](int32_t& inp) { inp = read(d); }); ensureAvailable(sizeof(nvinfer1::Dims3)); mCHW = read(d); ensureAvailable(static_cast(mNumInputs) * sizeof(size_t)); mCopySize.resize(mNumInputs); - std::for_each(mCopySize.begin(), mCopySize.end(), [&](size_t& inp) { inp = read(d); }); + std::ranges::for_each(mCopySize, [&](size_t& inp) { inp = read(d); }); PLUGIN_VALIDATE(d == a + length); } diff --git a/plugin/modulatedDeformConvPlugin/modulatedDeformConvPlugin.cpp b/plugin/modulatedDeformConvPlugin/modulatedDeformConvPlugin.cpp index 21ccf43c3e..884f6ef3bd 100644 --- a/plugin/modulatedDeformConvPlugin/modulatedDeformConvPlugin.cpp +++ b/plugin/modulatedDeformConvPlugin/modulatedDeformConvPlugin.cpp @@ -27,6 +27,7 @@ #include "modulatedDeformConvPlugin.h" #include #include +#include using namespace nvinfer1; using namespace nvinfer1::plugin; @@ -403,12 +404,12 @@ nvinfer1::PluginFieldCollection const* ModulatedDeformableConvPluginDynamicCreat return &mFC; } -// NOLINTNEXTLINE(readability-function-cognitive-complexity) nvinfer1::IPluginV3* ModulatedDeformableConvPluginDynamicCreator::createPlugin( char const* name, nvinfer1::PluginFieldCollection const* fc, nvinfer1::TensorRTPhase phase) noexcept { try { + using namespace std::string_view_literals; PLUGIN_VALIDATE(fc != nullptr); PLUGIN_VALIDATE(fc->fields != nullptr || fc->nbFields == 0); @@ -433,14 +434,14 @@ nvinfer1::IPluginV3* ModulatedDeformableConvPluginDynamicCreator::createPlugin( std::string const fieldName(field.name); - if (fieldName == "deformable_group") + if (fieldName == "deformable_group"sv) { PLUGIN_VALIDATE(field.type == PluginFieldType::kINT32); PLUGIN_VALIDATE(field.length == 1); deformableGroup = *static_cast(field.data); PLUGIN_VALIDATE(deformableGroup > 0); } - else if (fieldName == "group") + else if (fieldName == "group"sv) { PLUGIN_VALIDATE(field.type == PluginFieldType::kINT32); PLUGIN_VALIDATE(field.length == 1); @@ -449,8 +450,19 @@ nvinfer1::IPluginV3* ModulatedDeformableConvPluginDynamicCreator::createPlugin( } else if (bert::elem(fieldName, {"stride", "padding", "dilation"})) { - nvinfer1::Dims* dimsPtr - = (fieldName == "stride") ? &stride : ((fieldName == "padding") ? &padding : &dilation); + nvinfer1::Dims* dimsPtr; + if (fieldName == "stride"sv) + { + dimsPtr = &stride; + } + else if (fieldName == "padding"sv) + { + dimsPtr = &padding; + } + else + { + dimsPtr = &dilation; + } PluginFieldType const expectedFieldType = isBuildPhase ? PluginFieldType::kINT32 : PluginFieldType::kINT64; @@ -463,21 +475,19 @@ nvinfer1::IPluginV3* ModulatedDeformableConvPluginDynamicCreator::createPlugin( if (isBuildPhase) { // During build time, data is INT32, upcast to int64 for internal storage (Dims uses int64_t). - auto const* dataPtr = static_cast(field.data); - dimsPtr->d[0] = dataPtr[0]; - dimsPtr->d[1] = dataPtr[1]; + auto const dataPtr = std::span(static_cast(field.data), 2); + std::ranges::copy(dataPtr, dimsPtr->d); } else // Runtime phase { // During runtime, data is deserialized as INT64. PLUGIN_VALIDATE(phase == nvinfer1::TensorRTPhase::kRUNTIME); - auto const* dataPtr = static_cast(field.data); - dimsPtr->d[0] = dataPtr[0]; - dimsPtr->d[1] = dataPtr[1]; + auto const dataPtr = std::span(static_cast(field.data), 2); + std::ranges::copy(dataPtr, dimsPtr->d); } // Validate values - if (fieldName == "padding") + if (fieldName == "padding"sv) { PLUGIN_VALIDATE(dimsPtr->d[0] >= 0 && dimsPtr->d[1] >= 0); } diff --git a/plugin/multilevelProposeROI/multilevelProposeROIPlugin.cpp b/plugin/multilevelProposeROI/multilevelProposeROIPlugin.cpp index 57f53beced..cc8e189b16 100644 --- a/plugin/multilevelProposeROI/multilevelProposeROIPlugin.cpp +++ b/plugin/multilevelProposeROI/multilevelProposeROIPlugin.cpp @@ -350,7 +350,7 @@ void MultilevelProposeROI::check_valid_inputs(nvinfer1::Dims const* inputs, int3 size_t MultilevelProposeROI::getWorkspaceSize(int32_t batch_size) const noexcept { size_t total_size = 0; - PLUGIN_ASSERT(mAnchorsCnt.size() == static_cast(mFeatureCnt)); + PLUGIN_ASSERT(std::ssize(mAnchorsCnt) == mFeatureCnt); // workspace for propose on each feature map for (int32_t i = 0; i < mFeatureCnt; i++) @@ -419,7 +419,7 @@ void MultilevelProposeROI::generate_pyramid_anchors(nvinfer1::Dims const& imageS anchors.push_back(s_anchors); } - PLUGIN_VALIDATE(anchors.size() == static_cast(max_level - min_level + 1)); + PLUGIN_VALIDATE(std::ssize(anchors) == max_level - min_level + 1); } int32_t MultilevelProposeROI::enqueue( diff --git a/plugin/priorBoxPlugin/priorBoxPlugin.cpp b/plugin/priorBoxPlugin/priorBoxPlugin.cpp index ea78edf628..649e143e7c 100644 --- a/plugin/priorBoxPlugin/priorBoxPlugin.cpp +++ b/plugin/priorBoxPlugin/priorBoxPlugin.cpp @@ -224,7 +224,7 @@ void PriorBox::serialize(void* buffer) const noexcept auto writeArray = [&d](int32_t const size, float const* srcPtr, std::vector const& srcVec) { // srcVec is only used here to check that the size and srcPtr are correct. PLUGIN_VALIDATE(srcVec.data() == srcPtr); - PLUGIN_VALIDATE(srcVec.size() == static_cast(size)); + PLUGIN_VALIDATE(std::ssize(srcVec) == size); for (int32_t i = 0; i < size; i++) { write(d, srcPtr[i]); diff --git a/plugin/scatterElementsPlugin/scatterElementsPluginKernel.cu b/plugin/scatterElementsPlugin/scatterElementsPluginKernel.cu index 0f2c2f50dd..41c25b5b04 100644 --- a/plugin/scatterElementsPlugin/scatterElementsPluginKernel.cu +++ b/plugin/scatterElementsPlugin/scatterElementsPluginKernel.cu @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -23,25 +23,48 @@ #include "TensorInfo.cuh" #include "common/dimsHelpers.h" +#include "common/kernels/reducedMathPlugin.h" #include "reducer.cuh" #include "scatterElementsPluginKernel.h" #include +#include namespace nvinfer1 { namespace plugin { -#define THREADS 256 -#define BLOCKS(N) (N + THREADS - 1) / THREADS +namespace +{ +constexpr int64_t kTHREADS = 256; + +//! Output dimensions and divisors for decoding update coordinates. +struct ScatterElementsOutputLayout +{ + Dims outputDims; + ReducedDivisor indexDivisors[Dims::MAX_DIMS]; +}; + +//! Keep divisor parameters out of the flattened kernel specialization. +template +using ScatterElementsLayout = std::conditional_t; + +[[nodiscard]] constexpr int64_t calcBlocks(int64_t n) +{ + return (n + kTHREADS - 1) / kTHREADS; +} +} // namespace using detail::TensorInfo; using detail::getTensorInfo; using nvinfer1::pluginInternal::volume; -template -__global__ void scatterElements_kernel(const TScalar* updatesData, const TensorInfo indexInfo, - TScalar* outData, int32_t nE, int32_t nK, int32_t nN, int32_t nbElements) +//! Reduce updates at scatter indices using the output tensor's layout. +//! Specialize address generation to keep the coordinate loop's stack usage out of the flattened path. +template +__global__ void scatterElements_kernel(TScalar const* updatesData, TensorInfo const indexInfo, + ScatterElementsLayout const outputLayout, TScalar* outData, int32_t nE, int32_t nK, int32_t nN, + int32_t axis, int32_t nbElements) { int32_t thread_idx = blockIdx.x * blockDim.x + threadIdx.x; @@ -54,7 +77,38 @@ __global__ void scatterElements_kernel(const TScalar* updatesData, const TensorI int32_t offset = detail::IndexToOffset::get(thread_idx, indexInfo); int64_t idx = indexInfo.data[offset]; - Reducer::atomic_write(outData + b * nN * nK + idx * nK + k, updatesData[thread_idx]); + if (idx < -nN || idx >= nN) + { + return; + } + if (idx < 0) + { + idx += nN; + } + + int64_t outputOffset{}; + if constexpr (!UseOutputLayout) + { + outputOffset = static_cast(b) * nN * nK + idx * nK + k; + } + else + { + // Non-axis dimensions can be smaller for updates, so output strides must be computed separately. + int32_t linearIndex = thread_idx; + int64_t outputStride{1}; + for (int32_t dim = indexInfo.dims - 1; dim >= 0; --dim) + { + int32_t coordinate{}; + int32_t quotient{}; + outputLayout.indexDivisors[dim].divmod(linearIndex, quotient, coordinate); + linearIndex = quotient; + int64_t const outputCoordinate = dim == axis ? idx : coordinate; + outputOffset += outputCoordinate * outputStride; + outputStride *= outputLayout.outputDims.d[dim]; + } + } + + Reducer::atomic_write(outData + outputOffset, updatesData[thread_idx]); } } @@ -105,13 +159,36 @@ void dispatchScatterElementsKernel(void* outDataPtr, void const* dataDataPtr, vo auto nN = outDesc.dims.d[axis]; auto indexInfo = getTensorInfo(indicesDataPtr, indicesDesc); + bool sameNonAxisShape{true}; + for (int32_t dim = 0; dim < outDesc.dims.nbDims; ++dim) + { + if (dim != axis && updatesDesc.dims.d[dim] != outDesc.dims.d[dim]) + { + sameNonAxisShape = false; + break; + } + } auto updatesData = (TScalar*) updatesDataPtr; auto outData = (TScalar*) outDataPtr; AT_DISPATCH_REDUCTION_TYPES(reduction, [&] { - scatterElements_kernel - <<>>(updatesData, indexInfo, outData, nE, nK, nN, updatesNumEl); + if (sameNonAxisShape) + { + scatterElements_kernel<<>>( + updatesData, indexInfo, outDesc.dims, outData, nE, nK, nN, axis, updatesNumEl); + } + else + { + ScatterElementsOutputLayout outputLayout{outDesc.dims, {}}; + for (int32_t dim = 0; dim < indexInfo.dims; ++dim) + { + outputLayout.indexDivisors[dim] = ReducedDivisor(indexInfo.sizes[dim]); + } + scatterElements_kernel<<>>( + updatesData, indexInfo, outputLayout, outData, nE, nK, nN, axis, updatesNumEl); + } + PLUGIN_CUASSERT(cudaGetLastError()); }); } @@ -125,6 +202,12 @@ void runScatterElementsKernel(void* outDataPtr, void const* dataDataPtr, void co cudaStream_t stream) { + for (int32_t dim = 0; dim < outDesc.dims.nbDims; ++dim) + { + PLUGIN_VALIDATE(dim == axis || updatesDesc.dims.d[dim] <= outDesc.dims.d[dim], + "ScatterElements updates dimensions must not exceed output dimensions outside the axis"); + } + auto updatesNumEl = volume(updatesDesc.dims); auto outNumEl = volume(outDesc.dims); diff --git a/plugin/scatterPlugin/scatterPlugin.cpp b/plugin/scatterPlugin/scatterPlugin.cpp index 02788802aa..c462889f0a 100644 --- a/plugin/scatterPlugin/scatterPlugin.cpp +++ b/plugin/scatterPlugin/scatterPlugin.cpp @@ -111,9 +111,9 @@ void ScatterND::calculateTransformCoeff( nIndx *= dataTensorDims.d[i]; } - std::reverse(pitches.begin(), pitches.end()); // last dimension pitch is always one (assuming linear mem) + std::ranges::reverse(pitches); // last dimension pitch is always one (assuming linear mem) - std::copy(pitches.begin(), pitches.end(), transformCoeff); + std::ranges::copy(pitches, transformCoeff); } int32_t ScatterND::calculateCopySize(Dims const& dataDims) const noexcept diff --git a/python/docstrings/infer/pyCoreDoc.h b/python/docstrings/infer/pyCoreDoc.h index 4b7a30d8f4..d97d3fd762 100644 --- a/python/docstrings/infer/pyCoreDoc.h +++ b/python/docstrings/infer/pyCoreDoc.h @@ -181,6 +181,7 @@ constexpr char const* descr = R"trtdoc( :class:`IOptimizationProfile` implements :func:`__nonzero__` and :func:`__bool__` such that evaluating a profile as a :class:`bool` (e.g. ``if profile:``) will check whether the optimization profile can be passed to an IBuilderConfig object. This will perform partial validation, by e.g. checking that the maximum dimensions are at least as large as the optimum dimensions, and that the optimum dimensions are always as least as large as the minimum dimensions. Some validation steps require knowledge of the network definition and are deferred to engine build time. :ivar extra_memory_target: Additional memory that the builder should aim to maximally allocate for this profile, as a fraction of the memory it would use if the user did not impose any constraints on memory. This unconstrained case is the default; it corresponds to ``extra_memory_target`` == 1.0. If ``extra_memory_target`` == 0.0, the builder aims to create the new optimization profile without allocating any additional weight memory. Valid inputs lie between 0.0 and 1.0. This parameter is only a hint, and TensorRT does not guarantee that the ``extra_memory_target`` will be reached. This parameter is ignored for the first (default) optimization profile that is defined. + :ivar profile_stream: :class:`int` The handle for the CUDA stream used to profile this optimization profile. A nonzero value overrides :attr:`IBuilderConfig.profile_stream` for this profile. Setting it to 0 clears the override, so the profile uses the builder configuration's stream. If the stream belongs to a CUDA green context, TensorRT queries its SM/CGA limits, constrains and profiles this optimization profile for that partition, and stores a profile-specific execution-resource contract in the engine. CUDA green context use requires both a TensorRT build based on CUDA Toolkit 13.0 or newer and a CUDA driver compatible with CUDA 13.0 or newer. Keep the stream and its CUDA green context alive until engine building completes. )trtdoc"; constexpr char const* set_shape = R"trtdoc( @@ -557,6 +558,10 @@ constexpr char const* execute_async_v3 = R"trtdoc( Input tensors can be released after the :func:`set_input_consumed_event` whereas output tensors require stream synchronization. :arg stream_handle: The cuda stream on which the inference kernels will be enqueued. Using default stream may lead to performance issues due to additional cudaDeviceSynchronize() calls by TensorRT to ensure correct synchronizations. Please use non-default stream instead. + + When the stream belongs to a CUDA green context, TensorRT compares its resources with the active optimization profile's stored execution contract. TensorRT emits a warning if the stream has fewer SMs or a smaller CGA size, or if the profile was built without CUDA green context constraints, and continues execution. This configuration may cause enqueue crashes or other undefined behavior. CUDA green context execution requires both a TensorRT runtime built with CUDA Toolkit 13.0 or newer and a CUDA driver compatible with CUDA 13.0 or newer. + + When capturing this call in a CUDA graph, capture on a stream in the target CUDA green context and execute the graph in that same context. Keep the green context alive until each captured graph and every graph executable instantiated from it has been destroyed; another context with identical resources is not interchangeable. A graph captured on an ordinary stream does not acquire resource isolation when launched on a CUDA green context stream. )trtdoc"; constexpr char const* set_aux_streams = R"trtdoc( @@ -570,6 +575,13 @@ constexpr char const* set_aux_streams = R"trtdoc( The provided auxiliary streams must not be the default stream and must all be different to avoid deadlocks. + If any provided auxiliary stream used by TensorRT or the main stream passed to :func:`execute_async_v3` is associated with a CUDA green context, all provided auxiliary streams used by TensorRT and the main stream must belong to the same CUDA green context. Otherwise, :func:`execute_async_v3` reports an invalid-argument error and returns ``False``. When the main stream belongs to a CUDA green context, TensorRT creates any remaining auxiliary streams in that same context. + + CUDA green context use requires both a TensorRT runtime built with CUDA Toolkit 13.0 or newer and a CUDA driver compatible with CUDA 13.0 or newer. Ordinary auxiliary streams remain supported with older CUDA Toolkit and driver versions supported by TensorRT. + + .. warning:: + TensorRT does not take ownership of the provided auxiliary streams. Keep each provided auxiliary stream used by TensorRT, and any CUDA green context to which it belongs, alive while the stream is configured for use by :func:`execute_async_v3` and until all work enqueued on it has completed. If TensorRT creates auxiliary streams in a CUDA green context, keep that CUDA green context alive while the execution context retains those streams and until all work enqueued on them has completed. + :arg aux_streams: A list of cuda streams. If the length of the list is greater than engine.num_aux_streams, then only the first "engine.num_aux_streams" streams will be used. If the length is less than engine.num_aux_streams, such as an empty list, then TensorRT will use the provided streams for the first few auxiliary streams, and will create additional streams internally for the rest of the auxiliary streams. )trtdoc"; @@ -623,6 +635,8 @@ constexpr char const* set_communicator = R"trtdoc( The communicator must be uniform across all multi-device instances or undefined behavior occurs. + This is a collective call: it blocks until every rank sharing the communicator has called it. + :returns: True if the communicator was set successfully, False otherwise. )trtdoc"; @@ -743,6 +757,7 @@ constexpr char const* STRIPPED_WEIGHTS_SIZE = R"trtdoc(The stripped weight size in bytes for engines built with BuilderFlag::kSTRIP_PLAN.)trtdoc"; } // namespace EngineStatDoc + namespace ICudaEngineDoc { constexpr char const* descr = R"trtdoc( @@ -924,6 +939,7 @@ constexpr char const* create_serialization_config = R"trtdoc( Create a serialization configuration object. )trtdoc"; + constexpr char const* serialize_with_config = R"trtdoc( Serialize the network to a stream. )trtdoc"; @@ -1209,16 +1225,19 @@ constexpr char const* DLA_MANAGED_SRAM = R"trtdoc( The size of this pool must be at least 4 KiB and must be a power of 2. This defaults to 1 MiB. Orin has capacity of 1 MiB per core. + Each loadable is given the whole pool. )trtdoc"; constexpr char const* DLA_LOCAL_DRAM = R"trtdoc( DLA_LOCAL_DRAM is host RAM used by DLA to share intermediate tensor data across operations. The size of this pool must be at least 4 KiB and must be a power of 2. This defaults to 1 GiB. + Note: the compiled loadable may require less than this amount; at runtime, TensorRT will allocate only as much as is required. )trtdoc"; constexpr char const* DLA_GLOBAL_DRAM = R"trtdoc( DLA_GLOBAL_DRAM is host RAM used by DLA to store weights and metadata for execution. The size of this pool must be at least 4 KiB and must be a power of 2. This defaults to 512 MiB. + Note: the compiled loadable may require less than this amount; at runtime, TensorRT will allocate only as much as is required. )trtdoc"; constexpr char const* TACTIC_DRAM = R"trtdoc( TACTIC_DRAM is the host DRAM used by the optimizer to @@ -1538,7 +1557,7 @@ constexpr char const* descr = R"trtdoc( :ivar avg_timing_iterations: :class:`int` The number of averaging iterations used when timing layers. When timing layers, the builder minimizes over a set of average times for layer execution. This parameter controls the number of iterations used in averaging. By default the number of averaging iterations is 1. :ivar flags: :class:`int` The build mode flags to turn on builder options for this network. The flags are listed in the BuilderFlags enum. The flags set configuration options to build the network. This should be in integer consisting of one or more :class:`BuilderFlag` s, combined via binary OR. For example, ``1 << BuilderFlag.FP16 | 1 << BuilderFlag.DEBUG``. - :ivar profile_stream: :class:`int` The handle for the CUDA stream that is used to profile this network. + :ivar profile_stream: :class:`int` The handle for the default engine-level CUDA stream used to profile this network. An optimization profile's ``profile_stream`` overrides this stream. If an effective profile stream belongs to a CUDA green context, that profile is built with the queried SM/CGA contract. At runtime, TensorRT warns if a CUDA green context stream does not satisfy the active profile's contract, or if an unconstrained profile is enqueued on a CUDA green context stream, and continues execution. This configuration may cause enqueue crashes or other undefined behavior. CUDA green context use requires both a TensorRT build based on CUDA Toolkit 13.0 or newer and a CUDA driver compatible with CUDA 13.0 or newer. Ordinary profile streams remain supported with older CUDA Toolkit and driver versions supported by TensorRT. Keep the stream and its CUDA green context alive until engine building completes. :ivar num_optimization_profiles: :class:`int` The number of optimization profiles. :ivar default_device_type: :class:`tensorrt.DeviceType` The default DeviceType to be used by the Builder. :ivar DLA_core: :class:`int` The DLA core that the engine executes on. Must be between 0 and N-1 where N is the number of available DLA cores. @@ -2185,6 +2204,16 @@ constexpr char const* refit_cuda_engine_async = R"trtdoc( :returns: ``True`` on success, or ``False`` if new weights validation fails or get_missing_weights() != 0 before the call. )trtdoc"; +constexpr char const* release_refit_resources = R"trtdoc( + Release resources cached for the engine associated with this refitter. + + A later refit recreates the resources as needed. The application must ensure that all previously enqueued + asynchronous refit work from every refitter sharing the same engine has completed and that no such refitter + executes a refit operation concurrently with this call. + + :returns: ``True`` on success, ``False`` otherwise. +)trtdoc"; + constexpr char const* get_missing = R"trtdoc( Get description of missing weights. diff --git a/python/docstrings/infer/pyGraphDoc.h b/python/docstrings/infer/pyGraphDoc.h index c815b9c9c0..6fdcb824d4 100644 --- a/python/docstrings/infer/pyGraphDoc.h +++ b/python/docstrings/infer/pyGraphDoc.h @@ -3029,11 +3029,17 @@ constexpr char const* add_normalization_v2 = R"trtdoc( )trtdoc"; constexpr char const* add_squeeze = R"trtdoc( - Adds a Squeeze layer to the network. + Adds a Squeeze layer with an optional axes input to the network. See :class:`ISqueezeLayer` for more information. + With axes, removes the specified unit dimensions. Without axes, removes every dimension of the + input whose size is statically 1 in the network definition. Each dynamic dimension is + retained and must not be 1 at runtime, because removing it would change the output rank, + which is fixed at definition time. Zero-sized dimensions are retained, so empty tensors are supported. + :arg input: The input tensor to the layer. - :arg axes: The tensor containing axes to remove. Must be resolvable to a constant Int32 or Int64 1D shape tensor. + :arg axes: The tensor containing axes to remove, or :class:`None` to remove all unit + dimensions. Must be resolvable to a constant Int32 or Int64 1D shape tensor if provided. :returns: the new Squeeze layer, or :class:`None` if it could not be created. )trtdoc"; diff --git a/python/packaging/CMakeLists.txt b/python/packaging/CMakeLists.txt index bd3e4a351c..e9eed85d9a 100644 --- a/python/packaging/CMakeLists.txt +++ b/python/packaging/CMakeLists.txt @@ -140,7 +140,13 @@ endif() # \returns wheelPlatform The computed wheel platform name. Standalone wheels use manylinux tags. function(get_wheel_platform isStandalone outVar) if(WIN32) - set(wheelPlatform win_${TRT_LOWERCASE_CMAKE_SYSTEM_PROCESSOR}) + if(TRT_LOWERCASE_CMAKE_SYSTEM_PROCESSOR MATCHES "^(x86_64|amd64)$") + set(wheelPlatform win_amd64) + elseif(TRT_LOWERCASE_CMAKE_SYSTEM_PROCESSOR MATCHES "^(arm64|aarch64)$") + set(wheelPlatform win_arm64) + else() + set(wheelPlatform win_${TRT_LOWERCASE_CMAKE_SYSTEM_PROCESSOR}) + endif() elseif(isStandalone) if(NOT GLIBC_VERSION) # Determine glibc version for standalone wheels diff --git a/python/src/infer/pyCore.cpp b/python/src/infer/pyCore.cpp index 2a86605a76..f7fcf80eae 100644 --- a/python/src/infer/pyCore.cpp +++ b/python/src/infer/pyCore.cpp @@ -19,6 +19,7 @@ #include "ForwardDeclarations.h" #include "utils.h" #include +#include #include #include @@ -86,6 +87,23 @@ static auto const opt_profile_get_shape_input return shapes; }; +namespace +{ + +//! Return an optimization profile's CUDA stream as a Python-compatible handle. +size_t optProfileGetProfileStream(IOptimizationProfile const& self) +{ + return reinterpret_cast(self.getProfileStream()); +} + +//! Set an optimization profile's CUDA stream from a Python stream handle. +void optProfileSetProfileStream(IOptimizationProfile& self, size_t streamHandle) +{ + self.setProfileStream(reinterpret_cast(streamHandle)); +} + +} // namespace + // For IExecutionContext static auto const execute_v2 = [](IExecutionContext& self, std::vector& bindings) { @@ -181,6 +199,7 @@ static auto const reader_v2_read = [](IStreamReaderV2& self, void* destination, }; + // For ICudaEngine // TODO: Add slicing support? static auto const engine_getitem = [](ICudaEngine& self, int32_t pyIndex) { @@ -1058,6 +1077,7 @@ void bindCore(py::module& m) IOptimizationProfileDoc::get_shape_input) .def_property("extra_memory_target", &IOptimizationProfile::getExtraMemoryTarget, &IOptimizationProfile::setExtraMemoryTarget) + .def_property("profile_stream", lambdas::optProfileGetProfileStream, lambdas::optProfileSetProfileStream) .def("__nonzero__", &IOptimizationProfile::isValid) .def("__bool__", &IOptimizationProfile::isValid); @@ -1375,8 +1395,10 @@ void bindCore(py::module& m) IExecutionContextDoc::set_all_tensors_debug_state) .def_property("unfused_tensors_debug_state", &IExecutionContext::getUnfusedTensorsDebugState, &IExecutionContext::setUnfusedTensorsDebugState) + // setCommunicator blocks in ncclCommSplit, a collective every rank must reach. Holding the GIL + // across it deadlocks any single-process multi-rank program. .def("set_communicator", &IExecutionContext::setCommunicator, "communicator"_a, - IExecutionContextDoc::set_communicator) + IExecutionContextDoc::set_communicator, py::call_guard{}) .def("get_runtime_config", &IExecutionContext::getRuntimeConfig, IExecutionContextDoc::get_runtime_config, py::keep_alive<1, 0>{}, py::call_guard{}) ; @@ -1421,6 +1443,7 @@ void bindCore(py::module& m) .value("DEVICE", TensorLocation::kDEVICE, TensorLocationDoc::DEVICE) .value("HOST", TensorLocation::kHOST, TensorLocationDoc::HOST); // TensorLocation + py::enum_(m, "TensorIOMode", TensorIOModeDoc::descr, py::module_local()) .value("NONE", TensorIOMode::kNONE, TensorIOModeDoc::NONE) .value("INPUT", TensorIOMode::kINPUT, TensorIOModeDoc::INPUT) @@ -1916,6 +1939,8 @@ void bindCore(py::module& m) .def_property("weights_validation", &IRefitter::getWeightsValidation, &IRefitter::setWeightsValidation) .def("refit_cuda_engine_async", lambdas::refitter_refit_cuda_engine_async, "stream_handle"_a, RefitterDoc::refit_cuda_engine_async, py::call_guard{}) + .def("release_refit_resources", &IRefitter::releaseRefitResources, RefitterDoc::release_refit_resources, + py::call_guard{}) .def("get_weights_prototype", &IRefitter::getWeightsPrototype, "weights_name"_a, RefitterDoc::get_weights_prototype) .def("__del__", &utils::doNothingDel); diff --git a/python/src/infer/pyGraph.cpp b/python/src/infer/pyGraph.cpp index 689faa26b6..e39b2679d7 100644 --- a/python/src/infer/pyGraph.cpp +++ b/python/src/infer/pyGraph.cpp @@ -1226,7 +1226,9 @@ namespace tensorrt .def("is_debug_tensor", &INetworkDefinition::isDebugTensor, "tensor"_a, INetworkDefinitionDoc::is_debug_tensor) .def("mark_unfused_tensors_as_debug_tensors", &INetworkDefinition::markUnfusedTensorsAsDebugTensors, INetworkDefinitionDoc::mark_unfused_tensors_as_debug_tensors) .def("unmark_unfused_tensors_as_debug_tensors", &INetworkDefinition::unmarkUnfusedTensorsAsDebugTensors, INetworkDefinitionDoc::unmark_unfused_tensors_as_debug_tensors) - .def("add_squeeze", &INetworkDefinition::addSqueeze, "input"_a, "axes"_a, INetworkDefinitionDoc::add_squeeze, py::return_value_policy::reference_internal) + .def("add_squeeze", + py::overload_cast(&INetworkDefinition::addSqueeze), + "input"_a, "axes"_a = nullptr, INetworkDefinitionDoc::add_squeeze, py::return_value_policy::reference_internal) .def("add_unsqueeze", &INetworkDefinition::addUnsqueeze, "input"_a, "axes"_a, INetworkDefinitionDoc::add_unsqueeze, py::return_value_policy::reference_internal) .def("add_normalization_v2", &INetworkDefinition::addNormalizationV2, "input"_a, "scale"_a, "bias"_a, "axesMask"_a, INetworkDefinitionDoc::add_normalization_v2, py::return_value_policy::reference_internal) diff --git a/samples/CMakeLists.txt b/samples/CMakeLists.txt index b4c6607ffc..3be3cdbe89 100644 --- a/samples/CMakeLists.txt +++ b/samples/CMakeLists.txt @@ -15,7 +15,7 @@ # limitations under the License. # -# Target which holds all enabled samples, including trtexec. +# Target which holds all enabled samples. # If TRT_BUILD_SAMPLES is disabled, this target may not build many things. add_custom_target(tensorrt_samples) @@ -27,18 +27,18 @@ macro(add_sample) endmacro() if(TRT_SAFETY_INFERENCE_ONLY) - add_sample(trtSafeExec) add_sample(sampleSafeMNIST) add_sample(sampleSafePluginV3) + add_sample(trtSafeExec) else() - # Require the ONNX parser to build samples or trtexec. - if((${TRT_BUILD_SAMPLES} OR ${TRT_BUILD_TRTEXEC}) AND NOT TARGET nvonnxparser) - message(FATAL_ERROR "Building trtexec (TRT_BUILD_TRTEXEC=${TRT_BUILD_TRTEXEC}) and/or the other samples (TRT_BUILD_SAMPLES=${TRT_BUILD_SAMPLES}) requires the nvonnxparser target") + # Require the ONNX parser to build samples. + if(${TRT_BUILD_SAMPLES} AND NOT TARGET nvonnxparser) + message(FATAL_ERROR "Building the samples (TRT_BUILD_SAMPLES=${TRT_BUILD_SAMPLES}) requires the nvonnxparser target") endif() # Setup aliases for ease of swapping between static/dynamic TRT. if(${TRT_BUILD_SAMPLES_LINK_STATIC_TRT}) - add_library(TRT_SAMPLES::tensorrt INTERFACE IMPORTED GLOBAL) + add_library(TRT_SAMPLES::tensorrt INTERFACE IMPORTED) target_link_libraries(TRT_SAMPLES::tensorrt INTERFACE tensorrt_static) add_library(TRT_SAMPLES::onnxparser INTERFACE IMPORTED) @@ -50,7 +50,7 @@ else() set_property(GLOBAL PROPERTY JOB_POOLS three_jobs=3) set(CMAKE_JOB_POOL_LINK three_jobs) else() - add_library(TRT_SAMPLES::tensorrt INTERFACE IMPORTED GLOBAL) + add_library(TRT_SAMPLES::tensorrt INTERFACE IMPORTED) target_link_libraries(TRT_SAMPLES::tensorrt INTERFACE tensorrt) add_library(TRT_SAMPLES::onnxparser INTERFACE IMPORTED) @@ -59,16 +59,9 @@ else() # OSS samples need the ONNX parser path to be included when each sample is built add_subdirectory(common) - - if(${TRT_BUILD_TRTEXEC}) - add_sample(trtexec) - endif() + add_subdirectory(trtexecCommon) if(${TRT_BUILD_SAMPLES}) - if (NOT ${TRT_BUILD_TRTEXEC}) - message(WARNING "TRT_BUILD_SAMPLES is enabled but TRT_BUILD_TRTEXEC is not enabled. This may be unintended.") - endif() - # Public (OSS) Samples add_sample( sampleEditableTimingCache @@ -76,6 +69,7 @@ else() sampleNamedDimensions sampleOnnxMNIST sampleProgressMonitor + trtexec ) if (NOT ${TRT_PRODUCT_IS_RTX}) @@ -108,14 +102,9 @@ else() endif() endif() # TRT_BUILD_SAMPLES - if(${TRT_BUILD_ENABLE_UNIFIED_BUILDER} AND ${TRT_BUILD_TRTEXEC}) - add_sample(trtSafeExec) - add_sample(sampleSafeMNIST) - add_sample(sampleSafePluginV3) - elseif(BUILD_SAFE_SAMPLES) + if((${TRT_BUILD_ENABLE_UNIFIED_BUILDER} AND ${TRT_BUILD_SAMPLES}) OR BUILD_SAFE_SAMPLES) add_sample(sampleSafeMNIST) add_sample(sampleSafePluginV3) - add_sample(trtSafeExec) endif() endif() @@ -126,3 +115,4 @@ endif() foreach(FOLDER IN LISTS trtSampleFolders) add_subdirectory(${FOLDER}) endforeach() + diff --git a/samples/README.md b/samples/README.md index 606fc2c19b..8d4a423735 100644 --- a/samples/README.md +++ b/samples/README.md @@ -19,7 +19,6 @@ | [sampleNonZeroPlugin](sampleNonZeroPlugin) | C++ | INetwork | Adding plugin with data-dependent output shapes | | [sampleIOFormats](sampleIOFormats) | C++ | ONNX | Specifying TensorRT I/O Formats | | [sampleProgressMonitor](sampleProgressMonitor) | C++ | ONNX | Progress Monitor API usage | -| [trtexec](trtexec) | C++ | All | TensorRT Command-Line Wrapper: trtexec | | [engine_refit_onnx_bidaf](python/engine_refit_onnx_bidaf) | Python | ONNX | refitting an engine built from an ONNX model via parsers. | | [introductory_parser_samples](python/introductory_parser_samples) | Python | ONNX | Introduction To Importing Models Using TensorRT Parsers | | [onnx_packnet](python/onnx_packnet) | Python | ONNX | TensorRT Inference Of ONNX Models With Custom Layers | @@ -35,7 +34,6 @@ |---|---|---|---| | [sampleSafeMNIST](sampleSafeMNIST) | C++ | ONNX | Build a Safety Engine for MNIST | | [sampleSafePluginV3](sampleSafePluginV3) | C++ | ONNX | Use Safety-Supported Plugins With Safety Engines | -| [trtSafeExec](trtSafeExec) | C++ | ONNX | TensorRT Command-Line Wrapper With Safety Options | ## Preparing sample data diff --git a/samples/common/BatchStream.h b/samples/common/BatchStream.h deleted file mode 100644 index d12596e2c7..0000000000 --- a/samples/common/BatchStream.h +++ /dev/null @@ -1,381 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ -#ifndef BATCH_STREAM_H -#define BATCH_STREAM_H - -#include "NvInfer.h" -#include "common.h" -#include -#include -#include - -class IBatchStream -{ -public: - virtual void reset(int firstBatch) = 0; - virtual bool next() = 0; - virtual void skip(int skipCount) = 0; - virtual float* getBatch() = 0; - virtual float* getLabels() = 0; - virtual int getBatchesRead() const = 0; - virtual int getBatchSize() const = 0; - virtual nvinfer1::Dims getDims() const = 0; -}; - -class MNISTBatchStream : public IBatchStream -{ -public: - MNISTBatchStream(int batchSize, int maxBatches, const std::string& dataFile, const std::string& labelsFile, - const std::vector& directories) - : mBatchSize{batchSize} - , mMaxBatches{maxBatches} - , mDims{3, {1, 28, 28}} //!< We already know the dimensions of MNIST images. - { - readDataFile(samplesCommon::locateFile(dataFile, directories)); - readLabelsFile(samplesCommon::locateFile(labelsFile, directories)); - } - - void reset(int firstBatch) override - { - mBatchCount = firstBatch; - } - - bool next() override - { - if (mBatchCount >= mMaxBatches) - { - return false; - } - ++mBatchCount; - return true; - } - - void skip(int skipCount) override - { - mBatchCount += skipCount; - } - - float* getBatch() override - { - return mData.data() + (mBatchCount * mBatchSize * samplesCommon::volume(mDims)); - } - - float* getLabels() override - { - return mLabels.data() + (mBatchCount * mBatchSize); - } - - int getBatchesRead() const override - { - return mBatchCount; - } - - int getBatchSize() const override - { - return mBatchSize; - } - - nvinfer1::Dims getDims() const override - { - return nvinfer1::Dims{4, {mBatchSize, mDims.d[0], mDims.d[1], mDims.d[2]}}; - } - -private: - void readDataFile(const std::string& dataFilePath) - { - std::ifstream file{dataFilePath.c_str(), std::ios::binary}; - - int magicNumber, numImages, imageH, imageW; - file.read(reinterpret_cast(&magicNumber), sizeof(magicNumber)); - // All values in the MNIST files are big endian. - magicNumber = samplesCommon::swapEndianness(magicNumber); - ASSERT(magicNumber == 2051 && "Magic Number does not match the expected value for an MNIST image set"); - - // Read number of images and dimensions - file.read(reinterpret_cast(&numImages), sizeof(numImages)); - file.read(reinterpret_cast(&imageH), sizeof(imageH)); - file.read(reinterpret_cast(&imageW), sizeof(imageW)); - - numImages = samplesCommon::swapEndianness(numImages); - imageH = samplesCommon::swapEndianness(imageH); - imageW = samplesCommon::swapEndianness(imageW); - - // The MNIST data is made up of unsigned bytes, so we need to cast to float and normalize. - int numElements = numImages * imageH * imageW; - std::vector rawData(numElements); - file.read(reinterpret_cast(rawData.data()), numElements * sizeof(uint8_t)); - mData.resize(numElements); - std::transform( - rawData.begin(), rawData.end(), mData.begin(), [](uint8_t val) { return static_cast(val) / 255.F; }); - } - - void readLabelsFile(const std::string& labelsFilePath) - { - std::ifstream file{labelsFilePath.c_str(), std::ios::binary}; - int magicNumber, numImages; - file.read(reinterpret_cast(&magicNumber), sizeof(magicNumber)); - // All values in the MNIST files are big endian. - magicNumber = samplesCommon::swapEndianness(magicNumber); - ASSERT(magicNumber == 2049 && "Magic Number does not match the expected value for an MNIST labels file"); - - file.read(reinterpret_cast(&numImages), sizeof(numImages)); - numImages = samplesCommon::swapEndianness(numImages); - - std::vector rawLabels(numImages); - file.read(reinterpret_cast(rawLabels.data()), numImages * sizeof(uint8_t)); - mLabels.resize(numImages); - std::transform( - rawLabels.begin(), rawLabels.end(), mLabels.begin(), [](uint8_t val) { return static_cast(val); }); - } - - int mBatchSize{0}; - int mBatchCount{0}; //!< The batch that will be read on the next invocation of next() - int mMaxBatches{0}; - nvinfer1::Dims mDims{}; - std::vector mData{}; - std::vector mLabels{}; -}; - -class BatchStream : public IBatchStream -{ -public: - BatchStream(int batchSize, int maxBatches, std::string const& prefix, std::string const& suffix, - std::vector const& directories) - : mBatchSize(batchSize) - , mMaxBatches(maxBatches) - , mPrefix(prefix) - , mSuffix(suffix) - , mDataDir(directories) - { - std::ifstream file( - samplesCommon::locateFile(mPrefix + std::string("0") + mSuffix, mDataDir).c_str(), std::ios::binary); - ASSERT(file.good()); - int d[4]; - file.read(reinterpret_cast(d), 4 * sizeof(int32_t)); - mDims.nbDims = 4; // The number of dimensions. - mDims.d[0] = d[0]; // Batch Size - mDims.d[1] = d[1]; // Channels - mDims.d[2] = d[2]; // Height - mDims.d[3] = d[3]; // Width - ASSERT(mDims.d[0] > 0 && mDims.d[1] > 0 && mDims.d[2] > 0 && mDims.d[3] > 0); - - mImageSize = mDims.d[1] * mDims.d[2] * mDims.d[3]; - mBatch.resize(mBatchSize * mImageSize, 0); - mLabels.resize(mBatchSize, 0); - mFileBatch.resize(mDims.d[0] * mImageSize, 0); - mFileLabels.resize(mDims.d[0], 0); - } - - BatchStream(int batchSize, int maxBatches, std::string const& prefix, std::vector const& directories) - : BatchStream(batchSize, maxBatches, prefix, ".batch", directories) - { - } - - BatchStream(int batchSize, int maxBatches, nvinfer1::Dims const& dims, std::string const& listFile, - std::vector const& directories) - : mBatchSize(batchSize) - , mMaxBatches(maxBatches) - , mDims(dims) - , mListFile(listFile) - , mDataDir(directories) - { - mImageSize = mDims.d[1] * mDims.d[2] * mDims.d[3]; - mBatch.resize(mBatchSize * mImageSize, 0); - mLabels.resize(mBatchSize, 0); - mFileBatch.resize(mDims.d[0] * mImageSize, 0); - mFileLabels.resize(mDims.d[0], 0); - } - - // Resets data members - void reset(int firstBatch) override - { - mBatchCount = 0; - mFileCount = 0; - mFileBatchPos = mDims.d[0]; - skip(firstBatch); - } - - // Advance to next batch and return true, or return false if there is no batch left. - bool next() override - { - if (mBatchCount == mMaxBatches) - { - return false; - } - - for (int64_t csize = 1, batchPos = 0; batchPos < mBatchSize; batchPos += csize, mFileBatchPos += csize) - { - ASSERT(mFileBatchPos > 0 && mFileBatchPos <= mDims.d[0]); - if (mFileBatchPos == mDims.d[0] && !update()) - { - return false; - } - - // copy the smaller of: elements left to fulfill the request, or elements left in the file buffer. - csize = std::min(mBatchSize - batchPos, mDims.d[0] - mFileBatchPos); - std::copy_n( - getFileBatch() + mFileBatchPos * mImageSize, csize * mImageSize, getBatch() + batchPos * mImageSize); - std::copy_n(getFileLabels() + mFileBatchPos, csize, getLabels() + batchPos); - } - mBatchCount++; - return true; - } - - // Skips the batches - void skip(int skipCount) override - { - if (mBatchSize >= mDims.d[0] && mBatchSize % mDims.d[0] == 0 && mFileBatchPos == mDims.d[0]) - { - mFileCount += skipCount * mBatchSize / mDims.d[0]; - return; - } - - int x = mBatchCount; - for (int i = 0; i < skipCount; i++) - { - next(); - } - mBatchCount = x; - } - - float* getBatch() override - { - return mBatch.data(); - } - - float* getLabels() override - { - return mLabels.data(); - } - - int getBatchesRead() const override - { - return mBatchCount; - } - - int getBatchSize() const override - { - return mBatchSize; - } - - nvinfer1::Dims getDims() const override - { - return mDims; - } - -private: - float* getFileBatch() - { - return mFileBatch.data(); - } - - float* getFileLabels() - { - return mFileLabels.data(); - } - - bool update() - { - if (mListFile.empty()) - { - std::string inputFileName - = samplesCommon::locateFile(mPrefix + std::to_string(mFileCount++) + mSuffix, mDataDir); - std::ifstream file(inputFileName.c_str(), std::ios::binary); - if (!file) - { - return false; - } - int d[4]; - file.read(reinterpret_cast(d), 4 * sizeof(int32_t)); - ASSERT(mDims.d[0] == d[0] && mDims.d[1] == d[1] && mDims.d[2] == d[2] && mDims.d[3] == d[3]); - file.read(reinterpret_cast(getFileBatch()), sizeof(float) * mDims.d[0] * mImageSize); - file.read(reinterpret_cast(getFileLabels()), sizeof(float) * mDims.d[0]); - } - else - { - std::vector fNames; - std::ifstream file(samplesCommon::locateFile(mListFile, mDataDir), std::ios::binary); - if (!file) - { - return false; - } - - sample::gLogInfo << "Batch #" << mFileCount << std::endl; - file.seekg(((mBatchCount * mBatchSize)) * 7); - - for (int i = 1; i <= mBatchSize; i++) - { - std::string sName; - std::getline(file, sName); - sName = sName + ".ppm"; - sample::gLogInfo << "Calibrating with file " << sName << std::endl; - fNames.emplace_back(sName); - } - - mFileCount++; - - const int imageC = 3; - const int imageH = 300; - const int imageW = 300; - std::vector> ppms(fNames.size()); - for (uint32_t i = 0; i < fNames.size(); ++i) - { - readPPMFile(samplesCommon::locateFile(fNames[i], mDataDir), ppms[i]); - } - - std::vector data(samplesCommon::volume(mDims)); - const float scale = 2.0 / 255.0; - const float bias = 1.0; - long int volChl = mDims.d[2] * mDims.d[3]; - - // Normalize input data - for (int i = 0, volImg = mDims.d[1] * mDims.d[2] * mDims.d[3]; i < mBatchSize; ++i) - { - for (int c = 0; c < mDims.d[1]; ++c) - { - for (int j = 0; j < volChl; ++j) - { - data[i * volImg + c * volChl + j] = scale * float(ppms[i].buffer[j * mDims.d[1] + c]) - bias; - } - } - } - - std::copy_n(data.data(), mDims.d[0] * mImageSize, getFileBatch()); - } - - mFileBatchPos = 0; - return true; - } - - int64_t mBatchSize{0}; - int mMaxBatches{0}; - int mBatchCount{0}; - int mFileCount{0}; - int mFileBatchPos{0}; - int mImageSize{0}; - std::vector mBatch; //!< Data for the batch - std::vector mLabels; //!< Labels for the batch - std::vector mFileBatch; //!< List of image files - std::vector mFileLabels; //!< List of label files - std::string mPrefix; //!< Batch file name prefix - std::string mSuffix; //!< Batch file name suffix - nvinfer1::Dims mDims; //!< Input dimensions - std::string mListFile; //!< File name of the list of image names - std::vector mDataDir; //!< Directories where the files can be found -}; - -#endif diff --git a/samples/common/CMakeLists.txt b/samples/common/CMakeLists.txt index ad904de7bb..6e8478dbc2 100644 --- a/samples/common/CMakeLists.txt +++ b/samples/common/CMakeLists.txt @@ -15,53 +15,22 @@ # limitations under the License. # -if(NOT NLOHMANN_JSON_INCLUDE_DIRS) - include(FetchNlohmannJson) -endif() - add_library(trt_samples_common STATIC argsParser.h - BatchStream.h - bfloat16.cpp - bfloat16.h - bigInt.cpp - bigInt.h buffers.h - common.cpp common.h - debugTensorWriter.cpp - debugTensorWriter.h ErrorRecorder.h - getOptions.cpp - getOptions.h getoptWin.h - globalTimerKernel.cu - globalTimerKernel.h half.h logger.cpp logger.h logging.h parserOnnxConfig.h - sampleConfig.h - sampleDevice.cpp - sampleDevice.h - sampleEngines.cpp - sampleEngines.h sampleEntrypoints.h - sampleInference.cpp - sampleInference.h - sampleOptions.cpp - sampleOptions.h - sampleReporting.cpp - sampleReporting.h - sampleTuning.cpp - sampleTuning.h sampleUtils.cpp sampleUtils.h safeCommon.h - safeCudaAllocator.h safeErrorRecorder.h - streamReader.h ) if(MSVC) @@ -76,10 +45,7 @@ if (${TRT_BUILD_TESTING}) enable_testing() add_executable(trt_samples_common_test - bfloat16.test.cpp - getOptions.test.cpp half.test.cpp - sampleOptions.test.cpp sampleUtils.test.cpp ) @@ -99,7 +65,6 @@ endif() # TRT_BUILD_TESTING target_include_directories(trt_samples_common PUBLIC ${CMAKE_CURRENT_LIST_DIR} - ${NLOHMANN_JSON_INCLUDE_DIRS} ) target_link_libraries(trt_samples_common PUBLIC @@ -121,8 +86,6 @@ if(NOT WIN32 AND NOT ${CMAKE_SYSTEM_NAME} STREQUAL "QNX") target_link_libraries(trt_samples_common PUBLIC dl) endif() -target_link_libraries(trt_samples_common PUBLIC $) - # For statically-linked samples, we need to upgrade the link to always link TRT rather than letting the samples decide. if(${TRT_BUILD_SAMPLES_LINK_STATIC_TRT}) target_link_libraries(trt_samples_common PUBLIC diff --git a/samples/common/README.md b/samples/common/README.md index 0c87700aad..0f2c84a5dc 100644 --- a/samples/common/README.md +++ b/samples/common/README.md @@ -1,3 +1,3 @@ # samples/common -Shared utility library (`trt_samples_common`) used by all TensorRT C++ samples and `trtexec`. +Shared utility library (`trt_samples_common`) used by all TensorRT C++ samples. diff --git a/samples/common/common.h b/samples/common/common.h index d250d34dc3..22ee5df6a0 100644 --- a/samples/common/common.h +++ b/samples/common/common.h @@ -109,188 +109,11 @@ using namespace nvinfer1; #undef CHECK #define CHECK(status) CHECK_WITH_STREAM(status, std::cerr) -constexpr long double operator"" _GiB(long double val) -{ - return val * (1 << 30); -} -constexpr long double operator"" _MiB(long double val) -{ - return val * (1 << 20); -} -constexpr long double operator"" _KiB(long double val) -{ - return val * (1 << 10); -} - -struct SimpleProfiler : public nvinfer1::IProfiler -{ - struct Record - { - float time{0}; - int count{0}; - }; - - void reportLayerTime(const char* layerName, float ms) noexcept override - { - mProfile[layerName].count++; - mProfile[layerName].time += ms; - if (std::find(mLayerNames.begin(), mLayerNames.end(), layerName) == mLayerNames.end()) - { - mLayerNames.push_back(layerName); - } - } - - SimpleProfiler(const char* name, const std::vector& srcProfilers = std::vector()) - : mName(name) - { - for (const auto& srcProfiler : srcProfilers) - { - for (const auto& rec : srcProfiler.mProfile) - { - auto it = mProfile.find(rec.first); - if (it == mProfile.end()) - { - mProfile.insert(rec); - } - else - { - it->second.time += rec.second.time; - it->second.count += rec.second.count; - } - } - } - } - - friend std::ostream& operator<<(std::ostream& out, const SimpleProfiler& value) - { - out << "========== " << value.mName << " profile ==========" << std::endl; - float totalTime = 0; - std::string layerNameStr = "TensorRT layer name"; - int maxLayerNameLength = std::max(static_cast(layerNameStr.size()), 70); - for (const auto& elem : value.mProfile) - { - totalTime += elem.second.time; - maxLayerNameLength = std::max(maxLayerNameLength, static_cast(elem.first.size())); - } - - auto old_settings = out.flags(); - auto old_precision = out.precision(); - // Output header - { - out << std::setfill(' ') << std::setw(maxLayerNameLength) << layerNameStr << " "; - out << std::setw(12) << "Runtime, " - << "%" - << " "; - out << std::setw(12) << "Invocations" - << " "; - out << std::setw(12) << "Runtime, ms" << std::endl; - } - for (size_t i = 0; i < value.mLayerNames.size(); i++) - { - const std::string layerName = value.mLayerNames[i]; - auto elem = value.mProfile.at(layerName); - out << std::setw(maxLayerNameLength) << layerName << " "; - out << std::setw(12) << std::fixed << std::setprecision(1) << (elem.time * 100.0F / totalTime) << "%" - << " "; - out << std::setw(12) << elem.count << " "; - out << std::setw(12) << std::fixed << std::setprecision(2) << elem.time << std::endl; - } - out.flags(old_settings); - out.precision(old_precision); - out << "========== " << value.mName << " total runtime = " << totalTime << " ms ==========" << std::endl; - - return out; - } - -private: - std::string mName; - std::vector mLayerNames; - std::map mProfile; -}; - namespace samplesCommon { -using nvinfer1::utils::loadCacheFile; using nvinfer1::utils::buildTimingCacheFromFile; -using nvinfer1::utils::saveCacheFile; using nvinfer1::utils::updateTimingCacheFile; -//! \brief Swaps endianness of an integral type. -template >> -[[nodiscard]] T swapEndianness(T value) -{ - uint8_t bytes[sizeof(T)]; - std::memcpy(bytes, &value, sizeof(T)); - std::reverse(std::begin(bytes), std::end(bytes)); - std::memcpy(&value, bytes, sizeof(T)); - return value; -} - -class HostMemory -{ -public: - HostMemory() = delete; - virtual void* data() const noexcept - { - return mData; - } - virtual std::size_t size() const noexcept - { - return mSize; - } - virtual nvinfer1::DataType type() const noexcept - { - return mType; - } - virtual ~HostMemory() {} - -protected: - HostMemory(std::size_t size, nvinfer1::DataType type) - : mData{nullptr} - , mSize(size) - , mType(type) - { - } - void* mData; - std::size_t mSize; - nvinfer1::DataType mType; -}; - -template -class TypedHostMemory : public HostMemory -{ -public: - explicit TypedHostMemory(std::size_t size) - : HostMemory(size, dataType) - { - mData = new ElemType[size]; - } - ~TypedHostMemory() noexcept override - { - delete[] (ElemType*) mData; - } - ElemType* raw() noexcept - { - return static_cast(data()); - } -}; - -using FloatMemory = TypedHostMemory; -using HalfMemory = TypedHostMemory; -using ByteMemory = TypedHostMemory; - -inline void* safeCudaMalloc(size_t memSize) -{ - void* deviceMem; - CHECK(cudaMalloc(&deviceMem, memSize)); - if (deviceMem == nullptr) - { - std::cerr << "Out of memory" << std::endl; - exit(EXIT_FAILURE); - } - return deviceMem; -} - inline bool isDebug() { return std::getenv("TENSORRT_DEBUG") != nullptr; @@ -325,25 +148,6 @@ std::vector argMagnitudeSort(Iter begin, Iter end) return indices; } -inline bool readReferenceFile(const std::string& fileName, std::vector& refVector) -{ - std::ifstream infile(fileName); - if (!infile.is_open()) - { - std::cout << "ERROR: readReferenceFile: Attempting to read from a file that is not open." << std::endl; - return false; - } - std::string line; - while (std::getline(infile, line)) - { - if (line.empty()) - continue; - refVector.push_back(line); - } - infile.close(); - return true; -} - template std::vector classify( const std::vector& refVector, const std::vector& output, const size_t topK) @@ -358,122 +162,6 @@ std::vector classify( return result; } -// Returns indices of highest K magnitudes in v. -template -std::vector topKMagnitudes(const std::vector& v, const size_t k) -{ - std::vector indices = samplesCommon::argMagnitudeSort(v.cbegin(), v.cend()); - indices.resize(k); - return indices; -} - -template -bool readASCIIFile(const std::string& fileName, const size_t size, std::vector& out) -{ - std::ifstream infile(fileName); - if (!infile.is_open()) - { - std::cout << "ERROR readASCIIFile: Attempting to read from a file that is not open." << std::endl; - return false; - } - out.clear(); - out.reserve(size); - out.assign(std::istream_iterator(infile), std::istream_iterator()); - infile.close(); - return true; -} - -template -bool writeASCIIFile(const std::string& fileName, const std::vector& in) -{ - std::ofstream outfile(fileName); - if (!outfile.is_open()) - { - std::cout << "ERROR: writeASCIIFile: Attempting to write to a file that is not open." << std::endl; - return false; - } - for (auto fn : in) - { - outfile << fn << "\n"; - } - outfile.close(); - return true; -} - -inline void print_version() -{ - std::cout << " TensorRT version: " << NV_TENSORRT_MAJOR << "." << NV_TENSORRT_MINOR << "." << NV_TENSORRT_PATCH - << "." << NV_TENSORRT_BUILD << std::endl; -} - -inline std::string getFileType(const std::string& filepath) -{ - return filepath.substr(filepath.find_last_of(".") + 1); -} - -inline std::string toLower(const std::string& inp) -{ - std::string out = inp; - std::transform(out.begin(), out.end(), out.begin(), ::tolower); - return out; -} - -inline float getMaxValue(const float* buffer, int64_t size) -{ - assert(buffer != nullptr); - assert(size > 0); - return *std::max_element(buffer, buffer + size); -} - -#if !TRT_WINML && ENABLE_FEATURE_WEAK_TYPING -inline void setAllDynamicRanges(nvinfer1::INetworkDefinition* network, float inRange = 2.0F, float outRange = 4.0F) -{ - for (int i = 0; i < network->getNbLayers(); i++) - { - auto layer = network->getLayer(i); - for (int j = 0; j < layer->getNbInputs(); j++) - { - nvinfer1::ITensor* input{layer->getInput(j)}; - if (input != nullptr && !input->dynamicRangeIsSet()) - { - ASSERT(input->setDynamicRange(-inRange, inRange)); - } - } - } - - for (int i = 0; i < network->getNbLayers(); i++) - { - auto layer = network->getLayer(i); - for (int j = 0; j < layer->getNbOutputs(); j++) - { - nvinfer1::ITensor* output{layer->getOutput(j)}; - if (output != nullptr && !output->dynamicRangeIsSet()) - { - if (layer->getType() == nvinfer1::LayerType::kPOOLING) - { - ASSERT(output->setDynamicRange(-inRange, inRange)); - } - else - { - ASSERT(output->setDynamicRange(-outRange, outRange)); - } - } - } - } -} - -inline void setDummyInt8DynamicRanges(nvinfer1::IBuilderConfig const* c, nvinfer1::INetworkDefinition* n) -{ - if (c->getFlag(nvinfer1::BuilderFlag::kINT8)) - { - sample::gLogWarning << "No per-tensor dynamic range provided. Generating dummy values. Int8 accuracy " - "is not guaranteed." - << std::endl; - setAllDynamicRanges(n); - } -} -#endif // !TRT_WINML && ENABLE_FEATURE_WEAK_TYPING - inline void enableDLA( nvinfer1::IBuilder* builder, nvinfer1::IBuilderConfig* config, int useDLACore, bool allowGPUFallback = true) { @@ -500,24 +188,6 @@ inline void enableDLA( } } -//! \brief Matches a flag prefix in an argument, ignoring leading spaces. -//! \param arg The command-line argument to check. -//! \param flag The flag prefix to match (e.g., "--loadEngine="). -//! \return A string_view of the remainder after \p flag, or nullopt if \p flag isn't found. -[[nodiscard]] std::optional matchFlag(std::string_view arg, std::string_view flag); - -//! \overload std::optional matchFlag(std::string_view arg, std::string_view flag) to prevent -//! accidental use of `std::string&&` arguments which would produce a dangling view, but allow e.g., `char const*`. -template -[[nodiscard]] std::optional matchFlag(StringViewable&& arg, std::string_view flag) -{ - static_assert(!std::is_rvalue_reference_v, - "You don't want the above matchFlag with `std::string&&` arguments which would produce a dangling view."); - return matchFlag(std::string_view{arg}, flag); -} - -int32_t parseDLA(int32_t argc, char** argv); - inline size_t getNbBytes(nvinfer1::DataType t, int64_t vol) noexcept { switch (t) @@ -640,244 +310,6 @@ inline void readPGMFile(const std::string& fileName, uint8_t* buffer, int32_t in infile.seekg(1, infile.cur); infile.read(reinterpret_cast(buffer), inH * inW); } -template -struct PPM -{ - std::string magic, fileName; - int h, w, max; - uint8_t buffer[C * H * W]; -}; - -// New vPPM(variable sized PPM) class with variable dimensions. -struct vPPM -{ - std::string magic, fileName; - int h, w, max; - std::vector buffer; -}; - -struct BBox -{ - float x1, y1, x2, y2; -}; - -template -void readPPMFile(const std::string& filename, samplesCommon::PPM& ppm) -{ - ppm.fileName = filename; - std::ifstream infile(filename, std::ifstream::binary); - assert(infile.is_open() && "Attempting to read from a file that is not open."); - infile >> ppm.magic >> ppm.w >> ppm.h >> ppm.max; - infile.seekg(1, infile.cur); - infile.read(reinterpret_cast(ppm.buffer), ppm.w * ppm.h * 3); -} - -inline void readPPMFile(const std::string& filename, vPPM& ppm, std::vector& input_dir) -{ - ppm.fileName = filename; - std::ifstream infile(locateFile(filename, input_dir), std::ifstream::binary); - infile >> ppm.magic >> ppm.w >> ppm.h >> ppm.max; - infile.seekg(1, infile.cur); - - for (int i = 0; i < ppm.w * ppm.h * 3; ++i) - { - ppm.buffer.push_back(0); - } - - infile.read(reinterpret_cast(&ppm.buffer[0]), ppm.w * ppm.h * 3); -} - -template -void writePPMFileWithBBox(const std::string& filename, PPM& ppm, const BBox& bbox) -{ - std::ofstream outfile("./" + filename, std::ofstream::binary); - assert(!outfile.fail()); - outfile << "P6" - << "\n" - << ppm.w << " " << ppm.h << "\n" - << ppm.max << "\n"; - - auto round = [](float x) -> int { return int(std::floor(x + 0.5F)); }; - const int x1 = std::min(std::max(0, round(int(bbox.x1))), W - 1); - const int x2 = std::min(std::max(0, round(int(bbox.x2))), W - 1); - const int y1 = std::min(std::max(0, round(int(bbox.y1))), H - 1); - const int y2 = std::min(std::max(0, round(int(bbox.y2))), H - 1); - - for (int x = x1; x <= x2; ++x) - { - // bbox top border - ppm.buffer[(y1 * ppm.w + x) * 3] = 255; - ppm.buffer[(y1 * ppm.w + x) * 3 + 1] = 0; - ppm.buffer[(y1 * ppm.w + x) * 3 + 2] = 0; - // bbox bottom border - ppm.buffer[(y2 * ppm.w + x) * 3] = 255; - ppm.buffer[(y2 * ppm.w + x) * 3 + 1] = 0; - ppm.buffer[(y2 * ppm.w + x) * 3 + 2] = 0; - } - - for (int y = y1; y <= y2; ++y) - { - // bbox left border - ppm.buffer[(y * ppm.w + x1) * 3] = 255; - ppm.buffer[(y * ppm.w + x1) * 3 + 1] = 0; - ppm.buffer[(y * ppm.w + x1) * 3 + 2] = 0; - // bbox right border - ppm.buffer[(y * ppm.w + x2) * 3] = 255; - ppm.buffer[(y * ppm.w + x2) * 3 + 1] = 0; - ppm.buffer[(y * ppm.w + x2) * 3 + 2] = 0; - } - - outfile.write(reinterpret_cast(ppm.buffer), ppm.w * ppm.h * 3); -} - -inline void writePPMFileWithBBox(const std::string& filename, vPPM ppm, std::vector& dets) -{ - std::ofstream outfile("./" + filename, std::ofstream::binary); - assert(!outfile.fail()); - outfile << "P6" - << "\n" - << ppm.w << " " << ppm.h << "\n" - << ppm.max << "\n"; - auto round = [](float x) -> int { return int(std::floor(x + 0.5F)); }; - - for (auto bbox : dets) - { - for (int x = int(bbox.x1); x < int(bbox.x2); ++x) - { - // bbox top border - ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3] = 255; - ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3 + 1] = 0; - ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3 + 2] = 0; - // bbox bottom border - ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3] = 255; - ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3 + 1] = 0; - ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3 + 2] = 0; - } - - for (int y = int(bbox.y1); y < int(bbox.y2); ++y) - { - // bbox left border - ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3] = 255; - ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3 + 1] = 0; - ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3 + 2] = 0; - // bbox right border - ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3] = 255; - ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3 + 1] = 0; - ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3 + 2] = 0; - } - } - - outfile.write(reinterpret_cast(&ppm.buffer[0]), ppm.w * ppm.h * 3); -} - -class TimerBase -{ -public: - virtual void start() {} - virtual void stop() {} - float microseconds() const noexcept - { - return mMs * 1000.F; - } - float milliseconds() const noexcept - { - return mMs; - } - float seconds() const noexcept - { - return mMs / 1000.F; - } - void reset() noexcept - { - mMs = 0.F; - } - -protected: - float mMs{0.0F}; -}; - -class GpuTimer : public TimerBase -{ -public: - explicit GpuTimer(cudaStream_t stream) - : mStream(stream) - { - CHECK(cudaEventCreate(&mStart)); - CHECK(cudaEventCreate(&mStop)); - } - ~GpuTimer() - { - CHECK(cudaEventDestroy(mStart)); - CHECK(cudaEventDestroy(mStop)); - } - void start() override - { - CHECK(cudaEventRecord(mStart, mStream)); - } - void stop() override - { - CHECK(cudaEventRecord(mStop, mStream)); - float ms{0.0F}; - CHECK(cudaEventSynchronize(mStop)); - CHECK(cudaEventElapsedTime(&ms, mStart, mStop)); - mMs += ms; - } - -private: - cudaEvent_t mStart, mStop; - cudaStream_t mStream; -}; // class GpuTimer - -template -class CpuTimer : public TimerBase -{ -public: - using clock_type = Clock; - - void start() override - { - mStart = Clock::now(); - } - void stop() override - { - mStop = Clock::now(); - mMs += std::chrono::duration{mStop - mStart}.count(); - } - -private: - std::chrono::time_point mStart, mStop; -}; // class CpuTimer - -using PreciseCpuTimer = CpuTimer; - -inline std::vector splitString(std::string str, char delimiter = ',') -{ - std::vector splitVect; - std::stringstream ss(str); - std::string substr; - - while (ss.good()) - { - getline(ss, substr, delimiter); - splitVect.emplace_back(std::move(substr)); - } - return splitVect; -} - -inline int getC(nvinfer1::Dims const& d) -{ - return d.nbDims >= 3 ? d.d[d.nbDims - 3] : 1; -} - -inline int getH(const nvinfer1::Dims& d) -{ - return d.nbDims >= 2 ? d.d[d.nbDims - 2] : 1; -} - -inline int getW(const nvinfer1::Dims& d) -{ - return d.nbDims >= 1 ? d.d[d.nbDims - 1] : 1; -} //! Platform-agnostic wrapper around dynamic libraries. class DynamicLibrary @@ -1009,19 +441,6 @@ inline bool isSmSafe() || smVersion == 0x0A00 || smVersion == 0x0B00; } -inline int32_t getMaxPersistentCacheSize() -{ - int32_t deviceIndex{}; - CHECK(cudaGetDevice(&deviceIndex)); - - int32_t maxPersistentL2CacheSize{}; -#if CUDART_VERSION >= 11030 && !TRT_WINML - CHECK(cudaDeviceGetAttribute(&maxPersistentL2CacheSize, cudaDevAttrMaxPersistingL2CacheSize, deviceIndex)); -#endif - - return maxPersistentL2CacheSize; -} - } // namespace samplesCommon inline std::ostream& operator<<(std::ostream& os, const nvinfer1::Dims& dims) @@ -1034,32 +453,4 @@ inline std::ostream& operator<<(std::ostream& os, const nvinfer1::Dims& dims) return os << ")"; } -[[nodiscard]] inline std::string genFilenameSafeString(std::string_view s) -{ - std::string_view const kALLOWED{"._-,"}; - constexpr size_t kMAX_FILENAME_LENGTH = 150; // Leave some margin due to Windows path length limitation - constexpr size_t kELLIPSIS_LENGTH = 3; // Length of "..." - - auto processChar = [&kALLOWED](char c) { - return std::isalnum(static_cast(c)) || kALLOWED.find(c) != std::string_view::npos ? c : '_'; - }; - - std::string res; - if (s.length() <= kMAX_FILENAME_LENGTH) - { - res.reserve(s.size()); - std::transform(s.begin(), s.end(), std::back_inserter(res), processChar); - return res; - } - - res.reserve(kMAX_FILENAME_LENGTH); - size_t const halfLength = (kMAX_FILENAME_LENGTH - kELLIPSIS_LENGTH) / 2; - - std::transform(s.begin(), s.begin() + halfLength, std::back_inserter(res), processChar); - res += "..."; - std::transform(s.end() - halfLength, s.end(), std::back_inserter(res), processChar); - - return res; -} - #endif // TENSORRT_COMMON_H diff --git a/samples/common/dumpTFWts.py b/samples/common/dumpTFWts.py deleted file mode 100644 index 70770fbd80..0000000000 --- a/samples/common/dumpTFWts.py +++ /dev/null @@ -1,124 +0,0 @@ -#!/usr/bin/python -# -# SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# - -# Script to dump TensorFlow weights in TRT v1 and v2 dump format. -# The V1 format is for TensorRT 4.0. The V2 format is for TensorRT 4.0 and later. - -import sys -import struct -import argparse - -try: - import tensorflow as tf - from tensorflow.python import pywrap_tensorflow -except ImportError as err: - sys.stderr.write("""Error: Failed to import module ({})""".format(err)) - sys.exit() - -parser = argparse.ArgumentParser(description="TensorFlow Weight Dumper") - -parser.add_argument( - "-m", - "--model", - required=True, - help="The checkpoint file basename, example basename(model.ckpt-766908.data-00000-of-00001) -> model.ckpt-766908", -) -parser.add_argument("-o", "--output", required=True, help="The weight file to dump all the weights to.") -parser.add_argument("-1", "--wtsv1", required=False, default=False, type=bool, help="Dump the weights in the wts v1.") - -opt = parser.parse_args() - -if opt.wtsv1: - print("Outputting the trained weights in TensorRT's wts v1 format. This format is documented as:") - print("Line 0: ") - print("Line 1-Num: [buffer name] [buffer type] [buffer size] ") -else: - print("Outputting the trained weights in TensorRT's wts v2 format. This format is documented as:") - print("Line 0: ") - print("Line 1-Num: [buffer name] [buffer type] [(buffer shape{e.g. (1, 2, 3)}] ") - -inputbase = opt.model -outputbase = opt.output - - -def float_to_hex(f): - return hex(struct.unpack(" -#include -#include -#include -#include - -namespace nvinfer1::utility -{ - -namespace -{ - -using namespace std::string_view_literals; -using sample::gLogWarning; - -//! Matching for TRTOptions is defined as follows: -//! -//! If A and B both have longName set, A matches B if and only if A.longName == -//! B.longName and (A.shortName == B.shortName if both have short name set). -//! -//! If A only has shortName set and B only has longName set, then A does not -//! match B. It is assumed that when 2 TRTOptions are compared, one of them is -//! the definition of a TRTOption in the input to getOptions. As such, if the -//! definition only has shortName set, it will never be equal to a TRTOption -//! that does not have shortName set (and same for longName). -//! -//! If A and B both have shortName set but B does not have longName set, A -//! matches B if and only if A.shortName == B.shortName. -//! -//! If A has neither long or short name set, A matches B if and only if B has -//! neither long or short name set. -[[nodiscard]] bool matches(TRTOption const& a, TRTOption const& b) -{ - if (!a.longName.empty() && !b.longName.empty()) - { - if (a.shortName != '\0' && b.shortName != '\0') - { - return (a.longName == b.longName) && (a.shortName == b.shortName); - } - return a.longName == b.longName; - } - - // If only one of them is not set, this will return false anyway. - return a.shortName == b.shortName; -} - -//! getTRTOptionIndex returns the index of a TRTOption in a vector of -//! TRTOptions, -1 if not found. -[[nodiscard]] int32_t getTRTOptionIndex(std::vector const& options, TRTOption const& opt) -{ - auto it = std::find_if( - options.begin(), options.end(), [&opt](TRTOption const& option) { return matches(opt, option); }); - return it != options.end() ? static_cast(std::distance(options.begin(), it)) : -1; -} - -//! validateTRTOption will return a string containing an error message if options -//! contain non-numeric characters, or if there are duplicate option names found. -//! Otherwise, returns the empty string. -[[nodiscard]] std::string validateTRTOption( - std::set const& seenShortNames, std::set const& seenLongNames, TRTOption const& opt) -{ - if (opt.shortName != '\0') - { - if (!std::isalnum(opt.shortName)) - { - return "Short name '" + std::to_string(opt.shortName) + "' is non-alphanumeric"; - } - - if (seenShortNames.contains(opt.shortName)) - { - return "Short name '" + std::to_string(opt.shortName) + "' is a duplicate"; - } - } - - if (!opt.longName.empty()) - { - for (char const& c : opt.longName) - { - if (!std::isalnum(c) && c != '-' && c != '_') - { - return "Long name '" + opt.longName + "' contains characters that are not '-', '_', or alphanumeric"; - } - } - - if (seenLongNames.contains(opt.longName)) - { - return "Long name '" + opt.longName + "' is a duplicate"; - } - } - return ""; -} - -//! validateTRTOptions will return a string containing an error message if any -//! options contain non-numeric characters, or if there are duplicate option -//! names found. Otherwise, returns the empty string. -[[nodiscard]] std::string validateTRTOptions(std::vector const& options) -{ - std::set seenShortNames; - std::set seenLongNames; - for (size_t i = 0; i < options.size(); ++i) - { - std::string const errMsg = validateTRTOption(seenShortNames, seenLongNames, options[i]); - if (!errMsg.empty()) - { - return "Error '" + errMsg + "' at TRTOption " + std::to_string(i); - } - - seenShortNames.insert(options[i].shortName); - seenLongNames.insert(options[i].longName); - } - return ""; -} - -//! Structure to hold a parsed option and its inline value (if any) -struct ParsedOption -{ - TRTOption opt; - std::string inlineValue; -}; - -//! Parse an option string (starting with '-' or '--') into a TRTOption and optional inline value. -//! \param[in] argStr The option string to parse. -//! \param[out] result The parsed option and inline value. -//! \return error message if parsing fails, empty string otherwise. -[[nodiscard]] std::string parseOptionString(std::string_view const argStr, ParsedOption& result) -{ - // C++23: Return a `std::expected` instead. - if (argStr.size() < 2) - { - return "Option string is too short"; - } - if (argStr[1] != '-') - { - // Short option: must only have 1 char after the hyphen - if (argStr.size() > 2) - { - return "Short arg contains more than 1 character"; - } - result = ParsedOption{TRTOption{argStr[1]}}; - return {}; - } - else - { - // Long option: extract name and check for --foo=bar syntax - auto longName = argStr.substr(2); - size_t const eqIndex = longName.find('='); - - auto inlineValue = eqIndex != std::string_view::npos ? longName.substr(eqIndex + 1) : ""sv; - - // Note: If `eqIndex == std::string_view::npos`, then `longName.substr(0, eqIndex)` is the entire string_view. - result = ParsedOption{TRTOption{{}, std::string{longName.substr(0, eqIndex)}}, std::string{inlineValue}}; - return {}; - } -} - -//! Handle an option that requires a value. Returns error message if value cannot be obtained. -//! Updates currentArgIdx if a value is consumed from the next argument. -[[nodiscard]] std::string handleRequiredValue(TRTParsedArgs& parsedArgs, int32_t idx, std::string inlineValue, - std::string_view const argStr, int32_t& currentArgIdx, int32_t argc, char const* const* argv) -{ - // If we have an inline value (from --foo=bar), use it - if (!inlineValue.empty()) - { - parsedArgs.values[idx].addOccurrence(std::move(inlineValue)); - return {}; - } - - // Otherwise, consume the next argument as the value - if (currentArgIdx + 1 >= argc) - { - return "Last argument requires value, but none given"; - } - - std::string_view const nextArg(argv[currentArgIdx + 1]); - if (!nextArg.empty() && nextArg[0] == '-') - { - gLogWarning << "Warning: Using '" << nextArg << "' as a value for '" << argStr - << "', Should this be its own flag?" << std::endl; - } - - parsedArgs.values[idx].addOccurrence(std::string{nextArg}); - ++currentArgIdx; // Next argument consumed - return {}; -} - -//! parseArgs parses an argument list and returns a TRTParsedArgs with the -//! fields set accordingly. Assumes that options is validated. -//! ErrMsg will be set if: -//! - an argument is null -//! - an argument is empty -//! - an argument does not have option (i.e. "-" and "--") -//! - a short argument has more than 1 character -//! - the last argument in the list requires a value -[[nodiscard]] TRTParsedArgs parseArgs( - int32_t const argc, char const* const* const argv, std::vector const& options) -{ - TRTParsedArgs parsedArgs; - parsedArgs.values.resize(options.size()); - - for (int32_t i = 1; i < argc; ++i) // index of current command-line argument - { - if (argv[i] == nullptr) - { - return TRTParsedArgs{"Null argument at index " + std::to_string(i)}; - } - - std::string_view const argStr(argv[i]); - if (argStr.empty()) - { - return TRTParsedArgs{"Empty argument at index " + std::to_string(i)}; - } - - // No starting hyphen means it is a positional argument - if (argStr[0] != '-') - { - parsedArgs.positionalArgs.push_back(std::string{argStr}); - continue; - } - if (argStr == "-"sv || argStr == "--"sv) - { - return TRTParsedArgs{"Argument does not specify an option at index " + std::to_string(i)}; - } - - // Parse the option string - ParsedOption parsed; - if (std::string const parseErr = parseOptionString(argStr, parsed); !parseErr.empty()) - { - return TRTParsedArgs{parseErr + " at index " + std::to_string(i)}; - } - - // Find the option in the registered options list - int32_t const idx = getTRTOptionIndex(options, parsed.opt); - if (idx < 0) - { - continue; - } - - // Handle value-required options vs. flag options - if (options[idx].valueRequired) - { - if (std::string valueErr = handleRequiredValue(parsedArgs, idx, parsed.inlineValue, argStr, i, argc, argv); - !valueErr.empty()) - { - return TRTParsedArgs{std::move(valueErr)}; - } - } - else - { - parsedArgs.values[idx].addOccurrence(); - } - } - return parsedArgs; -} - -} // namespace - -TRTParsedArgs getOptions(int32_t argc, char const* const* argv, std::vector const& options) -{ - if (std::string errMsg = validateTRTOptions(options); !errMsg.empty()) - { - return TRTParsedArgs{std::move(errMsg)}; - } - - return parseArgs(argc, argv, options); -} -} // namespace nvinfer1::utility diff --git a/samples/common/getOptions.h b/samples/common/getOptions.h deleted file mode 100644 index 34dbee7cba..0000000000 --- a/samples/common/getOptions.h +++ /dev/null @@ -1,135 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -#ifndef TRT_GET_OPTIONS_H -#define TRT_GET_OPTIONS_H - -#include -#include -#include - -namespace nvinfer1::utility -{ - -//! TRTOption defines a command line option. At least 1 of shortName and longName -//! must be defined. -//! If bool initialization is undefined behavior on your system, valueRequired -//! must also be explicitly defined. -//! helpText is optional. -struct TRTOption -{ - char shortName{}; //!< Option name in short (single hyphen) form (e.g., -a, -b); '\0' for no short name. - std::string longName; //!< Option name in long (double hyphen) form (e.g., --foo, --bar); empty for no long name. - bool valueRequired{}; //!< True if a value is needed for an option (e.g., -N 4, --foo bar); false for not required. - std::string helpText; //!< Text to show when printing out the command usage -}; - -//! TRTParsedArgs is returned by getOptions after it has parsed a command line -//! argument list (argv). -struct TRTParsedArgs -{ - //! An error message if any errors occurred. Empty if no errors occurred. - std::string errMsg; - - //! A value for an option. - struct Value - { - //! The number of occurrences of the option (for value-required options, this equals `values.size()`). - int32_t occurrences{}; - //! The values for the option. For non-value args, will be empty. - std::vector values; - //! Increment the number of occurrences (for a non-value arg). - void addOccurrence() - { - ++occurrences; - } - //! Append \p value and set \p occurrences to the number of values. - void addOccurrence(std::string value) - { - values.push_back(std::move(value)); - occurrences = values.size(); - } - }; - //! A list of values for each option. - std::vector values; - //! Positional arguments that are passed in without an option (these must not start with a hyphen). - std::vector positionalArgs; -}; - -//! Parse the input arguments passed to main() and extract options as well as -//! positional arguments. -//! -//! Options are supposed to be passed to main() with a preceding hyphen '-'. -//! -//! If there is a single preceding hyphen, there should be exactly 1 character -//! after the hyphen, which is interpreted as the option. -//! -//! If there are 2 preceding hyphens, the entire argument (without the hyphens) -//! is interpreted as the option. -//! -//! If the option requires a value, the next argument is used as the value. -//! -//! Positional arguments must not start with a hyphen. -//! -//! If an argument requires a value, the next argument is interpreted as the -//! value, even if it is the form of a valid option (i.e. --foo --bar will store -//! "--bar" as a value for option "foo" if "foo" requires a value). -//! We also support --name=value syntax. In this case, 'value' would be used as -//! the value, NOT the next argument. -//! -//! For options: -//! { { 'a', "", false }, -//! { 'b', "", false }, -//! { 0, "cee", false }, -//! { 'd', "", true }, -//! { 'e', "", true }, -//! { 'f', "foo", true } } -//! -//! ./main hello world -a -a --cee -d 12 -f 34 -//! and -//! ./main hello world -a -a --cee -d 12 --foo 34 -//! -//! will result in: -//! -//! TRTParsedArgs { -//! errMsg: "", -//! values: { { 2, {} }, -//! { 0, {} }, -//! { 1, {} }, -//! { 1, {"12"} }, -//! { 0, {} }, -//! { 1, {"34"} } } -//! positionalArgs: {"hello", "world"}, -//! } -//! -//! Non-POSIX behavior: -//! - Does not support "-abcde" as a shorthand for "-a -b -c -d -e". Each -//! option must have its own hyphen prefix. -//! - Does not support -e12 as a shorthand for "-e 12". Values MUST be -//! whitespace-separated from the option it is for. -//! -//! @param[in] argc The number of arguments passed to main (including the -//! file name, which is disregarded) -//! @param[in] argv The arguments passed to main (including the file name, -//! which is disregarded) -//! @param[in] options List of TRTOptions to parse -//! @return TRTParsedArgs. See TRTParsedArgs documentation for descriptions of -//! the fields. -[[nodiscard]] TRTParsedArgs getOptions(int argc, char const* const* argv, std::vector const& options); -} // namespace nvinfer1::utility - -#endif // TRT_GET_OPTIONS_H diff --git a/samples/common/getOptions.test.cpp b/samples/common/getOptions.test.cpp deleted file mode 100644 index a6261e4d4c..0000000000 --- a/samples/common/getOptions.test.cpp +++ /dev/null @@ -1,145 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -#include "getOptions.h" -#include "ArgVec.test.h" - -#include - -#include - -using nvinfer1::utility::getOptions; -using nvinfer1::utility::TRTOption; -using namespace std::string_view_literals; - -using TestArgVec = ArgVec; - -// The options used by several tests below, matching the header's worked example. -static std::vector const kEXAMPLE_OPTIONS{ - {'a', "", false}, - {'b', "", false}, - {'\0', "cee", false}, - {'d', "", true}, - {'e', "", true}, - {'f', "foo", true}, -}; - -TEST(GetOptions, PositionalArgs) -{ - TestArgVec av{"hello", "world"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - ASSERT_EQ(result.positionalArgs.size(), 2U); - EXPECT_EQ(result.positionalArgs[0], "hello"sv); - EXPECT_EQ(result.positionalArgs[1], "world"sv); -} - -TEST(GetOptions, ShortFlag) -{ - TestArgVec av{"-a"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - EXPECT_EQ(result.values[0].occurrences, 1); // 'a' is index 0 -} - -TEST(GetOptions, ShortFlagRepeated) -{ - TestArgVec av{"-a", "-a"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - EXPECT_EQ(result.values[0].occurrences, 2); -} - -TEST(GetOptions, LongFlag) -{ - TestArgVec av{"--cee"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - EXPECT_EQ(result.values[2].occurrences, 1); // "cee" is index 2 -} - -TEST(GetOptions, ShortValueSpaceSeparated) -{ - TestArgVec av{"-d", "12"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - ASSERT_EQ(result.values[3].occurrences, 1); // 'd' is index 3 - EXPECT_EQ(result.values[3].values[0], "12"sv); -} - -TEST(GetOptions, LongValueEqualsSign) -{ - TestArgVec av{"--foo=34"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - ASSERT_EQ(result.values[5].occurrences, 1); // "foo" is index 5 - EXPECT_EQ(result.values[5].values[0], "34"sv); -} - -TEST(GetOptions, ExactExampleFromHeader) -{ - // ./main hello world -a -a --cee -d 12 -f 34 - TestArgVec av{"hello", "world", "-a", "-a", "--cee", "-d", "12", "-f", "34"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - EXPECT_EQ(result.values[0].occurrences, 2); // 'a' - EXPECT_EQ(result.values[1].occurrences, 0); // 'b' - EXPECT_EQ(result.values[2].occurrences, 1); // "cee" - ASSERT_EQ(result.values[3].occurrences, 1); // 'd' - EXPECT_EQ(result.values[3].values[0], "12"sv); - EXPECT_EQ(result.values[4].occurrences, 0); // 'e' - ASSERT_EQ(result.values[5].occurrences, 1); // "foo"/"f" - EXPECT_EQ(result.values[5].values[0], "34"sv); - ASSERT_EQ(result.positionalArgs.size(), 2U); - EXPECT_EQ(result.positionalArgs[0], "hello"sv); - EXPECT_EQ(result.positionalArgs[1], "world"sv); -} - -TEST(GetOptions, UnknownOptionsIgnored) -{ - TestArgVec av{"--unknown-flag"}; - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_TRUE(result.errMsg.empty()); - for (auto const& v : result.values) - { - EXPECT_EQ(v.occurrences, 0); - } -} - -TEST(GetOptions, MissingRequiredValue) -{ - TestArgVec av{"-d"}; // 'd' requires a value but none is given - auto const result = getOptions(av.argc(), av.argv(), kEXAMPLE_OPTIONS); - EXPECT_FALSE(result.errMsg.empty()); -} - -TEST(GetOptions, DuplicateShortName) -{ - std::vector const opts{{'a', "", false}, {'a', "other", false}}; - TestArgVec av{}; - auto const result = getOptions(av.argc(), av.argv(), opts); - EXPECT_FALSE(result.errMsg.empty()); -} - -TEST(GetOptions, EmptyOptions) -{ - TestArgVec av{"hello"}; - auto const result = getOptions(av.argc(), av.argv(), {}); - EXPECT_TRUE(result.errMsg.empty()); - ASSERT_EQ(result.positionalArgs.size(), 1U); - EXPECT_EQ(result.positionalArgs[0], "hello"sv); -} diff --git a/samples/common/logging.h b/samples/common/logging.h index cc3413e884..abaf7d132f 100644 --- a/samples/common/logging.h +++ b/samples/common/logging.h @@ -19,7 +19,6 @@ #define TENSORRT_LOGGING_H #include "NvInferRuntime.h" -#include "sampleOptions.h" #include #include #include @@ -35,6 +34,26 @@ namespace sample using Severity = nvinfer1::ILogger::Severity; +inline std::ostream& operator<<(std::ostream& os, nvinfer1::DataType dtype) +{ + switch (dtype) + { + case nvinfer1::DataType::kFLOAT: os << "fp32"; break; + case nvinfer1::DataType::kHALF: os << "fp16"; break; + case nvinfer1::DataType::kBF16: os << "bf16"; break; + case nvinfer1::DataType::kINT8: os << "int8"; break; + case nvinfer1::DataType::kINT32: os << "int32"; break; + case nvinfer1::DataType::kBOOL: os << "bool"; break; + case nvinfer1::DataType::kUINT8: os << "uint8"; break; + case nvinfer1::DataType::kFP8: os << "fp8"; break; + case nvinfer1::DataType::kINT64: os << "int64"; break; + case nvinfer1::DataType::kINT4: os << "int4"; break; + case nvinfer1::DataType::kFP4: os << "fp4"; break; + case nvinfer1::DataType::kE8M0: os << "e8m0"; break; + } + return os; +} + class LogStreamConsumerBuffer : public std::stringbuf { public: @@ -301,10 +320,7 @@ class Logger : public nvinfer1::ILogger kRUNNING, //!< The test is running kPASSED, //!< The test passed kFAILED, //!< The test failed - kWAIVED, //!< The test was waived - kTASK_BEGIN, //!< A sub-routine task has begun - kTASK_END, //!< A sub-routine task completed successfully - kTASK_ABORT //!< A sub-routine task was aborted (exception or validation failure) + kWAIVED //!< The test was waived }; //! @@ -352,11 +368,6 @@ class Logger : public nvinfer1::ILogger public: TestAtom(TestAtom&&) = default; - std::string getCmdline() const - { - return mCmdline; - } - private: friend class Logger; @@ -449,76 +460,6 @@ class Logger : public nvinfer1::ILogger return EXIT_FAILURE; } - static int32_t reportWaive(TestAtom const& testAtom) - { - reportTestEnd(testAtom, TestResult::kWAIVED); - return EXIT_SUCCESS; - } - - //! - //! \brief Report that a sub-routine task has begun. - //! - //! Used by the tuning loop to mark the start of each iteration so external - //! tooling can detect iteration boundaries in the trtexec log stream. - //! - static void reportTaskBegin(TestAtom const& testAtom) - { - reportTestResult(testAtom, TestResult::kTASK_BEGIN); - } - - //! - //! \brief Report that a sub-routine task has begun with iteration index and build route. - //! Prints a blank line before the banner for readability. - //! - //! Output example: - //! &&&& TASK_BEGIN [iter=0] BuildRoute = '-match_ragged_mha=on -copy_ppg=off' - //! - static void reportTaskBegin(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) - { - reportTaskWithBuildRoute( - TestResult::kTASK_BEGIN, index, buildRoute, /*blankBefore=*/true, /*blankAfter=*/false); - } - - //! - //! \brief Report that a sub-routine task completed successfully. - //! - static void reportTaskEnd(TestAtom const& testAtom) - { - reportTestResult(testAtom, TestResult::kTASK_END); - } - - //! - //! \brief Report that a sub-routine task completed successfully, with iteration info. - //! Prints a blank line after the banner for readability. - //! - static void reportTaskEnd(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) - { - reportTaskWithBuildRoute(TestResult::kTASK_END, index, buildRoute, /*blankBefore=*/false, /*blankAfter=*/true); - } - - //! - //! \brief Report that a sub-routine task was aborted (exception or validation failure). - //! - static void reportTaskAbort(TestAtom const& testAtom) - { - reportTestResult(testAtom, TestResult::kTASK_ABORT); - } - - //! - //! \brief Report that a sub-routine task was aborted, with iteration info. - //! Prints a blank line after the banner for readability. - //! - static void reportTaskAbort(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) - { - reportTaskWithBuildRoute( - TestResult::kTASK_ABORT, index, buildRoute, /*blankBefore=*/false, /*blankAfter=*/true); - } - - static int32_t reportTest(TestAtom const& testAtom, bool pass) - { - return pass ? reportPass(testAtom) : reportFail(testAtom); - } - Severity getReportableSeverity() const { return mReportableSeverity; @@ -552,35 +493,10 @@ class Logger : public nvinfer1::ILogger case TestResult::kPASSED: return "PASSED"; case TestResult::kFAILED: return "FAILED"; case TestResult::kWAIVED: return "WAIVED"; - case TestResult::kTASK_BEGIN: return "TASK_BEGIN"; - case TestResult::kTASK_END: return "TASK_END"; - case TestResult::kTASK_ABORT: return "TASK_ABORT"; default: assert(0); return ""; } } - //! - //! \brief Print a TASK_BEGIN/END/ABORT banner with iteration index and build route. - //! - //! Output format: - //! &&&& TASK_BEGIN [iter=0] BuildRoute = '-match_ragged_mha=on -copy_ppg=off' - //! - static void reportTaskWithBuildRoute( - TestResult result, std::string const& index, std::string const& buildRoute, bool blankBefore, bool blankAfter) - { - auto& os = severityOstream(Severity::kINFO); - if (blankBefore) - { - os << std::endl; - } - os << "&&&& " << testResultString(result) << " [iter=" << index << "] BuildRoute = '" << buildRoute << "'" - << std::endl; - if (blankAfter) - { - os << std::endl; - } - } - //! //! \brief returns an appropriate output stream (cout or cerr) to use with the given severity //! diff --git a/samples/common/safeCommon.h b/samples/common/safeCommon.h index 132eb1e37e..4e0027c57d 100644 --- a/samples/common/safeCommon.h +++ b/samples/common/safeCommon.h @@ -28,7 +28,6 @@ #include #include #include -#include #include #include #include @@ -569,20 +568,12 @@ inline bool parseBool(std::string const& arg, std::string const& name, std::opti return arg == "--" + name || (singleChar && arg == std::string{'-', *singleChar}); } -//! \brief Check whether \p internalOptions names \p name, with or without a value. -//! -//! Internal options reach TensorRT through the TRT_INTERNAL_OPTIONS environment variable, which samples -//! cannot query through the public API, so the string is parsed here instead. -//! -//! \return true when the option is present. -inline bool hasInternalOption(std::string const& internalOptions, std::string const& name) +inline bool hasCpuOnlyInternalOption(std::string const& internalOptions) { std::istringstream optionStream{internalOptions}; - auto const flag = "--" + name; - auto const assignment = flag + "="; for (std::string option; optionStream >> option;) { - if (option == flag || option.rfind(assignment, 0) == 0) + if (option == "--cpu_only" || option.rfind("--cpu_only=", 0) == 0) { return true; } @@ -590,40 +581,6 @@ inline bool hasInternalOption(std::string const& internalOptions, std::string co return false; } -inline bool hasCpuOnlyInternalOption(std::string const& internalOptions) -{ - return hasInternalOption(internalOptions, "cpu_only"); -} - -//! \brief Resolve the companion library to load beside a safe engine. -//! -//! An explicit path is passed through, so a caller who names a library that is not there gets an error -//! from the runtime instead of silence. The default beside the engine is offered only when that file -//! exists, so an engine built without a companion library still loads. Presence is the only question -//! asked of it: whether the file is loadable is dlopen's to answer, and it names the file when it is -//! not. An engine that needs a library and does not get one is caught by the runtime, which reads the -//! requirement out of the engine rather than from the caller. -//! -//! \param enginePath Path the engine was loaded from. -//! \param explicitPath Path given on the command line, or empty when none was. -//! -//! \return An absolute path to hand to createTRTGraph(), or std::nullopt when no companion library -//! applies. -//! -//! \throws std::filesystem::filesystem_error when the path cannot be inspected or made absolute, since -//! the safe runtime rejects a relative one and failing here names the path rather than leaving -//! graph creation to complain. A missing file is not an error and reports absent. -[[nodiscard]] inline std::optional resolveCompanionSoPath( - std::string const& enginePath, std::string const& explicitPath = "") -{ - std::filesystem::path const candidate{explicitPath.empty() ? enginePath + ".so" : explicitPath}; - if (explicitPath.empty() && !std::filesystem::is_regular_file(candidate)) - { - return std::nullopt; - } - return std::filesystem::absolute(candidate).lexically_normal().string(); -} - inline bool applyCpuOnlyMode() { #if !defined(_WIN32) diff --git a/samples/common/sampleConfig.h b/samples/common/sampleConfig.h deleted file mode 100644 index f5e421209a..0000000000 --- a/samples/common/sampleConfig.h +++ /dev/null @@ -1,291 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -#ifndef SampleConfig_H -#define SampleConfig_H - -#include -#include -#include - -#include "NvInfer.h" -#include "NvOnnxConfig.h" -class SampleConfig : public nvonnxparser::IOnnxConfig -{ -public: - enum class InputDataFormat : int - { - kASCII = 0, - kPPM = 1 - }; - -private: - std::string mModelFilename; - std::string mEngineFilename; - std::string mTextFilename; - std::string mFullTextFilename; - std::string mImageFilename; - std::string mReferenceFilename; - std::string mOutputFilename; - std::string mTimingCacheFilename; - int64_t mLabel{-1}; - int64_t mMaxBatchSize{32}; - int64_t mUseDLACore{-1}; - nvinfer1::DataType mModelDtype{nvinfer1::DataType::kFLOAT}; - bool mTF32{true}; - Verbosity mVerbosity{static_cast(nvinfer1::ILogger::Severity::kWARNING)}; - bool mPrintLayercInfo{false}; - bool mDebugBuilder{false}; - InputDataFormat mInputDataFormat{InputDataFormat::kASCII}; - uint64_t mTopK{0}; - float mFailurePercentage{-1.0F}; - float mTolerance{0.0F}; - float mAbsTolerance{1e-5F}; - -public: - SampleConfig() - { -#ifdef ONNX_DEBUG - if (isDebug()) - { - std::cout << " SampleConfig::ctor(): " << this << "\t" << std::endl; - } -#endif - } - - ~SampleConfig() override - { -#ifdef ONNX_DEBUG - if (isDebug()) - { - std::cout << "SampleConfig::dtor(): " << this << std::endl; - } -#endif - } - -public: - void setModelDtype(const nvinfer1::DataType mdt) noexcept override - { - mModelDtype = mdt; - } - - nvinfer1::DataType getModelDtype() const noexcept override - { - return mModelDtype; - } - - bool getTF32() const noexcept - { - return mTF32; - } - - void setTF32(bool enabled) noexcept - { - mTF32 = enabled; - } - - const char* getModelFileName() const noexcept override - { - return mModelFilename.c_str(); - } - - void setModelFileName(const char* onnxFilename) noexcept override - { - mModelFilename = std::string(onnxFilename); - } - Verbosity getVerbosityLevel() const noexcept override - { - return mVerbosity; - } - void addVerbosity() noexcept override - { - ++mVerbosity; - } - void reduceVerbosity() noexcept override - { - --mVerbosity; - } - void setVerbosityLevel(Verbosity v) noexcept override - { - mVerbosity = v; - } - const char* getEngineFileName() const noexcept - { - return mEngineFilename.c_str(); - } - void setEngineFileName(const char* engineFilename) noexcept - { - mEngineFilename = std::string(engineFilename); - } - const char* getTextFileName() const noexcept override - { - return mTextFilename.c_str(); - } - void setTextFileName(const char* textFilename) noexcept override - { - mTextFilename = std::string(textFilename); - } - const char* getFullTextFileName() const noexcept override - { - return mFullTextFilename.c_str(); - } - void setFullTextFileName(const char* fullTextFilename) noexcept override - { - mFullTextFilename = std::string(fullTextFilename); - } - void setLabel(int64_t label) noexcept - { - mLabel = label; - } //!< set the Label - - int64_t getLabel() const noexcept - { - return mLabel; - } //!< get the Label - - bool getPrintLayerInfo() const noexcept override - { - return mPrintLayercInfo; - } - - void setPrintLayerInfo(bool b) noexcept override - { - mPrintLayercInfo = b; - } //!< get the boolean variable corresponding to the Layer Info, see getPrintLayerInfo() - - void setMaxBatchSize(int64_t maxBatchSize) noexcept - { - mMaxBatchSize = maxBatchSize; - } //!< set the Max Batch Size - int64_t getMaxBatchSize() const noexcept - { - return mMaxBatchSize; - } //!< get the Max Batch Size - - void setUseDLACore(int64_t UseDLACore) noexcept - { - mUseDLACore = UseDLACore; - } //!< set the DLA core to use - int64_t getUseDLACore() const noexcept - { - return mUseDLACore; - } //!< get the DLA core to use - - void setDebugBuilder() noexcept - { - mDebugBuilder = true; - } //!< enable the Debug info, while building the engine. - bool getDebugBuilder() const noexcept - { - return mDebugBuilder; - } //!< get the boolean variable, corresponding to the debug builder - - const char* getImageFileName() const noexcept //!< set Image file name (PPM or ASCII) - { - return mImageFilename.c_str(); - } - void setImageFileName(const char* imageFilename) noexcept //!< get the Image file name - { - mImageFilename = std::string(imageFilename); - } - const char* getReferenceFileName() const noexcept - { - return mReferenceFilename.c_str(); - } - void setReferenceFileName(const char* referenceFilename) noexcept //!< set reference file name - { - mReferenceFilename = std::string(referenceFilename); - } - - void setInputDataFormat(InputDataFormat idt) noexcept - { - mInputDataFormat = idt; - } //!< specifies expected data format of the image file (PPM or ASCII) - InputDataFormat getInputDataFormat() const noexcept - { - return mInputDataFormat; - } //!< returns the expected data format of the image file. - - const char* getOutputFileName() const noexcept //!< specifies the file to save the results - { - return mOutputFilename.c_str(); - } - void setOutputFileName(const char* outputFilename) noexcept //!< get the output file name - { - mOutputFilename = std::string(outputFilename); - } - - uint64_t getTopK() const noexcept - { - return mTopK; - } - void setTopK(uint64_t topK) noexcept - { - mTopK = topK; - } //!< If this options is specified, return the K top probabilities. - - float getFailurePercentage() const noexcept - { - return mFailurePercentage; - } - - void setFailurePercentage(float f) noexcept - { - mFailurePercentage = f; - } - - float getAbsoluteTolerance() const noexcept - { - return mAbsTolerance; - } - - void setAbsoluteTolerance(float a) noexcept - { - mAbsTolerance = a; - } - - float getTolerance() const noexcept - { - return mTolerance; - } - - void setTolerance(float t) noexcept - { - mTolerance = t; - } - - const char* getTimingCacheFilename() const noexcept - { - return mTimingCacheFilename.c_str(); - } - - void setTimingCacheFileName(const char* timingCacheFilename) noexcept - { - mTimingCacheFilename = std::string(timingCacheFilename); - } - - bool isDebug() const noexcept - { -#if ONNX_DEBUG - return std::getenv("ONNX_DEBUG") != nullptr; -#else - return false; -#endif - } -}; // class SampleConfig - -#endif diff --git a/samples/common/sampleUtils.cpp b/samples/common/sampleUtils.cpp index 61782553c8..dd4738b9fc 100644 --- a/samples/common/sampleUtils.cpp +++ b/samples/common/sampleUtils.cpp @@ -16,745 +16,12 @@ */ #include "sampleUtils.h" -#include "bfloat16.h" -#include "common.h" -#include "half.h" -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#if CUDA_VERSION >= 11060 -#include -#endif - -using namespace nvinfer1; -using namespace std::string_view_literals; +#include "logger.h" namespace sample { -using TensorToLayer = std::unordered_map; -using LayerToTensor = std::unordered_map; -using TensorToTensor = std::unordered_map; - -int64_t volume(nvinfer1::Dims const& dims, nvinfer1::Dims const& strides, int32_t vecDim, int32_t comps, int32_t batch) -{ - int64_t maxNbElems = 1; - for (int32_t i = 0; i < dims.nbDims; ++i) - { - // Get effective length of axis. - int64_t d = dims.d[i]; - // Any dimension is 0, it is an empty tensor. - if (d == 0) - { - return 0; - } - if (i == vecDim) - { - d = samplesCommon::divUp(d, comps); - } - maxNbElems = std::max(maxNbElems, d * strides.d[i]); - } - return maxNbElems * batch * (vecDim < 0 ? 1 : comps); -} - -nvinfer1::Dims toDims(std::vector const& vec) -{ - int32_t limit = static_cast(nvinfer1::Dims::MAX_DIMS); - if (static_cast(vec.size()) > limit) - { - sample::gLogWarning << "Vector too long, only first 8 elements are used in dimension." << std::endl; - } - // Pick first nvinfer1::Dims::MAX_DIMS elements - nvinfer1::Dims dims{std::min(static_cast(vec.size()), limit), {}}; - std::copy_n(vec.begin(), dims.nbDims, std::begin(dims.d)); - return dims; -} - -void loadFromFile(std::string const& fileName, char* dst, size_t size) -{ - ASSERT(dst); - - std::ifstream file(fileName, std::ios::in | std::ios::binary); - if (file.is_open()) - { - file.seekg(0, std::ios::end); - int64_t fileSize = static_cast(file.tellg()); - // Due to change from int32_t to int64_t VC engines created with earlier versions - // may expect input of the half of the size - if (fileSize != static_cast(size) && fileSize != static_cast(size * 2)) - { - std::ostringstream msg; - msg << "Unexpected file size for input file: " << fileName << ". Note: Input binding size is: " << size - << " bytes but the file size is " << fileSize - << " bytes. Double check the size and datatype of the provided data."; - throw std::invalid_argument(msg.str()); - } - // Move file pointer back to the beginning after reading file size. - file.seekg(0, std::ios::beg); - file.read(dst, size); - size_t const nbBytesRead = file.gcount(); - file.close(); - if (nbBytesRead != size) - { - std::ostringstream msg; - msg << "Unexpected file size for input file: " << fileName << ". Note: Expected: " << size - << " bytes but only read: " << nbBytesRead << " bytes"; - throw std::invalid_argument(msg.str()); - } - } - else - { - std::ostringstream msg; - msg << "Cannot open file " << fileName << "!"; - throw std::invalid_argument(msg.str()); - } -} - -std::vector splitToStringVec(std::string const& s, char separator, int64_t maxSplit) -{ - std::vector splitted; - - for (size_t start = 0; start < s.length();) - { - // If maxSplit is specified and we have reached maxSplit, emplace back the rest of the string and break the - // loop. - if (maxSplit >= 0 && static_cast(splitted.size()) == maxSplit) - { - splitted.emplace_back(s.substr(start, s.length() - start)); - break; - } - - size_t separatorIndex = s.find(separator, start); - if (separatorIndex == std::string::npos) - { - separatorIndex = s.length(); - } - splitted.emplace_back(s.substr(start, separatorIndex - start)); - - // If the separator is the last character, then we should push an empty string at the end. - if (separatorIndex == s.length() - 1) - { - splitted.emplace_back(""); - } - - start = separatorIndex + 1; - } - - return splitted; -} - -bool broadcastIOFormats(std::vector const& formats, size_t nbBindings, bool isInput /*= true*/) -{ - bool broadcast = formats.size() == 1; - bool validFormatsCount = broadcast || (formats.size() == nbBindings); - if (!formats.empty() && !validFormatsCount) - { - if (isInput) - { - throw std::invalid_argument( - "The number of inputIOFormats must match network's inputs or be one for broadcasting."); - } - - throw std::invalid_argument( - "The number of outputIOFormats must match network's outputs or be one for broadcasting."); - } - return broadcast; -} - -// NOLINTNEXTLINE(readability-function-cognitive-complexity) -void sparsifyMatMulKernelWeights(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) -{ - // 1. Collect layers and tensors information from the network. - TensorToLayer matmulI2L; - TensorToLayer constO2L; - TensorToLayer shuffleI2L; - LayerToTensor shuffleL2O; - auto collectMappingInfo = [&](int32_t const idx) { - ILayer* l = network.getLayer(idx); - switch (l->getType()) - { - case nvinfer1::LayerType::kMATRIX_MULTIPLY: - { - // assume weights on the second input. - matmulI2L.insert({l->getInput(1), l}); - break; - } - case nvinfer1::LayerType::kCONSTANT: - { - DataType const dtype = static_cast(l)->getWeights().type; - if (dtype == nvinfer1::DataType::kFLOAT || dtype == nvinfer1::DataType::kHALF) - { - // Sparsify float only. - constO2L.insert({l->getOutput(0), l}); - } - break; - } - case nvinfer1::LayerType::kSHUFFLE: - { - shuffleI2L.insert({l->getInput(0), l}); - shuffleL2O.insert({l, l->getOutput(0)}); - break; - } - default: break; - } - }; - int32_t const nbLayers = network.getNbLayers(); - for (int32_t i = 0; i < nbLayers; ++i) - { - collectMappingInfo(i); - } - if (matmulI2L.size() == 0 || constO2L.size() == 0) - { - // No MatrixMultiply or Constant layer found, no weights to sparsify. - return; - } - - // Helper for analysis - auto isTranspose - = [](nvinfer1::Permutation const& perm) -> bool { return (perm.order[0] == 1 && perm.order[1] == 0); }; - auto is2D = [](nvinfer1::Dims const& dims) -> bool { return dims.nbDims == 2; }; - auto isIdenticalReshape = [](nvinfer1::Dims const& dims) -> bool { - for (int32_t i = 0; i < dims.nbDims; ++i) - { - if (dims.d[i] != i || dims.d[i] != -1) - { - return false; - } - } - return true; - }; - auto tensorReachedViaTranspose = [&](nvinfer1::ITensor* t, bool& needTranspose) -> ITensor* { - while (shuffleI2L.contains(t)) - { - nvinfer1::IShuffleLayer* s = static_cast(shuffleI2L.at(t)); - if (!is2D(s->getInput(0)->getDimensions()) || !is2D(s->getReshapeDimensions()) - || !isIdenticalReshape(s->getReshapeDimensions())) - { - break; - } - - if (isTranspose(s->getFirstTranspose())) - { - needTranspose = !needTranspose; - } - if (isTranspose(s->getSecondTranspose())) - { - needTranspose = !needTranspose; - } - - t = shuffleL2O.at(s); - } - return t; - }; - - // 2. Forward analysis to collect the Constant layers connected to MatMul via Transpose - std::unordered_map constantLayerToSparse; - for (auto& o2l : constO2L) - { - // If need to transpose the weights of the Constant layer. - // Need to transpose by default due to semantic difference. - bool needTranspose{true}; - ITensor* t = tensorReachedViaTranspose(o2l.first, needTranspose); - if (!matmulI2L.contains(t)) - { - continue; - } - - // check MatMul params... - IMatrixMultiplyLayer* mm = static_cast(matmulI2L.at(t)); - bool const twoInputs = mm->getNbInputs() == 2; - bool const all2D = is2D(mm->getInput(0)->getDimensions()) && is2D(mm->getInput(1)->getDimensions()); - bool const isSimple = mm->getOperation(0) == nvinfer1::MatrixOperation::kNONE - && mm->getOperation(1) != nvinfer1::MatrixOperation::kVECTOR; - if (!(twoInputs && all2D && isSimple)) - { - continue; - } - if (mm->getOperation(1) == nvinfer1::MatrixOperation::kTRANSPOSE) - { - needTranspose = !needTranspose; - } - - constantLayerToSparse.insert({static_cast(o2l.second), needTranspose}); - } - - // 3. Finally, sparsify the weights - auto sparsifyConstantWeights = [&sparseWeights](nvinfer1::IConstantLayer* layer, bool const needTranspose) { - Dims dims = layer->getOutput(0)->getDimensions(); - ASSERT(dims.nbDims == 2); - int32_t const idxN = needTranspose ? 1 : 0; - int32_t const n = dims.d[idxN]; - int32_t const k = dims.d[1 - idxN]; - sparseWeights.emplace_back(); - std::vector& spw = sparseWeights.back(); - Weights w = layer->getWeights(); - DataType const dtype = w.type; - ASSERT(dtype == nvinfer1::DataType::kFLOAT - || dtype == nvinfer1::DataType::kHALF); // non-float weights should have been ignored. - - if (needTranspose) - { - if (dtype == nvinfer1::DataType::kFLOAT) - { - spw.resize(w.count * sizeof(float)); - transpose2DWeights(spw.data(), w.values, k, n); - } - else if (dtype == nvinfer1::DataType::kHALF) - { - spw.resize(w.count * sizeof(half_float::half)); - transpose2DWeights(spw.data(), w.values, k, n); - } - - w.values = spw.data(); - std::vector tmpW; - sparsify(w, n, 1, tmpW); - - if (dtype == nvinfer1::DataType::kFLOAT) - { - transpose2DWeights(spw.data(), tmpW.data(), n, k); - } - else if (dtype == nvinfer1::DataType::kHALF) - { - transpose2DWeights(spw.data(), tmpW.data(), n, k); - } - } - else - { - sparsify(w, n, 1, spw); - } - - w.values = spw.data(); - layer->setWeights(w); - }; - for (auto& l : constantLayerToSparse) - { - sparsifyConstantWeights(l.first, l.second); - } -} - -template -void setSparseWeights(L& l, int32_t k, int32_t trs, std::vector& sparseWeights) -{ - auto weights = l.getKernelWeights(); - sparsify(weights, k, trs, sparseWeights); - weights.values = sparseWeights.data(); - l.setKernelWeights(weights); -} - -// Explicit instantiation -template void setSparseWeights( - IConvolutionLayer& l, int32_t k, int32_t trs, std::vector& sparseWeights); - -//! \brief Sparsify conv weights fed via Q/DQ chains (companion to sparsifyMatMulKernelWeights). -//! -//! Strongly-typed Q/DQ networks attach the conv weight as a tensor input rather than -//! static kernelWeights. Walks the chain forward from each FP Constant: -//! Constant -> Shuffle* -> Q? -> Shuffle* -> DQ -> Shuffle* -> Conv.input(1) -//! If the chain terminates at a Conv weight input, sparsify the constant in place. -// NOLINTNEXTLINE(readability-function-cognitive-complexity) -void sparsifyQDQConvKernelWeights( - nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) -{ - TensorToLayer convWeightI2L; - TensorToLayer constO2L; - TensorToTensor dqI2O; - TensorToTensor qI2O; - TensorToTensor shuffleI2O; - auto collectMappingInfo = [&](ILayer& l) { - switch (l.getType()) - { - case nvinfer1::LayerType::kCONVOLUTION: - // Conv with weights as a tensor input (vs. static kernelWeights). - if (l.getNbInputs() >= 2 && l.getInput(1) != nullptr) - { - convWeightI2L.try_emplace(l.getInput(1), &l); - } - break; - case nvinfer1::LayerType::kCONSTANT: - { - DataType const dtype = static_cast(l).getWeights().type; - auto const floatDTypes = {nvinfer1::DataType::kFLOAT, nvinfer1::DataType::kHALF, nvinfer1::DataType::kBF16}; - if (std::any_of(floatDTypes.begin(), floatDTypes.end(), [dtype](auto t) { return t == dtype; })) - { - constO2L.try_emplace(l.getOutput(0), &l); - } - break; - } - case nvinfer1::LayerType::kDEQUANTIZE: dqI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; - case nvinfer1::LayerType::kQUANTIZE: qI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; - case nvinfer1::LayerType::kSHUFFLE: shuffleI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; - default: break; - } - }; - int32_t const nbLayers = network.getNbLayers(); - for (int32_t i = 0; i < nbLayers; ++i) - { - collectMappingInfo(*network.getLayer(i)); - } - if (convWeightI2L.size() == 0 || constO2L.size() == 0 || dqI2O.size() == 0) - { - return; - } - - //! Skip past any Shuffle layers consuming t and return the tensor at the chain's end. - //! Returns t unchanged if no Shuffle reads it. - auto walkShuffleChain = [&](nvinfer1::ITensor* t) -> ITensor* { - while (true) - { - auto const it = shuffleI2O.find(t); - if (it == shuffleI2O.end()) - { - break; - } - t = it->second; - } - return t; - }; - - //! Follow Constant -> Shuffle* -> Q? -> Shuffle* -> DQ -> Shuffle* -> Conv.input(1) chain. - //! Returns the terminating IConvolutionLayer*, or nullptr if the chain breaks. - auto walkShuffleQDQChain = [&](nvinfer1::ITensor* t) -> IConvolutionLayer* { - t = walkShuffleChain(t); - if (auto const qI2OIt = qI2O.find(t); qI2OIt != qI2O.end()) - { - t = walkShuffleChain(qI2OIt->second); - } - auto const dqI2OIt = dqI2O.find(t); - if (dqI2OIt == dqI2O.end()) - { - return nullptr; - } - t = walkShuffleChain(dqI2OIt->second); - auto const convWeightI2LIt = convWeightI2L.find(t); - if (convWeightI2LIt == convWeightI2L.end()) - { - return nullptr; - } - ASSERT(convWeightI2LIt->second->getType() == nvinfer1::LayerType::kCONVOLUTION); - return static_cast(convWeightI2LIt->second); - }; - - for (auto& o2l : constO2L) - { - IConvolutionLayer* const conv = walkShuffleQDQChain(o2l.first); - if (conv == nullptr) - { - continue; - } - ASSERT(o2l.second->getType() == nvinfer1::LayerType::kCONSTANT); - IConstantLayer* constLayer = static_cast(o2l.second); - Weights w = constLayer->getWeights(); - if (w.count == 0) - { - continue; - } - Dims const kernelDims = conv->getKernelSizeNd(); - int32_t const k = conv->getNbOutputMaps(); - int64_t const trs = samplesCommon::volume(kernelDims); - // sparsify() reconstructs c (input channels) via c = count / (k*trs); fail loudly if - // the constant's element count doesn't match the KCRS layout this routine assumes. - ASSERT(k > 0 && 0 < trs && trs <= std::numeric_limits::max() - && w.count % (static_cast(k) * trs) == 0); - sparseWeights.emplace_back(); - sparsify(w, k, static_cast(trs), sparseWeights.back()); - w.values = sparseWeights.back().data(); - constLayer->setWeights(w); - } -} - -void sparsify(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) -{ - for (int32_t l = 0; l < network.getNbLayers(); ++l) - { - auto* layer = network.getLayer(l); - auto const t = layer->getType(); - if (t == nvinfer1::LayerType::kCONVOLUTION) - { - auto& conv = *static_cast(layer); - auto const& dims = conv.getKernelSizeNd(); - ASSERT(dims.nbDims == 2 || dims.nbDims == 3); - auto const k = conv.getNbOutputMaps(); - auto const trs = std::accumulate(dims.d, dims.d + dims.nbDims, 1, std::multiplies()); - sparseWeights.emplace_back(); - setSparseWeights(conv, k, trs, sparseWeights.back()); - } - } - - sparsifyMatMulKernelWeights(network, sparseWeights); - sparsifyQDQConvKernelWeights(network, sparseWeights); - sample::gLogVerbose << "--sparsity=force pruned " << sparseWeights.size() << " weights to be sparsity pattern." - << std::endl; - sample::gLogVerbose << "--sparsity=force has been deprecated. Please use to rewrite the " - "weights to a sparsity pattern and then run with --sparsity=enable" - << std::endl; -} - -void sparsify(Weights const& weights, int32_t k, int32_t trs, std::vector& sparseWeights) -{ - switch (weights.type) - { - case DataType::kFLOAT: - sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); - break; - case DataType::kHALF: - sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); - break; - case DataType::kBF16: - sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); - break; - case DataType::kINT8: - case DataType::kINT32: - case DataType::kUINT8: - case DataType::kBOOL: - case DataType::kINT4: - case DataType::kFP8: - case DataType::kINT64: - case DataType::kFP4: ASSERT(false && "Unsupported data type"); - case DataType::kE8M0: ASSERT(false && "E8M0 is not supported"); - } -} - -template -void print(std::ostream& os, T v) -{ - os << v; -} - -void print(std::ostream& os, int8_t v) -{ - os << static_cast(v); -} - -void print(std::ostream& os, uint8_t v) -{ - os << static_cast(v); -} - -void print(std::ostream& os, __half v) -{ - os << static_cast(v); -} - -#if CUDA_VERSION >= 11060 -void print(std::ostream& os, __nv_fp8_e4m3 v) -{ - os << static_cast(v); -} -#endif - -int32_t dataOffsetFromDims(int64_t v, Dims const& dims, Dims const& strides, int32_t vectorDim, int32_t spv) -{ - int32_t dataOffset = 0; - for (int32_t dimIndex = dims.nbDims - 1; dimIndex >= 0; --dimIndex) - { - int32_t dimVal = v % dims.d[dimIndex]; - if (dimIndex == vectorDim) - { - dataOffset += (dimVal / spv) * strides.d[dimIndex] * spv + dimVal % spv; - } - else - { - dataOffset += dimVal * strides.d[dimIndex] * (vectorDim == -1 ? 1 : spv); - } - v /= dims.d[dimIndex]; - ASSERT(v >= 0); - } - - return dataOffset; -} - -template -void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv) -{ - auto const vol = volume(dims); - T const* typedBuffer = static_cast(buffer); - for (int64_t v = 0; v < vol; ++v) - { - int32_t dataOffset = dataOffsetFromDims(v, dims, strides, vectorDim, spv); - if (v > 0) - { - os << separator; - } - print(os, typedBuffer[dataOffset]); - } -} - -void dumpInt4Buffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv) -{ - auto const vol = volume(dims); - uint8_t const* typedBuffer = static_cast(buffer); - for (int64_t v = 0; v < vol; ++v) - { - int32_t dataOffset = dataOffsetFromDims(v, dims, strides, vectorDim, spv); - if (v > 0) - { - os << separator; - } - - auto value = typedBuffer[dataOffset / 2]; - if (dataOffset % 2 == 0) - { - // Cast to int8_t before right shift, so right-shift will sign-extend. - // Left shift on int8_t can be undefined behaviour, must perform left shift on uint8_t. - os << (static_cast(value << 4) >> 4); - } - else - { - os << (static_cast(value) >> 4); - } - } -} - -// Explicit instantiation -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer<__half>(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -#if CUDA_VERSION >= 11060 -template void dumpBuffer<__nv_fp8_e4m3>(void const* buffer, std::string const& separator, std::ostream& os, - Dims const& dims, Dims const& strides, int32_t vectorDim, int32_t spv); -#endif -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); -template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); - -template -void sparsify(T const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights) -{ - auto const c = count / (k * trs); - sparseWeights.resize(count * sizeof(T)); - auto* sparseValues = reinterpret_cast(sparseWeights.data()); - - constexpr int32_t window = 4; - constexpr int32_t nonzeros = 2; - - int32_t const crs = c * trs; - auto const getIndex = [=](int32_t ki, int32_t ci, int32_t rsi) { return ki * crs + ci * trs + rsi; }; - - for (int64_t ki = 0; ki < k; ++ki) - { - for (int64_t rsi = 0; rsi < trs; ++rsi) - { - int32_t w = 0; - int32_t nz = 0; - for (int64_t ci = 0; ci < c; ++ci) - { - auto const index = getIndex(ki, ci, rsi); - if (nz < nonzeros) - { - sparseValues[index] = values[index]; - ++nz; - } - else - { - sparseValues[index] = 0; - } - if (++w == window) - { - w = 0; - nz = 0; - } - } - } - } -} - -// Explicit instantiation -template void sparsify( - float const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights); -template void sparsify( - half_float::half const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights); - -template -void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n) -{ - ASSERT(dst != src); - T* tdst = reinterpret_cast(dst); - T const* tsrc = reinterpret_cast(src); - for (int32_t mi = 0; mi < m; ++mi) - { - for (int32_t ni = 0; ni < n; ++ni) - { - int32_t const isrc = mi * n + ni; - int32_t const idst = ni * m + mi; - tdst[idst] = tsrc[isrc]; - } - } -} - -// Explicit instantiation -template void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); -template void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); - -template ::value, bool>::type> -void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max) -{ - T* typedBuffer = static_cast(buffer); - std::default_random_engine engine; - std::uniform_int_distribution distribution(min, max); - auto generator = [&engine, &distribution]() { return static_cast(distribution(engine)); }; - std::generate(typedBuffer, typedBuffer + volume, generator); -} - -template ::value, bool>::type> -void fillBuffer(void* buffer, int64_t volume, float min, float max) -{ - T* typedBuffer = static_cast(buffer); - std::default_random_engine engine; - std::uniform_real_distribution distribution(min, max); - auto generator = [&engine, &distribution]() { return static_cast(distribution(engine)); }; - std::generate(typedBuffer, typedBuffer + volume, generator); -} - -// Explicit instantiation -template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); -template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); -template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); -template void fillBuffer(void* buffer, int64_t volume, float min, float max); -template void fillBuffer<__half>(void* buffer, int64_t volume, float min, float max); -template void fillBuffer(void* buffer, int64_t volume, float min, float max); -#if CUDA_VERSION >= 11060 -template void fillBuffer<__nv_fp8_e4m3>(void* buffer, int64_t volume, float min, float max); -#endif -template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); -template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); - -bool matchStringWithOneWildcard(std::string const& pattern, std::string const& target) -{ - auto const splitPattern = splitToStringVec(pattern, '*', 1); - - // If there is no wildcard, return if the two strings match exactly. - if (splitPattern.size() == 1) - { - return pattern == target; - } - - // Otherwise, target must follow prefix+anything+postfix pattern. - return target.size() >= (splitPattern[0].size() + splitPattern[1].size()) && target.find(splitPattern[0]) == 0 - && target.rfind(splitPattern[1]) == (target.size() - splitPattern[1].size()); -} - -//! @brief Sanitizes the remote target config string by removing sensitive credentials +//! @brief Sanitizes the remote auto tuning config string by removing sensitive credentials //! //! This function removes usernames and passwords from URL-style configuration strings //! to prevent sensitive authentication information from appearing in logs or debug output. @@ -769,7 +36,7 @@ bool matchStringWithOneWildcard(std::string const& pattern, std::string const& t //! //! @param config The configuration string to sanitize //! @return Sanitized configuration string with passwords and usernames replaced by *** -std::string sanitizeRemoteConfig(std::string const& config) +std::string sanitizeRemoteAutoTuningConfig(std::string const& config) { if (config.empty()) { @@ -807,12 +74,12 @@ std::string sanitizeRemoteConfig(std::string const& config) } catch (std::exception const& e) { - sample::gLogError << "Exception in sanitizeRemoteConfig: " << e.what() << std::endl; + sample::gLogError << "Exception in sanitizeRemoteAutoTuningConfig: " << e.what() << std::endl; return config; // Return original on error } catch (...) { - sample::gLogError << "Unknown exception in sanitizeRemoteConfig" << std::endl; + sample::gLogError << "Unknown exception in sanitizeRemoteAutoTuningConfig" << std::endl; return config; // Return original on error } } @@ -827,11 +94,11 @@ bool validateNonEmpty(std::string const& value, std::string const& flagName) return true; } -bool validateRemoteConfig(std::string const& config) +bool validateRemoteAutoTuningConfig(std::string const& config) { if (config.find("://") == std::string::npos) { - sample::gLogError << "Invalid remote target config format. Expected format: " + sample::gLogError << "Invalid remote auto tuning config format. Expected format: " "protocol://username[:password]@hostname[:port]?param1=value1¶m2=value2" << std::endl; return false; @@ -841,10 +108,6 @@ bool validateRemoteConfig(std::string const& config) std::vector sanitizeArgv(int32_t argc, char** argv) { - // --remoteAutoTuningConfig is an alias of --remoteConfig; both carry credentials. - static constexpr std::array kREMOTE_CONFIG_FLAGS{ - "--remoteConfig=", "--remoteAutoTuningConfig="}; - std::vector sanitizedArgs; sanitizedArgs.reserve(argc); @@ -852,13 +115,11 @@ std::vector sanitizeArgv(int32_t argc, char** argv) { std::string arg = argv[i]; - for (auto const flag : kREMOTE_CONFIG_FLAGS) + // Sanitize remoteAutoTuningConfig argument + if (auto const flag = std::string("--remoteAutoTuningConfig="); + arg.size() > flag.size() && arg.substr(0, flag.size()) == flag) { - if (arg.size() > flag.size() && arg.compare(0, flag.size(), flag) == 0) - { - arg = std::string(flag) + sanitizeRemoteConfig(arg.substr(flag.size())); - break; - } + arg = std::string(flag) + sanitizeRemoteAutoTuningConfig(arg.substr(flag.size())); } sanitizedArgs.push_back(arg); @@ -867,549 +128,4 @@ std::vector sanitizeArgv(int32_t argc, char** argv) return sanitizedArgs; } -// ============================================================================ -// Accuracy Validator Implementations -// ============================================================================ - -template -double L0AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) -{ - // Uses PyTorch/NumPy allclose formula: |a - b| <= atol + rtol * |b| - // See: https://docs.pytorch.org/docs/stable/generated/torch.allclose.html - // and infer_ref_check/infer_ref_check.cpp::torchIsClose() - ASSERT(actual.size() == reference.size()); - ASSERT(actual.size() != 0); - int64_t mismatchCount = 0; - for (uint64_t i = 0; i < actual.size(); ++i) - { - double const absDiff = std::abs(static_cast(actual[i]) - static_cast(reference[i])); - double const refAbs = std::abs(static_cast(reference[i])); - double const tolerance = mAtol + mRtol * refAbs; - if (absDiff > tolerance) - { - mismatchCount++; - } - } - return static_cast(mismatchCount) / actual.size(); -} - -template -double L1AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) -{ - ASSERT(actual.size() == reference.size()); - ASSERT(actual.size() != 0); - double sum = 0.0; - for (uint64_t i = 0; i < actual.size(); ++i) - { - sum += std::abs(static_cast(actual[i]) - static_cast(reference[i])); - } - return sum / actual.size(); -} - -template -double L2AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) -{ - ASSERT(actual.size() == reference.size()); - ASSERT(actual.size() != 0); - double sum = 0.0; - for (uint64_t i = 0; i < actual.size(); ++i) - { - double diff = static_cast(actual[i]) - static_cast(reference[i]); - sum += diff * diff; - } - return sum / actual.size(); -} - -template -double LInfAccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) -{ - ASSERT(actual.size() == reference.size()); - ASSERT(actual.size() != 0); - double maxDiff = 0.0; - for (uint64_t i = 0; i < actual.size(); ++i) - { - double diff = std::abs(static_cast(actual[i]) - static_cast(reference[i])); - maxDiff = std::max(maxDiff, diff); - } - return maxDiff; -} - -template -double CosineSimilarityValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) -{ - ASSERT(actual.size() == reference.size()); - ASSERT(actual.size() != 0); - double dotProduct = 0.0; - double normActual = 0.0; - double normRef = 0.0; - for (uint64_t i = 0; i < actual.size(); ++i) - { - double a = static_cast(actual[i]); - double r = static_cast(reference[i]); - dotProduct += a * r; - normActual += a * a; - normRef += r * r; - } - double denominator = std::sqrt(normActual) * std::sqrt(normRef); - if (denominator < 1e-12) - { - return 1.0; // Handle zero vectors - } - double cosineSim = dotProduct / denominator; - return 1.0 - cosineSim; // Return as cost (0 = perfect match) -} - -// Explicit template instantiations for supported types -template class L0AccuracyValidator; -template class L0AccuracyValidator; -template class L0AccuracyValidator; -template class L0AccuracyValidator; - -template class L1AccuracyValidator; -template class L1AccuracyValidator; -template class L1AccuracyValidator; -template class L1AccuracyValidator; - -template class L2AccuracyValidator; -template class L2AccuracyValidator; -template class L2AccuracyValidator; -template class L2AccuracyValidator; - -template class LInfAccuracyValidator; -template class LInfAccuracyValidator; -template class LInfAccuracyValidator; -template class LInfAccuracyValidator; - -template class CosineSimilarityValidator; -template class CosineSimilarityValidator; -template class CosineSimilarityValidator; -template class CosineSimilarityValidator; - -bool peekArg(int32_t argc, char** argv, char const* flag) -{ - auto const flagLen = std::strlen(flag); - for (int32_t i = 1; i < argc; ++i) - { - if (argv[i] == nullptr) - { - continue; - } - // Match either bare flag (--continue) or flag=value (--tuneBuildRoutes=...). - if (std::strncmp(argv[i], flag, flagLen) == 0 && (argv[i][flagLen] == '\0' || argv[i][flagLen] == '=')) - { - return true; - } - } - return false; -} - -std::string buildShellQuotedCmdLine(int32_t argc, char** argv) -{ - std::string cmdLine; - for (int32_t i = 0; i < argc; ++i) - { - if (i > 0) - { - cmdLine += " "; - } - std::string arg = argv[i]; - bool const needsQuoting = arg.find_first_of(" \t|[]{}()&;'\"\\") != std::string::npos; - if (needsQuoting) - { - std::string escaped; - for (char c : arg) - { - if (c == '\'') - { - escaped += "'\\''"; - } - else - { - escaped += c; - } - } - cmdLine += "'" + escaped + "'"; - } - else - { - cmdLine += arg; - } - } - return cmdLine; -} - -//! \brief Resolve file paths in argv to absolute for cache storage. -//! -//! File-path flags that get resolved: --onnx=, --saveEngine=, --loadInputs=, -//! --loadRefOutputs=, --tuneBuildRouteFile=, --loadEngine=. All others are stored as-is. -//! --loadInputs and --loadRefOutputs have format "name:path,name:path" so each -//! path component is resolved separately. -namespace -{ -// NOLINTNEXTLINE(readability-function-cognitive-complexity) -std::vector resolveArgvPaths(int32_t argc, char** argv) -{ - static std::vector const kSIMPLE_PATH_FLAGS - = {"--onnx=", "--saveEngine=", "--tuneBuildRouteFile=", "--loadEngine=", "--loadCheckerBlob="}; - static std::vector const kMAPPED_PATH_FLAGS = {"--loadInputs=", "--loadRefOutputs="}; - - std::vector result; - for (int32_t i = 0; i < argc; ++i) - { - std::string arg(argv[i]); - - // Check simple path flags (--flag=path -> --flag=) - bool resolved = false; - for (auto const& prefix : kSIMPLE_PATH_FLAGS) - { - if (arg.starts_with(prefix)) - { - result.push_back(prefix + resolveAbsolutePath(arg.substr(prefix.size()))); - resolved = true; - break; - } - } - if (resolved) - { - continue; - } - - // Check mapped path flags (--flag=name:path,name:path -> resolve each path) - for (auto const& prefix : kMAPPED_PATH_FLAGS) - { - if (arg.starts_with(prefix)) - { - // Split on ',' to get individual name:path pairs - auto pairs = splitToStringVec(arg.substr(prefix.size()), ','); - std::string resolvedValue; - for (uint64_t p = 0; p < pairs.size(); ++p) - { - if (p > 0) - { - resolvedValue += ","; - } - // Split each pair on ':' to separate name from path - auto nameAndPath = splitToStringVec(pairs[p], ':', 1); - if (nameAndPath.size() == 2) - { - resolvedValue += nameAndPath[0] + ":" + resolveAbsolutePath(nameAndPath[1]); - } - else - { - resolvedValue += pairs[p]; // Malformed pair, keep as-is - } - } - result.push_back(prefix + resolvedValue); - resolved = true; - break; - } - } - if (resolved) - { - continue; - } - - result.push_back(arg); - } - return result; -} -} // anonymous namespace - -void writeTuningCacheHeader(std::string const& cacheFilePath, AllOptions const& options, int32_t argc, char** argv, - std::string const& tunerVersion, std::string const& defaultBuildRoute) -{ - // Use ordered_json to preserve insertion order matching best_config.json.example: - // tuner_version, accuracy_algorithm, accuracy_parameter, searching_algorithm, - // command_line, default_build_route, tuning_expr, files, argv - nlohmann::ordered_json header; - - header["tuner_version"] = tunerVersion; - header["accuracy_algorithm"] = getAlgorithmName(options.inference.accuracyValidationAlgorithm); - - nlohmann::ordered_json accParam; - accParam["atol"] = options.inference.atol; - accParam["rtol"] = options.inference.rtol; - accParam["epsilon"] = options.inference.accuracyThresholdEndToEnd; - header["accuracy_parameter"] = accParam; - - header["searching_algorithm"] = toString(options.tuning.tuningSearchAlgorithm); - - // Reconstruct command line for reference, with shell-safe quoting for arguments - // that contain spaces or metacharacters (e.g. --tuneBuildRoutes values). - std::string cmdLine = buildShellQuotedCmdLine(argc, argv); - header["command_line"] = cmdLine; - header["default_build_route"] = defaultBuildRoute; - - // Store the expanded tuning expression. This is the already-expanded string - // (handles --tuneBuildRouteFile case where the file may not exist at resume time). - header["tuning_expr"] = options.tuning.tuningExpr; - - // Store absolute paths to all file-based options for human readability and - // as a cross-check. The authoritative source for --continue reconstruction - // is the "argv" field below. - { - nlohmann::ordered_json files; - if (!options.model.baseModel.model.empty()) - { - files["onnx"] = resolveAbsolutePath(options.model.baseModel.model); - } - if (!options.build.engine.empty()) - { - files["save_engine"] = resolveAbsolutePath(options.build.engine); - } - // Input files: map of tensor_name → absolute path - if (!options.inference.refPairs.empty()) - { - nlohmann::ordered_json inputs; - for (auto const& [name, path] : options.inference.refPairs[0].first) - { - inputs[name] = resolveAbsolutePath(path); - } - if (!inputs.empty()) - { - files["inputs"] = inputs; - } - - nlohmann::ordered_json refOutputs; - for (auto const& [name, path] : options.inference.refPairs[0].second) - { - refOutputs[name] = resolveAbsolutePath(path); - } - if (!refOutputs.empty()) - { - files["ref_outputs"] = refOutputs; - } - } - header["files"] = files; - } - - // Store argv with file-path arguments resolved to absolute paths. - // This is the machine-readable source of truth for --continue reconstruction. - // When resuming, the stored argv is replayed to reconstruct all options - // (--iterations, --duration, --fp16, etc.) without enumerating each one. - { - auto resolvedArgv = resolveArgvPaths(argc, argv); - nlohmann::ordered_json argvArray(resolvedArgv); - header["argv"] = argvArray; - } - - std::ofstream file(cacheFilePath, std::ios::trunc); - if (!file) - { - sample::gLogError << "Cannot open tuning cache file for writing header: " << cacheFilePath << std::endl; - return; - } - file << header.dump() << std::endl; -} - -void writeTuningCacheIteration(std::string const& cacheFilePath, uint64_t iter, std::string const& buildRoute, - bool crashed, std::string const& errorMessage, std::unordered_map const& accuracyLossValues, - double gpuTimeMs) -{ - // Use ordered_json to preserve insertion order matching best_config.json.example: - // iter, build_route, crash, error_message, accuracy_loss, gpu_time - nlohmann::ordered_json result; - result[tuningCache::kIter] = iter; - result[tuningCache::kBuildRoute] = buildRoute; - result[tuningCache::kCrash] = crashed; - result[tuningCache::kErrorMessage] = errorMessage; - - // accuracy_loss is a per-output map: {"output_name": accuracy_value, ...} - // When crashed, accuracy values are unavailable so we write null. - if (crashed || accuracyLossValues.empty()) - { - result[tuningCache::kAccuracyLoss] = nullptr; - } - else - { - nlohmann::ordered_json accMap; - for (auto const& [name, value] : accuracyLossValues) - { - accMap[name] = value; - } - result[tuningCache::kAccuracyLoss] = accMap; - } - result[tuningCache::kGpuTime] = crashed ? nlohmann::ordered_json(nullptr) : nlohmann::ordered_json(gpuTimeMs); - - std::ofstream file(cacheFilePath, std::ios::app); - if (!file) - { - sample::gLogError << "Cannot open tuning cache file to append iteration " << iter << ": " << cacheFilePath - << std::endl; - return; - } - file << result.dump() << std::endl; -} - -std::vector reconstructArgvFromCacheHeader( - TuningCacheHeader const& header, std::string const& currentExePath, std::string const& cacheFilePath) -{ - std::vector newArgv; - - // Use current executable path as argv[0], not the one stored in the cache - // (the binary may have been rebuilt or moved since the original run). - newArgv.push_back(currentExePath); - - // Iterate over stored argv (skip stored argv[0]). - for (uint64_t i = 1; i < header.argv.size(); ++i) - { - std::string const& arg = header.argv[i]; - - // Replace --tuneBuildRoutes or --tuneBuildRouteFile with the stored tuning_expr. - // This handles the case where --tuneBuildRouteFile was used originally but the - // file no longer exists — the expanded expression is stored in tuning_expr. - if (arg.starts_with("--tuneBuildRoutes=") || arg.starts_with("--tuneBuildRouteFile=")) - { - continue; // Will be re-added below with the stored tuning_expr. - } - - // Remove --continue and --tuningCacheFile from the stored argv to avoid - // recursion (the stored run may itself have been a --continue run). - if (arg == "--continue"sv || arg.starts_with("--tuningCacheFile=")) - { - continue; - } - - newArgv.push_back(arg); - } - - // Add back the tuning expression and cache file path. - newArgv.push_back("--tuneBuildRoutes=" + header.tuningExpr); - newArgv.push_back("--tuningCacheFile=" + cacheFilePath); - - return newArgv; -} - -std::string resolveAbsolutePath(std::string const& path) -{ - if (path.empty()) - { - return path; - } -#if defined(_WIN32) - // On Windows, path resolution is not needed (tuning features are not supported on Windows). - // Return the path unchanged so the code compiles. - return path; -#else - // POSIX realpath() resolves symlinks and relative components to an absolute path. - // Returns nullptr if the file does not exist or another error occurs. - char resolved[PATH_MAX]; - if (realpath(path.c_str(), resolved) != nullptr) - { - return std::string(resolved); - } - return path; -#endif -} - -std::optional readTuningCacheHeader(std::string const& cacheFilePath) -{ - std::ifstream file(cacheFilePath); - if (!file.is_open()) - { - return std::nullopt; - } - - // First line is the JSON header. - std::string headerLine; - if (!std::getline(file, headerLine) || headerLine.empty()) - { - return std::nullopt; - } - - try - { - auto headerJson = nlohmann::json::parse(headerLine); - - TuningCacheHeader header; - - // Extract argv array → vector - if (headerJson.contains("argv") && headerJson["argv"].is_array()) - { - for (auto const& elem : headerJson["argv"]) - { - header.argv.push_back(elem.get()); - } - } - else - { - // argv field is required for --continue reconstruction. - sample::gLogError << "Tuning cache header missing 'argv' field" << std::endl; - return std::nullopt; - } - - // Extract tuning_expr string. - if (headerJson.contains("tuning_expr") && headerJson["tuning_expr"].is_string()) - { - header.tuningExpr = headerJson["tuning_expr"].get(); - } - else - { - sample::gLogError << "Tuning cache header missing 'tuning_expr' field" << std::endl; - return std::nullopt; - } - - // Count remaining non-empty lines as completed iterations. - header.completedIterations = 0; - std::string line; - while (std::getline(file, line)) - { - if (!line.empty()) - { - ++header.completedIterations; - } - } - - return header; - } - catch (nlohmann::json::exception const& e) - { - sample::gLogError << "Failed to parse tuning cache header: " << e.what() << std::endl; - return std::nullopt; - } -} - -std::vector readCachedIterationResults(std::string const& cacheFilePath, int64_t maxIterations) -{ - std::vector results; - std::ifstream file(cacheFilePath); - if (!file.is_open()) - { - return results; - } - - std::string line; - // Skip header line. - if (!std::getline(file, line)) - { - return results; - } - - // Read iteration lines, extracting crash and gpu_time fields. - while (std::getline(file, line) && static_cast(results.size()) < maxIterations) - { - if (line.empty()) - { - continue; - } - try - { - auto j = nlohmann::json::parse(line); - CachedIterationResult r; - r.crashed = j.value(tuningCache::kCrash, true); - r.gpuTimeMs = j.contains(tuningCache::kGpuTime) && j[tuningCache::kGpuTime].is_number() - ? j[tuningCache::kGpuTime].get() - : 0.0; - results.push_back(r); - } - catch (nlohmann::json::exception const&) - { - // Malformed line — treat as crashed. - results.push_back({true, 0.0}); - } - } - - return results; -} - } // namespace sample diff --git a/samples/common/sampleUtils.h b/samples/common/sampleUtils.h index fe18ac722a..5cd0549c34 100644 --- a/samples/common/sampleUtils.h +++ b/samples/common/sampleUtils.h @@ -18,124 +18,20 @@ #ifndef TRT_SAMPLE_UTILS_H #define TRT_SAMPLE_UTILS_H -#include -#include -#include -#include -#include -#include +#include #include -#include #include -#include -#include -#include - -#include "NvInfer.h" - -#include "common.h" -#include "logger.h" -#include "logging.h" -#include "sampleOptions.h" - -#define SMP_RETVAL_IF_FALSE(condition, msg, retval, err) \ - { \ - if ((condition) == false) \ - { \ - (err) << (msg) << std::endl; \ - return retval; \ - } \ - } - namespace sample { -template -inline T roundUp(T m, T n) -{ - return ((m + n - 1) / n) * n; -} - -//! comps is the number of components in a vector. Ignored if vecDim < 0. -int64_t volume(nvinfer1::Dims const& dims, nvinfer1::Dims const& strides, int32_t vecDim, int32_t comps, int32_t batch); - -using samplesCommon::volume; - -nvinfer1::Dims toDims(std::vector const& vec); - -template ::value, bool>::type = true> -void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); - -template ::value, bool>::type = true> -void fillBuffer(void* buffer, int64_t volume, float min, float max); - -template -void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, nvinfer1::Dims const& dims, - nvinfer1::Dims const& strides, int32_t vectorDim, int32_t spv); - -void dumpInt4Buffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, - Dims const& strides, int32_t vectorDim, int32_t spv); - -void loadFromFile(std::string const& fileName, char* dst, size_t size); - -std::vector splitToStringVec(std::string const& option, char separator, int64_t maxSplit = -1); - -bool broadcastIOFormats(std::vector const& formats, size_t nbBindings, bool isInput = true); - -#if !TRT_WINML -int32_t getCudaDriverVersion(); - -int32_t getCudaRuntimeVersion(); -#endif - -void sparsify(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights); -void sparsify(nvinfer1::Weights const& weights, int32_t k, int32_t rs, std::vector& sparseWeights); - -// Walk the weights elements and overwrite (at most) 2 out of 4 elements to 0. -template -void sparsify(T const* values, int64_t count, int32_t k, int32_t rs, std::vector& sparseWeights); - -template -void setSparseWeights(L& l, int32_t k, int32_t rs, std::vector& sparseWeights); - -// Sparsify the weights of Constant layers that are fed to MatMul via Shuffle layers. -// Forward analysis on the API graph to determine which weights to sparsify. -void sparsifyMatMulKernelWeights( - nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights); - -template -void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); - -//! A helper function to match a target string with a pattern where the pattern can contain up to one wildcard ('*') -//! character that matches to any strings. -bool matchStringWithOneWildcard(std::string const& pattern, std::string const& target); - -//! A helper method to find an item from an unordered_map. If the exact match exists, this is identical to -//! map.find(target). If the exact match does not exist, it returns the first plausible match, taking up to one wildcard -//! into account. If there is no plausible match, then it returns map.end(). -template -typename std::unordered_map::const_iterator findPlausible( - std::unordered_map const& map, std::string const& target) -{ - auto res = map.find(target); - if (res == map.end()) - { - res = std::find_if( - map.begin(), map.end(), [&](typename std::unordered_map::value_type const& item) { - return matchStringWithOneWildcard(item.first, target); - }); - } - return res; -} - // ==== Common argument parsing utilities ==== //! Validate that a value is not empty, log error if it is bool validateNonEmpty(std::string const& value, std::string const& flagName); -//! Validate remote target config format -bool validateRemoteConfig(std::string const& config); +//! Validate remote auto tuning config format +bool validateRemoteAutoTuningConfig(std::string const& config); //! Ensure directory path ends with '/' inline std::string normalizeDirectoryPath(std::string const& dirPath) @@ -148,193 +44,17 @@ inline std::string normalizeDirectoryPath(std::string const& dirPath) return result; } -//! Sanitizes the remote target config string by removing sensitive credentials +//! Sanitizes the remote auto tuning config string by removing sensitive credentials //! Removes usernames and passwords from URL-style config strings for security. //! Example: "ssh://user:pass@host:22" becomes "ssh://***:***@host:22" -std::string sanitizeRemoteConfig(std::string const& config); +std::string sanitizeRemoteAutoTuningConfig(std::string const& config); //! Sanitizes command line arguments for logging, removing sensitive credentials -//! Processes argv array and sanitizes sensitive arguments like remoteConfig +//! Processes argv array and sanitizes sensitive arguments like remoteAutoTuningConfig //! @param argc Number of arguments //! @param argv Array of argument strings //! @return Vector of sanitized argument strings std::vector sanitizeArgv(int32_t argc, char** argv); -//! Interface for accuracy validation -//! This interface provides a way to calculate the accuracy gap between the actual and reference outputs. -//! Since all the return value is a "loss value", the lower the return value, the better accuracy it is. -template -class IAccuracyValidator -{ -public: - virtual ~IAccuracyValidator() = default; - virtual double calculateAccuracy(std::vector const& actual, std::vector const& reference) = 0; -}; - -//! L0 accuracy validator calculates element-wise accuracy using the PyTorch/NumPy allclose formula. -//! An element matches if: |actual[i] - ref[i]| <= atol + rtol * |ref[i]| -//! accuracy = (number of mismatching elements) / N -//! Returns the mismatch ratio (0.0 means perfect match, 1.0 means all elements mismatch). -template -class L0AccuracyValidator : public IAccuracyValidator -{ -public: - L0AccuracyValidator(double atol, double rtol) - : mAtol(atol) - , mRtol(rtol) - { - } - - double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; - -private: - double mAtol; - double mRtol; -}; - -//! L1 accuracy validator calculates mean absolute error. -//! accuracy = Sum(|actual[i] - ref[i]|) / N -//! Returns the mean absolute error (0.0 means perfect match). -template -class L1AccuracyValidator : public IAccuracyValidator -{ -public: - double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; -}; - -//! L2 accuracy validator calculates mean squared error. -//! accuracy = Sum(|actual[i] - ref[i]|^2) / N -//! Returns the mean squared error (0.0 means perfect match). -template -class L2AccuracyValidator : public IAccuracyValidator -{ -public: - double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; -}; - -//! LInf accuracy validator calculates maximum absolute error. -//! accuracy = Max(|actual[i] - ref[i]|) -//! Returns the max absolute error (0.0 means perfect match). -template -class LInfAccuracyValidator : public IAccuracyValidator -{ -public: - double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; -}; - -//! Cosine similarity validator calculates 1 - cosine_similarity. -//! cosine_sim = Sum(actual[i] * ref[i]) / (sqrt(Sum(actual[i]^2)) * sqrt(Sum(ref[i]^2))) -//! accuracy loss = 1 - cosine_sim -//! Returns 1 - cosine_similarity (0.0 means perfect match). -template -class CosineSimilarityValidator : public IAccuracyValidator -{ -public: - double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; -}; - -//! \brief Get human-readable name string for an accuracy validation algorithm. -//! \param[in] algorithm The accuracy validation algorithm enum value. -//! \return Name string (e.g., "L0", "L1", "Cosine"). -inline std::string getAlgorithmName(AccuracyValidationAlgorithm algorithm) -{ - switch (algorithm) - { - case AccuracyValidationAlgorithm::kL0: return "L0"; - case AccuracyValidationAlgorithm::kL1: return "L1"; - case AccuracyValidationAlgorithm::kL2: return "L2"; - case AccuracyValidationAlgorithm::kLInf: return "LInf"; - case AccuracyValidationAlgorithm::kCosineSimilarity: return "Cosine"; - default: return "Unknown"; - } -} - -//! \brief Factory function to create an accuracy validator based on algorithm type. -//! \param[in] algorithm The accuracy validation algorithm to use. -//! \param[in] atol Absolute tolerance (only used by L0 algorithm). -//! \param[in] rtol Relative tolerance (only used by L0 algorithm). -//! \return Unique pointer to the appropriate IAccuracyValidator implementation. -template -std::unique_ptr> createAccuracyValidator( - AccuracyValidationAlgorithm algorithm, float atol = 1e-5F, float rtol = 1e-5F) -{ - switch (algorithm) - { - case AccuracyValidationAlgorithm::kL0: return std::make_unique>(atol, rtol); - case AccuracyValidationAlgorithm::kL1: return std::make_unique>(); - case AccuracyValidationAlgorithm::kL2: return std::make_unique>(); - case AccuracyValidationAlgorithm::kLInf: return std::make_unique>(); - case AccuracyValidationAlgorithm::kCosineSimilarity: return std::make_unique>(); - } - ASSERT(false && "Unknown Accuracy Validation Algorithm"); - return nullptr; -} - -//! \brief Cheap argv pre-scan. Returns true if some `argv[i]` exactly equals `flag` -//! or starts with `flag` + "=". Used by main() before option parsing to dispatch -//! between trtexec single-run mode and the tuning loop. -[[nodiscard]] bool peekArg(int32_t argc, char** argv, char const* flag); - -// ============================================================================ -// Tuning cache I/O (used by --tuneBuildRoutes / --continue). -// Header is a single JSON object on line 1; iterations are JSON Lines after. -// ============================================================================ - -//! \brief Reconstruct a shell-safe command line string from argc/argv. -std::string buildShellQuotedCmdLine(int32_t argc, char** argv); - -//! \brief Resolve a file path to an absolute path using POSIX realpath(). -//! Empty input or realpath() failure returns the input unchanged. -std::string resolveAbsolutePath(std::string const& path); - -//! \brief Write the tuning cache file header (line 1, JSON object). -void writeTuningCacheHeader(std::string const& cacheFilePath, AllOptions const& options, int32_t argc, char** argv, - std::string const& tunerVersion, std::string const& defaultBuildRoute); - -//! \brief Append one iteration line to the cache file. Fields: iter, build_route, crash, -//! error_message, accuracy_loss, gpu_time. Crashed iterations have null accuracy/gpu. -void writeTuningCacheIteration(std::string const& cacheFilePath, uint64_t iter, std::string const& buildRoute, - bool crashed, std::string const& errorMessage, std::unordered_map const& accuracyLossValues, - double gpuTimeMs); - -//! \struct TuningCacheHeader -//! \brief Parsed contents of the cache header, returned by readTuningCacheHeader(). -struct TuningCacheHeader -{ - std::vector argv; //!< Original command line with file paths absolute. - std::string tuningExpr; //!< Expanded --tuneBuildRoutes expression. - int64_t completedIterations{0}; //!< Number of iteration lines after the header. -}; - -//! \brief Read and parse the cache file's header line + count completed iteration lines. -std::optional readTuningCacheHeader(std::string const& cacheFilePath); - -//! \brief Rebuild argv for a --continue resume. argv[0] is replaced with currentExePath; -//! --tuneBuildRoutes is set to the cached expanded expression; --continue and -//! --tuningCacheFile are stripped from the stored argv and the cache path is re-appended. -std::vector reconstructArgvFromCacheHeader( - TuningCacheHeader const& header, std::string const& currentExePath, std::string const& cacheFilePath); - -//! \struct CachedIterationResult -//! \brief Minimal per-iteration fields from the cache, used to reconstruct mixed-mode positive knobs. -struct CachedIterationResult -{ - bool crashed{true}; - double gpuTimeMs{0.0}; -}; - -//! \brief Read up to maxIterations iteration lines from the cache and extract (crashed, gpu_time). -std::vector readCachedIterationResults(std::string const& cacheFilePath, int64_t maxIterations); - -namespace tuningCache -{ -constexpr char const* kIter = "iter"; -constexpr char const* kBuildRoute = "build_route"; -constexpr char const* kCrash = "crash"; -constexpr char const* kErrorMessage = "error_message"; -constexpr char const* kAccuracyLoss = "accuracy_loss"; -constexpr char const* kGpuTime = "gpu_time"; -} // namespace tuningCache - } // namespace sample #endif // TRT_SAMPLE_UTILS_H diff --git a/samples/common/sampleUtils.test.cpp b/samples/common/sampleUtils.test.cpp index 709fd8f2b5..7632ff3f32 100644 --- a/samples/common/sampleUtils.test.cpp +++ b/samples/common/sampleUtils.test.cpp @@ -16,7 +16,6 @@ */ #include "sampleUtils.h" -#include "ArgVec.test.h" #include @@ -25,94 +24,6 @@ using namespace sample; using namespace std::string_view_literals; -TEST(RoundUp, ExactMultiple) -{ - EXPECT_EQ(roundUp(4, 4), 4); - EXPECT_EQ(roundUp(8, 4), 8); - EXPECT_EQ(roundUp(0, 4), 0); -} - -TEST(RoundUp, NeedsRounding) -{ - EXPECT_EQ(roundUp(1, 4), 4); - EXPECT_EQ(roundUp(5, 4), 8); - EXPECT_EQ(roundUp(7, 4), 8); -} - -TEST(SplitToStringVec, SingleToken) -{ - auto const v = splitToStringVec("hello", ','); - ASSERT_EQ(v.size(), 1U); - EXPECT_EQ(v[0], "hello"sv); -} - -TEST(SplitToStringVec, MultipleTokens) -{ - auto const v = splitToStringVec("a,b,c", ','); - ASSERT_EQ(v.size(), 3U); - EXPECT_EQ(v[0], "a"sv); - EXPECT_EQ(v[1], "b"sv); - EXPECT_EQ(v[2], "c"sv); -} - -TEST(SplitToStringVec, EmptyString) -{ - auto const v = splitToStringVec("", ','); - EXPECT_TRUE(v.empty()); -} - -TEST(SplitToStringVec, MaxSplit) -{ - // maxSplit=1 means at most one split; the rest of the string is the second element. - auto const v = splitToStringVec("a:b:c", ':', 1); - ASSERT_EQ(v.size(), 2U); - EXPECT_EQ(v[0], "a"sv); - EXPECT_EQ(v[1], "b:c"sv); -} - -TEST(SplitToStringVec, TrailingSeparator) -{ - auto const v = splitToStringVec("a,b,", ','); - ASSERT_EQ(v.size(), 3U); - EXPECT_EQ(v[0], "a"sv); - EXPECT_EQ(v[1], "b"sv); - EXPECT_EQ(v[2], ""sv); -} - -TEST(MatchStringWithOneWildcard, ExactMatch) -{ - EXPECT_TRUE(matchStringWithOneWildcard("hello", "hello")); - EXPECT_FALSE(matchStringWithOneWildcard("hello", "world")); - EXPECT_FALSE(matchStringWithOneWildcard("hello", "hello2")); -} - -TEST(MatchStringWithOneWildcard, WildcardMatchesAnything) -{ - EXPECT_TRUE(matchStringWithOneWildcard("*", "anything")); - EXPECT_TRUE(matchStringWithOneWildcard("*", "")); -} - -TEST(MatchStringWithOneWildcard, PrefixWildcard) -{ - EXPECT_TRUE(matchStringWithOneWildcard("hello*", "hello")); - EXPECT_TRUE(matchStringWithOneWildcard("hello*", "hello world")); - EXPECT_FALSE(matchStringWithOneWildcard("hello*", "world")); -} - -TEST(MatchStringWithOneWildcard, SuffixWildcard) -{ - EXPECT_TRUE(matchStringWithOneWildcard("*world", "world")); - EXPECT_TRUE(matchStringWithOneWildcard("*world", "hello world")); - EXPECT_FALSE(matchStringWithOneWildcard("*world", "hello")); -} - -TEST(MatchStringWithOneWildcard, MiddleWildcard) -{ - EXPECT_TRUE(matchStringWithOneWildcard("he*ld", "held")); - EXPECT_TRUE(matchStringWithOneWildcard("he*ld", "hello world")); - EXPECT_FALSE(matchStringWithOneWildcard("he*ld", "hello")); -} - TEST(NormalizeDirectoryPath, AlreadyNormalized) { EXPECT_EQ(normalizeDirectoryPath("/some/path/"), "/some/path/"sv); @@ -128,55 +39,29 @@ TEST(NormalizeDirectoryPath, EmptyString) EXPECT_EQ(normalizeDirectoryPath(""), ""sv); } -TEST(SanitizeRemoteConfig, Empty) +TEST(SanitizeRemoteAutoTuningConfig, Empty) { - EXPECT_EQ(sanitizeRemoteConfig(""), ""sv); + EXPECT_EQ(sanitizeRemoteAutoTuningConfig(""), ""sv); } -TEST(SanitizeRemoteConfig, NoCredentials) +TEST(SanitizeRemoteAutoTuningConfig, NoCredentials) { // No @ means no credentials section; returned as-is. - EXPECT_EQ(sanitizeRemoteConfig("ssh://host:22"), "ssh://host:22"sv); -} - -TEST(SanitizeRemoteConfig, UsernameOnly) -{ - EXPECT_EQ(sanitizeRemoteConfig("ssh://user@host:22"), "ssh://***@host:22"sv); -} - -TEST(SanitizeRemoteConfig, UsernameAndPassword) -{ - EXPECT_EQ(sanitizeRemoteConfig("ssh://user:pass@host:22"), "ssh://***@host:22"sv); -} - -TEST(SanitizeRemoteConfig, WithQueryParams) -{ - EXPECT_EQ( - sanitizeRemoteConfig("ssh://admin:secret@server.com:22?timeout=30"), "ssh://***@server.com:22?timeout=30"sv); + EXPECT_EQ(sanitizeRemoteAutoTuningConfig("ssh://host:22"), "ssh://host:22"sv); } -TEST(SanitizeArgv, MasksRemoteConfigCredentials) +TEST(SanitizeRemoteAutoTuningConfig, UsernameOnly) { - ArgVec av{"--remoteConfig=ssh://user:pass@host:22", "--safe"}; - auto const sanitized = sanitizeArgv(av.argc(), av.argv()); - ASSERT_EQ(sanitized.size(), 3U); - EXPECT_EQ(sanitized[1], "--remoteConfig=ssh://***@host:22"sv); - EXPECT_EQ(sanitized[2], "--safe"sv); + EXPECT_EQ(sanitizeRemoteAutoTuningConfig("ssh://user@host:22"), "ssh://***@host:22"sv); } -TEST(SanitizeArgv, MasksAliasCredentials) +TEST(SanitizeRemoteAutoTuningConfig, UsernameAndPassword) { - ArgVec av{"--remoteAutoTuningConfig=ssh://user:pass@host:22"}; - auto const sanitized = sanitizeArgv(av.argc(), av.argv()); - ASSERT_EQ(sanitized.size(), 2U); - EXPECT_EQ(sanitized[1], "--remoteAutoTuningConfig=ssh://***@host:22"sv); + EXPECT_EQ(sanitizeRemoteAutoTuningConfig("ssh://user:pass@host:22"), "ssh://***@host:22"sv); } -TEST(SanitizeArgv, LeavesOtherArgumentsUntouched) +TEST(SanitizeRemoteAutoTuningConfig, WithQueryParams) { - ArgVec av{"--onnx=model.onnx", "--remoteConfig="}; - auto const sanitized = sanitizeArgv(av.argc(), av.argv()); - ASSERT_EQ(sanitized.size(), 3U); - EXPECT_EQ(sanitized[1], "--onnx=model.onnx"sv); - EXPECT_EQ(sanitized[2], "--remoteConfig="sv); + EXPECT_EQ(sanitizeRemoteAutoTuningConfig("ssh://admin:secret@server.com:22?timeout=30"), + "ssh://***@server.com:22?timeout=30"sv); } diff --git a/samples/sampleDistCollective/sampleDistCollective.cpp b/samples/sampleDistCollective/sampleDistCollective.cpp index 4be211d37f..f7af9f2024 100644 --- a/samples/sampleDistCollective/sampleDistCollective.cpp +++ b/samples/sampleDistCollective/sampleDistCollective.cpp @@ -40,11 +40,6 @@ #include "common.h" #include "logger.h" -#include "sampleDevice.h" -#include "sampleEngines.h" -#include "sampleInference.h" -#include "sampleOptions.h" -#include "sampleReporting.h" using namespace nvinfer1; using namespace sample; @@ -123,10 +118,6 @@ std::function pCreateInferBuilderInternal{}; std::function pCreateNvOnnxParserInternal{}; std::function pCreateNvOnnxRefitterInternal{}; -//! Track runtime used for the execution of trtexec. -//! Must be tracked as a global variable due to how library init functions APIs are organized. -RuntimeMode gUseRuntime = RuntimeMode::kFULL; - #if !TRT_STATIC template bool initLibrary(LibraryPtr& libPtr, std::string const& libName, FetchPtrs fetchFunc) @@ -173,12 +164,9 @@ bool initNvinfer() sample::gLogWarning << "Could not load function createInferRefitter_INTERNAL : " << e.what() << std::endl; } - if (gUseRuntime == RuntimeMode::kFULL) - { - pCreateInferBuilderInternal = l->symbolAddress("createInferBuilder_INTERNAL"); - } + pCreateInferBuilderInternal = l->symbolAddress("createInferBuilder_INTERNAL"); }; - return initLibrary(libnvinferPtr, sample::getRuntimeLibraryName(gUseRuntime), fetchPtrs); + return initLibrary(libnvinferPtr, kNVINFER_LIBNAME, fetchPtrs); #else pCreateInferRuntimeInternal = createInferRuntime_INTERNAL; pCreateInferRefitterInternal = createInferRefitter_INTERNAL; diff --git a/samples/sampleDynamicReshape/sampleDynamicReshape.cpp b/samples/sampleDynamicReshape/sampleDynamicReshape.cpp index d5c3f221f1..6dfbd8b98e 100644 --- a/samples/sampleDynamicReshape/sampleDynamicReshape.cpp +++ b/samples/sampleDynamicReshape/sampleDynamicReshape.cpp @@ -27,7 +27,6 @@ // Define TRT entrypoints used in common code #define DEFINE_TRT_ENTRYPOINTS 1 -#include "BatchStream.h" #include "argsParser.h" #include "buffers.h" #include "common.h" diff --git a/samples/sampleIOFormats/sampleIOFormats.cpp b/samples/sampleIOFormats/sampleIOFormats.cpp index b9f53590ea..f70f291932 100644 --- a/samples/sampleIOFormats/sampleIOFormats.cpp +++ b/samples/sampleIOFormats/sampleIOFormats.cpp @@ -33,7 +33,6 @@ #include "half.h" #include "logger.h" #include "parserOnnxConfig.h" -#include "sampleOptions.h" #include "NvInfer.h" #include "NvOnnxParser.h" diff --git a/samples/sampleSafeMNIST/README.md b/samples/sampleSafeMNIST/README.md index fbd6856d72..d273db0499 100644 --- a/samples/sampleSafeMNIST/README.md +++ b/samples/sampleSafeMNIST/README.md @@ -67,9 +67,6 @@ The Convolution layer computes a 2D (channel, height, and width) convolution, wi ``` This sample generates `safe_mnist.engine`, which is a binary file that contains the serialized engine data. - When the build produces a companion library holding the engine's generated host code, it is written - beside the engine as `safe_mnist.engine.so`. The infer part loads it from there, so keep the two together. - This sample reads ONNX model to build the network: - `safe_mnist.onnx` - The ONNX model that contains the network design. diff --git a/samples/sampleSafeMNIST/sampleSafeMNISTBuild.cpp b/samples/sampleSafeMNIST/sampleSafeMNISTBuild.cpp index 68925e15c5..4397f81ff9 100644 --- a/samples/sampleSafeMNIST/sampleSafeMNISTBuild.cpp +++ b/samples/sampleSafeMNIST/sampleSafeMNISTBuild.cpp @@ -245,24 +245,12 @@ bool SampleSafeMNIST::build() config->setRemoteAutoTuningConfig(mParams.remoteAutoTuningConfig.c_str()); } -#if ENABLE_UNIFIED_BUILDER - // A safety engine's generated host code lives in a companion library, and only this entry point hands - // back the two together. The other entry points return a single blob and so refuse a safety build - // once companion libraries are required. - auto const artifacts = std::unique_ptr( - builder->buildSerializedSafeNetwork(*network, *config, /*emitCheckerBlob=*/false)); - if (!artifacts) - { - return false; - } - auto const* const buffer = artifacts->getSerializedNetwork(); -#else auto buffer = std::unique_ptr(builder->buildSerializedNetwork(*network, *config)); + if (!buffer) { return false; } -#endif // ENABLE_UNIFIED_BUILDER ASSERT(network->getNbInputs() == 1); mInputDims = network->getInput(0)->getDimensions(); @@ -279,23 +267,6 @@ bool SampleSafeMNIST::build() file.write(reinterpret_cast(buffer->data()), buffer->size()); file.close(); -#if ENABLE_UNIFIED_BUILDER - // The engine cannot be loaded without its companion library, and the inference sample looks for it - // beside the engine. - if (auto const* const companionSo = artifacts->getCompanionSo()) - { - std::string const companionSoFile = engineFile + ".so"; - std::ofstream soFile(companionSoFile, std::ios::binary); - if (!soFile) - { - sample::gLogError << "Failed to open file to save companion library: " << companionSoFile << std::endl; - return false; - } - soFile.write(reinterpret_cast(companionSo->data()), companionSo->size()); - soFile.close(); - } -#endif // ENABLE_UNIFIED_BUILDER - return true; } @@ -365,7 +336,7 @@ int main(int argc, char** argv) if (!args.remoteAutoTuningConfig.empty()) { sample::gLogInfo << "Remote auto tuning config specified: " - << sample::sanitizeRemoteConfig(args.remoteAutoTuningConfig) << std::endl; + << sample::sanitizeRemoteAutoTuningConfig(args.remoteAutoTuningConfig) << std::endl; sample::gLogInfo << "This is a safety sample and will build in remote mode automatically." << std::endl; } diff --git a/samples/sampleSafeMNIST/sampleSafeMNISTInfer.cpp b/samples/sampleSafeMNIST/sampleSafeMNISTInfer.cpp index 8f18caa63d..5886e47890 100644 --- a/samples/sampleSafeMNIST/sampleSafeMNISTInfer.cpp +++ b/samples/sampleSafeMNIST/sampleSafeMNISTInfer.cpp @@ -337,9 +337,8 @@ bool doInference(SampleSafeMNISTInferArgs const& args) // Configure executor(s) std::vector graphs(nbThreads); - auto const companionSoPath = samplesSafeCommon::resolveCompanionSoPath(args.engineFileName); - SAFE_API_CALL(nvinfer2::safe::createTRTGraph(graphs[0], gieModelStream.data(), engineFileSize, - companionSoPath ? companionSoPath->c_str() : nullptr, *recorders[0], true), + SAFE_API_CALL(nvinfer2::safe::createTRTGraphWithSo( + graphs[0], gieModelStream.data(), engineFileSize, nullptr, *recorders[0], true), *recorders[0]); for (int32_t i = 1; i < nbThreads; ++i) @@ -449,20 +448,7 @@ int32_t main(int32_t argc, char** argv) return EXIT_SUCCESS; } - TestResult result = TestResult::kPASSED; - try - { - if (!doInference(args)) - { - result = TestResult::kFAILED; - } - } - catch (std::runtime_error& e) - { - SAFE_LOG << e.what() << std::endl; - result = TestResult::kFAILED; - } - + TestResult result = doInference(args) ? TestResult::kPASSED : TestResult::kFAILED; reportTestResult("TensorRT.sample_mnist_safe_infer", result, argc, argv); return EXIT_SUCCESS; diff --git a/samples/sampleSafePluginV3/README.md b/samples/sampleSafePluginV3/README.md index 90fea0f1e8..51f4ab1e89 100755 --- a/samples/sampleSafePluginV3/README.md +++ b/samples/sampleSafePluginV3/README.md @@ -81,9 +81,6 @@ See [Preparing sample data](../README.md#preparing-sample-data) in the main samp This sample generates `safe_plugin.engine`, which is a binary file that contains the serialized engine data. - When the build produces a companion library holding the engine's generated host code, it is written - beside the engine as `safe_plugin.engine.so`. The infer part loads it from there, so keep the two together. - This sample reads ONNX model to build the network: - `mnist_safe_plugin.onnx` - The ONNX model that contains the network design with maxPoolPlugin, version 1, namespace "" diff --git a/samples/sampleSafePluginV3/maxPoolPlugin.h b/samples/sampleSafePluginV3/maxPoolPlugin.h index da17f8dc77..096f98e717 100644 --- a/samples/sampleSafePluginV3/maxPoolPlugin.h +++ b/samples/sampleSafePluginV3/maxPoolPlugin.h @@ -20,6 +20,7 @@ #include #include +#include #include @@ -84,9 +85,15 @@ namespace nvinfer2::safe::consistency class MaxPoolPluginChecker : public IPluginChecker { public: +#if defined(NV_INFER_PLUGIN_CHECKER_POINTER_COUNT_API) bool validate(nvinfer2::safe::TensorDescriptor const* /*inputs*/, int32_t /*nbInputs*/, nvinfer2::safe::TensorDescriptor const* /*outputs*/, int32_t /*nbOutputs*/, nvinfer1::PluginFieldCollection* /*fc*/) noexcept override +#else + bool validate(std::vector const& /*inputs*/, + std::vector const& /*outputs*/, + nvinfer1::PluginFieldCollection* /*fc*/) noexcept override +#endif { // Always return true return true; diff --git a/samples/sampleSafePluginV3/sampleSafePluginBuild.cpp b/samples/sampleSafePluginV3/sampleSafePluginBuild.cpp index b6babe6994..7cabaafaa6 100644 --- a/samples/sampleSafePluginV3/sampleSafePluginBuild.cpp +++ b/samples/sampleSafePluginV3/sampleSafePluginBuild.cpp @@ -77,7 +77,8 @@ bool parseSampleSafePluginBuildArgs(SampleSafePluginBuildArgs& args, int32_t arg } else if (auto value = parseString(arg, "remoteAutoTuningConfig")) { - if (!sample::validateNonEmpty(*value, "Remote auto tuning config") || !sample::validateRemoteConfig(*value)) + if (!sample::validateNonEmpty(*value, "Remote auto tuning config") + || !sample::validateRemoteAutoTuningConfig(*value)) { return false; } @@ -266,24 +267,11 @@ bool SampleSafePlugin::build() config->setRemoteAutoTuningConfig(mParams.remoteAutoTuningConfig.c_str()); } -#if ENABLE_UNIFIED_BUILDER - // A safety engine's generated host code lives in a companion library, and only this entry point hands - // back the two together. The other entry points return a single blob and so refuse a safety build - // once companion libraries are required. - auto const artifacts = std::unique_ptr( - builder->buildSerializedSafeNetwork(*network, *config, /*emitCheckerBlob=*/false)); - if (!artifacts) - { - return false; - } - auto const* const buffer = artifacts->getSerializedNetwork(); -#else auto buffer = std::unique_ptr(builder->buildSerializedNetwork(*network, *config)); if (!buffer) { return false; } -#endif // ENABLE_UNIFIED_BUILDER ASSERT(network->getNbInputs() == 1); mInputDims = network->getInput(0)->getDimensions(); @@ -300,23 +288,6 @@ bool SampleSafePlugin::build() file.write(reinterpret_cast(buffer->data()), buffer->size()); file.close(); -#if ENABLE_UNIFIED_BUILDER - // The engine cannot be loaded without its companion library, and the inference sample looks for it - // beside the engine. - if (auto const* const companionSo = artifacts->getCompanionSo()) - { - std::string const companionSoFile = engineFile + ".so"; - std::ofstream soFile(companionSoFile, std::ios::binary); - if (!soFile) - { - sample::gLogError << "Failed to open file to save companion library: " << companionSoFile << std::endl; - return false; - } - soFile.write(reinterpret_cast(companionSo->data()), companionSo->size()); - soFile.close(); - } -#endif // ENABLE_UNIFIED_BUILDER - return true; } @@ -383,7 +354,7 @@ int main(int argc, char** argv) if (!args.remoteAutoTuningConfig.empty()) { sample::gLogInfo << "Remote auto tuning config specified: " - << sample::sanitizeRemoteConfig(args.remoteAutoTuningConfig) << std::endl; + << sample::sanitizeRemoteAutoTuningConfig(args.remoteAutoTuningConfig) << std::endl; sample::gLogInfo << "This is a safety sample and will build in remote mode automatically." << std::endl; } diff --git a/samples/sampleSafePluginV3/sampleSafePluginInfer.cpp b/samples/sampleSafePluginV3/sampleSafePluginInfer.cpp index 213f5d84a0..5800dc821b 100644 --- a/samples/sampleSafePluginV3/sampleSafePluginInfer.cpp +++ b/samples/sampleSafePluginV3/sampleSafePluginInfer.cpp @@ -238,9 +238,7 @@ bool doInference(SampleSafePluginInferArgs const& args) ITRTGraph* graph = nullptr; getSafePluginRegistry(g_recorder)->registerCreator(creator, "", g_recorder); - auto const companionSoPath = samplesSafeCommon::resolveCompanionSoPath(args.engineFileName); - createTRTGraph(graph, engineFile.data(), engineFileSize, companionSoPath ? companionSoPath->c_str() : nullptr, - g_recorder, true, nullptr); + createTRTGraphWithSo(graph, engineFile.data(), engineFileSize, nullptr, g_recorder, true, nullptr); SAFE_ASSERT(graph != nullptr); // Setup as many auxiliary streams as the graph requires - destroyed at scope end. diff --git a/samples/trtSafeExec/CMakeLists.txt b/samples/trtSafeExec/CMakeLists.txt index 3c36a93691..702145f2ae 100644 --- a/samples/trtSafeExec/CMakeLists.txt +++ b/samples/trtSafeExec/CMakeLists.txt @@ -20,20 +20,21 @@ add_executable(trtexec_safe trtSafeExec.cpp delayStreamKernel.cu) if(TRT_SAFETY_INFERENCE_ONLY) target_link_libraries(trtexec_safe PRIVATE trt_global_definitions tensorrt_headers) target_include_directories(trtexec_safe PRIVATE - ${CMAKE_CURRENT_SOURCE_DIR}/../common + ${CMAKE_CURRENT_SOURCE_DIR}/../trtexecCommon ) - # Gate out the CUDA kernel in delayStreamKernel.cu that uses __cudaLaunchKernel - # (not available in SafeCUDA); SafetyInferenceOnly.cmake sets the CMake variable - # but does not propagate it as a preprocessor define. - target_compile_definitions(trtexec_safe PRIVATE TRT_SAFETY_INFERENCE_ONLY) else() - target_link_libraries(trtexec_safe PRIVATE trt_samples_common TRTSAFE::nvinfer_safe_debug) + target_link_libraries(trtexec_safe PRIVATE trtexec_common TRTSAFE::nvinfer_safe_debug) endif() -add_dependencies(tensorrt_samples trtexec_safe) +if(TARGET tensorrt_tools) + add_dependencies(tensorrt_tools trtexec_safe) +elseif(TARGET tensorrt_samples) + add_dependencies(tensorrt_samples trtexec_safe) +endif() installLibraries( TARGETS trtexec_safe OPTIONAL COMPONENT external + EXPORT TensorRT ) diff --git a/samples/trtSafeExec/delayStreamKernel.cu b/samples/trtSafeExec/delayStreamKernel.cu index 47f00803be..3b88f23ac3 100644 --- a/samples/trtSafeExec/delayStreamKernel.cu +++ b/samples/trtSafeExec/delayStreamKernel.cu @@ -17,33 +17,48 @@ #include "delayStreamKernel.h" -#include +#include -#ifndef TRT_SAFETY_INFERENCE_ONLY namespace { -__global__ void delayKernel(long long nanoSeconds) +__device__ __forceinline__ uint64_t readGlobalTimer() { - // It is supported with compute capability 7.0 or higher. - __nanosleep(nanoSeconds); + uint64_t value; + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(value)); + return value; +} + +__global__ void delayKernel(uint64_t nanoSeconds) +{ + uint64_t const start{readGlobalTimer()}; + while (readGlobalTimer() - start < nanoSeconds) + { + // Busy-wait so that subsequent work can be submitted to the stream while this kernel is running. + } } } // namespace -#endif // TRT_SAFETY_INFERENCE_ONLY namespace nvinfer1 { -cudaError_t delayStream(cudaStream_t stream, float timeInMsec) noexcept +cudaError_t delayStream(cudaStream_t stream, std::chrono::duration duration) noexcept { -#ifndef TRT_SAFETY_INFERENCE_ONLY - auto nanoSeconds = static_cast(1000000 * timeInMsec); - delayKernel<<<1, 1, 0, stream>>>(nanoSeconds); + using FloatMilliseconds = std::chrono::duration; + if (duration < FloatMilliseconds::zero()) + { + return cudaErrorInvalidValue; + } + if (duration == FloatMilliseconds::zero()) + { + return cudaSuccess; + } + constexpr double kNANOSECONDS_PER_MILLISECOND{1000000.0}; + auto const nanoSeconds = kNANOSECONDS_PER_MILLISECOND * static_cast(duration.count()); + if (!(nanoSeconds < static_cast(std::numeric_limits::max()))) + { + // This comparison also rejects NaN and infinity before the integer conversion. + return cudaErrorInvalidValue; + } + delayKernel<<<1, 1, 0, stream>>>(static_cast(nanoSeconds)); return cudaGetLastError(); -#else - // QNX SafeCUDA does not support PTX JIT or __cudaLaunchKernel; the delay is - // optional timing-measurement padding and has no effect on inference correctness. - std::ignore = stream; - std::ignore = timeInMsec; - return cudaSuccess; -#endif // TRT_SAFETY_INFERENCE_ONLY } } // namespace nvinfer1 diff --git a/samples/trtSafeExec/delayStreamKernel.h b/samples/trtSafeExec/delayStreamKernel.h index d26472c6a2..1be1131493 100644 --- a/samples/trtSafeExec/delayStreamKernel.h +++ b/samples/trtSafeExec/delayStreamKernel.h @@ -1,5 +1,5 @@ /* -* SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -18,11 +18,16 @@ #ifndef DELAY_STREAM_KERNEL_H #define DELAY_STREAM_KERNEL_H +#include #include #include namespace nvinfer1 { -cudaError_t delayStream(cudaStream_t stream, float timeInMsec) noexcept; +//! \brief Launch a kernel on \p stream that busy-waits for \p duration, delaying subsequent work on the stream. +//! +//! \return cudaSuccess on success, cudaErrorInvalidValue if \p duration is negative or not representable as a +//! nanosecond count, otherwise the error from the kernel launch. +cudaError_t delayStream(cudaStream_t stream, std::chrono::duration duration) noexcept; } // namespace nvinfer1 #endif // DELAY_STREAM_KERNEL_H diff --git a/samples/trtSafeExec/trtSafeExec.cpp b/samples/trtSafeExec/trtSafeExec.cpp index 66ac572155..13f7db5f8f 100644 --- a/samples/trtSafeExec/trtSafeExec.cpp +++ b/samples/trtSafeExec/trtSafeExec.cpp @@ -23,14 +23,13 @@ #include "safeErrorRecorder.h" #include #include +#include #include #include #include -#include #include #include #include -#include #include #include #include @@ -61,15 +60,14 @@ class SafeExecArgs { public: std::string engineFile{"sample.engine"}; - std::string companionSoFile; int32_t iterations{10}; int32_t avgRuns{10}; int32_t warmUp{1}; int32_t device{0}; int32_t streams{1}; - float idle{0.F}; - float duration{3.F}; - float sleep{2.F}; + std::chrono::duration idle{0.F}; + std::chrono::duration duration{3.F}; + std::chrono::duration sleep{2.F}; float percentile{99.F}; bool spin{false}; bool verbose{false}; @@ -104,9 +102,6 @@ namespace //! Default alignment for memory allocations constexpr uint64_t kDEFAULT_ALIGNMENT{256U}; -//! Maximum number of bytes requested by one engine file read. -constexpr size_t kFILE_READ_CHUNK_BYTES{64U * 1024U * 1024U}; - //! //! \brief RAII wrapper for SafeMemAllocator to ensure automatic cleanup. //! @@ -269,9 +264,9 @@ SafePerformanceResult getSafePerformanceResult(TimingMetrics const& times, int32 SafePerformanceResult result; result.min = newTimes[0][metricIndex]; result.max = newTimes[newTimes.size() - 1][metricIndex]; - result.mean = std::accumulate(newTimes.begin(), newTimes.end(), 0.F, - [metricIndex](float acc, TimingMetric& a) { return acc + a[metricIndex]; }) - / newTimes.size(); + result.mean = std::accumulate(newTimes.begin(), newTimes.end(), 0.F, [metricIndex](float acc, TimingMetric& a) { + return acc + a[metricIndex]; + }) / newTimes.size(); size_t const medianIndex = newTimes.size() / 2ULL; result.median = newTimes.size() % 2ULL ? newTimes[medianIndex][metricIndex] @@ -501,7 +496,7 @@ bool parseSafetyPluginLibrary( } // Use template to allow volume for either nvinfer1::Dims or nvinfer2::safe::PhysicalDims -template +template int64_t volume(TDims const& dims, TDims const& strides, uint64_t bytesPerComponent) { if (dims.nbDims == 0 || strides.nbDims == 0) @@ -535,11 +530,6 @@ bool parseSafeExecArgs(SafeExecArgs& args, int32_t argc, char* argv[]) for (int32_t i = 1; i < argc; ++i) { std::string const arg = argv[i]; - if (auto value = loggedParseString(arg, "loadEngineSo")) - { - args.companionSoFile = std::move(*value); - continue; - } if (auto value = loggedParseString(arg, "loadEngine")) { args.engineFile = std::move(*value); @@ -570,15 +560,15 @@ bool parseSafeExecArgs(SafeExecArgs& args, int32_t argc, char* argv[]) } else if (auto const value = loggedParseString(arg, "idleTime")) { - args.idle = stof(*value); + args.idle = std::chrono::duration(stof(*value)); } else if (auto const value = loggedParseString(arg, "duration")) { - args.duration = stof(*value); + args.duration = std::chrono::duration(stof(*value)); } else if (auto const value = loggedParseString(arg, "sleepTime")) { - args.sleep = stof(*value); + args.sleep = std::chrono::duration(stof(*value)); } else if (loggedParseBool(arg, "spin")) { @@ -599,8 +589,8 @@ bool parseSafeExecArgs(SafeExecArgs& args, int32_t argc, char* argv[]) else if (loggedParseBool(arg, "useCudaGraph")) { // Deprecated: CUDA graph is now enabled by default. - safeLogWarning(*gSafeRecorder, - "--useCudaGraph is deprecated (now enabled by default). Use --noCudaGraph to disable."); + safeLogWarning( + *gSafeRecorder, "--useCudaGraph is deprecated (now enabled by default). Use --noCudaGraph to disable."); } else if (loggedParseBool(arg, "noCudaGraph")) { @@ -664,9 +654,6 @@ void printHelpInfo() std::cout << R"(Usage: trtexec_safe --loadEngine= [options] Required params: --loadEngine=FILE Load the serialized engine from FILE. - --loadEngineSo=FILE - Load the engine's companion library from FILE. Without it, .so is - used when that file exists. General optional params: --help or -h Display help information @@ -697,16 +684,16 @@ Perf measurement params: --warmUp=N Run N iterations before actual perf measurement (default = )" << defArgs.warmUp << R"() --idleTime=F Sleep F milliseconds between two continuous iterations (default = )" - << defArgs.idle << R"() + << defArgs.idle.count() << R"() --percentile=P For each iteration, report the percentile time at P percentage (0<=P<=100, with 0 representing min, and 100 representing max; default = )" << defArgs.percentile << R"(%) --noCudaGraph Disable CUDA graph capture and launch (default = CUDA graph enabled) --useCudaGraph [Deprecated] CUDA graph is now enabled by default. This flag is a no-op. --duration=F Run performance measurements for at least F seconds of wallclock time (default = )" - << defArgs.duration << R"(s) + << defArgs.duration.count() << R"(s) --sleepTime=F Delay inference start with a gap of F msec between launch and compute (default = )" - << defArgs.sleep << R"() + << defArgs.sleep.count() << R"() --separateProfileRun [Deprecated] Separate profile run is now always enabled. This flag is a no-op. @@ -778,52 +765,24 @@ void registerSafetyPlugins(nvinfer2::safe::ISafeRecorder& recorder, SafetyPlugin } //! -//! \brief Load a prebuilt TensorRT safe engine using bounded read requests. -//! \param engineFile Path to the serialized engine. -//! \return Buffer containing the complete serialized engine. -//! \throws std::runtime_error if the file size is invalid or any read is incomplete. +//! \brief Load a prebuilt TensorRT safe engine. +//! std::vector loadEngine(std::string const& engineFile) { - std::ifstream file(engineFile, std::ios::binary | std::ios::ate); - if (!file) - { - throw std::runtime_error("Failed to open engine file: " + engineFile); - } - - std::streamoff const fileSize{file.tellg()}; - if (fileSize <= 0) - { - throw std::runtime_error("Engine file is empty or has an invalid size: " + engineFile); - } - if (static_cast(fileSize) > std::numeric_limits::max()) - { - throw std::runtime_error("Engine file is too large to load into memory: " + engineFile); - } - - file.seekg(0, std::ios::beg); - if (!file) - { - throw std::runtime_error("Failed to seek to the beginning of engine file: " + engineFile); - } - - size_t const size{static_cast(fileSize)}; - std::vector modelBuffer(size); - size_t offset{0U}; - while (offset < size) - { - size_t const bytesToRead{std::min(kFILE_READ_CHUNK_BYTES, size - offset)}; - file.read(modelBuffer.data() + offset, static_cast(bytesToRead)); - std::streamsize const bytesRead{file.gcount()}; - if (bytesRead != static_cast(bytesToRead) || file.fail() || file.bad()) - { - std::ostringstream message; - message << "Failed to read complete engine file: " << engineFile << " at offset " << offset << ". Expected " - << bytesToRead << " bytes, read " << bytesRead << " (eof=" << file.eof() << ", fail=" << file.fail() - << ", bad=" << file.bad() << ")"; - throw std::runtime_error(message.str()); - } - offset += bytesToRead; - } + std::string const& filename = engineFile; + std::vector modelBuffer; + std::ifstream file(filename, std::ios::binary); + if (!file.good()) + { + safeLogError(*gSafeRecorder, "Could not open input engine file or file is empty. File name: " + filename); + return modelBuffer; + } + file.seekg(0, std::ifstream::end); + auto size = file.tellg(); + file.seekg(0, std::ifstream::beg); + modelBuffer.resize(size); + file.read(modelBuffer.data(), size); + file.close(); return modelBuffer; } @@ -1077,7 +1036,7 @@ bool task(SafeExecArgs const& args, nvinfer2::safe::ITRTGraph* graph, nvinfer2:: // GPU, host and enqueue times TimingMetrics totalTimes; using floatDurationMS = std::chrono::duration; - floatDurationMS const maxDurationMs = floatDurationMS(args.duration * 1000); + floatDurationMS const maxDurationMs = args.duration; floatDurationMS durationMs{0}; for (int32_t i = 0; i < nbIterations || durationMs.count() < maxDurationMs.count(); i++) @@ -1160,7 +1119,7 @@ bool task(SafeExecArgs const& args, nvinfer2::safe::ITRTGraph* graph, nvinfer2:: totalEnqueueTime += enqueueTime; // Mimic waiting for user input data (default = 0) - std::this_thread::sleep_for(std::chrono::duration(args.idle)); + std::this_thread::sleep_for(args.idle); } if (isProfileRun) @@ -1292,10 +1251,8 @@ bool doInference(SafeExecArgs const& args, std::chrono::high_resolution_clock::t // Configure executor(s) std::vector graphs(numThreads); std::vector scratchs(numThreads); - auto const companionSoPath = samplesSafeCommon::resolveCompanionSoPath(args.engineFile, args.companionSoFile); - SAFE_API_CALL(nvinfer2::safe::createTRTGraph(graphs[0], blob.data(), blob.size(), - companionSoPath ? companionSoPath->c_str() : nullptr, *recorders[0], !args.useScratchMemory, - &nvinfer2::safe::getSafeMemAllocator()), + SAFE_API_CALL(nvinfer2::safe::createTRTGraphWithSo(graphs[0], blob.data(), blob.size(), nullptr, *recorders[0], + !args.useScratchMemory, &nvinfer2::safe::getSafeMemAllocator()), *recorders[0]); SAFE_API_CALL(graphs[0]->setIOProfile(args.ioProfile), *recorders[0]); diff --git a/samples/trtexec/CMakeLists.txt b/samples/trtexec/CMakeLists.txt index 2bf75650cb..c9ca0bb769 100644 --- a/samples/trtexec/CMakeLists.txt +++ b/samples/trtexec/CMakeLists.txt @@ -15,8 +15,10 @@ # limitations under the License. # add_executable(trtexec trtexec_main.cpp trtexec.cpp) -target_link_libraries(trtexec PRIVATE trt_samples_common) -if (TRT_BUILD_SAMPLES) +target_link_libraries(trtexec PRIVATE trtexec_common) +if(TARGET tensorrt_tools) + add_dependencies(tensorrt_tools trtexec) +elseif(TARGET tensorrt_samples) add_dependencies(tensorrt_samples trtexec) endif() @@ -42,6 +44,7 @@ installLibraries( TARGETS trtexec OPTIONAL COMPONENT external + EXPORT TensorRT ) # When statically linked, trtexec requires the plugins library. @@ -66,8 +69,10 @@ set_target_properties(trtexec # In this mode, we build an additional binary trtexec_static that always links tensorrt_static. if(${TRT_BUILD_TRTEXEC_STATIC}) add_executable(trtexec_static trtexec_main.cpp trtexec.cpp) - target_link_libraries(trtexec_static PRIVATE trt_samples_common) - if (TRT_BUILD_SAMPLES) + target_link_libraries(trtexec_static PRIVATE trtexec_common) + if(TARGET tensorrt_tools) + add_dependencies(tensorrt_tools trtexec_static) + elseif(TARGET tensorrt_samples) add_dependencies(tensorrt_samples trtexec_static) endif() @@ -93,7 +98,7 @@ endif() if(TRT_BUILD_WINML_PLUGIN) set(internal_sample_name "tensorrt_rtx_internal") add_executable(${internal_sample_name} trtexec_main.cpp trtexec.cpp) - target_link_libraries(${internal_sample_name} PRIVATE trt_samples_common) + target_link_libraries(${internal_sample_name} PRIVATE trtexec_common) target_compile_definitions(${internal_sample_name} PRIVATE TRT_WINML_PLUGIN=1) installLibraries( diff --git a/samples/trtexec/prn_utils.py b/samples/trtexec/prn_utils.py index 6b0abf9fb3..f78e2410fb 100755 --- a/samples/trtexec/prn_utils.py +++ b/samples/trtexec/prn_utils.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 # -# SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 # # Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/trtexec/trtexec.cpp b/samples/trtexec/trtexec.cpp index 5075528287..f8f22086ba 100644 --- a/samples/trtexec/trtexec.cpp +++ b/samples/trtexec/trtexec.cpp @@ -23,18 +23,14 @@ #include "trtexec.h" -#if ENABLE_UNIFIED_BUILDER -#include "safeCommon.h" -#endif - #include #include #include #include -#include #include #include #include +#include #include #include #include @@ -344,9 +340,8 @@ int32_t printBuildRouteHelp(std::string const& knobName) // Database stores names with a leading '-'; accept the user's input both with and without // the dash by comparing the unprefixed substrings. - auto stripDash = [](std::string const& s) -> std::string { - return (!s.empty() && s.front() == '-') ? s.substr(1) : s; - }; + auto stripDash + = [](std::string const& s) -> std::string { return (!s.empty() && s.front() == '-') ? s.substr(1) : s; }; std::string const wantedKnob = stripDash(knobName); nlohmann::ordered_json filteredOptions = nlohmann::ordered_json::array(); @@ -615,6 +610,12 @@ int32_t runOnceBuildAndInfer( new BuildEnvironment(options.build.safe, options.build.versionCompatible, options.system.DLACore, options.build.tempdir, options.build.tempfileControls, options.build.leanDLLPath, sampleTest.getCmdline())); + if (!bEnv->greenContexts.initialize(options.build, options.system.device, sample::gLogError)) + { + sample::gLogError << "CUDA green context set up failed" << std::endl; + return EXIT_FAILURE; + } + #if !TRT_WINML bEnv->engine.setDLAWorkspaceAllocationStrategy(options.system.dlaWorkspaceAllocationStrategy); #endif // !TRT_WINML @@ -907,8 +908,8 @@ int32_t runOnceBuildAndInfer( std::ofstream out(options.tuning.tuningResultFile); if (!out) { - sample::gLogError << "Cannot open --tuningResultFile for writing: " - << options.tuning.tuningResultFile << std::endl; + sample::gLogError << "Cannot open --tuningResultFile for writing: " << options.tuning.tuningResultFile + << std::endl; } else { @@ -1015,8 +1016,7 @@ IterationResult readChildResult(std::string const& jsonPath) if (!in) { r.crashed = true; - r.errorMessage = "missing tuning result file " + jsonPath - + " (child likely crashed before writing)"; + r.errorMessage = "missing tuning result file " + jsonPath + " (child likely crashed before writing)"; return r; } try @@ -1286,9 +1286,9 @@ void emitDryRunListing(TuningContext const& ctx) //! nested lambda. struct PhaseState { - AllOptions const& options; //!< Parsed options for this run. - Logger::TestAtom const& sampleTest; //!< For TASK_BEGIN/END/ABORT banners. - pid_t const ppid{}; //!< Parent PID, used in temp filenames. + AllOptions const& options; //!< Parsed options for this run. + Logger::TestAtom const& sampleTest; //!< For TASK_BEGIN/END/ABORT banners. + pid_t const ppid{}; //!< Parent PID, used in temp filenames. std::chrono::steady_clock::time_point const startTime; //!< Loop start, for --tuningTimeOut. int32_t const argc{}; //!< Parent argv (passed verbatim to children). char** const argv{}; //!< Parent argv (passed verbatim to children). @@ -1332,7 +1332,8 @@ std::string makeIterationEnginePath(PhaseState const& state, char const* phaseLa bool runOnePhase(PhaseState& state, TuningContext const& phaseCtx, char const* phaseLabel, std::vector* positiveKnobs, double* baselineGpuTimeMsOut, int64_t skipUntil) { - sample::gLogInfo << "Tuning " << phaseLabel << ": " << phaseCtx.totalCount.toString() << " iterations." << std::endl; + sample::gLogInfo << "Tuning " << phaseLabel << ": " << phaseCtx.totalCount.toString() << " iterations." + << std::endl; double baselineGpuTimeMs = std::numeric_limits::infinity(); for (BigInt i{0}; i < phaseCtx.totalCount; ++i) { @@ -1343,20 +1344,21 @@ bool runOnePhase(PhaseState& state, TuningContext const& phaseCtx, char const* p } if (state.options.tuning.timeout > 0) { - auto const elapsedS = std::chrono::duration_cast( - std::chrono::steady_clock::now() - state.startTime).count(); + auto const elapsedS + = std::chrono::duration_cast(std::chrono::steady_clock::now() - state.startTime) + .count(); if (elapsedS >= state.options.tuning.timeout) { - sample::gLogInfo << "Tuning timeout reached (" << state.options.tuning.timeout - << "s); stopping early." << std::endl; + sample::gLogInfo << "Tuning timeout reached (" << state.options.tuning.timeout << "s); stopping early." + << std::endl; return false; } } std::string const route = phaseCtx.getPathAtIndex(i); std::string const enginePath = makeIterationEnginePath(state, phaseLabel, i); - std::string const jsonPath = "/tmp/trtexec_tuning_" + std::to_string(state.ppid) - + "_iter" + i.toString() + ".json"; + std::string const jsonPath + = "/tmp/trtexec_tuning_" + std::to_string(state.ppid) + "_iter" + i.toString() + ".json"; sample::gLogger.reportTaskBegin(state.sampleTest, i.toString(), route); @@ -1380,8 +1382,9 @@ bool runOnePhase(PhaseState& state, TuningContext const& phaseCtx, char const* p } else { - sample::gLogWarning << "Iteration [" << i.toString() << "] failed: " - << (result.errorMessage.empty() ? "(no message)" : result.errorMessage) << std::endl; + sample::gLogWarning << "Iteration [" << i.toString() + << "] failed: " << (result.errorMessage.empty() ? "(no message)" : result.errorMessage) + << std::endl; sample::gLogger.reportTaskAbort(state.sampleTest, i.toString(), route); } // For mixed-mode phase 1, collect knobs that beat the baseline. @@ -1517,21 +1520,21 @@ int32_t sample::runTuningLoop(int32_t argc, char** argv) std::vector positiveKnobs; double phase1BaselineMs{std::numeric_limits::infinity()}; bool const isMixed = options.tuning.tuningSearchAlgorithm == TuningSearchAlgorithm::kMIXED; - bool const phase1Completed - = runOnePhase(state, ctx, "phase1", isMixed ? &positiveKnobs : nullptr, &phase1BaselineMs, resume.resumeFromIter); + bool const phase1Completed = runOnePhase( + state, ctx, "phase1", isMixed ? &positiveKnobs : nullptr, &phase1BaselineMs, resume.resumeFromIter); if (phase1Completed && isMixed && positiveKnobs.size() > 1) { - sample::gLogInfo << "Mixed search: " << positiveKnobs.size() - << " positive knobs identified; entering phase 2." << std::endl; + sample::gLogInfo << "Mixed search: " << positiveKnobs.size() << " positive knobs identified; entering phase 2." + << std::endl; TuningContext const phase2Ctx = buildMixedPhase2Context(ctx, positiveKnobs); // Phase 2 always starts fresh (no resume mid-phase-2). (void) runOnePhase(state, phase2Ctx, "phase2", nullptr, nullptr, 0); } else if (isMixed) { - sample::gLogInfo << "Mixed search: " << positiveKnobs.size() - << " positive knob(s); skipping phase 2 (need >1)." << std::endl; + sample::gLogInfo << "Mixed search: " << positiveKnobs.size() << " positive knob(s); skipping phase 2 (need >1)." + << std::endl; } // 7. Promote the best iteration's engine to the user's --saveEngine path. diff --git a/samples/trtexec/trtexec_main.cpp b/samples/trtexec/trtexec_main.cpp index 10c06f8459..414156c2af 100644 --- a/samples/trtexec/trtexec_main.cpp +++ b/samples/trtexec/trtexec_main.cpp @@ -23,8 +23,7 @@ //! for its child fork+execs. Each branch does its own parseArgs. int main(int argc, char** argv) { - if (sample::peekArg(argc, argv, "--tuneBuildRoutes") - || sample::peekArg(argc, argv, "--tuneBuildRouteFile") + if (sample::peekArg(argc, argv, "--tuneBuildRoutes") || sample::peekArg(argc, argv, "--tuneBuildRouteFile") || sample::peekArg(argc, argv, "--continue")) { return sample::runTuningLoop(argc, argv); diff --git a/samples/trtexecCommon/.clang-tidy b/samples/trtexecCommon/.clang-tidy new file mode 100644 index 0000000000..78bc3f53ce --- /dev/null +++ b/samples/trtexecCommon/.clang-tidy @@ -0,0 +1,20 @@ +# These files moved here from samples/ unchanged. The gate lints the diff, so the move presents this +# long-standing code as newly added and its pre-existing violations all fire at once. Suppress the +# checks it trips so the move stays a pure move; cleaning them up is tracked separately. +InheritParentConfig: true +Checks: >- + -cert-dcl21-cpp, + -modernize-loop-convert, + -modernize-pass-by-value, + -modernize-raw-string-literal, + -modernize-type-traits, + -modernize-unary-static-assert, + -modernize-use-emplace, + -modernize-use-nodiscard, + -modernize-use-starts-ends-with, + -modernize-use-transparent-functors, + -readability-identifier-naming, + -readability-isolate-declaration, + -readability-redundant-inline-specifier, + -readability-redundant-string-init, + -readability-static-definition-in-anonymous-namespace diff --git a/samples/common/ArgVec.test.h b/samples/trtexecCommon/ArgVec.test.h similarity index 100% rename from samples/common/ArgVec.test.h rename to samples/trtexecCommon/ArgVec.test.h diff --git a/samples/trtexecCommon/CMakeLists.txt b/samples/trtexecCommon/CMakeLists.txt new file mode 100644 index 0000000000..4dea308f01 --- /dev/null +++ b/samples/trtexecCommon/CMakeLists.txt @@ -0,0 +1,116 @@ +# +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +if(NOT NLOHMANN_JSON_INCLUDE_DIRS) + include(FetchNlohmannJson) +endif() + +add_library(trtexec_common STATIC + bfloat16.cpp + bfloat16.h + bigInt.cpp + bigInt.h + buffers.h + common.cpp + common.h + debugTensorWriter.cpp + debugTensorWriter.h + ErrorRecorder.h + globalTimerKernel.cu + globalTimerKernel.h + half.h + logger.cpp + logger.h + logging.h + sampleDevice.cpp + sampleDevice.h + sampleEngines.cpp + sampleEngines.h + sampleEntrypoints.h + sampleInference.cpp + sampleInference.h + sampleOptions.cpp + sampleOptions.h + sampleReporting.cpp + sampleReporting.h + sampleTuning.cpp + sampleTuning.h + sampleUtils.cpp + sampleUtils.h + safeCommon.h + safeCudaAllocator.h + safeErrorRecorder.h + streamReader.h +) + +if (${TRT_BUILD_TESTING}) + include(GoogleTest) + enable_testing() + + add_executable(trtexec_common_test + bfloat16.test.cpp + half.test.cpp + sampleOptions.test.cpp + sampleUtils.test.cpp + ) + + target_link_libraries(trtexec_common_test PRIVATE + gtest_main + trtexec_common + ) + + gtest_discover_tests(trtexec_common_test DISCOVERY_MODE ${TRT_GTEST_DISCOVERY_MODE}) + + if(NOT ${TRT_BUILD_SAMPLES}) + set_target_properties(trtexec_common_test PROPERTIES + EXCLUDE_FROM_ALL TRUE + ) + endif() +endif() # TRT_BUILD_TESTING + +target_include_directories(trtexec_common PUBLIC + ${CMAKE_CURRENT_LIST_DIR} + ${NLOHMANN_JSON_INCLUDE_DIRS} +) + +target_link_libraries(trtexec_common PUBLIC + tensorrt_headers + trt_shared + trt_global_definitions + Threads::Threads + $ # Each sample individually must determine its linkage to TRT. + TRT_SAMPLES::onnxparser +) + +if(TARGET TRT::cudart) + target_link_libraries(trtexec_common PUBLIC TRT::cudart) +else() + target_link_libraries(trtexec_common PUBLIC CUDA::cudart_static) +endif() + +if(NOT WIN32 AND NOT ${CMAKE_SYSTEM_NAME} STREQUAL "QNX") + target_link_libraries(trtexec_common PUBLIC dl) +endif() + +target_link_libraries(trtexec_common PUBLIC $) + +# For statically-linked samples, we need to upgrade the link to always link TRT rather than letting the samples decide. +if(${TRT_BUILD_SAMPLES_LINK_STATIC_TRT}) + target_link_libraries(trtexec_common PUBLIC + $ # Has to be whole archive so we keep the builder resources correctly. + ) +endif() diff --git a/samples/trtexecCommon/ErrorRecorder.h b/samples/trtexecCommon/ErrorRecorder.h new file mode 100644 index 0000000000..7732aa2432 --- /dev/null +++ b/samples/trtexecCommon/ErrorRecorder.h @@ -0,0 +1,138 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef ERROR_RECORDER_H +#define ERROR_RECORDER_H +#include "NvInferRuntimeBase.h" +#include "logger.h" +#include +#include +#include +#include +#include + +using nvinfer1::IErrorRecorder; +using nvinfer1::ErrorCode; + +//! +//! A simple implementation of the IErrorRecorder interface for +//! use by samples. This interface also can be used as a reference +//! implementation. +//! The sample Error recorder is based on a vector that pairs the error +//! code and the error string into a single element. It also uses +//! standard mutex's and atomics in order to make sure that the code +//! works in a multi-threaded environment. +//! +class SampleErrorRecorder : public IErrorRecorder +{ + using errorPair = std::pair; + using errorStack = std::vector; + +public: + SampleErrorRecorder() = default; + + ~SampleErrorRecorder() noexcept override {} + int32_t getNbErrors() const noexcept final + { + return mErrorStack.size(); + } + ErrorCode getErrorCode(int32_t errorIdx) const noexcept final + { + return invalidIndexCheck(errorIdx) ? ErrorCode::kINVALID_ARGUMENT : (*this)[errorIdx].first; + } + IErrorRecorder::ErrorDesc getErrorDesc(int32_t errorIdx) const noexcept final + { + return invalidIndexCheck(errorIdx) ? "errorIdx out of range." : (*this)[errorIdx].second.c_str(); + } + // This class can never overflow since we have dynamic resize via std::vector usage. + bool hasOverflowed() const noexcept final + { + return false; + } + + // Empty the errorStack. + void clear() noexcept final + { + try + { + // grab a lock so that there is no addition while clearing. + std::lock_guard guard(mStackLock); + mErrorStack.clear(); + } + catch (std::exception const& e) + { + sample::gLogFatal << "Internal Error: " << e.what() << std::endl; + } + } + + //! Simple helper function that + bool empty() const noexcept + { + return mErrorStack.empty(); + } + + bool reportError(ErrorCode val, IErrorRecorder::ErrorDesc desc) noexcept final + { + try + { + std::lock_guard guard(mStackLock); + sample::gLogError << "Error[" << static_cast(val) << "]: " << desc << std::endl; + mErrorStack.push_back(errorPair(val, desc)); + } + catch (std::exception const& e) + { + sample::gLogFatal << "Internal Error: " << e.what() << std::endl; + } + // All errors are considered fatal. + return true; + } + + // Atomically increment or decrement the ref counter. + IErrorRecorder::RefCount incRefCount() noexcept final + { + return ++mRefCount; + } + IErrorRecorder::RefCount decRefCount() noexcept final + { + return --mRefCount; + } + +private: + // Simple helper functions. + errorPair const& operator[](size_t index) const noexcept + { + return mErrorStack[index]; + } + + bool invalidIndexCheck(int32_t index) const noexcept + { + // By converting signed to unsigned, we only need a single check since + // negative numbers turn into large positive greater than the size. + size_t sIndex = index; + return sIndex >= mErrorStack.size(); + } + // Mutex to hold when locking mErrorStack. + std::mutex mStackLock; + + // Reference count of the class. Destruction of the class when mRefCount + // is not zero causes undefined behavior. + std::atomic mRefCount{0}; + + // The error stack that holds the errors recorded by TensorRT. + errorStack mErrorStack; +}; // class SampleErrorRecorder +#endif // ERROR_RECORDER_H diff --git a/samples/trtexecCommon/README.md b/samples/trtexecCommon/README.md new file mode 100644 index 0000000000..52bd8fbe14 --- /dev/null +++ b/samples/trtexecCommon/README.md @@ -0,0 +1,9 @@ +# tools/trtexecCommon + +Shared utility library (`trtexec_common`) backing `trtexec` and `trtexec_safe`. + +`trtexec_safe` uses only `safeCommon.h` and `safeErrorRecorder.h`; everything else here is +trtexec's engine, option model and reporting. + +The TensorRT samples live in the OSS repo and carry their own copy of the utilities they +need under `samples/common`. diff --git a/samples/common/bfloat16.cpp b/samples/trtexecCommon/bfloat16.cpp similarity index 96% rename from samples/common/bfloat16.cpp rename to samples/trtexecCommon/bfloat16.cpp index 8222826ae4..67de9dc200 100644 --- a/samples/common/bfloat16.cpp +++ b/samples/trtexecCommon/bfloat16.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/bfloat16.h b/samples/trtexecCommon/bfloat16.h similarity index 94% rename from samples/common/bfloat16.h rename to samples/trtexecCommon/bfloat16.h index 0d0ab92229..a71d870812 100644 --- a/samples/common/bfloat16.h +++ b/samples/trtexecCommon/bfloat16.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/bfloat16.test.cpp b/samples/trtexecCommon/bfloat16.test.cpp similarity index 100% rename from samples/common/bfloat16.test.cpp rename to samples/trtexecCommon/bfloat16.test.cpp diff --git a/samples/common/bigInt.cpp b/samples/trtexecCommon/bigInt.cpp similarity index 98% rename from samples/common/bigInt.cpp rename to samples/trtexecCommon/bigInt.cpp index cdff2151f3..041c76f6c3 100644 --- a/samples/common/bigInt.cpp +++ b/samples/trtexecCommon/bigInt.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2024-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/bigInt.h b/samples/trtexecCommon/bigInt.h similarity index 99% rename from samples/common/bigInt.h rename to samples/trtexecCommon/bigInt.h index e1045094e9..b90fc893b4 100644 --- a/samples/common/bigInt.h +++ b/samples/trtexecCommon/bigInt.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2024-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/trtexecCommon/buffers.h b/samples/trtexecCommon/buffers.h new file mode 100644 index 0000000000..b9bca321d2 --- /dev/null +++ b/samples/trtexecCommon/buffers.h @@ -0,0 +1,427 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef TENSORRT_BUFFERS_H +#define TENSORRT_BUFFERS_H + +#include "NvInfer.h" +#include "common.h" +#include "half.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace samplesCommon +{ + +//! +//! \brief The GenericBuffer class is a templated class for buffers. +//! +//! \details This templated RAII (Resource Acquisition Is Initialization) class handles the allocation, +//! deallocation, querying of buffers on both the device and the host. +//! It can handle data of arbitrary types because it stores byte buffers. +//! The template parameters AllocFunc and FreeFunc are used for the +//! allocation and deallocation of the buffer. +//! AllocFunc must be a functor that takes in (void** ptr, size_t size) +//! and returns bool. ptr is a pointer to where the allocated buffer address should be stored. +//! size is the amount of memory in bytes to allocate. +//! The boolean indicates whether or not the memory allocation was successful. +//! FreeFunc must be a functor that takes in (void* ptr) and returns void. +//! ptr is the allocated buffer address. It must work with nullptr input. +//! +template +class GenericBuffer +{ +public: + //! + //! \brief Construct an empty buffer. + //! + GenericBuffer(nvinfer1::DataType type = nvinfer1::DataType::kFLOAT) + : mSize(0) + , mCapacity(0) + , mType(type) + , mBuffer(nullptr) + { + } + + //! + //! \brief Construct a buffer with the specified allocation size in bytes. + //! + GenericBuffer(size_t size, nvinfer1::DataType type) + : mSize(size) + , mCapacity(size) + , mType(type) + { + if (!allocFn(&mBuffer, this->nbBytes())) + { + throw std::bad_alloc(); + } + } + + GenericBuffer(GenericBuffer&& buf) + : mSize(buf.mSize) + , mCapacity(buf.mCapacity) + , mType(buf.mType) + , mBuffer(buf.mBuffer) + { + buf.mSize = 0; + buf.mCapacity = 0; + buf.mType = nvinfer1::DataType::kFLOAT; + buf.mBuffer = nullptr; + } + + GenericBuffer& operator=(GenericBuffer&& buf) + { + if (this != &buf) + { + freeFn(mBuffer); + mSize = buf.mSize; + mCapacity = buf.mCapacity; + mType = buf.mType; + mBuffer = buf.mBuffer; + // Reset buf. + buf.mSize = 0; + buf.mCapacity = 0; + buf.mBuffer = nullptr; + } + return *this; + } + + //! + //! \brief Returns pointer to underlying array. + //! + void* data() + { + return mBuffer; + } + + //! + //! \brief Returns pointer to underlying array. + //! + void const* data() const + { + return mBuffer; + } + + //! + //! \brief Returns the size (in number of elements) of the buffer. + //! + size_t size() const + { + return mSize; + } + + //! + //! \brief Returns the size (in bytes) of the buffer. + //! + size_t nbBytes() const + { + return samplesCommon::getNbBytes(mType, size()); + } + + //! + //! \brief Resizes the buffer. This is a no-op if the new size is smaller than or equal to the current capacity. + //! + void resize(size_t newSize) + { + mSize = newSize; + if (mCapacity < newSize) + { + freeFn(mBuffer); + if (!allocFn(&mBuffer, this->nbBytes())) + { + throw std::bad_alloc{}; + } + mCapacity = newSize; + } + } + + //! + //! \brief Overload of resize that accepts Dims + //! + void resize(nvinfer1::Dims const& dims) + { + return this->resize(samplesCommon::volume(dims)); + } + + ~GenericBuffer() + { + freeFn(mBuffer); + } + +private: + size_t mSize{0}, mCapacity{0}; + nvinfer1::DataType mType; + void* mBuffer; + AllocFunc allocFn; + FreeFunc freeFn; +}; + +class DeviceAllocator +{ +public: + bool operator()(void** ptr, size_t size) const + { + return cudaMalloc(ptr, size) == cudaSuccess; + } +}; + +class DeviceFree +{ +public: + void operator()(void* ptr) const + { + cudaFree(ptr); + } +}; + +class HostAllocator +{ +public: + bool operator()(void** ptr, size_t size) const + { + *ptr = malloc(size); + return *ptr != nullptr; + } +}; + +class HostFree +{ +public: + void operator()(void* ptr) const + { + free(ptr); + } +}; + +using DeviceBuffer = GenericBuffer; +using HostBuffer = GenericBuffer; + +//! +//! \brief The ManagedBuffer class groups together a pair of corresponding device and host buffers. +//! +class ManagedBuffer +{ +public: + DeviceBuffer deviceBuffer; + HostBuffer hostBuffer; +}; + +//! +//! \brief The BufferManager class handles host and device buffer allocation and deallocation. +//! +//! \details This RAII class handles host and device buffer allocation and deallocation, +//! memcpy between host and device buffers to aid with inference, +//! and debugging dumps to validate inference. The BufferManager class is meant to be +//! used to simplify buffer management and any interactions between buffers and the engine. +//! +class BufferManager +{ +public: + static const size_t kINVALID_SIZE_VALUE = ~size_t(0); + + //! + //! \brief Create a BufferManager for handling buffer interactions with engine, when the I/O tensor volumes + //! are provided + //! + BufferManager( + std::shared_ptr engine, std::vector const& volumes, int32_t batchSize = 0) + : mEngine(engine) + , mBatchSize(batchSize) + { + // Create host and device buffers + for (int32_t i = 0; i < mEngine->getNbIOTensors(); i++) + { + auto const name = engine->getIOTensorName(i); + mNames[name] = i; + + nvinfer1::DataType type = mEngine->getTensorDataType(name); + + std::unique_ptr manBuf{new ManagedBuffer()}; + manBuf->deviceBuffer = DeviceBuffer(volumes[i], type); + manBuf->hostBuffer = HostBuffer(volumes[i], type); + void* deviceBuffer = manBuf->deviceBuffer.data(); + mDeviceBindings.emplace_back(deviceBuffer); + mManagedBuffers.emplace_back(std::move(manBuf)); + } + } + + //! + //! \brief Create a BufferManager for handling buffer interactions with engine. + //! + BufferManager(std::shared_ptr engine, int32_t const batchSize = 0, + nvinfer1::IExecutionContext const* context = nullptr) + : mEngine(engine) + , mBatchSize(batchSize) + { + // Create host and device buffers + for (int32_t i = 0, e = mEngine->getNbIOTensors(); i < e; i++) + { + auto const name = engine->getIOTensorName(i); + mNames[name] = i; + + auto dims = context ? context->getTensorShape(name) : mEngine->getTensorShape(name); + size_t vol = context || !mBatchSize ? 1 : static_cast(mBatchSize); + nvinfer1::DataType type = mEngine->getTensorDataType(name); + int32_t vecDim = mEngine->getTensorVectorizedDim(name); + if (-1 != vecDim) // i.e., 0 != lgScalarsPerVector + { + int32_t scalarsPerVec = mEngine->getTensorComponentsPerElement(name); + dims.d[vecDim] = divUp(dims.d[vecDim], scalarsPerVec); + vol *= scalarsPerVec; + } + vol *= samplesCommon::volume(dims); + std::unique_ptr manBuf{new ManagedBuffer()}; + manBuf->deviceBuffer = DeviceBuffer(vol, type); + manBuf->hostBuffer = HostBuffer(vol, type); + void* deviceBuffer = manBuf->deviceBuffer.data(); + mDeviceBindings.emplace_back(deviceBuffer); + mManagedBuffers.emplace_back(std::move(manBuf)); + } + } + + //! + //! \brief Returns a vector of device buffers that you can use directly as + //! bindings for the execute and enqueue methods of IExecutionContext. + //! + std::vector& getDeviceBindings() + { + return mDeviceBindings; + } + + //! + //! \brief Returns a vector of device buffers. + //! + std::vector const& getDeviceBindings() const + { + return mDeviceBindings; + } + + //! + //! \brief Returns the device buffer corresponding to tensorName. + //! Returns nullptr if no such tensor can be found. + //! + void* getDeviceBuffer(std::string const& tensorName) const + { + return getBuffer(false, tensorName); + } + + //! + //! \brief Returns the host buffer corresponding to tensorName. + //! Returns nullptr if no such tensor can be found. + //! + void* getHostBuffer(std::string const& tensorName) const + { + return getBuffer(true, tensorName); + } + + //! + //! \brief Returns the size of the host and device buffers that correspond to tensorName. + //! Returns kINVALID_SIZE_VALUE if no such tensor can be found. + //! + size_t size(std::string const& tensorName) const + { + auto record = mNames.find(tensorName); + if (record == mNames.end()) + return kINVALID_SIZE_VALUE; + return mManagedBuffers[record->second]->hostBuffer.nbBytes(); + } + + //! + //! \brief Copy the contents of input host buffers to input device buffers synchronously. + //! + void copyInputToDevice() + { + memcpyBuffers(true, false, false); + } + + //! + //! \brief Copy the contents of output device buffers to output host buffers synchronously. + //! + void copyOutputToHost() + { + memcpyBuffers(false, true, false); + } + + //! + //! \brief Copy the contents of input host buffers to input device buffers asynchronously. + //! + void copyInputToDeviceAsync(cudaStream_t const& stream = 0) + { + memcpyBuffers(true, false, true, stream); + } + + //! + //! \brief Copy the contents of output device buffers to output host buffers asynchronously. + //! + void copyOutputToHostAsync(cudaStream_t const& stream = 0) + { + memcpyBuffers(false, true, true, stream); + } + + ~BufferManager() = default; + +private: + void* getBuffer(bool const isHost, std::string const& tensorName) const + { + auto record = mNames.find(tensorName); + if (record == mNames.end()) + return nullptr; + return (isHost ? mManagedBuffers[record->second]->hostBuffer.data() + : mManagedBuffers[record->second]->deviceBuffer.data()); + } + + bool tensorIsInput(std::string const& tensorName) const + { + return mEngine->getTensorIOMode(tensorName.c_str()) == nvinfer1::TensorIOMode::kINPUT; + } + + void memcpyBuffers(bool const copyInput, bool const deviceToHost, bool const async, cudaStream_t const& stream = 0) + { + for (auto const& n : mNames) + { + void* dstPtr = deviceToHost ? mManagedBuffers[n.second]->hostBuffer.data() + : mManagedBuffers[n.second]->deviceBuffer.data(); + void const* srcPtr = deviceToHost ? mManagedBuffers[n.second]->deviceBuffer.data() + : mManagedBuffers[n.second]->hostBuffer.data(); + size_t const byteSize = mManagedBuffers[n.second]->hostBuffer.nbBytes(); + const cudaMemcpyKind memcpyType = deviceToHost ? cudaMemcpyDeviceToHost : cudaMemcpyHostToDevice; + if ((copyInput && tensorIsInput(n.first)) || (!copyInput && !tensorIsInput(n.first))) + { + if (async) + CHECK(cudaMemcpyAsync(dstPtr, srcPtr, byteSize, memcpyType, stream)); + else + CHECK(cudaMemcpy(dstPtr, srcPtr, byteSize, memcpyType)); + } + } + } + + std::shared_ptr mEngine; //!< The pointer to the engine + int mBatchSize; //!< The batch size for legacy networks, 0 otherwise. + std::vector> mManagedBuffers; //!< The vector of pointers to managed buffers + std::vector mDeviceBindings; //!< The vector of device buffers needed for engine execution + std::unordered_map mNames; //!< The map of tensor name and index pairs +}; + +} // namespace samplesCommon + +#endif // TENSORRT_BUFFERS_H diff --git a/samples/common/common.cpp b/samples/trtexecCommon/common.cpp similarity index 100% rename from samples/common/common.cpp rename to samples/trtexecCommon/common.cpp diff --git a/samples/trtexecCommon/common.h b/samples/trtexecCommon/common.h new file mode 100644 index 0000000000..7db46d14f9 --- /dev/null +++ b/samples/trtexecCommon/common.h @@ -0,0 +1,1065 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef TENSORRT_COMMON_H +#define TENSORRT_COMMON_H +#include "NvInfer.h" +#if !TRT_WINML +#include "NvInferPlugin.h" +#endif +#include "logger.h" +#include "sampleEntrypoints.h" +#include "utils/cacheUtils.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#ifdef _MSC_VER +// For loadLibrary +// Needed so that the max/min definitions in windows.h do not conflict with std::max/min. +#define NOMINMAX +#include +#undef NOMINMAX +#else +#include +#endif + +#ifdef _MSC_VER +#define FN_NAME __FUNCTION__ +#else +#define FN_NAME __func__ +#endif + +#if defined(__aarch64__) || defined(__QNX__) +#define ENABLE_DLA_API 1 +#endif + +using namespace nvinfer1; + +#define CHECK_RETURN_W_MSG(status, val, errMsg) \ + do \ + { \ + if (!(status)) \ + { \ + sample::gLogError << errMsg << " Error in " << __FILE__ << ", function " << FN_NAME << "(), line " \ + << __LINE__ << std::endl; \ + return val; \ + } \ + } while (0) + +#undef ASSERT +#define ASSERT(condition) \ + do \ + { \ + if (!(condition)) \ + { \ + sample::gLogError << "Assertion failure: " << #condition << std::endl; \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +#define CHECK_RETURN(status, val) CHECK_RETURN_W_MSG(status, val, "") + +#undef CHECK_WITH_STREAM +#define CHECK_WITH_STREAM(status, stream) \ + do \ + { \ + if ((status) != cudaSuccess) \ + { \ + stream << "Cuda failure at " << __FILE__ << ":" << __LINE__ << ": " << cudaGetErrorString(status) \ + << std::endl; \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +#undef CHECK +#define CHECK(status) CHECK_WITH_STREAM(status, std::cerr) + +constexpr long double operator"" _GiB(long double val) +{ + return val * (1 << 30); +} +constexpr long double operator"" _MiB(long double val) +{ + return val * (1 << 20); +} +constexpr long double operator"" _KiB(long double val) +{ + return val * (1 << 10); +} + +struct SimpleProfiler : public nvinfer1::IProfiler +{ + struct Record + { + float time{0}; + int count{0}; + }; + + void reportLayerTime(char const* layerName, float ms) noexcept override + { + mProfile[layerName].count++; + mProfile[layerName].time += ms; + if (std::find(mLayerNames.begin(), mLayerNames.end(), layerName) == mLayerNames.end()) + { + mLayerNames.push_back(layerName); + } + } + + SimpleProfiler(char const* name, std::vector const& srcProfilers = std::vector()) + : mName(name) + { + for (auto const& srcProfiler : srcProfilers) + { + for (auto const& rec : srcProfiler.mProfile) + { + auto it = mProfile.find(rec.first); + if (it == mProfile.end()) + { + mProfile.insert(rec); + } + else + { + it->second.time += rec.second.time; + it->second.count += rec.second.count; + } + } + } + } + + friend std::ostream& operator<<(std::ostream& out, SimpleProfiler const& value) + { + out << "========== " << value.mName << " profile ==========" << std::endl; + float totalTime = 0; + std::string layerNameStr = "TensorRT layer name"; + int maxLayerNameLength = std::max(static_cast(layerNameStr.size()), 70); + for (auto const& elem : value.mProfile) + { + totalTime += elem.second.time; + maxLayerNameLength = std::max(maxLayerNameLength, static_cast(elem.first.size())); + } + + auto old_settings = out.flags(); + auto old_precision = out.precision(); + // Output header + { + out << std::setfill(' ') << std::setw(maxLayerNameLength) << layerNameStr << " "; + out << std::setw(12) << "Runtime, " + << "%" + << " "; + out << std::setw(12) << "Invocations" + << " "; + out << std::setw(12) << "Runtime, ms" << std::endl; + } + for (size_t i = 0; i < value.mLayerNames.size(); i++) + { + const std::string layerName = value.mLayerNames[i]; + auto elem = value.mProfile.at(layerName); + out << std::setw(maxLayerNameLength) << layerName << " "; + out << std::setw(12) << std::fixed << std::setprecision(1) << (elem.time * 100.0F / totalTime) << "%" + << " "; + out << std::setw(12) << elem.count << " "; + out << std::setw(12) << std::fixed << std::setprecision(2) << elem.time << std::endl; + } + out.flags(old_settings); + out.precision(old_precision); + out << "========== " << value.mName << " total runtime = " << totalTime << " ms ==========" << std::endl; + + return out; + } + +private: + std::string mName; + std::vector mLayerNames; + std::map mProfile; +}; + +namespace samplesCommon +{ +using nvinfer1::utils::loadCacheFile; +using nvinfer1::utils::buildTimingCacheFromFile; +using nvinfer1::utils::saveCacheFile; +using nvinfer1::utils::updateTimingCacheFile; + +//! \brief Swaps endianness of an integral type. +template >> +[[nodiscard]] T swapEndianness(T value) +{ + uint8_t bytes[sizeof(T)]; + std::memcpy(bytes, &value, sizeof(T)); + std::reverse(std::begin(bytes), std::end(bytes)); + std::memcpy(&value, bytes, sizeof(T)); + return value; +} + +class HostMemory +{ +public: + HostMemory() = delete; + virtual void* data() const noexcept + { + return mData; + } + virtual std::size_t size() const noexcept + { + return mSize; + } + virtual nvinfer1::DataType type() const noexcept + { + return mType; + } + virtual ~HostMemory() {} + +protected: + HostMemory(std::size_t size, nvinfer1::DataType type) + : mData{nullptr} + , mSize(size) + , mType(type) + { + } + void* mData; + std::size_t mSize; + nvinfer1::DataType mType; +}; + +template +class TypedHostMemory : public HostMemory +{ +public: + explicit TypedHostMemory(std::size_t size) + : HostMemory(size, dataType) + { + mData = new ElemType[size]; + } + ~TypedHostMemory() noexcept override + { + delete[](ElemType*) mData; + } + ElemType* raw() noexcept + { + return static_cast(data()); + } +}; + +using FloatMemory = TypedHostMemory; +using HalfMemory = TypedHostMemory; +using ByteMemory = TypedHostMemory; + +inline void* safeCudaMalloc(size_t memSize) +{ + void* deviceMem; + CHECK(cudaMalloc(&deviceMem, memSize)); + if (deviceMem == nullptr) + { + std::cerr << "Out of memory" << std::endl; + exit(EXIT_FAILURE); + } + return deviceMem; +} + +inline bool isDebug() +{ + return std::getenv("TENSORRT_DEBUG") != nullptr; +} + +static auto StreamDeleter = [](cudaStream_t* pStream) { + if (pStream) + { + static_cast(cudaStreamDestroy(*pStream)); + delete pStream; + } +}; + +inline std::unique_ptr makeCudaStream() +{ + std::unique_ptr pStream(new cudaStream_t, StreamDeleter); + if (cudaStreamCreateWithFlags(pStream.get(), cudaStreamNonBlocking) != cudaSuccess) + { + pStream.reset(nullptr); + } + + return pStream; +} + +//! Return vector of indices that puts magnitudes of sequence in descending order. +template +std::vector argMagnitudeSort(Iter begin, Iter end) +{ + std::vector indices(end - begin); + std::iota(indices.begin(), indices.end(), 0); + std::ranges::sort(indices, std::greater<>{}, [&begin](size_t i) { return std::abs(begin[i]); }); + return indices; +} + +inline bool readReferenceFile(std::string const& fileName, std::vector& refVector) +{ + std::ifstream infile(fileName); + if (!infile.is_open()) + { + std::cout << "ERROR: readReferenceFile: Attempting to read from a file that is not open." << std::endl; + return false; + } + std::string line; + while (std::getline(infile, line)) + { + if (line.empty()) + continue; + refVector.push_back(line); + } + infile.close(); + return true; +} + +template +std::vector classify( + std::vector const& refVector, std::vector const& output, const size_t topK) +{ + auto const inds = samplesCommon::argMagnitudeSort(output.cbegin(), output.cend()); + std::vector result; + result.reserve(topK); + for (size_t k = 0; k < topK; ++k) + { + result.push_back(refVector[inds[k]]); + } + return result; +} + +// Returns indices of highest K magnitudes in v. +template +std::vector topKMagnitudes(std::vector const& v, const size_t k) +{ + std::vector indices = samplesCommon::argMagnitudeSort(v.cbegin(), v.cend()); + indices.resize(k); + return indices; +} + +template +bool readASCIIFile(std::string const& fileName, const size_t size, std::vector& out) +{ + std::ifstream infile(fileName); + if (!infile.is_open()) + { + std::cout << "ERROR readASCIIFile: Attempting to read from a file that is not open." << std::endl; + return false; + } + out.clear(); + out.reserve(size); + out.assign(std::istream_iterator(infile), std::istream_iterator()); + infile.close(); + return true; +} + +template +bool writeASCIIFile(std::string const& fileName, std::vector const& in) +{ + std::ofstream outfile(fileName); + if (!outfile.is_open()) + { + std::cout << "ERROR: writeASCIIFile: Attempting to write to a file that is not open." << std::endl; + return false; + } + for (auto fn : in) + { + outfile << fn << "\n"; + } + outfile.close(); + return true; +} + +inline void print_version() +{ + std::cout << " TensorRT version: " << NV_TENSORRT_MAJOR << "." << NV_TENSORRT_MINOR << "." << NV_TENSORRT_PATCH + << "." << NV_TENSORRT_BUILD << std::endl; +} + +inline std::string getFileType(std::string const& filepath) +{ + return filepath.substr(filepath.find_last_of(".") + 1); +} + +inline std::string toLower(std::string const& inp) +{ + std::string out = inp; + std::transform(out.begin(), out.end(), out.begin(), ::tolower); + return out; +} + +inline float getMaxValue(float const* buffer, int64_t size) +{ + assert(buffer != nullptr); + assert(size > 0); + return *std::max_element(buffer, buffer + size); +} + +#if !TRT_WINML && ENABLE_FEATURE_WEAK_TYPING +inline void setAllDynamicRanges(nvinfer1::INetworkDefinition* network, float inRange = 2.0F, float outRange = 4.0F) +{ + for (int i = 0; i < network->getNbLayers(); i++) + { + auto layer = network->getLayer(i); + for (int j = 0; j < layer->getNbInputs(); j++) + { + nvinfer1::ITensor* input{layer->getInput(j)}; + if (input != nullptr && !input->dynamicRangeIsSet()) + { + ASSERT(input->setDynamicRange(-inRange, inRange)); + } + } + } + + for (int i = 0; i < network->getNbLayers(); i++) + { + auto layer = network->getLayer(i); + for (int j = 0; j < layer->getNbOutputs(); j++) + { + nvinfer1::ITensor* output{layer->getOutput(j)}; + if (output != nullptr && !output->dynamicRangeIsSet()) + { + if (layer->getType() == nvinfer1::LayerType::kPOOLING) + { + ASSERT(output->setDynamicRange(-inRange, inRange)); + } + else + { + ASSERT(output->setDynamicRange(-outRange, outRange)); + } + } + } + } +} + +inline void setDummyInt8DynamicRanges(nvinfer1::IBuilderConfig const* c, nvinfer1::INetworkDefinition* n) +{ + if (c->getFlag(nvinfer1::BuilderFlag::kINT8)) + { + sample::gLogWarning << "No per-tensor dynamic range provided. Generating dummy values. Int8 accuracy " + "is not guaranteed." + << std::endl; + setAllDynamicRanges(n); + } +} +#endif // !TRT_WINML && ENABLE_FEATURE_WEAK_TYPING + +inline void enableDLA( + nvinfer1::IBuilder* builder, nvinfer1::IBuilderConfig* config, int useDLACore, bool allowGPUFallback = true) +{ + if (useDLACore >= 0) + { + if (builder->getNbDLACores() == 0) + { + std::cerr << "Trying to use DLA core " << useDLACore << " on a platform that doesn't have any DLA cores" + << std::endl; + assert("Error: use DLA core on a platfrom that doesn't have any DLA cores" && false); + } + if (allowGPUFallback) + { + config->setFlag(nvinfer1::BuilderFlag::kGPU_FALLBACK); + } +#if ENABLE_FEATURE_WEAK_TYPING + if (!config->getFlag(nvinfer1::BuilderFlag::kINT8)) + { + config->setFlag(nvinfer1::BuilderFlag::kFP16); + } +#endif // ENABLE_FEATURE_WEAK_TYPING + config->setDefaultDeviceType(nvinfer1::DeviceType::kDLA); + config->setDLACore(useDLACore); + } +} + +//! \brief Matches a flag prefix in an argument, ignoring leading spaces. +//! \param arg The command-line argument to check. +//! \param flag The flag prefix to match (e.g., "--loadEngine="). +//! \return A string_view of the remainder after \p flag, or nullopt if \p flag isn't found. +[[nodiscard]] std::optional matchFlag(std::string_view arg, std::string_view flag); + +//! \overload std::optional matchFlag(std::string_view arg, std::string_view flag) to prevent +//! accidental use of `std::string&&` arguments which would produce a dangling view, but allow e.g., `char const*`. +template +[[nodiscard]] std::optional matchFlag(StringViewable&& arg, std::string_view flag) +{ + static_assert(!std::is_rvalue_reference_v, + "You don't want the above matchFlag with `std::string&&` arguments which would produce a dangling view."); + return matchFlag(std::string_view{arg}, flag); +} + +int32_t parseDLA(int32_t argc, char** argv); + +inline size_t getNbBytes(nvinfer1::DataType t, int64_t vol) noexcept +{ + switch (t) + { + case nvinfer1::DataType::kINT64: return 8 * vol; + case nvinfer1::DataType::kINT32: + case nvinfer1::DataType::kFLOAT: return 4 * vol; + case nvinfer1::DataType::kBF16: + case nvinfer1::DataType::kHALF: return 2 * vol; + case nvinfer1::DataType::kBOOL: + case nvinfer1::DataType::kUINT8: + case nvinfer1::DataType::kINT8: return vol; + case nvinfer1::DataType::kFP8: +#if CUDA_VERSION < 11060 + ASSERT(false && "FP8 is not supported"); +#else + return vol; +#endif + case nvinfer1::DataType::kE8M0: +#if CUDA_VERSION < 12080 + ASSERT(false && "E8M0 is not supported"); +#else + return vol; +#endif // CUDA_VERSION < 12080 + case nvinfer1::DataType::kINT4: + case nvinfer1::DataType::kFP4: return (vol + 1) / 2; + } + ASSERT(false && "Unknown element type"); +} + +// Return least integer no less than exact value of m/n. +template +inline auto divUp(A m, B n) -> std::enable_if_t::value && std::is_integral::value, A> +{ + ASSERT(n > 0); + return (m + n - 1) / n; +} + +inline int64_t volume(nvinfer1::Dims const& d) +{ + return std::accumulate(d.d, d.d + d.nbDims, int64_t{1}, std::multiplies{}); +} + +inline int64_t volume(nvinfer1::Dims const& dims, int32_t start, int32_t stop) +{ + ASSERT(start >= 0); + ASSERT(start <= stop); + ASSERT(stop <= dims.nbDims); + ASSERT(std::all_of(dims.d + start, dims.d + stop, [](int32_t x) { return x >= 0; })); + return std::accumulate(dims.d + start, dims.d + stop, int64_t{1}, std::multiplies{}); +} + +//! Locate path to file, given its filename or filepath suffix and possible dirs it might lie in. +//! Function will also walk back MAX_DEPTH dirs from CWD to check for such a file path. +inline std::string locateFile( + std::string const& filepathSuffix, std::vector const& directories, bool reportError = true) +{ + int const MAX_DEPTH{10}; + bool found{false}; + std::string filepath; + + for (auto& dir : directories) + { + if (!dir.empty() && dir.back() != '/') + { +#ifdef _MSC_VER + filepath = dir + "\\" + filepathSuffix; +#else + filepath = dir + "/" + filepathSuffix; +#endif + } + else + { + filepath = dir + filepathSuffix; + } + + for (int i = 0; i < MAX_DEPTH && !found; i++) + { + const std::ifstream checkFile(filepath); + found = checkFile.is_open(); + if (found) + { + break; + } + + filepath = "../" + filepath; // Try again in parent dir + } + + if (found) + { + break; + } + + filepath.clear(); + } + + // Could not find the file + if (filepath.empty()) + { + const std::string dirList = std::accumulate(directories.begin() + 1, directories.end(), directories.front(), + [](std::string const& a, std::string const& b) { return a + "\n\t" + b; }); + std::cout << "Could not find " << filepathSuffix << " in data directories:\n\t" << dirList << std::endl; + + if (reportError) + { + std::cout << "&&&& FAILED" << std::endl; + exit(EXIT_FAILURE); + } + } + + return filepath; +} + +inline void readPGMFile(std::string const& fileName, uint8_t* buffer, int32_t inH, int32_t inW) +{ + std::ifstream infile(fileName, std::ifstream::binary); + ASSERT(infile.is_open() && "Attempting to read from a file that is not open."); + std::string magic, w, h, max; + infile >> magic >> w >> h >> max; + infile.seekg(1, infile.cur); + infile.read(reinterpret_cast(buffer), inH * inW); +} +template +struct PPM +{ + std::string magic, fileName; + int h, w, max; + uint8_t buffer[C * H * W]; +}; + +// New vPPM(variable sized PPM) class with variable dimensions. +struct vPPM +{ + std::string magic, fileName; + int h, w, max; + std::vector buffer; +}; + +struct BBox +{ + float x1, y1, x2, y2; +}; + +template +void readPPMFile(std::string const& filename, samplesCommon::PPM& ppm) +{ + ppm.fileName = filename; + std::ifstream infile(filename, std::ifstream::binary); + assert(infile.is_open() && "Attempting to read from a file that is not open."); + infile >> ppm.magic >> ppm.w >> ppm.h >> ppm.max; + infile.seekg(1, infile.cur); + infile.read(reinterpret_cast(ppm.buffer), ppm.w * ppm.h * 3); +} + +inline void readPPMFile(std::string const& filename, vPPM& ppm, std::vector& input_dir) +{ + ppm.fileName = filename; + std::ifstream infile(locateFile(filename, input_dir), std::ifstream::binary); + infile >> ppm.magic >> ppm.w >> ppm.h >> ppm.max; + infile.seekg(1, infile.cur); + + for (int i = 0; i < ppm.w * ppm.h * 3; ++i) + { + ppm.buffer.push_back(0); + } + + infile.read(reinterpret_cast(&ppm.buffer[0]), ppm.w * ppm.h * 3); +} + +template +void writePPMFileWithBBox(std::string const& filename, PPM& ppm, BBox const& bbox) +{ + std::ofstream outfile("./" + filename, std::ofstream::binary); + assert(!outfile.fail()); + outfile << "P6" + << "\n" + << ppm.w << " " << ppm.h << "\n" + << ppm.max << "\n"; + + auto round = [](float x) -> int { return int(std::floor(x + 0.5F)); }; + int const x1 = std::min(std::max(0, round(int(bbox.x1))), W - 1); + int const x2 = std::min(std::max(0, round(int(bbox.x2))), W - 1); + int const y1 = std::min(std::max(0, round(int(bbox.y1))), H - 1); + int const y2 = std::min(std::max(0, round(int(bbox.y2))), H - 1); + + for (int x = x1; x <= x2; ++x) + { + // bbox top border + ppm.buffer[(y1 * ppm.w + x) * 3] = 255; + ppm.buffer[(y1 * ppm.w + x) * 3 + 1] = 0; + ppm.buffer[(y1 * ppm.w + x) * 3 + 2] = 0; + // bbox bottom border + ppm.buffer[(y2 * ppm.w + x) * 3] = 255; + ppm.buffer[(y2 * ppm.w + x) * 3 + 1] = 0; + ppm.buffer[(y2 * ppm.w + x) * 3 + 2] = 0; + } + + for (int y = y1; y <= y2; ++y) + { + // bbox left border + ppm.buffer[(y * ppm.w + x1) * 3] = 255; + ppm.buffer[(y * ppm.w + x1) * 3 + 1] = 0; + ppm.buffer[(y * ppm.w + x1) * 3 + 2] = 0; + // bbox right border + ppm.buffer[(y * ppm.w + x2) * 3] = 255; + ppm.buffer[(y * ppm.w + x2) * 3 + 1] = 0; + ppm.buffer[(y * ppm.w + x2) * 3 + 2] = 0; + } + + outfile.write(reinterpret_cast(ppm.buffer), ppm.w * ppm.h * 3); +} + +inline void writePPMFileWithBBox(std::string const& filename, vPPM ppm, std::vector& dets) +{ + std::ofstream outfile("./" + filename, std::ofstream::binary); + assert(!outfile.fail()); + outfile << "P6" + << "\n" + << ppm.w << " " << ppm.h << "\n" + << ppm.max << "\n"; + auto round = [](float x) -> int { return int(std::floor(x + 0.5F)); }; + + for (auto bbox : dets) + { + for (int x = int(bbox.x1); x < int(bbox.x2); ++x) + { + // bbox top border + ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3] = 255; + ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3 + 1] = 0; + ppm.buffer[(round(bbox.y1) * ppm.w + x) * 3 + 2] = 0; + // bbox bottom border + ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3] = 255; + ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3 + 1] = 0; + ppm.buffer[(round(bbox.y2) * ppm.w + x) * 3 + 2] = 0; + } + + for (int y = int(bbox.y1); y < int(bbox.y2); ++y) + { + // bbox left border + ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3] = 255; + ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3 + 1] = 0; + ppm.buffer[(y * ppm.w + round(bbox.x1)) * 3 + 2] = 0; + // bbox right border + ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3] = 255; + ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3 + 1] = 0; + ppm.buffer[(y * ppm.w + round(bbox.x2)) * 3 + 2] = 0; + } + } + + outfile.write(reinterpret_cast(&ppm.buffer[0]), ppm.w * ppm.h * 3); +} + +class TimerBase +{ +public: + virtual void start() {} + virtual void stop() {} + float microseconds() const noexcept + { + return mMs * 1000.F; + } + float milliseconds() const noexcept + { + return mMs; + } + float seconds() const noexcept + { + return mMs / 1000.F; + } + void reset() noexcept + { + mMs = 0.F; + } + +protected: + float mMs{0.0F}; +}; + +class GpuTimer : public TimerBase +{ +public: + explicit GpuTimer(cudaStream_t stream) + : mStream(stream) + { + CHECK(cudaEventCreate(&mStart)); + CHECK(cudaEventCreate(&mStop)); + } + ~GpuTimer() + { + CHECK(cudaEventDestroy(mStart)); + CHECK(cudaEventDestroy(mStop)); + } + void start() override + { + CHECK(cudaEventRecord(mStart, mStream)); + } + void stop() override + { + CHECK(cudaEventRecord(mStop, mStream)); + float ms{0.0F}; + CHECK(cudaEventSynchronize(mStop)); + CHECK(cudaEventElapsedTime(&ms, mStart, mStop)); + mMs += ms; + } + +private: + cudaEvent_t mStart, mStop; + cudaStream_t mStream; +}; // class GpuTimer + +template +class CpuTimer : public TimerBase +{ +public: + using clock_type = Clock; + + void start() override + { + mStart = Clock::now(); + } + void stop() override + { + mStop = Clock::now(); + mMs += std::chrono::duration{mStop - mStart}.count(); + } + +private: + std::chrono::time_point mStart, mStop; +}; // class CpuTimer + +using PreciseCpuTimer = CpuTimer; + +inline std::vector splitString(std::string str, char delimiter = ',') +{ + std::vector splitVect; + std::stringstream ss(str); + std::string substr; + + while (ss.good()) + { + getline(ss, substr, delimiter); + splitVect.emplace_back(std::move(substr)); + } + return splitVect; +} + +inline int getC(nvinfer1::Dims const& d) +{ + return d.nbDims >= 3 ? d.d[d.nbDims - 3] : 1; +} + +inline int getH(nvinfer1::Dims const& d) +{ + return d.nbDims >= 2 ? d.d[d.nbDims - 2] : 1; +} + +inline int getW(nvinfer1::Dims const& d) +{ + return d.nbDims >= 1 ? d.d[d.nbDims - 1] : 1; +} + +//! Platform-agnostic wrapper around dynamic libraries. +class DynamicLibrary +{ +public: + explicit DynamicLibrary(std::string name) + : mLibName{std::move(name)} + { +#if defined(_WIN32) + mHandle = LoadLibraryA(mLibName.c_str()); +#else // defined(_WIN32) + int32_t flags{RTLD_LAZY}; +#if ENABLE_ASAN + // https://github.com/google/sanitizers/issues/89 + // asan doesn't handle module unloading correctly and there are no plans on doing + // so. In order to get proper stack traces, don't delete the shared library on + // close so that asan can resolve the symbols correctly. + flags |= RTLD_NODELETE; +#endif // ENABLE_ASAN + + mHandle = dlopen(mLibName.c_str(), flags); +#endif // defined(_WIN32) + + if (mHandle == nullptr) + { + std::string errorStr{}; +#if !defined(_WIN32) + errorStr = std::string{" due to "} + std::string{dlerror()}; +#endif + throw std::runtime_error("Unable to open library: " + mLibName + errorStr); + } + } + + DynamicLibrary(DynamicLibrary const&) = delete; + DynamicLibrary(DynamicLibrary const&&) = delete; + + //! \return a pointer to a symbol from the DynamicLibrary with the given function signature. + //! + //! \throw std::invalid_argument if loading the symbol failed. + //! + //! \note Type checking is not possible, so if `Signature` is incorrect, the behavior is undefined. + + template + [[nodiscard]] Signature& symbolAddress(char const* name) const + { + static_assert(std::is_function_v, "Signature must be a function type."); + if (mHandle == nullptr) + { + throw std::runtime_error("Handle to library is nullptr."); + } + void* const ret = +#if defined(_MSC_VER) + static_cast(GetProcAddress(static_cast(mHandle), name)); +#else + dlsym(mHandle, name); +#endif + if (ret == nullptr) + { + throw std::invalid_argument(mLibName + ": error loading symbol: " + std::string(name)); + } + return *reinterpret_cast(ret); + } + + ~DynamicLibrary() + { + try + { +#if defined(_WIN32) + ASSERT(static_cast(FreeLibrary(static_cast(mHandle)))); +#else + ASSERT(dlclose(mHandle) == 0); +#endif + } + catch (...) + { + sample::gLogError << "Unable to close library: " << mLibName << std::endl; + } + } + +private: + std::string mLibName{}; //!< Name of the DynamicLibrary + void* mHandle{}; //!< Handle to the DynamicLibrary +}; + +[[nodiscard]] inline std::unique_ptr loadLibrary(std::string name) +{ + return std::make_unique(std::move(name)); +} + +//! Represents the compute capability of a device. +//! This pertains to virtual architectures represented by the intermediate PTX format. +//! This is distinct from the SM version. +//! See https://forums.developer.nvidia.com/t/how-should-i-use-correctly-the-sm-xx-and-compute-xx/219160 +struct ComputeCapability +{ + int32_t major{}; + int32_t minor{}; + + //! \return the compute capability of the CUDA device with the given \p deviceIndex. + [[nodiscard]] static ComputeCapability forDevice(int32_t deviceIndex) + { + int32_t major{0}; + int32_t minor{0}; + CHECK(cudaDeviceGetAttribute(&major, cudaDevAttrComputeCapabilityMajor, deviceIndex)); + CHECK(cudaDeviceGetAttribute(&minor, cudaDevAttrComputeCapabilityMinor, deviceIndex)); + // Redirect 12.1 to 12.0 to since dependencies do not support 12.1 yet and 12.1 can reuse 12.0 cubins to save + // lib size/compile time.. + if (major == 12 && minor == 1) + { + minor = 0; + } + return {major, minor}; + } +}; + +inline int32_t getSmVersion() +{ + int32_t deviceIndex = 0; + CHECK(cudaGetDevice(&deviceIndex)); + + auto const cc = ComputeCapability::forDevice(deviceIndex); + return ((cc.major << 8) | cc.minor); +} + +inline bool isSmSafe() +{ + int32_t const smVersion = getSmVersion(); + return smVersion == 0x0705 || smVersion == 0x0800 || smVersion == 0x0806 || smVersion == 0x0807 + || smVersion == 0x0A00 || smVersion == 0x0B00; +} + +inline int32_t getMaxPersistentCacheSize() +{ + int32_t deviceIndex{}; + CHECK(cudaGetDevice(&deviceIndex)); + + int32_t maxPersistentL2CacheSize{}; +#if CUDART_VERSION >= 11030 && !TRT_WINML + CHECK(cudaDeviceGetAttribute(&maxPersistentL2CacheSize, cudaDevAttrMaxPersistingL2CacheSize, deviceIndex)); +#endif + + return maxPersistentL2CacheSize; +} + +} // namespace samplesCommon + +inline std::ostream& operator<<(std::ostream& os, nvinfer1::Dims const& dims) +{ + os << "("; + for (int i = 0; i < dims.nbDims; ++i) + { + os << (i ? ", " : "") << dims.d[i]; + } + return os << ")"; +} + +[[nodiscard]] inline std::string genFilenameSafeString(std::string_view s) +{ + std::string_view const kALLOWED{"._-,"}; + constexpr size_t kMAX_FILENAME_LENGTH = 150; // Leave some margin due to Windows path length limitation + constexpr size_t kELLIPSIS_LENGTH = 3; // Length of "..." + + auto processChar = [&kALLOWED](char c) { + return std::isalnum(static_cast(c)) || kALLOWED.find(c) != std::string_view::npos ? c : '_'; + }; + + std::string res; + if (s.length() <= kMAX_FILENAME_LENGTH) + { + res.reserve(s.size()); + std::transform(s.begin(), s.end(), std::back_inserter(res), processChar); + return res; + } + + res.reserve(kMAX_FILENAME_LENGTH); + size_t const halfLength = (kMAX_FILENAME_LENGTH - kELLIPSIS_LENGTH) / 2; + + std::transform(s.begin(), s.begin() + halfLength, std::back_inserter(res), processChar); + res += "..."; + std::transform(s.end() - halfLength, s.end(), std::back_inserter(res), processChar); + + return res; +} + +#endif // TENSORRT_COMMON_H diff --git a/samples/common/debugTensorWriter.cpp b/samples/trtexecCommon/debugTensorWriter.cpp similarity index 99% rename from samples/common/debugTensorWriter.cpp rename to samples/trtexecCommon/debugTensorWriter.cpp index f1417f328f..c628274a79 100644 --- a/samples/common/debugTensorWriter.cpp +++ b/samples/trtexecCommon/debugTensorWriter.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/debugTensorWriter.h b/samples/trtexecCommon/debugTensorWriter.h similarity index 97% rename from samples/common/debugTensorWriter.h rename to samples/trtexecCommon/debugTensorWriter.h index 4123216d9c..181120178f 100644 --- a/samples/common/debugTensorWriter.h +++ b/samples/trtexecCommon/debugTensorWriter.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/globalTimerKernel.cu b/samples/trtexecCommon/globalTimerKernel.cu similarity index 100% rename from samples/common/globalTimerKernel.cu rename to samples/trtexecCommon/globalTimerKernel.cu diff --git a/samples/common/globalTimerKernel.h b/samples/trtexecCommon/globalTimerKernel.h similarity index 100% rename from samples/common/globalTimerKernel.h rename to samples/trtexecCommon/globalTimerKernel.h diff --git a/samples/trtexecCommon/half.h b/samples/trtexecCommon/half.h new file mode 100644 index 0000000000..2977c6f36b --- /dev/null +++ b/samples/trtexecCommon/half.h @@ -0,0 +1,4307 @@ +// half - IEEE 754-based half-precision floating point library. +// +// Copyright (c) 2012-2017 Christian Rau +// +// Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated +// documentation files (the "Software"), to deal in the Software without restriction, including without limitation the +// rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to +// permit persons to whom the Software is furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all copies or substantial portions of the +// Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE +// WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR +// COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR +// OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// Version 1.12.0 +// +// Note: this library defaults to HALF_ROUND_TIES_TO_EVEN=0 (ties away from zero), which differs +// from IEEE 754 round-to-nearest-even. Define HALF_ROUND_TIES_TO_EVEN=1 before including this +// header to opt into IEEE-conformant rounding. parsers/onnxOpenSource/half.h does this. + +/// \file +/// Main header file for half precision functionality. + +#ifndef HALF_HALF_HPP +#define HALF_HALF_HPP + +/// Combined gcc version number. +#define HALF_GNUC_VERSION (__GNUC__ * 100 + __GNUC_MINOR__) + +// check C++11 language features +#if defined(__clang__) // clang +#if __has_feature(cxx_static_assert) && !defined(HALF_ENABLE_CPP11_STATIC_ASSERT) +#define HALF_ENABLE_CPP11_STATIC_ASSERT 1 +#endif +#if __has_feature(cxx_constexpr) && !defined(HALF_ENABLE_CPP11_CONSTEXPR) +#define HALF_ENABLE_CPP11_CONSTEXPR 1 +#endif +#if __has_feature(cxx_noexcept) && !defined(HALF_ENABLE_CPP11_NOEXCEPT) +#define HALF_ENABLE_CPP11_NOEXCEPT 1 +#endif +#if __has_feature(cxx_user_literals) && !defined(HALF_ENABLE_CPP11_USER_LITERALS) +#define HALF_ENABLE_CPP11_USER_LITERALS 1 +#endif +#if (defined(__GXX_EXPERIMENTAL_CXX0X__) || __cplusplus >= 201103L) && !defined(HALF_ENABLE_CPP11_LONG_LONG) +#define HALF_ENABLE_CPP11_LONG_LONG 1 +#endif +/*#elif defined(__INTEL_COMPILER) //Intel C++ + #if __INTEL_COMPILER >= 1100 && !defined(HALF_ENABLE_CPP11_STATIC_ASSERT) ???????? + #define HALF_ENABLE_CPP11_STATIC_ASSERT 1 + #endif + #if __INTEL_COMPILER >= 1300 && !defined(HALF_ENABLE_CPP11_CONSTEXPR) ???????? + #define HALF_ENABLE_CPP11_CONSTEXPR 1 + #endif + #if __INTEL_COMPILER >= 1300 && !defined(HALF_ENABLE_CPP11_NOEXCEPT) ???????? + #define HALF_ENABLE_CPP11_NOEXCEPT 1 + #endif + #if __INTEL_COMPILER >= 1100 && !defined(HALF_ENABLE_CPP11_LONG_LONG) ???????? + #define HALF_ENABLE_CPP11_LONG_LONG 1 + #endif*/ +#elif defined(__GNUC__) // gcc +#if defined(__GXX_EXPERIMENTAL_CXX0X__) || __cplusplus >= 201103L +#if HALF_GNUC_VERSION >= 403 && !defined(HALF_ENABLE_CPP11_STATIC_ASSERT) +#define HALF_ENABLE_CPP11_STATIC_ASSERT 1 +#endif +#if HALF_GNUC_VERSION >= 406 && !defined(HALF_ENABLE_CPP11_CONSTEXPR) +#define HALF_ENABLE_CPP11_CONSTEXPR 1 +#endif +#if HALF_GNUC_VERSION >= 406 && !defined(HALF_ENABLE_CPP11_NOEXCEPT) +#define HALF_ENABLE_CPP11_NOEXCEPT 1 +#endif +#if HALF_GNUC_VERSION >= 407 && !defined(HALF_ENABLE_CPP11_USER_LITERALS) +#define HALF_ENABLE_CPP11_USER_LITERALS 1 +#endif +#if !defined(HALF_ENABLE_CPP11_LONG_LONG) +#define HALF_ENABLE_CPP11_LONG_LONG 1 +#endif +#endif +#elif defined(_MSC_VER) // Visual C++ +#if _MSC_VER >= 1900 && !defined(HALF_ENABLE_CPP11_CONSTEXPR) +#define HALF_ENABLE_CPP11_CONSTEXPR 1 +#endif +#if _MSC_VER >= 1900 && !defined(HALF_ENABLE_CPP11_NOEXCEPT) +#define HALF_ENABLE_CPP11_NOEXCEPT 1 +#endif +#if _MSC_VER >= 1900 && !defined(HALF_ENABLE_CPP11_USER_LITERALS) +#define HALF_ENABLE_CPP11_USER_LITERALS 1 +#endif +#if _MSC_VER >= 1600 && !defined(HALF_ENABLE_CPP11_STATIC_ASSERT) +#define HALF_ENABLE_CPP11_STATIC_ASSERT 1 +#endif +#if _MSC_VER >= 1310 && !defined(HALF_ENABLE_CPP11_LONG_LONG) +#define HALF_ENABLE_CPP11_LONG_LONG 1 +#endif +#define HALF_POP_WARNINGS 1 +#pragma warning(push) +#pragma warning(disable : 4099 4127 4146) // struct vs class, constant in if, negative unsigned +#endif + +// check C++11 library features +#include +#if defined(_LIBCPP_VERSION) // libc++ +#if defined(__GXX_EXPERIMENTAL_CXX0X__) || __cplusplus >= 201103 +#ifndef HALF_ENABLE_CPP11_TYPE_TRAITS +#define HALF_ENABLE_CPP11_TYPE_TRAITS 1 +#endif +#ifndef HALF_ENABLE_CPP11_CSTDINT +#define HALF_ENABLE_CPP11_CSTDINT 1 +#endif +#ifndef HALF_ENABLE_CPP11_CMATH +#define HALF_ENABLE_CPP11_CMATH 1 +#endif +#ifndef HALF_ENABLE_CPP11_HASH +#define HALF_ENABLE_CPP11_HASH 1 +#endif +#endif +#elif defined(__GLIBCXX__) // libstdc++ +#if defined(__GXX_EXPERIMENTAL_CXX0X__) || __cplusplus >= 201103 +#ifdef __clang__ +#if __GLIBCXX__ >= 20080606 && !defined(HALF_ENABLE_CPP11_TYPE_TRAITS) +#define HALF_ENABLE_CPP11_TYPE_TRAITS 1 +#endif +#if __GLIBCXX__ >= 20080606 && !defined(HALF_ENABLE_CPP11_CSTDINT) +#define HALF_ENABLE_CPP11_CSTDINT 1 +#endif +#if __GLIBCXX__ >= 20080606 && !defined(HALF_ENABLE_CPP11_CMATH) +#define HALF_ENABLE_CPP11_CMATH 1 +#endif +#if __GLIBCXX__ >= 20080606 && !defined(HALF_ENABLE_CPP11_HASH) +#define HALF_ENABLE_CPP11_HASH 1 +#endif +#else +#if HALF_GNUC_VERSION >= 403 && !defined(HALF_ENABLE_CPP11_CSTDINT) +#define HALF_ENABLE_CPP11_CSTDINT 1 +#endif +#if HALF_GNUC_VERSION >= 403 && !defined(HALF_ENABLE_CPP11_CMATH) +#define HALF_ENABLE_CPP11_CMATH 1 +#endif +#if HALF_GNUC_VERSION >= 403 && !defined(HALF_ENABLE_CPP11_HASH) +#define HALF_ENABLE_CPP11_HASH 1 +#endif +#endif +#endif +#elif defined(_CPPLIB_VER) // Dinkumware/Visual C++ +#if _CPPLIB_VER >= 520 +#ifndef HALF_ENABLE_CPP11_TYPE_TRAITS +#define HALF_ENABLE_CPP11_TYPE_TRAITS 1 +#endif +#ifndef HALF_ENABLE_CPP11_CSTDINT +#define HALF_ENABLE_CPP11_CSTDINT 1 +#endif +#ifndef HALF_ENABLE_CPP11_HASH +#define HALF_ENABLE_CPP11_HASH 1 +#endif +#endif +#if _CPPLIB_VER >= 610 +#ifndef HALF_ENABLE_CPP11_CMATH +#define HALF_ENABLE_CPP11_CMATH 1 +#endif +#endif +#endif +#undef HALF_GNUC_VERSION + +// support constexpr +#if HALF_ENABLE_CPP11_CONSTEXPR +#define HALF_CONSTEXPR constexpr +#define HALF_CONSTEXPR_CONST constexpr +#else +#define HALF_CONSTEXPR +#define HALF_CONSTEXPR_CONST const +#endif + +// support noexcept +#if HALF_ENABLE_CPP11_NOEXCEPT +#define HALF_NOEXCEPT noexcept +#define HALF_NOTHROW noexcept +#else +#define HALF_NOEXCEPT +#define HALF_NOTHROW throw() +#endif + +#include +#include +#include +#include +#include +#include +#if HALF_ENABLE_CPP11_TYPE_TRAITS +#include +#endif +#if HALF_ENABLE_CPP11_CSTDINT +#include +#endif +#if HALF_ENABLE_CPP11_HASH +#include +#endif + +/// Default rounding mode. +/// This specifies the rounding mode used for all conversions between [half](\ref half_float::half)s and `float`s as +/// well as for the half_cast() if not specifying a rounding mode explicitly. It can be redefined (before including +/// half.hpp) to one of the standard rounding modes using their respective constants or the equivalent values of +/// `std::float_round_style`: +/// +/// `std::float_round_style` | value | rounding +/// ---------------------------------|-------|------------------------- +/// `std::round_indeterminate` | -1 | fastest (default) +/// `std::round_toward_zero` | 0 | toward zero +/// `std::round_to_nearest` | 1 | to nearest +/// `std::round_toward_infinity` | 2 | toward positive infinity +/// `std::round_toward_neg_infinity` | 3 | toward negative infinity +/// +/// By default this is set to `-1` (`std::round_indeterminate`), which uses truncation (round toward zero, but with +/// overflows set to infinity) and is the fastest rounding mode possible. It can even be set to +/// `std::numeric_limits::round_style` to synchronize the rounding mode with that of the underlying +/// single-precision implementation. +#ifndef HALF_ROUND_STYLE +#define HALF_ROUND_STYLE 1 // = std::round_to_nearest +#endif + +/// Tie-breaking behaviour for round to nearest. +/// This specifies if ties in round to nearest should be resolved by rounding to the nearest even value. By default this +/// is defined to `0` resulting in the faster but slightly more biased behaviour of rounding away from zero in half-way +/// cases (and thus equal to the round() function), but can be redefined to `1` (before including half.hpp) if more +/// IEEE-conformant behaviour is needed. +#ifndef HALF_ROUND_TIES_TO_EVEN +#define HALF_ROUND_TIES_TO_EVEN 0 // ties away from zero +#endif + +/// Value signaling overflow. +/// In correspondence with `HUGE_VAL[F|L]` from `` this symbol expands to a positive value signaling the overflow +/// of an operation, in particular it just evaluates to positive infinity. +#define HUGE_VALH std::numeric_limits::infinity() + +/// Fast half-precision fma function. +/// This symbol is only defined if the fma() function generally executes as fast as, or faster than, a separate +/// half-precision multiplication followed by an addition. Due to the internal single-precision implementation of all +/// arithmetic operations, this is in fact always the case. +#define FP_FAST_FMAH 1 + +#ifndef FP_ILOGB0 +#define FP_ILOGB0 INT_MIN +#endif +#ifndef FP_ILOGBNAN +#define FP_ILOGBNAN INT_MAX +#endif +#ifndef FP_SUBNORMAL +#define FP_SUBNORMAL 0 +#endif +#ifndef FP_ZERO +#define FP_ZERO 1 +#endif +#ifndef FP_NAN +#define FP_NAN 2 +#endif +#ifndef FP_INFINITE +#define FP_INFINITE 3 +#endif +#ifndef FP_NORMAL +#define FP_NORMAL 4 +#endif + +/// Main namespace for half precision functionality. +/// This namespace contains all the functionality provided by the library. +namespace half_float +{ +class half; + +#if HALF_ENABLE_CPP11_USER_LITERALS +/// Library-defined half-precision literals. +/// Import this namespace to enable half-precision floating point literals: +/// ~~~~{.cpp} +/// using namespace half_float::literal; +/// half_float::half = 4.2_h; +/// ~~~~ +namespace literal +{ +half operator"" _h(long double); +} +#endif + +/// \internal +/// \brief Implementation details. +namespace detail +{ +#if HALF_ENABLE_CPP11_TYPE_TRAITS +/// Conditional type. +template +struct conditional : std::conditional +{ +}; + +/// Helper for tag dispatching. +template +struct bool_type : std::integral_constant +{ +}; +using std::false_type; +using std::true_type; + +/// Type traits for floating point types. +template +struct is_float : std::is_floating_point +{ +}; +#else +/// Conditional type. +template +struct conditional +{ + typedef T type; +}; +template +struct conditional +{ + typedef F type; +}; + +/// Helper for tag dispatching. +template +struct bool_type +{ +}; +typedef bool_type true_type; +typedef bool_type false_type; + +/// Type traits for floating point types. +template +struct is_float : false_type +{ +}; +template +struct is_float : is_float +{ +}; +template +struct is_float : is_float +{ +}; +template +struct is_float : is_float +{ +}; +template <> +struct is_float : true_type +{ +}; +template <> +struct is_float : true_type +{ +}; +template <> +struct is_float : true_type +{ +}; +#endif + +/// Type traits for floating point bits. +template +struct bits +{ + typedef unsigned char type; +}; +template +struct bits : bits +{ +}; +template +struct bits : bits +{ +}; +template +struct bits : bits +{ +}; + +#if HALF_ENABLE_CPP11_CSTDINT +/// Unsigned integer of (at least) 16 bits width. +typedef std::uint_least16_t uint16; + +/// Unsigned integer of (at least) 32 bits width. +template <> +struct bits +{ + typedef std::uint_least32_t type; +}; + +/// Unsigned integer of (at least) 64 bits width. +template <> +struct bits +{ + typedef std::uint_least64_t type; +}; +#else +/// Unsigned integer of (at least) 16 bits width. +typedef unsigned short uint16; + +/// Unsigned integer of (at least) 32 bits width. +template <> +struct bits : conditional::digits >= 32, unsigned int, unsigned long> +{ +}; + +#if HALF_ENABLE_CPP11_LONG_LONG +/// Unsigned integer of (at least) 64 bits width. +template <> +struct bits : conditional::digits >= 64, unsigned long, unsigned long long> +{ +}; +#else +/// Unsigned integer of (at least) 64 bits width. +template <> +struct bits +{ + typedef unsigned long type; +}; +#endif +#endif + +/// Tag type for binary construction. +struct binary_t +{ +}; + +/// Tag for binary construction. +HALF_CONSTEXPR_CONST binary_t binary = binary_t(); + +/// Temporary half-precision expression. +/// This class represents a half-precision expression which just stores a single-precision value internally. +struct expr +{ + /// Conversion constructor. + /// \param f single-precision value to convert + explicit HALF_CONSTEXPR expr(float f) HALF_NOEXCEPT : value_(f) {} + + /// Conversion to single-precision. + /// \return single precision value representing expression value + HALF_CONSTEXPR operator float() const HALF_NOEXCEPT + { + return value_; + } + +private: + /// Internal expression value stored in single-precision. + float value_; +}; + +/// SFINAE helper for generic half-precision functions. +/// This class template has to be specialized for each valid combination of argument types to provide a corresponding +/// `type` member equivalent to \a T. +/// \tparam T type to return +template +struct enable +{ +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; +template +struct enable +{ + typedef T type; +}; + +/// Return type for specialized generic 2-argument half-precision functions. +/// This class template has to be specialized for each valid combination of argument types to provide a corresponding +/// `type` member denoting the appropriate return type. +/// \tparam T first argument type +/// \tparam U first argument type +template +struct result : enable +{ +}; +template <> +struct result +{ + typedef half type; +}; + +/// \name Classification helpers +/// \{ + +/// Check for infinity. +/// \tparam T argument type (builtin floating point type) +/// \param arg value to query +/// \retval true if infinity +/// \retval false else +template +bool builtin_isinf(T arg) +{ +#if HALF_ENABLE_CPP11_CMATH + return std::isinf(arg); +#elif defined(_MSC_VER) + return !::_finite(static_cast(arg)) && !::_isnan(static_cast(arg)); +#else + return arg == std::numeric_limits::infinity() || arg == -std::numeric_limits::infinity(); +#endif +} + +/// Check for NaN. +/// \tparam T argument type (builtin floating point type) +/// \param arg value to query +/// \retval true if not a number +/// \retval false else +template +bool builtin_isnan(T arg) +{ +#if HALF_ENABLE_CPP11_CMATH + return std::isnan(arg); +#elif defined(_MSC_VER) + return ::_isnan(static_cast(arg)) != 0; +#else + return arg != arg; +#endif +} + +/// Check sign. +/// \tparam T argument type (builtin floating point type) +/// \param arg value to query +/// \retval true if signbit set +/// \retval false else +template +bool builtin_signbit(T arg) +{ +#if HALF_ENABLE_CPP11_CMATH + return std::signbit(arg); +#else + return arg < T() || (arg == T() && T(1) / arg < T()); +#endif +} + +/// \} +/// \name Conversion +/// \{ + +/// Convert IEEE single-precision to half-precision. +/// Credit for this goes to [Jeroen van der Zijp](ftp://ftp.fox-toolkit.org/pub/fasthalffloatconversion.pdf). +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \param value single-precision value +/// \return binary representation of half-precision value +template +uint16 float2half_impl(float value, true_type) +{ + typedef bits::type uint32; + uint32 bits; // = *reinterpret_cast(&value); //violating strict aliasing! + std::memcpy(&bits, &value, sizeof(float)); + /* uint16 hbits = (bits>>16) & 0x8000; + bits &= 0x7FFFFFFF; + int exp = bits >> 23; + if(exp == 255) + return hbits | 0x7C00 | (0x3FF&-static_cast((bits&0x7FFFFF)!=0)); + if(exp > 142) + { + if(R == std::round_toward_infinity) + return hbits | 0x7C00 - (hbits>>15); + if(R == std::round_toward_neg_infinity) + return hbits | 0x7BFF + (hbits>>15); + return hbits | 0x7BFF + (R!=std::round_toward_zero); + } + int g, s; + if(exp > 112) + { + g = (bits>>12) & 1; + s = (bits&0xFFF) != 0; + hbits |= ((exp-112)<<10) | ((bits>>13)&0x3FF); + } + else if(exp > 101) + { + int i = 125 - exp; + bits = (bits&0x7FFFFF) | 0x800000; + g = (bits>>i) & 1; + s = (bits&((1L<> (i+1); + } + else + { + g = 0; + s = bits != 0; + } + if(R == std::round_to_nearest) + #if HALF_ROUND_TIES_TO_EVEN + hbits += g & (s|hbits); + #else + hbits += g; + #endif + else if(R == std::round_toward_infinity) + hbits += ~(hbits>>15) & (s|g); + else if(R == std::round_toward_neg_infinity) + hbits += (hbits>>15) & (g|s); + */ + static const uint16 base_table[512] = {0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, + 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0000, 0x0001, 0x0002, 0x0004, 0x0008, + 0x0010, 0x0020, 0x0040, 0x0080, 0x0100, 0x0200, 0x0400, 0x0800, 0x0C00, 0x1000, 0x1400, 0x1800, 0x1C00, 0x2000, + 0x2400, 0x2800, 0x2C00, 0x3000, 0x3400, 0x3800, 0x3C00, 0x4000, 0x4400, 0x4800, 0x4C00, 0x5000, 0x5400, 0x5800, + 0x5C00, 0x6000, 0x6400, 0x6800, 0x6C00, 0x7000, 0x7400, 0x7800, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, + 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x7C00, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, 0x8000, + 0x8001, 0x8002, 0x8004, 0x8008, 0x8010, 0x8020, 0x8040, 0x8080, 0x8100, 0x8200, 0x8400, 0x8800, 0x8C00, 0x9000, + 0x9400, 0x9800, 0x9C00, 0xA000, 0xA400, 0xA800, 0xAC00, 0xB000, 0xB400, 0xB800, 0xBC00, 0xC000, 0xC400, 0xC800, + 0xCC00, 0xD000, 0xD400, 0xD800, 0xDC00, 0xE000, 0xE400, 0xE800, 0xEC00, 0xF000, 0xF400, 0xF800, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, + 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00, 0xFC00}; + static unsigned char const shift_table[512] = {24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 23, 22, 21, 20, 19, 18, 17, 16, 15, 14, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, + 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 13, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 23, 22, 21, 20, 19, 18, 17, 16, 15, 14, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, + 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, + 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 13}; + uint16 hbits = base_table[bits >> 23] + static_cast((bits & 0x7FFFFF) >> shift_table[bits >> 23]); + if (R == std::round_to_nearest) + hbits += (((bits & 0x7FFFFF) >> (shift_table[bits >> 23] - 1)) | (((bits >> 23) & 0xFF) == 102)) + & ((hbits & 0x7C00) != 0x7C00) +#if HALF_ROUND_TIES_TO_EVEN + & (((((static_cast(1) << (shift_table[bits >> 23] - 1)) - 1) & bits) != 0) | hbits) +#endif + ; + else if (R == std::round_toward_zero) + hbits -= ((hbits & 0x7FFF) == 0x7C00) & ~shift_table[bits >> 23]; + else if (R == std::round_toward_infinity) + hbits += ((((bits & 0x7FFFFF & ((static_cast(1) << (shift_table[bits >> 23])) - 1)) != 0) + | (((bits >> 23) <= 102) & ((bits >> 23) != 0))) + & (hbits < 0x7C00)) + - ((hbits == 0xFC00) & ((bits >> 23) != 511)); + else if (R == std::round_toward_neg_infinity) + hbits += ((((bits & 0x7FFFFF & ((static_cast(1) << (shift_table[bits >> 23])) - 1)) != 0) + | (((bits >> 23) <= 358) & ((bits >> 23) != 256))) + & (hbits < 0xFC00) & (hbits >> 15)) + - ((hbits == 0x7C00) & ((bits >> 23) != 255)); + return hbits; +} + +/// Convert IEEE double-precision to half-precision. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \param value double-precision value +/// \return binary representation of half-precision value +template +uint16 float2half_impl(double value, true_type) +{ + typedef bits::type uint32; + typedef bits::type uint64; + uint64 bits; // = *reinterpret_cast(&value); //violating strict aliasing! + std::memcpy(&bits, &value, sizeof(double)); + uint32 hi = bits >> 32, lo = bits & 0xFFFFFFFF; + uint16 hbits = (hi >> 16) & 0x8000; + hi &= 0x7FFFFFFF; + int exp = hi >> 20; + if (exp == 2047) + return hbits | 0x7C00 | (0x3FF & -static_cast((bits & 0xFFFFFFFFFFFFF) != 0)); + if (exp > 1038) + { + if (R == std::round_toward_infinity) + return hbits | 0x7C00 - (hbits >> 15); + if (R == std::round_toward_neg_infinity) + return hbits | 0x7BFF + (hbits >> 15); + return hbits | 0x7BFF + (R != std::round_toward_zero); + } + int g, s = lo != 0; + if (exp > 1008) + { + g = (hi >> 9) & 1; + s |= (hi & 0x1FF) != 0; + hbits |= ((exp - 1008) << 10) | ((hi >> 10) & 0x3FF); + } + else if (exp > 997) + { + int i = 1018 - exp; + hi = (hi & 0xFFFFF) | 0x100000; + g = (hi >> i) & 1; + s |= (hi & ((1L << i) - 1)) != 0; + hbits |= hi >> (i + 1); + } + else + { + g = 0; + s |= hi != 0; + } + if (R == std::round_to_nearest) +#if HALF_ROUND_TIES_TO_EVEN + hbits += g & (s | hbits); +#else + hbits += g; +#endif + else if (R == std::round_toward_infinity) + hbits += ~(hbits >> 15) & (s | g); + else if (R == std::round_toward_neg_infinity) + hbits += (hbits >> 15) & (g | s); + return hbits; +} + +/// Convert non-IEEE floating point to half-precision. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam T source type (builtin floating point type) +/// \param value floating point value +/// \return binary representation of half-precision value +template +uint16 float2half_impl(T value, ...) +{ + uint16 hbits = static_cast(builtin_signbit(value)) << 15; + if (value == T()) + return hbits; + if (builtin_isnan(value)) + return hbits | 0x7FFF; + if (builtin_isinf(value)) + return hbits | 0x7C00; + int exp; + std::frexp(value, &exp); + if (exp > 16) + { + if (R == std::round_toward_infinity) + return hbits | (0x7C00 - (hbits >> 15)); + else if (R == std::round_toward_neg_infinity) + return hbits | (0x7BFF + (hbits >> 15)); + return hbits | (0x7BFF + (R != std::round_toward_zero)); + } + if (exp < -13) + value = std::ldexp(value, 24); + else + { + value = std::ldexp(value, 11 - exp); + hbits |= ((exp + 13) << 10); + } + T ival, frac = std::modf(value, &ival); + hbits += static_cast(std::abs(static_cast(ival))); + if (R == std::round_to_nearest) + { + frac = std::abs(frac); +#if HALF_ROUND_TIES_TO_EVEN + hbits += (frac > T(0.5)) | ((frac == T(0.5)) & hbits); +#else + hbits += frac >= T(0.5); +#endif + } + else if (R == std::round_toward_infinity) + hbits += frac > T(); + else if (R == std::round_toward_neg_infinity) + hbits += frac < T(); + return hbits; +} + +/// Convert floating point to half-precision. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam T source type (builtin floating point type) +/// \param value floating point value +/// \return binary representation of half-precision value +template +uint16 float2half(T value) +{ + return float2half_impl( + value, bool_type < std::numeric_limits::is_iec559 && sizeof(typename bits::type) == sizeof(T) > ()); +} + +/// Convert integer to half-precision floating point. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam S `true` if value negative, `false` else +/// \tparam T type to convert (builtin integer type) +/// \param value non-negative integral value +/// \return binary representation of half-precision value +template +uint16 int2half_impl(T value) +{ +#if HALF_ENABLE_CPP11_STATIC_ASSERT && HALF_ENABLE_CPP11_TYPE_TRAITS + static_assert(std::is_integral::value, "int to half conversion only supports builtin integer types"); +#endif + if (S) + value = -value; + uint16 bits = S << 15; + if (value > 0xFFFF) + { + if (R == std::round_toward_infinity) + bits |= 0x7C00 - S; + else if (R == std::round_toward_neg_infinity) + bits |= 0x7BFF + S; + else + bits |= 0x7BFF + (R != std::round_toward_zero); + } + else if (value) + { + uint32_t m = value, exp = 24; + for (; m < 0x400; m <<= 1, --exp) + ; + for (; m > 0x7FF; m >>= 1, ++exp) + ; + bits |= (exp << 10) + m; + if (exp > 24) + { + if (R == std::round_to_nearest) + bits += (value >> (exp - 25)) & 1 +#if HALF_ROUND_TIES_TO_EVEN + & (((((1 << (exp - 25)) - 1) & value) != 0) | bits) +#endif + ; + else if (R == std::round_toward_infinity) + bits += ((value & ((1 << (exp - 24)) - 1)) != 0) & !S; + else if (R == std::round_toward_neg_infinity) + bits += ((value & ((1 << (exp - 24)) - 1)) != 0) & S; + } + } + return bits; +} + +/// Convert integer to half-precision floating point. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam T type to convert (builtin integer type) +/// \param value integral value +/// \return binary representation of half-precision value +template +uint16 int2half(T value) +{ + return (value < 0) ? int2half_impl(value) : int2half_impl(value); +} + +/// Convert half-precision to IEEE single-precision. +/// Credit for this goes to [Jeroen van der Zijp](ftp://ftp.fox-toolkit.org/pub/fasthalffloatconversion.pdf). +/// \param value binary representation of half-precision value +/// \return single-precision value +inline float half2float_impl(uint16 value, float, true_type) +{ + typedef bits::type uint32; + /* uint32 bits = static_cast(value&0x8000) << 16; + int abs = value & 0x7FFF; + if(abs) + { + bits |= 0x38000000 << static_cast(abs>=0x7C00); + for(; abs<0x400; abs<<=1,bits-=0x800000) ; + bits += static_cast(abs) << 13; + } + */ + static const uint32 mantissa_table[2048] = {0x00000000, 0x33800000, 0x34000000, 0x34400000, 0x34800000, 0x34A00000, + 0x34C00000, 0x34E00000, 0x35000000, 0x35100000, 0x35200000, 0x35300000, 0x35400000, 0x35500000, 0x35600000, + 0x35700000, 0x35800000, 0x35880000, 0x35900000, 0x35980000, 0x35A00000, 0x35A80000, 0x35B00000, 0x35B80000, + 0x35C00000, 0x35C80000, 0x35D00000, 0x35D80000, 0x35E00000, 0x35E80000, 0x35F00000, 0x35F80000, 0x36000000, + 0x36040000, 0x36080000, 0x360C0000, 0x36100000, 0x36140000, 0x36180000, 0x361C0000, 0x36200000, 0x36240000, + 0x36280000, 0x362C0000, 0x36300000, 0x36340000, 0x36380000, 0x363C0000, 0x36400000, 0x36440000, 0x36480000, + 0x364C0000, 0x36500000, 0x36540000, 0x36580000, 0x365C0000, 0x36600000, 0x36640000, 0x36680000, 0x366C0000, + 0x36700000, 0x36740000, 0x36780000, 0x367C0000, 0x36800000, 0x36820000, 0x36840000, 0x36860000, 0x36880000, + 0x368A0000, 0x368C0000, 0x368E0000, 0x36900000, 0x36920000, 0x36940000, 0x36960000, 0x36980000, 0x369A0000, + 0x369C0000, 0x369E0000, 0x36A00000, 0x36A20000, 0x36A40000, 0x36A60000, 0x36A80000, 0x36AA0000, 0x36AC0000, + 0x36AE0000, 0x36B00000, 0x36B20000, 0x36B40000, 0x36B60000, 0x36B80000, 0x36BA0000, 0x36BC0000, 0x36BE0000, + 0x36C00000, 0x36C20000, 0x36C40000, 0x36C60000, 0x36C80000, 0x36CA0000, 0x36CC0000, 0x36CE0000, 0x36D00000, + 0x36D20000, 0x36D40000, 0x36D60000, 0x36D80000, 0x36DA0000, 0x36DC0000, 0x36DE0000, 0x36E00000, 0x36E20000, + 0x36E40000, 0x36E60000, 0x36E80000, 0x36EA0000, 0x36EC0000, 0x36EE0000, 0x36F00000, 0x36F20000, 0x36F40000, + 0x36F60000, 0x36F80000, 0x36FA0000, 0x36FC0000, 0x36FE0000, 0x37000000, 0x37010000, 0x37020000, 0x37030000, + 0x37040000, 0x37050000, 0x37060000, 0x37070000, 0x37080000, 0x37090000, 0x370A0000, 0x370B0000, 0x370C0000, + 0x370D0000, 0x370E0000, 0x370F0000, 0x37100000, 0x37110000, 0x37120000, 0x37130000, 0x37140000, 0x37150000, + 0x37160000, 0x37170000, 0x37180000, 0x37190000, 0x371A0000, 0x371B0000, 0x371C0000, 0x371D0000, 0x371E0000, + 0x371F0000, 0x37200000, 0x37210000, 0x37220000, 0x37230000, 0x37240000, 0x37250000, 0x37260000, 0x37270000, + 0x37280000, 0x37290000, 0x372A0000, 0x372B0000, 0x372C0000, 0x372D0000, 0x372E0000, 0x372F0000, 0x37300000, + 0x37310000, 0x37320000, 0x37330000, 0x37340000, 0x37350000, 0x37360000, 0x37370000, 0x37380000, 0x37390000, + 0x373A0000, 0x373B0000, 0x373C0000, 0x373D0000, 0x373E0000, 0x373F0000, 0x37400000, 0x37410000, 0x37420000, + 0x37430000, 0x37440000, 0x37450000, 0x37460000, 0x37470000, 0x37480000, 0x37490000, 0x374A0000, 0x374B0000, + 0x374C0000, 0x374D0000, 0x374E0000, 0x374F0000, 0x37500000, 0x37510000, 0x37520000, 0x37530000, 0x37540000, + 0x37550000, 0x37560000, 0x37570000, 0x37580000, 0x37590000, 0x375A0000, 0x375B0000, 0x375C0000, 0x375D0000, + 0x375E0000, 0x375F0000, 0x37600000, 0x37610000, 0x37620000, 0x37630000, 0x37640000, 0x37650000, 0x37660000, + 0x37670000, 0x37680000, 0x37690000, 0x376A0000, 0x376B0000, 0x376C0000, 0x376D0000, 0x376E0000, 0x376F0000, + 0x37700000, 0x37710000, 0x37720000, 0x37730000, 0x37740000, 0x37750000, 0x37760000, 0x37770000, 0x37780000, + 0x37790000, 0x377A0000, 0x377B0000, 0x377C0000, 0x377D0000, 0x377E0000, 0x377F0000, 0x37800000, 0x37808000, + 0x37810000, 0x37818000, 0x37820000, 0x37828000, 0x37830000, 0x37838000, 0x37840000, 0x37848000, 0x37850000, + 0x37858000, 0x37860000, 0x37868000, 0x37870000, 0x37878000, 0x37880000, 0x37888000, 0x37890000, 0x37898000, + 0x378A0000, 0x378A8000, 0x378B0000, 0x378B8000, 0x378C0000, 0x378C8000, 0x378D0000, 0x378D8000, 0x378E0000, + 0x378E8000, 0x378F0000, 0x378F8000, 0x37900000, 0x37908000, 0x37910000, 0x37918000, 0x37920000, 0x37928000, + 0x37930000, 0x37938000, 0x37940000, 0x37948000, 0x37950000, 0x37958000, 0x37960000, 0x37968000, 0x37970000, + 0x37978000, 0x37980000, 0x37988000, 0x37990000, 0x37998000, 0x379A0000, 0x379A8000, 0x379B0000, 0x379B8000, + 0x379C0000, 0x379C8000, 0x379D0000, 0x379D8000, 0x379E0000, 0x379E8000, 0x379F0000, 0x379F8000, 0x37A00000, + 0x37A08000, 0x37A10000, 0x37A18000, 0x37A20000, 0x37A28000, 0x37A30000, 0x37A38000, 0x37A40000, 0x37A48000, + 0x37A50000, 0x37A58000, 0x37A60000, 0x37A68000, 0x37A70000, 0x37A78000, 0x37A80000, 0x37A88000, 0x37A90000, + 0x37A98000, 0x37AA0000, 0x37AA8000, 0x37AB0000, 0x37AB8000, 0x37AC0000, 0x37AC8000, 0x37AD0000, 0x37AD8000, + 0x37AE0000, 0x37AE8000, 0x37AF0000, 0x37AF8000, 0x37B00000, 0x37B08000, 0x37B10000, 0x37B18000, 0x37B20000, + 0x37B28000, 0x37B30000, 0x37B38000, 0x37B40000, 0x37B48000, 0x37B50000, 0x37B58000, 0x37B60000, 0x37B68000, + 0x37B70000, 0x37B78000, 0x37B80000, 0x37B88000, 0x37B90000, 0x37B98000, 0x37BA0000, 0x37BA8000, 0x37BB0000, + 0x37BB8000, 0x37BC0000, 0x37BC8000, 0x37BD0000, 0x37BD8000, 0x37BE0000, 0x37BE8000, 0x37BF0000, 0x37BF8000, + 0x37C00000, 0x37C08000, 0x37C10000, 0x37C18000, 0x37C20000, 0x37C28000, 0x37C30000, 0x37C38000, 0x37C40000, + 0x37C48000, 0x37C50000, 0x37C58000, 0x37C60000, 0x37C68000, 0x37C70000, 0x37C78000, 0x37C80000, 0x37C88000, + 0x37C90000, 0x37C98000, 0x37CA0000, 0x37CA8000, 0x37CB0000, 0x37CB8000, 0x37CC0000, 0x37CC8000, 0x37CD0000, + 0x37CD8000, 0x37CE0000, 0x37CE8000, 0x37CF0000, 0x37CF8000, 0x37D00000, 0x37D08000, 0x37D10000, 0x37D18000, + 0x37D20000, 0x37D28000, 0x37D30000, 0x37D38000, 0x37D40000, 0x37D48000, 0x37D50000, 0x37D58000, 0x37D60000, + 0x37D68000, 0x37D70000, 0x37D78000, 0x37D80000, 0x37D88000, 0x37D90000, 0x37D98000, 0x37DA0000, 0x37DA8000, + 0x37DB0000, 0x37DB8000, 0x37DC0000, 0x37DC8000, 0x37DD0000, 0x37DD8000, 0x37DE0000, 0x37DE8000, 0x37DF0000, + 0x37DF8000, 0x37E00000, 0x37E08000, 0x37E10000, 0x37E18000, 0x37E20000, 0x37E28000, 0x37E30000, 0x37E38000, + 0x37E40000, 0x37E48000, 0x37E50000, 0x37E58000, 0x37E60000, 0x37E68000, 0x37E70000, 0x37E78000, 0x37E80000, + 0x37E88000, 0x37E90000, 0x37E98000, 0x37EA0000, 0x37EA8000, 0x37EB0000, 0x37EB8000, 0x37EC0000, 0x37EC8000, + 0x37ED0000, 0x37ED8000, 0x37EE0000, 0x37EE8000, 0x37EF0000, 0x37EF8000, 0x37F00000, 0x37F08000, 0x37F10000, + 0x37F18000, 0x37F20000, 0x37F28000, 0x37F30000, 0x37F38000, 0x37F40000, 0x37F48000, 0x37F50000, 0x37F58000, + 0x37F60000, 0x37F68000, 0x37F70000, 0x37F78000, 0x37F80000, 0x37F88000, 0x37F90000, 0x37F98000, 0x37FA0000, + 0x37FA8000, 0x37FB0000, 0x37FB8000, 0x37FC0000, 0x37FC8000, 0x37FD0000, 0x37FD8000, 0x37FE0000, 0x37FE8000, + 0x37FF0000, 0x37FF8000, 0x38000000, 0x38004000, 0x38008000, 0x3800C000, 0x38010000, 0x38014000, 0x38018000, + 0x3801C000, 0x38020000, 0x38024000, 0x38028000, 0x3802C000, 0x38030000, 0x38034000, 0x38038000, 0x3803C000, + 0x38040000, 0x38044000, 0x38048000, 0x3804C000, 0x38050000, 0x38054000, 0x38058000, 0x3805C000, 0x38060000, + 0x38064000, 0x38068000, 0x3806C000, 0x38070000, 0x38074000, 0x38078000, 0x3807C000, 0x38080000, 0x38084000, + 0x38088000, 0x3808C000, 0x38090000, 0x38094000, 0x38098000, 0x3809C000, 0x380A0000, 0x380A4000, 0x380A8000, + 0x380AC000, 0x380B0000, 0x380B4000, 0x380B8000, 0x380BC000, 0x380C0000, 0x380C4000, 0x380C8000, 0x380CC000, + 0x380D0000, 0x380D4000, 0x380D8000, 0x380DC000, 0x380E0000, 0x380E4000, 0x380E8000, 0x380EC000, 0x380F0000, + 0x380F4000, 0x380F8000, 0x380FC000, 0x38100000, 0x38104000, 0x38108000, 0x3810C000, 0x38110000, 0x38114000, + 0x38118000, 0x3811C000, 0x38120000, 0x38124000, 0x38128000, 0x3812C000, 0x38130000, 0x38134000, 0x38138000, + 0x3813C000, 0x38140000, 0x38144000, 0x38148000, 0x3814C000, 0x38150000, 0x38154000, 0x38158000, 0x3815C000, + 0x38160000, 0x38164000, 0x38168000, 0x3816C000, 0x38170000, 0x38174000, 0x38178000, 0x3817C000, 0x38180000, + 0x38184000, 0x38188000, 0x3818C000, 0x38190000, 0x38194000, 0x38198000, 0x3819C000, 0x381A0000, 0x381A4000, + 0x381A8000, 0x381AC000, 0x381B0000, 0x381B4000, 0x381B8000, 0x381BC000, 0x381C0000, 0x381C4000, 0x381C8000, + 0x381CC000, 0x381D0000, 0x381D4000, 0x381D8000, 0x381DC000, 0x381E0000, 0x381E4000, 0x381E8000, 0x381EC000, + 0x381F0000, 0x381F4000, 0x381F8000, 0x381FC000, 0x38200000, 0x38204000, 0x38208000, 0x3820C000, 0x38210000, + 0x38214000, 0x38218000, 0x3821C000, 0x38220000, 0x38224000, 0x38228000, 0x3822C000, 0x38230000, 0x38234000, + 0x38238000, 0x3823C000, 0x38240000, 0x38244000, 0x38248000, 0x3824C000, 0x38250000, 0x38254000, 0x38258000, + 0x3825C000, 0x38260000, 0x38264000, 0x38268000, 0x3826C000, 0x38270000, 0x38274000, 0x38278000, 0x3827C000, + 0x38280000, 0x38284000, 0x38288000, 0x3828C000, 0x38290000, 0x38294000, 0x38298000, 0x3829C000, 0x382A0000, + 0x382A4000, 0x382A8000, 0x382AC000, 0x382B0000, 0x382B4000, 0x382B8000, 0x382BC000, 0x382C0000, 0x382C4000, + 0x382C8000, 0x382CC000, 0x382D0000, 0x382D4000, 0x382D8000, 0x382DC000, 0x382E0000, 0x382E4000, 0x382E8000, + 0x382EC000, 0x382F0000, 0x382F4000, 0x382F8000, 0x382FC000, 0x38300000, 0x38304000, 0x38308000, 0x3830C000, + 0x38310000, 0x38314000, 0x38318000, 0x3831C000, 0x38320000, 0x38324000, 0x38328000, 0x3832C000, 0x38330000, + 0x38334000, 0x38338000, 0x3833C000, 0x38340000, 0x38344000, 0x38348000, 0x3834C000, 0x38350000, 0x38354000, + 0x38358000, 0x3835C000, 0x38360000, 0x38364000, 0x38368000, 0x3836C000, 0x38370000, 0x38374000, 0x38378000, + 0x3837C000, 0x38380000, 0x38384000, 0x38388000, 0x3838C000, 0x38390000, 0x38394000, 0x38398000, 0x3839C000, + 0x383A0000, 0x383A4000, 0x383A8000, 0x383AC000, 0x383B0000, 0x383B4000, 0x383B8000, 0x383BC000, 0x383C0000, + 0x383C4000, 0x383C8000, 0x383CC000, 0x383D0000, 0x383D4000, 0x383D8000, 0x383DC000, 0x383E0000, 0x383E4000, + 0x383E8000, 0x383EC000, 0x383F0000, 0x383F4000, 0x383F8000, 0x383FC000, 0x38400000, 0x38404000, 0x38408000, + 0x3840C000, 0x38410000, 0x38414000, 0x38418000, 0x3841C000, 0x38420000, 0x38424000, 0x38428000, 0x3842C000, + 0x38430000, 0x38434000, 0x38438000, 0x3843C000, 0x38440000, 0x38444000, 0x38448000, 0x3844C000, 0x38450000, + 0x38454000, 0x38458000, 0x3845C000, 0x38460000, 0x38464000, 0x38468000, 0x3846C000, 0x38470000, 0x38474000, + 0x38478000, 0x3847C000, 0x38480000, 0x38484000, 0x38488000, 0x3848C000, 0x38490000, 0x38494000, 0x38498000, + 0x3849C000, 0x384A0000, 0x384A4000, 0x384A8000, 0x384AC000, 0x384B0000, 0x384B4000, 0x384B8000, 0x384BC000, + 0x384C0000, 0x384C4000, 0x384C8000, 0x384CC000, 0x384D0000, 0x384D4000, 0x384D8000, 0x384DC000, 0x384E0000, + 0x384E4000, 0x384E8000, 0x384EC000, 0x384F0000, 0x384F4000, 0x384F8000, 0x384FC000, 0x38500000, 0x38504000, + 0x38508000, 0x3850C000, 0x38510000, 0x38514000, 0x38518000, 0x3851C000, 0x38520000, 0x38524000, 0x38528000, + 0x3852C000, 0x38530000, 0x38534000, 0x38538000, 0x3853C000, 0x38540000, 0x38544000, 0x38548000, 0x3854C000, + 0x38550000, 0x38554000, 0x38558000, 0x3855C000, 0x38560000, 0x38564000, 0x38568000, 0x3856C000, 0x38570000, + 0x38574000, 0x38578000, 0x3857C000, 0x38580000, 0x38584000, 0x38588000, 0x3858C000, 0x38590000, 0x38594000, + 0x38598000, 0x3859C000, 0x385A0000, 0x385A4000, 0x385A8000, 0x385AC000, 0x385B0000, 0x385B4000, 0x385B8000, + 0x385BC000, 0x385C0000, 0x385C4000, 0x385C8000, 0x385CC000, 0x385D0000, 0x385D4000, 0x385D8000, 0x385DC000, + 0x385E0000, 0x385E4000, 0x385E8000, 0x385EC000, 0x385F0000, 0x385F4000, 0x385F8000, 0x385FC000, 0x38600000, + 0x38604000, 0x38608000, 0x3860C000, 0x38610000, 0x38614000, 0x38618000, 0x3861C000, 0x38620000, 0x38624000, + 0x38628000, 0x3862C000, 0x38630000, 0x38634000, 0x38638000, 0x3863C000, 0x38640000, 0x38644000, 0x38648000, + 0x3864C000, 0x38650000, 0x38654000, 0x38658000, 0x3865C000, 0x38660000, 0x38664000, 0x38668000, 0x3866C000, + 0x38670000, 0x38674000, 0x38678000, 0x3867C000, 0x38680000, 0x38684000, 0x38688000, 0x3868C000, 0x38690000, + 0x38694000, 0x38698000, 0x3869C000, 0x386A0000, 0x386A4000, 0x386A8000, 0x386AC000, 0x386B0000, 0x386B4000, + 0x386B8000, 0x386BC000, 0x386C0000, 0x386C4000, 0x386C8000, 0x386CC000, 0x386D0000, 0x386D4000, 0x386D8000, + 0x386DC000, 0x386E0000, 0x386E4000, 0x386E8000, 0x386EC000, 0x386F0000, 0x386F4000, 0x386F8000, 0x386FC000, + 0x38700000, 0x38704000, 0x38708000, 0x3870C000, 0x38710000, 0x38714000, 0x38718000, 0x3871C000, 0x38720000, + 0x38724000, 0x38728000, 0x3872C000, 0x38730000, 0x38734000, 0x38738000, 0x3873C000, 0x38740000, 0x38744000, + 0x38748000, 0x3874C000, 0x38750000, 0x38754000, 0x38758000, 0x3875C000, 0x38760000, 0x38764000, 0x38768000, + 0x3876C000, 0x38770000, 0x38774000, 0x38778000, 0x3877C000, 0x38780000, 0x38784000, 0x38788000, 0x3878C000, + 0x38790000, 0x38794000, 0x38798000, 0x3879C000, 0x387A0000, 0x387A4000, 0x387A8000, 0x387AC000, 0x387B0000, + 0x387B4000, 0x387B8000, 0x387BC000, 0x387C0000, 0x387C4000, 0x387C8000, 0x387CC000, 0x387D0000, 0x387D4000, + 0x387D8000, 0x387DC000, 0x387E0000, 0x387E4000, 0x387E8000, 0x387EC000, 0x387F0000, 0x387F4000, 0x387F8000, + 0x387FC000, 0x38000000, 0x38002000, 0x38004000, 0x38006000, 0x38008000, 0x3800A000, 0x3800C000, 0x3800E000, + 0x38010000, 0x38012000, 0x38014000, 0x38016000, 0x38018000, 0x3801A000, 0x3801C000, 0x3801E000, 0x38020000, + 0x38022000, 0x38024000, 0x38026000, 0x38028000, 0x3802A000, 0x3802C000, 0x3802E000, 0x38030000, 0x38032000, + 0x38034000, 0x38036000, 0x38038000, 0x3803A000, 0x3803C000, 0x3803E000, 0x38040000, 0x38042000, 0x38044000, + 0x38046000, 0x38048000, 0x3804A000, 0x3804C000, 0x3804E000, 0x38050000, 0x38052000, 0x38054000, 0x38056000, + 0x38058000, 0x3805A000, 0x3805C000, 0x3805E000, 0x38060000, 0x38062000, 0x38064000, 0x38066000, 0x38068000, + 0x3806A000, 0x3806C000, 0x3806E000, 0x38070000, 0x38072000, 0x38074000, 0x38076000, 0x38078000, 0x3807A000, + 0x3807C000, 0x3807E000, 0x38080000, 0x38082000, 0x38084000, 0x38086000, 0x38088000, 0x3808A000, 0x3808C000, + 0x3808E000, 0x38090000, 0x38092000, 0x38094000, 0x38096000, 0x38098000, 0x3809A000, 0x3809C000, 0x3809E000, + 0x380A0000, 0x380A2000, 0x380A4000, 0x380A6000, 0x380A8000, 0x380AA000, 0x380AC000, 0x380AE000, 0x380B0000, + 0x380B2000, 0x380B4000, 0x380B6000, 0x380B8000, 0x380BA000, 0x380BC000, 0x380BE000, 0x380C0000, 0x380C2000, + 0x380C4000, 0x380C6000, 0x380C8000, 0x380CA000, 0x380CC000, 0x380CE000, 0x380D0000, 0x380D2000, 0x380D4000, + 0x380D6000, 0x380D8000, 0x380DA000, 0x380DC000, 0x380DE000, 0x380E0000, 0x380E2000, 0x380E4000, 0x380E6000, + 0x380E8000, 0x380EA000, 0x380EC000, 0x380EE000, 0x380F0000, 0x380F2000, 0x380F4000, 0x380F6000, 0x380F8000, + 0x380FA000, 0x380FC000, 0x380FE000, 0x38100000, 0x38102000, 0x38104000, 0x38106000, 0x38108000, 0x3810A000, + 0x3810C000, 0x3810E000, 0x38110000, 0x38112000, 0x38114000, 0x38116000, 0x38118000, 0x3811A000, 0x3811C000, + 0x3811E000, 0x38120000, 0x38122000, 0x38124000, 0x38126000, 0x38128000, 0x3812A000, 0x3812C000, 0x3812E000, + 0x38130000, 0x38132000, 0x38134000, 0x38136000, 0x38138000, 0x3813A000, 0x3813C000, 0x3813E000, 0x38140000, + 0x38142000, 0x38144000, 0x38146000, 0x38148000, 0x3814A000, 0x3814C000, 0x3814E000, 0x38150000, 0x38152000, + 0x38154000, 0x38156000, 0x38158000, 0x3815A000, 0x3815C000, 0x3815E000, 0x38160000, 0x38162000, 0x38164000, + 0x38166000, 0x38168000, 0x3816A000, 0x3816C000, 0x3816E000, 0x38170000, 0x38172000, 0x38174000, 0x38176000, + 0x38178000, 0x3817A000, 0x3817C000, 0x3817E000, 0x38180000, 0x38182000, 0x38184000, 0x38186000, 0x38188000, + 0x3818A000, 0x3818C000, 0x3818E000, 0x38190000, 0x38192000, 0x38194000, 0x38196000, 0x38198000, 0x3819A000, + 0x3819C000, 0x3819E000, 0x381A0000, 0x381A2000, 0x381A4000, 0x381A6000, 0x381A8000, 0x381AA000, 0x381AC000, + 0x381AE000, 0x381B0000, 0x381B2000, 0x381B4000, 0x381B6000, 0x381B8000, 0x381BA000, 0x381BC000, 0x381BE000, + 0x381C0000, 0x381C2000, 0x381C4000, 0x381C6000, 0x381C8000, 0x381CA000, 0x381CC000, 0x381CE000, 0x381D0000, + 0x381D2000, 0x381D4000, 0x381D6000, 0x381D8000, 0x381DA000, 0x381DC000, 0x381DE000, 0x381E0000, 0x381E2000, + 0x381E4000, 0x381E6000, 0x381E8000, 0x381EA000, 0x381EC000, 0x381EE000, 0x381F0000, 0x381F2000, 0x381F4000, + 0x381F6000, 0x381F8000, 0x381FA000, 0x381FC000, 0x381FE000, 0x38200000, 0x38202000, 0x38204000, 0x38206000, + 0x38208000, 0x3820A000, 0x3820C000, 0x3820E000, 0x38210000, 0x38212000, 0x38214000, 0x38216000, 0x38218000, + 0x3821A000, 0x3821C000, 0x3821E000, 0x38220000, 0x38222000, 0x38224000, 0x38226000, 0x38228000, 0x3822A000, + 0x3822C000, 0x3822E000, 0x38230000, 0x38232000, 0x38234000, 0x38236000, 0x38238000, 0x3823A000, 0x3823C000, + 0x3823E000, 0x38240000, 0x38242000, 0x38244000, 0x38246000, 0x38248000, 0x3824A000, 0x3824C000, 0x3824E000, + 0x38250000, 0x38252000, 0x38254000, 0x38256000, 0x38258000, 0x3825A000, 0x3825C000, 0x3825E000, 0x38260000, + 0x38262000, 0x38264000, 0x38266000, 0x38268000, 0x3826A000, 0x3826C000, 0x3826E000, 0x38270000, 0x38272000, + 0x38274000, 0x38276000, 0x38278000, 0x3827A000, 0x3827C000, 0x3827E000, 0x38280000, 0x38282000, 0x38284000, + 0x38286000, 0x38288000, 0x3828A000, 0x3828C000, 0x3828E000, 0x38290000, 0x38292000, 0x38294000, 0x38296000, + 0x38298000, 0x3829A000, 0x3829C000, 0x3829E000, 0x382A0000, 0x382A2000, 0x382A4000, 0x382A6000, 0x382A8000, + 0x382AA000, 0x382AC000, 0x382AE000, 0x382B0000, 0x382B2000, 0x382B4000, 0x382B6000, 0x382B8000, 0x382BA000, + 0x382BC000, 0x382BE000, 0x382C0000, 0x382C2000, 0x382C4000, 0x382C6000, 0x382C8000, 0x382CA000, 0x382CC000, + 0x382CE000, 0x382D0000, 0x382D2000, 0x382D4000, 0x382D6000, 0x382D8000, 0x382DA000, 0x382DC000, 0x382DE000, + 0x382E0000, 0x382E2000, 0x382E4000, 0x382E6000, 0x382E8000, 0x382EA000, 0x382EC000, 0x382EE000, 0x382F0000, + 0x382F2000, 0x382F4000, 0x382F6000, 0x382F8000, 0x382FA000, 0x382FC000, 0x382FE000, 0x38300000, 0x38302000, + 0x38304000, 0x38306000, 0x38308000, 0x3830A000, 0x3830C000, 0x3830E000, 0x38310000, 0x38312000, 0x38314000, + 0x38316000, 0x38318000, 0x3831A000, 0x3831C000, 0x3831E000, 0x38320000, 0x38322000, 0x38324000, 0x38326000, + 0x38328000, 0x3832A000, 0x3832C000, 0x3832E000, 0x38330000, 0x38332000, 0x38334000, 0x38336000, 0x38338000, + 0x3833A000, 0x3833C000, 0x3833E000, 0x38340000, 0x38342000, 0x38344000, 0x38346000, 0x38348000, 0x3834A000, + 0x3834C000, 0x3834E000, 0x38350000, 0x38352000, 0x38354000, 0x38356000, 0x38358000, 0x3835A000, 0x3835C000, + 0x3835E000, 0x38360000, 0x38362000, 0x38364000, 0x38366000, 0x38368000, 0x3836A000, 0x3836C000, 0x3836E000, + 0x38370000, 0x38372000, 0x38374000, 0x38376000, 0x38378000, 0x3837A000, 0x3837C000, 0x3837E000, 0x38380000, + 0x38382000, 0x38384000, 0x38386000, 0x38388000, 0x3838A000, 0x3838C000, 0x3838E000, 0x38390000, 0x38392000, + 0x38394000, 0x38396000, 0x38398000, 0x3839A000, 0x3839C000, 0x3839E000, 0x383A0000, 0x383A2000, 0x383A4000, + 0x383A6000, 0x383A8000, 0x383AA000, 0x383AC000, 0x383AE000, 0x383B0000, 0x383B2000, 0x383B4000, 0x383B6000, + 0x383B8000, 0x383BA000, 0x383BC000, 0x383BE000, 0x383C0000, 0x383C2000, 0x383C4000, 0x383C6000, 0x383C8000, + 0x383CA000, 0x383CC000, 0x383CE000, 0x383D0000, 0x383D2000, 0x383D4000, 0x383D6000, 0x383D8000, 0x383DA000, + 0x383DC000, 0x383DE000, 0x383E0000, 0x383E2000, 0x383E4000, 0x383E6000, 0x383E8000, 0x383EA000, 0x383EC000, + 0x383EE000, 0x383F0000, 0x383F2000, 0x383F4000, 0x383F6000, 0x383F8000, 0x383FA000, 0x383FC000, 0x383FE000, + 0x38400000, 0x38402000, 0x38404000, 0x38406000, 0x38408000, 0x3840A000, 0x3840C000, 0x3840E000, 0x38410000, + 0x38412000, 0x38414000, 0x38416000, 0x38418000, 0x3841A000, 0x3841C000, 0x3841E000, 0x38420000, 0x38422000, + 0x38424000, 0x38426000, 0x38428000, 0x3842A000, 0x3842C000, 0x3842E000, 0x38430000, 0x38432000, 0x38434000, + 0x38436000, 0x38438000, 0x3843A000, 0x3843C000, 0x3843E000, 0x38440000, 0x38442000, 0x38444000, 0x38446000, + 0x38448000, 0x3844A000, 0x3844C000, 0x3844E000, 0x38450000, 0x38452000, 0x38454000, 0x38456000, 0x38458000, + 0x3845A000, 0x3845C000, 0x3845E000, 0x38460000, 0x38462000, 0x38464000, 0x38466000, 0x38468000, 0x3846A000, + 0x3846C000, 0x3846E000, 0x38470000, 0x38472000, 0x38474000, 0x38476000, 0x38478000, 0x3847A000, 0x3847C000, + 0x3847E000, 0x38480000, 0x38482000, 0x38484000, 0x38486000, 0x38488000, 0x3848A000, 0x3848C000, 0x3848E000, + 0x38490000, 0x38492000, 0x38494000, 0x38496000, 0x38498000, 0x3849A000, 0x3849C000, 0x3849E000, 0x384A0000, + 0x384A2000, 0x384A4000, 0x384A6000, 0x384A8000, 0x384AA000, 0x384AC000, 0x384AE000, 0x384B0000, 0x384B2000, + 0x384B4000, 0x384B6000, 0x384B8000, 0x384BA000, 0x384BC000, 0x384BE000, 0x384C0000, 0x384C2000, 0x384C4000, + 0x384C6000, 0x384C8000, 0x384CA000, 0x384CC000, 0x384CE000, 0x384D0000, 0x384D2000, 0x384D4000, 0x384D6000, + 0x384D8000, 0x384DA000, 0x384DC000, 0x384DE000, 0x384E0000, 0x384E2000, 0x384E4000, 0x384E6000, 0x384E8000, + 0x384EA000, 0x384EC000, 0x384EE000, 0x384F0000, 0x384F2000, 0x384F4000, 0x384F6000, 0x384F8000, 0x384FA000, + 0x384FC000, 0x384FE000, 0x38500000, 0x38502000, 0x38504000, 0x38506000, 0x38508000, 0x3850A000, 0x3850C000, + 0x3850E000, 0x38510000, 0x38512000, 0x38514000, 0x38516000, 0x38518000, 0x3851A000, 0x3851C000, 0x3851E000, + 0x38520000, 0x38522000, 0x38524000, 0x38526000, 0x38528000, 0x3852A000, 0x3852C000, 0x3852E000, 0x38530000, + 0x38532000, 0x38534000, 0x38536000, 0x38538000, 0x3853A000, 0x3853C000, 0x3853E000, 0x38540000, 0x38542000, + 0x38544000, 0x38546000, 0x38548000, 0x3854A000, 0x3854C000, 0x3854E000, 0x38550000, 0x38552000, 0x38554000, + 0x38556000, 0x38558000, 0x3855A000, 0x3855C000, 0x3855E000, 0x38560000, 0x38562000, 0x38564000, 0x38566000, + 0x38568000, 0x3856A000, 0x3856C000, 0x3856E000, 0x38570000, 0x38572000, 0x38574000, 0x38576000, 0x38578000, + 0x3857A000, 0x3857C000, 0x3857E000, 0x38580000, 0x38582000, 0x38584000, 0x38586000, 0x38588000, 0x3858A000, + 0x3858C000, 0x3858E000, 0x38590000, 0x38592000, 0x38594000, 0x38596000, 0x38598000, 0x3859A000, 0x3859C000, + 0x3859E000, 0x385A0000, 0x385A2000, 0x385A4000, 0x385A6000, 0x385A8000, 0x385AA000, 0x385AC000, 0x385AE000, + 0x385B0000, 0x385B2000, 0x385B4000, 0x385B6000, 0x385B8000, 0x385BA000, 0x385BC000, 0x385BE000, 0x385C0000, + 0x385C2000, 0x385C4000, 0x385C6000, 0x385C8000, 0x385CA000, 0x385CC000, 0x385CE000, 0x385D0000, 0x385D2000, + 0x385D4000, 0x385D6000, 0x385D8000, 0x385DA000, 0x385DC000, 0x385DE000, 0x385E0000, 0x385E2000, 0x385E4000, + 0x385E6000, 0x385E8000, 0x385EA000, 0x385EC000, 0x385EE000, 0x385F0000, 0x385F2000, 0x385F4000, 0x385F6000, + 0x385F8000, 0x385FA000, 0x385FC000, 0x385FE000, 0x38600000, 0x38602000, 0x38604000, 0x38606000, 0x38608000, + 0x3860A000, 0x3860C000, 0x3860E000, 0x38610000, 0x38612000, 0x38614000, 0x38616000, 0x38618000, 0x3861A000, + 0x3861C000, 0x3861E000, 0x38620000, 0x38622000, 0x38624000, 0x38626000, 0x38628000, 0x3862A000, 0x3862C000, + 0x3862E000, 0x38630000, 0x38632000, 0x38634000, 0x38636000, 0x38638000, 0x3863A000, 0x3863C000, 0x3863E000, + 0x38640000, 0x38642000, 0x38644000, 0x38646000, 0x38648000, 0x3864A000, 0x3864C000, 0x3864E000, 0x38650000, + 0x38652000, 0x38654000, 0x38656000, 0x38658000, 0x3865A000, 0x3865C000, 0x3865E000, 0x38660000, 0x38662000, + 0x38664000, 0x38666000, 0x38668000, 0x3866A000, 0x3866C000, 0x3866E000, 0x38670000, 0x38672000, 0x38674000, + 0x38676000, 0x38678000, 0x3867A000, 0x3867C000, 0x3867E000, 0x38680000, 0x38682000, 0x38684000, 0x38686000, + 0x38688000, 0x3868A000, 0x3868C000, 0x3868E000, 0x38690000, 0x38692000, 0x38694000, 0x38696000, 0x38698000, + 0x3869A000, 0x3869C000, 0x3869E000, 0x386A0000, 0x386A2000, 0x386A4000, 0x386A6000, 0x386A8000, 0x386AA000, + 0x386AC000, 0x386AE000, 0x386B0000, 0x386B2000, 0x386B4000, 0x386B6000, 0x386B8000, 0x386BA000, 0x386BC000, + 0x386BE000, 0x386C0000, 0x386C2000, 0x386C4000, 0x386C6000, 0x386C8000, 0x386CA000, 0x386CC000, 0x386CE000, + 0x386D0000, 0x386D2000, 0x386D4000, 0x386D6000, 0x386D8000, 0x386DA000, 0x386DC000, 0x386DE000, 0x386E0000, + 0x386E2000, 0x386E4000, 0x386E6000, 0x386E8000, 0x386EA000, 0x386EC000, 0x386EE000, 0x386F0000, 0x386F2000, + 0x386F4000, 0x386F6000, 0x386F8000, 0x386FA000, 0x386FC000, 0x386FE000, 0x38700000, 0x38702000, 0x38704000, + 0x38706000, 0x38708000, 0x3870A000, 0x3870C000, 0x3870E000, 0x38710000, 0x38712000, 0x38714000, 0x38716000, + 0x38718000, 0x3871A000, 0x3871C000, 0x3871E000, 0x38720000, 0x38722000, 0x38724000, 0x38726000, 0x38728000, + 0x3872A000, 0x3872C000, 0x3872E000, 0x38730000, 0x38732000, 0x38734000, 0x38736000, 0x38738000, 0x3873A000, + 0x3873C000, 0x3873E000, 0x38740000, 0x38742000, 0x38744000, 0x38746000, 0x38748000, 0x3874A000, 0x3874C000, + 0x3874E000, 0x38750000, 0x38752000, 0x38754000, 0x38756000, 0x38758000, 0x3875A000, 0x3875C000, 0x3875E000, + 0x38760000, 0x38762000, 0x38764000, 0x38766000, 0x38768000, 0x3876A000, 0x3876C000, 0x3876E000, 0x38770000, + 0x38772000, 0x38774000, 0x38776000, 0x38778000, 0x3877A000, 0x3877C000, 0x3877E000, 0x38780000, 0x38782000, + 0x38784000, 0x38786000, 0x38788000, 0x3878A000, 0x3878C000, 0x3878E000, 0x38790000, 0x38792000, 0x38794000, + 0x38796000, 0x38798000, 0x3879A000, 0x3879C000, 0x3879E000, 0x387A0000, 0x387A2000, 0x387A4000, 0x387A6000, + 0x387A8000, 0x387AA000, 0x387AC000, 0x387AE000, 0x387B0000, 0x387B2000, 0x387B4000, 0x387B6000, 0x387B8000, + 0x387BA000, 0x387BC000, 0x387BE000, 0x387C0000, 0x387C2000, 0x387C4000, 0x387C6000, 0x387C8000, 0x387CA000, + 0x387CC000, 0x387CE000, 0x387D0000, 0x387D2000, 0x387D4000, 0x387D6000, 0x387D8000, 0x387DA000, 0x387DC000, + 0x387DE000, 0x387E0000, 0x387E2000, 0x387E4000, 0x387E6000, 0x387E8000, 0x387EA000, 0x387EC000, 0x387EE000, + 0x387F0000, 0x387F2000, 0x387F4000, 0x387F6000, 0x387F8000, 0x387FA000, 0x387FC000, 0x387FE000}; + static const uint32 exponent_table[64] = {0x00000000, 0x00800000, 0x01000000, 0x01800000, 0x02000000, 0x02800000, + 0x03000000, 0x03800000, 0x04000000, 0x04800000, 0x05000000, 0x05800000, 0x06000000, 0x06800000, 0x07000000, + 0x07800000, 0x08000000, 0x08800000, 0x09000000, 0x09800000, 0x0A000000, 0x0A800000, 0x0B000000, 0x0B800000, + 0x0C000000, 0x0C800000, 0x0D000000, 0x0D800000, 0x0E000000, 0x0E800000, 0x0F000000, 0x47800000, 0x80000000, + 0x80800000, 0x81000000, 0x81800000, 0x82000000, 0x82800000, 0x83000000, 0x83800000, 0x84000000, 0x84800000, + 0x85000000, 0x85800000, 0x86000000, 0x86800000, 0x87000000, 0x87800000, 0x88000000, 0x88800000, 0x89000000, + 0x89800000, 0x8A000000, 0x8A800000, 0x8B000000, 0x8B800000, 0x8C000000, 0x8C800000, 0x8D000000, 0x8D800000, + 0x8E000000, 0x8E800000, 0x8F000000, 0xC7800000}; + static unsigned short const offset_table[64] = {0, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, + 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, + 1024, 1024, 0, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, + 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024, 1024}; + uint32 bits = mantissa_table[offset_table[value >> 10] + (value & 0x3FF)] + exponent_table[value >> 10]; + // return *reinterpret_cast(&bits); //violating strict aliasing! + float out; + std::memcpy(&out, &bits, sizeof(float)); + return out; +} + +/// Convert half-precision to IEEE double-precision. +/// \param value binary representation of half-precision value +/// \return double-precision value +inline double half2float_impl(uint16 value, double, true_type) +{ + typedef bits::type uint32; + typedef bits::type uint64; + uint32 hi = static_cast(value & 0x8000) << 16; + int abs = value & 0x7FFF; + if (abs) + { + hi |= 0x3F000000 << static_cast(abs >= 0x7C00); + for (; abs < 0x400; abs <<= 1, hi -= 0x100000) + ; + hi += static_cast(abs) << 10; + } + uint64 bits = static_cast(hi) << 32; + // return *reinterpret_cast(&bits); //violating strict aliasing! + double out; + std::memcpy(&out, &bits, sizeof(double)); + return out; +} + +/// Convert half-precision to non-IEEE floating point. +/// \tparam T type to convert to (builtin integer type) +/// \param value binary representation of half-precision value +/// \return floating point value +template +T half2float_impl(uint16 value, T, ...) +{ + T out; + int abs = value & 0x7FFF; + if (abs > 0x7C00) + out = std::numeric_limits::has_quiet_NaN ? std::numeric_limits::quiet_NaN() : T(); + else if (abs == 0x7C00) + out = std::numeric_limits::has_infinity ? std::numeric_limits::infinity() : std::numeric_limits::max(); + else if (abs > 0x3FF) + out = std::ldexp(static_cast((abs & 0x3FF) | 0x400), (abs >> 10) - 25); + else + out = std::ldexp(static_cast(abs), -24); + return (value & 0x8000) ? -out : out; +} + +/// Convert half-precision to floating point. +/// \tparam T type to convert to (builtin integer type) +/// \param value binary representation of half-precision value +/// \return floating point value +template +T half2float(uint16 value) +{ + return half2float_impl( + value, T(), bool_type < std::numeric_limits::is_iec559 && sizeof(typename bits::type) == sizeof(T) > ()); +} + +/// Convert half-precision floating point to integer. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam E `true` for round to even, `false` for round away from zero +/// \tparam T type to convert to (buitlin integer type with at least 16 bits precision, excluding any implicit sign +/// bits) \param value binary representation of half-precision value \return integral value +template +T half2int_impl(uint16 value) +{ +#if HALF_ENABLE_CPP11_STATIC_ASSERT && HALF_ENABLE_CPP11_TYPE_TRAITS + static_assert(std::is_integral::value, "half to int conversion only supports builtin integer types"); +#endif + uint32_t e = value & 0x7FFF; + if (e >= 0x7C00) + return (value & 0x8000) ? std::numeric_limits::min() : std::numeric_limits::max(); + if (e < 0x3800) + { + if (R == std::round_toward_infinity) + return T(~(value >> 15) & (e != 0)); + else if (R == std::round_toward_neg_infinity) + return -T(value > 0x8000); + return T(); + } + uint32_t m = (value & 0x3FF) | 0x400; + e >>= 10; + if (e < 25) + { + if (R == std::round_to_nearest) + m += (1 << (24 - e)) - (~(m >> (25 - e)) & E); + else if (R == std::round_toward_infinity) + m += ((value >> 15) - 1) & ((1 << (25 - e)) - 1U); + else if (R == std::round_toward_neg_infinity) + m += -(value >> 15) & ((1 << (25 - e)) - 1U); + m >>= 25 - e; + } + else + m <<= e - 25; + return (value & 0x8000) ? -static_cast(m) : static_cast(m); +} + +/// Convert half-precision floating point to integer. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam T type to convert to (buitlin integer type with at least 16 bits precision, excluding any implicit sign +/// bits) \param value binary representation of half-precision value \return integral value +template +T half2int(uint16 value) +{ + return half2int_impl(value); +} + +/// Convert half-precision floating point to integer using round-to-nearest-away-from-zero. +/// \tparam T type to convert to (buitlin integer type with at least 16 bits precision, excluding any implicit sign +/// bits) \param value binary representation of half-precision value \return integral value +template +T half2int_up(uint16 value) +{ + return half2int_impl(value); +} + +/// Round half-precision number to nearest integer value. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \tparam E `true` for round to even, `false` for round away from zero +/// \param value binary representation of half-precision value +/// \return half-precision bits for nearest integral value +template +uint16 round_half_impl(uint16 value) +{ + uint32_t e = value & 0x7FFF; + uint16 result = value; + if (e < 0x3C00) + { + result &= 0x8000; + if (R == std::round_to_nearest) + result |= 0x3C00U & -(e >= (0x3800 + E)); + else if (R == std::round_toward_infinity) + result |= 0x3C00U & -(~(value >> 15) & (e != 0)); + else if (R == std::round_toward_neg_infinity) + result |= 0x3C00U & -(value > 0x8000); + } + else if (e < 0x6400) + { + e = 25 - (e >> 10); + uint32_t mask = (1 << e) - 1; + if (R == std::round_to_nearest) + result += (1 << (e - 1)) - (~(result >> e) & E); + else if (R == std::round_toward_infinity) + result += mask & ((value >> 15) - 1); + else if (R == std::round_toward_neg_infinity) + result += mask & -(value >> 15); + result &= ~mask; + } + return result; +} + +/// Round half-precision number to nearest integer value. +/// \tparam R rounding mode to use, `std::round_indeterminate` for fastest rounding +/// \param value binary representation of half-precision value +/// \return half-precision bits for nearest integral value +template +uint16 round_half(uint16 value) +{ + return round_half_impl(value); +} + +/// Round half-precision number to nearest integer value using round-to-nearest-away-from-zero. +/// \param value binary representation of half-precision value +/// \return half-precision bits for nearest integral value +inline uint16 round_half_up(uint16 value) +{ + return round_half_impl(value); +} +/// \} + +struct functions; +template +struct unary_specialized; +template +struct binary_specialized; +template +struct half_caster; +} // namespace detail + +/// Half-precision floating point type. +/// This class implements an IEEE-conformant half-precision floating point type with the usual arithmetic operators and +/// conversions. It is implicitly convertible to single-precision floating point, which makes artihmetic expressions and +/// functions with mixed-type operands to be of the most precise operand type. Additionally all arithmetic operations +/// (and many mathematical functions) are carried out in single-precision internally. All conversions from single- to +/// half-precision are done using the library's default rounding mode, but temporary results inside chained arithmetic +/// expressions are kept in single-precision as long as possible (while of course still maintaining a strong +/// half-precision type). +/// +/// According to the C++98/03 definition, the half type is not a POD type. But according to C++11's less strict and +/// extended definitions it is both a standard layout type and a trivially copyable type (even if not a POD type), which +/// means it can be standard-conformantly copied using raw binary copies. But in this context some more words about the +/// actual size of the type. Although the half is representing an IEEE 16-bit type, it does not neccessarily have to be +/// of exactly 16-bits size. But on any reasonable implementation the actual binary representation of this type will +/// most probably not ivolve any additional "magic" or padding beyond the simple binary representation of the underlying +/// 16-bit IEEE number, even if not strictly guaranteed by the standard. But even then it only has an actual size of 16 +/// bits if your C++ implementation supports an unsigned integer type of exactly 16 bits width. But this should be the +/// case on nearly any reasonable platform. +/// +/// So if your C++ implementation is not totally exotic or imposes special alignment requirements, it is a reasonable +/// assumption that the data of a half is just comprised of the 2 bytes of the underlying IEEE representation. +class half +{ + friend struct detail::functions; + friend struct detail::unary_specialized; + friend struct detail::binary_specialized; + template + friend struct detail::half_caster; + friend class std::numeric_limits; +#if HALF_ENABLE_CPP11_HASH + friend struct std::hash; +#endif +#if HALF_ENABLE_CPP11_USER_LITERALS + friend half literal::operator"" _h(long double); +#endif + +public: + /// Default constructor. + /// This initializes the half to 0. Although this does not match the builtin types' default-initialization semantics + /// and may be less efficient than no initialization, it is needed to provide proper value-initialization semantics. + HALF_CONSTEXPR half() HALF_NOEXCEPT : data_() {} + + /// Copy constructor. + /// \tparam T type of concrete half expression + /// \param rhs half expression to copy from + half(detail::expr rhs) + : data_(detail::float2half(static_cast(rhs))) + { + } + + /// Conversion constructor. + /// \param rhs float to convert + explicit half(float rhs) + : data_(detail::float2half(rhs)) + { + } + + /// Conversion to single-precision. + /// \return single precision value representing expression value + operator float() const + { + return detail::half2float(data_); + } + + /// Assignment operator. + /// \tparam T type of concrete half expression + /// \param rhs half expression to copy from + /// \return reference to this half + half& operator=(detail::expr rhs) + { + return *this = static_cast(rhs); + } + + /// Arithmetic assignment. + /// \tparam T type of concrete half expression + /// \param rhs half expression to add + /// \return reference to this half + template + typename detail::enable::type operator+=(T rhs) + { + return *this += static_cast(rhs); + } + + /// Arithmetic assignment. + /// \tparam T type of concrete half expression + /// \param rhs half expression to subtract + /// \return reference to this half + template + typename detail::enable::type operator-=(T rhs) + { + return *this -= static_cast(rhs); + } + + /// Arithmetic assignment. + /// \tparam T type of concrete half expression + /// \param rhs half expression to multiply with + /// \return reference to this half + template + typename detail::enable::type operator*=(T rhs) + { + return *this *= static_cast(rhs); + } + + /// Arithmetic assignment. + /// \tparam T type of concrete half expression + /// \param rhs half expression to divide by + /// \return reference to this half + template + typename detail::enable::type operator/=(T rhs) + { + return *this /= static_cast(rhs); + } + + /// Assignment operator. + /// \param rhs single-precision value to copy from + /// \return reference to this half + half& operator=(float rhs) + { + data_ = detail::float2half(rhs); + return *this; + } + + /// Arithmetic assignment. + /// \param rhs single-precision value to add + /// \return reference to this half + half& operator+=(float rhs) + { + data_ = detail::float2half(detail::half2float(data_) + rhs); + return *this; + } + + /// Arithmetic assignment. + /// \param rhs single-precision value to subtract + /// \return reference to this half + half& operator-=(float rhs) + { + data_ = detail::float2half(detail::half2float(data_) - rhs); + return *this; + } + + /// Arithmetic assignment. + /// \param rhs single-precision value to multiply with + /// \return reference to this half + half& operator*=(float rhs) + { + data_ = detail::float2half(detail::half2float(data_) * rhs); + return *this; + } + + /// Arithmetic assignment. + /// \param rhs single-precision value to divide by + /// \return reference to this half + half& operator/=(float rhs) + { + data_ = detail::float2half(detail::half2float(data_) / rhs); + return *this; + } + + /// Prefix increment. + /// \return incremented half value + half& operator++() + { + return *this += 1.0F; + } + + /// Prefix decrement. + /// \return decremented half value + half& operator--() + { + return *this -= 1.0F; + } + + /// Postfix increment. + /// \return non-incremented half value + half operator++(int) + { + half out(*this); + ++*this; + return out; + } + + /// Postfix decrement. + /// \return non-decremented half value + half operator--(int) + { + half out(*this); + --*this; + return out; + } + +private: + /// Rounding mode to use + static const std::float_round_style round_style = (std::float_round_style)(HALF_ROUND_STYLE); + + /// Constructor. + /// \param bits binary representation to set half to + HALF_CONSTEXPR half(detail::binary_t, detail::uint16 bits) HALF_NOEXCEPT : data_(bits) {} + + /// Internal binary representation + detail::uint16 data_; +}; + +#if HALF_ENABLE_CPP11_USER_LITERALS +namespace literal +{ +/// Half literal. +/// While this returns an actual half-precision value, half literals can unfortunately not be constant expressions due +/// to rather involved conversions. +/// \param value literal value +/// \return half with given value (if representable) +inline half operator"" _h(long double value) +{ + return half(detail::binary, detail::float2half(value)); +} +} // namespace literal +#endif + +namespace detail +{ +/// Wrapper implementing unspecialized half-precision functions. +struct functions +{ + /// Addition implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision sum stored in single-precision + static expr plus(float x, float y) + { + return expr(x + y); + } + + /// Subtraction implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision difference stored in single-precision + static expr minus(float x, float y) + { + return expr(x - y); + } + + /// Multiplication implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision product stored in single-precision + static expr multiplies(float x, float y) + { + return expr(x * y); + } + + /// Division implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision quotient stored in single-precision + static expr divides(float x, float y) + { + return expr(x / y); + } + + /// Output implementation. + /// \param out stream to write to + /// \param arg value to write + /// \return reference to stream + template + static std::basic_ostream& write(std::basic_ostream& out, float arg) + { + return out << arg; + } + + /// Input implementation. + /// \param in stream to read from + /// \param arg half to read into + /// \return reference to stream + template + static std::basic_istream& read(std::basic_istream& in, half& arg) + { + float f; + if (in >> f) + arg = f; + return in; + } + + /// Modulo implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision division remainder stored in single-precision + static expr fmod(float x, float y) + { + return expr(std::fmod(x, y)); + } + + /// Remainder implementation. + /// \param x first operand + /// \param y second operand + /// \return Half-precision division remainder stored in single-precision + static expr remainder(float x, float y) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::remainder(x, y)); +#else + if (builtin_isnan(x) || builtin_isnan(y)) + return expr(std::numeric_limits::quiet_NaN()); + float ax = std::fabs(x), ay = std::fabs(y); + if (ax >= 65536.0f || ay < std::ldexp(1.0f, -24)) + return expr(std::numeric_limits::quiet_NaN()); + if (ay >= 65536.0f) + return expr(x); + if (ax == ay) + return expr(builtin_signbit(x) ? -0.0f : 0.0f); + ax = std::fmod(ax, ay + ay); + float y2 = 0.5f * ay; + if (ax > y2) + { + ax -= ay; + if (ax >= y2) + ax -= ay; + } + return expr(builtin_signbit(x) ? -ax : ax); +#endif + } + + /// Remainder implementation. + /// \param x first operand + /// \param y second operand + /// \param quo address to store quotient bits at + /// \return Half-precision division remainder stored in single-precision + static expr remquo(float x, float y, int* quo) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::remquo(x, y, quo)); +#else + if (builtin_isnan(x) || builtin_isnan(y)) + return expr(std::numeric_limits::quiet_NaN()); + bool sign = builtin_signbit(x), qsign = static_cast(sign ^ builtin_signbit(y)); + float ax = std::fabs(x), ay = std::fabs(y); + if (ax >= 65536.0f || ay < std::ldexp(1.0f, -24)) + return expr(std::numeric_limits::quiet_NaN()); + if (ay >= 65536.0f) + return expr(x); + if (ax == ay) + return *quo = qsign ? -1 : 1, expr(sign ? -0.0f : 0.0f); + ax = std::fmod(ax, 8.0f * ay); + int cquo = 0; + if (ax >= 4.0f * ay) + { + ax -= 4.0f * ay; + cquo += 4; + } + if (ax >= 2.0f * ay) + { + ax -= 2.0f * ay; + cquo += 2; + } + float y2 = 0.5f * ay; + if (ax > y2) + { + ax -= ay; + ++cquo; + if (ax >= y2) + { + ax -= ay; + ++cquo; + } + } + return *quo = qsign ? -cquo : cquo, expr(sign ? -ax : ax); +#endif + } + + /// Positive difference implementation. + /// \param x first operand + /// \param y second operand + /// \return Positive difference stored in single-precision + static expr fdim(float x, float y) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::fdim(x, y)); +#else + return expr((x <= y) ? 0.0f : (x - y)); +#endif + } + + /// Fused multiply-add implementation. + /// \param x first operand + /// \param y second operand + /// \param z third operand + /// \return \a x * \a y + \a z stored in single-precision + static expr fma(float x, float y, float z) + { +#if HALF_ENABLE_CPP11_CMATH && defined(FP_FAST_FMAF) + return expr(std::fma(x, y, z)); +#else + return expr(x * y + z); +#endif + } + + /// Get NaN. + /// \return Half-precision quiet NaN + static half nanh() + { + return half(binary, 0x7FFF); + } + + /// Exponential implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr exp(float arg) + { + return expr(std::exp(arg)); + } + + /// Exponential implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr expm1(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::expm1(arg)); +#else + return expr(static_cast(std::exp(static_cast(arg)) - 1.0)); +#endif + } + + /// Binary exponential implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr exp2(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::exp2(arg)); +#else + return expr(static_cast(std::exp(arg * 0.69314718055994530941723212145818))); +#endif + } + + /// Logarithm implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr log(float arg) + { + return expr(std::log(arg)); + } + + /// Common logarithm implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr log10(float arg) + { + return expr(std::log10(arg)); + } + + /// Logarithm implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr log1p(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::log1p(arg)); +#else + return expr(static_cast(std::log(1.0 + arg))); +#endif + } + + /// Binary logarithm implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr log2(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::log2(arg)); +#else + return expr(static_cast(std::log(static_cast(arg)) * 1.4426950408889634073599246810019)); +#endif + } + + /// Square root implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr sqrt(float arg) + { + return expr(std::sqrt(arg)); + } + + /// Cubic root implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr cbrt(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::cbrt(arg)); +#else + if (builtin_isnan(arg) || builtin_isinf(arg)) + return expr(arg); + return expr(builtin_signbit(arg) ? -static_cast(std::pow(-static_cast(arg), 1.0 / 3.0)) + : static_cast(std::pow(static_cast(arg), 1.0 / 3.0))); +#endif + } + + /// Hypotenuse implementation. + /// \param x first argument + /// \param y second argument + /// \return function value stored in single-preicision + static expr hypot(float x, float y) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::hypot(x, y)); +#else + return expr((builtin_isinf(x) || builtin_isinf(y)) + ? std::numeric_limits::infinity() + : static_cast(std::sqrt(static_cast(x) * x + static_cast(y) * y))); +#endif + } + + /// Power implementation. + /// \param base value to exponentiate + /// \param exp power to expontiate to + /// \return function value stored in single-preicision + static expr pow(float base, float exp) + { + return expr(std::pow(base, exp)); + } + + /// Sine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr sin(float arg) + { + return expr(std::sin(arg)); + } + + /// Cosine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr cos(float arg) + { + return expr(std::cos(arg)); + } + + /// Tan implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr tan(float arg) + { + return expr(std::tan(arg)); + } + + /// Arc sine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr asin(float arg) + { + return expr(std::asin(arg)); + } + + /// Arc cosine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr acos(float arg) + { + return expr(std::acos(arg)); + } + + /// Arc tangent implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr atan(float arg) + { + return expr(std::atan(arg)); + } + + /// Arc tangent implementation. + /// \param x first argument + /// \param y second argument + /// \return function value stored in single-preicision + static expr atan2(float x, float y) + { + return expr(std::atan2(x, y)); + } + + /// Hyperbolic sine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr sinh(float arg) + { + return expr(std::sinh(arg)); + } + + /// Hyperbolic cosine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr cosh(float arg) + { + return expr(std::cosh(arg)); + } + + /// Hyperbolic tangent implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr tanh(float arg) + { + return expr(std::tanh(arg)); + } + + /// Hyperbolic area sine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr asinh(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::asinh(arg)); +#else + return expr((arg == -std::numeric_limits::infinity()) + ? arg + : static_cast(std::log(arg + std::sqrt(arg * arg + 1.0)))); +#endif + } + + /// Hyperbolic area cosine implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr acosh(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::acosh(arg)); +#else + return expr((arg < -1.0f) ? std::numeric_limits::quiet_NaN() + : static_cast(std::log(arg + std::sqrt(arg * arg - 1.0)))); +#endif + } + + /// Hyperbolic area tangent implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr atanh(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::atanh(arg)); +#else + return expr(static_cast(0.5 * std::log((1.0 + arg) / (1.0 - arg)))); +#endif + } + + /// Error function implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr erf(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::erf(arg)); +#else + return expr(static_cast(erf(static_cast(arg)))); +#endif + } + + /// Complementary implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr erfc(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::erfc(arg)); +#else + return expr(static_cast(1.0 - erf(static_cast(arg)))); +#endif + } + + /// Gamma logarithm implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr lgamma(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::lgamma(arg)); +#else + if (builtin_isinf(arg)) + return expr(std::numeric_limits::infinity()); + if (arg < 0.0f) + { + float i, f = std::modf(-arg, &i); + if (f == 0.0f) + return expr(std::numeric_limits::infinity()); + return expr(static_cast(1.1447298858494001741434273513531 + - std::log(std::abs(std::sin(3.1415926535897932384626433832795 * f))) - lgamma(1.0 - arg))); + } + return expr(static_cast(lgamma(static_cast(arg)))); +#endif + } + + /// Gamma implementation. + /// \param arg function argument + /// \return function value stored in single-preicision + static expr tgamma(float arg) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::tgamma(arg)); +#else + if (arg == 0.0f) + return builtin_signbit(arg) ? expr(-std::numeric_limits::infinity()) + : expr(std::numeric_limits::infinity()); + if (arg < 0.0f) + { + float i, f = std::modf(-arg, &i); + if (f == 0.0f) + return expr(std::numeric_limits::quiet_NaN()); + double value = 3.1415926535897932384626433832795 + / (std::sin(3.1415926535897932384626433832795 * f) * std::exp(lgamma(1.0 - arg))); + return expr(static_cast((std::fmod(i, 2.0f) == 0.0f) ? -value : value)); + } + if (builtin_isinf(arg)) + return expr(arg); + return expr(static_cast(std::exp(lgamma(static_cast(arg))))); +#endif + } + + /// Floor implementation. + /// \param arg value to round + /// \return rounded value + static half floor(half arg) + { + return half(binary, round_half(arg.data_)); + } + + /// Ceiling implementation. + /// \param arg value to round + /// \return rounded value + static half ceil(half arg) + { + return half(binary, round_half(arg.data_)); + } + + /// Truncation implementation. + /// \param arg value to round + /// \return rounded value + static half trunc(half arg) + { + return half(binary, round_half(arg.data_)); + } + + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static half round(half arg) + { + return half(binary, round_half_up(arg.data_)); + } + + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static long lround(half arg) + { + return detail::half2int_up(arg.data_); + } + + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static half rint(half arg) + { + return half(binary, round_half(arg.data_)); + } + + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static long lrint(half arg) + { + return detail::half2int(arg.data_); + } + +#if HALF_ENABLE_CPP11_LONG_LONG + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static long long llround(half arg) + { + return detail::half2int_up(arg.data_); + } + + /// Nearest integer implementation. + /// \param arg value to round + /// \return rounded value + static long long llrint(half arg) + { + return detail::half2int(arg.data_); + } +#endif + + /// Decompression implementation. + /// \param arg number to decompress + /// \param exp address to store exponent at + /// \return normalized significant + static half frexp(half arg, int* exp) + { + int m = arg.data_ & 0x7FFF, e = -14; + if (m >= 0x7C00 || !m) + return *exp = 0, arg; + for (; m < 0x400; m <<= 1, --e) + ; + return *exp = e + (m >> 10), half(binary, (arg.data_ & 0x8000) | 0x3800 | (m & 0x3FF)); + } + + /// Decompression implementation. + /// \param arg number to decompress + /// \param iptr address to store integer part at + /// \return fractional part + static half modf(half arg, half* iptr) + { + uint32_t e = arg.data_ & 0x7FFF; + if (e >= 0x6400) + return *iptr = arg, half(binary, arg.data_ & (0x8000U | -(e > 0x7C00))); + if (e < 0x3C00) + return iptr->data_ = arg.data_ & 0x8000, arg; + e >>= 10; + uint32_t mask = (1 << (25 - e)) - 1, m = arg.data_ & mask; + iptr->data_ = arg.data_ & ~mask; + if (!m) + return half(binary, arg.data_ & 0x8000); + for (; m < 0x400; m <<= 1, --e) + ; + return half(binary, static_cast((arg.data_ & 0x8000) | (e << 10) | (m & 0x3FF))); + } + + /// Scaling implementation. + /// \param arg number to scale + /// \param exp power of two to scale by + /// \return scaled number + static half scalbln(half arg, long exp) + { + uint32_t m = arg.data_ & 0x7FFF; + if (m >= 0x7C00 || !m) + return arg; + for (; m < 0x400; m <<= 1, --exp) + ; + exp += m >> 10; + uint16 value = arg.data_ & 0x8000; + if (exp > 30) + { + if (half::round_style == std::round_toward_zero) + value |= 0x7BFF; + else if (half::round_style == std::round_toward_infinity) + value |= 0x7C00 - (value >> 15); + else if (half::round_style == std::round_toward_neg_infinity) + value |= 0x7BFF + (value >> 15); + else + value |= 0x7C00; + } + else if (exp > 0) + value |= (exp << 10) | (m & 0x3FF); + else if (exp > -11) + { + m = (m & 0x3FF) | 0x400; + if (half::round_style == std::round_to_nearest) + { + m += 1 << -exp; +#if HALF_ROUND_TIES_TO_EVEN + m -= (m >> (1 - exp)) & 1; +#endif + } + else if (half::round_style == std::round_toward_infinity) + m += ((value >> 15) - 1) & ((1 << (1 - exp)) - 1U); + else if (half::round_style == std::round_toward_neg_infinity) + m += -(value >> 15) & ((1 << (1 - exp)) - 1U); + value |= m >> (1 - exp); + } + else if (half::round_style == std::round_toward_infinity) + value -= (value >> 15) - 1; + else if (half::round_style == std::round_toward_neg_infinity) + value += value >> 15; + return half(binary, value); + } + + /// Exponent implementation. + /// \param arg number to query + /// \return floating point exponent + static int ilogb(half arg) + { + int abs = arg.data_ & 0x7FFF; + if (!abs) + return FP_ILOGB0; + if (abs < 0x7C00) + { + int exp = (abs >> 10) - 15; + if (abs < 0x400) + for (; abs < 0x200; abs <<= 1, --exp) + ; + return exp; + } + if (abs > 0x7C00) + return FP_ILOGBNAN; + return INT_MAX; + } + + /// Exponent implementation. + /// \param arg number to query + /// \return floating point exponent + static half logb(half arg) + { + int abs = arg.data_ & 0x7FFF; + if (!abs) + return half(binary, 0xFC00); + if (abs < 0x7C00) + { + int exp = (abs >> 10) - 15; + if (abs < 0x400) + for (; abs < 0x200; abs <<= 1, --exp) + ; + uint16 bits = (exp < 0) << 15; + if (exp) + { + uint32_t m = std::abs(exp) << 6, e = 18; + for (; m < 0x400; m <<= 1, --e) + ; + bits |= (e << 10) + m; + } + return half(binary, bits); + } + if (abs > 0x7C00) + return arg; + return half(binary, 0x7C00); + } + + /// Enumeration implementation. + /// \param from number to increase/decrease + /// \param to direction to enumerate into + /// \return next representable number + static half nextafter(half from, half to) + { + uint16 fabs = from.data_ & 0x7FFF, tabs = to.data_ & 0x7FFF; + if (fabs > 0x7C00) + return from; + if (tabs > 0x7C00 || from.data_ == to.data_ || !(fabs | tabs)) + return to; + if (!fabs) + return half(binary, (to.data_ & 0x8000) + 1); + bool lt = ((fabs == from.data_) ? static_cast(fabs) : -static_cast(fabs)) + < ((tabs == to.data_) ? static_cast(tabs) : -static_cast(tabs)); + return half(binary, from.data_ + (((from.data_ >> 15) ^ static_cast(lt)) << 1) - 1); + } + + /// Enumeration implementation. + /// \param from number to increase/decrease + /// \param to direction to enumerate into + /// \return next representable number + static half nexttoward(half from, long double to) + { + if (isnan(from)) + return from; + long double lfrom = static_cast(from); + if (builtin_isnan(to) || lfrom == to) + return half(static_cast(to)); + if (!(from.data_ & 0x7FFF)) + return half(binary, (static_cast(builtin_signbit(to)) << 15) + 1); + return half(binary, from.data_ + (((from.data_ >> 15) ^ static_cast(lfrom < to)) << 1) - 1); + } + + /// Sign implementation + /// \param x first operand + /// \param y second operand + /// \return composed value + static half copysign(half x, half y) + { + return half(binary, x.data_ ^ ((x.data_ ^ y.data_) & 0x8000)); + } + + /// Classification implementation. + /// \param arg value to classify + /// \retval true if infinite number + /// \retval false else + static int fpclassify(half arg) + { + uint32_t abs = arg.data_ & 0x7FFF; + return abs + ? ((abs > 0x3FF) ? ((abs >= 0x7C00) ? ((abs > 0x7C00) ? FP_NAN : FP_INFINITE) : FP_NORMAL) : FP_SUBNORMAL) + : FP_ZERO; + } + + /// Classification implementation. + /// \param arg value to classify + /// \retval true if finite number + /// \retval false else + static bool isfinite(half arg) + { + return (arg.data_ & 0x7C00) != 0x7C00; + } + + /// Classification implementation. + /// \param arg value to classify + /// \retval true if infinite number + /// \retval false else + static bool isinf(half arg) + { + return (arg.data_ & 0x7FFF) == 0x7C00; + } + + /// Classification implementation. + /// \param arg value to classify + /// \retval true if not a number + /// \retval false else + static bool isnan(half arg) + { + return (arg.data_ & 0x7FFF) > 0x7C00; + } + + /// Classification implementation. + /// \param arg value to classify + /// \retval true if normal number + /// \retval false else + static bool isnormal(half arg) + { + return ((arg.data_ & 0x7C00) != 0) & ((arg.data_ & 0x7C00) != 0x7C00); + } + + /// Sign bit implementation. + /// \param arg value to check + /// \retval true if signed + /// \retval false if unsigned + static bool signbit(half arg) + { + return (arg.data_ & 0x8000) != 0; + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if operands equal + /// \retval false else + static bool isequal(half x, half y) + { + return (x.data_ == y.data_ || !((x.data_ | y.data_) & 0x7FFF)) && !isnan(x); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if operands not equal + /// \retval false else + static bool isnotequal(half x, half y) + { + return (x.data_ != y.data_ && ((x.data_ | y.data_) & 0x7FFF)) || isnan(x); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if \a x > \a y + /// \retval false else + static bool isgreater(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + return xabs <= 0x7C00 && yabs <= 0x7C00 + && (((xabs == x.data_) ? xabs : -xabs) > ((yabs == y.data_) ? yabs : -yabs)); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if \a x >= \a y + /// \retval false else + static bool isgreaterequal(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + return xabs <= 0x7C00 && yabs <= 0x7C00 + && (((xabs == x.data_) ? xabs : -xabs) >= ((yabs == y.data_) ? yabs : -yabs)); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if \a x < \a y + /// \retval false else + static bool isless(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + return xabs <= 0x7C00 && yabs <= 0x7C00 + && (((xabs == x.data_) ? xabs : -xabs) < ((yabs == y.data_) ? yabs : -yabs)); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if \a x <= \a y + /// \retval false else + static bool islessequal(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + return xabs <= 0x7C00 && yabs <= 0x7C00 + && (((xabs == x.data_) ? xabs : -xabs) <= ((yabs == y.data_) ? yabs : -yabs)); + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if either \a x > \a y nor \a x < \a y + /// \retval false else + static bool islessgreater(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + if (xabs > 0x7C00 || yabs > 0x7C00) + return false; + int a = (xabs == x.data_) ? xabs : -xabs, b = (yabs == y.data_) ? yabs : -yabs; + return a < b || a > b; + } + + /// Comparison implementation. + /// \param x first operand + /// \param y second operand + /// \retval true if operand unordered + /// \retval false else + static bool isunordered(half x, half y) + { + return isnan(x) || isnan(y); + } + +private: + static double erf(double arg) + { + if (builtin_isinf(arg)) + return (arg < 0.0) ? -1.0 : 1.0; + double x2 = arg * arg, ax2 = 0.147 * x2, + value = std::sqrt(1.0 - std::exp(-x2 * (1.2732395447351626861510701069801 + ax2) / (1.0 + ax2))); + return builtin_signbit(arg) ? -value : value; + } + + static double lgamma(double arg) + { + double v = 1.0; + for (; arg < 8.0; ++arg) + v *= arg; + double w = 1.0 / (arg * arg); + return (((((((-0.02955065359477124183006535947712 * w + 0.00641025641025641025641025641026) * w + + -0.00191752691752691752691752691753) + * w + + 8.4175084175084175084175084175084e-4) + * w + + -5.952380952380952380952380952381e-4) + * w + + 7.9365079365079365079365079365079e-4) + * w + + -0.00277777777777777777777777777778) + * w + + 0.08333333333333333333333333333333) + / arg + + 0.91893853320467274178032973640562 - std::log(v) - arg + (arg - 0.5) * std::log(arg); + } +}; + +/// Wrapper for unary half-precision functions needing specialization for individual argument types. +/// \tparam T argument type +template +struct unary_specialized +{ + /// Negation implementation. + /// \param arg value to negate + /// \return negated value + static HALF_CONSTEXPR half negate(half arg) + { + return half(binary, arg.data_ ^ 0x8000); + } + + /// Absolute value implementation. + /// \param arg function argument + /// \return absolute value + static half fabs(half arg) + { + return half(binary, arg.data_ & 0x7FFF); + } +}; +template <> +struct unary_specialized +{ + static HALF_CONSTEXPR expr negate(float arg) + { + return expr(-arg); + } + static expr fabs(float arg) + { + return expr(std::fabs(arg)); + } +}; + +/// Wrapper for binary half-precision functions needing specialization for individual argument types. +/// \tparam T first argument type +/// \tparam U first argument type +template +struct binary_specialized +{ + /// Minimum implementation. + /// \param x first operand + /// \param y second operand + /// \return minimum value + static expr fmin(float x, float y) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::fmin(x, y)); +#else + if (builtin_isnan(x)) + return expr(y); + if (builtin_isnan(y)) + return expr(x); + return expr(std::min(x, y)); +#endif + } + + /// Maximum implementation. + /// \param x first operand + /// \param y second operand + /// \return maximum value + static expr fmax(float x, float y) + { +#if HALF_ENABLE_CPP11_CMATH + return expr(std::fmax(x, y)); +#else + if (builtin_isnan(x)) + return expr(y); + if (builtin_isnan(y)) + return expr(x); + return expr(std::max(x, y)); +#endif + } +}; +template <> +struct binary_specialized +{ + static half fmin(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + if (xabs > 0x7C00) + return y; + if (yabs > 0x7C00) + return x; + return (((xabs == x.data_) ? xabs : -xabs) > ((yabs == y.data_) ? yabs : -yabs)) ? y : x; + } + static half fmax(half x, half y) + { + int xabs = x.data_ & 0x7FFF, yabs = y.data_ & 0x7FFF; + if (xabs > 0x7C00) + return y; + if (yabs > 0x7C00) + return x; + return (((xabs == x.data_) ? xabs : -xabs) < ((yabs == y.data_) ? yabs : -yabs)) ? y : x; + } +}; + +/// Helper class for half casts. +/// This class template has to be specialized for all valid cast argument to define an appropriate static `cast` member +/// function and a corresponding `type` member denoting its return type. +/// \tparam T destination type +/// \tparam U source type +/// \tparam R rounding mode to use +template +struct half_caster +{ +}; +template +struct half_caster +{ +#if HALF_ENABLE_CPP11_STATIC_ASSERT && HALF_ENABLE_CPP11_TYPE_TRAITS + static_assert(std::is_arithmetic::value, "half_cast from non-arithmetic type unsupported"); +#endif + + static half cast(U arg) + { + return cast_impl(arg, is_float()); + } + +private: + static half cast_impl(U arg, true_type) + { + return half(binary, float2half(arg)); + } + static half cast_impl(U arg, false_type) + { + return half(binary, int2half(arg)); + } +}; +template +struct half_caster +{ +#if HALF_ENABLE_CPP11_STATIC_ASSERT && HALF_ENABLE_CPP11_TYPE_TRAITS + static_assert(std::is_arithmetic::value, "half_cast to non-arithmetic type unsupported"); +#endif + + static T cast(half arg) + { + return cast_impl(arg, is_float()); + } + +private: + static T cast_impl(half arg, true_type) + { + return half2float(arg.data_); + } + static T cast_impl(half arg, false_type) + { + return half2int(arg.data_); + } +}; +template +struct half_caster +{ +#if HALF_ENABLE_CPP11_STATIC_ASSERT && HALF_ENABLE_CPP11_TYPE_TRAITS + static_assert(std::is_arithmetic::value, "half_cast to non-arithmetic type unsupported"); +#endif + + static T cast(expr arg) + { + return cast_impl(arg, is_float()); + } + +private: + static T cast_impl(float arg, true_type) + { + return static_cast(arg); + } + static T cast_impl(half arg, false_type) + { + return half2int(arg.data_); + } +}; +template +struct half_caster +{ + static half cast(half arg) + { + return arg; + } +}; +template +struct half_caster : half_caster +{ +}; + +/// \name Comparison operators +/// \{ + +/// Comparison for equality. +/// \param x first operand +/// \param y second operand +/// \retval true if operands equal +/// \retval false else +template +typename enable::type operator==(T x, U y) +{ + return functions::isequal(x, y); +} + +/// Comparison for inequality. +/// \param x first operand +/// \param y second operand +/// \retval true if operands not equal +/// \retval false else +template +typename enable::type operator!=(T x, U y) +{ + return functions::isnotequal(x, y); +} + +/// Comparison for less than. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x less than \a y +/// \retval false else +template +typename enable::type operator<(T x, U y) +{ + return functions::isless(x, y); +} + +/// Comparison for greater than. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x greater than \a y +/// \retval false else +template +typename enable::type operator>(T x, U y) +{ + return functions::isgreater(x, y); +} + +/// Comparison for less equal. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x less equal \a y +/// \retval false else +template +typename enable::type operator<=(T x, U y) +{ + return functions::islessequal(x, y); +} + +/// Comparison for greater equal. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x greater equal \a y +/// \retval false else +template +typename enable::type operator>=(T x, U y) +{ + return functions::isgreaterequal(x, y); +} + +/// \} +/// \name Arithmetic operators +/// \{ + +/// Add halfs. +/// \param x left operand +/// \param y right operand +/// \return sum of half expressions +template +typename enable::type operator+(T x, U y) +{ + return functions::plus(x, y); +} + +/// Subtract halfs. +/// \param x left operand +/// \param y right operand +/// \return difference of half expressions +template +typename enable::type operator-(T x, U y) +{ + return functions::minus(x, y); +} + +/// Multiply halfs. +/// \param x left operand +/// \param y right operand +/// \return product of half expressions +template +typename enable::type operator*(T x, U y) +{ + return functions::multiplies(x, y); +} + +/// Divide halfs. +/// \param x left operand +/// \param y right operand +/// \return quotient of half expressions +template +typename enable::type operator/(T x, U y) +{ + return functions::divides(x, y); +} + +/// Identity. +/// \param arg operand +/// \return uncahnged operand +template +HALF_CONSTEXPR typename enable::type operator+(T arg) +{ + return arg; +} + +/// Negation. +/// \param arg operand +/// \return negated operand +template +HALF_CONSTEXPR typename enable::type operator-(T arg) +{ + return unary_specialized::negate(arg); +} + +/// \} +/// \name Input and output +/// \{ + +/// Output operator. +/// \param out output stream to write into +/// \param arg half expression to write +/// \return reference to output stream +template +typename enable&, T>::type operator<<(std::basic_ostream& out, T arg) +{ + return functions::write(out, arg); +} + +/// Input operator. +/// \param in input stream to read from +/// \param arg half to read into +/// \return reference to input stream +template +std::basic_istream& operator>>(std::basic_istream& in, half& arg) +{ + return functions::read(in, arg); +} + +/// \} +/// \name Basic mathematical operations +/// \{ + +/// Absolute value. +/// \param arg operand +/// \return absolute value of \a arg +// template typename enable::type abs(T arg) { return unary_specialized::fabs(arg); } +inline half abs(half arg) +{ + return unary_specialized::fabs(arg); +} +inline expr abs(expr arg) +{ + return unary_specialized::fabs(arg); +} + +/// Absolute value. +/// \param arg operand +/// \return absolute value of \a arg +// template typename enable::type fabs(T arg) { return unary_specialized::fabs(arg); } +inline half fabs(half arg) +{ + return unary_specialized::fabs(arg); +} +inline expr fabs(expr arg) +{ + return unary_specialized::fabs(arg); +} + +/// Remainder of division. +/// \param x first operand +/// \param y second operand +/// \return remainder of floating point division. +// template typename enable::type fmod(T x, U y) { return functions::fmod(x, y); } +inline expr fmod(half x, half y) +{ + return functions::fmod(x, y); +} +inline expr fmod(half x, expr y) +{ + return functions::fmod(x, y); +} +inline expr fmod(expr x, half y) +{ + return functions::fmod(x, y); +} +inline expr fmod(expr x, expr y) +{ + return functions::fmod(x, y); +} + +/// Remainder of division. +/// \param x first operand +/// \param y second operand +/// \return remainder of floating point division. +// template typename enable::type remainder(T x, U y) { return +// functions::remainder(x, y); } +inline expr remainder(half x, half y) +{ + return functions::remainder(x, y); +} +inline expr remainder(half x, expr y) +{ + return functions::remainder(x, y); +} +inline expr remainder(expr x, half y) +{ + return functions::remainder(x, y); +} +inline expr remainder(expr x, expr y) +{ + return functions::remainder(x, y); +} + +/// Remainder of division. +/// \param x first operand +/// \param y second operand +/// \param quo address to store some bits of quotient at +/// \return remainder of floating point division. +// template typename enable::type remquo(T x, U y, int *quo) { return +// functions::remquo(x, y, quo); } +inline expr remquo(half x, half y, int* quo) +{ + return functions::remquo(x, y, quo); +} +inline expr remquo(half x, expr y, int* quo) +{ + return functions::remquo(x, y, quo); +} +inline expr remquo(expr x, half y, int* quo) +{ + return functions::remquo(x, y, quo); +} +inline expr remquo(expr x, expr y, int* quo) +{ + return functions::remquo(x, y, quo); +} + +/// Fused multiply add. +/// \param x first operand +/// \param y second operand +/// \param z third operand +/// \return ( \a x * \a y ) + \a z rounded as one operation. +// template typename enable::type fma(T x, U y, V z) { return +// functions::fma(x, y, z); } +inline expr fma(half x, half y, half z) +{ + return functions::fma(x, y, z); +} +inline expr fma(half x, half y, expr z) +{ + return functions::fma(x, y, z); +} +inline expr fma(half x, expr y, half z) +{ + return functions::fma(x, y, z); +} +inline expr fma(half x, expr y, expr z) +{ + return functions::fma(x, y, z); +} +inline expr fma(expr x, half y, half z) +{ + return functions::fma(x, y, z); +} +inline expr fma(expr x, half y, expr z) +{ + return functions::fma(x, y, z); +} +inline expr fma(expr x, expr y, half z) +{ + return functions::fma(x, y, z); +} +inline expr fma(expr x, expr y, expr z) +{ + return functions::fma(x, y, z); +} + +/// Maximum of half expressions. +/// \param x first operand +/// \param y second operand +/// \return maximum of operands +// template typename result::type fmax(T x, U y) { return +// binary_specialized::fmax(x, y); } +inline half fmax(half x, half y) +{ + return binary_specialized::fmax(x, y); +} +inline expr fmax(half x, expr y) +{ + return binary_specialized::fmax(x, y); +} +inline expr fmax(expr x, half y) +{ + return binary_specialized::fmax(x, y); +} +inline expr fmax(expr x, expr y) +{ + return binary_specialized::fmax(x, y); +} + +/// Minimum of half expressions. +/// \param x first operand +/// \param y second operand +/// \return minimum of operands +// template typename result::type fmin(T x, U y) { return +// binary_specialized::fmin(x, y); } +inline half fmin(half x, half y) +{ + return binary_specialized::fmin(x, y); +} +inline expr fmin(half x, expr y) +{ + return binary_specialized::fmin(x, y); +} +inline expr fmin(expr x, half y) +{ + return binary_specialized::fmin(x, y); +} +inline expr fmin(expr x, expr y) +{ + return binary_specialized::fmin(x, y); +} + +/// Positive difference. +/// \param x first operand +/// \param y second operand +/// \return \a x - \a y or 0 if difference negative +// template typename enable::type fdim(T x, U y) { return functions::fdim(x, y); } +inline expr fdim(half x, half y) +{ + return functions::fdim(x, y); +} +inline expr fdim(half x, expr y) +{ + return functions::fdim(x, y); +} +inline expr fdim(expr x, half y) +{ + return functions::fdim(x, y); +} +inline expr fdim(expr x, expr y) +{ + return functions::fdim(x, y); +} + +/// Get NaN value. +/// \return quiet NaN +inline half nanh(char const*) +{ + return functions::nanh(); +} + +/// \} +/// \name Exponential functions +/// \{ + +/// Exponential function. +/// \param arg function argument +/// \return e raised to \a arg +// template typename enable::type exp(T arg) { return functions::exp(arg); } +inline expr exp(half arg) +{ + return functions::exp(arg); +} +inline expr exp(expr arg) +{ + return functions::exp(arg); +} + +/// Exponential minus one. +/// \param arg function argument +/// \return e raised to \a arg subtracted by 1 +// template typename enable::type expm1(T arg) { return functions::expm1(arg); } +inline expr expm1(half arg) +{ + return functions::expm1(arg); +} +inline expr expm1(expr arg) +{ + return functions::expm1(arg); +} + +/// Binary exponential. +/// \param arg function argument +/// \return 2 raised to \a arg +// template typename enable::type exp2(T arg) { return functions::exp2(arg); } +inline expr exp2(half arg) +{ + return functions::exp2(arg); +} +inline expr exp2(expr arg) +{ + return functions::exp2(arg); +} + +/// Natural logorithm. +/// \param arg function argument +/// \return logarithm of \a arg to base e +// template typename enable::type log(T arg) { return functions::log(arg); } +inline expr log(half arg) +{ + return functions::log(arg); +} +inline expr log(expr arg) +{ + return functions::log(arg); +} + +/// Common logorithm. +/// \param arg function argument +/// \return logarithm of \a arg to base 10 +// template typename enable::type log10(T arg) { return functions::log10(arg); } +inline expr log10(half arg) +{ + return functions::log10(arg); +} +inline expr log10(expr arg) +{ + return functions::log10(arg); +} + +/// Natural logorithm. +/// \param arg function argument +/// \return logarithm of \a arg plus 1 to base e +// template typename enable::type log1p(T arg) { return functions::log1p(arg); } +inline expr log1p(half arg) +{ + return functions::log1p(arg); +} +inline expr log1p(expr arg) +{ + return functions::log1p(arg); +} + +/// Binary logorithm. +/// \param arg function argument +/// \return logarithm of \a arg to base 2 +// template typename enable::type log2(T arg) { return functions::log2(arg); } +inline expr log2(half arg) +{ + return functions::log2(arg); +} +inline expr log2(expr arg) +{ + return functions::log2(arg); +} + +/// \} +/// \name Power functions +/// \{ + +/// Square root. +/// \param arg function argument +/// \return square root of \a arg +// template typename enable::type sqrt(T arg) { return functions::sqrt(arg); } +inline expr sqrt(half arg) +{ + return functions::sqrt(arg); +} +inline expr sqrt(expr arg) +{ + return functions::sqrt(arg); +} + +/// Cubic root. +/// \param arg function argument +/// \return cubic root of \a arg +// template typename enable::type cbrt(T arg) { return functions::cbrt(arg); } +inline expr cbrt(half arg) +{ + return functions::cbrt(arg); +} +inline expr cbrt(expr arg) +{ + return functions::cbrt(arg); +} + +/// Hypotenuse function. +/// \param x first argument +/// \param y second argument +/// \return square root of sum of squares without internal over- or underflows +// template typename enable::type hypot(T x, U y) { return functions::hypot(x, y); +//} +inline expr hypot(half x, half y) +{ + return functions::hypot(x, y); +} +inline expr hypot(half x, expr y) +{ + return functions::hypot(x, y); +} +inline expr hypot(expr x, half y) +{ + return functions::hypot(x, y); +} +inline expr hypot(expr x, expr y) +{ + return functions::hypot(x, y); +} + +/// Power function. +/// \param base first argument +/// \param exp second argument +/// \return \a base raised to \a exp +// template typename enable::type pow(T base, U exp) { return functions::pow(base, +// exp); } +inline expr pow(half base, half exp) +{ + return functions::pow(base, exp); +} +inline expr pow(half base, expr exp) +{ + return functions::pow(base, exp); +} +inline expr pow(expr base, half exp) +{ + return functions::pow(base, exp); +} +inline expr pow(expr base, expr exp) +{ + return functions::pow(base, exp); +} + +/// \} +/// \name Trigonometric functions +/// \{ + +/// Sine function. +/// \param arg function argument +/// \return sine value of \a arg +// template typename enable::type sin(T arg) { return functions::sin(arg); } +inline expr sin(half arg) +{ + return functions::sin(arg); +} +inline expr sin(expr arg) +{ + return functions::sin(arg); +} + +/// Cosine function. +/// \param arg function argument +/// \return cosine value of \a arg +// template typename enable::type cos(T arg) { return functions::cos(arg); } +inline expr cos(half arg) +{ + return functions::cos(arg); +} +inline expr cos(expr arg) +{ + return functions::cos(arg); +} + +/// Tangent function. +/// \param arg function argument +/// \return tangent value of \a arg +// template typename enable::type tan(T arg) { return functions::tan(arg); } +inline expr tan(half arg) +{ + return functions::tan(arg); +} +inline expr tan(expr arg) +{ + return functions::tan(arg); +} + +/// Arc sine. +/// \param arg function argument +/// \return arc sine value of \a arg +// template typename enable::type asin(T arg) { return functions::asin(arg); } +inline expr asin(half arg) +{ + return functions::asin(arg); +} +inline expr asin(expr arg) +{ + return functions::asin(arg); +} + +/// Arc cosine function. +/// \param arg function argument +/// \return arc cosine value of \a arg +// template typename enable::type acos(T arg) { return functions::acos(arg); } +inline expr acos(half arg) +{ + return functions::acos(arg); +} +inline expr acos(expr arg) +{ + return functions::acos(arg); +} + +/// Arc tangent function. +/// \param arg function argument +/// \return arc tangent value of \a arg +// template typename enable::type atan(T arg) { return functions::atan(arg); } +inline expr atan(half arg) +{ + return functions::atan(arg); +} +inline expr atan(expr arg) +{ + return functions::atan(arg); +} + +/// Arc tangent function. +/// \param x first argument +/// \param y second argument +/// \return arc tangent value +// template typename enable::type atan2(T x, U y) { return functions::atan2(x, y); +//} +inline expr atan2(half x, half y) +{ + return functions::atan2(x, y); +} +inline expr atan2(half x, expr y) +{ + return functions::atan2(x, y); +} +inline expr atan2(expr x, half y) +{ + return functions::atan2(x, y); +} +inline expr atan2(expr x, expr y) +{ + return functions::atan2(x, y); +} + +/// \} +/// \name Hyperbolic functions +/// \{ + +/// Hyperbolic sine. +/// \param arg function argument +/// \return hyperbolic sine value of \a arg +// template typename enable::type sinh(T arg) { return functions::sinh(arg); } +inline expr sinh(half arg) +{ + return functions::sinh(arg); +} +inline expr sinh(expr arg) +{ + return functions::sinh(arg); +} + +/// Hyperbolic cosine. +/// \param arg function argument +/// \return hyperbolic cosine value of \a arg +// template typename enable::type cosh(T arg) { return functions::cosh(arg); } +inline expr cosh(half arg) +{ + return functions::cosh(arg); +} +inline expr cosh(expr arg) +{ + return functions::cosh(arg); +} + +/// Hyperbolic tangent. +/// \param arg function argument +/// \return hyperbolic tangent value of \a arg +// template typename enable::type tanh(T arg) { return functions::tanh(arg); } +inline expr tanh(half arg) +{ + return functions::tanh(arg); +} +inline expr tanh(expr arg) +{ + return functions::tanh(arg); +} + +/// Hyperbolic area sine. +/// \param arg function argument +/// \return area sine value of \a arg +// template typename enable::type asinh(T arg) { return functions::asinh(arg); } +inline expr asinh(half arg) +{ + return functions::asinh(arg); +} +inline expr asinh(expr arg) +{ + return functions::asinh(arg); +} + +/// Hyperbolic area cosine. +/// \param arg function argument +/// \return area cosine value of \a arg +// template typename enable::type acosh(T arg) { return functions::acosh(arg); } +inline expr acosh(half arg) +{ + return functions::acosh(arg); +} +inline expr acosh(expr arg) +{ + return functions::acosh(arg); +} + +/// Hyperbolic area tangent. +/// \param arg function argument +/// \return area tangent value of \a arg +// template typename enable::type atanh(T arg) { return functions::atanh(arg); } +inline expr atanh(half arg) +{ + return functions::atanh(arg); +} +inline expr atanh(expr arg) +{ + return functions::atanh(arg); +} + +/// \} +/// \name Error and gamma functions +/// \{ + +/// Error function. +/// \param arg function argument +/// \return error function value of \a arg +// template typename enable::type erf(T arg) { return functions::erf(arg); } +inline expr erf(half arg) +{ + return functions::erf(arg); +} +inline expr erf(expr arg) +{ + return functions::erf(arg); +} + +/// Complementary error function. +/// \param arg function argument +/// \return 1 minus error function value of \a arg +// template typename enable::type erfc(T arg) { return functions::erfc(arg); } +inline expr erfc(half arg) +{ + return functions::erfc(arg); +} +inline expr erfc(expr arg) +{ + return functions::erfc(arg); +} + +/// Natural logarithm of gamma function. +/// \param arg function argument +/// \return natural logarith of gamma function for \a arg +// template typename enable::type lgamma(T arg) { return functions::lgamma(arg); } +inline expr lgamma(half arg) +{ + return functions::lgamma(arg); +} +inline expr lgamma(expr arg) +{ + return functions::lgamma(arg); +} + +/// Gamma function. +/// \param arg function argument +/// \return gamma function value of \a arg +// template typename enable::type tgamma(T arg) { return functions::tgamma(arg); } +inline expr tgamma(half arg) +{ + return functions::tgamma(arg); +} +inline expr tgamma(expr arg) +{ + return functions::tgamma(arg); +} + +/// \} +/// \name Rounding +/// \{ + +/// Nearest integer not less than half value. +/// \param arg half to round +/// \return nearest integer not less than \a arg +// template typename enable::type ceil(T arg) { return functions::ceil(arg); } +inline half ceil(half arg) +{ + return functions::ceil(arg); +} +inline half ceil(expr arg) +{ + return functions::ceil(arg); +} + +/// Nearest integer not greater than half value. +/// \param arg half to round +/// \return nearest integer not greater than \a arg +// template typename enable::type floor(T arg) { return functions::floor(arg); } +inline half floor(half arg) +{ + return functions::floor(arg); +} +inline half floor(expr arg) +{ + return functions::floor(arg); +} + +/// Nearest integer not greater in magnitude than half value. +/// \param arg half to round +/// \return nearest integer not greater in magnitude than \a arg +// template typename enable::type trunc(T arg) { return functions::trunc(arg); } +inline half trunc(half arg) +{ + return functions::trunc(arg); +} +inline half trunc(expr arg) +{ + return functions::trunc(arg); +} + +/// Nearest integer. +/// \param arg half to round +/// \return nearest integer, rounded away from zero in half-way cases +// template typename enable::type round(T arg) { return functions::round(arg); } +inline half round(half arg) +{ + return functions::round(arg); +} +inline half round(expr arg) +{ + return functions::round(arg); +} + +/// Nearest integer. +/// \param arg half to round +/// \return nearest integer, rounded away from zero in half-way cases +// template typename enable::type lround(T arg) { return functions::lround(arg); } +inline long lround(half arg) +{ + return functions::lround(arg); +} +inline long lround(expr arg) +{ + return functions::lround(arg); +} + +/// Nearest integer using half's internal rounding mode. +/// \param arg half expression to round +/// \return nearest integer using default rounding mode +// template typename enable::type nearbyint(T arg) { return functions::nearbyint(arg); } +inline half nearbyint(half arg) +{ + return functions::rint(arg); +} +inline half nearbyint(expr arg) +{ + return functions::rint(arg); +} + +/// Nearest integer using half's internal rounding mode. +/// \param arg half expression to round +/// \return nearest integer using default rounding mode +// template typename enable::type rint(T arg) { return functions::rint(arg); } +inline half rint(half arg) +{ + return functions::rint(arg); +} +inline half rint(expr arg) +{ + return functions::rint(arg); +} + +/// Nearest integer using half's internal rounding mode. +/// \param arg half expression to round +/// \return nearest integer using default rounding mode +// template typename enable::type lrint(T arg) { return functions::lrint(arg); } +inline long lrint(half arg) +{ + return functions::lrint(arg); +} +inline long lrint(expr arg) +{ + return functions::lrint(arg); +} +#if HALF_ENABLE_CPP11_LONG_LONG +/// Nearest integer. +/// \param arg half to round +/// \return nearest integer, rounded away from zero in half-way cases +// template typename enable::type llround(T arg) { return functions::llround(arg); } +inline long long llround(half arg) +{ + return functions::llround(arg); +} +inline long long llround(expr arg) +{ + return functions::llround(arg); +} + +/// Nearest integer using half's internal rounding mode. +/// \param arg half expression to round +/// \return nearest integer using default rounding mode +// template typename enable::type llrint(T arg) { return functions::llrint(arg); } +inline long long llrint(half arg) +{ + return functions::llrint(arg); +} +inline long long llrint(expr arg) +{ + return functions::llrint(arg); +} +#endif + +/// \} +/// \name Floating point manipulation +/// \{ + +/// Decompress floating point number. +/// \param arg number to decompress +/// \param exp address to store exponent at +/// \return significant in range [0.5, 1) +// template typename enable::type frexp(T arg, int *exp) { return functions::frexp(arg, exp); } +inline half frexp(half arg, int* exp) +{ + return functions::frexp(arg, exp); +} +inline half frexp(expr arg, int* exp) +{ + return functions::frexp(arg, exp); +} + +/// Multiply by power of two. +/// \param arg number to modify +/// \param exp power of two to multiply with +/// \return \a arg multplied by 2 raised to \a exp +// template typename enable::type ldexp(T arg, int exp) { return functions::scalbln(arg, exp); +//} +inline half ldexp(half arg, int exp) +{ + return functions::scalbln(arg, exp); +} +inline half ldexp(expr arg, int exp) +{ + return functions::scalbln(arg, exp); +} + +/// Extract integer and fractional parts. +/// \param arg number to decompress +/// \param iptr address to store integer part at +/// \return fractional part +// template typename enable::type modf(T arg, half *iptr) { return functions::modf(arg, iptr); +//} +inline half modf(half arg, half* iptr) +{ + return functions::modf(arg, iptr); +} +inline half modf(expr arg, half* iptr) +{ + return functions::modf(arg, iptr); +} + +/// Multiply by power of two. +/// \param arg number to modify +/// \param exp power of two to multiply with +/// \return \a arg multplied by 2 raised to \a exp +// template typename enable::type scalbn(T arg, int exp) { return functions::scalbln(arg, exp); +//} +inline half scalbn(half arg, int exp) +{ + return functions::scalbln(arg, exp); +} +inline half scalbn(expr arg, int exp) +{ + return functions::scalbln(arg, exp); +} + +/// Multiply by power of two. +/// \param arg number to modify +/// \param exp power of two to multiply with +/// \return \a arg multplied by 2 raised to \a exp +// template typename enable::type scalbln(T arg, long exp) { return functions::scalbln(arg, +// exp); +//} +inline half scalbln(half arg, long exp) +{ + return functions::scalbln(arg, exp); +} +inline half scalbln(expr arg, long exp) +{ + return functions::scalbln(arg, exp); +} + +/// Extract exponent. +/// \param arg number to query +/// \return floating point exponent +/// \retval FP_ILOGB0 for zero +/// \retval FP_ILOGBNAN for NaN +/// \retval MAX_INT for infinity +// template typename enable::type ilogb(T arg) { return functions::ilogb(arg); } +inline int ilogb(half arg) +{ + return functions::ilogb(arg); +} +inline int ilogb(expr arg) +{ + return functions::ilogb(arg); +} + +/// Extract exponent. +/// \param arg number to query +/// \return floating point exponent +// template typename enable::type logb(T arg) { return functions::logb(arg); } +inline half logb(half arg) +{ + return functions::logb(arg); +} +inline half logb(expr arg) +{ + return functions::logb(arg); +} + +/// Next representable value. +/// \param from value to compute next representable value for +/// \param to direction towards which to compute next value +/// \return next representable value after \a from in direction towards \a to +// template typename enable::type nextafter(T from, U to) { return +// functions::nextafter(from, to); } +inline half nextafter(half from, half to) +{ + return functions::nextafter(from, to); +} +inline half nextafter(half from, expr to) +{ + return functions::nextafter(from, to); +} +inline half nextafter(expr from, half to) +{ + return functions::nextafter(from, to); +} +inline half nextafter(expr from, expr to) +{ + return functions::nextafter(from, to); +} + +/// Next representable value. +/// \param from value to compute next representable value for +/// \param to direction towards which to compute next value +/// \return next representable value after \a from in direction towards \a to +// template typename enable::type nexttoward(T from, long double to) { return +// functions::nexttoward(from, to); } +inline half nexttoward(half from, long double to) +{ + return functions::nexttoward(from, to); +} +inline half nexttoward(expr from, long double to) +{ + return functions::nexttoward(from, to); +} + +/// Take sign. +/// \param x value to change sign for +/// \param y value to take sign from +/// \return value equal to \a x in magnitude and to \a y in sign +// template typename enable::type copysign(T x, U y) { return +// functions::copysign(x, y); } +inline half copysign(half x, half y) +{ + return functions::copysign(x, y); +} +inline half copysign(half x, expr y) +{ + return functions::copysign(x, y); +} +inline half copysign(expr x, half y) +{ + return functions::copysign(x, y); +} +inline half copysign(expr x, expr y) +{ + return functions::copysign(x, y); +} + +/// \} +/// \name Floating point classification +/// \{ + +/// Classify floating point value. +/// \param arg number to classify +/// \retval FP_ZERO for positive and negative zero +/// \retval FP_SUBNORMAL for subnormal numbers +/// \retval FP_INFINITY for positive and negative infinity +/// \retval FP_NAN for NaNs +/// \retval FP_NORMAL for all other (normal) values +// template typename enable::type fpclassify(T arg) { return functions::fpclassify(arg); } +inline int fpclassify(half arg) +{ + return functions::fpclassify(arg); +} +inline int fpclassify(expr arg) +{ + return functions::fpclassify(arg); +} + +/// Check if finite number. +/// \param arg number to check +/// \retval true if neither infinity nor NaN +/// \retval false else +// template typename enable::type isfinite(T arg) { return functions::isfinite(arg); } +inline bool isfinite(half arg) +{ + return functions::isfinite(arg); +} +inline bool isfinite(expr arg) +{ + return functions::isfinite(arg); +} + +/// Check for infinity. +/// \param arg number to check +/// \retval true for positive or negative infinity +/// \retval false else +// template typename enable::type isinf(T arg) { return functions::isinf(arg); } +inline bool isinf(half arg) +{ + return functions::isinf(arg); +} +inline bool isinf(expr arg) +{ + return functions::isinf(arg); +} + +/// Check for NaN. +/// \param arg number to check +/// \retval true for NaNs +/// \retval false else +// template typename enable::type isnan(T arg) { return functions::isnan(arg); } +inline bool isnan(half arg) +{ + return functions::isnan(arg); +} +inline bool isnan(expr arg) +{ + return functions::isnan(arg); +} + +/// Check if normal number. +/// \param arg number to check +/// \retval true if normal number +/// \retval false if either subnormal, zero, infinity or NaN +// template typename enable::type isnormal(T arg) { return functions::isnormal(arg); } +inline bool isnormal(half arg) +{ + return functions::isnormal(arg); +} +inline bool isnormal(expr arg) +{ + return functions::isnormal(arg); +} + +/// Check sign. +/// \param arg number to check +/// \retval true for negative number +/// \retval false for positive number +// template typename enable::type signbit(T arg) { return functions::signbit(arg); } +inline bool signbit(half arg) +{ + return functions::signbit(arg); +} +inline bool signbit(expr arg) +{ + return functions::signbit(arg); +} + +/// \} +/// \name Comparison +/// \{ + +/// Comparison for greater than. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x greater than \a y +/// \retval false else +// template typename enable::type isgreater(T x, U y) { return +// functions::isgreater(x, y); } +inline bool isgreater(half x, half y) +{ + return functions::isgreater(x, y); +} +inline bool isgreater(half x, expr y) +{ + return functions::isgreater(x, y); +} +inline bool isgreater(expr x, half y) +{ + return functions::isgreater(x, y); +} +inline bool isgreater(expr x, expr y) +{ + return functions::isgreater(x, y); +} + +/// Comparison for greater equal. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x greater equal \a y +/// \retval false else +// template typename enable::type isgreaterequal(T x, U y) { return +// functions::isgreaterequal(x, y); } +inline bool isgreaterequal(half x, half y) +{ + return functions::isgreaterequal(x, y); +} +inline bool isgreaterequal(half x, expr y) +{ + return functions::isgreaterequal(x, y); +} +inline bool isgreaterequal(expr x, half y) +{ + return functions::isgreaterequal(x, y); +} +inline bool isgreaterequal(expr x, expr y) +{ + return functions::isgreaterequal(x, y); +} + +/// Comparison for less than. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x less than \a y +/// \retval false else +// template typename enable::type isless(T x, U y) { return functions::isless(x, +// y); +//} +inline bool isless(half x, half y) +{ + return functions::isless(x, y); +} +inline bool isless(half x, expr y) +{ + return functions::isless(x, y); +} +inline bool isless(expr x, half y) +{ + return functions::isless(x, y); +} +inline bool isless(expr x, expr y) +{ + return functions::isless(x, y); +} + +/// Comparison for less equal. +/// \param x first operand +/// \param y second operand +/// \retval true if \a x less equal \a y +/// \retval false else +// template typename enable::type islessequal(T x, U y) { return +// functions::islessequal(x, y); } +inline bool islessequal(half x, half y) +{ + return functions::islessequal(x, y); +} +inline bool islessequal(half x, expr y) +{ + return functions::islessequal(x, y); +} +inline bool islessequal(expr x, half y) +{ + return functions::islessequal(x, y); +} +inline bool islessequal(expr x, expr y) +{ + return functions::islessequal(x, y); +} + +/// Comarison for less or greater. +/// \param x first operand +/// \param y second operand +/// \retval true if either less or greater +/// \retval false else +// template typename enable::type islessgreater(T x, U y) { return +// functions::islessgreater(x, y); } +inline bool islessgreater(half x, half y) +{ + return functions::islessgreater(x, y); +} +inline bool islessgreater(half x, expr y) +{ + return functions::islessgreater(x, y); +} +inline bool islessgreater(expr x, half y) +{ + return functions::islessgreater(x, y); +} +inline bool islessgreater(expr x, expr y) +{ + return functions::islessgreater(x, y); +} + +/// Check if unordered. +/// \param x first operand +/// \param y second operand +/// \retval true if unordered (one or two NaN operands) +/// \retval false else +// template typename enable::type isunordered(T x, U y) { return +// functions::isunordered(x, y); } +inline bool isunordered(half x, half y) +{ + return functions::isunordered(x, y); +} +inline bool isunordered(half x, expr y) +{ + return functions::isunordered(x, y); +} +inline bool isunordered(expr x, half y) +{ + return functions::isunordered(x, y); +} +inline bool isunordered(expr x, expr y) +{ + return functions::isunordered(x, y); +} + +/// \name Casting +/// \{ + +/// Cast to or from half-precision floating point number. +/// This casts between [half](\ref half_float::half) and any built-in arithmetic type. The values are converted +/// directly using the given rounding mode, without any roundtrip over `float` that a `static_cast` would otherwise do. +/// It uses the default rounding mode. +/// +/// Using this cast with neither of the two types being a [half](\ref half_float::half) or with any of the two types +/// not being a built-in arithmetic type (apart from [half](\ref half_float::half), of course) results in a compiler +/// error and casting between [half](\ref half_float::half)s is just a no-op. +/// \tparam T destination type (half or built-in arithmetic type) +/// \tparam U source type (half or built-in arithmetic type) +/// \param arg value to cast +/// \return \a arg converted to destination type +template +T half_cast(U arg) +{ + return half_caster::cast(arg); +} + +/// Cast to or from half-precision floating point number. +/// This casts between [half](\ref half_float::half) and any built-in arithmetic type. The values are converted +/// directly using the given rounding mode, without any roundtrip over `float` that a `static_cast` would otherwise do. +/// +/// Using this cast with neither of the two types being a [half](\ref half_float::half) or with any of the two types +/// not being a built-in arithmetic type (apart from [half](\ref half_float::half), of course) results in a compiler +/// error and casting between [half](\ref half_float::half)s is just a no-op. +/// \tparam T destination type (half or built-in arithmetic type) +/// \tparam R rounding mode to use. +/// \tparam U source type (half or built-in arithmetic type) +/// \param arg value to cast +/// \return \a arg converted to destination type +template +T half_cast(U arg) +{ + return half_caster::cast(arg); +} +/// \} +} // namespace detail + +using detail::operator==; +using detail::operator!=; +using detail::operator<; +using detail::operator>; +using detail::operator<=; +using detail::operator>=; +using detail::operator+; +using detail::operator-; +using detail::operator*; +using detail::operator/; +using detail::operator<<; +using detail::operator>>; + +using detail::abs; +using detail::acos; +using detail::acosh; +using detail::asin; +using detail::asinh; +using detail::atan; +using detail::atan2; +using detail::atanh; +using detail::cbrt; +using detail::ceil; +using detail::cos; +using detail::cosh; +using detail::erf; +using detail::erfc; +using detail::exp; +using detail::exp2; +using detail::expm1; +using detail::fabs; +using detail::fdim; +using detail::floor; +using detail::fma; +using detail::fmax; +using detail::fmin; +using detail::fmod; +using detail::hypot; +using detail::lgamma; +using detail::log; +using detail::log10; +using detail::log1p; +using detail::log2; +using detail::lrint; +using detail::lround; +using detail::nanh; +using detail::nearbyint; +using detail::pow; +using detail::remainder; +using detail::remquo; +using detail::rint; +using detail::round; +using detail::sin; +using detail::sinh; +using detail::sqrt; +using detail::tan; +using detail::tanh; +using detail::tgamma; +using detail::trunc; +#if HALF_ENABLE_CPP11_LONG_LONG +using detail::llrint; +using detail::llround; +#endif +using detail::copysign; +using detail::fpclassify; +using detail::frexp; +using detail::ilogb; +using detail::isfinite; +using detail::isgreater; +using detail::isgreaterequal; +using detail::isinf; +using detail::isless; +using detail::islessequal; +using detail::islessgreater; +using detail::isnan; +using detail::isnormal; +using detail::isunordered; +using detail::ldexp; +using detail::logb; +using detail::modf; +using detail::nextafter; +using detail::nexttoward; +using detail::scalbln; +using detail::scalbn; +using detail::signbit; + +using detail::half_cast; +} // namespace half_float + +/// Extensions to the C++ standard library. +namespace std +{ +/// Numeric limits for half-precision floats. +/// Because of the underlying single-precision implementation of many operations, it inherits some properties from +/// `std::numeric_limits`. +template <> +class numeric_limits : public numeric_limits +{ +public: + /// Supports signed values. + static HALF_CONSTEXPR_CONST bool is_signed = true; + + /// Is not exact. + static HALF_CONSTEXPR_CONST bool is_exact = false; + + /// Doesn't provide modulo arithmetic. + static HALF_CONSTEXPR_CONST bool is_modulo = false; + + /// IEEE conformant. + static HALF_CONSTEXPR_CONST bool is_iec559 = true; + + /// Supports infinity. + static HALF_CONSTEXPR_CONST bool has_infinity = true; + + /// Supports quiet NaNs. + static HALF_CONSTEXPR_CONST bool has_quiet_NaN = true; + + /// Supports subnormal values. + static HALF_CONSTEXPR_CONST float_denorm_style has_denorm = denorm_present; + + /// Rounding mode. + /// Due to the mix of internal single-precision computations (using the rounding mode of the underlying + /// single-precision implementation) with the rounding mode of the single-to-half conversions, the actual rounding + /// mode might be `std::round_indeterminate` if the default half-precision rounding mode doesn't match the + /// single-precision rounding mode. + static HALF_CONSTEXPR_CONST float_round_style round_style + = (std::numeric_limits::round_style == half_float::half::round_style) ? half_float::half::round_style + : round_indeterminate; + + /// Significant digits. + static HALF_CONSTEXPR_CONST int digits = 11; + + /// Significant decimal digits. + static HALF_CONSTEXPR_CONST int digits10 = 3; + + /// Required decimal digits to represent all possible values. + static HALF_CONSTEXPR_CONST int max_digits10 = 5; + + /// Number base. + static HALF_CONSTEXPR_CONST int radix = 2; + + /// One more than smallest exponent. + static HALF_CONSTEXPR_CONST int min_exponent = -13; + + /// Smallest normalized representable power of 10. + static HALF_CONSTEXPR_CONST int min_exponent10 = -4; + + /// One more than largest exponent + static HALF_CONSTEXPR_CONST int max_exponent = 16; + + /// Largest finitely representable power of 10. + static HALF_CONSTEXPR_CONST int max_exponent10 = 4; + + /// Smallest positive normal value. + static HALF_CONSTEXPR half_float::half min() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x0400); + } + + /// Smallest finite value. + static HALF_CONSTEXPR half_float::half lowest() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0xFBFF); + } + + /// Largest finite value. + static HALF_CONSTEXPR half_float::half max() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x7BFF); + } + + /// Difference between one and next representable value. + static HALF_CONSTEXPR half_float::half epsilon() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x1400); + } + + /// Maximum rounding error. + static HALF_CONSTEXPR half_float::half round_error() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, (round_style == std::round_to_nearest) ? 0x3800 : 0x3C00); + } + + /// Positive infinity. + static HALF_CONSTEXPR half_float::half infinity() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x7C00); + } + + /// Quiet NaN. + static HALF_CONSTEXPR half_float::half quiet_NaN() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x7FFF); + } + + /// Signalling NaN. + static HALF_CONSTEXPR half_float::half signaling_NaN() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x7DFF); + } + + /// Smallest positive subnormal value. + static HALF_CONSTEXPR half_float::half denorm_min() HALF_NOTHROW + { + return half_float::half(half_float::detail::binary, 0x0001); + } +}; + +#if HALF_ENABLE_CPP11_HASH +/// Hash function for half-precision floats. +/// This is only defined if C++11 `std::hash` is supported and enabled. +template <> +struct hash //: unary_function +{ + /// Type of function argument. + typedef half_float::half argument_type; + + /// Function return type. + typedef size_t result_type; + + /// Compute hash function. + /// \param arg half to hash + /// \return hash value + result_type operator()(argument_type arg) const + { + return hash()(static_cast(arg.data_) & -(arg.data_ != 0x8000)); + } +}; +#endif +} // namespace std + +#undef HALF_CONSTEXPR +#undef HALF_CONSTEXPR_CONST +#undef HALF_NOEXCEPT +#undef HALF_NOTHROW +#ifdef HALF_POP_WARNINGS +#pragma warning(pop) +#undef HALF_POP_WARNINGS +#endif + +#endif diff --git a/samples/trtexecCommon/half.test.cpp b/samples/trtexecCommon/half.test.cpp new file mode 100644 index 0000000000..d06140f8b7 --- /dev/null +++ b/samples/trtexecCommon/half.test.cpp @@ -0,0 +1,162 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "half.h" + +#include + +#include +#include +#include +#include +#include + +using half_float::half; +using NLF32 = std::numeric_limits; +using NLHalf = std::numeric_limits; + +TEST(Half, Type) +{ + static_assert(sizeof(half) == sizeof(uint16_t), "half should be 16 bits!"); + static_assert(alignof(half) == alignof(uint16_t), "half should be 16 bit aligned!"); + EXPECT_EQ(half{}.operator float(), 0.0F); +} + +TEST(Half, Constructors) +{ + EXPECT_EQ(half{}, 0.0F); + EXPECT_EQ(half{1.0F}, 1.0F); + EXPECT_EQ(half{-1.0F}, -1.0F); + EXPECT_EQ(half{0.0F}, 0.0F); + EXPECT_EQ(half{0.5F}, 0.5F); + // Preserve sign bit, even for zero. + EXPECT_EQ(std::signbit(static_cast(half{-0.0F})), std::signbit(-1.0F)); + EXPECT_EQ(std::signbit(static_cast(-half{0.0F})), std::signbit(-1.0F)); + EXPECT_EQ(std::signbit(static_cast(half{0.0F})), std::signbit(1.0F)); +} + +TEST(Half, UnaryMinus) +{ + half const pos(2.5F); + half const neg = -pos; + EXPECT_EQ(neg, -2.5F); + + half const negInput(-3.0F); + half const posResult = -negInput; + EXPECT_EQ(posResult, 3.0F); + + half const zero(0.0F); + half const negZero = -zero; + EXPECT_EQ(negZero, 0.0F); +} + +TEST(Half, Addition) +{ + half const a(1.5F); + half const b(2.5F); + EXPECT_EQ(a + b, 4.0F); + + half const c(-1.0F); + half const d(3.0F); + EXPECT_EQ(c + d, 2.0F); + + half const e(0.0F); + half const f(5.0F); + EXPECT_EQ(e + f, 5.0F); +} + +TEST(Half, FloatConversion) +{ + EXPECT_EQ(static_cast(half{3.14159F}), 3.140625F); + EXPECT_EQ(static_cast(half{1000.0F}), 1000.0F); + EXPECT_EQ(static_cast(half{0.001F}), 0.0010004043F); + // Out-of-bounds conversion rounds to infinity. + EXPECT_TRUE(std::isinf(static_cast(half{NLF32::max()}))); +} + +TEST(Half, StreamOutput) +{ + auto toStr = [](auto const& value) { + std::stringstream ss; + ss << value; + return ss.str(); + }; + using namespace std::string_view_literals; + EXPECT_EQ(toStr(half(2.718F)), "2.71875"sv); + EXPECT_EQ(toStr(half(0.0F)), "0"sv); + // half should match float stringification for special values. + EXPECT_EQ(toStr(half{NLF32::infinity()}), std::to_string(NLF32::infinity())); + EXPECT_EQ(toStr(-half{NLF32::infinity()}), std::to_string(-NLF32::infinity())); + EXPECT_EQ(toStr(half{NLF32::quiet_NaN()}), std::to_string(NLF32::quiet_NaN())); +} + +TEST(Half, NumericLimits) +{ + static_assert(std::numeric_limits::is_specialized); + EXPECT_FALSE(std::numeric_limits::is_integer); + EXPECT_TRUE(std::numeric_limits::has_infinity); + + constexpr auto kINF = NLHalf::infinity(); + EXPECT_TRUE(isinf(kINF)); + EXPECT_LT(0.0F, kINF); + EXPECT_TRUE(isinf(-kINF)); + EXPECT_LT(-kINF, 0.0F); + + constexpr auto kMAX_VAL = NLHalf::max(); + EXPECT_FALSE(isinf(kMAX_VAL)); + EXPECT_EQ(kMAX_VAL, 65504.0F); + EXPECT_EQ(half(2.0F * 65505.0F), NLHalf::infinity()); + + constexpr auto kNAN = NLHalf::quiet_NaN(); + EXPECT_TRUE(isnan(kNAN)); + EXPECT_NE(kNAN, kNAN); // NaN is not equal to itself. +} + +TEST(Half, TypeTraits) +{ + static_assert(!std::is_arithmetic_v); + static_assert(!std::is_scalar_v); +} + +TEST(Half, NaN) +{ + EXPECT_TRUE(isnan(half{NLF32::quiet_NaN()})); + EXPECT_TRUE(isnan(half{} / half{})); + EXPECT_TRUE(isnan(half{half{} / half{}})); + EXPECT_TRUE(isnan(half{NLF32::quiet_NaN()} / half{})); + EXPECT_TRUE(isnan(half{NLF32::infinity()} - half{NLF32::infinity()})); +} + +TEST(Half, EdgeCases) +{ + auto const f16Large = half(1e38F); + EXPECT_TRUE(std::isinf(static_cast(f16Large))); + EXPECT_LT(NLHalf::max(), f16Large); + + auto const f16Small = half(0x1p-24F); // smallest subnormal + EXPECT_LT(0.0F, f16Small); + // Note: this library uses HALF_ROUND_TIES_TO_EVEN=0 (ties away from zero) by default. + // 0x1p-24 / 2 = 0x1p-25, which is equidistant between 0 and 0x1p-24. Ties-away-from-zero + // rounds up to 0x1p-24 rather than down to 0 as IEEE round-to-nearest-even would. + EXPECT_EQ(f16Small, half(f16Small / 2.0F)); +} + +TEST(Half, PrecisionAndRounding) +{ + constexpr float kPRECISE_VALUE = 1.0F + 1e-7F; + EXPECT_NEAR(static_cast(half{kPRECISE_VALUE}), kPRECISE_VALUE, 1e-3F); +} diff --git a/samples/trtexecCommon/logger.cpp b/samples/trtexecCommon/logger.cpp new file mode 100644 index 0000000000..7254a3d0f3 --- /dev/null +++ b/samples/trtexecCommon/logger.cpp @@ -0,0 +1,41 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "logger.h" +#include "ErrorRecorder.h" +#include "logging.h" +using namespace nvinfer1; +SampleErrorRecorder gRecorder; +namespace sample +{ +Logger gLogger{Logger::Severity::kINFO}; +LogStreamConsumer gLogVerbose{LOG_VERBOSE(gLogger)}; +LogStreamConsumer gLogInfo{LOG_INFO(gLogger)}; +LogStreamConsumer gLogWarning{LOG_WARN(gLogger)}; +LogStreamConsumer gLogError{LOG_ERROR(gLogger)}; +LogStreamConsumer gLogFatal{LOG_FATAL(gLogger)}; + +void setReportableSeverity(Logger::Severity severity) +{ + gLogger.setReportableSeverity(severity); + gLogVerbose.setReportableSeverity(severity); + gLogInfo.setReportableSeverity(severity); + gLogWarning.setReportableSeverity(severity); + gLogError.setReportableSeverity(severity); + gLogFatal.setReportableSeverity(severity); +} +} // namespace sample diff --git a/samples/trtexecCommon/logger.h b/samples/trtexecCommon/logger.h new file mode 100644 index 0000000000..396d1359ea --- /dev/null +++ b/samples/trtexecCommon/logger.h @@ -0,0 +1,37 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef LOGGER_H +#define LOGGER_H + +#include "logging.h" + +class SampleErrorRecorder; +extern SampleErrorRecorder gRecorder; +namespace sample +{ +extern Logger gLogger; +extern LogStreamConsumer gLogVerbose; +extern LogStreamConsumer gLogInfo; +extern LogStreamConsumer gLogWarning; +extern LogStreamConsumer gLogError; +extern LogStreamConsumer gLogFatal; + +void setReportableSeverity(Logger::Severity severity); +} // namespace sample + +#endif // LOGGER_H diff --git a/samples/trtexecCommon/logging.h b/samples/trtexecCommon/logging.h new file mode 100644 index 0000000000..2ffe55a37d --- /dev/null +++ b/samples/trtexecCommon/logging.h @@ -0,0 +1,687 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef TENSORRT_LOGGING_H +#define TENSORRT_LOGGING_H + +#include "NvInferRuntime.h" +#include "sampleOptions.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace sample +{ + +using Severity = nvinfer1::ILogger::Severity; + +class LogStreamConsumerBuffer : public std::stringbuf +{ +public: + LogStreamConsumerBuffer(std::ostream& stream, std::string const& prefix, bool shouldLog) + : mOutput(stream) + , mPrefix(prefix) + , mShouldLog(shouldLog) + { + } + + LogStreamConsumerBuffer(LogStreamConsumerBuffer&& other) noexcept + : mOutput(other.mOutput) + , mPrefix(other.mPrefix) + , mShouldLog(other.mShouldLog) + { + } + LogStreamConsumerBuffer(LogStreamConsumerBuffer const& other) = delete; + LogStreamConsumerBuffer() = delete; + LogStreamConsumerBuffer& operator=(LogStreamConsumerBuffer const&) = delete; + LogStreamConsumerBuffer& operator=(LogStreamConsumerBuffer&&) = delete; + + ~LogStreamConsumerBuffer() override + { + // std::streambuf::pbase() gives a pointer to the beginning of the buffered part of the output sequence + // std::streambuf::pptr() gives a pointer to the current position of the output sequence + // if the pointer to the beginning is not equal to the pointer to the current position, + // call putOutput() to log the output to the stream + if (pbase() != pptr()) + { + putOutput(); + } + } + + //! + //! synchronizes the stream buffer and returns 0 on success + //! synchronizing the stream buffer consists of inserting the buffer contents into the stream, + //! resetting the buffer and flushing the stream + //! + int32_t sync() override + { + putOutput(); + return 0; + } + + void putOutput() + { + if (mShouldLog) + { + // prepend timestamp + std::time_t timestamp = std::time(nullptr); + tm* tm_local = std::localtime(×tamp); + mOutput << "["; + mOutput << std::setw(2) << std::setfill('0') << 1 + tm_local->tm_mon << "/"; + mOutput << std::setw(2) << std::setfill('0') << tm_local->tm_mday << "/"; + mOutput << std::setw(4) << std::setfill('0') << 1900 + tm_local->tm_year << "-"; + mOutput << std::setw(2) << std::setfill('0') << tm_local->tm_hour << ":"; + mOutput << std::setw(2) << std::setfill('0') << tm_local->tm_min << ":"; + mOutput << std::setw(2) << std::setfill('0') << tm_local->tm_sec << "] "; + // std::stringbuf::str() gets the string contents of the buffer + // insert the buffer contents pre-appended by the appropriate prefix into the stream + mOutput << mPrefix << str(); + } + // set the buffer to empty + str(""); + // flush the stream + mOutput.flush(); + } + + void setShouldLog(bool shouldLog) + { + mShouldLog = shouldLog; + } + +private: + std::ostream& mOutput; + std::string mPrefix; + bool mShouldLog{}; +}; // class LogStreamConsumerBuffer + +//! +//! \class LogStreamConsumerBase +//! \brief Convenience object used to initialize LogStreamConsumerBuffer before std::ostream in LogStreamConsumer +//! +class LogStreamConsumerBase +{ +public: + LogStreamConsumerBase(std::ostream& stream, std::string const& prefix, bool shouldLog) + : mBuffer(stream, prefix, shouldLog) + { + } + +protected: + std::mutex mLogMutex; + LogStreamConsumerBuffer mBuffer; +}; // class LogStreamConsumerBase + +//! +//! \class LogStreamConsumer +//! \brief Convenience object used to facilitate use of C++ stream syntax when logging messages. +//! Order of base classes is LogStreamConsumerBase and then std::ostream. +//! This is because the LogStreamConsumerBase class is used to initialize the LogStreamConsumerBuffer member field +//! in LogStreamConsumer and then the address of the buffer is passed to std::ostream. +//! This is necessary to prevent the address of an uninitialized buffer from being passed to std::ostream. +//! Please do not change the order of the parent classes. +//! +class LogStreamConsumer : protected LogStreamConsumerBase, public std::ostream +{ +public: + //! + //! \brief Creates a LogStreamConsumer which logs messages with level severity. + //! Reportable severity determines if the messages are severe enough to be logged. + //! + LogStreamConsumer(nvinfer1::ILogger::Severity reportableSeverity, nvinfer1::ILogger::Severity severity) + : LogStreamConsumerBase(severityOstream(severity), severityPrefix(severity), severity <= reportableSeverity) + , std::ostream(&mBuffer) // links the stream buffer with the stream + , mShouldLog(severity <= reportableSeverity) + , mSeverity(severity) + { + } + + LogStreamConsumer(LogStreamConsumer&& other) noexcept + : LogStreamConsumerBase(severityOstream(other.mSeverity), severityPrefix(other.mSeverity), other.mShouldLog) + , std::ostream(&mBuffer) // links the stream buffer with the stream + , mShouldLog(other.mShouldLog) + , mSeverity(other.mSeverity) + { + } + LogStreamConsumer(LogStreamConsumer const& other) = delete; + LogStreamConsumer() = delete; + ~LogStreamConsumer() override = default; + LogStreamConsumer& operator=(LogStreamConsumer const&) = delete; + LogStreamConsumer& operator=(LogStreamConsumer&&) = delete; + + void setReportableSeverity(Severity reportableSeverity) + { + mShouldLog = mSeverity <= reportableSeverity; + mBuffer.setShouldLog(mShouldLog); + } + + std::mutex& getMutex() + { + return mLogMutex; + } + + bool getShouldLog() const + { + return mShouldLog; + } + +private: + static std::ostream& severityOstream(Severity severity) + { + return severity >= Severity::kINFO ? std::cout : std::cerr; + } + + static std::string severityPrefix(Severity severity) + { + switch (severity) + { + case Severity::kINTERNAL_ERROR: return "[F] "; + case Severity::kERROR: return "[E] "; + case Severity::kWARNING: return "[W] "; + case Severity::kINFO: return "[I] "; + case Severity::kVERBOSE: return "[V] "; + default: assert(0); return ""; + } + } + + bool mShouldLog; + Severity mSeverity; +}; // class LogStreamConsumer + +template +LogStreamConsumer& operator<<(LogStreamConsumer& logger, const T& obj) +{ + if (logger.getShouldLog()) + { + std::lock_guard guard(logger.getMutex()); + auto& os = static_cast(logger); + os << obj; + } + return logger; +} + +//! +//! Special handling std::endl +//! +inline LogStreamConsumer& operator<<(LogStreamConsumer& logger, std::ostream& (*f)(std::ostream&) ) +{ + if (logger.getShouldLog()) + { + std::lock_guard guard(logger.getMutex()); + auto& os = static_cast(logger); + os << f; + } + return logger; +} + +inline LogStreamConsumer& operator<<(LogStreamConsumer& logger, nvinfer1::Dims const& dims) +{ + if (logger.getShouldLog()) + { + std::lock_guard guard(logger.getMutex()); + auto& os = static_cast(logger); + for (int32_t i = 0; i < dims.nbDims; ++i) + { + os << (i ? "x" : "") << dims.d[i]; + } + } + return logger; +} + +template +inline LogStreamConsumer& operator<<(LogStreamConsumer& logger, std::pair const& value) +{ + if (logger.getShouldLog()) + { + std::lock_guard guard(logger.getMutex()); + auto& os = static_cast(logger); + os << "(" << value.first << ", " << value.second << ")"; + } + return logger; +} + +//! +//! \class Logger +//! +//! \brief Class which manages logging of TensorRT tools and samples +//! +//! \details This class provides a common interface for TensorRT tools and samples to log information to the console, +//! and supports logging two types of messages: +//! +//! - Debugging messages with an associated severity (info, warning, error, or internal error/fatal) +//! - Test pass/fail messages +//! +//! The advantage of having all samples use this class for logging as opposed to emitting directly to stdout/stderr is +//! that the logic for controlling the verbosity and formatting of sample output is centralized in one location. +//! +//! In the future, this class could be extended to support dumping test results to a file in some standard format +//! (for example, JUnit XML), and providing additional metadata (e.g. timing the duration of a test run). +//! +//! TODO: For backwards compatibility with existing samples, this class inherits directly from the nvinfer1::ILogger +//! interface, which is problematic since there isn't a clean separation between messages coming from the TensorRT +//! library and messages coming from the sample. +//! +//! In the future (once all samples are updated to use Logger::getTRTLogger() to access the ILogger) we can refactor the +//! class to eliminate the inheritance and instead make the nvinfer1::ILogger implementation a member of the Logger +//! object. +//! +class Logger : public nvinfer1::ILogger +{ +public: + explicit Logger(Severity severity = Severity::kWARNING) + : mReportableSeverity(severity) + { + } + + //! + //! \enum TestResult + //! \brief Represents the state of a given test + //! + enum class TestResult + { + kRUNNING, //!< The test is running + kPASSED, //!< The test passed + kFAILED, //!< The test failed + kWAIVED, //!< The test was waived + kTASK_BEGIN, //!< A sub-routine task has begun + kTASK_END, //!< A sub-routine task completed successfully + kTASK_ABORT //!< A sub-routine task was aborted (exception or validation failure) + }; + + //! + //! \brief Forward-compatible method for retrieving the nvinfer1::ILogger associated with this Logger + //! \return The nvinfer1::ILogger associated with this Logger + //! + //! TODO Once all samples are updated to use this method to register the logger with TensorRT, + //! we can eliminate the inheritance of Logger from ILogger + //! + nvinfer1::ILogger& getTRTLogger() noexcept + { + return *this; + } + + //! + //! \brief Implementation of the nvinfer1::ILogger::log() virtual method + //! + //! Note samples should not be calling this function directly; it will eventually go away once we eliminate the + //! inheritance from nvinfer1::ILogger + //! + void log(Severity severity, char const* msg) noexcept override + { + LogStreamConsumer(mReportableSeverity, severity) << "[TRT] " << std::string(msg) << std::endl; + } + + //! + //! \brief Method for controlling the verbosity of logging output + //! + //! \param severity The logger will only emit messages that have severity of this level or higher. + //! + void setReportableSeverity(Severity severity) noexcept + { + mReportableSeverity = severity; + } + + //! + //! \brief Opaque handle that holds logging information for a particular test + //! + //! This object is an opaque handle to information used by the Logger to print test results. + //! The sample must call Logger::defineTest() in order to obtain a TestAtom that can be used + //! with Logger::reportTest{Start,End}(). + //! + class TestAtom + { + public: + TestAtom(TestAtom&&) = default; + + std::string getCmdline() const + { + return mCmdline; + } + + private: + friend class Logger; + + TestAtom(bool started, std::string const& name, std::string const& cmdline) + : mStarted(started) + , mName(name) + , mCmdline(cmdline) + { + } + + bool mStarted; + std::string mName; + std::string mCmdline; + }; + + //! + //! \brief Define a test for logging + //! + //! \param[in] name The name of the test. This should be a string starting with + //! "TensorRT" and containing dot-separated strings containing + //! the characters [A-Za-z0-9_]. + //! For example, "TensorRT.sample_googlenet" + //! \param[in] cmdline The command line used to reproduce the test + // + //! \return a TestAtom that can be used in Logger::reportTest{Start,End}(). + //! + static TestAtom defineTest(std::string const& name, std::string const& cmdline) + { + return TestAtom(false, name, cmdline); + } + + //! + //! \brief A convenience overloaded version of defineTest() that accepts an array of command-line arguments + //! as input + //! + //! \param[in] name The name of the test + //! \param[in] argc The number of command-line arguments + //! \param[in] argv The array of command-line arguments (given as C strings) + //! + //! \return a TestAtom that can be used in Logger::reportTest{Start,End}(). + //! + static TestAtom defineTest(std::string const& name, int32_t argc, char const* const* argv) + { + // Append TensorRT version as info + const std::string vname = name + " [TensorRT v" + std::to_string(NV_TENSORRT_VERSION) + "] [b" + + std::to_string(NV_TENSORRT_BUILD) + "]"; + auto cmdline = genCmdlineString(argc, argv); + return defineTest(vname, cmdline); + } + + //! + //! \brief Report that a test has started. + //! + //! \pre reportTestStart() has not been called yet for the given testAtom + //! + //! \param[in] testAtom The handle to the test that has started + //! + static void reportTestStart(TestAtom& testAtom) + { + reportTestResult(testAtom, TestResult::kRUNNING); + assert(!testAtom.mStarted); + testAtom.mStarted = true; + } + + //! + //! \brief Report that a test has ended. + //! + //! \pre reportTestStart() has been called for the given testAtom + //! + //! \param[in] testAtom The handle to the test that has ended + //! \param[in] result The result of the test. Should be one of TestResult::kPASSED, + //! TestResult::kFAILED, TestResult::kWAIVED + //! + static void reportTestEnd(TestAtom const& testAtom, TestResult result) + { + assert(result != TestResult::kRUNNING); + assert(testAtom.mStarted); + reportTestResult(testAtom, result); + } + + static int32_t reportPass(TestAtom const& testAtom) + { + reportTestEnd(testAtom, TestResult::kPASSED); + return EXIT_SUCCESS; + } + + static int32_t reportFail(TestAtom const& testAtom) + { + reportTestEnd(testAtom, TestResult::kFAILED); + return EXIT_FAILURE; + } + + static int32_t reportWaive(TestAtom const& testAtom) + { + reportTestEnd(testAtom, TestResult::kWAIVED); + return EXIT_SUCCESS; + } + + //! + //! \brief Report that a sub-routine task has begun. + //! + //! Used by the tuning loop to mark the start of each iteration so external + //! tooling can detect iteration boundaries in the trtexec log stream. + //! + static void reportTaskBegin(TestAtom const& testAtom) + { + reportTestResult(testAtom, TestResult::kTASK_BEGIN); + } + + //! + //! \brief Report that a sub-routine task has begun with iteration index and build route. + //! Prints a blank line before the banner for readability. + //! + //! Output example: + //! &&&& TASK_BEGIN [iter=0] BuildRoute = '-match_ragged_mha=on -copy_ppg=off' + //! + static void reportTaskBegin(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) + { + reportTaskWithBuildRoute( + TestResult::kTASK_BEGIN, index, buildRoute, /*blankBefore=*/true, /*blankAfter=*/false); + } + + //! + //! \brief Report that a sub-routine task completed successfully. + //! + static void reportTaskEnd(TestAtom const& testAtom) + { + reportTestResult(testAtom, TestResult::kTASK_END); + } + + //! + //! \brief Report that a sub-routine task completed successfully, with iteration info. + //! Prints a blank line after the banner for readability. + //! + static void reportTaskEnd(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) + { + reportTaskWithBuildRoute(TestResult::kTASK_END, index, buildRoute, /*blankBefore=*/false, /*blankAfter=*/true); + } + + //! + //! \brief Report that a sub-routine task was aborted (exception or validation failure). + //! + static void reportTaskAbort(TestAtom const& testAtom) + { + reportTestResult(testAtom, TestResult::kTASK_ABORT); + } + + //! + //! \brief Report that a sub-routine task was aborted, with iteration info. + //! Prints a blank line after the banner for readability. + //! + static void reportTaskAbort(TestAtom const& /*testAtom*/, std::string const& index, std::string const& buildRoute) + { + reportTaskWithBuildRoute( + TestResult::kTASK_ABORT, index, buildRoute, /*blankBefore=*/false, /*blankAfter=*/true); + } + + static int32_t reportTest(TestAtom const& testAtom, bool pass) + { + return pass ? reportPass(testAtom) : reportFail(testAtom); + } + + Severity getReportableSeverity() const + { + return mReportableSeverity; + } + +private: + //! + //! \brief returns an appropriate string for prefixing a log message with the given severity + //! + static char const* severityPrefix(Severity severity) + { + switch (severity) + { + case Severity::kINTERNAL_ERROR: return "[F] "; + case Severity::kERROR: return "[E] "; + case Severity::kWARNING: return "[W] "; + case Severity::kINFO: return "[I] "; + case Severity::kVERBOSE: return "[V] "; + default: assert(0); return ""; + } + } + + //! + //! \brief returns an appropriate string for prefixing a test result message with the given result + //! + static char const* testResultString(TestResult result) + { + switch (result) + { + case TestResult::kRUNNING: return "RUNNING"; + case TestResult::kPASSED: return "PASSED"; + case TestResult::kFAILED: return "FAILED"; + case TestResult::kWAIVED: return "WAIVED"; + case TestResult::kTASK_BEGIN: return "TASK_BEGIN"; + case TestResult::kTASK_END: return "TASK_END"; + case TestResult::kTASK_ABORT: return "TASK_ABORT"; + default: assert(0); return ""; + } + } + + //! + //! \brief Print a TASK_BEGIN/END/ABORT banner with iteration index and build route. + //! + //! Output format: + //! &&&& TASK_BEGIN [iter=0] BuildRoute = '-match_ragged_mha=on -copy_ppg=off' + //! + static void reportTaskWithBuildRoute( + TestResult result, std::string const& index, std::string const& buildRoute, bool blankBefore, bool blankAfter) + { + auto& os = severityOstream(Severity::kINFO); + if (blankBefore) + { + os << std::endl; + } + os << "&&&& " << testResultString(result) << " [iter=" << index << "] BuildRoute = '" << buildRoute << "'" + << std::endl; + if (blankAfter) + { + os << std::endl; + } + } + + //! + //! \brief returns an appropriate output stream (cout or cerr) to use with the given severity + //! + static std::ostream& severityOstream(Severity severity) + { + return severity >= Severity::kINFO ? std::cout : std::cerr; + } + + //! + //! \brief method that implements logging test results + //! + static void reportTestResult(TestAtom const& testAtom, TestResult result) + { + severityOstream(Severity::kINFO) << "&&&& " << testResultString(result) << " " << testAtom.mName << " # " + << testAtom.mCmdline << std::endl; + } + + //! + //! \brief generate a command line string from the given (argc, argv) values + //! Note: It simply joins the arguments without proper escaping. If spaces is part + //! of an argument, they will be joined with single space. + //! + static std::string genCmdlineString(int32_t argc, char const* const* argv) + { + std::stringstream ss; + for (int32_t i = 0; i < argc; i++) + { + if (i > 0) + { + ss << " "; + } + ss << argv[i]; + } + return ss.str(); + } + + Severity mReportableSeverity; +}; // class Logger + +namespace +{ +//! +//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kVERBOSE +//! +//! Example usage: +//! +//! LOG_VERBOSE(logger) << "hello world" << std::endl; +//! +inline LogStreamConsumer LOG_VERBOSE(Logger const& logger) +{ + return LogStreamConsumer(logger.getReportableSeverity(), Severity::kVERBOSE); +} + +//! +//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kINFO +//! +//! Example usage: +//! +//! LOG_INFO(logger) << "hello world" << std::endl; +//! +inline LogStreamConsumer LOG_INFO(Logger const& logger) +{ + return LogStreamConsumer(logger.getReportableSeverity(), Severity::kINFO); +} + +//! +//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kWARNING +//! +//! Example usage: +//! +//! LOG_WARN(logger) << "hello world" << std::endl; +//! +inline LogStreamConsumer LOG_WARN(Logger const& logger) +{ + return LogStreamConsumer(logger.getReportableSeverity(), Severity::kWARNING); +} + +//! +//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kERROR +//! +//! Example usage: +//! +//! LOG_ERROR(logger) << "hello world" << std::endl; +//! +inline LogStreamConsumer LOG_ERROR(Logger const& logger) +{ + return LogStreamConsumer(logger.getReportableSeverity(), Severity::kERROR); +} + +//! +//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kINTERNAL_ERROR +//! ("fatal" severity) +//! +//! Example usage: +//! +//! LOG_FATAL(logger) << "hello world" << std::endl; +//! +inline LogStreamConsumer LOG_FATAL(Logger const& logger) +{ + return LogStreamConsumer(logger.getReportableSeverity(), Severity::kINTERNAL_ERROR); +} +} // anonymous namespace +} // namespace sample +#endif // TENSORRT_LOGGING_H diff --git a/samples/trtexecCommon/safeCommon.h b/samples/trtexecCommon/safeCommon.h new file mode 100644 index 0000000000..7be45e13f4 --- /dev/null +++ b/samples/trtexecCommon/safeCommon.h @@ -0,0 +1,666 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef TENSORRT_SAFE_COMMON_H +#define TENSORRT_SAFE_COMMON_H + +#include "NvInferRuntimeBase.h" +#include "NvInferSafeRecorder.h" +#include "NvInferSafeRuntime.h" +#include "cuda_runtime.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// For safeLoadLibrary +#ifdef _MSC_VER +// Needed so that the max/min definitions in windows.h do not conflict with std::max/min. +#define NOMINMAX +#include +#undef NOMINMAX +#else +#include +#endif +#if IS_QNX_SAFE +#include +#include +#endif // IS_QNX_SAFE + +using namespace nvinfer1; + +#undef CHECK_WITH_STREAM +#define CHECK_WITH_STREAM(status, stream) \ + do \ + { \ + if ((status) != cudaSuccess) \ + { \ + stream << "Cuda failure at " << __FILE__ << ":" << __LINE__ << ": " << cudaGetErrorString(status) \ + << std::endl; \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +#undef CUDA_CHECK +#define CUDA_CHECK(status) CHECK_WITH_STREAM(status, std::cerr) + +#define SAFE_LOG std::cerr + +inline std::string getTimestampStr() +{ + std::time_t timestamp = std::time(nullptr); + tm* tm_local = std::localtime(×tamp); + std::stringstream ss; + ss << "["; + ss << std::setw(2) << std::setfill('0') << 1 + tm_local->tm_mon << "/"; + ss << std::setw(2) << std::setfill('0') << tm_local->tm_mday << "/"; + ss << std::setw(4) << std::setfill('0') << 1900 + tm_local->tm_year << "-"; + ss << std::setw(2) << std::setfill('0') << tm_local->tm_hour << ":"; + ss << std::setw(2) << std::setfill('0') << tm_local->tm_min << ":"; + ss << std::setw(2) << std::setfill('0') << tm_local->tm_sec << "] "; + return ss.str(); +} + +inline void safeLogDebug(nvinfer2::safe::ISafeRecorder& recorder, std::string desc) +{ + desc = getTimestampStr() + "[D] " + desc; + recorder.reportDebug(desc.c_str()); +} + +inline void safeLogVerbose(nvinfer2::safe::ISafeRecorder& recorder, std::string desc) +{ + desc = getTimestampStr() + "[V] " + desc; + recorder.reportVerbose(desc.c_str()); +} + +inline void safeLogInfo(nvinfer2::safe::ISafeRecorder& recorder, std::string desc) +{ + desc = getTimestampStr() + "[I] " + desc; + recorder.reportInfo(desc.c_str()); +} + +inline void safeLogWarning(nvinfer2::safe::ISafeRecorder& recorder, std::string desc) +{ + desc = getTimestampStr() + "[W] " + desc; + recorder.reportWarn(desc.c_str()); +} + +inline void safeLogError( + nvinfer2::safe::ISafeRecorder& recorder, std::string desc, ErrorCode val = ErrorCode::kFAILED_EXECUTION) +{ + desc = getTimestampStr() + "[E] " + desc; + recorder.reportError(val, desc.c_str()); +} + +#undef SAFE_ASSERT +#define SAFE_ASSERT(condition) \ + do \ + { \ + if (!(condition)) \ + { \ + std::cerr << "Assertion failure: " << #condition << std::endl; \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +#define SAFE_API_CALL(api_call, recorder) \ + do \ + { \ + std::stringstream ss; \ + const ErrorCode ret = (api_call); \ + if (ret != ErrorCode::kSUCCESS) \ + { \ + ss << "SAFE API Error: [" << #api_call << "]: " << toString(ret); \ + safeLogError(recorder, ss.str(), ret); \ + throw ErrorCode{ret}; \ + } \ + ss << "SAFE API:[" << #api_call << "]: PASSED"; \ + safeLogVerbose(recorder, ss.str()); \ + } while (0) + +#define CUDA_CALL(cuda_api_call, recorder) \ + do \ + { \ + std::stringstream ss; \ + cudaError_t error = (cuda_api_call); \ + if (error != cudaSuccess) \ + { \ + ss << "CUDA Error: [" << #cuda_api_call << "]: " << cudaGetErrorString(error); \ + auto ret = ErrorCode::kFAILED_EXECUTION; \ + safeLogError(recorder, ss.str(), ret); \ + throw ErrorCode{ret}; \ + } \ + ss << "CUDA:[" << #cuda_api_call << "]: PASSED"; \ + safeLogVerbose(recorder, ss.str()); \ + } while (0) + +inline std::string toString(ErrorCode ec) +{ + static auto const ecStrings = [] { + std::unordered_map result; +#define INSERT_ELEMENT(p, s) result.emplace(p, s); + INSERT_ELEMENT(ErrorCode::kSUCCESS, "SUCCESS") + INSERT_ELEMENT(ErrorCode::kUNSPECIFIED_ERROR, "UNSPECIFIED_ERROR") + INSERT_ELEMENT(ErrorCode::kINTERNAL_ERROR, "INTERNAL_ERROR") + INSERT_ELEMENT(ErrorCode::kINVALID_ARGUMENT, "INVALID_ARGUMENT") + INSERT_ELEMENT(ErrorCode::kINVALID_CONFIG, "INVALID_CONFIG") + INSERT_ELEMENT(ErrorCode::kFAILED_ALLOCATION, "FAILED_ALLOCATION") + INSERT_ELEMENT(ErrorCode::kFAILED_INITIALIZATION, "FAILED_INITIALIZATION") + INSERT_ELEMENT(ErrorCode::kFAILED_EXECUTION, "FAILED_EXECUTION") + INSERT_ELEMENT(ErrorCode::kFAILED_COMPUTATION, "FAILED_COMPUTATION") + INSERT_ELEMENT(ErrorCode::kINVALID_STATE, "INVALID_STATE") + INSERT_ELEMENT(ErrorCode::kUNSUPPORTED_STATE, "UNSUPPORTED_STATE") +#undef INSERT_ELEMENT + return result; + }(); + return ecStrings.at(ec); +} + +//! Locate path to file, given its filename or filepath suffix and possible dirs it might lie in. +//! Function will also walk back MAX_DEPTH dirs from CWD to check for such a file path. +inline std::string locateFile( + std::string const& filepathSuffix, std::vector const& directories, bool reportError = true) +{ + int const MAX_DEPTH{10}; + bool found{false}; + std::string filepath; + + for (auto& dir : directories) + { + if (!dir.empty() && dir.back() != '/') + { +#ifdef _MSC_VER + filepath = dir + "\\" + filepathSuffix; +#else + filepath = dir + "/" + filepathSuffix; +#endif + } + else + { + filepath = dir + filepathSuffix; + } + + for (int i = 0; i < MAX_DEPTH && !found; i++) + { + const std::ifstream checkFile(filepath); + found = checkFile.is_open(); + if (found) + { + break; + } + + filepath = "../" + filepath; // Try again in parent dir + } + + if (found) + { + break; + } + + filepath.clear(); + } + + // Could not find the file + if (filepath.empty()) + { + const std::string dirList = std::accumulate(directories.begin() + 1, directories.end(), directories.front(), + [](std::string const& a, std::string const& b) { return a + "\n\t" + b; }); + std::cout << "Could not find " << filepathSuffix << " in data directories:\n\t" << dirList << std::endl; + + if (reportError) + { + std::cout << "&&&& FAILED" << std::endl; + exit(EXIT_FAILURE); + } + } + + return filepath; +} + +inline void readPGMFile(std::string const& fileName, uint8_t* buffer, int32_t inH, int32_t inW) +{ + std::ifstream infile(fileName, std::ifstream::binary); + SAFE_ASSERT(infile.is_open() && "Attempting to read from a file that is not open."); + std::string magic, w, h, max; + infile >> magic >> w >> h >> max; + infile.seekg(1, infile.cur); + infile.read(reinterpret_cast(buffer), inH * inW); +} + +namespace samplesSafeCommon +{ +#if !TRT_WINML +//! Represents the compute capability of a device. +//! This pertains to virtual architectures represented by the intermediate PTX format. +//! This is distinct from the SM version. +//! See https://forums.developer.nvidia.com/t/how-should-i-use-correctly-the-sm-xx-and-compute-xx/219160 +struct ComputeCapability +{ + int32_t major{}; + int32_t minor{}; + + //! \return the compute capability of the CUDA device with the given \p deviceIndex. + [[nodiscard]] static ComputeCapability forDevice(int32_t deviceIndex) + { + int32_t major{0}; + int32_t minor{0}; + CUDA_CHECK(cudaDeviceGetAttribute(&major, cudaDevAttrComputeCapabilityMajor, deviceIndex)); + CUDA_CHECK(cudaDeviceGetAttribute(&minor, cudaDevAttrComputeCapabilityMinor, deviceIndex)); + return {major, minor}; + } +}; + +inline int32_t getSmVersion() +{ + int32_t deviceIndex{}; + CUDA_CHECK(cudaGetDevice(&deviceIndex)); + + auto const cc = ComputeCapability::forDevice(deviceIndex); + return ((cc.major << 8) | cc.minor); +} + +inline bool isSmSafe() +{ + int32_t const smVersion = getSmVersion(); + return smVersion == 0x0705 || smVersion == 0x0800 || smVersion == 0x0806 || smVersion == 0x0807 + || smVersion == 0x0A00 || smVersion == 0x0B00; +} +#endif + +inline int32_t calculateSoftmax(float* const prob, int32_t const numDigits) +{ + SAFE_ASSERT(prob != nullptr); + SAFE_ASSERT(numDigits == 10); + float sum{0.0F}; + std::transform(prob, prob + numDigits, prob, [&sum](float v) -> float { + sum += exp(v); + return exp(v); + }); + + SAFE_ASSERT(sum != 0.0F); + std::transform(prob, prob + numDigits, prob, [sum](float v) -> float { return v / sum; }); + int32_t idx = std::max_element(prob, prob + numDigits) - prob; + return idx; +} + +//! +//! \brief generate a command line string from the given (argc, argv) values +//! Note: It simply joins the arguments without proper escaping. If spaces is part +//! of an argument, they will be joined with single space. +//! +static std::string genCmdlineString(int32_t argc, char const* const* argv) +{ + std::stringstream ss; + for (int32_t i = 0; i < argc; i++) + { + if (i > 0) + { + ss << " "; + } + ss << argv[i]; + } + return ss.str(); +} + +//! +//! \enum TestResult +//! \brief Represents the state of a given test +//! +enum class TestResult +{ + kFAILED, //!< The test failed + kPASSED, //!< The test passed +}; + +//! +//! \brief method that implements logging test start +//! +inline void reportTestStart(std::string testName, int32_t argc, char const* const* argv) +{ + SAFE_LOG << "&&&& RUNNING " << testName << " [TensorRT v" << std::to_string(NV_TENSORRT_VERSION) << "] [b" + << std::to_string(NV_TENSORRT_BUILD) << "]" + << " # " << genCmdlineString(argc, argv) << std::endl; +} + +//! +//! \brief method that implements logging test results +//! +inline void reportTestResult(std::string testName, TestResult result, int32_t argc, char const* const* argv) +{ + SAFE_LOG << "&&&& " << (result == TestResult::kPASSED ? "PASSED" : "FAILED") << " " << testName << " [TensorRT v" + << std::to_string(NV_TENSORRT_VERSION) << "] [b" << std::to_string(NV_TENSORRT_BUILD) << "]" + << " # " << genCmdlineString(argc, argv) << std::endl; +} + +//! +//! \class TrtCudaGraphSafe +//! \brief Managed CUDA graph +//! +class TrtCudaGraphSafe +{ +public: + explicit TrtCudaGraphSafe() = default; + + TrtCudaGraphSafe(TrtCudaGraphSafe const&) = delete; + + TrtCudaGraphSafe& operator=(TrtCudaGraphSafe const&) = delete; + + TrtCudaGraphSafe(TrtCudaGraphSafe&&) = delete; + + TrtCudaGraphSafe& operator=(TrtCudaGraphSafe&&) = delete; + + ~TrtCudaGraphSafe() + { + if (mGraphExec) + { + cudaGraphExecDestroy(mGraphExec); + } + } + + void beginCapture(cudaStream_t& stream) + { + CUDA_CHECK(cudaStreamBeginCapture(stream, cudaStreamCaptureModeThreadLocal)); + } + + bool launch(cudaStream_t& stream) + { + return cudaGraphLaunch(mGraphExec, stream) == cudaSuccess; + } + + void endCapture(cudaStream_t& stream) + { + CUDA_CHECK(cudaStreamEndCapture(stream, &mGraph)); + CUDA_CHECK(cudaGraphInstantiate(&mGraphExec, mGraph, nullptr, nullptr, 0)); + CUDA_CHECK(cudaGraphDestroy(mGraph)); + } + + void endCaptureOnError(cudaStream_t& stream) + { + // There are two possibilities why stream capture would fail: + // (1) stream is in cudaErrorStreamCaptureInvalidated state. + // (2) TRT reports a failure. + // In case (1), the returning mGraph should be nullptr. + // In case (2), the returning mGraph is not nullptr, but it should not be used. + auto const ret = cudaStreamEndCapture(stream, &mGraph); + if (ret == cudaErrorStreamCaptureInvalidated) + { + SAFE_ASSERT(mGraph == nullptr); + } + else + { + SAFE_ASSERT(ret == cudaSuccess); + SAFE_ASSERT(mGraph != nullptr); + CUDA_CHECK(cudaGraphDestroy(mGraph)); + mGraph = nullptr; + } + // Clean up any CUDA error. + cudaGetLastError(); + SAFE_LOG << "The CUDA graph capture on the stream has failed." << std::endl; + } + +private: + cudaGraph_t mGraph{}; + cudaGraphExec_t mGraphExec{}; +}; + +inline void* safeLoadLibrary(std::string const& path) +{ +#ifdef _MSC_VER + void* handle = LoadLibraryA(path.c_str()); +#else + int32_t flags{RTLD_LAZY}; + void* handle = dlopen(path.c_str(), flags); +#endif + if (handle == nullptr) + { +#ifdef _MSC_VER + sample::gLogError << "Could not load plugin library: " << path << std::endl; +#else + SAFE_LOG << "Could not load plugin library: " << path << ", due to: " << dlerror() << std::endl; +#endif + } + return handle; +} + +//! +//! \class SafetyPluginAttribute +//! \brief Represents a safety plugin with its namespace and name +//! +class SafetyPluginAttribute +{ +public: + std::string pluginNamespace; //!< Plugin namespace (optional, can be empty) + std::string pluginName; //!< Plugin name +}; + +//! +//! \class SafetyPluginLibraryArgument +//! \brief Represents a safety plugin library with its name and associated plugin attributes +//! Used for parsing command line arguments in the format: libraryName[namespace::pluginName1,pluginName2] +//! +class SafetyPluginLibraryArgument +{ +public: + std::string libraryName; //!< Name of the plugin library + std::vector pluginAttrs; //!< Vector of plugin attributes contained in this library +}; + +inline std::vector safeSplitString(std::string str, char delimiter = ',') +{ + std::vector splitVect; + std::stringstream ss(str); + std::string substr; + + while (ss.good()) + { + getline(ss, substr, delimiter); + splitVect.emplace_back(std::move(substr)); + } + return splitVect; +} + +// Safety plugin cmd argument example: safetyPluginLibrary[namespace::pluginName1,pluginName2] +inline bool parseSafetyPluginArgument(std::string const& option, SafetyPluginLibraryArgument& args) +{ + auto const leftBracketIdx = option.find('['); + auto const rightBracketIdx = option.find(']'); + if (leftBracketIdx == std::string::npos || rightBracketIdx == std::string::npos || leftBracketIdx > rightBracketIdx) + { + SAFE_LOG << "Invalid safety plugin argument: " << option << std::endl; + return false; + } + args.libraryName = option.substr(0, leftBracketIdx); + auto const pluginOptionStr = option.substr(leftBracketIdx + 1, rightBracketIdx - leftBracketIdx - 1); + auto const pluginOptions = safeSplitString(pluginOptionStr, ','); + if (args.libraryName.empty() || pluginOptions.empty()) + { + SAFE_LOG << "Invalid safety plugin argument: " << option << std::endl; + return false; + } + + auto parsePluginOption = [](std::string const& pluginOption) { + SafetyPluginAttribute attr{}; + // Check if namespace is used, leave as empty if not exist + auto const sepratorIdx = pluginOption.find("::"); + if (sepratorIdx == std::string::npos) + { + attr.pluginName = pluginOption; + } + else + { + attr.pluginNamespace = pluginOption.substr(0, sepratorIdx); + attr.pluginName = pluginOption.substr(sepratorIdx + 2, pluginOption.length() - sepratorIdx - 2); + } + return attr; + }; + + for (auto const& pluginOption : pluginOptions) + { + auto attr = parsePluginOption(pluginOption); + if (!attr.pluginName.empty()) + { + args.pluginAttrs.push_back(attr); + } + } + + return true; +} + +//! \brief Check if arg is a command-line option with a value, and get the value. +//! If arg matches the form --OPTION=VALUE (or -C=VALUE if singleChar is provided) +//! and OPTION matches the provided name (or C matches the provided singleChar), +//! extract the option VALUE. +//! \param arg The argument to attempt to parse +//! \param name The option name to match +//! \param singleChar Single-character option, or std::nullopt if full name is required. +//! \return If name matched, parsed string VALUE; otherwise std::nullopt. +inline std::optional parseString( + std::string const& arg, std::string const& name, std::optional singleChar = std::nullopt) +{ + for (std::string const& prefix : { + "--" + name + "=", + singleChar ? std::string{'-', *singleChar, '='} : "", + }) + { + if (!prefix.empty() && prefix == arg.substr(0, prefix.size())) + { + return arg.substr(prefix.size()); + } + } + return std::nullopt; +} + +//! \brief Check if arg is command-line option without a value. +//! Check whether arg matches the form --OPT (or -C if singleChar is provided) +//! and OPT matches the provided name. +//! \param arg The argument to attempt to parse +//! \param name The option name to match +//! \param singleChar Single-character version of OPT, or std::nullopt if full name is required. +//! \return true on match, false otherwise. +inline bool parseBool(std::string const& arg, std::string const& name, std::optional singleChar = std::nullopt) +{ + return arg == "--" + name || (singleChar && arg == std::string{'-', *singleChar}); +} + +inline bool hasCpuOnlyInternalOption(std::string const& internalOptions) +{ + std::istringstream optionStream{internalOptions}; + for (std::string option; optionStream >> option;) + { + if (option == "--cpu_only" || option.rfind("--cpu_only=", 0) == 0) + { + return true; + } + } + return false; +} + +inline bool applyCpuOnlyMode() +{ +#if !defined(_WIN32) + // The use of TRT_INTERNAL_OPTIONS is special to TensorRT 11.0 and will disappear in later releases. + char const* internalOptions = std::getenv("TRT_INTERNAL_OPTIONS"); + std::string internalOptionsStr{internalOptions != nullptr ? internalOptions : ""}; + if (hasCpuOnlyInternalOption(internalOptionsStr)) + { + return true; + } + + internalOptionsStr += " --cpu_only=1"; + if (setenv("TRT_INTERNAL_OPTIONS", internalOptionsStr.c_str(), 1) != 0) + { + SAFE_LOG << "Failed to set TRT_INTERNAL_OPTIONS for CPU-only mode: " << std::strerror(errno) << std::endl; + return false; + } +#endif + return true; +} + +//! \brief Allocate \p graph's auxiliary CUDA streams and register them via setAuxStreams. +//! +//! The returned scope guard owns the streams; when the last reference is dropped, each +//! stream is destroyed via cudaStreamDestroy. The pointed-to value (nullptr) is purposely +//! opaque; only the deleter matters. The caller must keep the guard alive during the graph +//! inference runs. +[[nodiscard]] inline std::shared_ptr setUpAuxStreamsOn( + nvinfer2::safe::ITRTGraph& graph, nvinfer2::safe::ISafeRecorder& recorder) +{ + int32_t nbAuxStreams{}; + SAFE_API_CALL(graph.getNbAuxStreams(nbAuxStreams), recorder); + std::vector streams(nbAuxStreams); + for (auto& s : streams) + { + CUDA_CHECK(cudaStreamCreateWithFlags(&s, cudaStreamNonBlocking)); + } + SAFE_API_CALL(graph.setAuxStreams(streams.data(), nbAuxStreams), recorder); + return {nullptr, [streams = std::move(streams)](void*) { + for (cudaStream_t s : streams) + { + if (s) + { + (void) cudaStreamDestroy(s); + } + } + }}; +} + +} // namespace samplesSafeCommon + +namespace safetyCompliance +{ +inline void initSafeCuda() +{ + // According to CUDA initialization in NVIDIA CUDA SAFETY API REFERENCE FOR DRIVE OS + // We will need to do the following in order + // 1. Initialize the calling thread with CUDA specific information (Call any CUDA RT API identified as init) + // 2. Query/Configure and choose the desired CUDA device + // 3. CUDA context initialization. (Call cudaDeviceGetLimit or cuCtxCreate) + size_t stackSizeLimit = 0; + int32_t deviceIndex = 0; + CUDA_CHECK(cudaGetDevice(&deviceIndex)); + CUDA_CHECK(cudaDeviceGetLimit(&stackSizeLimit, cudaLimitStackSize)); +#if IS_QNX_SAFE + CUDA_CHECK(cudaSafeExSelectAPIMode(cudaSafeExAPIModeAsilB)); +#endif // IS_QNX_SAFE +} + +inline void setPromgrAbility() +{ +#if IS_QNX_SAFE + // Comply with DEEPLRN_RES_117 on QNX-safe by dropping PROCMGR_AID_MEM_PHYS ability and locking out any further + // changes + procmgr_ability( + 0, PROCMGR_ADN_NONROOT | PROCMGR_AOP_DENY | PROCMGR_AOP_LOCK | PROCMGR_AID_MEM_PHYS, PROCMGR_AID_EOL); +#endif // IS_QNX_SAFE +} + +} // namespace safetyCompliance + +#endif // TENSORRT_SAFE_COMMON_H diff --git a/samples/common/safeCudaAllocator.h b/samples/trtexecCommon/safeCudaAllocator.h similarity index 100% rename from samples/common/safeCudaAllocator.h rename to samples/trtexecCommon/safeCudaAllocator.h diff --git a/samples/trtexecCommon/safeErrorRecorder.h b/samples/trtexecCommon/safeErrorRecorder.h new file mode 100644 index 0000000000..90f0f8951c --- /dev/null +++ b/samples/trtexecCommon/safeErrorRecorder.h @@ -0,0 +1,294 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef SAFE_ERROR_RECORDER_H +#define SAFE_ERROR_RECORDER_H + +#include "NvInferSafeRecorder.h" +#include "safeCommon.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#if ENABLE_NVLOG +#include +#endif // ENABLE_NVLOG + +namespace sample +{ +using namespace nvinfer2::safe; + +namespace detail +{ +//! Copy the contents of a std::string_view into a buffer, truncating if it doesn't fit, and always null-terminating the +//! result (unless dst == nullptr or dstSize == 0). Returns the number of bytes written, including the null terminator. +//! Behavior is undefined if dstSize < 0. +//! Behavior is undefined if dst == nullpltr but dstSize > 0. +inline int64_t truncatingCopyAsCString(std::string_view const src, AsciiChar* const dst, int64_t const dstSize) +{ + SAFE_ASSERT(0 <= dstSize); + SAFE_ASSERT(dst != nullptr || dstSize == 0); + if (dstSize == 0) + { + return 0; + } + SAFE_ASSERT(0 < dstSize); //< We at least have room for a null terminator. + auto const toWrite = std::min(static_cast(src.size()), dstSize - 1); + std::copy_n(src.data(), toWrite, dst); + dst[toWrite] = '\0'; + return toWrite + 1; +} + +//! Copy the contents of a std::string_view into a fixed-size buffer, null terminated (unless the buffer has zero size) +//! and return it. +template +[[nodiscard]] constexpr TArray truncatedCopyAsCString(std::string_view const src) +{ + TArray result{}; + // Expecting TArray to be a std::array or similar, constructing to its size and having a static size: + SAFE_ASSERT(std::tuple_size::value == result.size()); + truncatingCopyAsCString(src, result.data(), result.size()); + return result; +} + +#if ENABLE_NVLOG +//! \return a severity number corresponding to an `nvinfer2::safe::Severity`. +//! Unlisted enumerators map to `NVOS_LOG_SEVERITY_INFO` +[[nodiscard]] constexpr uint8_t toNvOsLogSeverity(nvinfer2::safe::Severity sev) +{ + switch (sev) + { + case nvinfer2::safe::Severity::kINFO: return NVOS_LOG_SEVERITY_INFO; + case nvinfer2::safe::Severity::kWARNING: return NVOS_LOG_SEVERITY_WARNING; + case nvinfer2::safe::Severity::kVERBOSE: return NVOS_LOG_SEVERITY_DEBUG1; + default: return NVOS_LOG_SEVERITY_INFO; + } +} +#endif // ENABLE_NVLOG +} // namespace detail + +//! The SampleSafeRecorder implementation of the ISafeRecorder interface. +class SampleSafeRecorder : public ISafeRecorder +{ + using DescHolder = std::array; + using errorPair = std::pair; + +public: + SampleSafeRecorder(nvinfer2::safe::Severity severity = nvinfer2::safe::Severity::kINFO, int32_t index = -1, + char const* filename = "TRTErrors.log") + : ISafeRecorder(severity, index) + { + fatalErrorLogFile = fopen(filename, "w+"); + if (fatalErrorLogFile == nullptr) + { + std::cerr << "Failed to open error log file: " << filename << std::endl; + } + } + + virtual ~SampleSafeRecorder() noexcept + { + if (mRefCount != 0) + { + reportError(ErrorCode::kINTERNAL_ERROR, "Non-zero reference count for recorder upon deallocation."); + } + } + + int32_t getNbErrors() const noexcept final + { + return nbErrors; + } + ErrorCode getErrorCode(int32_t errorIdx) const noexcept final + { + return invalidIndexCheck(errorIdx) ? ErrorCode::kINVALID_ARGUMENT : (*this)[errorIdx].first; + }; + ErrorDesc getErrorDesc(int32_t errorIdx) const noexcept final + { + return invalidIndexCheck(errorIdx) ? "ErrorIdx is out of range." : (*this)[errorIdx].second.data(); + } + + bool hasOverflowed() const noexcept final + { + return (nbErrors >= kMAX_NB_ERRORS); + } + + int32_t getMaxNbErrors() const + { + return kMAX_NB_ERRORS; + } + + // Empty the errorStack. + void clear() noexcept final + { + try + { + // grab a lock so that there is no addition while clearing. + std::lock_guard guard(mStackLock); + nbErrors = 0; + } + catch (std::exception const& e) + { +#if ENABLE_NVLOG + NvOsDebugPrintStr(NVOS_LOG_CODE_START, NVOS_LOG_SEVERITY_ERROR, e.what()); +#else + std::cerr << "Internal Error: " << e.what() << std::endl; +#endif // ENABLE_NVLOG + } + }; + + //! Simple helper function that checks if the error stack is empty. + bool empty() const noexcept + { + return (nbErrors == 0); + } + + bool reportError(ErrorCode val, ErrorDesc desc) noexcept final + { + try + { + std::string_view const descView = desc; //< This implicitly calls strlen once. + std::cerr << descView << std::endl; + DescHolder descArr = detail::truncatedCopyAsCString(descView); + + std::lock_guard guard(mStackLock); +#if ENABLE_NVLOG + NvOsDebugPrintStrInt( + NVOS_LOG_CODE_START, NVOS_LOG_SEVERITY_ERROR, descArr.data(), static_cast(val)); +#else + // Only write to the array if there's space available + if (nbErrors < kMAX_NB_ERRORS) + { + mErrorStack.at(nbErrors) = errorPair(val, descArr); + nbErrors++; + } +#endif // ENABLE_NVLOG + } + catch (std::exception const& e) + { +#if ENABLE_NVLOG + NvOsDebugPrintStr(NVOS_LOG_CODE_START, NVOS_LOG_SEVERITY_ERROR, e.what()); +#else + // `std::ofstream` uses heap allocation which is not allowed for safe samples + // Hence, C functions are used here to write data to file. + if (fatalErrorLogFile != nullptr) + { + setbuf(fatalErrorLogFile, NULL); + fwrite(e.what(), strlen(e.what()), 1, fatalErrorLogFile); + fwrite("\n", 1, 1, fatalErrorLogFile); + fflush(fatalErrorLogFile); + } + std::cerr << e.what() << std::endl; +#endif // ENABLE_NVLOG + } + // All errors are considered fatal. + return true; + } + + bool reportInfo(ErrorDesc desc) noexcept final + { + return reportIfSevere(nvinfer2::safe::Severity::kINFO, desc); + } + + bool reportWarn(ErrorDesc desc) noexcept final + { + return reportIfSevere(nvinfer2::safe::Severity::kWARNING, desc); + } + + bool reportVerbose(ErrorDesc desc) noexcept final + { + return reportIfSevere(nvinfer2::safe::Severity::kVERBOSE, desc); + } + + bool reportDebug(ErrorDesc desc) noexcept final + { + return reportIfSevere(nvinfer2::safe::Severity::kDEBUG, desc); + } + + // Atomically increment or decrement the ref counter. + RefCount incRefCount() noexcept final + { + return ++mRefCount; + } + RefCount decRefCount() noexcept final + { + return --mRefCount; + } + +private: + // Simple helper functions. + errorPair const& operator[](int32_t index) const noexcept + { + return mErrorStack[index]; + } + + bool invalidIndexCheck(int32_t index) const noexcept + { + return index >= nbErrors; + } + + bool reportIfSevere(nvinfer2::safe::Severity msgSev, ErrorDesc desc) noexcept + { +#if ENABLE_NVLOG + if (mSeverity >= msgSev) + { + auto const severity = detail::toNvOsLogSeverity(msgSev); + std::lock_guard guard(mStackLock); + NvOsDebugPrintStr(NVOS_LOG_CODE_START, severity, desc); + return true; + } + return false; +#else + if (mSeverity >= msgSev) + { + std::lock_guard guard(mStackLock); + std::cout << desc << std::endl; + return true; + } + return false; +#endif // ENABLE_NVLOG + } + + // Used to store the logs that are Fatal + FILE* fatalErrorLogFile; + + // Mutex to hold when locking mErrorStack for thread safety. + std::mutex mStackLock; + + // Reference count of the class. Destruction of the class when mRefCount + // is not zero causes undefined behavior. + std::atomic mRefCount{0}; + + // Number of errors that occurred so far. + int32_t nbErrors{0}; + + // Maximum number of errors that can be stored. + static constexpr int32_t kMAX_NB_ERRORS = 10; + + // The error stack that holds the errors recorded by TensorRT. + std::array mErrorStack; +}; // class SampleRecorder + +} // namespace sample + +#endif // SAFE_ERROR_RECORDER_H diff --git a/samples/common/sampleDevice.cpp b/samples/trtexecCommon/sampleDevice.cpp similarity index 100% rename from samples/common/sampleDevice.cpp rename to samples/trtexecCommon/sampleDevice.cpp diff --git a/samples/common/sampleDevice.h b/samples/trtexecCommon/sampleDevice.h similarity index 93% rename from samples/common/sampleDevice.h rename to samples/trtexecCommon/sampleDevice.h index 4efb232f2e..9e3d2696f0 100644 --- a/samples/common/sampleDevice.h +++ b/samples/trtexecCommon/sampleDevice.h @@ -58,13 +58,24 @@ class TrtCudaStream { public: TrtCudaStream() + : TrtCudaStream(nullptr) { - CHECK(cudaStreamCreate(&mStream)); } - TrtCudaStream(const TrtCudaStream&) = delete; + //! \brief Uses \p stream without taking ownership, or creates an owned stream when it is nullptr. + explicit TrtCudaStream(cudaStream_t stream) + : mStream(stream) + , mOwnsStream(stream == nullptr) + { + if (mOwnsStream) + { + CHECK(cudaStreamCreate(&mStream)); + } + } - TrtCudaStream& operator=(const TrtCudaStream&) = delete; + TrtCudaStream(TrtCudaStream const&) = delete; + + TrtCudaStream& operator=(TrtCudaStream const&) = delete; TrtCudaStream(TrtCudaStream&&) = delete; @@ -72,7 +83,10 @@ class TrtCudaStream ~TrtCudaStream() { - CHECK(cudaStreamDestroy(mStream)); + if (mOwnsStream) + { + CHECK(cudaStreamDestroy(mStream)); + } } cudaStream_t get() const @@ -94,6 +108,7 @@ class TrtCudaStream private: cudaStream_t mStream{}; + bool mOwnsStream; }; //! @@ -150,7 +165,7 @@ class TrtCudaEvent } // Returns time elapsed time in milliseconds - float operator-(const TrtCudaEvent& e) const + float operator-(TrtCudaEvent const& e) const { // Synchronize both events to ensure they have completed before calculating elapsed time synchronize(); @@ -195,9 +210,9 @@ class TrtCudaGraph public: explicit TrtCudaGraph() = default; - TrtCudaGraph(const TrtCudaGraph&) = delete; + TrtCudaGraph(TrtCudaGraph const&) = delete; - TrtCudaGraph& operator=(const TrtCudaGraph&) = delete; + TrtCudaGraph& operator=(TrtCudaGraph const&) = delete; TrtCudaGraph(TrtCudaGraph&&) = delete; @@ -235,7 +250,7 @@ class TrtCudaGraph // (2) TRT reports a failure. // In case (1), the returning mGraph should be nullptr. // In case (2), the returning mGraph is not nullptr, but it should not be used. - const auto ret = cudaStreamEndCapture(stream.get(), &mGraph); + auto const ret = cudaStreamEndCapture(stream.get(), &mGraph); if (ret == cudaErrorStreamCaptureInvalidated) { assert(mGraph == nullptr); @@ -267,9 +282,9 @@ class TrtCudaBuffer public: TrtCudaBuffer() = default; - TrtCudaBuffer(const TrtCudaBuffer&) = delete; + TrtCudaBuffer(TrtCudaBuffer const&) = delete; - TrtCudaBuffer& operator=(const TrtCudaBuffer&) = delete; + TrtCudaBuffer& operator=(TrtCudaBuffer const&) = delete; TrtCudaBuffer(TrtCudaBuffer&& rhs) { @@ -580,8 +595,7 @@ class OutputAllocator : public nvinfer1::IOutputAllocator void* reallocateOutput( char const* tensorName, void* currentMemory, uint64_t size, uint64_t alignment) noexcept override #else - void* reallocateOutput( - char const* tensorName, void* currentMemory, uint64_t size, uint64_t alignment) noexcept + void* reallocateOutput(char const* tensorName, void* currentMemory, uint64_t size, uint64_t alignment) noexcept #endif // !TRT_WINML { // Some memory allocators return nullptr when allocating zero bytes, but TensorRT requires a non-null ptr diff --git a/samples/common/sampleEngines.cpp b/samples/trtexecCommon/sampleEngines.cpp similarity index 79% rename from samples/common/sampleEngines.cpp rename to samples/trtexecCommon/sampleEngines.cpp index 0cd0a4c6b5..c69926d137 100644 --- a/samples/common/sampleEngines.cpp +++ b/samples/trtexecCommon/sampleEngines.cpp @@ -18,6 +18,11 @@ #include #include #include +#include +#include +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME +#include +#endif #include #include #include @@ -42,10 +47,6 @@ #include "sampleOptions.h" #include "sampleUtils.h" -#if ENABLE_UNIFIED_BUILDER -#include "safeCommon.h" -#endif - #if ENABLE_UNIFIED_BUILDER #include "NvInferConsistency.h" #include "NvInferReference.h" @@ -91,8 +92,441 @@ class FileStreamWriter final : public nvinfer1::IStreamWriter } }; +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME +bool checkCudaDriver(CUresult result, char const* operation, std::ostream& err) +{ + if (result == CUDA_SUCCESS) + { + return true; + } + + err << operation << " failed with CUDA driver error " << static_cast(result) << "." << std::endl; + return false; +} + +template +bool loadCudaDriverEntryPoint(char const* symbol, Function& function, uint32_t requestedVersion, std::ostream& err) +{ + void* entryPoint{nullptr}; + cudaDriverEntryPointQueryResult queryResult{}; + cudaError_t const result + = cudaGetDriverEntryPointByVersion(symbol, &entryPoint, requestedVersion, cudaEnableDefault, &queryResult); + if (result != cudaSuccess || queryResult != cudaDriverEntryPointSuccess || entryPoint == nullptr) + { + err << "CUDA driver entry point " << symbol << " is unavailable"; + if (result != cudaSuccess) + { + err << ": " << cudaGetErrorString(result); + } + err << "." << std::endl; + return false; + } + function = reinterpret_cast(entryPoint); + return true; +} + +struct GreenContextDriverApi +{ + [[nodiscard]] bool load(int32_t driverVersion, std::ostream& err) + { + bool const loaded = loadCudaDriverEntryPoint("cuDeviceGetDevResource", deviceGetDevResource, 12040U, err) + && loadCudaDriverEntryPoint("cuDeviceGet", deviceGet, 2000U, err) + && loadCudaDriverEntryPoint("cuDevResourceGenerateDesc", devResourceGenerateDesc, 12040U, err) + && loadCudaDriverEntryPoint("cuGreenCtxCreate", greenCtxCreate, 12040U, err) + && loadCudaDriverEntryPoint("cuGreenCtxDestroy", greenCtxDestroy, 12040U, err) + && loadCudaDriverEntryPoint("cuGreenCtxGetDevResource", greenCtxGetDevResource, 12040U, err) + && loadCudaDriverEntryPoint("cuGreenCtxStreamCreate", greenCtxStreamCreate, 12050U, err) + && loadCudaDriverEntryPoint("cuStreamDestroy", streamDestroy, 4000U, err); + if (!loaded) + { + return false; + } +#if CUDA_VERSION >= 13010 + if (driverVersion >= 13010) + { + return loadCudaDriverEntryPoint("cuDevSmResourceSplit", devSmResourceSplit, 13010U, err); + } +#else + static_cast(driverVersion); +#endif + return loadCudaDriverEntryPoint("cuDevSmResourceSplitByCount", devSmResourceSplitByCount, 12040U, err); + } + + PFN_cuDeviceGetDevResource_v12040 deviceGetDevResource{}; + PFN_cuDeviceGet_v2000 deviceGet{}; + PFN_cuDevSmResourceSplitByCount_v12040 devSmResourceSplitByCount{}; +#if CUDA_VERSION >= 13010 + PFN_cuDevSmResourceSplit_v13010 devSmResourceSplit{}; +#endif + PFN_cuDevResourceGenerateDesc_v12040 devResourceGenerateDesc{}; + PFN_cuGreenCtxCreate_v12040 greenCtxCreate{}; + PFN_cuGreenCtxDestroy_v12040 greenCtxDestroy{}; + PFN_cuGreenCtxGetDevResource_v12040 greenCtxGetDevResource{}; + PFN_cuGreenCtxStreamCreate_v12050 greenCtxStreamCreate{}; + PFN_cuStreamDestroy_v4000 streamDestroy{}; +}; +#endif + } // namespace +struct GreenContextManager::Impl +{ +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + struct Context + { + explicit Context(GreenContextDriverApi const& api) + : driverApi(&api) + { + } + Context(Context const&) = delete; + Context& operator=(Context const&) = delete; + + ~Context() + { + clearInferenceStreams(); + if (buildStream != nullptr) + { + static_cast(driverApi->streamDestroy(buildStream)); + } + if (greenContext != nullptr) + { + static_cast(driverApi->greenCtxDestroy(greenContext)); + } + } + + void clearInferenceStreams() noexcept + { + for (CUstream stream : inferenceStreams) + { + static_cast(driverApi->streamDestroy(stream)); + } + inferenceStreams.clear(); + } + + [[nodiscard]] bool createInferenceStreams(int32_t streamCount, std::ostream& err) + { + inferenceStreams.reserve(static_cast(streamCount)); + for (int32_t index = 0; index < streamCount; ++index) + { + CUstream stream{nullptr}; + if (!checkCudaDriver(driverApi->greenCtxStreamCreate(&stream, greenContext, CU_STREAM_NON_BLOCKING, 0), + "Creating a CUDA green context inference stream", err)) + { + clearInferenceStreams(); + return false; + } + inferenceStreams.push_back(stream); + } + return true; + } + + GreenContextDriverApi const* driverApi; + CUgreenCtx greenContext{nullptr}; + CUstream buildStream{nullptr}; + std::vector inferenceStreams; + uint32_t smCount{}; + uint32_t coscheduledSmCount{}; + }; + + [[nodiscard]] static std::unique_ptr createContext( + GreenContextSpec const& spec, CUdevice device, GreenContextDriverApi const& driverApi, std::ostream& err) + { + if (spec.smCount <= 0 || spec.coscheduledSmCount < 0) + { + err << "Invalid CUDA green context resource specification." << std::endl; + return nullptr; + } + + CUdevResource deviceSms{}; + if (!checkCudaDriver(driverApi.deviceGetDevResource(device, &deviceSms, CU_DEV_RESOURCE_TYPE_SM), + "Querying CUDA SM resources", err)) + { + return nullptr; + } + if (static_cast(spec.smCount) > deviceSms.sm.smCount) + { + err << "CUDA green context requests " << spec.smCount << " SMs, but device " << device << " has only " + << deviceSms.sm.smCount << " SMs." << std::endl; + return nullptr; + } + + CUdevResource group{}; +#if CUDA_VERSION >= 13010 + if (driverApi.devSmResourceSplit != nullptr) + { + CU_DEV_SM_RESOURCE_GROUP_PARAMS params{}; + params.smCount = static_cast(spec.smCount); + params.coscheduledSmCount = static_cast(spec.coscheduledSmCount); + if (!checkCudaDriver(driverApi.devSmResourceSplit(&group, 1U, &deviceSms, nullptr, 0U, ¶ms), + "Partitioning CUDA SM resources", err)) + { + return nullptr; + } + } + else +#endif + { + uint32_t groupCount{1U}; + if (!checkCudaDriver(driverApi.devSmResourceSplitByCount( + &group, &groupCount, &deviceSms, nullptr, 0U, static_cast(spec.smCount)), + "Partitioning CUDA SM resources", err)) + { + return nullptr; + } + if (groupCount != 1U) + { + err << "CUDA created " << groupCount << " SM resource groups instead of one." << std::endl; + return nullptr; + } + } + + CUdevResourceDesc descriptor{}; + if (!checkCudaDriver( + driverApi.devResourceGenerateDesc(&descriptor, &group, 1U), "Creating a CUDA resource descriptor", err)) + { + return nullptr; + } + + auto context = std::make_unique(driverApi); + if (!checkCudaDriver( + driverApi.greenCtxCreate(&context->greenContext, descriptor, device, CU_GREEN_CTX_DEFAULT_STREAM), + "Creating a CUDA green context", err)) + { + return nullptr; + } + CUdevResource contextSms{}; + if (!checkCudaDriver( + driverApi.greenCtxGetDevResource(context->greenContext, &contextSms, CU_DEV_RESOURCE_TYPE_SM), + "Querying CUDA green context resources", err)) + { + return nullptr; + } + context->smCount = contextSms.sm.smCount; + context->coscheduledSmCount = contextSms.sm.smCoscheduledAlignment; + bool const smCountMatches = context->smCount == static_cast(spec.smCount); + bool const coscheduledSmCountMatches = spec.coscheduledSmCount == 0 + || context->coscheduledSmCount == static_cast(spec.coscheduledSmCount); + if (!smCountMatches || !coscheduledSmCountMatches) + { + err << "CUDA created a green context with SMs=" << context->smCount + << " and co-scheduled SMs=" << context->coscheduledSmCount + << ", which does not match the requested SMs=" << spec.smCount; + if (spec.coscheduledSmCount != 0) + { + err << " and co-scheduled SMs=" << spec.coscheduledSmCount; + } + err << ". Use trtexec built with CUDA Toolkit 13.1 or newer and CUDA driver 13.1 or newer for exact " + "resource configuration." + << std::endl; + return nullptr; + } + if (!checkCudaDriver( + driverApi.greenCtxStreamCreate(&context->buildStream, context->greenContext, CU_STREAM_NON_BLOCKING, 0), + "Creating a CUDA green context build stream", err)) + { + return nullptr; + } + return context; + } + + void clearInferenceStreams() noexcept + { + if (globalContext != nullptr) + { + globalContext->clearInferenceStreams(); + } + for (auto& context : profileContexts) + { + if (context != nullptr) + { + context->clearInferenceStreams(); + } + } + inferenceContext = nullptr; + } + + GreenContextDriverApi driverApi; + std::unique_ptr globalContext; + std::vector> profileContexts; + Context* inferenceContext{nullptr}; +#endif +}; + +GreenContextManager::GreenContextManager() + : mImpl(std::make_unique()) +{ +} + +GreenContextManager::~GreenContextManager() = default; + +GreenContextManager::GreenContextManager(GreenContextManager&& other) noexcept = default; + +bool GreenContextManager::initialize(BuildOptions const& build, int32_t device, std::ostream& err) +{ + bool const requested = build.greenContext.has_value() + || std::any_of(build.profileGreenContexts.begin(), build.profileGreenContexts.end(), + [](auto const& context) { return context.has_value(); }); + if (!requested) + { + return true; + } + +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + if (build.cpuOnly) + { + err << "CUDA green contexts cannot be used with --cpuOnly." << std::endl; + return false; + } + if (build.safe) + { + err << "CUDA green contexts are not supported with --safe." << std::endl; + return false; + } + + cudaError_t const runtimeResult = cudaFree(nullptr); + if (runtimeResult != cudaSuccess) + { + err << "Initializing CUDA for green context creation failed: " << cudaGetErrorString(runtimeResult) + << std::endl; + return false; + } + + int32_t driverVersion{}; + cudaError_t const driverVersionResult = cudaDriverGetVersion(&driverVersion); + if (driverVersionResult != cudaSuccess) + { + err << "Querying the CUDA driver version failed: " << cudaGetErrorString(driverVersionResult) << std::endl; + return false; + } + if (driverVersion < 13000) + { + err << "--greenContext requires CUDA driver 13.0 or newer." << std::endl; + return false; + } + + if (!mImpl->driverApi.load(driverVersion, err)) + { + return false; + } + CUdevice cudaDevice{}; + if (!checkCudaDriver(mImpl->driverApi.deviceGet(&cudaDevice, device), "Getting the CUDA device", err)) + { + return false; + } + + if (build.greenContext.has_value()) + { + mImpl->globalContext = Impl::createContext(*build.greenContext, cudaDevice, mImpl->driverApi, err); + if (mImpl->globalContext == nullptr) + { + return false; + } + sample::gLogInfo << "Created global CUDA green context: SMs=" << mImpl->globalContext->smCount + << ", co-scheduled SMs=" << mImpl->globalContext->coscheduledSmCount << std::endl; + } + + mImpl->profileContexts.resize(build.profileGreenContexts.size()); + for (size_t profileIndex = 0; profileIndex < build.profileGreenContexts.size(); ++profileIndex) + { + if (!build.profileGreenContexts[profileIndex].has_value()) + { + continue; + } + mImpl->profileContexts[profileIndex] + = Impl::createContext(*build.profileGreenContexts[profileIndex], cudaDevice, mImpl->driverApi, err); + if (mImpl->profileContexts[profileIndex] == nullptr) + { + return false; + } + auto const& context = mImpl->profileContexts[profileIndex]; + sample::gLogInfo << "Created CUDA green context for optimization profile " << profileIndex + << ": SMs=" << context->smCount << ", co-scheduled SMs=" << context->coscheduledSmCount + << std::endl; + } + return true; +#else + static_cast(device); + err << "--greenContext requires CUDA Toolkit 13.0 or newer and is unavailable on this platform." << std::endl; + return false; +#endif +} + +cudaStream_t GreenContextManager::globalBuildStream() const noexcept +{ +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + return mImpl != nullptr && mImpl->globalContext != nullptr ? mImpl->globalContext->buildStream : nullptr; +#else + return nullptr; +#endif +} + +cudaStream_t GreenContextManager::profileBuildStream(size_t profileIndex) const noexcept +{ +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + if (mImpl == nullptr || profileIndex >= mImpl->profileContexts.size() + || mImpl->profileContexts[profileIndex] == nullptr) + { + return nullptr; + } + return mImpl->profileContexts[profileIndex]->buildStream; +#else + static_cast(profileIndex); + return nullptr; +#endif +} + +bool GreenContextManager::prepareInferenceStreams(size_t profileIndex, int32_t streamCount, std::ostream& err) +{ +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + if (mImpl == nullptr || streamCount < 0) + { + err << "Invalid CUDA green context inference stream count." << std::endl; + return false; + } + + mImpl->clearInferenceStreams(); + Impl::Context* context{nullptr}; + if (profileIndex < mImpl->profileContexts.size() && mImpl->profileContexts[profileIndex] != nullptr) + { + context = mImpl->profileContexts[profileIndex].get(); + } + else if (mImpl->globalContext != nullptr) + { + context = mImpl->globalContext.get(); + } + if (context == nullptr) + { + return true; + } + if (!context->createInferenceStreams(streamCount, err)) + { + return false; + } + mImpl->inferenceContext = context; + return true; +#else + static_cast(profileIndex); + static_cast(streamCount); + static_cast(err); + return true; +#endif +} + +cudaStream_t GreenContextManager::inferenceStream(size_t streamIndex) const noexcept +{ +#if CUDA_VERSION >= 13000 && !TRT_WINML && !IS_QNX_SAFE && !HOS_RUNTIME + if (mImpl == nullptr || mImpl->inferenceContext == nullptr + || streamIndex >= mImpl->inferenceContext->inferenceStreams.size()) + { + return nullptr; + } + return mImpl->inferenceContext->inferenceStreams[streamIndex]; +#else + static_cast(streamIndex); + return nullptr; +#endif +} + nvinfer1::ICudaEngine* LazilyDeserializedEngine::get() { SMP_RETVAL_IF_FALSE( @@ -617,17 +1051,30 @@ void setPreviewFeatures(IBuilderConfig& config, BuildOptions const& build) return true; } +//! Creates optimization profiles and assigns profile-specific CUDA green context streams. +//! +//! \return The optimization profiles created by \p builder. +[[nodiscard]] std::vector createOptimizationProfiles( + BuildOptions const& build, IBuilder& builder, GreenContextManager const& greenContexts) +{ + std::vector profiles(build.optProfiles.size()); + for (size_t profileIndex = 0; profileIndex < profiles.size(); ++profileIndex) + { + profiles[profileIndex] = builder.createOptimizationProfile(); + if (auto const stream = greenContexts.profileBuildStream(profileIndex); stream != nullptr) + { + profiles[profileIndex]->setProfileStream(stream); + } + } + return profiles; +} + // NOLINTNEXTLINE(readability-function-cognitive-complexity, readability-function-size) bool setupNetworkAndConfig(BuildOptions const& build, SystemOptions const& sys, IBuilder& builder, INetworkDefinition& network, IBuilderConfig& config, std::ostream& err, - std::vector>& sparseWeights) + std::vector>& sparseWeights, GreenContextManager const& greenContexts) { - std::vector profiles{}; - profiles.resize(build.optProfiles.size()); - for (auto& profile : profiles) - { - profile = builder.createOptimizationProfile(); - } + auto const profiles = createOptimizationProfiles(build, builder, greenContexts); bool hasDynamicShapes{false}; @@ -1142,52 +1589,13 @@ bool buildSerializedEngine(BuildOptions const& build, SystemOptions const& sys, INetworkDefinition& network, IBuilderConfig& config, BuildEnvironment& env, std::ostream& err) { std::unique_ptr serializedEngine; -#if ENABLE_UNIFIED_BUILDER - //! Engine bytes copied out of the safe artifacts, which free their own memory when they go out of scope. - std::vector safeEngineBytes; -#endif // ENABLE_UNIFIED_BUILDER #if !TRT_WINML -#if ENABLE_UNIFIED_BUILDER - if (build.safe) - { - // A safe engine may need a companion library, and only this entry point returns the two together. - // The others reject EngineCapability::kSAFETY once companion libraries are required, so safe builds - // go through here whether or not one is produced. - auto const artifacts = std::unique_ptr{ - builder.buildSerializedSafeNetwork(network, config, build.dumpCheckerBlob)}; - SMP_RETVAL_IF_FALSE(artifacts != nullptr, "Engine could not be created from network", false, err); - - // Every getter hands back memory owned by the artifacts, so each one is copied before they die. - auto const copyOut = [](IHostMemory const* src) { - auto const* const bytes = static_cast(src->data()); - return std::vector(bytes, bytes + src->size()); - }; - - safeEngineBytes = copyOut(artifacts->getSerializedNetwork()); - - if (auto const* const companionSo = artifacts->getCompanionSo()) - { - sample::gLogInfo << "Created companion library with size: " << (companionSo->size() / 1.0_MiB) << " MiB" - << std::endl; - env.companionSo.setBlob(copyOut(companionSo)); - } - - if (build.dumpCheckerBlob) - { - auto const* const checkerBlob = artifacts->getCheckerBlob(); - SMP_RETVAL_IF_FALSE(checkerBlob != nullptr, "Failed to create the checker blob.", false, err); - sample::gLogInfo << "Created checker blob with size: " << (checkerBlob->size() / 1.0_MiB) << " MiB" - << std::endl; - env.checkerBlob.setBlob(copyOut(checkerBlob)); - } - } - else -#endif // ENABLE_UNIFIED_BUILDER if (build.safe && build.save && build.dumpCheckerBlob) { - IHostMemory* kernelTextPtr{nullptr}; - serializedEngine = std::unique_ptr{builder.buildSerializedNetwork(network, config, kernelTextPtr)}; - auto checkerBlob = std::unique_ptr{kernelTextPtr}; + IHostMemory* checkerBlobPtr{nullptr}; + serializedEngine + = std::unique_ptr{builder.buildSerializedNetwork(network, config, checkerBlobPtr)}; + auto checkerBlob = std::unique_ptr{checkerBlobPtr}; if (checkerBlob != nullptr) { auto const checkerBlobSize = checkerBlob->size(); @@ -1202,10 +1610,9 @@ bool buildSerializedEngine(BuildOptions const& build, SystemOptions const& sys, sample::gLogInfo << "Created empty checker blob." << std::endl; } } - else + else if (serializedEngine != nullptr) { - sample::gLogError << "Failed to create the checker blob." << std::endl; - return false; + sample::gLogWarning << "No checker blob was produced: the engine has no checkable kernels." << std::endl; } } else @@ -1213,44 +1620,23 @@ bool buildSerializedEngine(BuildOptions const& build, SystemOptions const& sys, { serializedEngine = std::unique_ptr{builder.buildSerializedNetwork(network, config)}; } - void const* engineData{nullptr}; - int64_t engineSize{0}; -#if ENABLE_UNIFIED_BUILDER - if (!safeEngineBytes.empty()) - { - engineData = safeEngineBytes.data(); - engineSize = static_cast(safeEngineBytes.size()); - } - else -#endif // ENABLE_UNIFIED_BUILDER - { - SMP_RETVAL_IF_FALSE(serializedEngine != nullptr, "Engine could not be created from network", false, err); - engineData = serializedEngine->data(); - engineSize = static_cast(serializedEngine->size()); - } - sample::gLogInfo << "Created engine with size: " << (engineSize / 1.0_MiB) << " MiB" << std::endl; + SMP_RETVAL_IF_FALSE(serializedEngine != nullptr, "Engine could not be created from network", false, err); + sample::gLogInfo << "Created engine with size: " << (serializedEngine->size() / 1.0_MiB) << " MiB" << std::endl; if (build.safe && build.consistency) { - std::vector pluginBuildLibPaths; + std::vector pluginBuildLibPaths; #if ENABLE_UNIFIED_BUILDER pluginBuildLibPaths.reserve(sys.safetyPlugins.size()); std::transform(sys.safetyPlugins.begin(), sys.safetyPlugins.end(), std::back_inserter(pluginBuildLibPaths), - [](auto const& sp) { return sp.libraryName.c_str(); }); + [](auto const& sp) { return sp.libraryName; }); #endif - if (!checkSafeEngine(engineData, engineSize, pluginBuildLibPaths.data(), - static_cast(pluginBuildLibPaths.size()))) + if (!checkSafeEngine( + serializedEngine->data(), static_cast(serializedEngine->size()), pluginBuildLibPaths)) { return false; } } -#if ENABLE_UNIFIED_BUILDER - if (!safeEngineBytes.empty()) - { - env.engine.setBlob(std::move(safeEngineBytes)); - return true; - } -#endif // ENABLE_UNIFIED_BUILDER env.engine.setBlob(std::move(serializedEngine)); return true; } @@ -1260,14 +1646,15 @@ bool buildSerializedEngine(BuildOptions const& build, SystemOptions const& sys, //! //! \return Whether the engine creation succeeds or fails. //! -bool networkToSerializedEngine( - BuildOptions const& build, SystemOptions const& sys, BuildEnvironment& env, std::ostream& err, PostConfigCallback const& postConfigHook) +bool networkToSerializedEngine(BuildOptions const& build, SystemOptions const& sys, BuildEnvironment& env, + std::ostream& err, PostConfigCallback const& postConfigHook) { IBuilder& builder = *env.builder; IBuilderConfig& config = *env.builderConfig; INetworkDefinition& network = *env.network; std::vector> sparseWeights; - SMP_RETVAL_IF_FALSE(setupNetworkAndConfig(build, sys, builder, network, config, err, sparseWeights), + SMP_RETVAL_IF_FALSE( + setupNetworkAndConfig(build, sys, builder, network, config, err, sparseWeights, env.greenContexts), "Network And Config setup failed", false, err); if (postConfigHook) @@ -1284,13 +1671,15 @@ bool networkToSerializedEngine( // CUDA stream used for profiling by the builder. #if !TRT_WINML - auto profileStream = build.cpuOnly + auto const greenProfileStream = env.greenContexts.globalBuildStream(); + auto profileStream = build.cpuOnly || greenProfileStream != nullptr ? std::unique_ptr{nullptr, samplesCommon::StreamDeleter} : samplesCommon::makeCudaStream(); if (!build.cpuOnly) { - SMP_RETVAL_IF_FALSE(profileStream != nullptr, "Cuda stream creation failed", false, err); - config.setProfileStream(*profileStream); + SMP_RETVAL_IF_FALSE( + greenProfileStream != nullptr || profileStream != nullptr, "Cuda stream creation failed", false, err); + config.setProfileStream(greenProfileStream != nullptr ? greenProfileStream : *profileStream); } #endif auto const tBegin = std::chrono::high_resolution_clock::now(); @@ -1332,8 +1721,8 @@ bool networkToSerializedEngine( //! //! \brief Parse a given model, create a network and an engine. //! -bool modelToBuildEnv( - ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, std::ostream& err, PostConfigCallback const& postConfigHook) +bool modelToBuildEnv(ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, + std::ostream& err, PostConfigCallback const& postConfigHook) { env.builder.reset(createBuilder()); SMP_RETVAL_IF_FALSE(env.builder != nullptr, "Builder creation failed", false, err); @@ -1375,8 +1764,8 @@ bool modelToBuildEnv( std::vector vcPluginLibrariesUsed; SMP_RETVAL_IF_FALSE(env.network != nullptr, "Network creation failed", false, err); - env.parser - = modelToNetwork(model, build, *env.network, err, build.versionCompatible ? &vcPluginLibrariesUsed : nullptr, *env.builderConfig); + env.parser = modelToNetwork(model, build, *env.network, err, + build.versionCompatible ? &vcPluginLibrariesUsed : nullptr, *env.builderConfig); SMP_RETVAL_IF_FALSE(env.parser.operator bool(), "Parsing model failed", false, err); #if !TRT_WINML @@ -1509,10 +1898,13 @@ bool loadEngineToBuildEnv(std::string const& filepath, BuildEnvironment& env, st std::ifstream engineFile(filepath, std::ios::binary); SMP_RETVAL_IF_FALSE(engineFile.good(), "", false, err << "Error opening engine file: " << filepath); engineFile.seekg(0, std::ifstream::end); - int64_t fsize = engineFile.tellg(); + // tellg() returns -1 on stream failure; cast-to-size_t would otherwise produce a huge allocation. + std::streamoff const fsizeRaw = engineFile.tellg(); + SMP_RETVAL_IF_FALSE(fsizeRaw > 0, "", false, err << "Error determining engine file size: " << filepath); + int64_t const fsize = static_cast(fsizeRaw); engineFile.seekg(0, std::ifstream::beg); - std::vector engineBlob(fsize); + std::vector engineBlob(static_cast(fsize)); engineFile.read(reinterpret_cast(engineBlob.data()), fsize); SMP_RETVAL_IF_FALSE(engineFile.good(), "", false, err << "Error loading engine file: " << filepath); auto const tEnd = std::chrono::high_resolution_clock::now(); @@ -1522,14 +1914,13 @@ bool loadEngineToBuildEnv(std::string const& filepath, BuildEnvironment& env, st if (enableConsistency) { - std::vector pluginBuildLibPaths; + std::vector pluginBuildLibPaths; #if ENABLE_UNIFIED_BUILDER pluginBuildLibPaths.reserve(sys.safetyPlugins.size()); std::transform(sys.safetyPlugins.begin(), sys.safetyPlugins.end(), std::back_inserter(pluginBuildLibPaths), - [](auto const& sp) { return sp.libraryName.c_str(); }); + [](auto const& sp) { return sp.libraryName; }); #endif - if (!checkSafeEngine(engineBlob.data(), static_cast(fsize), pluginBuildLibPaths.data(), - static_cast(pluginBuildLibPaths.size()))) + if (!checkSafeEngine(engineBlob.data(), fsize, pluginBuildLibPaths)) { sample::gLogError << "Consistency validation is not enabled." << std::endl; return false; @@ -1598,9 +1989,9 @@ bool printPlanVersion(BuildEnvironment& env, std::ostream& err) case 0U: { // Blob index to store the plan version may depend on the serialization version. - sample::gLogInfo << "Plan was created with TensorRT version " << static_cast(blob[24]) - << "." << static_cast(blob[25]) << "." << static_cast(blob[26]) - << "." << static_cast(blob[27]) << std::endl; + sample::gLogInfo << "Plan was created with TensorRT version " << static_cast(blob[24]) << "." + << static_cast(blob[25]) << "." << static_cast(blob[26]) << "." + << static_cast(blob[27]) << std::endl; return true; } } @@ -1664,8 +2055,8 @@ bool saveEngine(ICudaEngine const& engine, std::string const& fileName, std::ost } // NOLINTNEXTLINE(readability-function-cognitive-complexity) -bool getEngineBuildEnv( - ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, std::ostream& err, PostConfigCallback const& postConfigHook) +bool getEngineBuildEnv(ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, + std::ostream& err, PostConfigCallback const& postConfigHook) { bool createEngineSuccess{false}; @@ -1674,19 +2065,6 @@ bool getEngineBuildEnv( if (build.safe) { createEngineSuccess = loadEngineToBuildEnv(build.engine, env, err, sys, build.safe && build.consistency); -#if ENABLE_UNIFIED_BUILDER - // An explicit path is taken as given, so naming a library that is not there is an error. - // Otherwise the default beside the engine is used when it exists; an engine built without a - // companion library has none to pair with and loads on its own. - env.companionSoPath = samplesSafeCommon::resolveCompanionSoPath(build.engine, build.loadEngineSo); - if (env.companionSoPath) - { - std::ifstream companionSoFile(*env.companionSoPath, std::ios::binary); - SMP_RETVAL_IF_FALSE(companionSoFile.good(), - ("Companion library not found: " + *env.companionSoPath).c_str(), false, err); - sample::gLogInfo << "Using companion library " << *env.companionSoPath << std::endl; - } -#endif // ENABLE_UNIFIED_BUILDER } else { @@ -1787,26 +2165,6 @@ bool getEngineBuildEnv( << std::endl; } } -#if ENABLE_UNIFIED_BUILDER - // The runtime cannot load the engine without its companion library, so it is saved whenever the - // build produced one rather than behind a flag of its own. - if (build.safe && !env.companionSo.hasBlob() && !build.saveEngineSo.empty()) - { - sample::gLogWarning << "Ignoring --saveEngineSo: the build produced no companion library." << std::endl; - } - if (build.safe && env.companionSo.hasBlob()) - { - auto const companionSoFileName - = build.saveEngineSo.empty() ? build.engine + ".so" : build.saveEngineSo; - auto const companionSoBlob = env.companionSo.getBlobOrEmpty(); - std::ofstream companionSoFile(companionSoFileName, std::ios::binary); - companionSoFile.write(static_cast(companionSoBlob.data), companionSoBlob.size); - SMP_RETVAL_IF_FALSE(!companionSoFile.fail(), "Saving companion library to file failed.", false, err); - companionSoFile.close(); - env.companionSoPath = companionSoFileName; - sample::gLogInfo << "Saved companion library to " << companionSoFileName << std::endl; - } -#endif // ENABLE_UNIFIED_BUILDER } return true; @@ -2097,12 +2455,22 @@ static constexpr auto kREFERENCE_CHECKER_LIBRARY = nullptr; #if ENABLE_UNIFIED_BUILDER +//! \brief Create a consistency checker for a serialized safe engine. +//! +//! \param recorder Error recorder the checker reports through. +//! \param serializedEngine Serialized engine to validate. Must not be null. +//! \param engineSize Size of \p serializedEngine in bytes. Must be positive and representable as size_t. +//! \param pluginBuildLibs Plugin libraries the engine was built against, loaded by the checker. +//! +//! \return The checker, or nullptr if the arguments are invalid or the checker library is unavailable. std::unique_ptr createConsistencyChecker( nvinfer2::safe::ISafeRecorder& recorder, void const* serializedEngine, int64_t const engineSize, - char const* const* pluginBuildLibs, int64_t const nbPluginBuildLibs) noexcept + std::vector const& pluginBuildLibs) noexcept { - - if (serializedEngine == nullptr || engineSize <= 0) + // The checker takes a size_t, which is narrower than int64_t on 32-bit targets. Round-tripping the cast + // rejects sizes that would truncate, which would otherwise validate only part of the engine. + if (serializedEngine == nullptr || engineSize <= 0 + || static_cast(static_cast(engineSize)) != engineSize) { return nullptr; } @@ -2110,16 +2478,14 @@ std::unique_ptr createConsiste #if !defined(_WIN32) if (hasSafeRuntime()) { - using CreateCheckerFn = nvinfer2::safe::ErrorCode (*)( - nvinfer2::safe::consistency::IConsistencyChecker*& checker, nvinfer2::safe::ISafeRecorder& recorder, - void const* data, int64_t size, char const* const* pluginBuildLibs, int64_t nbPluginBuildLibs) noexcept; + // Derive the signature from the header so an ABI change in the checker library is a compile error + // rather than a segfault at the call. + using CreateCheckerFn = decltype(&nvinfer2::safe::consistency::createConsistencyChecker); if (auto const createFn = reinterpret_cast(dlsym(kCONSISTENCY_CHECKER_LIBRARY.get(), "createConsistencyChecker"))) { - nvinfer2::safe::consistency::IConsistencyChecker* checker{nullptr}; - auto const result - = createFn(checker, recorder, serializedEngine, engineSize, pluginBuildLibs, nbPluginBuildLibs); - if (result == nvinfer2::safe::ErrorCode::kSUCCESS) + if (nvinfer2::safe::consistency::IConsistencyChecker * checker{nullptr}; nvinfer2::safe::ErrorCode::kSUCCESS + == createFn(checker, recorder, serializedEngine, static_cast(engineSize), pluginBuildLibs)) { return std::unique_ptr{checker}; } @@ -2151,8 +2517,7 @@ std::unique_ptr createReferenceChe if (auto const createFn = reinterpret_cast(dlsym(kREFERENCE_CHECKER_LIBRARY.get(), symbolName))) { - if (nvinfer2::safe::reference::IReferenceChecker * checker{nullptr}; - ErrorCode::kSUCCESS + if (nvinfer2::safe::reference::IReferenceChecker * checker{nullptr}; ErrorCode::kSUCCESS == createFn(checker, recorder, serializedEngine, engineSize, checkerBlob.data, checkerBlob.size, remoteConfig.c_str())) { @@ -2175,13 +2540,13 @@ bool hasConsistencyChecker() return kCONSISTENCY_CHECKER_LIBRARY != nullptr; } -bool hasReferenceChecker() +[[nodiscard]] bool hasReferenceChecker() { return kREFERENCE_CHECKER_LIBRARY != nullptr; } -bool checkSafeEngine(void const* serializedEngine, int64_t const engineSize, char const* const* pluginBuildLibs, - int64_t const nbPluginBuildLibs) +bool checkSafeEngine( + void const* serializedEngine, int64_t const engineSize, std::vector const& pluginBuildLibs) { #if !ENABLE_UNIFIED_BUILDER return false; @@ -2194,7 +2559,7 @@ bool checkSafeEngine(void const* serializedEngine, int64_t const engineSize, cha sample::SampleSafeRecorder recorder{nvinfer2::safe::Severity::kINFO}; std::unique_ptr checker - = createConsistencyChecker(recorder, serializedEngine, engineSize, pluginBuildLibs, nbPluginBuildLibs); + = createConsistencyChecker(recorder, serializedEngine, engineSize, pluginBuildLibs); if (checker == nullptr) { sample::gLogError << "Failed to create consistency checker." << std::endl; diff --git a/samples/common/sampleEngines.h b/samples/trtexecCommon/sampleEngines.h similarity index 83% rename from samples/common/sampleEngines.h rename to samples/trtexecCommon/sampleEngines.h index e11a0a34ee..4593f16568 100644 --- a/samples/common/sampleEngines.h +++ b/samples/trtexecCommon/sampleEngines.h @@ -29,7 +29,7 @@ #include #include #include -#include +#include #include namespace sample @@ -37,8 +37,8 @@ namespace sample //! \brief Callback invoked after standard builder configuration, before engine build. //! Custom tools can use this to apply additional builder configuration on top of trtexec's. -using PostConfigCallback = std::function; +using PostConfigCallback + = std::function; #if TRT_BUILD_ONNX_PARSER struct Parser @@ -301,6 +301,47 @@ class LazilyDeserializedEngine //!@} }; +//! \brief Owns CUDA green contexts and their streams for engine building and inference. +class GreenContextManager +{ +public: + //! Construct an empty CUDA green context manager. + GreenContextManager(); + + //! Destroy all owned streams and CUDA green contexts. + ~GreenContextManager(); + + GreenContextManager(GreenContextManager const&) = delete; + GreenContextManager& operator=(GreenContextManager const&) = delete; + + //! Move the CUDA green context resources owned by \p other. + GreenContextManager(GreenContextManager&& other) noexcept; + + //! \brief Creates the CUDA green contexts requested by \p build on \p device. + //! + //! \return True on success, including when no CUDA green context was requested. + bool initialize(BuildOptions const& build, int32_t device, std::ostream& err); + + //! \return The builder-level CUDA green context stream, or nullptr when none was requested. + [[nodiscard]] cudaStream_t globalBuildStream() const noexcept; + + //! \return The CUDA green context stream explicitly assigned to \p profileIndex, or nullptr. + [[nodiscard]] cudaStream_t profileBuildStream(size_t profileIndex) const noexcept; + + //! \brief Creates \p streamCount inference streams in the effective CUDA green context for \p profileIndex. + //! + //! A profile-specific CUDA green context takes precedence over the global context. + //! \return True on success, including when the effective profile has no CUDA green context. + bool prepareInferenceStreams(size_t profileIndex, int32_t streamCount, std::ostream& err); + + //! \return Inference stream \p streamIndex, or nullptr when ordinary CUDA streams should be used. + [[nodiscard]] cudaStream_t inferenceStream(size_t streamIndex) const noexcept; + +private: + struct Impl; + std::unique_ptr mImpl; +}; + struct BuildEnvironment { BuildEnvironment() = delete; @@ -311,13 +352,13 @@ struct BuildEnvironment std::string const& cmdline = "") : engine(isSafe, versionCompatible, DLACore, tempdir, tempfileControls, leanDLLPath) , checkerBlob(false, false, -1, "", tempfileControls, "") -#if ENABLE_UNIFIED_BUILDER - , companionSo(false, false, -1, "", tempfileControls, "") -#endif // ENABLE_UNIFIED_BUILDER , cmdline(cmdline) { } + //! CUDA green contexts and streams used by the builder and runtime. + GreenContextManager greenContexts; + //! \name Owned TensorRT objects //! Per TensorRT object lifetime requirements as outlined in the developer guide, //! factory objects must remain live while the objects created by those factories @@ -347,16 +388,6 @@ struct BuildEnvironment //! the reference checker replays. LazilyDeserializedEngine checkerBlob; -#if ENABLE_UNIFIED_BUILDER - //! The companion library holding the safe engine's generated host code. Loading the engine needs it, so - //! it is saved beside the engine and handed back to the runtime at load. - LazilyDeserializedEngine companionSo; -#endif // ENABLE_UNIFIED_BUILDER - - //! Path to the engine's companion library on disk, std::nullopt when the engine needs none. The - //! runtime loads the library by path, so it has to exist as a file before inference. - std::optional companionSoPath; - //! The command line string. std::string cmdline; //!@} @@ -386,15 +417,15 @@ bool saveEngine(nvinfer1::ICudaEngine const& engine, std::string const& fileName //! //! \return Pointer to the engine created or nullptr if the creation failed //! -bool getEngineBuildEnv( - ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, std::ostream& err, PostConfigCallback const& postConfigHook = nullptr); +bool getEngineBuildEnv(ModelOptions const& model, BuildOptions const& build, SystemOptions& sys, BuildEnvironment& env, + std::ostream& err, PostConfigCallback const& postConfigHook = nullptr); //! //! \brief Create a serialized network //! //! \return Pointer to a host memory for a serialized network //! -nvinfer1::IHostMemory* networkToSerialized(const BuildOptions& build, const SystemOptions& sys, +nvinfer1::IHostMemory* networkToSerialized(BuildOptions const& build, SystemOptions const& sys, nvinfer1::IBuilder& builder, nvinfer1::INetworkDefinition& network, std::ostream& err); //! @@ -403,7 +434,7 @@ nvinfer1::IHostMemory* networkToSerialized(const BuildOptions& build, const Syst //! \return Pointer to a host memory for a serialized network //! nvinfer1::IHostMemory* modelToSerialized( - const ModelOptions& model, const BuildOptions& build, const SystemOptions& sys, std::ostream& err); + ModelOptions const& model, BuildOptions const& build, SystemOptions const& sys, std::ostream& err); //! //! \brief Serialize network and save it into a file @@ -411,7 +442,7 @@ nvinfer1::IHostMemory* modelToSerialized( //! \return boolean Return true if the network was successfully serialized and saved //! bool serializeAndSave( - const ModelOptions& model, const BuildOptions& build, const SystemOptions& sys, std::ostream& err); + ModelOptions const& model, BuildOptions const& build, SystemOptions const& sys, std::ostream& err); #if TRT_BUILD_ONNX_PARSER //! @@ -432,13 +463,16 @@ bool timeRefit(nvinfer1::INetworkDefinition const& network, nvinfer1::ICudaEngin //! \brief Check if safe runtime is loaded. [[nodiscard]] bool hasSafeRuntime(); +//! \brief Run a consistency check on a serialized safe engine. //! -//! \brief Run consistency check on serialized engine. +//! \param serializedEngine Serialized engine to validate. Must not be null. +//! \param engineSize Size of \p serializedEngine in bytes. Must be positive. +//! \param pluginBuildLibs Plugin libraries the engine was built against, loaded by the checker. //! -[[nodiscard]] bool checkSafeEngine(void const* serializedEngine, int64_t const engineSize, - char const* const* pluginBuildLibs, int64_t const nbPluginBuildLibs); +//! \return True if the engine passes the consistency check, false if it fails or if no checker is available. +[[nodiscard]] bool checkSafeEngine( + void const* serializedEngine, int64_t const engineSize, std::vector const& pluginBuildLibs); -//! //! \brief Run the per-kernel reference check on a serialized safe engine. //! //! Confirms that the per-kernel metadata a reference check is derived from describes this engine, and @@ -450,7 +484,8 @@ bool timeRefit(nvinfer1::INetworkDefinition const& network, nvinfer1::ICudaEngin //! empty: an engine of nothing but static library kernels needs none, and the checker reports the //! omission for any engine that does. //! \param remoteConfig The remote target connection token (from --remoteConfig). -//! +//! \return true when every kernel the checker examined matched its reference; false when any kernel +//! mismatched, or when the check could not be run at all. [[nodiscard]] bool referenceCheckEngine(void const* serializedEngine, int64_t const engineSize, EngineBlob const& checkerBlob, std::string const& remoteConfig); @@ -467,7 +502,6 @@ bool loadStreamingEngineToBuildEnv(std::string const& engine, BuildEnvironment& bool loadEngineToBuildEnv(std::string const& engine, BuildEnvironment& env, std::ostream& err, SystemOptions const& sys, bool const enableConsistency); -//! //! \brief Load the checker blob that carries the reference-check metadata into \p env. //! //! Reads the path given by --loadCheckerBlob. Omitting the flag is not an error: only the checker can @@ -477,7 +511,6 @@ bool loadEngineToBuildEnv(std::string const& engine, BuildEnvironment& env, std: //! \param env The environment to load into. //! \param err Stream for diagnostics. //! \return false if --loadCheckerBlob names a file that could not be read. -//! [[nodiscard]] bool loadCheckerBlobToBuildEnv(BuildOptions const& build, BuildEnvironment& env, std::ostream& err); } // namespace sample diff --git a/samples/trtexecCommon/sampleEntrypoints.h b/samples/trtexecCommon/sampleEntrypoints.h new file mode 100644 index 0000000000..9aeb22b134 --- /dev/null +++ b/samples/trtexecCommon/sampleEntrypoints.h @@ -0,0 +1,113 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef TRT_SAMPLE_ENTRYPOINTS_H +#define TRT_SAMPLE_ENTRYPOINTS_H + +//! \file sampleEntrypoints.h +//! +//! Declares and conditionally defines entrypoints needed to create base TensorRT objects, depending +//! on whether the given sample uses TRT at link time or dynamically. Since common code is built once +//! and shared across all samples (both link-time and dynamic TRT), it does not define these entrypoints, +//! so each sample must define them individually. +//! +//! Samples that use TRT at link time can define DEFINE_TRT_ENTRYPOINTS before including this header to +//! pick up the definitions here. + +#include "NvInfer.h" +#if TRT_BUILD_ONNX_PARSER +#include "NvOnnxParser.h" +#endif +#include "logger.h" + +extern nvinfer1::IBuilder* createBuilder(); +extern nvinfer1::IRuntime* createRuntime(); +extern nvinfer1::IRefitter* createRefitter(nvinfer1::ICudaEngine& engine); +#if TRT_BUILD_ONNX_PARSER +extern nvonnxparser::IParser* createONNXParser(nvinfer1::INetworkDefinition& network); +extern nvonnxparser::IParserRefitter* createONNXRefitter(nvinfer1::IRefitter& refitter); +#endif + +#if !defined(DEFINE_TRT_ENTRYPOINTS) +#define DEFINE_TRT_ENTRYPOINTS 0 +#endif + +// Allow opting out of individual entrypoints that are unused by the sample +#if !defined(DEFINE_TRT_BUILDER_ENTRYPOINT) +#define DEFINE_TRT_BUILDER_ENTRYPOINT 1 +#endif +#if !defined(DEFINE_TRT_RUNTIME_ENTRYPOINT) +#define DEFINE_TRT_RUNTIME_ENTRYPOINT 1 +#endif +#if !defined(DEFINE_TRT_REFITTER_ENTRYPOINT) +#define DEFINE_TRT_REFITTER_ENTRYPOINT 1 +#endif +#if !defined(DEFINE_TRT_ONNX_PARSER_ENTRYPOINT) +#define DEFINE_TRT_ONNX_PARSER_ENTRYPOINT 1 +#endif + +#if DEFINE_TRT_ENTRYPOINTS +nvinfer1::IBuilder* createBuilder() +{ +#if DEFINE_TRT_BUILDER_ENTRYPOINT + return nvinfer1::createInferBuilder(sample::gLogger.getTRTLogger()); +#else + return {}; +#endif +} + +nvinfer1::IRuntime* createRuntime() +{ +#if DEFINE_TRT_RUNTIME_ENTRYPOINT + return nvinfer1::createInferRuntime(sample::gLogger.getTRTLogger()); +#else + return {}; +#endif +} + +nvinfer1::IRefitter* createRefitter(nvinfer1::ICudaEngine& engine) +{ +#if DEFINE_TRT_REFITTER_ENTRYPOINT + return nvinfer1::createInferRefitter(engine, sample::gLogger.getTRTLogger()); +#else + return {}; +#endif +} + +#if TRT_BUILD_ONNX_PARSER +nvonnxparser::IParser* createONNXParser(nvinfer1::INetworkDefinition& network) +{ +#if DEFINE_TRT_ONNX_PARSER_ENTRYPOINT + return nvonnxparser::createParser(network, sample::gLogger.getTRTLogger()); +#else + return {}; +#endif +} + +nvonnxparser::IParserRefitter* createONNXRefitter(nvinfer1::IRefitter& refitter) +{ +#if DEFINE_TRT_ONNX_PARSER_ENTRYPOINT + return nvonnxparser::createParserRefitter(refitter, sample::gLogger.getTRTLogger()); +#else + return {}; +#endif +} +#endif // TRT_BUILD_ONNX_PARSER + +#endif // DEFINE_TRT_ENTRYPOINTS + +#endif // TRT_SAMPLE_ENTRYPOINTS_H diff --git a/samples/common/sampleInference.cpp b/samples/trtexecCommon/sampleInference.cpp similarity index 97% rename from samples/common/sampleInference.cpp rename to samples/trtexecCommon/sampleInference.cpp index 79973db9d1..c2d8f89bed 100644 --- a/samples/common/sampleInference.cpp +++ b/samples/trtexecCommon/sampleInference.cpp @@ -86,12 +86,11 @@ namespace safe { namespace { -//! Function pointer to the safe runtime's `createTRTGraph` symbol. -//! Bound by `initNvinferSafe()` when `gUseRuntime == RuntimeMode::kSAFE`; stays empty until then, or if -//! the library fails to load. +//! Function pointer to the safe runtime's `createTRTGraphWithSo` symbol. +//! Bound by `initNvinferSafe()`; stays empty until then, or if the library fails to load. std::function - sCreateTrtGraphInternal{}; + sCreateTrtGraphWithSoInternal{}; //! Function pointer to the safe runtime's `destroyTRTGraph` symbol. Bound as above. std::function sDestroyTrtGraphInternal{}; @@ -120,7 +119,7 @@ namespace //! The function performs the following operations: //! - Dynamically loads the safe TensorRT runtime library //! - Retrieves and stores function pointers for: -//! - createTRTGraph: Creates a safe TRT graph from serialized engine data +//! - createTRTGraphWithSo: Creates a safe TRT graph from serialized engine data //! - destroyTRTGraph: Destroys a safe TRT graph and releases resources //! - getSafePluginRegistry: Gets the safe plugin registry for loading plugins //! @@ -132,9 +131,9 @@ bool initNvinferSafe(SafeRuntimeSettings const& settings) #if !TRT_STATIC static LibraryPtr libnvinfersafePtr{}; auto fetchPtrs = [](samplesCommon::DynamicLibrary& l) { - sCreateTrtGraphInternal = l.symbolAddress( - "createTRTGraph"); + sCreateTrtGraphWithSoInternal + = l.symbolAddress("createTRTGraphWithSo"); sDestroyTrtGraphInternal = l.symbolAddress("destroyTRTGraph"); @@ -155,16 +154,15 @@ bool initNvinferSafe(SafeRuntimeSettings const& settings) //! a safe TRT graph for inference with safety-certified TensorRT engines. //! nvinfer1::ErrorCode createSafeTRTGraph(nvinfer2::safe::ITRTGraph*& graph, void const* blob, int64_t size, - nvinfer2::safe::AsciiChar const* companionSoPath, ISafeRecorder& recorder, bool useManaged, - ISafeMemAllocator* allocator, SafeRuntimeSettings const& settings) + ISafeRecorder& recorder, bool useManaged, ISafeMemAllocator* allocator, SafeRuntimeSettings const& settings) { if (!initNvinferSafe(settings)) { return nvinfer1::ErrorCode::kINTERNAL_ERROR; } - ASSERT(sCreateTrtGraphInternal != nullptr); - // A null path is valid: it means this engine needs no companion library. - return sCreateTrtGraphInternal(graph, blob, size, companionSoPath, recorder, useManaged, allocator); + ASSERT(sCreateTrtGraphWithSoInternal != nullptr); + // No companion .so is bound yet; a null path matches the behavior of the legacy createTRTGraph. + return sCreateTrtGraphWithSoInternal(graph, blob, size, nullptr, recorder, useManaged, allocator); } //! @@ -398,7 +396,7 @@ bool allocateContextMemory(InferenceEnvironmentStd& iEnv, InferenceOptions const else { size_t sizeToAlloc{0}; - const char* allocReason{nullptr}; + char const* allocReason{nullptr}; if (inference.memoryAllocationStrategy == MemoryAllocationStrategy::kPROFILE) { auto const p = inference.optProfileIndex; @@ -493,20 +491,23 @@ IRuntimeConfig* setJITRuntimeConfig(nvinfer1::ICudaEngine* engine, InferenceOpti if (!inference.runtimeCacheFile.empty()) { nvinfer1::IRuntimeCache* runtimeCache = runtimeConfig->createRuntimeCache(); - // deserialize runtime cache from file + // load runtime cache from file + auto const rtcLoadBegin = std::chrono::steady_clock::now(); std::vector loadedCacheBytes = samplesCommon::loadCacheFile(sample::gLogger, inference.runtimeCacheFile); std::vector runtimeCacheBytes(loadedCacheBytes.begin(), loadedCacheBytes.end()); + auto const rtcLoadEnd = std::chrono::steady_clock::now(); if (!loadedCacheBytes.empty()) { - std::vector runtimeCacheBytes(loadedCacheBytes.begin(), loadedCacheBytes.end()); + // deserialize runtime cache auto const rtcDeserializeBegin = std::chrono::steady_clock::now(); runtimeCache->deserialize(runtimeCacheBytes.data(), runtimeCacheBytes.size()); auto const rtcDeserializeEnd = std::chrono::steady_clock::now(); sample::gLogInfo << "Runtime Cache deserialized in " - << std::chrono::duration(rtcDeserializeEnd - rtcDeserializeBegin).count() << " ms." - << std::endl; + << std::chrono::duration(rtcDeserializeEnd - rtcDeserializeBegin).count() + << " ms (File loaded in " << std::chrono::duration(rtcLoadEnd - rtcLoadBegin).count() + << " ms)." << std::endl; } // The runtime cache is portable only within a matching environment: a loaded cache is // rejected (and ignored for execution) if the GPU device/SKU, the TensorRT-RTX version, or @@ -579,8 +580,8 @@ bool populateRuntimeCacheForDeferredJit(nvinfer1::ICudaEngine& engine, Inference return false; } auto const tEnd = std::chrono::high_resolution_clock::now(); - sample::gLogInfo << "Deferred JIT compilation in " << std::chrono::duration(tEnd - tBegin).count() - << " sec." << std::endl; + sample::gLogInfo << "Deferred JIT compilation in " << std::chrono::duration(tEnd - tBegin).count() << " sec." + << std::endl; return serializeRuntimeCache(context.get(), inference); } @@ -605,7 +606,7 @@ void getSafeTensorInfo(uint32_t profileIndex, nvinfer2::safe::ITRTGraph* safeGra { nvinfer2::safe::TensorDescriptor desc; auto const b = tensorInfo.bindingIndex; - const char* name = nullptr; + char const* name = nullptr; safeGraph->getIOTensorName(name, b); tensorInfo.name = name; safeGraph->getIOTensorDescriptor(desc, name); @@ -638,9 +639,8 @@ bool setUpSafeInference(InferenceEnvironmentSafe& iEnv, InferenceOptions const& bool const useManagedMemory{inference.useManaged}; nvinfer2::safe::ITRTGraph* tempGraph = nullptr; - auto const* const companionSoPath = iEnv.companionSoPath ? iEnv.companionSoPath->c_str() : nullptr; - if (sample::safe::createSafeTRTGraph(tempGraph, safeEngineBlob.data, safeEngineBlob.size, companionSoPath, - *gSafeRecorder, useManagedMemory, nullptr, iEnv.safeRuntimeSettings) + if (sample::safe::createSafeTRTGraph(tempGraph, safeEngineBlob.data, safeEngineBlob.size, *gSafeRecorder, + useManagedMemory, nullptr, iEnv.safeRuntimeSettings) != nvinfer2::safe::ErrorCode::kSUCCESS) { sample::gLogError << "Create Safe TRT Graph Failed." << std::endl; @@ -854,6 +854,12 @@ bool setUpStdInference(InferenceEnvironmentStd& iEnv, InferenceOptions const& in << std::endl; } + if (!iEnv.greenContexts.prepareInferenceStreams( + static_cast(inference.optProfileIndex), inference.infStreams, sample::gLogError)) + { + return false; + } + for (int32_t s = 0; s < inference.infStreams; ++s) { IExecutionContext* ec = setupExecutionContext(iEnv, engine, inference, properties); @@ -1130,7 +1136,7 @@ class EnqueueExplicit : private Enqueue } return result; } - catch (const std::exception&) + catch (std::exception const&) { return false; } @@ -1172,7 +1178,7 @@ class EnqueueExplicitSafe : private SafeEnqueue bool const result = (mGraph.executeAsync(stream.get()) == nvinfer1::ErrorCode::kSUCCESS); return result; } - catch (const std::exception&) + catch (std::exception const&) { return false; } @@ -1274,11 +1280,13 @@ class IterationBase { public: - explicit IterationBase(int32_t id, InferenceOptions const& inference, BindingsBase& bindings) + explicit IterationBase( + int32_t id, InferenceOptions const& inference, BindingsBase& bindings, cudaStream_t computeStream = nullptr) : mBindings(bindings) , mStreamId(id) , mDepth(1 + inference.overlap) , mActive(mDepth) + , mStream{{TrtCudaStream{}, TrtCudaStream{computeStream}, TrtCudaStream{}}} , mEvents(mDepth) , mEnqueueTimes(mDepth) { @@ -1460,9 +1468,9 @@ class IterationBase class IterationStd : public IterationBase { public: - explicit IterationStd( - int32_t id, InferenceOptions const& inference, nvinfer1::IExecutionContext& context, BindingsStd& bindings) - : IterationBase(id, inference, bindings) + explicit IterationStd(int32_t id, InferenceOptions const& inference, nvinfer1::IExecutionContext& context, + BindingsStd& bindings, cudaStream_t computeStream) + : IterationBase(id, inference, bindings, computeStream) { createEnqueueFunction(inference, context, bindings); } @@ -1924,7 +1932,8 @@ void inferenceExecution(InferenceOptions const& inference, InferenceEnvironmentB int32_t const streamId{threadIdx * streamsPerThread + s}; auto iteration = std::make_unique(streamId, inference, *static_cast(iEnv).getContext(streamId), - *static_cast(iEnv).bindings[streamId]); + *static_cast(iEnv).bindings[streamId], + iEnv.greenContexts.inferenceStream(static_cast(streamId))); if (!inference.includeTransfers) { iteration->setInputData(true); @@ -2068,9 +2077,8 @@ bool runMultiTasksInference(std::vectoriOptions, *(tEnv->iEnv), sync, /*threadIdx*/ 0, /*streamsPerThread*/ 1, tEnv->device, tEnv->trace, - tEnv->rOptions)); + threads.emplace_back(makeThread(tEnv->iOptions, *(tEnv->iEnv), sync, /*threadIdx*/ 0, /*streamsPerThread*/ 1, + tEnv->device, tEnv->trace, tEnv->rOptions)); } for (auto& th : threads) { diff --git a/samples/common/sampleInference.h b/samples/trtexecCommon/sampleInference.h similarity index 97% rename from samples/common/sampleInference.h rename to samples/trtexecCommon/sampleInference.h index 128dee09a9..4938985f17 100644 --- a/samples/common/sampleInference.h +++ b/samples/trtexecCommon/sampleInference.h @@ -28,11 +28,12 @@ #include #include #include -#include #include #include #if ENABLE_UNIFIED_BUILDER +// Also pulls in safeErrorRecorder.h, whose `using namespace nvinfer2::safe` is what lets the +// declarations below name ISafeRecorder, ISafeMemAllocator and ITRTGraph unqualified. #include "safeCudaAllocator.h" #endif namespace sample @@ -136,7 +137,6 @@ bool initNvinferSafe(SafeRuntimeSettings const& settings); //! \param graph: Pointer to the safe TRT graph to be created //! \param blob: Pointer to the serialized engine data //! \param size: Size of the serialized engine data -//! \param companionSoPath: Path to the engine's companion library, or nullptr when it needs none //! \param recorder: Reference to the safe recorder //! \param useManaged: Flag indicating whether to use managed memory //! \param allocator: Pointer to the safe memory allocator @@ -144,8 +144,7 @@ bool initNvinferSafe(SafeRuntimeSettings const& settings); //! \return Error code indicating the success or failure of the operation //! nvinfer1::ErrorCode createSafeTRTGraph(nvinfer2::safe::ITRTGraph*& graph, void const* blob, int64_t size, - nvinfer2::safe::AsciiChar const* companionSoPath, ISafeRecorder& recorder, bool useManaged, - ISafeMemAllocator* allocator, SafeRuntimeSettings const& settings); + ISafeRecorder& recorder, bool useManaged, ISafeMemAllocator* allocator, SafeRuntimeSettings const& settings); //! //! \brief Destroy a safe TRT graph and release resources @@ -181,17 +180,15 @@ struct InferenceEnvironmentBase InferenceEnvironmentBase(InferenceEnvironmentBase const& other) = delete; InferenceEnvironmentBase(InferenceEnvironmentBase&& other) = delete; InferenceEnvironmentBase(BuildEnvironment& bEnv) - : engine(std::move(bEnv.engine)) - , companionSoPath(bEnv.companionSoPath) + : greenContexts(std::move(bEnv.greenContexts)) + , engine(std::move(bEnv.engine)) , safe(bEnv.engine.isSafe()) , cmdline(bEnv.cmdline) { } + GreenContextManager greenContexts; LazilyDeserializedEngine engine; - - //! Path to the engine's companion library, std::nullopt when the engine needs none. - std::optional companionSoPath; std::unique_ptr profiler; std::vector deviceMemory; //< Device memory used for inference when the allocation strategy is not static. @@ -567,8 +564,7 @@ class BindingsSafe : public BindingsBase struct TaskInferenceEnvironment { TaskInferenceEnvironment(std::string engineFile, InferenceOptions const& inference, - ReportingOptions const& reporting, int32_t deviceId = 0, - int32_t DLACore = -1, int32_t bs = batchNotProvided); + ReportingOptions const& reporting, int32_t deviceId = 0, int32_t DLACore = -1, int32_t bs = batchNotProvided); InferenceOptions iOptions{}; ReportingOptions rOptions{}; int32_t device{defaultDevice}; diff --git a/samples/common/sampleOptions.cpp b/samples/trtexecCommon/sampleOptions.cpp similarity index 95% rename from samples/common/sampleOptions.cpp rename to samples/trtexecCommon/sampleOptions.cpp index 49ce61499a..d97140e475 100644 --- a/samples/common/sampleOptions.cpp +++ b/samples/trtexecCommon/sampleOptions.cpp @@ -17,9 +17,12 @@ #include #include +#include #include #include #include +#include +#include #include #include #include @@ -253,7 +256,8 @@ IOFormat stringToValue(std::string const& option) template <> SparsityFlag stringToValue(std::string const& option) { - std::unordered_map const table{ + std::unordered_map const table + { {"disable", SparsityFlag::kDISABLE}, {"enable", SparsityFlag::kENABLE}, #if !TRT_WINML { @@ -791,6 +795,157 @@ void processShapes(BuildOptions::ShapeProfile& shapes, bool minShapes, bool optS shapes = newShapes; } +//! Throw a consistently formatted error for an invalid green context specification. +[[noreturn]] void throwInvalidGreenContext(std::string const& spec, std::string const& reason) +{ + throw std::invalid_argument("Invalid --greenContext specification '" + spec + "': " + reason); +} + +//! Parse a positive SM count from a green context specification. +[[nodiscard]] int32_t parseGreenContextCount( + std::string_view const value, std::string const& name, std::string const& spec) +{ + uint32_t count{}; + char const* const end = value.data() + value.size(); + auto const [ptr, error] = std::from_chars(value.data(), end, count); + if (error == std::errc::result_out_of_range || count > static_cast(std::numeric_limits::max())) + { + throwInvalidGreenContext( + spec, name + " must not exceed " + std::to_string(std::numeric_limits::max()) + "."); + } + if (error != std::errc{} || ptr != end) + { + throwInvalidGreenContext(spec, name + " must be a positive decimal integer."); + } + if (count == 0) + { + throwInvalidGreenContext(spec, name + " must be greater than zero."); + } + return static_cast(count); +} + +//! Parse the command-line grammar for one green context. +[[nodiscard]] GreenContextSpec parseGreenContextSpec(std::string const& value) +{ + GreenContextSpec spec{}; + if (value.find(':') == std::string::npos) + { + spec.smCount = parseGreenContextCount(value, "SM count", value); + return spec; + } + + bool foundSmCount{false}; + bool foundCoscheduledSmCount{false}; + for (std::string const& fieldString : splitToStringVec(value, ',')) + { + std::string_view const field{fieldString}; + size_t const colon = field.find(':'); + if (field.empty() || colon == std::string_view::npos || colon == 0 || colon + 1 == field.size() + || field.find(':', colon + 1) != std::string_view::npos) + { + throwInvalidGreenContext(value, "expected sm:[,cosched:] or a bare SM count."); + } + + std::string_view const key = field.substr(0, colon); + std::string_view const count = field.substr(colon + 1); + if (key == "sm") + { + if (foundSmCount) + { + throwInvalidGreenContext(value, "sm may be specified only once."); + } + spec.smCount = parseGreenContextCount(count, "SM count", value); + foundSmCount = true; + } + else if (key == "cosched") + { + if (!foundSmCount) + { + throwInvalidGreenContext(value, "cosched must follow sm."); + } + if (foundCoscheduledSmCount) + { + throwInvalidGreenContext(value, "cosched may be specified only once."); + } + spec.coscheduledSmCount = parseGreenContextCount(count, "co-scheduled SM count", value); + foundCoscheduledSmCount = true; + } + else + { + throwInvalidGreenContext(value, "unknown key '" + std::string(key) + "'."); + } + } + + if (!foundSmCount) + { + throwInvalidGreenContext(value, "sm is required."); + } + return spec; +} + +//! Assign each green context to the nearest preceding optimization profile, or to the engine if none precedes it. +void getGreenContexts(Arguments const& arguments, std::optional& greenContext, + std::vector>& profileGreenContexts) +{ + auto const greenContextRange = arguments.equal_range("--greenContext"); + if (greenContextRange.first == greenContextRange.second) + { + return; + } + + struct PositionedProfile + { + int32_t position; + size_t index; + }; + struct PositionedGreenContext + { + int32_t position; + GreenContextSpec spec; + }; + + std::vector profiles; + auto const profileRange = arguments.equal_range("--profile"); + std::transform(profileRange.first, profileRange.second, std::back_inserter(profiles), [](auto const& entry) { + return PositionedProfile{entry.second.second, stringToValue(entry.second.first)}; + }); + std::ranges::sort(profiles, {}, &PositionedProfile::position); + + std::vector contexts; + std::transform( + greenContextRange.first, greenContextRange.second, std::back_inserter(contexts), [](auto const& entry) { + return PositionedGreenContext{entry.second.second, parseGreenContextSpec(entry.second.first)}; + }); + std::ranges::sort(contexts, {}, &PositionedGreenContext::position); + + for (auto const& context : contexts) + { + auto const nextProfile = std::ranges::upper_bound(profiles, context.position, {}, &PositionedProfile::position); + if (nextProfile == profiles.begin()) + { + if (greenContext) + { + throw std::invalid_argument( + "--greenContext can be specified at most once outside an optimization profile block."); + } + greenContext = context.spec; + continue; + } + + size_t const profileIndex = std::prev(nextProfile)->index; + if (profileIndex >= profileGreenContexts.size()) + { + profileGreenContexts.resize(profileIndex + 1); + } + if (profileGreenContexts[profileIndex]) + { + throw std::invalid_argument("--greenContext can be specified at most once for optimization profile " + + std::to_string(profileIndex) + "."); + } + profileGreenContexts[profileIndex] = context.spec; + } +} + bool getOptimizationProfiles( Arguments& arguments, std::vector& optProfiles, char const* argument) { @@ -813,8 +968,9 @@ bool getOptimizationProfiles( while (getAndDelOptionWithPosition(arguments, argument, profileIndex, pos)) { BuildOptions::ShapeProfile optProfile{}; - bool minShapes{false}, maxShapes{false}, optShapes{false}; - for (int32_t i = 0; i < nvinfer1::EnumMax(); i++, pos++) + bool minShapes{false}, maxShapes{false}, optShapes{false}, greenContext{false}; + constexpr int32_t kMAX_PROFILE_OPTIONS{nvinfer1::EnumMax() + 1}; + for (int32_t i = 0; i < kMAX_PROFILE_OPTIONS; i++, pos++) { std::string value; @@ -833,6 +989,10 @@ bool getOptimizationProfiles( optShapes = true; getShapes(optProfile, value, nvinfer1::OptProfileSelector::kOPT); } + else if (!greenContext && getAndDelOptionBehind(arguments, "--greenContext", pos, value)) + { + greenContext = true; + } else { break; @@ -1021,8 +1181,7 @@ std::ostream& printPrecision(std::ostream& os, BuildOptions const& options) std::ostream& printTempfileControls(std::ostream& os, TempfileControlFlags const tempfileControls) { - auto getFlag = [&](TempfileControlFlag f) -> char const* - { + auto getFlag = [&](TempfileControlFlag f) -> char const* { bool allowed = !!(tempfileControls & (1U << static_cast(f))); return allowed ? "allow" : "deny"; }; @@ -1059,8 +1218,7 @@ std::ostream& printSparsity(std::ostream& os, BuildOptions const& options) std::ostream& printMemoryPools(std::ostream& os, BuildOptions const& options) { - auto const printValueOrDefault = [&os](double const val, char const* unit = "MiB") - { + auto const printValueOrDefault = [&os](double const val, char const* unit = "MiB") { if (val >= 0) { os << val << " " << unit; @@ -1248,6 +1406,7 @@ void BuildOptions::parse(Arguments& arguments) getFormats(inputFormats, "--inputIOFormats"); getFormats(outputFormats, "--outputIOFormats"); #endif // TRT_WINML + getGreenContexts(arguments, greenContext, profileGreenContexts); if (!getOptimizationProfiles(arguments, optProfiles, "--profile")) { ShapeProfile shapes; @@ -1268,6 +1427,8 @@ void BuildOptions::parse(Arguments& arguments) processShapes(shapes, minShapes, optShapes, maxShapes); optProfiles.emplace_back(shapes); } + arguments.erase("--greenContext"); + profileGreenContexts.resize(optProfiles.size()); BuildOptions::ShapeProfile dummyShapes; bool remainingMinShapes = getShapesBuild(arguments, dummyShapes, "--minShapes", nvinfer1::OptProfileSelector::kMIN); @@ -1394,15 +1555,6 @@ void BuildOptions::parse(Arguments& arguments) { throw std::invalid_argument("--dumpCheckerBlob requires --safe to be enabled."); } - -#if ENABLE_UNIFIED_BUILDER - getAndDelOption(arguments, "--saveEngineSo", saveEngineSo); - getAndDelOption(arguments, "--loadEngineSo", loadEngineSo); - if ((!saveEngineSo.empty() || !loadEngineSo.empty()) && !safe) - { - throw std::invalid_argument("--saveEngineSo and --loadEngineSo require --safe to be enabled."); - } -#endif // ENABLE_UNIFIED_BUILDER getAndDelOption(arguments, "--buildDLAStandalone", buildDLAStandalone); getAndDelOption(arguments, "--allowGPUFallback", allowGPUFallback); getAndDelOption(arguments, "--consistency", consistency); @@ -2641,6 +2793,16 @@ std::ostream& operator<<(std::ostream& os, IOFormat const& format) return os; } +std::ostream& operator<<(std::ostream& os, GreenContextSpec const& spec) +{ + os << "sm:" << spec.smCount; + if (spec.coscheduledSmCount != 0) + { + os << ",cosched:" << spec.coscheduledSmCount; + } + return os; +} + std::ostream& operator<<(std::ostream& os, nvinfer1::DeviceType devType) { switch (devType) @@ -2855,11 +3017,26 @@ std::ostream& operator<<(std::ostream& os, BuildOptions const& options) } }; + os << "Green Context: "; + if (options.greenContext) + { + os << *options.greenContext; + } + else + { + os << "Disabled"; + } + os << std::endl; + printIOFormats(os, "Input(s)", options.inputFormats); printIOFormats(os, "Output(s)", options.outputFormats); for (size_t i = 0; i < options.optProfiles.size(); i++) { printShapes(os, "build", options.optProfiles[i], i); + if (i < options.profileGreenContexts.size() && options.profileGreenContexts[i]) + { + os << "Optimization Profile " << i << " Green Context: " << *options.profileGreenContexts[i] << std::endl; + } } return os; } @@ -3250,12 +3427,6 @@ void BuildOptions::help(std::ostream& os) " sources the kernel checker analyses and the metadata that the" "\n" " reference checker replays." "\n" " --dumpKernelText is accepted as an alias." "\n" -#if ENABLE_UNIFIED_BUILDER - " --saveEngineSo= Save the safe engine's companion library to file, defaulting to" "\n" - " .so. Ignored when the build produces no companion library." "\n" - " --loadEngineSo= Load the safe engine's companion library from file. Without it," "\n" - " .so is used when that file exists." "\n" -#endif // ENABLE_UNIFIED_BUILDER " --buildDLAStandalone Enable build DLA standalone loadable which can be loaded by cuDLA, when this option is enabled, " "\n" " --allowGPUFallback is disallowed and --skipInference is enabled by default. Additionally, " "\n" #if ENABLE_FEATURE_WEAK_TYPING @@ -3341,6 +3512,20 @@ void BuildOptions::help(std::ostream& os) " --profile Build with dynamic shapes using a profile with the min/max/opt shapes provided. Can be specified" "\n" " multiple times to create multiple profiles with contiguous index." "\n" " (ex: --profile=0 --minShapes= --optShapes= --maxShapes= --profile=1 ...)" "\n" +#if !TRT_WINML + " --greenContext=spec Build and run under a CUDA green context that trtexec creates." "\n" + " Requires CUDA Toolkit and driver 13.0 or newer. CUDA 13.1 or newer guarantees exact resource" "\n" + " configuration; CUDA 13.0 accepts only requests that its default resource alignment can satisfy" "\n" + " exactly." "\n" + " Green Context: spec ::= sm:[,cosched:] | " "\n" + " sm: Number of SMs in the partition (required)." "\n" + " cosched: Co-scheduled SM count (optional). Defaults to the architecture alignment:" "\n" + " 2 on compute capability 7.x/8.x and Tegra; 8 on compute capability 9.0+." "\n" + " Without a preceding --profile, the setting applies to the whole engine." "\n" + " When placed after --profile=N, the setting applies to optimization profile N." "\n" + " A profile-specific setting overrides the global setting. Profiles without one inherit the" "\n" + " global setting, or use standard device behavior when no global setting exists." "\n" +#endif // !TRT_WINML #if !TRT_WINML " --allowWeightStreaming Enable a weight streaming engine. TensorRT will disable" "\n" #else // !TRT_WINML diff --git a/samples/common/sampleOptions.h b/samples/trtexecCommon/sampleOptions.h similarity index 94% rename from samples/common/sampleOptions.h rename to samples/trtexecCommon/sampleOptions.h index 48f1404a05..567fc71c34 100644 --- a/samples/common/sampleOptions.h +++ b/samples/trtexecCommon/sampleOptions.h @@ -196,6 +196,14 @@ struct IOFormat nvinfer1::TensorFormats formats{}; }; +//! A CUDA green context resource specification. +struct GreenContextSpec +{ + int32_t smCount{}; + //! Zero means unspecified; CUDA selects the architecture-specific co-scheduling alignment. + int32_t coscheduledSmCount{}; +}; + using ShapeRange = std::array, nvinfer1::EnumMax()>; #if ENABLE_FEATURE_WEAK_TYPING @@ -299,12 +307,6 @@ class BuildOptions : public Options bool reference{false}; bool dumpCheckerBlob{false}; std::string checkerBlob; -#if ENABLE_UNIFIED_BUILDER - //! Where to write, or read, the safe engine's companion library. Empty means the default beside the - //! engine, which only applies where companion libraries are enabled. - std::string saveEngineSo; - std::string loadEngineSo; -#endif // ENABLE_UNIFIED_BUILDER bool buildDLAStandalone{false}; bool allowGPUFallback{false}; bool skipInference{false}; @@ -339,6 +341,8 @@ class BuildOptions : public Options std::string engine; using ShapeProfile = std::unordered_map; std::vector optProfiles; + std::optional greenContext; + std::vector> profileGreenContexts; std::vector inputFormats; std::vector outputFormats; nvinfer1::TacticSources enabledTactics{0}; @@ -525,20 +529,20 @@ class TuningOptions : public Options { public: std::string tuningCacheFile{"best_config.json"}; - std::string tuningExpr{}; //!< --tuneBuildRoutes - std::string tuningExprFile{}; //!< --tuneBuildRouteFile + std::string tuningExpr{}; //!< --tuneBuildRoutes + std::string tuningExprFile{}; //!< --tuneBuildRouteFile TuningSearchAlgorithm tuningSearchAlgorithm{TuningSearchAlgorithm::kFAST}; - int64_t timeout{-1}; //!< --tuningTimeOut (s); -1 = no timeout - bool helpBuildRoute{false}; //!< --helpBuildRoute (short-circuit) - std::string helpBuildRouteKnob{}; //!< --helpBuildRoute= filter - bool continueFromCache{false}; //!< --continue - bool dryRun{false}; //!< --dryRun (enumerate, don't build) + int64_t timeout{-1}; //!< --tuningTimeOut (s); -1 = no timeout + bool helpBuildRoute{false}; //!< --helpBuildRoute (short-circuit) + std::string helpBuildRouteKnob{}; //!< --helpBuildRoute= filter + bool continueFromCache{false}; //!< --continue + bool dryRun{false}; //!< --dryRun (enumerate, don't build) //! \brief Hidden parent->child IPC channel. //! //! When set, runOnceBuildAndInfer writes a small JSON to this path containing //! gpu_time_ms, accuracy_failed, and per-tensor accuracy_loss before returning. //! Injected into the child's argv by the tuning loop; never shown in --help. - std::string tuningResultFile{}; //!< --tuningResultFile= + std::string tuningResultFile{}; //!< --tuningResultFile= void parse(Arguments& arguments) override; static void help(std::ostream& out); @@ -589,6 +593,8 @@ std::ostream& operator<<(std::ostream& os, BaseModelOptions const& options); std::ostream& operator<<(std::ostream& os, IOFormat const& format); +std::ostream& operator<<(std::ostream& os, GreenContextSpec const& spec); + std::ostream& operator<<(std::ostream& os, ShapeRange const& dims); std::ostream& operator<<(std::ostream& os, ModelOptions const& options); diff --git a/samples/common/sampleOptions.test.cpp b/samples/trtexecCommon/sampleOptions.test.cpp similarity index 100% rename from samples/common/sampleOptions.test.cpp rename to samples/trtexecCommon/sampleOptions.test.cpp diff --git a/samples/common/sampleReporting.cpp b/samples/trtexecCommon/sampleReporting.cpp similarity index 97% rename from samples/common/sampleReporting.cpp rename to samples/trtexecCommon/sampleReporting.cpp index 8d5fd09d01..931d643975 100644 --- a/samples/common/sampleReporting.cpp +++ b/samples/trtexecCommon/sampleReporting.cpp @@ -108,7 +108,7 @@ float findCoeffOfVariance(std::vector const& timings, T const& to return std::sqrt(variance) / mean * 100.F; } -inline InferenceTime traceToTiming(const InferenceTrace& a) +inline InferenceTime traceToTiming(InferenceTrace const& a) { return InferenceTime( (a.enqEnd - a.enqStart), (a.h2dEnd - a.h2dStart), (a.computeEnd - a.computeStart), (a.d2hEnd - a.d2hStart)); @@ -247,7 +247,7 @@ void printEpilog(std::vector const& timings, float walltimeMs, st auto const getD2h = [](InferenceTime const& t) { return t.d2h; }; auto const d2hResult = getPerformanceResult(timings, getD2h, percentiles); - auto const toPerfString = [&](const PerformanceResult& r) { + auto const toPerfString = [&](PerformanceResult const& r) { std::stringstream s; s << "min = " << r.min << " ms, max = " << r.max << " ms, mean = " << r.mean << " ms, " << "median = " << r.median << " ms"; @@ -325,7 +325,7 @@ void printPerformanceReport(std::vector const& trace, ReportingO { int32_t batchSize = infOpts.batch; float const warmupMs = infOpts.warmup; - auto const isNotWarmup = [&warmupMs](const InferenceTrace& a) { return a.computeStart >= warmupMs; }; + auto const isNotWarmup = [&warmupMs](InferenceTrace const& a) { return a.computeStart >= warmupMs; }; auto const noWarmup = std::find_if(trace.begin(), trace.end(), isNotWarmup); int32_t const warmups = noWarmup - trace.begin(); float const benchTime = trace.back().d2hEnd - noWarmup->h2dStart; @@ -663,26 +663,26 @@ void printOutput(ReportingOptions const& reporting, InferenceEnvironmentBase con if (iEnv.safe) { #if ENABLE_UNIFIED_BUILDER - auto const& binding = static_cast(iEnv).bindings.at(0); + auto const& binding = static_cast(iEnv).bindings.at(0); if (!binding) { sample::gLogError << "Empty bindings! Skip printing outputs." << std::endl; return; } - auto const& graph = static_cast(iEnv).mClonedGraphs.at(0); + auto const& graph = static_cast(iEnv).mClonedGraphs.at(0); details::safeDump(graph, binding, reporting, batch); #else sample::gLogWarning << "Safe mode is not supported! Skip printing outputs." << std::endl; #endif return; } - auto const& binding = static_cast(iEnv).bindings.at(0); + auto const& binding = static_cast(iEnv).bindings.at(0); if (!binding) { sample::gLogError << "Empty bindings! Skip printing outputs." << std::endl; return; } - auto const& context = static_cast(iEnv).contexts.at(0); + auto const& context = static_cast(iEnv).contexts.at(0); details::dump(context, binding, reporting, batch); } diff --git a/samples/common/sampleReporting.h b/samples/trtexecCommon/sampleReporting.h similarity index 100% rename from samples/common/sampleReporting.h rename to samples/trtexecCommon/sampleReporting.h diff --git a/samples/common/sampleTuning.cpp b/samples/trtexecCommon/sampleTuning.cpp similarity index 99% rename from samples/common/sampleTuning.cpp rename to samples/trtexecCommon/sampleTuning.cpp index 007953b62d..03c1fde9c9 100644 --- a/samples/common/sampleTuning.cpp +++ b/samples/trtexecCommon/sampleTuning.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/common/sampleTuning.h b/samples/trtexecCommon/sampleTuning.h similarity index 99% rename from samples/common/sampleTuning.h rename to samples/trtexecCommon/sampleTuning.h index 40a9de8b16..ee8eb13bb4 100644 --- a/samples/common/sampleTuning.h +++ b/samples/trtexecCommon/sampleTuning.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); diff --git a/samples/trtexecCommon/sampleUtils.cpp b/samples/trtexecCommon/sampleUtils.cpp new file mode 100644 index 0000000000..3ed0c88657 --- /dev/null +++ b/samples/trtexecCommon/sampleUtils.cpp @@ -0,0 +1,1413 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "sampleUtils.h" +#include "bfloat16.h" +#include "common.h" +#include "half.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#if CUDA_VERSION >= 11060 +#include +#endif + +using namespace nvinfer1; +using namespace std::string_view_literals; + +namespace sample +{ + +using TensorToLayer = std::unordered_map; +using LayerToTensor = std::unordered_map; +using TensorToTensor = std::unordered_map; + +int64_t volume(nvinfer1::Dims const& dims, nvinfer1::Dims const& strides, int32_t vecDim, int32_t comps, int32_t batch) +{ + int64_t maxNbElems = 1; + for (int32_t i = 0; i < dims.nbDims; ++i) + { + // Get effective length of axis. + int64_t d = dims.d[i]; + // Any dimension is 0, it is an empty tensor. + if (d == 0) + { + return 0; + } + if (i == vecDim) + { + d = samplesCommon::divUp(d, comps); + } + maxNbElems = std::max(maxNbElems, d * strides.d[i]); + } + return maxNbElems * batch * (vecDim < 0 ? 1 : comps); +} + +nvinfer1::Dims toDims(std::vector const& vec) +{ + int32_t limit = static_cast(nvinfer1::Dims::MAX_DIMS); + if (static_cast(vec.size()) > limit) + { + sample::gLogWarning << "Vector too long, only first 8 elements are used in dimension." << std::endl; + } + // Pick first nvinfer1::Dims::MAX_DIMS elements + nvinfer1::Dims dims{std::min(static_cast(vec.size()), limit), {}}; + std::copy_n(vec.begin(), dims.nbDims, std::begin(dims.d)); + return dims; +} + +void loadFromFile(std::string const& fileName, char* dst, size_t size) +{ + ASSERT(dst); + + std::ifstream file(fileName, std::ios::in | std::ios::binary); + if (file.is_open()) + { + file.seekg(0, std::ios::end); + int64_t fileSize = static_cast(file.tellg()); + // Due to change from int32_t to int64_t VC engines created with earlier versions + // may expect input of the half of the size + if (fileSize != static_cast(size) && fileSize != static_cast(size * 2)) + { + std::ostringstream msg; + msg << "Unexpected file size for input file: " << fileName << ". Note: Input binding size is: " << size + << " bytes but the file size is " << fileSize + << " bytes. Double check the size and datatype of the provided data."; + throw std::invalid_argument(msg.str()); + } + // Move file pointer back to the beginning after reading file size. + file.seekg(0, std::ios::beg); + file.read(dst, size); + size_t const nbBytesRead = file.gcount(); + file.close(); + if (nbBytesRead != size) + { + std::ostringstream msg; + msg << "Unexpected file size for input file: " << fileName << ". Note: Expected: " << size + << " bytes but only read: " << nbBytesRead << " bytes"; + throw std::invalid_argument(msg.str()); + } + } + else + { + std::ostringstream msg; + msg << "Cannot open file " << fileName << "!"; + throw std::invalid_argument(msg.str()); + } +} + +std::vector splitToStringVec(std::string const& s, char separator, int64_t maxSplit) +{ + std::vector splitted; + + for (size_t start = 0; start < s.length();) + { + // If maxSplit is specified and we have reached maxSplit, emplace back the rest of the string and break the + // loop. + if (maxSplit >= 0 && static_cast(splitted.size()) == maxSplit) + { + splitted.emplace_back(s.substr(start, s.length() - start)); + break; + } + + size_t separatorIndex = s.find(separator, start); + if (separatorIndex == std::string::npos) + { + separatorIndex = s.length(); + } + splitted.emplace_back(s.substr(start, separatorIndex - start)); + + // If the separator is the last character, then we should push an empty string at the end. + if (separatorIndex == s.length() - 1) + { + splitted.emplace_back(""); + } + + start = separatorIndex + 1; + } + + return splitted; +} + +bool broadcastIOFormats(std::vector const& formats, size_t nbBindings, bool isInput /*= true*/) +{ + bool broadcast = formats.size() == 1; + bool validFormatsCount = broadcast || (formats.size() == nbBindings); + if (!formats.empty() && !validFormatsCount) + { + if (isInput) + { + throw std::invalid_argument( + "The number of inputIOFormats must match network's inputs or be one for broadcasting."); + } + + throw std::invalid_argument( + "The number of outputIOFormats must match network's outputs or be one for broadcasting."); + } + return broadcast; +} + +// NOLINTNEXTLINE(readability-function-cognitive-complexity) +void sparsifyMatMulKernelWeights(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) +{ + // 1. Collect layers and tensors information from the network. + TensorToLayer matmulI2L; + TensorToLayer constO2L; + TensorToLayer shuffleI2L; + LayerToTensor shuffleL2O; + auto collectMappingInfo = [&](int32_t const idx) { + ILayer* l = network.getLayer(idx); + switch (l->getType()) + { + case nvinfer1::LayerType::kMATRIX_MULTIPLY: + { + // assume weights on the second input. + matmulI2L.insert({l->getInput(1), l}); + break; + } + case nvinfer1::LayerType::kCONSTANT: + { + DataType const dtype = static_cast(l)->getWeights().type; + if (dtype == nvinfer1::DataType::kFLOAT || dtype == nvinfer1::DataType::kHALF) + { + // Sparsify float only. + constO2L.insert({l->getOutput(0), l}); + } + break; + } + case nvinfer1::LayerType::kSHUFFLE: + { + shuffleI2L.insert({l->getInput(0), l}); + shuffleL2O.insert({l, l->getOutput(0)}); + break; + } + default: break; + } + }; + int32_t const nbLayers = network.getNbLayers(); + for (int32_t i = 0; i < nbLayers; ++i) + { + collectMappingInfo(i); + } + if (matmulI2L.size() == 0 || constO2L.size() == 0) + { + // No MatrixMultiply or Constant layer found, no weights to sparsify. + return; + } + + // Helper for analysis + auto isTranspose + = [](nvinfer1::Permutation const& perm) -> bool { return (perm.order[0] == 1 && perm.order[1] == 0); }; + auto is2D = [](nvinfer1::Dims const& dims) -> bool { return dims.nbDims == 2; }; + auto isIdenticalReshape = [](nvinfer1::Dims const& dims) -> bool { + for (int32_t i = 0; i < dims.nbDims; ++i) + { + if (dims.d[i] != i || dims.d[i] != -1) + { + return false; + } + } + return true; + }; + auto tensorReachedViaTranspose = [&](nvinfer1::ITensor* t, bool& needTranspose) -> ITensor* { + while (shuffleI2L.contains(t)) + { + nvinfer1::IShuffleLayer* s = static_cast(shuffleI2L.at(t)); + if (!is2D(s->getInput(0)->getDimensions()) || !is2D(s->getReshapeDimensions()) + || !isIdenticalReshape(s->getReshapeDimensions())) + { + break; + } + + if (isTranspose(s->getFirstTranspose())) + { + needTranspose = !needTranspose; + } + if (isTranspose(s->getSecondTranspose())) + { + needTranspose = !needTranspose; + } + + t = shuffleL2O.at(s); + } + return t; + }; + + // 2. Forward analysis to collect the Constant layers connected to MatMul via Transpose + std::unordered_map constantLayerToSparse; + for (auto& o2l : constO2L) + { + // If need to transpose the weights of the Constant layer. + // Need to transpose by default due to semantic difference. + bool needTranspose{true}; + ITensor* t = tensorReachedViaTranspose(o2l.first, needTranspose); + if (!matmulI2L.contains(t)) + { + continue; + } + + // check MatMul params... + IMatrixMultiplyLayer* mm = static_cast(matmulI2L.at(t)); + bool const twoInputs = mm->getNbInputs() == 2; + bool const all2D = is2D(mm->getInput(0)->getDimensions()) && is2D(mm->getInput(1)->getDimensions()); + bool const isSimple = mm->getOperation(0) == nvinfer1::MatrixOperation::kNONE + && mm->getOperation(1) != nvinfer1::MatrixOperation::kVECTOR; + if (!(twoInputs && all2D && isSimple)) + { + continue; + } + if (mm->getOperation(1) == nvinfer1::MatrixOperation::kTRANSPOSE) + { + needTranspose = !needTranspose; + } + + constantLayerToSparse.insert({static_cast(o2l.second), needTranspose}); + } + + // 3. Finally, sparsify the weights + auto sparsifyConstantWeights = [&sparseWeights](nvinfer1::IConstantLayer* layer, bool const needTranspose) { + Dims dims = layer->getOutput(0)->getDimensions(); + ASSERT(dims.nbDims == 2); + int32_t const idxN = needTranspose ? 1 : 0; + int32_t const n = dims.d[idxN]; + int32_t const k = dims.d[1 - idxN]; + sparseWeights.emplace_back(); + std::vector& spw = sparseWeights.back(); + Weights w = layer->getWeights(); + DataType const dtype = w.type; + ASSERT(dtype == nvinfer1::DataType::kFLOAT + || dtype == nvinfer1::DataType::kHALF); // non-float weights should have been ignored. + + if (needTranspose) + { + if (dtype == nvinfer1::DataType::kFLOAT) + { + spw.resize(w.count * sizeof(float)); + transpose2DWeights(spw.data(), w.values, k, n); + } + else if (dtype == nvinfer1::DataType::kHALF) + { + spw.resize(w.count * sizeof(half_float::half)); + transpose2DWeights(spw.data(), w.values, k, n); + } + + w.values = spw.data(); + std::vector tmpW; + sparsify(w, n, 1, tmpW); + + if (dtype == nvinfer1::DataType::kFLOAT) + { + transpose2DWeights(spw.data(), tmpW.data(), n, k); + } + else if (dtype == nvinfer1::DataType::kHALF) + { + transpose2DWeights(spw.data(), tmpW.data(), n, k); + } + } + else + { + sparsify(w, n, 1, spw); + } + + w.values = spw.data(); + layer->setWeights(w); + }; + for (auto& l : constantLayerToSparse) + { + sparsifyConstantWeights(l.first, l.second); + } +} + +template +void setSparseWeights(L& l, int32_t k, int32_t trs, std::vector& sparseWeights) +{ + auto weights = l.getKernelWeights(); + sparsify(weights, k, trs, sparseWeights); + weights.values = sparseWeights.data(); + l.setKernelWeights(weights); +} + +// Explicit instantiation +template void setSparseWeights( + IConvolutionLayer& l, int32_t k, int32_t trs, std::vector& sparseWeights); + +//! \brief Sparsify conv weights fed via Q/DQ chains (companion to sparsifyMatMulKernelWeights). +//! +//! Strongly-typed Q/DQ networks attach the conv weight as a tensor input rather than +//! static kernelWeights. Walks the chain forward from each FP Constant: +//! Constant -> Shuffle* -> Q? -> Shuffle* -> DQ -> Shuffle* -> Conv.input(1) +//! If the chain terminates at a Conv weight input, sparsify the constant in place. +// NOLINTNEXTLINE(readability-function-cognitive-complexity) +void sparsifyQDQConvKernelWeights( + nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) +{ + TensorToLayer convWeightI2L; + TensorToLayer constO2L; + TensorToTensor dqI2O; + TensorToTensor qI2O; + TensorToTensor shuffleI2O; + auto collectMappingInfo = [&](ILayer& l) { + switch (l.getType()) + { + case nvinfer1::LayerType::kCONVOLUTION: + // Conv with weights as a tensor input (vs. static kernelWeights). + if (l.getNbInputs() >= 2 && l.getInput(1) != nullptr) + { + convWeightI2L.try_emplace(l.getInput(1), &l); + } + break; + case nvinfer1::LayerType::kCONSTANT: + { + DataType const dtype = static_cast(l).getWeights().type; + auto const floatDTypes = {nvinfer1::DataType::kFLOAT, nvinfer1::DataType::kHALF, nvinfer1::DataType::kBF16}; + if (std::any_of(floatDTypes.begin(), floatDTypes.end(), [dtype](auto t) { return t == dtype; })) + { + constO2L.try_emplace(l.getOutput(0), &l); + } + break; + } + case nvinfer1::LayerType::kDEQUANTIZE: dqI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; + case nvinfer1::LayerType::kQUANTIZE: qI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; + case nvinfer1::LayerType::kSHUFFLE: shuffleI2O.try_emplace(l.getInput(0), l.getOutput(0)); break; + default: break; + } + }; + int32_t const nbLayers = network.getNbLayers(); + for (int32_t i = 0; i < nbLayers; ++i) + { + collectMappingInfo(*network.getLayer(i)); + } + if (convWeightI2L.size() == 0 || constO2L.size() == 0 || dqI2O.size() == 0) + { + return; + } + + //! Skip past any Shuffle layers consuming t and return the tensor at the chain's end. + //! Returns t unchanged if no Shuffle reads it. + auto walkShuffleChain = [&](nvinfer1::ITensor* t) -> ITensor* { + while (true) + { + auto const it = shuffleI2O.find(t); + if (it == shuffleI2O.end()) + { + break; + } + t = it->second; + } + return t; + }; + + //! Follow Constant -> Shuffle* -> Q? -> Shuffle* -> DQ -> Shuffle* -> Conv.input(1) chain. + //! Returns the terminating IConvolutionLayer*, or nullptr if the chain breaks. + auto walkShuffleQDQChain = [&](nvinfer1::ITensor* t) -> IConvolutionLayer* { + t = walkShuffleChain(t); + if (auto const qI2OIt = qI2O.find(t); qI2OIt != qI2O.end()) + { + t = walkShuffleChain(qI2OIt->second); + } + auto const dqI2OIt = dqI2O.find(t); + if (dqI2OIt == dqI2O.end()) + { + return nullptr; + } + t = walkShuffleChain(dqI2OIt->second); + auto const convWeightI2LIt = convWeightI2L.find(t); + if (convWeightI2LIt == convWeightI2L.end()) + { + return nullptr; + } + ASSERT(convWeightI2LIt->second->getType() == nvinfer1::LayerType::kCONVOLUTION); + return static_cast(convWeightI2LIt->second); + }; + + for (auto& o2l : constO2L) + { + IConvolutionLayer* const conv = walkShuffleQDQChain(o2l.first); + if (conv == nullptr) + { + continue; + } + ASSERT(o2l.second->getType() == nvinfer1::LayerType::kCONSTANT); + IConstantLayer* constLayer = static_cast(o2l.second); + Weights w = constLayer->getWeights(); + if (w.count == 0) + { + continue; + } + Dims const kernelDims = conv->getKernelSizeNd(); + int32_t const k = conv->getNbOutputMaps(); + int64_t const trs = samplesCommon::volume(kernelDims); + // sparsify() reconstructs c (input channels) via c = count / (k*trs); fail loudly if + // the constant's element count doesn't match the KCRS layout this routine assumes. + ASSERT(k > 0 && 0 < trs && trs <= std::numeric_limits::max() + && w.count % (static_cast(k) * trs) == 0); + sparseWeights.emplace_back(); + sparsify(w, k, static_cast(trs), sparseWeights.back()); + w.values = sparseWeights.back().data(); + constLayer->setWeights(w); + } +} + +void sparsify(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights) +{ + for (int32_t l = 0; l < network.getNbLayers(); ++l) + { + auto* layer = network.getLayer(l); + auto const t = layer->getType(); + if (t == nvinfer1::LayerType::kCONVOLUTION) + { + auto& conv = *static_cast(layer); + auto const& dims = conv.getKernelSizeNd(); + ASSERT(dims.nbDims == 2 || dims.nbDims == 3); + auto const k = conv.getNbOutputMaps(); + auto const trs = std::accumulate(dims.d, dims.d + dims.nbDims, 1, std::multiplies()); + sparseWeights.emplace_back(); + setSparseWeights(conv, k, trs, sparseWeights.back()); + } + } + + sparsifyMatMulKernelWeights(network, sparseWeights); + sparsifyQDQConvKernelWeights(network, sparseWeights); + sample::gLogVerbose << "--sparsity=force pruned " << sparseWeights.size() << " weights to be sparsity pattern." + << std::endl; + sample::gLogVerbose << "--sparsity=force has been deprecated. Please use to rewrite the " + "weights to a sparsity pattern and then run with --sparsity=enable" + << std::endl; +} + +void sparsify(Weights const& weights, int32_t k, int32_t trs, std::vector& sparseWeights) +{ + switch (weights.type) + { + case DataType::kFLOAT: + sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); + break; + case DataType::kHALF: + sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); + break; + case DataType::kBF16: + sparsify(static_cast(weights.values), weights.count, k, trs, sparseWeights); + break; + case DataType::kINT8: + case DataType::kINT32: + case DataType::kUINT8: + case DataType::kBOOL: + case DataType::kINT4: + case DataType::kFP8: + case DataType::kINT64: + case DataType::kFP4: ASSERT(false && "Unsupported data type"); + case DataType::kE8M0: ASSERT(false && "E8M0 is not supported"); + } +} + +template +void print(std::ostream& os, T v) +{ + os << v; +} + +void print(std::ostream& os, int8_t v) +{ + os << static_cast(v); +} + +void print(std::ostream& os, uint8_t v) +{ + os << static_cast(v); +} + +void print(std::ostream& os, __half v) +{ + os << static_cast(v); +} + +#if CUDA_VERSION >= 11060 +void print(std::ostream& os, __nv_fp8_e4m3 v) +{ + os << static_cast(v); +} +#endif + +int32_t dataOffsetFromDims(int64_t v, Dims const& dims, Dims const& strides, int32_t vectorDim, int32_t spv) +{ + int32_t dataOffset = 0; + for (int32_t dimIndex = dims.nbDims - 1; dimIndex >= 0; --dimIndex) + { + int32_t dimVal = v % dims.d[dimIndex]; + if (dimIndex == vectorDim) + { + dataOffset += (dimVal / spv) * strides.d[dimIndex] * spv + dimVal % spv; + } + else + { + dataOffset += dimVal * strides.d[dimIndex] * (vectorDim == -1 ? 1 : spv); + } + v /= dims.d[dimIndex]; + ASSERT(v >= 0); + } + + return dataOffset; +} + +template +void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv) +{ + auto const vol = volume(dims); + T const* typedBuffer = static_cast(buffer); + for (int64_t v = 0; v < vol; ++v) + { + int32_t dataOffset = dataOffsetFromDims(v, dims, strides, vectorDim, spv); + if (v > 0) + { + os << separator; + } + print(os, typedBuffer[dataOffset]); + } +} + +void dumpInt4Buffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv) +{ + auto const vol = volume(dims); + uint8_t const* typedBuffer = static_cast(buffer); + for (int64_t v = 0; v < vol; ++v) + { + int32_t dataOffset = dataOffsetFromDims(v, dims, strides, vectorDim, spv); + if (v > 0) + { + os << separator; + } + + auto value = typedBuffer[dataOffset / 2]; + if (dataOffset % 2 == 0) + { + // Cast to int8_t before right shift, so right-shift will sign-extend. + // Left shift on int8_t can be undefined behaviour, must perform left shift on uint8_t. + os << (static_cast(value << 4) >> 4); + } + else + { + os << (static_cast(value) >> 4); + } + } +} + +// Explicit instantiation +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer<__half>(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +#if CUDA_VERSION >= 11060 +template void dumpBuffer<__nv_fp8_e4m3>(void const* buffer, std::string const& separator, std::ostream& os, + Dims const& dims, Dims const& strides, int32_t vectorDim, int32_t spv); +#endif +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); +template void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); + +template +void sparsify(T const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights) +{ + auto const c = count / (k * trs); + sparseWeights.resize(count * sizeof(T)); + auto* sparseValues = reinterpret_cast(sparseWeights.data()); + + constexpr int32_t window = 4; + constexpr int32_t nonzeros = 2; + + int32_t const crs = c * trs; + auto const getIndex = [=](int32_t ki, int32_t ci, int32_t rsi) { return ki * crs + ci * trs + rsi; }; + + for (int64_t ki = 0; ki < k; ++ki) + { + for (int64_t rsi = 0; rsi < trs; ++rsi) + { + int32_t w = 0; + int32_t nz = 0; + for (int64_t ci = 0; ci < c; ++ci) + { + auto const index = getIndex(ki, ci, rsi); + if (nz < nonzeros) + { + sparseValues[index] = values[index]; + ++nz; + } + else + { + sparseValues[index] = 0; + } + if (++w == window) + { + w = 0; + nz = 0; + } + } + } + } +} + +// Explicit instantiation +template void sparsify( + float const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights); +template void sparsify( + half_float::half const* values, int64_t count, int32_t k, int32_t trs, std::vector& sparseWeights); + +template +void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n) +{ + ASSERT(dst != src); + T* tdst = reinterpret_cast(dst); + T const* tsrc = reinterpret_cast(src); + for (int32_t mi = 0; mi < m; ++mi) + { + for (int32_t ni = 0; ni < n; ++ni) + { + int32_t const isrc = mi * n + ni; + int32_t const idst = ni * m + mi; + tdst[idst] = tsrc[isrc]; + } + } +} + +// Explicit instantiation +template void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); +template void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); + +template ::value, bool>::type> +void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max) +{ + T* typedBuffer = static_cast(buffer); + std::default_random_engine engine; + std::uniform_int_distribution distribution(min, max); + auto generator = [&engine, &distribution]() { return static_cast(distribution(engine)); }; + std::generate(typedBuffer, typedBuffer + volume, generator); +} + +template ::value, bool>::type> +void fillBuffer(void* buffer, int64_t volume, float min, float max) +{ + T* typedBuffer = static_cast(buffer); + std::default_random_engine engine; + std::uniform_real_distribution distribution(min, max); + auto generator = [&engine, &distribution]() { return static_cast(distribution(engine)); }; + std::generate(typedBuffer, typedBuffer + volume, generator); +} + +// Explicit instantiation +template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); +template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); +template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); +template void fillBuffer(void* buffer, int64_t volume, float min, float max); +template void fillBuffer<__half>(void* buffer, int64_t volume, float min, float max); +template void fillBuffer(void* buffer, int64_t volume, float min, float max); +#if CUDA_VERSION >= 11060 +template void fillBuffer<__nv_fp8_e4m3>(void* buffer, int64_t volume, float min, float max); +#endif +template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); +template void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); + +bool matchStringWithOneWildcard(std::string const& pattern, std::string const& target) +{ + auto const splitPattern = splitToStringVec(pattern, '*', 1); + + // If there is no wildcard, return if the two strings match exactly. + if (splitPattern.size() == 1) + { + return pattern == target; + } + + // Otherwise, target must follow prefix+anything+postfix pattern. + return target.size() >= (splitPattern[0].size() + splitPattern[1].size()) && target.find(splitPattern[0]) == 0 + && target.rfind(splitPattern[1]) == (target.size() - splitPattern[1].size()); +} + +//! @brief Sanitizes the remote target config string by removing sensitive credentials +//! +//! This function removes usernames and passwords from URL-style configuration strings +//! to prevent sensitive authentication information from appearing in logs or debug output. +//! The credentials section (username:password) is replaced with "***" for security. +//! +//! Config format: protocol://username[:password]@hostname[:port]?param1=value1¶m2=value2 +//! Supported protocols: ssh, http, https, etc. A config that omits the protocol is still redacted, so a +//! malformed value cannot smuggle credentials into the log ahead of validateRemoteConfig rejecting it. +//! +//! Examples: +//! Input: "ssh://admin:secretpass@server.com:22?timeout=30" +//! Output: "ssh://***@server.com:22?timeout=30" +//! +//! @param config The configuration string to sanitize +//! @return Sanitized configuration string with passwords and usernames replaced by *** +std::string sanitizeRemoteConfig(std::string const& config) +{ + if (config.empty()) + { + return config; + } + + try + { + size_t const protocolEnd = config.find("://"); + size_t const credentialsStart = (protocolEnd == std::string::npos) ? 0U : protocolEnd + 3U; + if (credentialsStart >= config.length()) + { + return config; // Truncated after protocol + } + + // A password may itself contain '@', so redact through the last one rather than the first. Bound + // the search to the authority: a later '@' in a path or query belongs to neither credentials nor host. + size_t const authorityEnd = config.find_first_of("/?#", credentialsStart); + size_t const searchEnd = (authorityEnd == std::string::npos) ? config.length() : authorityEnd; + if (searchEnd <= credentialsStart) + { + return config; // Empty authority, nothing to redact + } + + size_t const credentialsEnd = config.rfind('@', searchEnd - 1U); + if (credentialsEnd == std::string::npos || credentialsEnd <= credentialsStart) + { + return config; // No credentials, return as is + } + + return config.substr(0, credentialsStart) + "***" + config.substr(credentialsEnd); + } + catch (std::exception const& e) + { + sample::gLogError << "Exception in sanitizeRemoteConfig: " << e.what() << std::endl; + return config; // Return original on error + } + catch (...) + { + sample::gLogError << "Unknown exception in sanitizeRemoteConfig" << std::endl; + return config; // Return original on error + } +} + +bool validateNonEmpty(std::string const& value, std::string const& flagName) +{ + if (value.empty()) + { + sample::gLogError << flagName << " cannot be empty" << std::endl; + return false; + } + return true; +} + +bool validateRemoteConfig(std::string const& config) +{ + if (config.find("://") == std::string::npos) + { + sample::gLogError << "Invalid remote target config format. Expected format: " + "protocol://username[:password]@hostname[:port]?param1=value1¶m2=value2" + << std::endl; + return false; + } + return true; +} + +std::vector sanitizeArgv(int32_t argc, char** argv) +{ + // --remoteAutoTuningConfig is an alias of --remoteConfig; both carry credentials. + static constexpr std::array kREMOTE_CONFIG_FLAGS{ + "--remoteConfig=", "--remoteAutoTuningConfig="}; + + std::vector sanitizedArgs; + sanitizedArgs.reserve(argc); + + for (int32_t i = 0; i < argc; ++i) + { + std::string arg = argv[i]; + + for (auto const flag : kREMOTE_CONFIG_FLAGS) + { + if (arg.size() > flag.size() && arg.compare(0, flag.size(), flag) == 0) + { + arg = std::string(flag) + sanitizeRemoteConfig(arg.substr(flag.size())); + break; + } + } + + sanitizedArgs.push_back(arg); + } + + return sanitizedArgs; +} + +// ============================================================================ +// Accuracy Validator Implementations +// ============================================================================ + +template +double L0AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) +{ + // Uses PyTorch/NumPy allclose formula: |a - b| <= atol + rtol * |b| + // See: https://docs.pytorch.org/docs/stable/generated/torch.allclose.html + // and tools/common/accuracyComparison.h. + ASSERT(actual.size() == reference.size()); + ASSERT(actual.size() != 0); + int64_t mismatchCount = 0; + for (uint64_t i = 0; i < actual.size(); ++i) + { + double const absDiff = std::abs(static_cast(actual[i]) - static_cast(reference[i])); + double const refAbs = std::abs(static_cast(reference[i])); + double const tolerance = mAtol + mRtol * refAbs; + if (absDiff > tolerance) + { + mismatchCount++; + } + } + return static_cast(mismatchCount) / actual.size(); +} + +template +double L1AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) +{ + ASSERT(actual.size() == reference.size()); + ASSERT(actual.size() != 0); + double sum = 0.0; + for (uint64_t i = 0; i < actual.size(); ++i) + { + sum += std::abs(static_cast(actual[i]) - static_cast(reference[i])); + } + return sum / actual.size(); +} + +template +double L2AccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) +{ + ASSERT(actual.size() == reference.size()); + ASSERT(actual.size() != 0); + double sum = 0.0; + for (uint64_t i = 0; i < actual.size(); ++i) + { + double diff = static_cast(actual[i]) - static_cast(reference[i]); + sum += diff * diff; + } + return sum / actual.size(); +} + +template +double LInfAccuracyValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) +{ + ASSERT(actual.size() == reference.size()); + ASSERT(actual.size() != 0); + double maxDiff = 0.0; + for (uint64_t i = 0; i < actual.size(); ++i) + { + double diff = std::abs(static_cast(actual[i]) - static_cast(reference[i])); + maxDiff = std::max(maxDiff, diff); + } + return maxDiff; +} + +template +double CosineSimilarityValidator::calculateAccuracy(std::vector const& actual, std::vector const& reference) +{ + ASSERT(actual.size() == reference.size()); + ASSERT(actual.size() != 0); + double dotProduct = 0.0; + double normActual = 0.0; + double normRef = 0.0; + for (uint64_t i = 0; i < actual.size(); ++i) + { + double a = static_cast(actual[i]); + double r = static_cast(reference[i]); + dotProduct += a * r; + normActual += a * a; + normRef += r * r; + } + double denominator = std::sqrt(normActual) * std::sqrt(normRef); + if (denominator < 1e-12) + { + return 1.0; // Handle zero vectors + } + double cosineSim = dotProduct / denominator; + return 1.0 - cosineSim; // Return as cost (0 = perfect match) +} + +// Explicit template instantiations for supported types +template class L0AccuracyValidator; +template class L0AccuracyValidator; +template class L0AccuracyValidator; +template class L0AccuracyValidator; + +template class L1AccuracyValidator; +template class L1AccuracyValidator; +template class L1AccuracyValidator; +template class L1AccuracyValidator; + +template class L2AccuracyValidator; +template class L2AccuracyValidator; +template class L2AccuracyValidator; +template class L2AccuracyValidator; + +template class LInfAccuracyValidator; +template class LInfAccuracyValidator; +template class LInfAccuracyValidator; +template class LInfAccuracyValidator; + +template class CosineSimilarityValidator; +template class CosineSimilarityValidator; +template class CosineSimilarityValidator; +template class CosineSimilarityValidator; + +bool peekArg(int32_t argc, char** argv, char const* flag) +{ + auto const flagLen = std::strlen(flag); + for (int32_t i = 1; i < argc; ++i) + { + if (argv[i] == nullptr) + { + continue; + } + // Match either bare flag (--continue) or flag=value (--tuneBuildRoutes=...). + if (std::strncmp(argv[i], flag, flagLen) == 0 && (argv[i][flagLen] == '\0' || argv[i][flagLen] == '=')) + { + return true; + } + } + return false; +} + +std::string buildShellQuotedCmdLine(int32_t argc, char** argv) +{ + std::string cmdLine; + for (int32_t i = 0; i < argc; ++i) + { + if (i > 0) + { + cmdLine += " "; + } + std::string arg = argv[i]; + bool const needsQuoting = arg.find_first_of(" \t|[]{}()&;'\"\\") != std::string::npos; + if (needsQuoting) + { + std::string escaped; + for (char c : arg) + { + if (c == '\'') + { + escaped += "'\\''"; + } + else + { + escaped += c; + } + } + cmdLine += "'" + escaped + "'"; + } + else + { + cmdLine += arg; + } + } + return cmdLine; +} + +//! \brief Resolve file paths in argv to absolute for cache storage. +//! +//! File-path flags that get resolved: --onnx=, --saveEngine=, --loadInputs=, +//! --loadRefOutputs=, --tuneBuildRouteFile=, --loadEngine=. All others are stored as-is. +//! --loadInputs and --loadRefOutputs have format "name:path,name:path" so each +//! path component is resolved separately. +namespace +{ +// NOLINTNEXTLINE(readability-function-cognitive-complexity) +std::vector resolveArgvPaths(int32_t argc, char** argv) +{ + static std::vector const kSIMPLE_PATH_FLAGS + = {"--onnx=", "--saveEngine=", "--tuneBuildRouteFile=", "--loadEngine=", "--loadCheckerBlob="}; + static std::vector const kMAPPED_PATH_FLAGS = {"--loadInputs=", "--loadRefOutputs="}; + + std::vector result; + for (int32_t i = 0; i < argc; ++i) + { + std::string arg(argv[i]); + + // Check simple path flags (--flag=path -> --flag=) + bool resolved = false; + for (auto const& prefix : kSIMPLE_PATH_FLAGS) + { + if (arg.starts_with(prefix)) + { + result.push_back(prefix + resolveAbsolutePath(arg.substr(prefix.size()))); + resolved = true; + break; + } + } + if (resolved) + { + continue; + } + + // Check mapped path flags (--flag=name:path,name:path -> resolve each path) + for (auto const& prefix : kMAPPED_PATH_FLAGS) + { + if (arg.starts_with(prefix)) + { + // Split on ',' to get individual name:path pairs + auto pairs = splitToStringVec(arg.substr(prefix.size()), ','); + std::string resolvedValue; + for (uint64_t p = 0; p < pairs.size(); ++p) + { + if (p > 0) + { + resolvedValue += ","; + } + // Split each pair on ':' to separate name from path + auto nameAndPath = splitToStringVec(pairs[p], ':', 1); + if (nameAndPath.size() == 2) + { + resolvedValue += nameAndPath[0] + ":" + resolveAbsolutePath(nameAndPath[1]); + } + else + { + resolvedValue += pairs[p]; // Malformed pair, keep as-is + } + } + result.push_back(prefix + resolvedValue); + resolved = true; + break; + } + } + if (resolved) + { + continue; + } + + result.push_back(arg); + } + return result; +} +} // anonymous namespace + +void writeTuningCacheHeader(std::string const& cacheFilePath, AllOptions const& options, int32_t argc, char** argv, + std::string const& tunerVersion, std::string const& defaultBuildRoute) +{ + // Use ordered_json to preserve insertion order matching best_config.json.example: + // tuner_version, accuracy_algorithm, accuracy_parameter, searching_algorithm, + // command_line, default_build_route, tuning_expr, files, argv + nlohmann::ordered_json header; + + header["tuner_version"] = tunerVersion; + header["accuracy_algorithm"] = getAlgorithmName(options.inference.accuracyValidationAlgorithm); + + nlohmann::ordered_json accParam; + accParam["atol"] = options.inference.atol; + accParam["rtol"] = options.inference.rtol; + accParam["epsilon"] = options.inference.accuracyThresholdEndToEnd; + header["accuracy_parameter"] = accParam; + + header["searching_algorithm"] = toString(options.tuning.tuningSearchAlgorithm); + + // Reconstruct command line for reference, with shell-safe quoting for arguments + // that contain spaces or metacharacters (e.g. --tuneBuildRoutes values). + std::string cmdLine = buildShellQuotedCmdLine(argc, argv); + header["command_line"] = cmdLine; + header["default_build_route"] = defaultBuildRoute; + + // Store the expanded tuning expression. This is the already-expanded string + // (handles --tuneBuildRouteFile case where the file may not exist at resume time). + header["tuning_expr"] = options.tuning.tuningExpr; + + // Store absolute paths to all file-based options for human readability and + // as a cross-check. The authoritative source for --continue reconstruction + // is the "argv" field below. + { + nlohmann::ordered_json files; + if (!options.model.baseModel.model.empty()) + { + files["onnx"] = resolveAbsolutePath(options.model.baseModel.model); + } + if (!options.build.engine.empty()) + { + files["save_engine"] = resolveAbsolutePath(options.build.engine); + } + // Input files: map of tensor_name → absolute path + if (!options.inference.refPairs.empty()) + { + nlohmann::ordered_json inputs; + for (auto const& [name, path] : options.inference.refPairs[0].first) + { + inputs[name] = resolveAbsolutePath(path); + } + if (!inputs.empty()) + { + files["inputs"] = inputs; + } + + nlohmann::ordered_json refOutputs; + for (auto const& [name, path] : options.inference.refPairs[0].second) + { + refOutputs[name] = resolveAbsolutePath(path); + } + if (!refOutputs.empty()) + { + files["ref_outputs"] = refOutputs; + } + } + header["files"] = files; + } + + // Store argv with file-path arguments resolved to absolute paths. + // This is the machine-readable source of truth for --continue reconstruction. + // When resuming, the stored argv is replayed to reconstruct all options + // (--iterations, --duration, --fp16, etc.) without enumerating each one. + { + auto resolvedArgv = resolveArgvPaths(argc, argv); + nlohmann::ordered_json argvArray(resolvedArgv); + header["argv"] = argvArray; + } + + std::ofstream file(cacheFilePath, std::ios::trunc); + if (!file) + { + sample::gLogError << "Cannot open tuning cache file for writing header: " << cacheFilePath << std::endl; + return; + } + file << header.dump() << std::endl; +} + +void writeTuningCacheIteration(std::string const& cacheFilePath, uint64_t iter, std::string const& buildRoute, + bool crashed, std::string const& errorMessage, std::unordered_map const& accuracyLossValues, + double gpuTimeMs) +{ + // Use ordered_json to preserve insertion order matching best_config.json.example: + // iter, build_route, crash, error_message, accuracy_loss, gpu_time + nlohmann::ordered_json result; + result[tuningCache::kIter] = iter; + result[tuningCache::kBuildRoute] = buildRoute; + result[tuningCache::kCrash] = crashed; + result[tuningCache::kErrorMessage] = errorMessage; + + // accuracy_loss is a per-output map: {"output_name": accuracy_value, ...} + // When crashed, accuracy values are unavailable so we write null. + if (crashed || accuracyLossValues.empty()) + { + result[tuningCache::kAccuracyLoss] = nullptr; + } + else + { + nlohmann::ordered_json accMap; + for (auto const& [name, value] : accuracyLossValues) + { + accMap[name] = value; + } + result[tuningCache::kAccuracyLoss] = accMap; + } + result[tuningCache::kGpuTime] = crashed ? nlohmann::ordered_json(nullptr) : nlohmann::ordered_json(gpuTimeMs); + + std::ofstream file(cacheFilePath, std::ios::app); + if (!file) + { + sample::gLogError << "Cannot open tuning cache file to append iteration " << iter << ": " << cacheFilePath + << std::endl; + return; + } + file << result.dump() << std::endl; +} + +std::vector reconstructArgvFromCacheHeader( + TuningCacheHeader const& header, std::string const& currentExePath, std::string const& cacheFilePath) +{ + std::vector newArgv; + + // Use current executable path as argv[0], not the one stored in the cache + // (the binary may have been rebuilt or moved since the original run). + newArgv.push_back(currentExePath); + + // Iterate over stored argv (skip stored argv[0]). + for (uint64_t i = 1; i < header.argv.size(); ++i) + { + std::string const& arg = header.argv[i]; + + // Replace --tuneBuildRoutes or --tuneBuildRouteFile with the stored tuning_expr. + // This handles the case where --tuneBuildRouteFile was used originally but the + // file no longer exists — the expanded expression is stored in tuning_expr. + if (arg.starts_with("--tuneBuildRoutes=") || arg.starts_with("--tuneBuildRouteFile=")) + { + continue; // Will be re-added below with the stored tuning_expr. + } + + // Remove --continue and --tuningCacheFile from the stored argv to avoid + // recursion (the stored run may itself have been a --continue run). + if (arg == "--continue"sv || arg.starts_with("--tuningCacheFile=")) + { + continue; + } + + newArgv.push_back(arg); + } + + // Add back the tuning expression and cache file path. + newArgv.push_back("--tuneBuildRoutes=" + header.tuningExpr); + newArgv.push_back("--tuningCacheFile=" + cacheFilePath); + + return newArgv; +} + +std::string resolveAbsolutePath(std::string const& path) +{ + if (path.empty()) + { + return path; + } +#if defined(_WIN32) + // On Windows, path resolution is not needed (tuning features are not supported on Windows). + // Return the path unchanged so the code compiles. + return path; +#else + // POSIX realpath() resolves symlinks and relative components to an absolute path. + // Returns nullptr if the file does not exist or another error occurs. + char resolved[PATH_MAX]; + if (realpath(path.c_str(), resolved) != nullptr) + { + return std::string(resolved); + } + return path; +#endif +} + +std::optional readTuningCacheHeader(std::string const& cacheFilePath) +{ + std::ifstream file(cacheFilePath); + if (!file.is_open()) + { + return std::nullopt; + } + + // First line is the JSON header. + std::string headerLine; + if (!std::getline(file, headerLine) || headerLine.empty()) + { + return std::nullopt; + } + + try + { + auto headerJson = nlohmann::json::parse(headerLine); + + TuningCacheHeader header; + + // Extract argv array → vector + if (headerJson.contains("argv") && headerJson["argv"].is_array()) + { + for (auto const& elem : headerJson["argv"]) + { + header.argv.push_back(elem.get()); + } + } + else + { + // argv field is required for --continue reconstruction. + sample::gLogError << "Tuning cache header missing 'argv' field" << std::endl; + return std::nullopt; + } + + // Extract tuning_expr string. + if (headerJson.contains("tuning_expr") && headerJson["tuning_expr"].is_string()) + { + header.tuningExpr = headerJson["tuning_expr"].get(); + } + else + { + sample::gLogError << "Tuning cache header missing 'tuning_expr' field" << std::endl; + return std::nullopt; + } + + // Count remaining non-empty lines as completed iterations. + header.completedIterations = 0; + std::string line; + while (std::getline(file, line)) + { + if (!line.empty()) + { + ++header.completedIterations; + } + } + + return header; + } + catch (nlohmann::json::exception const& e) + { + sample::gLogError << "Failed to parse tuning cache header: " << e.what() << std::endl; + return std::nullopt; + } +} + +std::vector readCachedIterationResults(std::string const& cacheFilePath, int64_t maxIterations) +{ + std::vector results; + std::ifstream file(cacheFilePath); + if (!file.is_open()) + { + return results; + } + + std::string line; + // Skip header line. + if (!std::getline(file, line)) + { + return results; + } + + // Read iteration lines, extracting crash and gpu_time fields. + while (std::getline(file, line) && static_cast(results.size()) < maxIterations) + { + if (line.empty()) + { + continue; + } + try + { + auto j = nlohmann::json::parse(line); + CachedIterationResult r; + r.crashed = j.value(tuningCache::kCrash, true); + r.gpuTimeMs = j.contains(tuningCache::kGpuTime) && j[tuningCache::kGpuTime].is_number() + ? j[tuningCache::kGpuTime].get() + : 0.0; + results.push_back(r); + } + catch (nlohmann::json::exception const&) + { + // Malformed line — treat as crashed. + results.push_back({true, 0.0}); + } + } + + return results; +} + +} // namespace sample diff --git a/samples/trtexecCommon/sampleUtils.h b/samples/trtexecCommon/sampleUtils.h new file mode 100644 index 0000000000..90ccf82e04 --- /dev/null +++ b/samples/trtexecCommon/sampleUtils.h @@ -0,0 +1,340 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef TRT_SAMPLE_UTILS_H +#define TRT_SAMPLE_UTILS_H + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +#include "NvInfer.h" + +#include "common.h" +#include "logger.h" +#include "logging.h" +#include "sampleOptions.h" + +#define SMP_RETVAL_IF_FALSE(condition, msg, retval, err) \ + { \ + if ((condition) == false) \ + { \ + (err) << (msg) << std::endl; \ + return retval; \ + } \ + } + +namespace sample +{ + +template +inline T roundUp(T m, T n) +{ + return ((m + n - 1) / n) * n; +} + +//! comps is the number of components in a vector. Ignored if vecDim < 0. +int64_t volume(nvinfer1::Dims const& dims, nvinfer1::Dims const& strides, int32_t vecDim, int32_t comps, int32_t batch); + +using samplesCommon::volume; + +nvinfer1::Dims toDims(std::vector const& vec); + +template ::value, bool>::type = true> +void fillBuffer(void* buffer, int64_t volume, int32_t min, int32_t max); + +template ::value, bool>::type = true> +void fillBuffer(void* buffer, int64_t volume, float min, float max); + +template +void dumpBuffer(void const* buffer, std::string const& separator, std::ostream& os, nvinfer1::Dims const& dims, + nvinfer1::Dims const& strides, int32_t vectorDim, int32_t spv); + +void dumpInt4Buffer(void const* buffer, std::string const& separator, std::ostream& os, Dims const& dims, + Dims const& strides, int32_t vectorDim, int32_t spv); + +void loadFromFile(std::string const& fileName, char* dst, size_t size); + +std::vector splitToStringVec(std::string const& option, char separator, int64_t maxSplit = -1); + +bool broadcastIOFormats(std::vector const& formats, size_t nbBindings, bool isInput = true); + +#if !TRT_WINML +int32_t getCudaDriverVersion(); + +int32_t getCudaRuntimeVersion(); +#endif + +void sparsify(nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights); +void sparsify(nvinfer1::Weights const& weights, int32_t k, int32_t rs, std::vector& sparseWeights); + +// Walk the weights elements and overwrite (at most) 2 out of 4 elements to 0. +template +void sparsify(T const* values, int64_t count, int32_t k, int32_t rs, std::vector& sparseWeights); + +template +void setSparseWeights(L& l, int32_t k, int32_t rs, std::vector& sparseWeights); + +// Sparsify the weights of Constant layers that are fed to MatMul via Shuffle layers. +// Forward analysis on the API graph to determine which weights to sparsify. +void sparsifyMatMulKernelWeights( + nvinfer1::INetworkDefinition& network, std::vector>& sparseWeights); + +template +void transpose2DWeights(void* dst, void const* src, int32_t const m, int32_t const n); + +//! A helper function to match a target string with a pattern where the pattern can contain up to one wildcard ('*') +//! character that matches to any strings. +bool matchStringWithOneWildcard(std::string const& pattern, std::string const& target); + +//! A helper method to find an item from an unordered_map. If the exact match exists, this is identical to +//! map.find(target). If the exact match does not exist, it returns the first plausible match, taking up to one wildcard +//! into account. If there is no plausible match, then it returns map.end(). +template +typename std::unordered_map::const_iterator findPlausible( + std::unordered_map const& map, std::string const& target) +{ + auto res = map.find(target); + if (res == map.end()) + { + res = std::find_if( + map.begin(), map.end(), [&](typename std::unordered_map::value_type const& item) { + return matchStringWithOneWildcard(item.first, target); + }); + } + return res; +} + +// ==== Common argument parsing utilities ==== + +//! Validate that a value is not empty, log error if it is +bool validateNonEmpty(std::string const& value, std::string const& flagName); + +//! Validate remote target config format +bool validateRemoteConfig(std::string const& config); + +//! Ensure directory path ends with '/' +inline std::string normalizeDirectoryPath(std::string const& dirPath) +{ + std::string result = dirPath; + if (!result.empty() && result.back() != '/') + { + result.push_back('/'); + } + return result; +} + +//! Sanitizes the remote target config string by removing sensitive credentials +//! Removes usernames and passwords from URL-style config strings for security. +//! Example: "ssh://user:pass@host:22" becomes "ssh://***:***@host:22" +std::string sanitizeRemoteConfig(std::string const& config); + +//! Sanitizes command line arguments for logging, removing sensitive credentials +//! Processes argv array and sanitizes sensitive arguments like remoteConfig +//! @param argc Number of arguments +//! @param argv Array of argument strings +//! @return Vector of sanitized argument strings +std::vector sanitizeArgv(int32_t argc, char** argv); + +//! Interface for accuracy validation +//! This interface provides a way to calculate the accuracy gap between the actual and reference outputs. +//! Since all the return value is a "loss value", the lower the return value, the better accuracy it is. +template +class IAccuracyValidator +{ +public: + virtual ~IAccuracyValidator() = default; + virtual double calculateAccuracy(std::vector const& actual, std::vector const& reference) = 0; +}; + +//! L0 accuracy validator calculates element-wise accuracy using the PyTorch/NumPy allclose formula. +//! An element matches if: |actual[i] - ref[i]| <= atol + rtol * |ref[i]| +//! accuracy = (number of mismatching elements) / N +//! Returns the mismatch ratio (0.0 means perfect match, 1.0 means all elements mismatch). +template +class L0AccuracyValidator : public IAccuracyValidator +{ +public: + L0AccuracyValidator(double atol, double rtol) + : mAtol(atol) + , mRtol(rtol) + { + } + + double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; + +private: + double mAtol; + double mRtol; +}; + +//! L1 accuracy validator calculates mean absolute error. +//! accuracy = Sum(|actual[i] - ref[i]|) / N +//! Returns the mean absolute error (0.0 means perfect match). +template +class L1AccuracyValidator : public IAccuracyValidator +{ +public: + double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; +}; + +//! L2 accuracy validator calculates mean squared error. +//! accuracy = Sum(|actual[i] - ref[i]|^2) / N +//! Returns the mean squared error (0.0 means perfect match). +template +class L2AccuracyValidator : public IAccuracyValidator +{ +public: + double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; +}; + +//! LInf accuracy validator calculates maximum absolute error. +//! accuracy = Max(|actual[i] - ref[i]|) +//! Returns the max absolute error (0.0 means perfect match). +template +class LInfAccuracyValidator : public IAccuracyValidator +{ +public: + double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; +}; + +//! Cosine similarity validator calculates 1 - cosine_similarity. +//! cosine_sim = Sum(actual[i] * ref[i]) / (sqrt(Sum(actual[i]^2)) * sqrt(Sum(ref[i]^2))) +//! accuracy loss = 1 - cosine_sim +//! Returns 1 - cosine_similarity (0.0 means perfect match). +template +class CosineSimilarityValidator : public IAccuracyValidator +{ +public: + double calculateAccuracy(std::vector const& actual, std::vector const& reference) override; +}; + +//! \brief Get human-readable name string for an accuracy validation algorithm. +//! \param[in] algorithm The accuracy validation algorithm enum value. +//! \return Name string (e.g., "L0", "L1", "Cosine"). +inline std::string getAlgorithmName(AccuracyValidationAlgorithm algorithm) +{ + switch (algorithm) + { + case AccuracyValidationAlgorithm::kL0: return "L0"; + case AccuracyValidationAlgorithm::kL1: return "L1"; + case AccuracyValidationAlgorithm::kL2: return "L2"; + case AccuracyValidationAlgorithm::kLInf: return "LInf"; + case AccuracyValidationAlgorithm::kCosineSimilarity: return "Cosine"; + default: return "Unknown"; + } +} + +//! \brief Factory function to create an accuracy validator based on algorithm type. +//! \param[in] algorithm The accuracy validation algorithm to use. +//! \param[in] atol Absolute tolerance (only used by L0 algorithm). +//! \param[in] rtol Relative tolerance (only used by L0 algorithm). +//! \return Unique pointer to the appropriate IAccuracyValidator implementation. +template +std::unique_ptr> createAccuracyValidator( + AccuracyValidationAlgorithm algorithm, float atol = 1e-5F, float rtol = 1e-5F) +{ + switch (algorithm) + { + case AccuracyValidationAlgorithm::kL0: return std::make_unique>(atol, rtol); + case AccuracyValidationAlgorithm::kL1: return std::make_unique>(); + case AccuracyValidationAlgorithm::kL2: return std::make_unique>(); + case AccuracyValidationAlgorithm::kLInf: return std::make_unique>(); + case AccuracyValidationAlgorithm::kCosineSimilarity: return std::make_unique>(); + } + ASSERT(false && "Unknown Accuracy Validation Algorithm"); + return nullptr; +} + +//! \brief Cheap argv pre-scan. Returns true if some `argv[i]` exactly equals `flag` +//! or starts with `flag` + "=". Used by main() before option parsing to dispatch +//! between trtexec single-run mode and the tuning loop. +[[nodiscard]] bool peekArg(int32_t argc, char** argv, char const* flag); + +// ============================================================================ +// Tuning cache I/O (used by --tuneBuildRoutes / --continue). +// Header is a single JSON object on line 1; iterations are JSON Lines after. +// ============================================================================ + +//! \brief Reconstruct a shell-safe command line string from argc/argv. +std::string buildShellQuotedCmdLine(int32_t argc, char** argv); + +//! \brief Resolve a file path to an absolute path using POSIX realpath(). +//! Empty input or realpath() failure returns the input unchanged. +std::string resolveAbsolutePath(std::string const& path); + +//! \brief Write the tuning cache file header (line 1, JSON object). +void writeTuningCacheHeader(std::string const& cacheFilePath, AllOptions const& options, int32_t argc, char** argv, + std::string const& tunerVersion, std::string const& defaultBuildRoute); + +//! \brief Append one iteration line to the cache file. Fields: iter, build_route, crash, +//! error_message, accuracy_loss, gpu_time. Crashed iterations have null accuracy/gpu. +void writeTuningCacheIteration(std::string const& cacheFilePath, uint64_t iter, std::string const& buildRoute, + bool crashed, std::string const& errorMessage, std::unordered_map const& accuracyLossValues, + double gpuTimeMs); + +//! \struct TuningCacheHeader +//! \brief Parsed contents of the cache header, returned by readTuningCacheHeader(). +struct TuningCacheHeader +{ + std::vector argv; //!< Original command line with file paths absolute. + std::string tuningExpr; //!< Expanded --tuneBuildRoutes expression. + int64_t completedIterations{0}; //!< Number of iteration lines after the header. +}; + +//! \brief Read and parse the cache file's header line + count completed iteration lines. +std::optional readTuningCacheHeader(std::string const& cacheFilePath); + +//! \brief Rebuild argv for a --continue resume. argv[0] is replaced with currentExePath; +//! --tuneBuildRoutes is set to the cached expanded expression; --continue and +//! --tuningCacheFile are stripped from the stored argv and the cache path is re-appended. +std::vector reconstructArgvFromCacheHeader( + TuningCacheHeader const& header, std::string const& currentExePath, std::string const& cacheFilePath); + +//! \struct CachedIterationResult +//! \brief Minimal per-iteration fields from the cache, used to reconstruct mixed-mode positive knobs. +struct CachedIterationResult +{ + bool crashed{true}; + double gpuTimeMs{0.0}; +}; + +//! \brief Read up to maxIterations iteration lines from the cache and extract (crashed, gpu_time). +std::vector readCachedIterationResults(std::string const& cacheFilePath, int64_t maxIterations); + +namespace tuningCache +{ +constexpr char const* kIter = "iter"; +constexpr char const* kBuildRoute = "build_route"; +constexpr char const* kCrash = "crash"; +constexpr char const* kErrorMessage = "error_message"; +constexpr char const* kAccuracyLoss = "accuracy_loss"; +constexpr char const* kGpuTime = "gpu_time"; +} // namespace tuningCache + +} // namespace sample +#endif // TRT_SAMPLE_UTILS_H diff --git a/samples/trtexecCommon/sampleUtils.test.cpp b/samples/trtexecCommon/sampleUtils.test.cpp new file mode 100644 index 0000000000..099d4c34e4 --- /dev/null +++ b/samples/trtexecCommon/sampleUtils.test.cpp @@ -0,0 +1,218 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "sampleUtils.h" +#include "ArgVec.test.h" + +#include + +#include + +using namespace sample; +using namespace std::string_view_literals; + +TEST(RoundUp, ExactMultiple) +{ + EXPECT_EQ(roundUp(4, 4), 4); + EXPECT_EQ(roundUp(8, 4), 8); + EXPECT_EQ(roundUp(0, 4), 0); +} + +TEST(RoundUp, NeedsRounding) +{ + EXPECT_EQ(roundUp(1, 4), 4); + EXPECT_EQ(roundUp(5, 4), 8); + EXPECT_EQ(roundUp(7, 4), 8); +} + +TEST(SplitToStringVec, SingleToken) +{ + auto const v = splitToStringVec("hello", ','); + ASSERT_EQ(v.size(), 1U); + EXPECT_EQ(v[0], "hello"sv); +} + +TEST(SplitToStringVec, MultipleTokens) +{ + auto const v = splitToStringVec("a,b,c", ','); + ASSERT_EQ(v.size(), 3U); + EXPECT_EQ(v[0], "a"sv); + EXPECT_EQ(v[1], "b"sv); + EXPECT_EQ(v[2], "c"sv); +} + +TEST(SplitToStringVec, EmptyString) +{ + auto const v = splitToStringVec("", ','); + EXPECT_TRUE(v.empty()); +} + +TEST(SplitToStringVec, MaxSplit) +{ + // maxSplit=1 means at most one split; the rest of the string is the second element. + auto const v = splitToStringVec("a:b:c", ':', 1); + ASSERT_EQ(v.size(), 2U); + EXPECT_EQ(v[0], "a"sv); + EXPECT_EQ(v[1], "b:c"sv); +} + +TEST(SplitToStringVec, TrailingSeparator) +{ + auto const v = splitToStringVec("a,b,", ','); + ASSERT_EQ(v.size(), 3U); + EXPECT_EQ(v[0], "a"sv); + EXPECT_EQ(v[1], "b"sv); + EXPECT_EQ(v[2], ""sv); +} + +TEST(MatchStringWithOneWildcard, ExactMatch) +{ + EXPECT_TRUE(matchStringWithOneWildcard("hello", "hello")); + EXPECT_FALSE(matchStringWithOneWildcard("hello", "world")); + EXPECT_FALSE(matchStringWithOneWildcard("hello", "hello2")); +} + +TEST(MatchStringWithOneWildcard, WildcardMatchesAnything) +{ + EXPECT_TRUE(matchStringWithOneWildcard("*", "anything")); + EXPECT_TRUE(matchStringWithOneWildcard("*", "")); +} + +TEST(MatchStringWithOneWildcard, PrefixWildcard) +{ + EXPECT_TRUE(matchStringWithOneWildcard("hello*", "hello")); + EXPECT_TRUE(matchStringWithOneWildcard("hello*", "hello world")); + EXPECT_FALSE(matchStringWithOneWildcard("hello*", "world")); +} + +TEST(MatchStringWithOneWildcard, SuffixWildcard) +{ + EXPECT_TRUE(matchStringWithOneWildcard("*world", "world")); + EXPECT_TRUE(matchStringWithOneWildcard("*world", "hello world")); + EXPECT_FALSE(matchStringWithOneWildcard("*world", "hello")); +} + +TEST(MatchStringWithOneWildcard, MiddleWildcard) +{ + EXPECT_TRUE(matchStringWithOneWildcard("he*ld", "held")); + EXPECT_TRUE(matchStringWithOneWildcard("he*ld", "hello world")); + EXPECT_FALSE(matchStringWithOneWildcard("he*ld", "hello")); +} + +TEST(NormalizeDirectoryPath, AlreadyNormalized) +{ + EXPECT_EQ(normalizeDirectoryPath("/some/path/"), "/some/path/"sv); +} + +TEST(NormalizeDirectoryPath, MissingTrailingSlash) +{ + EXPECT_EQ(normalizeDirectoryPath("/some/path"), "/some/path/"sv); +} + +TEST(NormalizeDirectoryPath, EmptyString) +{ + EXPECT_EQ(normalizeDirectoryPath(""), ""sv); +} + +TEST(SanitizeRemoteConfig, Empty) +{ + EXPECT_EQ(sanitizeRemoteConfig(""), ""sv); +} + +TEST(SanitizeRemoteConfig, NoCredentials) +{ + // No @ means no credentials section; returned as-is. + EXPECT_EQ(sanitizeRemoteConfig("ssh://host:22"), "ssh://host:22"sv); +} + +TEST(SanitizeRemoteConfig, UsernameOnly) +{ + EXPECT_EQ(sanitizeRemoteConfig("ssh://user@host:22"), "ssh://***@host:22"sv); +} + +TEST(SanitizeRemoteConfig, UsernameAndPassword) +{ + EXPECT_EQ(sanitizeRemoteConfig("ssh://user:pass@host:22"), "ssh://***@host:22"sv); +} + +TEST(SanitizeRemoteConfig, WithQueryParams) +{ + EXPECT_EQ( + sanitizeRemoteConfig("ssh://admin:secret@server.com:22?timeout=30"), "ssh://***@server.com:22?timeout=30"sv); +} + +TEST(SanitizeRemoteConfig, MissingProtocol) +{ + EXPECT_EQ(sanitizeRemoteConfig("user:pass@host:22"), "***@host:22"sv); +} + +TEST(SanitizeRemoteConfig, EmptyCredentials) +{ + EXPECT_EQ(sanitizeRemoteConfig("ssh://@host:22"), "ssh://@host:22"sv); +} + +TEST(SanitizeRemoteConfig, PasswordContainingAt) +{ + EXPECT_EQ(sanitizeRemoteConfig("ssh://user:p@ss@host:22"), "ssh://***@host:22"sv); +} + +TEST(SanitizeRemoteConfig, AtInQueryIsNotCredentials) +{ + EXPECT_EQ(sanitizeRemoteConfig("ssh://host:22?user=a@b"), "ssh://host:22?user=a@b"sv); +} + +TEST(SanitizeArgv, MasksRemoteConfigCredentials) +{ + ArgVec av{"--remoteConfig=ssh://user:pass@host:22", "--safe"}; + auto const sanitized = sanitizeArgv(av.argc(), av.argv()); + ASSERT_EQ(sanitized.size(), 3U); + EXPECT_EQ(sanitized[1], "--remoteConfig=ssh://***@host:22"sv); + EXPECT_EQ(sanitized[2], "--safe"sv); +} + +TEST(SanitizeArgv, MasksAliasCredentials) +{ + ArgVec av{"--remoteAutoTuningConfig=ssh://user:pass@host:22"}; + auto const sanitized = sanitizeArgv(av.argc(), av.argv()); + ASSERT_EQ(sanitized.size(), 2U); + EXPECT_EQ(sanitized[1], "--remoteAutoTuningConfig=ssh://***@host:22"sv); +} + +TEST(SanitizeArgv, MasksMalformedCredentials) +{ + ArgVec av{"--remoteConfig=user:pass@host:22"}; + auto const sanitized = sanitizeArgv(av.argc(), av.argv()); + ASSERT_EQ(sanitized.size(), 2U); + EXPECT_EQ(sanitized[1], "--remoteConfig=***@host:22"sv); +} + +TEST(SanitizeArgv, MasksMalformedCredentialsForAlias) +{ + ArgVec av{"--remoteAutoTuningConfig=user:pass@host:22"}; + auto const sanitized = sanitizeArgv(av.argc(), av.argv()); + ASSERT_EQ(sanitized.size(), 2U); + EXPECT_EQ(sanitized[1], "--remoteAutoTuningConfig=***@host:22"sv); +} + +TEST(SanitizeArgv, LeavesOtherArgumentsUntouched) +{ + ArgVec av{"--onnx=model.onnx", "--remoteConfig="}; + auto const sanitized = sanitizeArgv(av.argc(), av.argv()); + ASSERT_EQ(sanitized.size(), 3U); + EXPECT_EQ(sanitized[1], "--onnx=model.onnx"sv); + EXPECT_EQ(sanitized[2], "--remoteConfig="sv); +} diff --git a/samples/common/streamReader.h b/samples/trtexecCommon/streamReader.h similarity index 100% rename from samples/common/streamReader.h rename to samples/trtexecCommon/streamReader.h