From 5731c07d6800fe4a87629726e77d86819469d0b4 Mon Sep 17 00:00:00 2001 From: tianruih Date: Fri, 21 Aug 2026 15:10:27 -0700 Subject: [PATCH 01/29] [None][feat] port NVFP4 cold-page compression to latest KVCM2 Signed-off-by: tianruih --- cpp/tensorrt_llm/CMakeLists.txt | 2 + .../kernels/nvfp4BoundaryKernels.cu | 997 +++++++++++++ .../kernels/nvfp4BoundaryKernels.h | 116 ++ .../kv_cache_compression/CMakeLists.txt | 9 + .../nvfp4ColdPageCodec.cpp | 409 +++++ .../kv_cache_compression/nvfp4ColdPageCodec.h | 90 ++ cpp/tensorrt_llm/nanobind/CMakeLists.txt | 1 + cpp/tensorrt_llm/nanobind/bindings.cpp | 4 + .../nanobind/kvCacheCompression/bindings.cpp | 66 + .../nanobind/kvCacheCompression/bindings.h | 27 + cpp/tests/unit_tests/CMakeLists.txt | 1 + cpp/tests/unit_tests/kernels/CMakeLists.txt | 16 + .../kernels/nvfp4BoundaryKernelsTest.cpp | 1310 +++++++++++++++++ .../kv_cache_compression/CMakeLists.txt | 22 + .../nvfp4ColdPageCodecTest.cpp | 572 +++++++ .../quantization_for_cold_page/__init__.py | 6 + .../quantization_for_cold_page.py | 175 +++ .../triattention/triattention.py | 18 +- tensorrt_llm/_torch/pyexecutor/_util.py | 131 +- .../_torch/pyexecutor/kv_cache_manager_v2.py | 23 +- .../_torch/pyexecutor/resource_manager.py | 49 +- tensorrt_llm/llmapi/__init__.py | 8 +- tensorrt_llm/llmapi/llm_args.py | 34 +- .../usage/llm_args_golden_manifest.json | 18 +- tensorrt_llm/usage/llmapi_config.py | 22 +- .../executor/test_kv_cache_budget_split.py | 2 +- .../test_kv_cache_compression_manager.py | 198 ++- .../executor/test_kv_cache_estimation.py | 7 +- .../executor/test_kv_cache_manager_v2.py | 30 + .../_torch/kv_cache_compression/conftest.py | 5 +- .../test_quantization_for_cold_page.py | 417 ++++++ .../test_triattention_pipeline.py | 11 +- .../api_stability/references/llm.yaml | 2 +- tests/unittest/llmapi/test_llm_args.py | 42 +- .../test_llmapi_config_telemetry_docs.py | 53 + 35 files changed, 4750 insertions(+), 143 deletions(-) create mode 100644 cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu create mode 100644 cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h create mode 100644 cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h create mode 100644 cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp create mode 100644 cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.h create mode 100644 cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp create mode 100644 cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt create mode 100644 cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp create mode 100644 tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py create mode 100644 tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py create mode 100644 tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py diff --git a/cpp/tensorrt_llm/CMakeLists.txt b/cpp/tensorrt_llm/CMakeLists.txt index d9a7853bc9d5..096783b735c5 100644 --- a/cpp/tensorrt_llm/CMakeLists.txt +++ b/cpp/tensorrt_llm/CMakeLists.txt @@ -144,6 +144,7 @@ add_subdirectory(common) add_subdirectory(kernels) add_subdirectory(layers) add_subdirectory(runtime) +add_subdirectory(kv_cache_compression) set(BATCH_MANAGER_TARGET tensorrt_llm_batch_manager_static) set(BATCH_MANAGER_TARGET_ARCH ${TARGET_ARCH}) @@ -198,6 +199,7 @@ set(TRTLLM_LINK_LIBS cute_dsl_src layers_src runtime_src + kv_cache_compression_src compressorKernels_src mhcKernels_src userbuffers_src diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu new file mode 100644 index 000000000000..71147fd15db3 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu @@ -0,0 +1,997 @@ +/* + * 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 "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" + +#include "tensorrt_llm/common/assert.h" +#include "tensorrt_llm/common/cudaUtils.h" +#include "tensorrt_llm/common/envUtils.h" +#include "tensorrt_llm/kernels/cudaAsyncOps.cuh" +#include "tensorrt_llm/kernels/quantization.cuh" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +TRTLLM_NAMESPACE_BEGIN + +namespace kernels +{ +namespace +{ + +// Match KVCM V2's mapped-Host copy CTA and one-split policy. +constexpr std::uint32_t kThreadsPerBlock = 128; +constexpr std::uint32_t kAsyncStages = 4; +// Mapped-Host reads use eight cp.async stages; GPU-resident input uses four. +constexpr std::uint32_t kHostLoadAsyncStages = 8; +constexpr std::uint32_t kHostMemorySplits = 1; +constexpr std::uint32_t kMaxTasksPerLaunch = 256; +// Bound by-value buffer metadata to CUDA's kernel-parameter limit. +constexpr std::uint32_t kMaxBuffersPerLaunch = kNvfp4BoundaryMaxBuffersPerLaunch; +constexpr std::uint32_t kElementsPerLane = 8; +constexpr std::uint32_t kElementsPerBlockScale = 16; +// Bound per-tile shared scale staging to 1 KiB. +constexpr std::uint32_t kTargetScaleTransferBytes = 1024; +constexpr std::uint32_t kHalfGroupsPerTransfer = 2U * kTargetScaleTransferBytes; +constexpr std::size_t kModernKernelParameterLimit = 32764; + +static_assert(kThreadsPerBlock % 2 == 0, "An NVFP4 scale group is shared by two lanes"); +static_assert(kTargetScaleTransferBytes > 0 && kTargetScaleTransferBytes % sizeof(uint4) == 0, + "Compact scale transfer must remain a positive 16-byte multiple"); +static_assert(kHalfGroupsPerTransfer % 2U == 0, "A transfer tile must not split an NVFP4 scale group"); +static_assert(std::is_trivially_copyable_v); +static_assert(std::is_trivially_copyable_v); +static_assert(std::is_trivially_copyable_v); +static_assert( + kMaxBuffersPerLaunch <= std::numeric_limits::max(), "Boundary buffer count must fit CUDA grid.y"); +static_assert(sizeof(std::array) + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + + 3U * sizeof(std::uint32_t) + <= kModernKernelParameterLimit); +static_assert(sizeof(std::array) + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + + 3U * sizeof(std::uint32_t) + <= kModernKernelParameterLimit); + +// Issue one predicated 16-byte cp.async load. +template +__device__ __forceinline__ void copyAsyncGlobalToShared(T* shared, T const* global, bool valid) +{ + static_assert(sizeof(T) == 16, "Boundary transfer grains must match batchedCopy's 16-byte width"); + if (valid) + { + asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" + : + : "l"(__cvta_generic_to_shared(shared)), "l"(global) + : "memory"); + } +} + +struct OffloadBufferTask +{ + std::uint8_t const* raw; + std::uint8_t* coldData; + std::uint8_t* coldScale; + std::uint8_t* coldPadding; +}; + +struct OnboardBufferTask +{ + std::uint8_t const* coldData; + std::uint8_t const* coldScale; + std::uint8_t* raw; +}; + +__device__ OffloadBufferTask resolveTask(Nvfp4BoundaryOffloadPageTask const& page, + Nvfp4BoundaryBufferPlan const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) +{ + std::size_t const gpuPage = static_cast(page.gpuPageIndex); + auto* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; + return {reinterpret_cast(buffer.rawBase + gpuPage * buffer.rawSlotBytes), + coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, coldPage + buffer.coldPaddingOffset}; +} + +__device__ OnboardBufferTask resolveTask(Nvfp4BoundaryOnboardPageTask const& page, + Nvfp4BoundaryBufferPlan const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) +{ + std::size_t const gpuPage = static_cast(page.gpuPageIndex); + auto const* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; + return {coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, + reinterpret_cast(buffer.rawBase + gpuPage * buffer.rawSlotBytes)}; +} + +__host__ __device__ constexpr std::uint32_t packedBytesPerBuffer(std::uint32_t halfGroups) +{ + return halfGroups * sizeof(std::uint32_t); +} + +__host__ __device__ constexpr std::uint32_t scaleBytesPerBuffer(std::uint32_t halfGroups) +{ + return halfGroups / 2U; +} + +// Shared staging padding is not stored in the cold record. +__host__ __device__ constexpr std::uint32_t packedStagingBytesPerBuffer(std::uint32_t halfGroups) +{ + return (packedBytesPerBuffer(halfGroups) + sizeof(uint4) - 1U) / sizeof(uint4) * sizeof(uint4); +} + +__host__ __device__ constexpr std::uint32_t compactStagingBytesPerBuffer(std::uint32_t halfGroups) +{ + return packedStagingBytesPerBuffer(halfGroups) + scaleBytesPerBuffer(halfGroups); +} + +__host__ __device__ constexpr std::uint32_t totalHalfGroupsPerBuffer(Nvfp4BoundaryKernelParams const& params) +{ + return static_cast(params.numKvHeads) * static_cast(params.tokensPerPage) + * (static_cast(params.headDim) / kElementsPerLane); +} + +// Tile at most 2,048 half-groups; headDim % 16 prevents splitting a scale group. +__host__ __device__ constexpr std::uint32_t compressedTransferHalfGroups(Nvfp4BoundaryKernelParams const& params) +{ + // Avoid std::min: its reference return ODR-uses this host/device constexpr. + auto const halfGroups = totalHalfGroupsPerBuffer(params); + return halfGroups < kHalfGroupsPerTransfer ? halfGroups : kHalfGroupsPerTransfer; +} + +// Flush vectorized packed values and scales, then copy remaining tails bytewise. +__device__ void flushCompactRangeToHost(std::uint8_t const* compactStages, OffloadBufferTask const& task, + std::uint32_t packedStageCapacityBytes, std::uint32_t packedDestinationOffset, std::uint32_t packedBytes, + std::uint32_t scaleDestinationOffset, std::uint32_t scaleBytes) +{ + auto const* packedSource = compactStages; + auto* packedDestination = task.coldData + packedDestinationOffset; + bool const alignedPacked = reinterpret_cast(packedSource) % sizeof(uint4) == 0 + && reinterpret_cast(packedDestination) % sizeof(uint4) == 0; + bool const alignedPackedPair = reinterpret_cast(packedSource) % sizeof(uint2) == 0 + && reinterpret_cast(packedDestination) % sizeof(uint2) == 0; + std::uint32_t const packedVectorBytes = alignedPacked ? packedBytes - packedBytes % sizeof(uint4) : 0; + std::uint32_t const packedPairBytes = alignedPackedPair ? packedBytes - packedBytes % sizeof(uint2) : 0; + for (std::uint32_t grain = threadIdx.x; grain < packedVectorBytes / sizeof(uint4); grain += blockDim.x) + { + reinterpret_cast(packedDestination)[grain] = reinterpret_cast(packedSource)[grain]; + } + for (std::uint32_t pair = packedVectorBytes / sizeof(uint2) + threadIdx.x; pair < packedPairBytes / sizeof(uint2); + pair += blockDim.x) + { + reinterpret_cast(packedDestination)[pair] = reinterpret_cast(packedSource)[pair]; + } + for (std::uint32_t byte = packedPairBytes + threadIdx.x; byte < packedBytes; byte += blockDim.x) + { + packedDestination[byte] = packedSource[byte]; + } + + auto const* scaleSource = compactStages + packedStageCapacityBytes; + auto* scaleDestination = task.coldScale + scaleDestinationOffset; + bool const alignedScale = reinterpret_cast(scaleSource) % sizeof(uint4) == 0 + && reinterpret_cast(scaleDestination) % sizeof(uint4) == 0; + std::uint32_t const scaleVectorBytes = alignedScale ? scaleBytes - scaleBytes % sizeof(uint4) : 0; + for (std::uint32_t grain = threadIdx.x; grain < scaleVectorBytes / sizeof(uint4); grain += blockDim.x) + { + reinterpret_cast(scaleDestination)[grain] = reinterpret_cast(scaleSource)[grain]; + } + for (std::uint32_t byte = scaleVectorBytes + threadIdx.x; byte < scaleBytes; byte += blockDim.x) + { + scaleDestination[byte] = scaleSource[byte]; + } +} + +// Zero codec-specified record padding so persisted cold Slots are deterministic. +__device__ void clearColdPadding(OffloadBufferTask const& task, Nvfp4BoundaryBufferPlan const& buffer) +{ + if (blockIdx.x != 0U) + { + return; + } + for (std::uint32_t byte = threadIdx.x; byte < buffer.coldPaddingBytes; byte += blockDim.x) + { + task.coldPadding[byte] = 0U; + } +} + +// CTA-uniform byte-exact copy used by lossless side buffers. Vectorize only when both +// endpoints permit it; arbitrary descriptor offsets and byte tails remain supported. +__device__ void copyLosslessBytes(std::uint8_t const* source, std::uint8_t* destination, std::size_t bytes) +{ + bool const aligned = reinterpret_cast(source) % sizeof(uint4) == 0 + && reinterpret_cast(destination) % sizeof(uint4) == 0; + std::size_t const vectorBytes = aligned ? bytes - bytes % sizeof(uint4) : 0U; + for (std::size_t grain = threadIdx.x; grain < vectorBytes / sizeof(uint4); grain += blockDim.x) + { + reinterpret_cast(destination)[grain] = reinterpret_cast(source)[grain]; + } + for (std::size_t byte = vectorBytes + threadIdx.x; byte < bytes; byte += blockDim.x) + { + destination[byte] = source[byte]; + } +} + +__device__ __forceinline__ uint4 collectTwoFp8Words(std::uint64_t first, std::uint64_t second) +{ + return make_uint4(static_cast(first), static_cast(first >> 32U), + static_cast(second), static_cast(second >> 32U)); +} + +// E2M1-to-FP16x2 PTX is adapted from arcquantFP4.cu and fusedMoeCommKernels.cu. +__device__ void unpackE2m1ToFloat(std::uint32_t packed, float2 (&values)[4]) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + std::uint32_t fp16Pairs[4]; + asm volatile( + "{\n" + ".reg .b8 b0, b1, b2, b3;\n" + "mov.b32 {b0, b1, b2, b3}, %4;\n" + "cvt.rn.f16x2.e2m1x2 %0, b0;\n" + "cvt.rn.f16x2.e2m1x2 %1, b1;\n" + "cvt.rn.f16x2.e2m1x2 %2, b2;\n" + "cvt.rn.f16x2.e2m1x2 %3, b3;\n" + "}\n" + : "=r"(fp16Pairs[0]), "=r"(fp16Pairs[1]), "=r"(fp16Pairs[2]), "=r"(fp16Pairs[3]) + : "r"(packed)); + +#pragma unroll + for (std::uint32_t i = 0; i < 4; ++i) + { + values[i] = __half22float2(reinterpret_cast<__half2&>(fp16Pairs[i])); + } +#endif +} + +template +__device__ void store16BitValues(T* output, std::uint32_t elementOffset, float2 const (&values)[4], float scale) +{ + // Store through uint4 to preserve STG.128; nvcc may scalarize PackedVec. + std::uint32_t outputWords[4]; +#pragma unroll + for (std::uint32_t i = 0; i < 4; ++i) + { + float2 const scaled = make_float2(values[i].x * scale, values[i].y * scale); + if constexpr (std::is_same_v) + { + half2 const pair = __float22half2_rn(scaled); + outputWords[i] = reinterpret_cast(pair); + } + else + { + __nv_bfloat162 const pair = __float22bfloat162_rn(scaled); + outputWords[i] = reinterpret_cast(pair); + } + } + uint4 const outputGrain = make_uint4(outputWords[0], outputWords[1], outputWords[2], outputWords[3]); + reinterpret_cast(output + elementOffset)[0] = outputGrain; +} + +// Preserve independent source-FP8 and destination-NVFP4 global scales. +// Restore in FP32, round pairs to FP16, then reduce each 16-value grain. +__device__ uint2 quantizeFp8GrainToNvfp4(PackedVec<__nv_fp8_e4m3> const& grain, float fp8ScaleQuantOrig, + float nvfp4ScaleOrigQuant, std::uint8_t* scaleOutput) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + // Use the production packed FP8x2 conversion surface. + PackedVec restored[2]; +#pragma unroll + for (std::uint32_t pair = 0; pair < 8; ++pair) + { + float2 values = static_cast(grain.elts[pair]); + values.x *= fp8ScaleQuantOrig; + values.y *= fp8ScaleQuantOrig; + restored[pair / 4U].elts[pair % 4U] = __float22half2_rn(values); + } + + auto firstHalfMax = cuda_abs(restored[0].elts[0]); + auto secondHalfMax = cuda_abs(restored[1].elts[0]); +#pragma unroll + for (std::uint32_t i = 1; i < 4; ++i) + { + firstHalfMax = cuda_max(firstHalfMax, cuda_abs(restored[0].elts[i])); + secondHalfMax = cuda_max(secondHalfMax, cuda_abs(restored[1].elts[i])); + } + auto const localMax = cuda_max(firstHalfMax, secondHalfMax); + float const vecMax = static_cast(cuda_max(localMax.x, localMax.y)); + + float scaleValue = nvfp4ScaleOrigQuant * (vecMax * reciprocal_approximate_ftz(6.0F)); + __nv_fp8_e4m3 const roundedScale(scaleValue); + *scaleOutput = roundedScale.__x; + scaleValue = static_cast(roundedScale); + float const outputScale = vecMax != 0.0F + ? reciprocal_approximate_ftz(scaleValue * reciprocal_approximate_ftz(nvfp4ScaleOrigQuant)) + : 0.0F; + + std::uint32_t packed[2]; +#pragma unroll + for (std::uint32_t halfGroup = 0; halfGroup < 2; ++halfGroup) + { + float2 values[4]; +#pragma unroll + for (std::uint32_t i = 0; i < 4; ++i) + { + values[i] = __half22float2(restored[halfGroup].elts[i]); + values[i].x *= outputScale; + values[i].y *= outputScale; + } + packed[halfGroup] = fp32_vec_to_e2m1(values); + } + return make_uint2(packed[0], packed[1]); +#else + static_cast(grain); + static_cast(fp8ScaleQuantOrig); + static_cast(nvfp4ScaleOrigQuant); + static_cast(scaleOutput); + return make_uint2(0U, 0U); +#endif +} + +// Restore one natural 16-value NVFP4 scale group. +template +__device__ float onboardDequantScale(std::uint8_t encodedScale, Nvfp4BoundaryKernelParams const& params) +{ + __nv_fp8_e4m3 blockScale; + blockScale.__x = encodedScale; + float scale = static_cast(blockScale) * params.nvfp4ScaleQuantOrig; + if constexpr (std::is_same_v) + { + scale *= params.fp8ScaleOrigQuant; + } + return scale; +} + +template +__device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firstHalfGroup, float dequantScale) +{ + std::uint32_t const packedWords[2] = {packedPair.x, packedPair.y}; + if constexpr (!std::is_same_v) + { +#pragma unroll + for (std::uint32_t laneInScale = 0; laneInScale < 2; ++laneInScale) + { + float2 values[4]; + unpackE2m1ToFloat(packedWords[laneInScale], values); + store16BitValues(output, (firstHalfGroup + laneInScale) * kElementsPerLane, values, dequantScale); + } + } + else + { + std::uint64_t packedFp8[2]; +#pragma unroll + for (std::uint32_t laneInScale = 0; laneInScale < 2; ++laneInScale) + { + float2 values[4]; + unpackE2m1ToFloat(packedWords[laneInScale], values); +#pragma unroll + for (std::uint32_t i = 0; i < 4; ++i) + { + values[i].x *= dequantScale; + values[i].y *= dequantScale; + } + packedFp8[laneInScale] = fp32_vec_to_e4m3(values); + } + reinterpret_cast(output)[firstHalfGroup / 2U] = collectTwoFp8Words(packedFp8[0], packedFp8[1]); + } +} + +// FP16/BF16 GPU Page -> mapped-Host NVFP4 in bounded tiles. +template +__global__ void offloadFrom16BitTiledKernel( + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, std::uint8_t* coldBase, + std::size_t coldPageBytes, std::uint32_t numBuffers) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + asm volatile("griddepcontrol.launch_dependents;\n"); + + std::uint32_t const bufferIndex = blockIdx.y; + assert(bufferIndex < numBuffers); + auto const& buffer = buffers[bufferIndex]; + auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + + asm volatile("griddepcontrol.wait;\n" : : : "memory"); + if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + { + copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); + clearColdPadding(task, buffer); + return; + } + + auto const params = buffer.params; + std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); + std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); + std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + + // Use a four-stage 16-byte cp.async ring and tile-bounded compact staging. + __shared__ __align__(16) PackedVec rawStages[kAsyncStages][kThreadsPerBlock]; + extern __shared__ __align__(16) std::uint8_t compactStages[]; + auto* packedStages = reinterpret_cast(compactStages); + auto* scaleStages = compactStages + packedStageCapacityBytes; + + for (std::uint32_t firstHalfGroup = blockIdx.x * tileHalfGroups; firstHalfGroup < halfGroupsPerBuffer; + firstHalfGroup += gridDim.x * tileHalfGroups) + { + std::uint32_t const halfGroups = std::min(tileHalfGroups, halfGroupsPerBuffer - firstHalfGroup); + std::uint32_t const iterations = (halfGroups + kThreadsPerBlock - 1U) / kThreadsPerBlock; + + for (std::uint32_t iteration = 0; iteration < iterations + kAsyncStages; ++iteration) + { + std::uint32_t const stage = iteration % kAsyncStages; + if (iteration >= kAsyncStages) + { + std::uint32_t const transformIteration = iteration - kAsyncStages; + std::uint32_t const localHalfGroup = kThreadsPerBlock * transformIteration + threadIdx.x; + cp_async_wait_group(); + if (localHalfGroup < halfGroups) + { + std::uint32_t const globalHalfGroup = firstHalfGroup + localHalfGroup; + std::uint32_t const laneInScale = globalHalfGroup & 1U; + + // Even tile boundaries preserve the 16-value scale groups. + std::uint32_t const localScaleOffset = localHalfGroup >> 1U; + std::uint8_t* scale = laneInScale == 0 ? scaleStages + localScaleOffset : nullptr; + PackedVec input = rawStages[stage][threadIdx.x]; + packedStages[localHalfGroup] = cvt_warp_fp16_to_fp4( + input, params.nvfp4ScaleOrigQuant, scale); + } + } + + std::uint32_t const localLoadHalfGroup = kThreadsPerBlock * iteration + threadIdx.x; + bool const valid = localLoadHalfGroup < halfGroups; + auto const* rawInput = reinterpret_cast const*>(task.raw); + auto const* source = valid ? rawInput + firstHalfGroup + localLoadHalfGroup : rawInput; + copyAsyncGlobalToShared(&rawStages[stage][threadIdx.x], source, valid); + cp_async_commit_group(); + } + + // Publish quant results and finish the flush before reusing staging. + cp_async_wait_group<0>(); + __syncthreads(); + flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), + packedBytesPerBuffer(halfGroups), firstHalfGroup / 2U, halfGroups / 2U); + __syncthreads(); + } + clearColdPadding(task, buffer); +#endif +} + +// FP8 E4M3 GPU Page -> mapped-Host NVFP4 in bounded tiles. +__global__ void offloadFromFp8TiledKernel( + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, std::uint8_t* coldBase, + std::size_t coldPageBytes, std::uint32_t numBuffers) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + asm volatile("griddepcontrol.launch_dependents;\n"); + + std::uint32_t const bufferIndex = blockIdx.y; + assert(bufferIndex < numBuffers); + auto const& buffer = buffers[bufferIndex]; + auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + + asm volatile("griddepcontrol.wait;\n" : : : "memory"); + if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + { + copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); + clearColdPadding(task, buffer); + return; + } + + auto const params = buffer.params; + std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); + std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); + std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + + // Each cp.async moves one 16-byte grain in the production PackedVec layout. + __shared__ __align__(16) PackedVec<__nv_fp8_e4m3> rawStages[kAsyncStages][kThreadsPerBlock]; + extern __shared__ __align__(16) std::uint8_t compactStages[]; + auto* packedStages = reinterpret_cast(compactStages); + auto* scaleStages = compactStages + packedStageCapacityBytes; + + for (std::uint32_t firstHalfGroup = blockIdx.x * tileHalfGroups; firstHalfGroup < halfGroupsPerBuffer; + firstHalfGroup += gridDim.x * tileHalfGroups) + { + std::uint32_t const halfGroups = std::min(tileHalfGroups, halfGroupsPerBuffer - firstHalfGroup); + std::uint32_t const firstGrain = firstHalfGroup / 2U; + std::uint32_t const grains = halfGroups / 2U; + std::uint32_t const iterations = (grains + kThreadsPerBlock - 1U) / kThreadsPerBlock; + + for (std::uint32_t iteration = 0; iteration < iterations + kAsyncStages; ++iteration) + { + std::uint32_t const stage = iteration % kAsyncStages; + if (iteration >= kAsyncStages) + { + std::uint32_t const transformIteration = iteration - kAsyncStages; + cp_async_wait_group(); + // One lane owns the complete 16-value scale group. + std::uint32_t const localGrain = kThreadsPerBlock * transformIteration + threadIdx.x; + if (localGrain < grains) + { + uint2 const packed = quantizeFp8GrainToNvfp4(rawStages[stage][threadIdx.x], + params.fp8ScaleQuantOrig, params.nvfp4ScaleOrigQuant, scaleStages + localGrain); + reinterpret_cast(packedStages)[localGrain] = packed; + } + } + + std::uint32_t const localLoadGrain = kThreadsPerBlock * iteration + threadIdx.x; + bool const valid = localLoadGrain < grains; + auto const* rawInput = reinterpret_cast const*>(task.raw); + auto const* source = valid ? rawInput + firstGrain + localLoadGrain : rawInput; + copyAsyncGlobalToShared(&rawStages[stage][threadIdx.x], source, valid); + cp_async_commit_group(); + } + + cp_async_wait_group<0>(); + __syncthreads(); + flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), + packedBytesPerBuffer(halfGroups), firstHalfGroup / 2U, halfGroups / 2U); + __syncthreads(); + } + clearColdPadding(task, buffer); +#endif +} + +// Load packed values and one-byte scales from mapped Host memory into shared. +// Use uint4/uint2 when aligned and byte tails otherwise. +__device__ void loadCompactRangeFromHost(std::uint8_t* compactStages, OnboardBufferTask const& task, + std::uint32_t packedStageCapacityBytes, std::uint32_t packedSourceOffset, std::uint32_t packedBytes, + std::uint32_t scaleSourceOffset, std::uint32_t scaleBytes) +{ + auto const* packedSource = task.coldData + packedSourceOffset; + auto* packedDestination = compactStages; + auto const* scaleSource = task.coldScale + scaleSourceOffset; + auto* scaleDestination = compactStages + packedStageCapacityBytes; + bool const alignedPacked = reinterpret_cast(packedSource) % sizeof(uint4) == 0 + && reinterpret_cast(packedDestination) % sizeof(uint4) == 0; + bool const alignedPackedPair = reinterpret_cast(packedSource) % sizeof(uint2) == 0 + && reinterpret_cast(packedDestination) % sizeof(uint2) == 0; + std::uint32_t const packedVectorBytes = alignedPacked ? packedBytes - packedBytes % sizeof(uint4) : 0; + std::uint32_t const packedPairBytes = alignedPackedPair ? packedBytes - packedBytes % sizeof(uint2) : 0; + bool const alignedScale = reinterpret_cast(scaleSource) % sizeof(uint4) == 0 + && reinterpret_cast(scaleDestination) % sizeof(uint4) == 0; + std::uint32_t const scaleVectorBytes = alignedScale ? scaleBytes - scaleBytes % sizeof(uint4) : 0; + std::uint32_t const packedGrains = packedVectorBytes / sizeof(uint4); + std::uint32_t const scaleGrains = scaleVectorBytes / sizeof(uint4); + std::uint32_t const totalGrains = packedGrains + scaleGrains; + std::uint32_t const iterations = (totalGrains + blockDim.x - 1U) / blockDim.x; + + for (std::uint32_t iteration = 0; iteration < iterations; ++iteration) + { + if (iteration >= kHostLoadAsyncStages) + { + cp_async_wait_group(); + } + std::uint32_t const grain = blockDim.x * iteration + threadIdx.x; + bool const valid = grain < totalGrains; + auto const* source = reinterpret_cast(packedSource); + auto* destination = reinterpret_cast(packedDestination); + if (valid && grain < packedGrains) + { + source += grain; + destination += grain; + } + else if (valid) + { + source = reinterpret_cast(scaleSource) + grain - packedGrains; + destination = reinterpret_cast(scaleDestination) + grain - packedGrains; + } + copyAsyncGlobalToShared(destination, source, valid); + cp_async_commit_group(); + } + cp_async_wait_group<0>(); + + // headDim % 16 makes packed intervals exact uint2 scale groups. + for (std::uint32_t pair = packedVectorBytes / sizeof(uint2) + threadIdx.x; pair < packedPairBytes / sizeof(uint2); + pair += blockDim.x) + { + reinterpret_cast(packedDestination)[pair] = reinterpret_cast(packedSource)[pair]; + } + for (std::uint32_t byte = packedPairBytes + threadIdx.x; byte < packedBytes; byte += blockDim.x) + { + packedDestination[byte] = packedSource[byte]; + } + for (std::uint32_t byte = scaleVectorBytes + threadIdx.x; byte < scaleBytes; byte += blockDim.x) + { + scaleDestination[byte] = scaleSource[byte]; + } + __syncthreads(); +} + +// Mapped-Host NVFP4 -> runtime GPU Page in bounded tiles. +template +__global__ void onboardTiledKernel( + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, + std::uint8_t const* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + asm volatile("griddepcontrol.launch_dependents;\n"); + + std::uint32_t const bufferIndex = blockIdx.y; + assert(bufferIndex < numBuffers); + auto const& buffer = buffers[bufferIndex]; + auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + + asm volatile("griddepcontrol.wait;\n" : : : "memory"); + if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + { + copyLosslessBytes(task.coldData, task.raw, buffer.rawBytes); + return; + } + + auto const params = buffer.params; + std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); + std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); + std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + extern __shared__ __align__(16) std::uint8_t compactStages[]; + auto* rawOutput = reinterpret_cast(task.raw); + + for (std::uint32_t firstHalfGroup = blockIdx.x * tileHalfGroups; firstHalfGroup < halfGroupsPerBuffer; + firstHalfGroup += gridDim.x * tileHalfGroups) + { + std::uint32_t const halfGroups = std::min(tileHalfGroups, halfGroupsPerBuffer - firstHalfGroup); + std::uint32_t const packedBytes = packedBytesPerBuffer(halfGroups); + std::uint32_t const scaleBytes = halfGroups / 2U; + + // Stage packed data and scales before dequantization. + loadCompactRangeFromHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), + packedBytes, firstHalfGroup / 2U, scaleBytes); + + auto const* packedStages = reinterpret_cast(compactStages); + auto const* scaleStages = compactStages + packedStageCapacityBytes; + std::uint32_t const packedGrains = packedBytes / sizeof(uint4); + for (std::uint32_t localGrain = threadIdx.x; localGrain < packedGrains; localGrain += blockDim.x) + { + uint4 const packedGrain = packedStages[localGrain]; + std::uint32_t const packedWords[4] = {packedGrain.x, packedGrain.y, packedGrain.z, packedGrain.w}; +#pragma unroll + for (std::uint32_t pair = 0; pair < 2; ++pair) + { + std::uint32_t const localScaleGroup = localGrain * 2U + pair; + std::uint32_t const firstPairHalfGroup = firstHalfGroup + localScaleGroup * 2U; + + restoreNvfp4Pair(make_uint2(packedWords[pair * 2U], packedWords[pair * 2U + 1U]), rawOutput, + firstPairHalfGroup, onboardDequantScale(scaleStages[localScaleGroup], params)); + } + } + if (packedBytes % sizeof(uint4) != 0U && threadIdx.x == 0) + { + std::uint32_t const localScaleGroup = packedGrains * 2U; + std::uint32_t const firstPairHalfGroup = firstHalfGroup + localScaleGroup * 2U; + restoreNvfp4Pair(reinterpret_cast(compactStages)[localScaleGroup], rawOutput, + firstPairHalfGroup, onboardDequantScale(scaleStages[localScaleGroup], params)); + } + + // Finish consumers before reusing shared memory. + __syncthreads(); + } +#endif +} + +void validateParams(Nvfp4BoundaryKernelParams const& params, bool useFp8) +{ + TLLM_CHECK_WITH_INFO(params.numKvHeads > 0, "numKvHeads must be positive"); + TLLM_CHECK_WITH_INFO(params.tokensPerPage > 0, "tokensPerPage must be positive"); + TLLM_CHECK_WITH_INFO(params.headDim > 0 && params.headDim % kElementsPerBlockScale == 0, + "headDim must be positive and divisible by 16, got %d", params.headDim); + + std::uint64_t const rows + = static_cast(params.numKvHeads) * static_cast(params.tokensPerPage); + constexpr std::uint64_t maxHalfGroups = std::numeric_limits::max() / kElementsPerLane; + std::uint64_t const halfGroupsPerRow = static_cast(params.headDim / kElementsPerLane); + TLLM_CHECK_WITH_INFO(rows <= maxHalfGroups / halfGroupsPerRow, + "Page geometry exceeds the 32-bit compact-offset range: " + "heads=%d, tokens=%d, headDim=%d", + params.numKvHeads, params.tokensPerPage, params.headDim); + + TLLM_CHECK_WITH_INFO(std::isfinite(params.nvfp4ScaleOrigQuant) && params.nvfp4ScaleOrigQuant > 0.0F, + "NVFP4 original-to-quantized scale must be finite and positive"); + TLLM_CHECK_WITH_INFO(std::isfinite(params.nvfp4ScaleQuantOrig) && params.nvfp4ScaleQuantOrig > 0.0F, + "NVFP4 quantized-to-original scale must be finite and positive"); + if (useFp8) + { + TLLM_CHECK_WITH_INFO(std::isfinite(params.fp8ScaleOrigQuant) && params.fp8ScaleOrigQuant > 0.0F, + "FP8 original-to-quantized scale must be finite and positive"); + TLLM_CHECK_WITH_INFO(std::isfinite(params.fp8ScaleQuantOrig) && params.fp8ScaleQuantOrig > 0.0F, + "FP8 quantized-to-original scale must be finite and positive"); + } +} + +struct ColdInterval +{ + std::size_t begin; + std::size_t end; +}; + +void addColdInterval(std::vector& intervals, std::size_t offset, std::size_t bytes, + std::size_t coldPageBytes, char const* label) +{ + if (bytes == 0U) + { + return; + } + TLLM_CHECK_WITH_INFO( + offset <= coldPageBytes && bytes <= coldPageBytes - offset, "%s exceeds the cold Page stride", label); + intervals.push_back({offset, offset + bytes}); +} + +void validateBufferPlan( + Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldPageBytes, bool useFp8, std::vector& intervals) +{ + TLLM_CHECK_WITH_INFO(buffer.rawBase != 0U, "rawBase must not be null"); + TLLM_CHECK_WITH_INFO(buffer.rawBytes > 0 && buffer.rawBytes <= buffer.rawSlotBytes, + "Raw buffer bytes must be positive and fit within the GPU Slot stride"); + + switch (buffer.transform) + { + case Nvfp4BoundaryTransform::kNvfp4: + { + validateParams(buffer.params, useFp8); + TLLM_CHECK_WITH_INFO( + buffer.rawBase % alignof(uint4) == 0, "rawBase must be aligned to %zu bytes", alignof(uint4)); + TLLM_CHECK_WITH_INFO(buffer.rawSlotBytes % alignof(uint4) == 0, + "GPU raw Slot stride must be aligned to %zu bytes for NVFP4", alignof(uint4)); + std::uint64_t const elements = static_cast(buffer.params.numKvHeads) + * static_cast(buffer.params.tokensPerPage) + * static_cast(buffer.params.headDim); + std::uint64_t const expectedRawBytes = elements * (useFp8 ? 1U : 2U); + TLLM_CHECK_WITH_INFO(buffer.rawBytes == static_cast(expectedRawBytes), + "Raw buffer size does not match NVFP4 geometry and runtime type"); + addColdInterval(intervals, buffer.coldDataOffset, static_cast(elements / 2U), coldPageBytes, + "NVFP4 packed-data interval"); + addColdInterval(intervals, buffer.coldScaleOffset, static_cast(elements / 16U), coldPageBytes, + "NVFP4 scale interval"); + break; + } + case Nvfp4BoundaryTransform::kLossless: + addColdInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes, "Lossless-data interval"); + break; + default: TLLM_THROW("Unsupported NVFP4 boundary buffer transform"); + } + addColdInterval( + intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes, "Cold-record padding interval"); +} + +// Launch the 256-descriptor raw-argument ABI with optional PDL. +// Only a partial chunk needs a zero-padded stack argument. +template +void launchBoundaryBatch(Kernel kernel, Task const* tasks, std::uint32_t count, dim3 grid, dim3 block, + std::uint32_t dynamicSmemBytes, std::array const& buffers, + ColdPointer coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers, cudaStream_t stream) +{ + static_assert(std::is_trivially_copyable_v>); + static_assert(sizeof(std::array) == sizeof(Task) * kMaxTasksPerLaunch); + TLLM_CHECK(count > 0 && count <= kMaxTasksPerLaunch); + + // Zero initialization keeps partial argument padding deterministic. + std::array tail{}; + void const* taskArgument = tasks; + if (count < kMaxTasksPerLaunch) + { + std::copy_n(tasks, count, tail.begin()); + taskArgument = tail.data(); + } + + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = common::getEnvEnablePDL() ? 1 : 0; + + cudaLaunchConfig_t config{}; + config.gridDim = grid; + config.blockDim = block; + config.dynamicSmemBytes = dynamicSmemBytes; + config.stream = stream; + config.attrs = &attribute; + config.numAttrs = 1; + + void* arguments[] = {const_cast(taskArgument), const_cast(buffers.data()), + &coldBase, &coldPageBytes, &numBuffers}; + TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); +} + +// Submit Page descriptors in fixed-capacity chunks. +template +void launchTaskBatches(std::vector const& tasks, LaunchBatch const& launchBatch) +{ + std::size_t offset = 0; + while (offset < tasks.size()) + { + std::uint32_t const count + = static_cast(std::min(tasks.size() - offset, kMaxTasksPerLaunch)); + launchBatch(tasks.data() + offset, count); + offset += count; + } +} + +// One tiled SM100 kernel family covers every runtime dtype, transform, and Page geometry. +// grid.y selects one independent buffer; geometry only changes the tile loop length. + +template +void launchOffloadFrom16Bit(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, std::uint8_t* coldBase, cudaStream_t stream) +{ + dim3 const block(kThreadsPerBlock); + launchTaskBatches(pages, + [&](Nvfp4BoundaryOffloadPageTask const* taskData, std::uint32_t count) + { + dim3 const grid(kHostMemorySplits, plan.numBuffers, count); + launchBoundaryBatch(offloadFrom16BitTiledKernel, taskData, count, grid, block, + compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, + plan.numBuffers, stream); + }); +} + +void launchOffloadFromFp8(std::vector const& pages, Nvfp4BoundaryPreparedPlan const& plan, + std::uint8_t* coldBase, cudaStream_t stream) +{ + dim3 const block(kThreadsPerBlock); + launchTaskBatches(pages, + [&](Nvfp4BoundaryOffloadPageTask const* taskData, std::uint32_t count) + { + dim3 const grid(kHostMemorySplits, plan.numBuffers, count); + launchBoundaryBatch(offloadFromFp8TiledKernel, taskData, count, grid, block, + compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, + plan.numBuffers, stream); + }); +} + +template +void launchOnboard(std::vector const& pages, Nvfp4BoundaryPreparedPlan const& plan, + std::uint8_t const* coldBase, cudaStream_t stream) +{ + dim3 const block(kThreadsPerBlock); + launchTaskBatches(pages, + [&](Nvfp4BoundaryOnboardPageTask const* taskData, std::uint32_t count) + { + dim3 const grid(kHostMemorySplits, plan.numBuffers, count); + launchBoundaryBatch(onboardTiledKernel, taskData, count, grid, block, + compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, + plan.numBuffers, stream); + }); +} + +// Drain earlier chunks after synchronous launch failure before Slots are recycled. +template +void launchAndDrainOnFailure(cudaStream_t stream, Launch const& launch) +{ + try + { + launch(); + } + catch (...) + { + cudaError_t const drainStatus = cudaStreamSynchronize(stream); + if (drainStatus != cudaSuccess) + { + // An asynchronous drain failure leaves Slot ownership unknown; fail-stop. + TLLM_LOG_ERROR("NVFP4 boundary rollback drain failed: %s", cudaGetErrorString(drainStatus)); + std::terminate(); + } + throw; + } +} + +} // namespace + +Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType) +{ + // TODO: Make this codec-private, then remove caller-guaranteed admission checks. + TLLM_CHECK_WITH_INFO(common::isSM100Family(), "NVFP4 boundary kernels require an SM100-family GPU"); + TLLM_CHECK_WITH_INFO(!buffers.empty(), "NVFP4 boundary launch requires at least one buffer"); + TLLM_CHECK_WITH_INFO(buffers.size() <= kMaxBuffersPerLaunch, + "NVFP4 boundary launch supports at most %u local buffers, got %zu", kMaxBuffersPerLaunch, buffers.size()); + TLLM_CHECK_WITH_INFO(coldPageBytes > 0, "Cold Page stride must be positive"); + TLLM_CHECK_WITH_INFO( + coldPageBytes % alignof(uint4) == 0, "Cold Page stride must be aligned to %zu bytes", alignof(uint4)); + + bool useFp8 = false; + switch (runtimeType) + { + case Nvfp4BoundaryRuntimeType::kFloat16: + case Nvfp4BoundaryRuntimeType::kBfloat16: break; + case Nvfp4BoundaryRuntimeType::kFp8E4m3: useFp8 = true; break; + default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); + } + + Nvfp4BoundaryPreparedPlan plan; + plan.numBuffers = static_cast(buffers.size()); + plan.coldPageBytes = coldPageBytes; + plan.runtimeType = runtimeType; + std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); + std::vector intervals; + intervals.reserve(3U * buffers.size()); + for (auto const& buffer : buffers) + { + validateBufferPlan(buffer, coldPageBytes, useFp8, intervals); + if (buffer.transform == Nvfp4BoundaryTransform::kNvfp4) + { + plan.maxTileHalfGroups = std::max(plan.maxTileHalfGroups, compressedTransferHalfGroups(buffer.params)); + } + } + std::sort(intervals.begin(), intervals.end(), + [](ColdInterval const& lhs, ColdInterval const& rhs) + { return lhs.begin < rhs.begin || (lhs.begin == rhs.begin && lhs.end < rhs.end); }); + for (std::size_t index = 1; index < intervals.size(); ++index) + { + TLLM_CHECK_WITH_INFO( + intervals[index - 1U].end <= intervals[index].begin, "Cold record intervals must not overlap"); + } + return plan; +} + +void invokeNvfp4BoundaryOffloadCompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream) +{ + if (pages.empty()) + { + return; + } + TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); + switch (plan.runtimeType) + { + case Nvfp4BoundaryRuntimeType::kFloat16: + launchAndDrainOnFailure( + stream, [&] { launchOffloadFrom16Bit(pages, plan, static_cast(coldBase), stream); }); + break; + case Nvfp4BoundaryRuntimeType::kBfloat16: + launchAndDrainOnFailure(stream, + [&] { launchOffloadFrom16Bit<__nv_bfloat16>(pages, plan, static_cast(coldBase), stream); }); + break; + case Nvfp4BoundaryRuntimeType::kFp8E4m3: + launchAndDrainOnFailure( + stream, [&] { launchOffloadFromFp8(pages, plan, static_cast(coldBase), stream); }); + break; + default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); + } +} + +void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream) +{ + if (pages.empty()) + { + return; + } + TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); + switch (plan.runtimeType) + { + case Nvfp4BoundaryRuntimeType::kFloat16: + launchAndDrainOnFailure( + stream, [&] { launchOnboard(pages, plan, static_cast(coldBase), stream); }); + break; + case Nvfp4BoundaryRuntimeType::kBfloat16: + launchAndDrainOnFailure(stream, + [&] { launchOnboard<__nv_bfloat16>(pages, plan, static_cast(coldBase), stream); }); + break; + case Nvfp4BoundaryRuntimeType::kFp8E4m3: + launchAndDrainOnFailure(stream, + [&] { launchOnboard<__nv_fp8_e4m3>(pages, plan, static_cast(coldBase), stream); }); + break; + default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); + } +} + +} // namespace kernels + +TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h new file mode 100644 index 000000000000..708e5835c2f0 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h @@ -0,0 +1,116 @@ +/* + * 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. + */ + +#pragma once + +#include "tensorrt_llm/common/config.h" + +#include +#include +#include +#include +#include + +TRTLLM_NAMESPACE_BEGIN + +namespace kernels +{ + +//! Active GPU representation at the cold-page boundary. +enum class Nvfp4BoundaryRuntimeType : std::uint8_t +{ + kFloat16, + kBfloat16, + kFp8E4m3, +}; + +//! One Base Page selected for GPU-to-Host transformation. +struct Nvfp4BoundaryOffloadPageTask +{ + std::int32_t gpuPageIndex; + std::int32_t coldPageIndex; +}; + +//! One Base Page selected for Host-to-GPU transformation. +struct Nvfp4BoundaryOnboardPageTask +{ + std::int32_t gpuPageIndex; + std::int32_t coldPageIndex; +}; + +//! Per-buffer geometry and scales for one NVFP4 record in HND order. +//! `headDim` is a multiple of 16; `*OrigQuant` encodes and `*QuantOrig` decodes this buffer. +struct Nvfp4BoundaryKernelParams +{ + std::int32_t numKvHeads; + std::int32_t tokensPerPage; + std::int32_t headDim; + float nvfp4ScaleOrigQuant; + float nvfp4ScaleQuantOrig; + float fp8ScaleOrigQuant; + float fp8ScaleQuantOrig; +}; + +//! Transformation applied to one independently addressed hot buffer. +enum class Nvfp4BoundaryTransform : std::uint8_t +{ + kNvfp4, + kLossless, +}; + +//! Immutable transform plan for one hot buffer and its fixed-offset cold record. +struct Nvfp4BoundaryBufferPlan +{ + std::uintptr_t rawBase; + std::size_t rawSlotBytes; + std::size_t rawBytes; + std::size_t coldDataOffset; + std::size_t coldScaleOffset; + std::size_t coldPaddingOffset; + std::uint32_t coldPaddingBytes; + Nvfp4BoundaryTransform transform; + Nvfp4BoundaryKernelParams params; +}; + +inline constexpr std::uint32_t kNvfp4BoundaryMaxBuffersPerLaunch = 256; + +//! Configure-time launch plan for one Attention lifecycle. +struct Nvfp4BoundaryPreparedPlan +{ + std::array buffers{}; + std::uint32_t numBuffers = 0; + std::uint32_t maxTileHalfGroups = 0; + std::size_t coldPageBytes = 0; + Nvfp4BoundaryRuntimeType runtimeType = Nvfp4BoundaryRuntimeType::kFloat16; +}; + +//! Validate and freeze one lifecycle's boundary-transform plan. +[[nodiscard]] Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType); + +//! Compress GPU Pages into mapped-Host NVFP4 records. +void invokeNvfp4BoundaryOffloadCompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream); + +//! Restore mapped-Host NVFP4 records into GPU Pages. +void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream); + +} // namespace kernels + +TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt new file mode 100644 index 000000000000..33e49bd5d927 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -0,0 +1,9 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# All rights reserved. SPDX-License-Identifier: Apache-2.0 + +add_library(kv_cache_compression_src OBJECT nvfp4ColdPageCodec.cpp) +target_include_directories( + kv_cache_compression_src + PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) +set_property(TARGET kv_cache_compression_src PROPERTY POSITION_INDEPENDENT_CODE + ON) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp new file mode 100644 index 000000000000..5ff5cba23636 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp @@ -0,0 +1,409 @@ +/* + * 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 "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" + +#include "tensorrt_llm/common/logger.h" +#include "tensorrt_llm/common/nvtxUtils.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +constexpr std::size_t kCompactAlignment = 16U; +constexpr std::size_t kElementsPerBlockScale = 16U; +constexpr std::size_t kPackedElementsPerByte = 2U; +constexpr char const* kKeyRole = "key"; +constexpr char const* kValueRole = "value"; + +std::size_t alignUp(std::size_t value, std::size_t alignment) +{ + if (value > std::numeric_limits::max() - (alignment - 1U)) + { + throw std::overflow_error("Cold Page size overflows size_t"); + } + return (value + alignment - 1U) / alignment * alignment; +} + +std::size_t checkedMul(std::size_t lhs, std::size_t rhs, char const* label) +{ + if (rhs != 0U && lhs > std::numeric_limits::max() / rhs) + { + throw std::overflow_error(label); + } + return lhs * rhs; +} + +std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) +{ + if (lhs > std::numeric_limits::max() - rhs) + { + throw std::overflow_error(label); + } + return lhs + rhs; +} + +std::size_t scalarCount(Nvfp4ColdPageLayerConfig const& config) +{ + if (config.numKvHeads <= 0 || config.tokensPerPage <= 0 || config.headDim <= 0) + { + throw std::invalid_argument("NVFP4 cold Page geometry must be positive"); + } + if (config.headDim % static_cast(kElementsPerBlockScale) != 0) + { + throw std::invalid_argument("NVFP4 cold Pages require headDim divisible by 16"); + } + auto const headsTimesTokens = checkedMul(static_cast(config.numKvHeads), + static_cast(config.tokensPerPage), "NVFP4 Page geometry overflows size_t"); + return checkedMul( + headsTimesTokens, static_cast(config.headDim), "NVFP4 Page geometry overflows size_t"); +} + +kernels::Nvfp4BoundaryKernelParams makeKernelParams(Nvfp4ColdPageLayerConfig const& config, std::size_t scaleIndex) +{ + kernels::Nvfp4BoundaryKernelParams params{}; + params.numKvHeads = config.numKvHeads; + params.tokensPerPage = config.tokensPerPage; + params.headDim = config.headDim; + params.nvfp4ScaleOrigQuant = config.nvfp4ScaleOrigQuant[scaleIndex]; + params.nvfp4ScaleQuantOrig = config.nvfp4ScaleQuantOrig[scaleIndex]; + params.fp8ScaleOrigQuant = config.fp8ScaleOrigQuant[scaleIndex]; + params.fp8ScaleQuantOrig = config.fp8ScaleQuantOrig[scaleIndex]; + return params; +} + +std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) +{ + if (offset > std::numeric_limits::max() - base) + { + throw std::overflow_error("GPU buffer address overflows uintptr_t"); + } + return base + offset; +} + +} // namespace + +Nvfp4ColdPageCodec::Nvfp4ColdPageCodec(std::vector layerConfigs) +{ + for (auto& config : layerConfigs) + { + static_cast(scalarCount(config)); + if (config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFloat16 + && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kBfloat16 + && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3) + { + throw std::invalid_argument("Nvfp4ColdPageCodec received an unsupported runtime type"); + } + auto const validScales = [](auto const& scales) { + return std::all_of( + scales.begin(), scales.end(), [](float value) { return std::isfinite(value) && value > 0; }); + }; + if (!validScales(config.nvfp4ScaleOrigQuant) || !validScales(config.nvfp4ScaleQuantOrig) + || (config.runtimeType == kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3 + && (!validScales(config.fp8ScaleOrigQuant) || !validScales(config.fp8ScaleQuantOrig)))) + { + throw std::invalid_argument("NVFP4 runtime scales must be finite and positive"); + } + auto const layerId = config.layerId; + if (!mLayerConfigs.emplace(layerId, std::move(config)).second) + { + throw std::invalid_argument("Nvfp4ColdPageCodec layer IDs must be unique"); + } + } +} + +bool Nvfp4ColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept +{ + try + { + struct BufferLocation + { + kv::PoolIndex poolIndex{0}; + std::size_t offset = 0; + std::size_t bytes = 0; + bool found = false; + }; + + struct AttentionBuffers + { + BufferLocation key; + BufferLocation value; + std::vector losslessBuffers; + }; + + // Use KVCM's default codec for non-Attention lifecycles. + auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); + if (!losslessCodec->configure(gpuDescs, numGpuDescs)) + { + throw std::invalid_argument("Default lossless codec rejected GPU layouts"); + } + + std::map pending; + std::set configuredAttentionLayers; + for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) + { + auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; + for (auto const& variant : gpuDesc.slotDesc.variants) + { + std::map attentionBuffers; + bool hasForeignLayerBuffer = false; + for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) + { + auto const& coalesced = variant.coalescedBuffers.at(poolIndex); + std::size_t offset = 0U; + for (auto const& bufferId : coalesced.bufferIds) + { + auto const config = mLayerConfigs.find(bufferId.layerId); + if (config == mLayerConfigs.end()) + { + hasForeignLayerBuffer = true; + } + else + { + configuredAttentionLayers.insert(bufferId.layerId); + auto& buffers = attentionBuffers[bufferId.layerId]; + if (bufferId.role == kKeyRole || bufferId.role == kValueRole) + { + auto& location = bufferId.role == kKeyRole ? buffers.key : buffers.value; + if (location.found) + { + throw std::invalid_argument("GPU lifecycle contains a duplicate K/V buffer"); + } + location = BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}; + } + else + { + buffers.losslessBuffers.push_back( + BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}); + } + } + offset + = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); + } + } + + LayerGroupState state; + if (!attentionBuffers.empty()) + { + if (hasForeignLayerBuffer) + { + throw std::invalid_argument("Attention lifecycle mixes configured and unconfigured layers"); + } + state.transform = Transform::kNvfp4Attention; + std::vector plans; + auto const runtimeType = mLayerConfigs.at(attentionBuffers.begin()->first).runtimeType; + for (auto const& [layerId, buffers] : attentionBuffers) + { + if (!buffers.key.found) + { + throw std::invalid_argument("Configured Attention layer is missing key"); + } + auto const& config = mLayerConfigs.at(layerId); + if (runtimeType != config.runtimeType) + { + throw std::invalid_argument("Attention lifecycle must use one runtime dtype"); + } + + auto const elements = scalarCount(config); + auto const rawElementBytes + = config.runtimeType == kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3 ? 1U : 2U; + auto const rawBytes + = checkedMul(elements, rawElementBytes, "Runtime KV Page size overflows size_t"); + if (buffers.key.bytes != rawBytes || (buffers.value.found && buffers.value.bytes != rawBytes)) + { + throw std::invalid_argument("GPU Attention buffer size does not match its geometry"); + } + + auto const layerOffset = state.coldPageBytes; + auto const packedBytesPerBuffer = elements / kPackedElementsPerByte; + auto const scaleBytesPerBuffer = elements / kElementsPerBlockScale; + auto const compressedBufferCount = buffers.value.found ? 2U : 1U; + auto const scaleOffset = checkedAdd(layerOffset, + checkedMul( + packedBytesPerBuffer, compressedBufferCount, "NVFP4 packed Page size overflows size_t"), + "NVFP4 cold Page size overflows size_t"); + + auto appendNvfp4 = [&](BufferLocation const& location, std::size_t scaleIndex, + std::size_t coldDataOffset, std::size_t coldScaleOffset) + { + auto const& pool = gpuDesc.pools.at(location.poolIndex); + plans.push_back( + kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), + pool.slotBytes, location.bytes, coldDataOffset, coldScaleOffset, 0U, 0U, + kernels::Nvfp4BoundaryTransform::kNvfp4, makeKernelParams(config, scaleIndex)}); + }; + + appendNvfp4(buffers.key, 0U, layerOffset, scaleOffset); + if (buffers.value.found) + { + appendNvfp4(buffers.value, 1U, + checkedAdd(layerOffset, packedBytesPerBuffer, "NVFP4 cold Page size overflows size_t"), + checkedAdd(scaleOffset, scaleBytesPerBuffer, "NVFP4 cold Page size overflows size_t")); + } + + auto coldOffset = checkedAdd(scaleOffset, + checkedMul( + scaleBytesPerBuffer, compressedBufferCount, "NVFP4 scale Page size overflows size_t"), + "NVFP4 cold Page size overflows size_t"); + for (auto const& location : buffers.losslessBuffers) + { + auto const& pool = gpuDesc.pools.at(location.poolIndex); + plans.push_back(kernels::Nvfp4BoundaryBufferPlan{ + checkedAddress(pool.baseAddress, location.offset), pool.slotBytes, location.bytes, + coldOffset, 0U, 0U, 0U, kernels::Nvfp4BoundaryTransform::kLossless, {}}); + coldOffset = checkedAdd( + coldOffset, location.bytes, "Lossless Attention side-buffer size overflows size_t"); + } + + auto const alignedEnd = alignUp(coldOffset, kCompactAlignment); + auto const paddingBytes = alignedEnd - coldOffset; + plans.back().coldPaddingOffset = coldOffset; + plans.back().coldPaddingBytes = static_cast(paddingBytes); + state.coldPageBytes = alignedEnd; + } + state.preparedPlan = kernels::prepareNvfp4BoundaryPlan(plans, state.coldPageBytes, runtimeType); + } + else + { + state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); + } + pending.emplace(variant.lifeCycleId, std::move(state)); + } + } + if (configuredAttentionLayers.size() != mLayerConfigs.size()) + { + throw std::invalid_argument("A configured Attention layer is absent from all GPU descriptors"); + } + + mLayerGroups = std::move(pending); + mLosslessCodec = std::move(losslessCodec); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("Nvfp4ColdPageCodec::configure rejected GPU layouts: %s", error.what()); + return false; + } +} + +std::size_t Nvfp4ColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept +{ + auto const* state = findLayerGroup(layerGroupId); + return state == nullptr ? 0U : state->coldPageBytes; +} + +kv::LayerGroupId Nvfp4ColdPageCodec::getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept +{ + return findLayerGroup(layerGroupId) == nullptr ? kv::LayerGroupId{-1} : layerGroupId; +} + +kv::PageIndexLocation Nvfp4ColdPageCodec::queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept +{ + return findLayerGroup(layerGroupId) == nullptr ? kv::PageIndexLocation::kBadLocation : kv::PageIndexLocation::kHost; +} + +Nvfp4ColdPageCodec::LayerGroupState const* Nvfp4ColdPageCodec::findLayerGroup( + kv::LayerGroupId layerGroupId) const noexcept +{ + auto const found = mLayerGroups.find(layerGroupId); + return found == mLayerGroups.end() ? nullptr : &found->second; +} + +bool Nvfp4ColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && pageIndices == nullptr)) + { + throw std::invalid_argument("encode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (state->transform == Transform::kLosslessConcat) + { + return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); + } + + NVTX3_SCOPED_RANGE(KVCC_OFFLOAD_COMPRESS_D2H); + thread_local std::vector pages; + pages.clear(); + pages.reserve(numBasePages); + for (std::size_t page = 0; page < numBasePages; ++page) + { + pages.push_back({pageIndices[page].src, pageIndices[page].dst}); + } + kernels::invokeNvfp4BoundaryOffloadCompress(pages, state->preparedPlan, dstBasePtr, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("Nvfp4ColdPageCodec::encode failed before completion fencing: %s", error.what()); + return false; + } +} + +bool Nvfp4ColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, + kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && pageIndices == nullptr)) + { + throw std::invalid_argument("decode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (state->transform == Transform::kLosslessConcat) + { + return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); + } + + NVTX3_SCOPED_RANGE(KVCC_ONBOARD_H2D_DECOMPRESS); + thread_local std::vector pages; + pages.clear(); + pages.reserve(numBasePages); + for (std::size_t page = 0; page < numBasePages; ++page) + { + pages.push_back({pageIndices[page].dst, pageIndices[page].src}); + } + kernels::invokeNvfp4BoundaryOnboardDecompress(pages, state->preparedPlan, srcBasePtr, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("Nvfp4ColdPageCodec::decode failed before completion fencing: %s", error.what()); + return false; + } +} + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h new file mode 100644 index 000000000000..ef129f858207 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h @@ -0,0 +1,90 @@ +/* + * 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. + */ + +#pragma once + +#include "kv_cache_manager_v2/coldPageCodec.h" +#include "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" + +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +//! Per-layer geometry and calibration for NVFP4 cold Pages. +struct Nvfp4ColdPageLayerConfig +{ + kv::LayerId layerId = 0; + kernels::Nvfp4BoundaryRuntimeType runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFloat16; + std::int32_t numKvHeads = 0; + std::int32_t tokensPerPage = 0; + std::int32_t headDim = 0; + std::array nvfp4ScaleOrigQuant{}; + std::array nvfp4ScaleQuantOrig{}; + std::array fp8ScaleOrigQuant{1.0F, 1.0F}; + std::array fp8ScaleQuantOrig{1.0F, 1.0F}; +}; + +//! NVFP4 codec for compact Attention records with lossless side-buffer spans. +class Nvfp4ColdPageCodec final : public kv::IKvCacheColdPageCodec +{ +public: + explicit Nvfp4ColdPageCodec(std::vector layerConfigs); + + bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; + + [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; + + bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + + bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + +private: + enum class Transform + { + kNvfp4Attention, + kLosslessConcat, + }; + + struct LayerGroupState + { + Transform transform = Transform::kLosslessConcat; + kernels::Nvfp4BoundaryPreparedPlan preparedPlan; + std::size_t coldPageBytes = 0; + }; + + [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; + + std::map mLayerConfigs; + std::map mLayerGroups; + std::unique_ptr mLosslessCodec; +}; + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/CMakeLists.txt b/cpp/tensorrt_llm/nanobind/CMakeLists.txt index 5dc1de84a308..982fe0369431 100755 --- a/cpp/tensorrt_llm/nanobind/CMakeLists.txt +++ b/cpp/tensorrt_llm/nanobind/CMakeLists.txt @@ -17,6 +17,7 @@ set(SRCS executor/bindings.cpp executor/executorConfig.cpp executor/request.cpp + kvCacheCompression/bindings.cpp process_group/bindings.cpp runtime/bindings.cpp runtime/hostfunc.cpp diff --git a/cpp/tensorrt_llm/nanobind/bindings.cpp b/cpp/tensorrt_llm/nanobind/bindings.cpp index a2054dbd7217..0371c48d5954 100644 --- a/cpp/tensorrt_llm/nanobind/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/bindings.cpp @@ -44,6 +44,7 @@ #include "tensorrt_llm/nanobind/batch_manager/llmRequest.h" #include "tensorrt_llm/nanobind/common/tllmExceptions.h" #include "tensorrt_llm/nanobind/executor/bindings.h" +#include "tensorrt_llm/nanobind/kvCacheCompression/bindings.h" #include "tensorrt_llm/nanobind/process_group/bindings.h" #include "tensorrt_llm/nanobind/runtime/bindings.h" #include "tensorrt_llm/nanobind/suffixAutomaton/bindings.h" @@ -138,6 +139,8 @@ NB_MODULE(TRTLLM_NB_MODULE, m) = mInternalBatchManager.def_submodule("kv_cache_manager_v2_utils", "KV Cache Manager V2 Utils bindings"); auto mInternalBatchManagerKvCacheV2 = mInternalBatchManager.def_submodule("kv_cache_manager_v2", "KV Cache Manager V2 bindings"); + auto mInternalKvCacheCompression + = mInternal.def_submodule("kv_cache_compression", "KV cache compression internal bindings"); tensorrt_llm::nanobind::batch_manager::KvCacheManagerV2Bindings::initBindings(mInternalBatchManagerKvCacheV2); auto mInternalThop = mInternal.def_submodule("thop", "Torch op internal bindings"); auto mExceptions = m.def_submodule("exceptions", "Exceptions internal bindings"); @@ -145,6 +148,7 @@ NB_MODULE(TRTLLM_NB_MODULE, m) tensorrt_llm::nanobind::executor::initBindings(mExecutor); tensorrt_llm::nanobind::runtime::initBindingsEarly(mInternalRuntime); tensorrt_llm::nanobind::common::initExceptionsBindings(mExceptions); + tensorrt_llm::nanobind::kv_cache_compression::initBindings(mInternalKvCacheCompression); tensorrt_llm::nanobind::thop::initBindings(mInternalThop); auto buildInfo = m.def_submodule("BuildInfo"); diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp new file mode 100644 index 000000000000..34eecd1fd38e --- /dev/null +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -0,0 +1,66 @@ +/* + * 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 "bindings.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" + +#include +#include +#include +#include + +#include +#include +#include + +namespace nb = nanobind; +namespace compression = tensorrt_llm::kv_cache_compression; +namespace kernels = tensorrt_llm::kernels; +namespace kv = tensorrt_llm::batch_manager::kv_cache_manager_v2; + +namespace tensorrt_llm::nanobind::kv_cache_compression +{ + +void initBindings(nb::module_& module) +{ + nb::enum_(module, "Nvfp4BoundaryRuntimeType") + .value("FLOAT16", kernels::Nvfp4BoundaryRuntimeType::kFloat16) + .value("BFLOAT16", kernels::Nvfp4BoundaryRuntimeType::kBfloat16) + .value("FP8_E4M3", kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3); + + nb::class_(module, "Nvfp4ColdPageLayerConfig") + .def(nb::init<>()) + .def_rw("layer_id", &compression::Nvfp4ColdPageLayerConfig::layerId) + .def_rw("runtime_type", &compression::Nvfp4ColdPageLayerConfig::runtimeType) + .def_rw("num_kv_heads", &compression::Nvfp4ColdPageLayerConfig::numKvHeads) + .def_rw("tokens_per_page", &compression::Nvfp4ColdPageLayerConfig::tokensPerPage) + .def_rw("head_dim", &compression::Nvfp4ColdPageLayerConfig::headDim) + .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageLayerConfig::nvfp4ScaleOrigQuant) + .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageLayerConfig::nvfp4ScaleQuantOrig) + .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageLayerConfig::fp8ScaleOrigQuant) + .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageLayerConfig::fp8ScaleQuantOrig); + + // Construct in C++ so ownership can transfer to KVCM as a unique_ptr codec. + module.def( + "create_nvfp4_cold_page_codec", + [](std::vector layerConfigs) + -> std::unique_ptr + { return std::make_unique(std::move(layerConfigs)); }, + nb::arg("layer_configs"), "Create an owning NVFP4 cold-page codec for one-time transfer into KVCacheManager."); +} + +} // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.h b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.h new file mode 100644 index 000000000000..f85d098097e8 --- /dev/null +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.h @@ -0,0 +1,27 @@ +/* + * 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. + */ + +#pragma once + +#include + +namespace tensorrt_llm::nanobind::kv_cache_compression +{ + +void initBindings(::nanobind::module_& module); + +} // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tests/unit_tests/CMakeLists.txt b/cpp/tests/unit_tests/CMakeLists.txt index 9d22bd03b52a..531bff2205cc 100644 --- a/cpp/tests/unit_tests/CMakeLists.txt +++ b/cpp/tests/unit_tests/CMakeLists.txt @@ -22,6 +22,7 @@ if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/executor) endif() add_subdirectory(common) +add_subdirectory(kv_cache_compression) add_subdirectory(kernels) add_subdirectory(multi_gpu) add_subdirectory(layers) diff --git a/cpp/tests/unit_tests/kernels/CMakeLists.txt b/cpp/tests/unit_tests/kernels/CMakeLists.txt index 6b0e5a118211..643c98dae575 100644 --- a/cpp/tests/unit_tests/kernels/CMakeLists.txt +++ b/cpp/tests/unit_tests/kernels/CMakeLists.txt @@ -46,6 +46,22 @@ if(USING_OSS_CUTLASS_MOE_GEMM) endif() add_gtest(ropeTest ropeTest.cu) +set(NVFP4_BOUNDARY_KERNEL_TEST_SRC + nvfp4BoundaryKernelsTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) +add_gtest(nvfp4BoundaryKernelsTest "${NVFP4_BOUNDARY_KERNEL_TEST_SRC}" + NO_TLLM_LINKAGE) +target_include_directories( + nvfp4BoundaryKernelsTest + PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) +target_link_libraries(nvfp4BoundaryKernelsTest PRIVATE CUDA::cudart + CUDA::cuda_driver) add_gtest(shiftKCacheKernelTest shiftKCacheKernelTest.cu) add_gtest(smoothQuantKernelTest smoothQuant/smoothQuantKernelTest.cpp) add_gtest(stopCriteriaKernelsTest stopCriteriaKernelsTest.cpp) diff --git a/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp new file mode 100644 index 000000000000..a2130e87cd81 --- /dev/null +++ b/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp @@ -0,0 +1,1310 @@ +/* + * 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 "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" +#include "tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.h" +#include "tensorrt_llm/common/cudaUtils.h" + +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ + +using tensorrt_llm::batch_manager::kv_cache_manager_v2::HostMem; +using tensorrt_llm::batch_manager::kv_cache_manager_v2::MemAddress; +using tensorrt_llm::kernels::Nvfp4BoundaryBufferPlan; +using tensorrt_llm::kernels::Nvfp4BoundaryKernelParams; +using tensorrt_llm::kernels::Nvfp4BoundaryOffloadPageTask; +using tensorrt_llm::kernels::Nvfp4BoundaryOnboardPageTask; +using tensorrt_llm::kernels::Nvfp4BoundaryPreparedPlan; +using tensorrt_llm::kernels::Nvfp4BoundaryRuntimeType; +using tensorrt_llm::kernels::Nvfp4BoundaryTransform; + +constexpr std::size_t kGuardBytes = 64; +constexpr std::uint8_t kCanary = 0xA5; +constexpr std::size_t kDefaultNumPages = 3; +constexpr std::size_t kCrossLaunchNumPages = 257; + +struct PageGeometry +{ + std::int32_t numHeads; + std::int32_t tokensPerPage; + std::int32_t headDim; +}; + +constexpr PageGeometry kDefaultGeometry{2, 8, 32}; +constexpr PageGeometry kMinimumCompactGeometry{1, 1, 16}; +constexpr PageGeometry kPackedBodyAndTailGeometry{1, 3, 16}; +constexpr PageGeometry kSmallVectorGeometry{1, 4, 16}; +constexpr PageGeometry kLinearScaleTailGeometry{1, 5, 32}; +constexpr PageGeometry kTiledLinearScaleTailGeometry{1, 4097, 16}; +constexpr PageGeometry kCrossRowTileGeometry{1, 343, 48}; +constexpr PageGeometry kLargeHeadDimTailGeometry{1, 1, 65552}; +constexpr PageGeometry kModelLikeGeometry{8, 64, 128}; +constexpr std::array kValidTokenCounts{1, 16, 17, 63, 64}; + +enum class RawKind +{ + kFloat16, + kBfloat16, + kFp8, +}; + +enum class InputPattern +{ + kDense, + kAllZero, + kSparseOutlier, + kRoundingMargins, +}; + +std::size_t roundUp(std::size_t value, std::size_t alignment) +{ + return (value + alignment - 1) / alignment * alignment; +} + +class CudaStream +{ +public: + CudaStream() + { + TLLM_CUDA_CHECK(cudaStreamCreateWithFlags(&mStream, cudaStreamNonBlocking)); + } + + ~CudaStream() + { + if (mStream != nullptr) + { + cudaStreamDestroy(mStream); + } + } + + operator cudaStream_t() const + { + return mStream; + } + +private: + cudaStream_t mStream{}; +}; + +//! Device allocation guarded by canaries to catch descriptor or vector-tail out-of-bounds writes. +class DeviceRegion +{ +public: + explicit DeviceRegion(std::size_t payloadBytes) + : mPayloadBytes(payloadBytes) + , mTotalBytes(payloadBytes + 2 * kGuardBytes) + { + TLLM_CUDA_CHECK(cudaMalloc(&mBase, mTotalBytes)); + TLLM_CUDA_CHECK(cudaMemset(mBase, kCanary, mTotalBytes)); + } + + ~DeviceRegion() + { + if (mBase != nullptr) + { + cudaFree(mBase); + } + } + + DeviceRegion(DeviceRegion const&) = delete; + DeviceRegion& operator=(DeviceRegion const&) = delete; + + void* data() const + { + return static_cast(mBase) + kGuardBytes; + } + + void copyFrom(std::vector const& bytes) + { + ASSERT_EQ(bytes.size(), mPayloadBytes); + ASSERT_EQ(cudaMemcpy(data(), bytes.data(), bytes.size(), cudaMemcpyHostToDevice), cudaSuccess); + } + + void copyFrom(std::size_t offset, std::vector const& bytes) + { + ASSERT_LE(offset + bytes.size(), mPayloadBytes); + ASSERT_EQ( + cudaMemcpy(static_cast(data()) + offset, bytes.data(), bytes.size(), cudaMemcpyHostToDevice), + cudaSuccess); + } + + std::vector copyToHost() const + { + std::vector bytes(mPayloadBytes); + EXPECT_EQ(cudaMemcpy(bytes.data(), data(), bytes.size(), cudaMemcpyDeviceToHost), cudaSuccess); + return bytes; + } + + std::vector copyToHost(std::size_t offset, std::size_t bytes) const + { + EXPECT_LE(offset + bytes, mPayloadBytes); + std::vector result(bytes); + EXPECT_EQ( + cudaMemcpy(result.data(), static_cast(data()) + offset, bytes, cudaMemcpyDeviceToHost), + cudaSuccess); + return result; + } + + void expectCanaries() const + { + std::vector bytes(mTotalBytes); + ASSERT_EQ(cudaMemcpy(bytes.data(), mBase, bytes.size(), cudaMemcpyDeviceToHost), cudaSuccess); + EXPECT_TRUE(std::all_of( + bytes.begin(), bytes.begin() + kGuardBytes, [](std::uint8_t value) { return value == kCanary; })); + EXPECT_TRUE( + std::all_of(bytes.end() - kGuardBytes, bytes.end(), [](std::uint8_t value) { return value == kCanary; })); + } + +private: + void* mBase{}; + std::size_t mPayloadBytes{}; + std::size_t mTotalBytes{}; +}; + +//! CUDA-mapped HostMem matching KVCM V2's Host carrier. +class MappedHostRegion +{ +public: + explicit MappedHostRegion(std::size_t payloadBytes) + : mMemory(roundUp(kGuardBytes + payloadBytes + kGuardBytes, HostMem::kAlignment)) + , mPayloadBytes(payloadBytes) + { + TLLM_CHECK_WITH_INFO(kGuardBytes + payloadBytes + kGuardBytes <= mMemory.size(), + "Mapped Host test allocation is too small for payload and canaries"); + std::memset(reinterpret_cast(mMemory.address()), kCanary, mMemory.size()); + } + + void* data() const + { + return reinterpret_cast(mMemory.address() + kGuardBytes); + } + + std::uint8_t* bytes() const + { + return static_cast(data()); + } + + std::vector payload() const + { + return {bytes(), bytes() + mPayloadBytes}; + } + + void expectCanaries() const + { + auto const* base = reinterpret_cast(mMemory.address()); + EXPECT_TRUE(std::all_of(base, base + kGuardBytes, [](std::uint8_t value) { return value == kCanary; })); + EXPECT_TRUE(std::all_of(base + kGuardBytes + mPayloadBytes, base + mMemory.size(), + [](std::uint8_t value) { return value == kCanary; })); + } + +private: + HostMem mMemory; + std::size_t mPayloadBytes{}; +}; + +struct LayerBuffers +{ + explicit LayerBuffers(std::size_t rawBytes) + : rawK(rawBytes) + , rawV(rawBytes) + { + } + + DeviceRegion rawK; + DeviceRegion rawV; +}; + +std::size_t numElements(PageGeometry const& geometry) +{ + return static_cast(geometry.numHeads) * geometry.tokensPerPage * geometry.headDim; +} + +std::size_t rawBytes(RawKind kind, PageGeometry const& geometry) +{ + return numElements(geometry) * (kind == RawKind::kFp8 ? 1 : 2); +} + +std::size_t rawElementBytes(RawKind kind) +{ + return kind == RawKind::kFp8 ? 1U : 2U; +} + +std::size_t packedBytes(PageGeometry const& geometry) +{ + return numElements(geometry) / 2; +} + +std::size_t scaleBytes(PageGeometry const& geometry) +{ + return numElements(geometry) / 16; +} + +Nvfp4BoundaryKernelParams makeParams(PageGeometry const& geometry = kDefaultGeometry, std::uint32_t role = 0U) +{ + Nvfp4BoundaryKernelParams params{}; + params.numKvHeads = geometry.numHeads; + params.tokensPerPage = geometry.tokensPerPage; + params.headDim = geometry.headDim; + params.nvfp4ScaleOrigQuant = role == 0U ? 1.0F : 2.0F; + params.nvfp4ScaleQuantOrig = role == 0U ? 1.0F : 0.5F; + params.fp8ScaleOrigQuant = role == 0U ? 2.0F : 4.0F; + params.fp8ScaleQuantOrig = role == 0U ? 0.5F : 0.25F; + return params; +} + +Nvfp4BoundaryRuntimeType runtimeType(RawKind kind) +{ + switch (kind) + { + case RawKind::kFloat16: return Nvfp4BoundaryRuntimeType::kFloat16; + case RawKind::kBfloat16: return Nvfp4BoundaryRuntimeType::kBfloat16; + case RawKind::kFp8: return Nvfp4BoundaryRuntimeType::kFp8E4m3; + } + return Nvfp4BoundaryRuntimeType::kFloat16; +} + +template +void storeScalar(std::vector& bytes, std::size_t index, T value) +{ + std::memcpy(bytes.data() + index * sizeof(T), &value, sizeof(T)); +} + +template +T loadScalar(std::vector const& bytes, std::size_t index) +{ + T value; + std::memcpy(&value, bytes.data() + index * sizeof(T), sizeof(T)); + return value; +} + +void storeRawValue(std::vector& bytes, RawKind kind, std::size_t index, float value, + Nvfp4BoundaryKernelParams const& params) +{ + switch (kind) + { + case RawKind::kFloat16: storeScalar(bytes, index, __float2half(value)); break; + case RawKind::kBfloat16: storeScalar(bytes, index, __float2bfloat16(value)); break; + case RawKind::kFp8: storeScalar(bytes, index, __nv_fp8_e4m3(value * params.fp8ScaleOrigQuant)); break; + } +} + +float loadRawValue( + std::vector const& bytes, RawKind kind, std::size_t index, Nvfp4BoundaryKernelParams const& params) +{ + switch (kind) + { + case RawKind::kFloat16: return __half2float(loadScalar(bytes, index)); + case RawKind::kBfloat16: return __bfloat162float(loadScalar<__nv_bfloat16>(bytes, index)); + case RawKind::kFp8: return static_cast(loadScalar<__nv_fp8_e4m3>(bytes, index)) * params.fp8ScaleQuantOrig; + } + return 0.0F; +} + +std::uint32_t linearScaleOffset(std::uint32_t row, std::uint32_t scaleInRow, PageGeometry const& geometry) +{ + std::uint32_t const scalesPerRow = static_cast(geometry.headDim) / 16; + return row * scalesPerRow + scaleInRow; +} + +float e2m1Value(std::uint8_t nibble) +{ + constexpr std::array levels{0.0F, 0.5F, 1.0F, 1.5F, 2.0F, 3.0F, 4.0F, 6.0F}; + float const value = levels[nibble & 0x7U]; + return (nibble & 0x8U) != 0 ? -value : value; +} + +//! Independent nearest-level oracle; fixtures avoid ties instead of duplicating production tie rules. +std::uint8_t quantizeE2m1(float value) +{ + constexpr std::array levels{0.0F, 0.5F, 1.0F, 1.5F, 2.0F, 3.0F, 4.0F, 6.0F}; + bool const negative = std::signbit(value); + float const magnitude = std::abs(value); + std::uint8_t best = 0; + float bestDistance = std::abs(magnitude - levels[0]); + for (std::uint8_t index = 1; index < levels.size(); ++index) + { + float const distance = std::abs(magnitude - levels[index]); + if (distance < bestDistance) + { + best = index; + bestDistance = distance; + } + } + return static_cast(best | (negative ? 0x8U : 0U)); +} + +//! Exactly representable E2M1 values and E4M3 scales keep byte comparisons deterministic. +std::vector makeRawPage(RawKind kind, std::size_t page, std::uint32_t role, + Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry, InputPattern inputPattern) +{ + constexpr std::array densePattern{ + 0.0F, 0.5F, -1.0F, 1.5F, -2.0F, 3.0F, -4.0F, 6.0F, -0.5F, 1.0F, -1.5F, 2.0F, -3.0F, 4.0F, -6.0F, 0.5F}; + constexpr std::array firstLaneOutlierPattern{ + 6.0F, -0.5F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.5F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F}; + constexpr std::array secondLaneOutlierPattern{ + 0.0F, -0.5F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.5F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, -6.0F}; + constexpr std::array roundingMarginsPattern{ + 6.0F, 0.20F, 0.30F, 0.65F, 0.85F, 1.10F, 1.40F, 1.60F, 1.90F, 2.30F, 2.70F, 3.20F, 3.80F, 4.50F, 5.20F, -0.30F}; + constexpr std::array blockScales{0.25F, 0.5F, 1.0F, 2.0F}; + + std::vector bytes(rawBytes(kind, geometry)); + for (std::size_t index = 0; index < numElements(geometry); ++index) + { + std::size_t const scaleGroup = index / 16; + float const blockScale = blockScales[(scaleGroup + page + role) % blockScales.size()]; + float normalizedValue = 0.0F; + if (inputPattern == InputPattern::kDense) + { + normalizedValue = densePattern[index % densePattern.size()]; + } + else if (inputPattern == InputPattern::kSparseOutlier) + { + auto const& pattern = (scaleGroup & 1U) == 0 ? firstLaneOutlierPattern : secondLaneOutlierPattern; + normalizedValue = pattern[index % pattern.size()]; + } + else if (inputPattern == InputPattern::kRoundingMargins) + { + normalizedValue = roundingMarginsPattern[index % roundingMarginsPattern.size()]; + } + if (((page / 32) & 1U) != 0) + { + normalizedValue = -normalizedValue; + } + float const value = normalizedValue * blockScale / params.nvfp4ScaleOrigQuant; + storeRawValue(bytes, kind, index, value, params); + } + return bytes; +} + +struct ReferenceNvfp4 +{ + std::vector packed; + std::vector scales; +}; + +ReferenceNvfp4 compressReference(std::vector const& raw, RawKind kind, + Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) +{ + ReferenceNvfp4 result{{}, {}}; + result.packed.resize(packedBytes(geometry)); + result.scales.resize(scaleBytes(geometry)); + std::uint32_t const scalesPerRow = static_cast(geometry.headDim) / 16; + std::uint32_t const rows = static_cast(geometry.numHeads * geometry.tokensPerPage); + + for (std::uint32_t row = 0; row < rows; ++row) + { + for (std::uint32_t scaleInRow = 0; scaleInRow < scalesPerRow; ++scaleInRow) + { + std::size_t const blockStart = static_cast(row) * geometry.headDim + scaleInRow * 16; + float amax = 0.0F; + for (std::uint32_t i = 0; i < 16; ++i) + { + amax = std::max(amax, std::abs(loadRawValue(raw, kind, blockStart + i, params))); + } + + __nv_fp8_e4m3 blockScale(params.nvfp4ScaleOrigQuant * amax / 6.0F); + result.scales[linearScaleOffset(row, scaleInRow, geometry)] = blockScale.__x; + float const blockScaleFloat = static_cast(blockScale); + float const outputScale = blockScaleFloat == 0.0F ? 0.0F : params.nvfp4ScaleOrigQuant / blockScaleFloat; + for (std::uint32_t i = 0; i < 16; i += 2) + { + std::uint8_t const lo = quantizeE2m1(loadRawValue(raw, kind, blockStart + i, params) * outputScale); + std::uint8_t const hi = quantizeE2m1(loadRawValue(raw, kind, blockStart + i + 1, params) * outputScale); + result.packed[(blockStart + i) / 2] = static_cast(lo | (hi << 4)); + } + } + } + return result; +} + +std::vector decompressReference(ReferenceNvfp4 const& compressed, RawKind kind, + Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) +{ + std::vector raw(rawBytes(kind, geometry)); + std::uint32_t const scalesPerRow = static_cast(geometry.headDim) / 16; + std::uint32_t const rows = static_cast(geometry.numHeads * geometry.tokensPerPage); + for (std::uint32_t row = 0; row < rows; ++row) + { + for (std::uint32_t scaleInRow = 0; scaleInRow < scalesPerRow; ++scaleInRow) + { + __nv_fp8_e4m3 blockScale; + blockScale.__x = compressed.scales[linearScaleOffset(row, scaleInRow, geometry)]; + float const dequantScale = static_cast(blockScale) * params.nvfp4ScaleQuantOrig; + std::size_t const blockStart = static_cast(row) * geometry.headDim + scaleInRow * 16; + for (std::uint32_t i = 0; i < 16; ++i) + { + std::uint8_t const byte = compressed.packed[(blockStart + i) / 2]; + std::uint8_t const nibble = (i & 1U) == 0 ? byte & 0xFU : byte >> 4; + storeRawValue(raw, kind, blockStart + i, e2m1Value(nibble) * dequantScale, params); + } + } + } + return raw; +} + +void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultGeometry, + std::size_t numPages = kDefaultNumPages, InputPattern inputPattern = InputPattern::kDense, + bool synchronizeBetweenDirections = true, bool repeatRoundTrip = false, std::size_t coldBaseOffset = 0) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; + CudaStream stream; + std::size_t const rawSlotBytes = rawBytes(kind, geometry); + // Align compact Slot strides while independently testing an arbitrary staging-base offset. + std::size_t const compactSlotBytes = roundUp(2U * (packedBytes(geometry) + scaleBytes(geometry)), alignof(uint4)); + // Use alternate Slots to cover non-contiguous KVCM Page indices. + std::size_t const slotCapacity = 2U * numPages; + DeviceRegion rawInputK(slotCapacity * rawSlotBytes); + DeviceRegion rawInputV(slotCapacity * rawSlotBytes); + DeviceRegion rawOutputK(slotCapacity * rawSlotBytes); + DeviceRegion rawOutputV(slotCapacity * rawSlotBytes); + MappedHostRegion compactPages(coldBaseOffset + slotCapacity * compactSlotBytes); + auto* compactBase = compactPages.bytes() + coldBaseOffset; + std::vector, 2>> rawHost(numPages); + std::vector offloadTasks; + offloadTasks.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + std::size_t const slot = 2U * page; + rawHost[page][0] = makeRawPage(kind, page, 0, params[0], geometry, inputPattern); + rawHost[page][1] = makeRawPage(kind, page, 1, params[1], geometry, inputPattern); + rawInputK.copyFrom(slot * rawSlotBytes, rawHost[page][0]); + rawInputV.copyFrom(slot * rawSlotBytes, rawHost[page][1]); + offloadTasks.push_back({static_cast(slot), static_cast(slot)}); + } + + std::size_t const packed = packedBytes(geometry); + std::size_t const scale = scaleBytes(geometry); + std::size_t const payloadBytes = 2U * (packed + scale); + std::uint32_t const paddingBytes = static_cast(compactSlotBytes - payloadBytes); + std::vector const inputBuffers{ + {reinterpret_cast(rawInputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, + Nvfp4BoundaryTransform::kNvfp4, params[0]}, + {reinterpret_cast(rawInputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, + payloadBytes, paddingBytes, Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + + auto const inputPlan + = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan(inputBuffers, compactSlotBytes, runtimeType(kind)); + std::vector> references(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + references[page][0] = compressReference(rawHost[page][0], kind, params[0], geometry); + references[page][1] = compressReference(rawHost[page][1], kind, params[1], geometry); + } + + auto const verifyCompressedPages = [&] + { + auto const payload = compactPages.payload(); + EXPECT_TRUE(std::all_of(payload.begin(), payload.begin() + static_cast(coldBaseOffset), + [](std::uint8_t value) { return value == kCanary; })); + for (std::size_t page = 0; page < numPages; ++page) + { + std::size_t const base = coldBaseOffset + 2U * page * compactSlotBytes; + auto const region = [&](std::size_t offset, std::size_t bytes) + { + return std::vector(payload.begin() + static_cast(base + offset), + payload.begin() + static_cast(base + offset + bytes)); + }; + EXPECT_EQ(region(0, packed), references[page][0].packed); + EXPECT_EQ(region(packed, packed), references[page][1].packed); + EXPECT_EQ(region(2 * packed, scale), references[page][0].scales); + EXPECT_EQ(region(2 * packed + scale, scale), references[page][1].scales); + auto const padding = region(2U * (packed + scale), compactSlotBytes - 2U * (packed + scale)); + EXPECT_TRUE(std::all_of(padding.begin(), padding.end(), [](std::uint8_t value) { return value == 0U; })); + + std::size_t const unusedBase = coldBaseOffset + (2U * page + 1U) * compactSlotBytes; + EXPECT_TRUE(std::all_of(payload.begin() + static_cast(unusedBase), + payload.begin() + static_cast(unusedBase + compactSlotBytes), + [](std::uint8_t value) { return value == kCanary; })); + } + }; + + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactBase, stream); + + if (synchronizeBetweenDirections) + { + // Read the Host Slot only after StorageManager-style event fencing. + cudaEvent_t offloadComplete{}; + ASSERT_EQ(cudaEventCreateWithFlags(&offloadComplete, cudaEventDisableTiming), cudaSuccess); + ASSERT_EQ(cudaEventRecord(offloadComplete, stream), cudaSuccess); + ASSERT_EQ(cudaEventSynchronize(offloadComplete), cudaSuccess); + ASSERT_EQ(cudaEventDestroy(offloadComplete), cudaSuccess); + verifyCompressedPages(); + + std::size_t const compactPayloadBytes = 2U * (packedBytes(geometry) + scaleBytes(geometry)); + if (compactPayloadBytes != compactSlotBytes) + { + // Re-encode poisoned recycled Slots to verify deterministic payload and padding bytes. + auto const firstSerialization = compactPages.payload(); + for (std::size_t page = 0; page < numPages; ++page) + { + std::memset(compactBase + 2U * page * compactSlotBytes, 0x5A, compactSlotBytes); + } + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactBase, stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + EXPECT_EQ(compactPages.payload(), firstSerialization); + verifyCompressedPages(); + } + } + + std::vector onboardTasks; + onboardTasks.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + std::size_t const slot = 2U * page; + onboardTasks.push_back({static_cast(slot), static_cast(slot)}); + } + std::vector const outputBuffers{ + {reinterpret_cast(rawOutputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, + Nvfp4BoundaryTransform::kNvfp4, params[0]}, + {reinterpret_cast(rawOutputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, + payloadBytes, paddingBytes, Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + auto const outputPlan + = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan(outputBuffers, compactSlotBytes, runtimeType(kind)); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, compactBase, stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + + if (!synchronizeBetweenDirections) + { + // Verify back-to-back offload/onboard without an intervening Host fence. + verifyCompressedPages(); + } + + if (repeatRoundTrip) + { + // A second lossy round trip catches stale descriptors and validates Q(D(Q(D(Q(x))))). + for (std::size_t page = 0; page < numPages; ++page) + { + for (std::uint32_t role = 0; role < 2; ++role) + { + auto const restored = decompressReference(references[page][role], kind, params[role], geometry); + references[page][role] = compressReference(restored, kind, params[role], geometry); + } + } + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, outputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, inputPlan, compactBase, stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + verifyCompressedPages(); + } + + for (std::size_t page = 0; page < numPages; ++page) + { + std::size_t const slotOffset = 2U * page * rawSlotBytes; + auto const& finalK = repeatRoundTrip ? rawInputK : rawOutputK; + auto const& finalV = repeatRoundTrip ? rawInputV : rawOutputV; + EXPECT_EQ(finalK.copyToHost(slotOffset, rawSlotBytes), + decompressReference(references[page][0], kind, params[0], geometry)); + EXPECT_EQ(finalV.copyToHost(slotOffset, rawSlotBytes), + decompressReference(references[page][1], kind, params[1], geometry)); + } + rawInputK.expectCanaries(); + rawInputV.expectCanaries(); + rawOutputK.expectCanaries(); + rawOutputV.expectCanaries(); + compactPages.expectCanaries(); +} + +std::vector makePartialRawPage(RawKind kind, std::int32_t validTokens, bool zeroTail, std::uint32_t role, + Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) +{ + std::vector bytes(rawBytes(kind, geometry)); + for (std::int32_t head = 0; head < geometry.numHeads; ++head) + { + for (std::int32_t token = 0; token < geometry.tokensPerPage; ++token) + { + for (std::int32_t dim = 0; dim < geometry.headDim; ++dim) + { + std::size_t const index + = (static_cast(head) * geometry.tokensPerPage + token) * geometry.headDim + dim; + float value = 0.0F; + if (token < validTokens) + { + // Zero-tail and stale-tail fixtures share valid rows with distinct K/V values and scales. + value = static_cast((dim % 13) - 6) * 0.125F + static_cast(head) * 0.03125F + + static_cast(token) * 0.0078125F + static_cast(role) * 0.0625F; + } + else if (!zeroTail) + { + // Poison inactive rows; 16-value groups stay within a token row and cannot affect the prefix. + if (dim % 31 == 0) + { + value = std::numeric_limits::quiet_NaN(); + } + else if (dim % 29 == 0) + { + value = std::numeric_limits::infinity(); + } + else + { + value = static_cast((dim % 9) - 4) * 0.25F + static_cast(token) * 0.015625F + + static_cast(head + 3 * role) * 0.046875F; + } + } + storeRawValue(bytes, kind, index, value, params); + } + } + } + return bytes; +} + +void expectSameValidPrefix(std::vector const& lhs, std::vector const& rhs, RawKind kind, + std::int32_t validTokens, PageGeometry const& geometry) +{ + std::size_t const rowBytes = static_cast(geometry.headDim) * rawElementBytes(kind); + for (std::int32_t head = 0; head < geometry.numHeads; ++head) + { + for (std::int32_t token = 0; token < validTokens; ++token) + { + std::size_t const offset = (static_cast(head) * geometry.tokensPerPage + token) * rowBytes; + EXPECT_EQ(std::memcmp(lhs.data() + offset, rhs.data() + offset, rowBytes), 0) + << "valid prefix differs at head=" << head << " token=" << token; + } + } +} + +void expectZeroTail( + std::vector const& bytes, RawKind kind, std::int32_t validTokens, PageGeometry const& geometry) +{ + std::size_t const rowBytes = static_cast(geometry.headDim) * rawElementBytes(kind); + std::vector const zero(rowBytes, 0); + for (std::int32_t head = 0; head < geometry.numHeads; ++head) + { + for (std::int32_t token = validTokens; token < geometry.tokensPerPage; ++token) + { + std::size_t const offset = (static_cast(head) * geometry.tokensPerPage + token) * rowBytes; + EXPECT_EQ(std::memcmp(bytes.data() + offset, zero.data(), rowBytes), 0) + << "zero tail changed at head=" << head << " token=" << token; + } + } +} + +void runPartialPageTailIsolation(RawKind kind) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + + PageGeometry constexpr geometry = kModelLikeGeometry; + std::size_t constexpr pageVariants = 2U; + std::size_t const numPages = pageVariants * kValidTokenCounts.size(); + std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; + std::size_t const rawSlotBytes = rawBytes(kind, geometry); + std::size_t const compactSlotBytes = roundUp(2U * (packedBytes(geometry) + scaleBytes(geometry)), alignof(uint4)); + + DeviceRegion rawInputK(numPages * rawSlotBytes); + DeviceRegion rawInputV(numPages * rawSlotBytes); + DeviceRegion rawOutputK(numPages * rawSlotBytes); + DeviceRegion rawOutputV(numPages * rawSlotBytes); + MappedHostRegion compactPages(numPages * compactSlotBytes); + std::vector offloadTasks; + std::vector onboardTasks; + offloadTasks.reserve(numPages); + onboardTasks.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + std::int32_t const validTokens = kValidTokenCounts[page / pageVariants]; + bool const zeroTail = page % pageVariants == 0U; + rawInputK.copyFrom( + page * rawSlotBytes, makePartialRawPage(kind, validTokens, zeroTail, 0, params[0], geometry)); + rawInputV.copyFrom( + page * rawSlotBytes, makePartialRawPage(kind, validTokens, zeroTail, 1, params[1], geometry)); + auto const pageIndex = static_cast(page); + offloadTasks.push_back({pageIndex, pageIndex}); + onboardTasks.push_back({pageIndex, pageIndex}); + } + + std::size_t const packed = packedBytes(geometry); + std::size_t const scale = scaleBytes(geometry); + std::size_t const payloadBytes = 2U * (packed + scale); + auto const makeBuffers = [&](DeviceRegion const& rawK, DeviceRegion const& rawV) + { + return std::vector{ + {reinterpret_cast(rawK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, + Nvfp4BoundaryTransform::kNvfp4, params[0]}, + {reinterpret_cast(rawV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, + payloadBytes, static_cast(compactSlotBytes - payloadBytes), + Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + }; + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + makeBuffers(rawInputK, rawInputV), compactSlotBytes, runtimeType(kind)); + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + makeBuffers(rawOutputK, rawOutputV), compactSlotBytes, runtimeType(kind)); + CudaStream stream; + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactPages.data(), stream); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, compactPages.data(), stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + + for (std::size_t pair = 0; pair < kValidTokenCounts.size(); ++pair) + { + std::int32_t const validTokens = kValidTokenCounts[pair]; + for (std::uint32_t role = 0; role < 2U; ++role) + { + auto const& output = role == 0U ? rawOutputK : rawOutputV; + auto const zeroOutput = output.copyToHost(pageVariants * pair * rawSlotBytes, rawSlotBytes); + auto const staleOutput = output.copyToHost((pageVariants * pair + 1U) * rawSlotBytes, rawSlotBytes); + expectSameValidPrefix(zeroOutput, staleOutput, kind, validTokens, geometry); + expectZeroTail(zeroOutput, kind, validTokens, geometry); + } + } + + rawInputK.expectCanaries(); + rawInputV.expectCanaries(); + rawOutputK.expectCanaries(); + rawOutputV.expectCanaries(); + compactPages.expectCanaries(); +} + +struct RoundTripCase +{ + char const* name; + RawKind kind; + PageGeometry geometry{kDefaultGeometry}; + std::size_t numPages{kDefaultNumPages}; + InputPattern inputPattern{InputPattern::kDense}; + bool synchronizeBetweenDirections{true}; + bool repeatRoundTrip{false}; + std::size_t coldBaseOffset{0}; +}; + +RoundTripCase constexpr kRoundTripCases[]{ + {"DefaultFloat16", RawKind::kFloat16}, + {"DefaultBfloat16", RawKind::kBfloat16}, + {"DefaultFp8IndependentScales", RawKind::kFp8}, + {"SmallVectorFloat16", RawKind::kFloat16, kSmallVectorGeometry, 1}, + {"MinimumFloat16", RawKind::kFloat16, kMinimumCompactGeometry, 1}, + {"MinimumBfloat16", RawKind::kBfloat16, kMinimumCompactGeometry, 1}, + {"MinimumFp8", RawKind::kFp8, kMinimumCompactGeometry, 1}, + {"PackedScaleTailFloat16", RawKind::kFloat16, kPackedBodyAndTailGeometry, 1}, + {"PackedScaleTailFp8", RawKind::kFp8, kPackedBodyAndTailGeometry, 1}, + {"ByteAlignedColdBaseBfloat16", RawKind::kBfloat16, kPackedBodyAndTailGeometry, 1, InputPattern::kDense, true, + false, 1}, + {"ByteAlignedColdBaseFp8", RawKind::kFp8, kPackedBodyAndTailGeometry, 1, InputPattern::kDense, true, false, 1}, + {"OddTokenFloat16", RawKind::kFloat16, kLinearScaleTailGeometry, 1}, + {"OddTokenBfloat16", RawKind::kBfloat16, kLinearScaleTailGeometry, 1}, + {"OddTokenFp8", RawKind::kFp8, kLinearScaleTailGeometry, 1}, + {"TiledTailsFloat16", RawKind::kFloat16, kTiledLinearScaleTailGeometry, 1}, + {"TiledTailsFp8", RawKind::kFp8, kTiledLinearScaleTailGeometry, 1}, + {"CrossRowTileBfloat16", RawKind::kBfloat16, kCrossRowTileGeometry, 1, InputPattern::kDense, true, false, 1}, + {"LargeHeadDimFp8", RawKind::kFp8, kLargeHeadDimTailGeometry, 1}, + {"ModelLikeBfloat16", RawKind::kBfloat16, kModelLikeGeometry, 1}, + {"ModelLikeFp8", RawKind::kFp8, kModelLikeGeometry, 1}, + {"ZeroGroupsFloat16", RawKind::kFloat16, kDefaultGeometry, 1, InputPattern::kAllZero}, + {"ZeroGroupsBfloat16", RawKind::kBfloat16, kDefaultGeometry, 1, InputPattern::kAllZero}, + {"ZeroGroupsFp8", RawKind::kFp8, kDefaultGeometry, 1, InputPattern::kAllZero}, + {"WarpLaneAmaxFloat16", RawKind::kFloat16, kDefaultGeometry, 1, InputPattern::kSparseOutlier}, + {"WarpLaneAmaxBfloat16", RawKind::kBfloat16, kDefaultGeometry, 1, InputPattern::kSparseOutlier}, + {"WarpLaneAmaxFp8", RawKind::kFp8, kDefaultGeometry, 1, InputPattern::kSparseOutlier}, + {"ReuseDefaultFloat16", RawKind::kFloat16, kDefaultGeometry, 3, InputPattern::kDense, true, true}, + {"ReuseDefaultBfloat16", RawKind::kBfloat16, kDefaultGeometry, 3, InputPattern::kDense, true, true}, + {"ReuseDefaultFp8", RawKind::kFp8, kDefaultGeometry, 3, InputPattern::kDense, true, true}, + {"ReuseModelLikeFloat16", RawKind::kFloat16, kModelLikeGeometry, 2, InputPattern::kDense, true, true}, + {"ReuseModelLikeBfloat16", RawKind::kBfloat16, kModelLikeGeometry, 2, InputPattern::kDense, true, true}, + {"ReuseModelLikeFp8", RawKind::kFp8, kModelLikeGeometry, 2, InputPattern::kDense, true, true}, + {"RoundingMarginsFloat16", RawKind::kFloat16, kDefaultGeometry, 1, InputPattern::kRoundingMargins}, + {"CrossLaunchBfloat16", RawKind::kBfloat16, kSmallVectorGeometry, kCrossLaunchNumPages}, + {"CrossLaunchFp8", RawKind::kFp8, kSmallVectorGeometry, kCrossLaunchNumPages}, + {"PdlBfloat16", RawKind::kBfloat16, kSmallVectorGeometry, 65, InputPattern::kDense, false}, + {"PdlFp8", RawKind::kFp8, kSmallVectorGeometry, 65, InputPattern::kDense, false}, +}; + +class Nvfp4BoundaryRoundTripTest : public testing::TestWithParam +{ +}; + +TEST_P(Nvfp4BoundaryRoundTripTest, MatchesReference) +{ + auto const& test = GetParam(); + runBoundaryRoundTrip(test.kind, test.geometry, test.numPages, test.inputPattern, test.synchronizeBetweenDirections, + test.repeatRoundTrip, test.coldBaseOffset); +} + +std::string roundTripCaseName(testing::TestParamInfo const& info) +{ + return info.param.name; +} + +INSTANTIATE_TEST_SUITE_P(Scenarios, Nvfp4BoundaryRoundTripTest, testing::ValuesIn(kRoundTripCases), roundTripCaseName); + +void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + + PageGeometry constexpr geometry{1, 64, 576}; + std::size_t constexpr numPages = 2; + std::size_t constexpr sideRawBytes = 64U * (128U + 4U); + std::size_t constexpr sideSlotBytes = sideRawBytes + 13U; + std::size_t constexpr coldBaseOffset = 1; + auto const params = makeParams(geometry); + std::size_t const mlaRawBytes = rawBytes(kind, geometry); + std::size_t const mlaPackedBytes = packedBytes(geometry); + std::size_t const mlaScaleBytes = scaleBytes(geometry); + std::size_t const mlaPayloadBytes = mlaPackedBytes + mlaScaleBytes; + std::size_t constexpr gapBeforeSide = 3; + std::size_t const sideColdOffset = mlaPayloadBytes + gapBeforeSide; + std::size_t const sideColdEnd = sideColdOffset + sideRawBytes; + std::size_t const coldPageBytes = roundUp(sideColdEnd, alignof(uint4)); + + DeviceRegion mlaInput(numPages * mlaRawBytes); + DeviceRegion mlaOutput(numPages * mlaRawBytes); + DeviceRegion sideInput(numPages * sideSlotBytes); + DeviceRegion sideOutput(numPages * sideSlotBytes); + MappedHostRegion coldStorage(coldBaseOffset + numPages * coldPageBytes); + auto* coldBase = coldStorage.bytes() + coldBaseOffset; + + std::array, numPages> mlaHost; + std::array, numPages> sideHost; + std::array references; + std::vector offloadTasks; + std::vector onboardTasks; + for (std::size_t page = 0; page < numPages; ++page) + { + mlaHost[page] = makeRawPage(kind, page, 0U, params, geometry, InputPattern::kDense); + references[page] = compressReference(mlaHost[page], kind, params, geometry); + sideHost[page].resize(sideRawBytes); + for (std::size_t byte = 0; byte < sideRawBytes; ++byte) + { + sideHost[page][byte] = static_cast((17U * byte + 53U * page + 11U) & 0xFFU); + } + mlaInput.copyFrom(page * mlaRawBytes, mlaHost[page]); + sideInput.copyFrom(page * sideSlotBytes, sideHost[page]); + auto const pageIndex = static_cast(page); + offloadTasks.push_back({pageIndex, pageIndex}); + onboardTasks.push_back({pageIndex, pageIndex}); + } + + auto const makePlans = [&](DeviceRegion const& mla, DeviceRegion const& side) + { + return std::vector{ + {reinterpret_cast(mla.data()), mlaRawBytes, mlaRawBytes, 0U, mlaPackedBytes, + mlaPayloadBytes, static_cast(gapBeforeSide), Nvfp4BoundaryTransform::kNvfp4, params}, + {reinterpret_cast(side.data()), sideSlotBytes, sideRawBytes, sideColdOffset, 0U, + sideColdEnd, static_cast(coldPageBytes - sideColdEnd), Nvfp4BoundaryTransform::kLossless, + {}}}; + }; + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + makePlans(mlaInput, sideInput), coldPageBytes, runtimeType(kind)); + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + makePlans(mlaOutput, sideOutput), coldPageBytes, runtimeType(kind)); + + CudaStream stream; + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, coldBase, stream); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, coldBase, stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + + auto const cold = coldStorage.payload(); + EXPECT_EQ(cold.front(), kCanary); + for (std::size_t page = 0; page < numPages; ++page) + { + std::size_t const coldPage = coldBaseOffset + page * coldPageBytes; + auto const coldRegion = [&](std::size_t offset, std::size_t bytes) + { + return std::vector(cold.begin() + static_cast(coldPage + offset), + cold.begin() + static_cast(coldPage + offset + bytes)); + }; + EXPECT_EQ(coldRegion(0U, mlaPackedBytes), references[page].packed); + EXPECT_EQ(coldRegion(mlaPackedBytes, mlaScaleBytes), references[page].scales); + EXPECT_TRUE(std::all_of(cold.begin() + static_cast(coldPage + mlaPayloadBytes), + cold.begin() + static_cast(coldPage + sideColdOffset), + [](std::uint8_t value) { return value == 0U; })); + EXPECT_EQ(coldRegion(sideColdOffset, sideRawBytes), sideHost[page]); + EXPECT_TRUE(std::all_of(cold.begin() + static_cast(coldPage + sideColdEnd), + cold.begin() + static_cast(coldPage + coldPageBytes), + [](std::uint8_t value) { return value == 0U; })); + + EXPECT_EQ(mlaOutput.copyToHost(page * mlaRawBytes, mlaRawBytes), + decompressReference(references[page], kind, params, geometry)); + EXPECT_EQ(sideOutput.copyToHost(page * sideSlotBytes, sideRawBytes), sideHost[page]); + for (auto const* side : {&sideInput, &sideOutput}) + { + auto const slotTail = side->copyToHost(page * sideSlotBytes + sideRawBytes, sideSlotBytes - sideRawBytes); + EXPECT_TRUE( + std::all_of(slotTail.begin(), slotTail.end(), [](std::uint8_t value) { return value == kCanary; })); + } + } + mlaInput.expectCanaries(); + mlaOutput.expectCanaries(); + sideInput.expectCanaries(); + sideOutput.expectCanaries(); + coldStorage.expectCanaries(); +} + +class Nvfp4BoundaryMlaSideTest : public testing::TestWithParam +{ +}; + +TEST_P(Nvfp4BoundaryMlaSideTest, MlaPageAndDefaultDsaIndexKeyRoundTripExactly) +{ + runUnaryMlaWithLosslessSideRoundTrip(GetParam()); +} + +INSTANTIATE_TEST_SUITE_P( + AllRuntimeTypes, Nvfp4BoundaryMlaSideTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); + +TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatch) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + + constexpr std::size_t numLayers = 2; + RawKind constexpr kind = RawKind::kBfloat16; + PageGeometry constexpr geometry = kMinimumCompactGeometry; + std::size_t const rawSlotBytes = rawBytes(kind, geometry); + std::size_t const layerRecordBytes = 2U * (packedBytes(geometry) + scaleBytes(geometry)); + std::size_t const layerRecordStride = roundUp(layerRecordBytes, alignof(uint4)); + std::size_t const coldPageBytes = numLayers * layerRecordStride; + + std::array, numLayers> rawInputK; + std::array, numLayers> rawInputV; + std::array, numLayers> rawOutputK; + std::array, numLayers> rawOutputV; + std::array, numLayers> params{ + std::array{makeParams(geometry, 0U), makeParams(geometry, 1U)}, + std::array{makeParams(geometry, 0U), makeParams(geometry, 1U)}}; + // Distinct per-layer K/V scales verify blockIdx.y selects immutable launch metadata. + params[1][0].nvfp4ScaleOrigQuant = 0.5F; + params[1][0].nvfp4ScaleQuantOrig = 2.0F; + params[1][1].nvfp4ScaleOrigQuant = 4.0F; + params[1][1].nvfp4ScaleQuantOrig = 0.25F; + + std::array, 2>, numLayers> rawHost; + std::array, numLayers> references; + std::vector inputPlans; + std::vector outputPlans; + inputPlans.reserve(2U * numLayers); + outputPlans.reserve(2U * numLayers); + for (std::size_t layer = 0; layer < numLayers; ++layer) + { + rawInputK[layer] = std::make_unique(rawSlotBytes); + rawInputV[layer] = std::make_unique(rawSlotBytes); + rawOutputK[layer] = std::make_unique(rawSlotBytes); + rawOutputV[layer] = std::make_unique(rawSlotBytes); + for (std::uint32_t role = 0; role < 2; ++role) + { + rawHost[layer][role] = makeRawPage(kind, layer, role, params[layer][role], geometry, InputPattern::kDense); + references[layer][role] = compressReference(rawHost[layer][role], kind, params[layer][role], geometry); + } + rawInputK[layer]->copyFrom(rawHost[layer][0]); + rawInputV[layer]->copyFrom(rawHost[layer][1]); + std::size_t const base = layer * layerRecordStride; + std::size_t const packed = packedBytes(geometry); + std::size_t const scale = scaleBytes(geometry); + auto const appendPlans = [&](auto& plans, DeviceRegion const& rawK, DeviceRegion const& rawV) + { + plans.push_back({reinterpret_cast(rawK.data()), rawSlotBytes, rawSlotBytes, base, + base + 2U * packed, 0U, 0U, Nvfp4BoundaryTransform::kNvfp4, params[layer][0]}); + plans.push_back({reinterpret_cast(rawV.data()), rawSlotBytes, rawSlotBytes, base + packed, + base + 2U * packed + scale, base + layerRecordBytes, + static_cast(layerRecordStride - layerRecordBytes), Nvfp4BoundaryTransform::kNvfp4, + params[layer][1]}); + }; + appendPlans(inputPlans, *rawInputK[layer], *rawInputV[layer]); + appendPlans(outputPlans, *rawOutputK[layer], *rawOutputV[layer]); + } + + MappedHostRegion compactPage(coldPageBytes); + CudaStream stream; + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + inputPlans, coldPageBytes, Nvfp4BoundaryRuntimeType::kBfloat16); + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + outputPlans, coldPageBytes, Nvfp4BoundaryRuntimeType::kBfloat16); + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({{0, 0}}, inputPlan, compactPage.data(), stream); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress({{0, 0}}, outputPlan, compactPage.data(), stream); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + + auto const compact = compactPage.payload(); + std::size_t const packed = packedBytes(geometry); + std::size_t const scale = scaleBytes(geometry); + auto const compactRegion = [&](std::size_t offset, std::size_t bytes) + { + return std::vector(compact.begin() + static_cast(offset), + compact.begin() + static_cast(offset + bytes)); + }; + for (std::size_t layer = 0; layer < numLayers; ++layer) + { + std::size_t const base = layer * layerRecordStride; + EXPECT_EQ(compactRegion(base, packed), references[layer][0].packed); + EXPECT_EQ(compactRegion(base + packed, packed), references[layer][1].packed); + EXPECT_EQ(compactRegion(base + 2U * packed, scale), references[layer][0].scales); + EXPECT_EQ(compactRegion(base + 2U * packed + scale, scale), references[layer][1].scales); + auto const padding = compactRegion(base + layerRecordBytes, layerRecordStride - layerRecordBytes); + EXPECT_TRUE(std::all_of(padding.begin(), padding.end(), [](std::uint8_t value) { return value == 0U; })); + EXPECT_EQ(rawOutputK[layer]->copyToHost(), + decompressReference(references[layer][0], kind, params[layer][0], geometry)); + EXPECT_EQ(rawOutputV[layer]->copyToHost(), + decompressReference(references[layer][1], kind, params[layer][1], geometry)); + } + compactPage.expectCanaries(); +} + +void expectWholePageLaunchTopology(std::size_t numPages, std::vector expectedGridZ) +{ + constexpr std::size_t numLayers = 2; + RawKind constexpr kind = RawKind::kBfloat16; + std::size_t const rawSlotBytes = rawBytes(kind, kSmallVectorGeometry); + std::size_t const recordBytes = 2U * (packedBytes(kSmallVectorGeometry) + scaleBytes(kSmallVectorGeometry)); + std::size_t const recordStride = roundUp(recordBytes, alignof(uint4)); + std::size_t const coldPageBytes = numLayers * recordStride; + + std::array, numLayers> rawK; + std::array, numLayers> rawV; + std::vector buffers; + buffers.reserve(2U * numLayers); + for (std::size_t layer = 0; layer < numLayers; ++layer) + { + rawK[layer] = std::make_unique(numPages * rawSlotBytes); + rawV[layer] = std::make_unique(numPages * rawSlotBytes); + auto kParams = makeParams(kSmallVectorGeometry, 0U); + auto const vParams = makeParams(kSmallVectorGeometry, 1U); + kParams.nvfp4ScaleOrigQuant *= static_cast(layer + 1U); + kParams.nvfp4ScaleQuantOrig /= static_cast(layer + 1U); + std::size_t const base = layer * recordStride; + std::size_t const packed = packedBytes(kSmallVectorGeometry); + std::size_t const scale = scaleBytes(kSmallVectorGeometry); + buffers.push_back({reinterpret_cast(rawK[layer]->data()), rawSlotBytes, rawSlotBytes, base, + base + 2U * packed, 0U, 0U, Nvfp4BoundaryTransform::kNvfp4, kParams}); + buffers.push_back({reinterpret_cast(rawV[layer]->data()), rawSlotBytes, rawSlotBytes, + base + packed, base + 2U * packed + scale, base + recordBytes, + static_cast(recordStride - recordBytes), Nvfp4BoundaryTransform::kNvfp4, vParams}); + } + + MappedHostRegion coldPages(numPages * coldPageBytes); + std::vector offloadPages; + std::vector onboardPages; + offloadPages.reserve(numPages); + onboardPages.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + auto const pageIndex = static_cast(page); + offloadPages.push_back({pageIndex, pageIndex}); + onboardPages.push_back({pageIndex, pageIndex}); + } + + CudaStream stream; + auto const plan + = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan(buffers, coldPageBytes, Nvfp4BoundaryRuntimeType::kBfloat16); + auto const expectWholePageKernels = [&](auto const& enqueue) + { + cudaGraph_t graph{}; + ASSERT_EQ(cudaStreamBeginCapture(stream, cudaStreamCaptureModeGlobal), cudaSuccess); + enqueue(); + ASSERT_EQ(cudaStreamEndCapture(stream, &graph), cudaSuccess); + + std::size_t numNodes = 0; + ASSERT_EQ(cudaGraphGetNodes(graph, nullptr, &numNodes), cudaSuccess); + std::vector nodes(numNodes); + ASSERT_EQ(cudaGraphGetNodes(graph, nodes.data(), &numNodes), cudaSuccess); + std::size_t kernelNodes = 0; + std::vector actualGridZ; + for (auto const node : nodes) + { + cudaGraphNodeType nodeType{}; + ASSERT_EQ(cudaGraphNodeGetType(node, &nodeType), cudaSuccess); + if (nodeType != cudaGraphNodeTypeKernel) + { + continue; + } + ++kernelNodes; + cudaKernelNodeParams nodeParams{}; + ASSERT_EQ(cudaGraphKernelNodeGetParams(node, &nodeParams), cudaSuccess); + EXPECT_EQ(nodeParams.gridDim.y, 2U * numLayers); + actualGridZ.push_back(nodeParams.gridDim.z); + } + std::sort(actualGridZ.begin(), actualGridZ.end()); + std::sort(expectedGridZ.begin(), expectedGridZ.end()); + EXPECT_EQ(kernelNodes, expectedGridZ.size()); + EXPECT_EQ(actualGridZ, expectedGridZ); + ASSERT_EQ(cudaGraphDestroy(graph), cudaSuccess); + }; + + expectWholePageKernels([&] + { tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadPages, plan, coldPages.data(), stream); }); + expectWholePageKernels([&] + { tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardPages, plan, coldPages.data(), stream); }); +} + +TEST(Nvfp4BoundaryWholePageTest, TwoHundredFiftySevenPagesUseExactlyTwoWholePageKernelsPerDirection) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + expectWholePageLaunchTopology(257, {1, 256}); +} + +class Nvfp4BoundaryTailTest : public testing::TestWithParam +{ +}; + +TEST_P(Nvfp4BoundaryTailTest, InactiveRowsDoNotAffectTheValidPrefix) +{ + runPartialPageTailIsolation(GetParam()); +} + +INSTANTIATE_TEST_SUITE_P( + AllRuntimeTypes, Nvfp4BoundaryTailTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); + +TEST(Nvfp4BoundaryValidationTest, EmptyBatchIsAnAsyncNoOp) +{ + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({}, Nvfp4BoundaryPreparedPlan{}, nullptr, nullptr); + tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress({}, Nvfp4BoundaryPreparedPlan{}, nullptr, nullptr); +} + +TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + + std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); + LayerBuffers buffers(rawSlotBytes); + std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); + auto const prepare = + [&](Nvfp4BoundaryKernelParams const& params, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + { + std::size_t activeRawBytes = rawSlotBytes; + std::size_t dataBytes = packedBytes(kDefaultGeometry); + std::size_t scales = scaleBytes(kDefaultGeometry); + if (params.numKvHeads > 0 && params.tokensPerPage > 0 && params.headDim > 0) + { + std::uint64_t const elements = static_cast(params.numKvHeads) + * static_cast(params.tokensPerPage) * static_cast(params.headDim); + std::uint64_t const candidateRawBytes = elements * (type == Nvfp4BoundaryRuntimeType::kFp8E4m3 ? 1U : 2U); + if (candidateRawBytes > 0U && candidateRawBytes <= rawSlotBytes) + { + activeRawBytes = static_cast(candidateRawBytes); + dataBytes = static_cast(elements / 2U); + scales = static_cast(elements / 16U); + } + } + Nvfp4BoundaryBufferPlan const buffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, + activeRawBytes, 0U, dataBytes, dataBytes + scales, + static_cast(coldPageBytes - dataBytes - scales), Nvfp4BoundaryTransform::kNvfp4, params}; + static_cast(tensorrt_llm::kernels::prepareNvfp4BoundaryPlan({buffer}, coldPageBytes, type)); + }; + auto const expectInvalid + = [&](char const* name, auto const& mutate, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + { + SCOPED_TRACE(name); + auto params = makeParams(); + mutate(params); + EXPECT_ANY_THROW(prepare(params, type)); + }; + + auto valid = makeParams(); + valid.tokensPerPage = 6; + EXPECT_NO_THROW(prepare(valid)); + EXPECT_NO_THROW(prepare(makeParams(PageGeometry{1, 1, 16}))); + + expectInvalid("zero heads", [](auto& params) { params.numKvHeads = 0; }); + expectInvalid("zero tokens", [](auto& params) { params.tokensPerPage = 0; }); + expectInvalid("zero head dimension", [](auto& params) { params.headDim = 0; }); + expectInvalid("unaligned head dimension", [](auto& params) { params.headDim = 24; }); + expectInvalid( + "element count overflow", [](auto& params) { params.numKvHeads = std::numeric_limits::max(); }); + expectInvalid("zero NVFP4 quant scale", [](auto& params) { params.nvfp4ScaleOrigQuant = 0.0F; }); + expectInvalid("negative NVFP4 dequant scale", [](auto& params) { params.nvfp4ScaleQuantOrig = -1.0F; }); + expectInvalid("NaN NVFP4 quant scale", + [](auto& params) { params.nvfp4ScaleOrigQuant = std::numeric_limits::quiet_NaN(); }); + expectInvalid("infinite NVFP4 dequant scale", + [](auto& params) { params.nvfp4ScaleQuantOrig = std::numeric_limits::infinity(); }); + expectInvalid( + "zero FP8 quant scale", [](auto& params) { params.fp8ScaleOrigQuant = 0.0F; }, + Nvfp4BoundaryRuntimeType::kFp8E4m3); + expectInvalid( + "infinite FP8 dequant scale", [](auto& params) + { params.fp8ScaleQuantOrig = std::numeric_limits::infinity(); }, Nvfp4BoundaryRuntimeType::kFp8E4m3); +} + +TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) +{ + ASSERT_EQ(cudaSetDevice(0), cudaSuccess); + if (!tensorrt_llm::common::isSM100Family()) + { + GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + } + + std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); + LayerBuffers buffers(rawSlotBytes); + Nvfp4BoundaryOffloadPageTask const validOffload{0, 0}; + std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); + std::size_t const packed = packedBytes(kDefaultGeometry); + std::size_t const scale = scaleBytes(kDefaultGeometry); + Nvfp4BoundaryBufferPlan const validBuffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, + rawSlotBytes, 0U, packed, packed + scale, static_cast(coldPageBytes - packed - scale), + Nvfp4BoundaryTransform::kNvfp4, makeParams()}; + auto const prepare = [&](Nvfp4BoundaryBufferPlan const& buffer, std::size_t pageBytes, + Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + { return tensorrt_llm::kernels::prepareNvfp4BoundaryPlan({buffer}, pageBytes, type); }; + auto const expectInvalid + = [&](char const* name, auto const& mutate, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + { + SCOPED_TRACE(name); + auto buffer = validBuffer; + mutate(buffer); + EXPECT_ANY_THROW(static_cast(prepare(buffer, coldPageBytes, type))); + }; + auto const validPlan = prepare(validBuffer, coldPageBytes); + + expectInvalid("unaligned raw base", [](auto& buffer) { buffer.rawBase += 1U; }); + EXPECT_ANY_THROW( + tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({validOffload}, validPlan, nullptr, nullptr)); + expectInvalid("unaligned raw stride", [](auto& buffer) { buffer.rawSlotBytes += alignof(uint4) / 2U; }); + expectInvalid("raw bytes exceed stride", [](auto& buffer) { buffer.rawBytes = buffer.rawSlotBytes + 1U; }); + expectInvalid("cold data interval exceeds page", [&](auto& buffer) { buffer.coldDataOffset = coldPageBytes; }); + expectInvalid("cold intervals overlap", [](auto& buffer) { buffer.coldScaleOffset = buffer.coldDataOffset; }); + EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes + alignof(uint4) / 2U))); + expectInvalid( + "unaligned FP8 raw base", [](auto& buffer) { buffer.rawBase += 1U; }, Nvfp4BoundaryRuntimeType::kFp8E4m3); + + auto const unsupportedType = static_cast(255); + EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes, unsupportedType))); +} + +} // namespace diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt new file mode 100644 index 000000000000..fdd5074c50ce --- /dev/null +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -0,0 +1,22 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# All rights reserved. SPDX-License-Identifier: Apache-2.0 + +set(NVFP4_COLD_PAGE_CODEC_TEST_SRC + nvfp4ColdPageCodecTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) + +add_gtest(nvfp4ColdPageCodecTest "${NVFP4_COLD_PAGE_CODEC_TEST_SRC}" + NO_TLLM_LINKAGE) +target_link_libraries(nvfp4ColdPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) +target_include_directories( + nvfp4ColdPageCodecTest + PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp new file mode 100644 index 000000000000..339a2b277c04 --- /dev/null +++ b/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp @@ -0,0 +1,572 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" + +#include + +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +static_assert(std::is_base_of_v); + +struct RecordedLaunch +{ + int offloadCalls = 0; + int onboardCalls = 0; + std::vector offloadPages; + std::vector onboardPages; + kernels::Nvfp4BoundaryPreparedPlan plan; + void const* coldBase = nullptr; + cudaStream_t stream{}; +}; + +RecordedLaunch gLaunch; + +constexpr std::uintptr_t kGpuKBase = 0x100000; +constexpr std::uintptr_t kGpuVBase = 0x200000; +constexpr std::uintptr_t kColdBase = 0x300000; +constexpr std::size_t kLayerRawBytes = 320; +constexpr std::size_t kLayerColdBytesAligned = 192; +constexpr std::size_t kNumAttentionLayers = 8; +constexpr std::size_t kGpuSlotBytes = kNumAttentionLayers * kLayerRawBytes; +constexpr std::size_t kColdSlotBytes = kNumAttentionLayers * kLayerColdBytesAligned; +constexpr std::size_t kMlaRawBytes = 64U * 576U * 2U; +constexpr std::size_t kMlaColdBytes = 64U * 576U / 2U + 64U * 576U / 16U; +constexpr std::size_t kDsaIndexKeyBytes = 64U * (128U + 4U); +constexpr std::uintptr_t kStreamValue = 0x7000; + +void resetLaunch() +{ + gLaunch = {}; +} + +std::vector makeLayers(std::size_t count = kNumAttentionLayers, int firstLayer = 0) +{ + std::vector layers; + layers.reserve(count); + for (std::size_t index = 0; index < count; ++index) + { + auto const scale = static_cast(index + 2U); + layers.push_back({firstLayer + static_cast(index), kernels::Nvfp4BoundaryRuntimeType::kFloat16, 1, 5, 32, + {scale, scale + 0.5F}, {1.0F / scale, 1.0F / (scale + 0.5F)}}); + } + return layers; +} + +std::vector makeMlaLayers(std::size_t count) +{ + std::vector layers; + layers.reserve(count); + for (std::size_t layer = 0; layer < count; ++layer) + { + layers.push_back({static_cast(layer), kernels::Nvfp4BoundaryRuntimeType::kFloat16, 1, 64, 576, + {1.0F, 1.0F}, {1.0F, 1.0F}}); + } + return layers; +} + +kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::PoolGroupIndex{0}, + kv::LayerGroupId lifeCycle = kv::LayerGroupId{3}, std::size_t count = kNumAttentionLayers, int firstLayer = 0, + std::uintptr_t keyBase = kGpuKBase, std::uintptr_t valueBase = kGpuVBase) +{ + kv::CoalescedBuffer keys{kLayerRawBytes, {}}; + kv::CoalescedBuffer values{kLayerRawBytes, {}}; + for (std::size_t index = 0; index < count; ++index) + { + auto const layerId = firstLayer + static_cast(index); + keys.bufferIds.push_back({layerId, "key"}); + values.bufferIds.push_back({layerId, "value"}); + } + + kv::SlotDescVariant variant{ + lifeCycle, kv::TypedVec{std::move(keys), std::move(values)}}; + auto const slotBytes = count * kLayerRawBytes; + return kv::PoolGroupDesc{poolGroupIndex, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, + kv::TypedVec{ + kv::PoolDesc{kv::PoolIndex{0}, keyBase, slotBytes}, kv::PoolDesc{kv::PoolIndex{1}, valueBase, slotBytes}}}; +} + +kv::PoolGroupDesc makeMlaDesc(std::vector const& ownsIndexer, kv::LayerGroupId lifeCycle = kv::LayerGroupId{0}, + int firstLayer = 0, std::size_t keyBytes = kLayerRawBytes, std::size_t indexBytes = 68U, + std::uintptr_t keyBase = kGpuKBase, std::uintptr_t indexBase = kGpuVBase) +{ + kv::CoalescedBuffer keys{keyBytes, {}}; + kv::CoalescedBuffer indexes{indexBytes, {}}; + for (std::size_t index = 0; index < ownsIndexer.size(); ++index) + { + auto const layerId = firstLayer + static_cast(index); + keys.bufferIds.push_back({layerId, "key"}); + if (ownsIndexer[index]) + { + indexes.bufferIds.push_back({layerId, "index_key"}); + } + } + + kv::TypedVec buffers; + buffers.push_back(std::move(keys)); + if (!indexes.bufferIds.empty()) + { + buffers.push_back(std::move(indexes)); + } + kv::SlotDescVariant variant{lifeCycle, std::move(buffers)}; + + kv::TypedVec pools; + pools.push_back(kv::PoolDesc{kv::PoolIndex{0}, keyBase, ownsIndexer.size() * keyBytes}); + auto const indexCount = static_cast(std::count(ownsIndexer.begin(), ownsIndexer.end(), true)); + if (indexCount != 0U) + { + pools.push_back(kv::PoolDesc{kv::PoolIndex{1}, indexBase, indexCount * indexBytes}); + } + return kv::PoolGroupDesc{ + kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, std::move(pools)}; +} + +bool configureOne(Nvfp4ColdPageCodec& codec, kv::PoolGroupDesc const& desc) +{ + return codec.configure(&desc, kv::PoolGroupIndex{1}); +} + +std::unique_ptr makeConfiguredAttentionCodec( + std::size_t count = kNumAttentionLayers, kv::LayerGroupId lifeCycle = kv::LayerGroupId{3}) +{ + auto codec = std::make_unique(makeLayers(count)); + EXPECT_TRUE(configureOne(*codec, makeAttentionDesc(kv::PoolGroupIndex{0}, lifeCycle, count))); + return codec; +} + +TEST(Nvfp4ColdPageCodecTest, OneCompletePageTaskCoversAllLayersWithDistinctScales) +{ + resetLaunch(); + auto codec = makeConfiguredAttentionCodec(); + EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{3}), kColdSlotBytes); + EXPECT_EQ(codec->queryPageIndexLocation(kv::LayerGroupId{3}), kv::PageIndexLocation::kHost); + + kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; + auto const stream = reinterpret_cast(kStreamValue); + ASSERT_TRUE( + codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + + EXPECT_EQ(gLaunch.offloadCalls, 1); + ASSERT_EQ(gLaunch.offloadPages.size(), 2U); + EXPECT_EQ(gLaunch.offloadPages[0].gpuPageIndex, 1); + EXPECT_EQ(gLaunch.offloadPages[0].coldPageIndex, 2); + EXPECT_EQ(gLaunch.offloadPages[1].gpuPageIndex, 3); + EXPECT_EQ(gLaunch.offloadPages[1].coldPageIndex, 5); + ASSERT_EQ(gLaunch.plan.numBuffers, 2U * kNumAttentionLayers); + for (std::size_t layer = 0; layer < kNumAttentionLayers; ++layer) + { + auto const& key = gLaunch.plan.buffers[2U * layer]; + auto const& value = gLaunch.plan.buffers[2U * layer + 1U]; + auto const layerOffset = layer * kLayerColdBytesAligned; + EXPECT_EQ(key.rawBase, kGpuKBase + layer * kLayerRawBytes); + EXPECT_EQ(value.rawBase, kGpuVBase + layer * kLayerRawBytes); + EXPECT_EQ(key.rawSlotBytes, kGpuSlotBytes); + EXPECT_EQ(value.rawSlotBytes, kGpuSlotBytes); + EXPECT_EQ(key.rawBytes, kLayerRawBytes); + EXPECT_EQ(value.rawBytes, kLayerRawBytes); + EXPECT_EQ(key.coldDataOffset, layerOffset); + EXPECT_EQ(value.coldDataOffset, layerOffset + 80U); + EXPECT_EQ(key.coldScaleOffset, layerOffset + 160U); + EXPECT_EQ(value.coldScaleOffset, layerOffset + 170U); + EXPECT_EQ(value.coldPaddingOffset, layerOffset + 180U); + EXPECT_EQ(value.coldPaddingBytes, 12U); + EXPECT_EQ(key.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); + EXPECT_EQ(value.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); + EXPECT_EQ(key.params.tokensPerPage, 5); + EXPECT_EQ(key.params.headDim, 32); + EXPECT_FLOAT_EQ(key.params.nvfp4ScaleOrigQuant, static_cast(layer + 2U)); + EXPECT_FLOAT_EQ(value.params.nvfp4ScaleOrigQuant, static_cast(layer + 2U) + 0.5F); + } + EXPECT_EQ(gLaunch.coldBase, reinterpret_cast(kColdBase)); + EXPECT_EQ(gLaunch.plan.coldPageBytes, kColdSlotBytes); + EXPECT_EQ(gLaunch.stream, stream); + + ASSERT_TRUE(codec->decode( + kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + EXPECT_EQ(gLaunch.onboardCalls, 1); + ASSERT_EQ(gLaunch.onboardPages.size(), 2U); + EXPECT_EQ(gLaunch.onboardPages[0].gpuPageIndex, 2); + EXPECT_EQ(gLaunch.onboardPages[0].coldPageIndex, 1); +} + +TEST(Nvfp4ColdPageCodecTest, KeyOnlyMlaUsesLatentPackedThenScaleLayout) +{ + resetLaunch(); + Nvfp4ColdPageCodec codec{makeLayers(1)}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc({false}))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 96U); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 1U); + auto const& latent = gLaunch.plan.buffers[0]; + EXPECT_EQ(latent.rawBase, kGpuKBase); + EXPECT_EQ(latent.rawSlotBytes, kLayerRawBytes); + EXPECT_EQ(latent.rawBytes, kLayerRawBytes); + EXPECT_EQ(latent.coldDataOffset, 0U); + EXPECT_EQ(latent.coldScaleOffset, 80U); + EXPECT_EQ(latent.coldPaddingOffset, 90U); + EXPECT_EQ(latent.coldPaddingBytes, 6U); + EXPECT_EQ(latent.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); +} + +TEST(Nvfp4ColdPageCodecTest, KeyAndIndexAppendsLosslessIndexWithinTheLayerRecord) +{ + resetLaunch(); + Nvfp4ColdPageCodec codec{makeLayers(1)}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc({true}))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 160U); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 2U); + auto const& latent = gLaunch.plan.buffers[0]; + auto const& index = gLaunch.plan.buffers[1]; + EXPECT_EQ(latent.coldDataOffset, 0U); + EXPECT_EQ(latent.coldScaleOffset, 80U); + EXPECT_EQ(index.rawBase, kGpuVBase); + EXPECT_EQ(index.rawSlotBytes, 68U); + EXPECT_EQ(index.rawBytes, 68U); + EXPECT_EQ(index.coldDataOffset, 90U); + EXPECT_EQ(index.coldPaddingOffset, 158U); + EXPECT_EQ(index.coldPaddingBytes, 2U); + EXPECT_EQ(index.transform, kernels::Nvfp4BoundaryTransform::kLossless); +} + +TEST(Nvfp4ColdPageCodecTest, FullAndSharedIndexerLayersHaveDistinctPerLayerRecords) +{ + resetLaunch(); + Nvfp4ColdPageCodec codec{makeLayers(3)}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc({true, false, true}))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 416U); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 5U); + EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); + EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kLayerRawBytes); + EXPECT_EQ(gLaunch.plan.buffers[3].rawBase, kGpuKBase + 2U * kLayerRawBytes); + EXPECT_EQ(gLaunch.plan.buffers[4].rawBase, kGpuVBase + 68U); + EXPECT_EQ(gLaunch.plan.buffers[0].rawSlotBytes, 3U * kLayerRawBytes); + EXPECT_EQ(gLaunch.plan.buffers[1].rawSlotBytes, 2U * 68U); + EXPECT_EQ(gLaunch.plan.buffers[0].coldDataOffset, 0U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); + EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, 160U); + EXPECT_EQ(gLaunch.plan.buffers[3].coldDataOffset, 256U); + EXPECT_EQ(gLaunch.plan.buffers[4].coldDataOffset, 346U); + EXPECT_EQ(gLaunch.plan.buffers[4].coldPaddingOffset, 414U); + EXPECT_EQ(gLaunch.plan.buffers[4].coldPaddingBytes, 2U); +} + +TEST(Nvfp4ColdPageCodecTest, DeepSeekV32AllIndexerLayoutFitsOneColdPagePlan) +{ + resetLaunch(); + std::vector const ownsIndexer(61, true); + Nvfp4ColdPageCodec codec{makeMlaLayers(ownsIndexer.size())}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc(ownsIndexer, kv::LayerGroupId{0}, 0, kMlaRawBytes, kDsaIndexKeyBytes))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 61U * (kMlaColdBytes + kDsaIndexKeyBytes)); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_EQ(gLaunch.plan.numBuffers, 122U); + EXPECT_EQ(gLaunch.plan.coldPageBytes, 1780224U); +} + +TEST(Nvfp4ColdPageCodecTest, Glm52MixedIndexerLayoutFitsOneColdPagePlan) +{ + resetLaunch(); + std::vector ownsIndexer(78, false); + for (std::size_t layer = 0; layer < ownsIndexer.size(); ++layer) + { + ownsIndexer[layer] = layer < 3U || (layer >= 6U && layer % 4U == 2U); + } + ASSERT_EQ(std::count(ownsIndexer.begin(), ownsIndexer.end(), true), 21); + + Nvfp4ColdPageCodec codec{makeMlaLayers(ownsIndexer.size())}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc(ownsIndexer, kv::LayerGroupId{0}, 0, kMlaRawBytes, kDsaIndexKeyBytes))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 78U * kMlaColdBytes + 21U * kDsaIndexKeyBytes); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_EQ(gLaunch.plan.numBuffers, 99U); + EXPECT_EQ(gLaunch.plan.coldPageBytes, 1794816U); +} + +TEST(Nvfp4ColdPageCodecTest, PreservesOneCodecSubmissionAcrossThe256PageKernelBoundary) +{ + resetLaunch(); + auto codec = makeConfiguredAttentionCodec(); + + std::vector indices(257); + for (std::size_t page = 0; page < indices.size(); ++page) + { + indices[page] = {static_cast(500U - page), static_cast(page * 2U)}; + } + ASSERT_TRUE(codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices.data(), indices.size(), + reinterpret_cast(kStreamValue))); + EXPECT_EQ(gLaunch.offloadCalls, 1); + EXPECT_EQ(gLaunch.offloadPages.size(), 257U); + EXPECT_EQ(gLaunch.plan.numBuffers, 2U * kNumAttentionLayers); +} + +TEST(Nvfp4ColdPageCodecTest, EmptyAttentionBatchIsValidAndDoesNotLaunch) +{ + resetLaunch(); + auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); + + EXPECT_TRUE(codec->encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + EXPECT_TRUE(codec->decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + EXPECT_EQ(gLaunch.offloadCalls, 0); + EXPECT_EQ(gLaunch.onboardCalls, 0); +} + +TEST(Nvfp4ColdPageCodecTest, DefaultStreamIsAccepted) +{ + resetLaunch(); + auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); + + kv::PageIndexPair const indices[]{{0, 0}}; + EXPECT_TRUE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_TRUE(codec->decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_EQ(gLaunch.offloadCalls, 1); + EXPECT_EQ(gLaunch.onboardCalls, 1); +} + +TEST(Nvfp4ColdPageCodecTest, NonEmptyAttentionBatchRequiresPageIndices) +{ + auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); + + auto const stream = reinterpret_cast(kStreamValue); + EXPECT_FALSE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), nullptr, 1U, stream)); + EXPECT_FALSE(codec->decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), nullptr, 1U, stream)); +} + +TEST(Nvfp4ColdPageCodecTest, OnlyFp8RuntimeRequiresFp8Scales) +{ + auto layers = makeLayers(1); + layers.front().fp8ScaleOrigQuant = {0.0F, 0.0F}; + layers.front().fp8ScaleQuantOrig = {0.0F, 0.0F}; + EXPECT_NO_THROW({ Nvfp4ColdPageCodec codec{layers}; }); + + layers.front().runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3; + EXPECT_THROW({ Nvfp4ColdPageCodec codec{layers}; }, std::invalid_argument); +} + +TEST(Nvfp4ColdPageCodecTest, DiscoversLifecycleMembershipAcrossPoolGroups) +{ + auto layers = makeLayers(2); + auto secondGroupLayers = makeLayers(2, 2); + layers.insert(layers.end(), secondGroupLayers.begin(), secondGroupLayers.end()); + Nvfp4ColdPageCodec codec{layers}; + std::array descs{makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2), + makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1}, 2, 2, 0x400000, 0x500000)}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 2U * kLayerColdBytesAligned); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kLayerColdBytesAligned); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{0}), kv::LayerGroupId{0}); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); +} + +TEST(Nvfp4ColdPageCodecTest, RejectsConfiguredAttentionLayerAbsentFromAllGpuDescriptors) +{ + Nvfp4ColdPageCodec codec{makeLayers(2)}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1))); +} + +TEST(Nvfp4ColdPageCodecTest, RejectsAttentionBufferWithMismatchedGeometry) +{ + Nvfp4ColdPageCodec codec{makeLayers()}; + auto desc = makeAttentionDesc(); + desc.slotDesc.variants.front().coalescedBuffers[kv::PoolIndex{0}].singleBufferSize += 16U; + desc.pools[kv::PoolIndex{0}].slotBytes += 16U * kNumAttentionLayers; + + EXPECT_FALSE(configureOne(codec, desc)); +} + +TEST(Nvfp4ColdPageCodecTest, CoalescedAttentionSideBufferUsesItsOwnBaseOffsetAndSlotStride) +{ + resetLaunch(); + Nvfp4ColdPageCodec codec{makeLayers(1)}; + auto desc = makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1); + auto& keys = desc.slotDesc.variants.front().coalescedBuffers[kv::PoolIndex{0}]; + keys.bufferIds.push_back({0, "index_key"}); + desc.pools[kv::PoolIndex{0}].slotBytes += keys.singleBufferSize; + + ASSERT_TRUE(configureOne(codec, desc)); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 512U); + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 3U); + auto const& side = gLaunch.plan.buffers[2]; + EXPECT_EQ(side.rawBase, kGpuKBase + kLayerRawBytes); + EXPECT_EQ(side.rawSlotBytes, 2U * kLayerRawBytes); + EXPECT_EQ(side.rawBytes, kLayerRawBytes); + EXPECT_EQ(side.coldDataOffset, 180U); + EXPECT_EQ(side.coldPaddingOffset, 500U); + EXPECT_EQ(side.coldPaddingBytes, 12U); + EXPECT_EQ(side.transform, kernels::Nvfp4BoundaryTransform::kLossless); +} + +TEST(Nvfp4ColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) +{ + auto codec = makeConfiguredAttentionCodec(); + EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{99}), 0U); + EXPECT_EQ(codec->getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); + EXPECT_EQ(codec->queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); +} + +TEST(Nvfp4ColdPageCodecTest, NonAttentionLifecycleUsesLosslessSingleBlob) +{ + resetLaunch(); + int deviceCount = 0; + if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) + { + GTEST_SKIP() << "CUDA device is required for the lossless copy data-plane test"; + } + + constexpr std::size_t kPoolBytes = 64; + constexpr std::size_t kSlots = 2; + std::byte* statePool = nullptr; + std::byte* convPool = nullptr; + std::byte* coldPool = nullptr; + cudaStream_t stream = nullptr; + ASSERT_EQ(cudaMalloc(reinterpret_cast(&statePool), kPoolBytes * kSlots), cudaSuccess); + ASSERT_EQ(cudaMalloc(reinterpret_cast(&convPool), kPoolBytes * kSlots), cudaSuccess); + ASSERT_EQ(cudaMalloc(reinterpret_cast(&coldPool), 2U * kPoolBytes * kSlots), cudaSuccess); + ASSERT_EQ(cudaStreamCreate(&stream), cudaSuccess); + + kv::SlotDescVariant variant; + variant.lifeCycleId = kv::LayerGroupId{1}; + variant.coalescedBuffers = kv::TypedVec{ + kv::CoalescedBuffer{kPoolBytes, {{10, "ssm_state"}}}, + kv::CoalescedBuffer{kPoolBytes, {{10, "conv_state"}}}, + }; + kv::PoolGroupDesc desc; + desc.poolGroupIndex = kv::PoolGroupIndex{1}; + desc.numSlots = kSlots; + desc.slotDesc.variants = {variant}; + desc.pools = kv::TypedVec{ + kv::PoolDesc{kv::PoolIndex{0}, reinterpret_cast(statePool), kPoolBytes}, + kv::PoolDesc{kv::PoolIndex{1}, reinterpret_cast(convPool), kPoolBytes}, + }; + + Nvfp4ColdPageCodec codec{makeLayers(1)}; + std::array descs{makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1), desc}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kPoolBytes); + + std::vector state(kPoolBytes); + std::vector conv(kPoolBytes); + for (std::size_t index = 0; index < kPoolBytes; ++index) + { + state[index] = static_cast(index + 1U); + conv[index] = static_cast(index + 65U); + } + ASSERT_EQ(cudaMemcpy(statePool + kPoolBytes, state.data(), kPoolBytes, cudaMemcpyHostToDevice), cudaSuccess); + ASSERT_EQ(cudaMemcpy(convPool + kPoolBytes, conv.data(), kPoolBytes, cudaMemcpyHostToDevice), cudaSuccess); + + kv::PageIndexPair const encodePair{0, 1}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{1}, coldPool, &encodePair, 1U, stream)); + ASSERT_EQ(cudaMemsetAsync(statePool, 0, kPoolBytes, stream), cudaSuccess); + ASSERT_EQ(cudaMemsetAsync(convPool, 0, kPoolBytes, stream), cudaSuccess); + kv::PageIndexPair const decodePair{0, 0}; + ASSERT_TRUE(codec.decode(kv::LayerGroupId{1}, coldPool, &decodePair, 1U, stream)); + ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); + + std::vector restoredState(kPoolBytes); + std::vector restoredConv(kPoolBytes); + ASSERT_EQ(cudaMemcpy(restoredState.data(), statePool, kPoolBytes, cudaMemcpyDeviceToHost), cudaSuccess); + ASSERT_EQ(cudaMemcpy(restoredConv.data(), convPool, kPoolBytes, cudaMemcpyDeviceToHost), cudaSuccess); + EXPECT_EQ(restoredState, state); + EXPECT_EQ(restoredConv, conv); + EXPECT_EQ(gLaunch.offloadCalls, 0); + + EXPECT_EQ(cudaStreamDestroy(stream), cudaSuccess); + EXPECT_EQ(cudaFree(coldPool), cudaSuccess); + EXPECT_EQ(cudaFree(convPool), cudaSuccess); + EXPECT_EQ(cudaFree(statePool), cudaSuccess); +} + +TEST(Nvfp4ColdPageCodecTest, AttentionAndSsmSharingOneHotPoolGroupUseDifferentTransforms) +{ + Nvfp4ColdPageCodec codec{makeLayers(1)}; + auto desc = makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1); + + kv::SlotDescVariant ssm; + ssm.lifeCycleId = kv::LayerGroupId{1}; + ssm.coalescedBuffers = kv::TypedVec{ + kv::CoalescedBuffer{kLayerRawBytes, {{10, "ssm_state"}}}, + kv::CoalescedBuffer{kLayerRawBytes, {{10, "conv_state"}}}, + }; + desc.slotDesc.variants.push_back(std::move(ssm)); + + ASSERT_TRUE(configureOne(codec, desc)); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), kLayerColdBytesAligned); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kLayerRawBytes); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{0}), kv::LayerGroupId{0}); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); +} + +} // namespace +} // namespace tensorrt_llm::kv_cache_compression + +namespace tensorrt_llm::kernels +{ + +Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType) +{ + if (buffers.empty() || buffers.size() > kNvfp4BoundaryMaxBuffersPerLaunch || coldPageBytes == 0U) + { + throw std::invalid_argument("invalid test launch plan"); + } + Nvfp4BoundaryPreparedPlan plan; + std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); + plan.numBuffers = static_cast(buffers.size()); + plan.coldPageBytes = coldPageBytes; + plan.runtimeType = runtimeType; + return plan; +} + +void invokeNvfp4BoundaryOffloadCompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream) +{ + auto& launch = kv_cache_compression::gLaunch; + ++launch.offloadCalls; + launch.offloadPages = pages; + launch.plan = plan; + launch.coldBase = coldBase; + launch.stream = stream; +} + +void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, + Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream) +{ + auto& launch = kv_cache_compression::gLaunch; + ++launch.onboardCalls; + launch.onboardPages = pages; + launch.plan = plan; + launch.coldBase = coldBase; + launch.stream = stream; +} + +} // namespace tensorrt_llm::kernels diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py new file mode 100644 index 000000000000..83cf2067bd62 --- /dev/null +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from .quantization_for_cold_page import ColdPageQuantizationCompression + +__all__ = ["ColdPageQuantizationCompression"] diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py new file mode 100644 index 000000000000..86d03cf0bca1 --- /dev/null +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -0,0 +1,175 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""NVFP4 compression for KVCM V2 cold pages. + +The compression manager owns optional ModelOpt K/V global scales and creates +the native storage-boundary codec before KVCM allocates cold Slots. Attention, +the active KV-cache format, and the normal model-loading path remain unaware of +the cold representation. +""" + +import json +import math +import os +import re +from pathlib import Path +from typing import TYPE_CHECKING, Sequence + +from tensorrt_llm.quantization.modelopt_config import ( + is_modelopt_quant_config, + read_modelopt_quant_config, +) + +from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager + +if TYPE_CHECKING: + from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig + +ScalePair = tuple[float, float] +LayerScales = tuple[ScalePair, ScalePair] + +_IDENTITY_NVFP4_SCALES: LayerScales = ((1.0, 1.0), (1.0, 1.0)) +_MODEL_OPT_KV_SCALE_KEY = re.compile( + r"(?:^|\.)layers\.(?P\d+)\.self_attn\." + r"(?P[kv])_proj\.(?P=kind)_scale$" +) + + +def _load_modelopt_nvfp4_scales( + checkpoint_path: str | None, +) -> dict[int, LayerScales]: + """Load optional ModelOpt NVFP4 K/V global scales by model layer. + + This mirrors the native NVFP4 KV loader's contract: no checkpoint means + identity scales, ``TRTLLM_LOAD_KV_SCALES`` can disable loading, and K/V + scalars from standard safetensors shards are reduced with ``max``. + """ + + if checkpoint_path is None or os.environ.get("TRTLLM_LOAD_KV_SCALES", "1") != "1": + return {} + + checkpoint_dir = Path(checkpoint_path) + weight_files = sorted(checkpoint_dir.glob("*.safetensors")) + ordinary_files = [path for path in weight_files if "consolidated" not in path.name] + weight_files = ordinary_files or weight_files + if not weight_files: + raise FileNotFoundError( + f"No safetensors files in ModelOpt scale checkpoint {checkpoint_dir}" + ) + + metadata_path = checkpoint_dir / "hf_quant_config.json" + if metadata_path.exists(): + metadata = json.loads(metadata_path.read_text()) + else: + config_path = checkpoint_dir / "config.json" + metadata = ( + json.loads(config_path.read_text()).get("quantization_config") + if config_path.exists() + else None + ) + if not is_modelopt_quant_config(metadata): + return {} + if read_modelopt_quant_config(metadata).get("kv_cache_quant_algo") != "NVFP4": + return {} + + from safetensors import safe_open + + values: dict[int, dict[str, list[float]]] = {} + for file_path in weight_files: + with safe_open(str(file_path), framework="pt", device="cpu") as checkpoint: + for tensor_name in checkpoint.keys(): + match = _MODEL_OPT_KV_SCALE_KEY.search(tensor_name) + if match is None: + continue + value = float(checkpoint.get_tensor(tensor_name).reshape([]).item()) + if not math.isfinite(value) or value <= 0.0: + raise ValueError( + f"ModelOpt KV scale {file_path}:{tensor_name} must be finite and positive" + ) + layer_values = values.setdefault(int(match.group("layer_id")), {"k": [], "v": []}) + layer_values[match.group("kind")].append(value) + + result: dict[int, LayerScales] = {} + for layer_id, layer_values in values.items(): + k_values, v_values = layer_values["k"], layer_values["v"] + if not k_values or not v_values: + raise ValueError(f"ModelOpt KV scales for layer {layer_id} must contain both K and V") + quant_orig = (max(k_values), max(v_values)) + result[layer_id] = ( + (1.0 / quant_orig[0], 1.0 / quant_orig[1]), + quant_orig, + ) + return result + + +class ColdPageQuantizationCompression(KVCacheCompressionManager): + """NVFP4 cold-page manager and owner of its optional model scales.""" + + uses_iteration_lifecycle = False + provides_cold_page_codec = True + + def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: + super().__init__(config) + self._model_nvfp4_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) + + def create_cold_page_codec( + self, + cache_config: object, + *, + runtime_dtype: DataType, + pp_layers: Sequence[int], + num_kv_heads_per_layer: Sequence[int], + head_dim_per_layer: Sequence[int], + ) -> object: + """Create the native codec that KVCM consumes exactly once.""" + + from tensorrt_llm.bindings.internal import kv_cache_compression as native + from tensorrt_llm.runtime.kv_cache_manager_v2 import SsmLayerConfig + + attention_layers = [] + for layer in cache_config.layers: + if isinstance(layer, SsmLayerConfig): + continue + roles = {buffer.role for buffer in layer.buffers} + if "key" not in roles: + raise NotImplementedError( + "NVFP4 cold-page compression requires an Attention key buffer" + ) + has_value = "value" in roles + attention_layers.append((layer, has_value)) + if not attention_layers: + return native.create_nvfp4_cold_page_codec([]) + + runtime_type = { + DataType.HALF: native.Nvfp4BoundaryRuntimeType.FLOAT16, + DataType.BF16: native.Nvfp4BoundaryRuntimeType.BFLOAT16, + DataType.FP8: native.Nvfp4BoundaryRuntimeType.FP8_E4M3, + }.get(runtime_dtype) + if runtime_type is None: + raise RuntimeError( + "NVFP4 cold-page compression supports FP16, BF16, or FP8 " + f"Attention KV, not {runtime_dtype}" + ) + + native_configs = [] + for layer, has_value in attention_layers: + layer_id = int(layer.layer_id) + if has_value: + orig_quant, quant_orig = self._model_nvfp4_scales.get( + int(pp_layers[layer_id]), _IDENTITY_NVFP4_SCALES + ) + else: + # ModelOpt K/V projection scales do not apply to MLA latent buffers. + orig_quant, quant_orig = _IDENTITY_NVFP4_SCALES + native_config = native.Nvfp4ColdPageLayerConfig() + native_config.layer_id = layer_id + native_config.runtime_type = runtime_type + native_config.num_kv_heads = int(num_kv_heads_per_layer[layer_id]) + native_config.tokens_per_page = int(cache_config.tokens_per_block) + native_config.head_dim = int(head_dim_per_layer[layer_id]) + native_config.nvfp4_scale_orig_quant = orig_quant + native_config.nvfp4_scale_quant_orig = quant_orig + native_configs.append(native_config) + + # The C++ factory is required for unique_ptr ownership transfer into KVCM. + return native.create_nvfp4_cold_page_codec(native_configs) diff --git a/tensorrt_llm/_torch/kv_cache_compression/triattention/triattention.py b/tensorrt_llm/_torch/kv_cache_compression/triattention/triattention.py index 8462a4bd8cc7..468c26eff699 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/triattention/triattention.py +++ b/tensorrt_llm/_torch/kv_cache_compression/triattention/triattention.py @@ -150,12 +150,10 @@ class TriAttentionCompressionManager(KVCacheCompressionManager): def __init__( self, config: "TriAttentionKvCacheCompressionConfig", - kv_cache_manager: KVCacheManagerV2, - draft_kv_cache_manager: Optional[KVCacheManagerV2] = None, *, pretrained_config: "PretrainedConfig", ) -> None: - super().__init__(config, kv_cache_manager, draft_kv_cache_manager) + super().__init__(config) self.budget = config.budget self.beta = config.beta self.eviction_mode = config.eviction_mode @@ -168,6 +166,14 @@ def __init__( self._load_calibration() self._prepared_generation_batch: Optional["ScheduledRequests"] = None + + def bind_kv_cache_managers( + self, + kv_cache_manager: KVCacheManagerV2, + draft_kv_cache_manager: Optional[KVCacheManagerV2] = None, + ) -> None: + """Finalize state whose geometry is owned by the constructed KVCMs.""" + super().bind_kv_cache_managers(kv_cache_manager, draft_kv_cache_manager) # Manager-lifetime constants. self._num_extra_kv_tokens = int(kv_cache_manager.num_extra_kv_tokens) self._protected_tail_capacity = ( @@ -420,7 +426,7 @@ def _evict_due_requests( # this request (pre-launch) instead of failing the batch. continue draft_cache = None - if self.draft_kv_cache_manager is not None: + if self.has_independent_draft_kv_cache: # A missing draft cache is a wiring bug: keep the precise KeyError. draft_cache = self.draft_kv_cache_manager.kv_cache_map[request_id] if not draft_cache.is_active: @@ -600,7 +606,7 @@ def cumulative_offsets(move_counts: List[int]) -> List[int]: if self._swa_window is not None: swa_offsets = cumulative_offsets([self._swa_window + tail for tail in tails]) draft_offsets = None - if self.draft_kv_cache_manager is not None: + if self.has_independent_draft_kv_cache: draft_offsets = cumulative_offsets( [self.budget + self._draft_protected_tail_capacity] * len(eviction_requests) ) @@ -681,7 +687,7 @@ def _initialize_eviction_state(self) -> None: """Create manager-lifetime state once.""" target_layout = self._create_kv_layout() draft_layout = ( - self._create_kv_layout(draft=True) if self.draft_kv_cache_manager is not None else None + self._create_kv_layout(draft=True) if self.has_independent_draft_kv_cache else None ) self._target_layout = target_layout self._draft_layout = draft_layout diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 5bedb58e526f..197f98ff992e 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1364,7 +1364,8 @@ def _create_kv_cache_manager( self, model_engine: PyTorchModelEngine, estimating_kv_cache: bool = False, - kv_cache_config_override: Optional[KvCacheConfig] = None + kv_cache_config_override: Optional[KvCacheConfig] = None, + cold_page_codec_provider: Optional[object] = None, ) -> KVCacheManager: mapping = self._mapping assert model_engine.model.model_config.is_generation, "Only construct KV cache for generation models." @@ -1402,6 +1403,7 @@ def _create_kv_cache_manager( execution_stream=self._execution_stream, layer_mask=spec_dec_layer_mask, is_disagg=self._is_disagg, + cold_page_codec_provider=cold_page_codec_provider, ) if not self._skip_est: @@ -1518,6 +1520,7 @@ def _create_one_model_draft_kv_cache_manager( max_seq_len: int, estimating_kv_cache: bool = False, kv_cache_config_override: Optional[KvCacheConfig] = None, + cold_page_codec_provider: Optional[object] = None, ) -> Optional[KVCacheManager]: """ Create a KV cache manager for draft model layers in one-model mode @@ -1589,6 +1592,7 @@ def _create_one_model_draft_kv_cache_manager( layer_mask=spec_dec_layer_mask, num_layers=num_draft_layers, is_disagg=self._is_disagg, + cold_page_codec_provider=cold_page_codec_provider, ) def _get_target_and_draft_cache_costs( @@ -2051,10 +2055,34 @@ def build_managers(self, budget_attr, self_kv_cache_config, draft_kv_cache_config)) + compression_manager = None + compression_config = self._llm_args.kv_cache_compression_config + if compression_config is not None: + create_compression_manager = True + if compression_config.algorithm == "quantization_for_cold_page": + if estimating_kv_cache and not self._skip_est: + create_compression_manager = False + elif _uses_nvfp4_kv_cache(self._model_engine): + logger.info( + "Skipping cold-page NVFP4 quantization because the active " + "KV cache already uses NVFP4; KVCM will migrate its native " + "data and block-scale buffers losslessly.") + create_compression_manager = False + if create_compression_manager: + model_config = self._model_engine.model.model_config + compression_manager = create_kv_cache_compression_manager( + compression_config, + pretrained_config=model_config.pretrained_config) + cold_page_codec_provider = ( + compression_manager if compression_manager is not None + and compression_manager.provides_cold_page_codec else None) + kv_cache_manager = self._create_kv_cache_manager( self._model_engine, estimating_kv_cache, - kv_cache_config_override=self_kv_cache_config) + kv_cache_config_override=self_kv_cache_config, + cold_page_codec_provider=cold_page_codec_provider, + ) # Carry the fp8 context-MLA workspace admission cap (computed in configure_kv_cache_capacity) onto # the real KV manager so the scheduler reads it directly instead of re-deriving from pool layout. @@ -2092,7 +2120,8 @@ def build_managers(self, draft_kv_cache_manager = self._create_one_model_draft_kv_cache_manager( original_max_seq_len, estimating_kv_cache, - kv_cache_config_override=draft_build_kv_cache_config) + kv_cache_config_override=draft_build_kv_cache_config, + cold_page_codec_provider=cold_page_codec_provider) # Encoder-decoder cross-attention pool cross_kv_cache_manager = None @@ -2106,9 +2135,16 @@ def build_managers(self, ResourceManagerType.DRAFT_KV_CACHE_MANAGER] = draft_kv_cache_manager resources[ ResourceManagerType.CROSS_KV_CACHE_MANAGER] = cross_kv_cache_manager + if compression_manager is not None: + resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] = ( + compression_manager) def teardown_managers(self, resources: Dict) -> None: """Clean up KV caches for model, draft model, and cross pool.""" + compression_manager = resources.pop( + ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER, None) + if compression_manager is not None: + compression_manager.shutdown() resources[ResourceManagerType.KV_CACHE_MANAGER].shutdown() del resources[ResourceManagerType.KV_CACHE_MANAGER] draft_kv_cache_manager = resources[ @@ -2220,11 +2256,18 @@ def _create_kv_cache_manager( num_kv_heads: Optional[Union[int, List[int]]] = None, head_dim: Optional[int] = None, kv_cache_type=None, - is_disagg: bool = False) -> KVCacheManager: + is_disagg: bool = False, + cold_page_codec_provider: Optional[object] = None) -> KVCacheManager: """ Returns: A KVCacheManager instance for the given model engine or model config """ + if cold_page_codec_provider is not None and not issubclass( + kv_cache_manager_cls, KVCacheManagerV2): + raise ValueError( + "Cold-page quantization requires the resolved KV cache manager " + f"to be KVCacheManagerV2; selected {kv_cache_manager_cls.__name__}") + if (estimating_kv_cache and issubclass(kv_cache_manager_cls, KVCacheManagerV2) and kv_cache_config.pool_ratio is None @@ -2351,6 +2394,8 @@ def _create_kv_cache_manager( manager_extra_kwargs = {} if issubclass(kv_cache_manager_cls, KVCacheManagerV2): manager_extra_kwargs["enable_stats"] = enable_kv_cache_stats + manager_extra_kwargs[ + "cold_page_codec_provider"] = cold_page_codec_provider if issubclass(kv_cache_manager_cls, MambaHybridCacheManagerV2): manager_extra_kwargs["is_disagg"] = is_disagg @@ -2752,6 +2797,18 @@ def validate_kv_cache_compression_compatibility( spec_config: Optional[SpeculativeConfig], ) -> None: """Reject unsupported KV-cache compression feature combinations.""" + if config.algorithm == "quantization_for_cold_page": + from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND + + if _BACKEND == "python": + raise ValueError( + "Cold-page quantization requires the C++ KVCacheManagerV2 backend" + ) + if not is_sm_100f(): + raise RuntimeError( + "NVFP4 cold-page compression requires an SM100-family device " + "(SM100 or SM103).") + if kv_cache_config.enable_block_reuse and not config.supports_block_reuse(): raise ValueError( f"KV-cache compression algorithm {config.algorithm!r} does not " @@ -2765,25 +2822,41 @@ def validate_kv_cache_compression_compatibility( "support speculative decoding with its current configuration; " "TriAttention requires eviction_mode='union'") mode = spec_config.spec_dec_mode + if config.algorithm == "quantization_for_cold_page": + if not (mode.is_eagle3_one_model() + or mode.is_mtp_eagle_one_model()): + raise ValueError( + "Cold-page quantization supports speculative decoding only " + f"with one-model MTP-EAGLE or EAGLE3, not {mode.name}") + return if not (mode.is_mtp_one_model() or mode.is_eagle3_one_model()): raise ValueError( f"KV-cache compression does not support speculative decoding " f"mode {mode.name}; use one-model MTP or EAGLE3") +def _uses_nvfp4_kv_cache(model_engine: PyTorchModelEngine) -> bool: + quant_config = model_engine.model.model_config.quant_config + return (quant_config is not None + and getattr(quant_config, "kv_cache_quant_algo", None) == "NVFP4") + + def create_kv_cache_compression_manager( config: KvCacheCompressionConfig, - kv_cache_manager: KVCacheManagerV2, - draft_kv_cache_manager: Optional[KVCacheManagerV2] = None, pretrained_config: Optional["transformers.PretrainedConfig"] = None, ) -> Optional[KVCacheCompressionManager]: - """Build the KV-cache compression manager for ``config.algorithm``, or return - None if no algorithm matches. + """Build a KV-cache compression manager before KVCM construction. - Called from ``create_py_executor`` and registered as a resource manager, - like the KV cache manager itself. Concrete algorithms add a dispatch branch - here. Feature compatibility is checked before resource-manager construction. + The caller may use the manager as a cold-page codec provider while building + KVCMs, then binds the completed target and optional draft managers. Only + iteration-driven implementations enter ``ResourceManager``. """ + if config.algorithm == "quantization_for_cold_page": + from ..kv_cache_compression.quantization_for_cold_page import \ + ColdPageQuantizationCompression + + return ColdPageQuantizationCompression(config) + if config.algorithm == "triattention": if not is_sm_100f(): raise RuntimeError( @@ -2795,8 +2868,6 @@ def create_kv_cache_compression_manager( return TriAttentionCompressionManager( config, - kv_cache_manager, - draft_kv_cache_manager=draft_kv_cache_manager, pretrained_config=pretrained_config, ) @@ -3067,24 +3138,15 @@ def create_py_executor_instance( resources[ResourceManagerType.SEQ_SLOT_MANAGER] = SeqSlotManager( max_num_sequences) - # Register the compression manager (if one is configured) with the other - # managers, before building ResourceManager, so it is part of the manager - # set from the start. Reads its own config, not the sparse-attention one. - kv_cache_compression_config = getattr(llm_args, - "kv_cache_compression_config", None) - if kv_cache_compression_config is not None: - draft_kv_cache_manager = resources.get( - ResourceManagerType.DRAFT_KV_CACHE_MANAGER) - compression_manager = create_kv_cache_compression_manager( - kv_cache_compression_config, - kv_cache_manager, - draft_kv_cache_manager=draft_kv_cache_manager, - pretrained_config=model_engine.model.model_config.pretrained_config, + compression_manager = resources.get( + ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER) + if compression_manager is not None: + compression_manager.bind_kv_cache_managers( + resources[ResourceManagerType.KV_CACHE_MANAGER], + resources.get(ResourceManagerType.DRAFT_KV_CACHE_MANAGER), ) - if compression_manager is not None: - resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] = ( - compression_manager) - + if not compression_manager.uses_iteration_lifecycle: + resources.pop(ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER) resource_manager = ResourceManager(resources) # KV cache manager runs last (others may depend on it), except the @@ -3097,7 +3159,8 @@ def create_py_executor_instance( if cross_kv_cache_manager is not None: resource_manager.resource_managers.move_to_end( ResourceManagerType.CROSS_KV_CACHE_MANAGER, last=True) - # Compression is the final reconciler after every native KV manager. + # Iteration-driven compression is the final reconciler after every native + # KV manager. Boundary quantization runs only at native storage migration. if (ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER in resource_manager.resource_managers): resource_manager.resource_managers.move_to_end( @@ -3586,7 +3649,11 @@ def _adjust_torch_mem_fraction(): def validate_feature_combination(llm_args, model_engine, sampler_type): # Validate the flags for features' combination compression_config = llm_args.kv_cache_compression_config - if compression_config is not None: + cold_compression_is_redundant = (compression_config is not None + and compression_config.algorithm + == "quantization_for_cold_page" + and _uses_nvfp4_kv_cache(model_engine)) + if compression_config is not None and not cold_compression_is_redundant: validate_kv_cache_compression_compatibility( compression_config, llm_args.kv_cache_config, diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index f9aa027ed44b..635a33dd3d4a 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -799,6 +799,7 @@ def __init__( enable_stats: bool = False, num_reserved_index_slots: int = 1, is_estimating_kv_cache: bool = False, + cold_page_codec_provider=None, **kwargs, ) -> None: self.mapping = mapping @@ -1119,8 +1120,26 @@ def append_to_kv_heads_per_layer( self.kv_cache_manager_py_config = config + def create_impl(cache_config): + manager_kwargs = {} + if cold_page_codec_provider is not None: + manager_kwargs["cold_page_codec"] = ( + cold_page_codec_provider.create_cold_page_codec( + cache_config, + runtime_dtype=self.dtype, + pp_layers=self.pp_layers, + num_kv_heads_per_layer=self.num_kv_heads_per_layer, + head_dim_per_layer=self.head_dim_per_layer, + ) + ) + return KVCacheManagerPy( + cache_config, + event_manager=self.event_manager, + **manager_kwargs, + ) + try: - self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) + self.impl = create_impl(config) except (CuError, KVCacheOutOfMemoryError): if has_host_cache_tier: logger.warning( @@ -1133,7 +1152,7 @@ def append_to_kv_heads_per_layer( ] config = replace(config, cache_tiers=cache_tiers_without_host) self.kv_cache_manager_py_config = config - self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) + self.impl = create_impl(config) else: raise self.can_evict = len(config.cache_tiers) > 1 diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 3cf6ab2e852a..737300f52ed9 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -2715,30 +2715,22 @@ def _free_blocks(self, block_list: list): class KVCacheCompressionManager(BaseResourceManager): - """Framework-level base class for all KV-cache compression managers. - - Inherits :class:`BaseResourceManager` so PyExecutor's main loop - auto-invokes ``prepare_resources`` / ``update_resources`` / - ``free_resources`` each iteration without any PyExecutor code changes; the - base implementations below translate those callbacks into the lifecycle - hooks. - - Concrete compression methods subclass this directly. The hooks default to - no-op; subclasses override what they need. The manager never inherits from - any cache manager because this layer decides *how* the physical KV is used, - not *what* physical KV exists. Subclasses hold ``KVCacheManagerV2`` as a tool. - - A subclass compacts through the ``KVCacheManagerV2`` it holds and records - the evicted count on ``LlmRequest.py_num_compressed_tokens``; the model - engine subtracts that count when building ``num_cached_tokens_per_seq``. - """ + """Base for iteration-driven and storage-boundary KV compression.""" - def __init__( + uses_iteration_lifecycle = True + provides_cold_page_codec = False + + def __init__(self, config: "KvCacheCompressionConfig") -> None: + self.config = config + self.kv_cache_manager: Optional["KVCacheManagerV2"] = None + self.draft_kv_cache_manager: Optional["KVCacheManagerV2"] = None + + def bind_kv_cache_managers( self, - config: "KvCacheCompressionConfig", kv_cache_manager: "KVCacheManagerV2", draft_kv_cache_manager: Optional["KVCacheManagerV2"] = None, - ): + ) -> None: + """Bind the target and optional draft KVCMs after their construction.""" from .kv_cache_manager_v2 import KVCacheManagerV2 if not isinstance(kv_cache_manager, KVCacheManagerV2): @@ -2750,16 +2742,27 @@ def __init__( self.kv_cache_manager = kv_cache_manager self.draft_kv_cache_manager = draft_kv_cache_manager kv_cache_manager.kv_compression_manages_history = ( - config.changes_physical_kv_length) + self.config.changes_physical_kv_length) if draft_kv_cache_manager is not None: - # The draft cache is compacted together with the target. draft_kv_cache_manager.kv_compression_manages_history = ( - config.changes_physical_kv_length) + self.config.changes_physical_kv_length) @property def has_independent_draft_kv_cache(self) -> bool: return self.draft_kv_cache_manager is not None + def create_cold_page_codec( + self, + cache_config: object, + *, + runtime_dtype: DataType, + pp_layers: Sequence[int], + num_kv_heads_per_layer: Sequence[int], + head_dim_per_layer: Sequence[int], + ) -> Optional[object]: + """Create a native storage-boundary codec when the algorithm provides one.""" + return None + # ================================================================== # # KV-cache lifecycle hooks (5, in temporal order). # # Subclasses override what they need; all default to no-op. # diff --git a/tensorrt_llm/llmapi/__init__.py b/tensorrt_llm/llmapi/__init__.py index aff2e9736977..e4520f2aadfa 100644 --- a/tensorrt_llm/llmapi/__init__.py +++ b/tensorrt_llm/llmapi/__init__.py @@ -8,9 +8,10 @@ # yapf: disable from .llm_args import (AttentionDpConfig, AutoDecodingConfig, BatchingType, BlockReuseConfig, CacheTransceiverConfig, CalibConfig, - CapacitySchedulerPolicy, ContextChunkingPolicy, - CudaGraphConfig, DecodeCudaGraphConfig, - DeepSeekSparseAttentionConfig, + CapacitySchedulerPolicy, + ColdPageQuantizationCompressionConfig, + ContextChunkingPolicy, CudaGraphConfig, + DecodeCudaGraphConfig, DeepSeekSparseAttentionConfig, DeepSeekV4SparseAttentionConfig, DFlashDecodingConfig, DraftTargetDecodingConfig, DSparkDecodingConfig, DynamicBatchConfig, Eagle3DecodingConfig, @@ -91,6 +92,7 @@ 'MiniMaxM3SparseAttentionConfig', 'SchedulingParams', 'SkipSoftmaxAttentionConfig', + 'ColdPageQuantizationCompressionConfig', 'TriAttentionKvCacheCompressionConfig', 'PrometheusMetricsConfig', 'PrefillCudaGraphBackend', diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 51103ea4767f..3253c6c2cae0 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3684,9 +3684,10 @@ class KvCacheCompressionConfig(StrictBaseModel): algorithm (e.g. periodic token eviction) alongside KVCacheManagerV2. Kept separate from SparseAttentionConfig by design -- compression changes - which KV is stored, not the attention computation. The manager is registered - as a resource manager in create_py_executor (_util.py), like the KV cache - manager itself. Concrete algorithms subclass this and add their parameters. + which KV is stored, not the attention computation. Iteration-driven methods + use the resource-manager cycle; storage-boundary managers provide a native + codec that KVCacheManagerV2 retains and invokes. Concrete algorithms + subclass this and add their parameters. """ changes_physical_kv_length: ClassVar[bool] = False @@ -3705,6 +3706,30 @@ def supports_speculative_decoding(self) -> bool: return False +class ColdPageQuantizationCompressionConfig(KvCacheCompressionConfig): + """Quantize Host and Disk KV pages without changing the active GPU cache.""" + + algorithm: Literal["quantization_for_cold_page"] = "quantization_for_cold_page" + quant: Literal["nvfp4"] = Field( + default="nvfp4", + description="Quantization format stored in the compressed cache tier.") + scale_checkpoint_path: Optional[str] = Field( + default=None, + min_length=1, + telemetry=False, + description= + "Optional local ModelOpt NVFP4 checkpoint directory supplying per-layer " + "K/V global scales. Omit it to use identity global scales.") + + def supports_block_reuse(self) -> bool: + # Compression changes representation and residency, not token identity. + return True + + def supports_speculative_decoding(self) -> bool: + # Each target or draft KVCM encodes its own pages at the storage boundary. + return True + + class TriAttentionKvCacheCompressionConfig(KvCacheCompressionConfig): """TriAttention KV-cache compression: periodic decode-time eviction. @@ -3757,7 +3782,8 @@ def supports_speculative_decoding(self) -> bool: KvCacheCompressionConfigType: TypeAlias = Annotated[ - Union[TriAttentionKvCacheCompressionConfig], + Union[ColdPageQuantizationCompressionConfig, + TriAttentionKvCacheCompressionConfig], Field(discriminator="algorithm"), ] diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index eaab417947d6..b42702f438b8 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -299,7 +299,7 @@ "decode", "encode" ], - "annotation": "Literal['decode']", + "annotation": "Literal['decode', 'encode']", "converter": "", "kind": "categorical", "path": "cuda_graph_config.mode" @@ -631,9 +631,10 @@ }, { "allowed_values": [ + "quantization_for_cold_page", "triattention" ], - "annotation": "Literal['triattention']", + "annotation": "Literal['quantization_for_cold_page', 'triattention']", "converter": "", "kind": "categorical", "path": "kv_cache_compression_config.algorithm" @@ -670,6 +671,15 @@ "kind": "value", "path": "kv_cache_compression_config.normalize_scores" }, + { + "allowed_values": [ + "nvfp4" + ], + "annotation": "Literal['nvfp4']", + "converter": "", + "kind": "categorical", + "path": "kv_cache_compression_config.quant" + }, { "allowed_values": [], "annotation": "", @@ -1548,7 +1558,7 @@ "rocket", "skip_softmax" ], - "annotation": "Literal['dsa']", + "annotation": "Literal['dsa', 'deepseek_v4', 'minimax_m3', 'rocket', 'skip_softmax']", "converter": "", "kind": "categorical", "path": "sparse_attention_config.algorithm" @@ -1872,7 +1882,7 @@ "SaveState", "User_Provided" ], - "annotation": "Literal['AUTO']", + "annotation": "Literal['AUTO', 'DFlash', 'DSpark', 'Draft_Target', 'Eagle3', 'Eagle', 'Lookahead', 'MTP', 'Medusa', 'NGram', 'PARD', 'SA', 'SaveState', 'User_Provided']", "converter": "", "kind": "categorical", "path": "speculative_config.decoding_type" diff --git a/tensorrt_llm/usage/llmapi_config.py b/tensorrt_llm/usage/llmapi_config.py index 25c661af8d61..7678a8308688 100644 --- a/tensorrt_llm/usage/llmapi_config.py +++ b/tensorrt_llm/usage/llmapi_config.py @@ -305,9 +305,10 @@ def build_capture_manifest(model_cls: type[BaseModel]) -> list[_ManifestEntry]: The single source of truth. Type-safe annotations auto-enroll; str/Any allowlist escape hatches opt in; telemetry=False opts out. Recurses into statically reachable nested BaseModels with a cycle guard. Collapses - duplicate keys (shared union-arm base fields): keeps the first by - (key, defining_class), unions allowed_values across arms, and FAILS if two - arms give a key a different kind. + duplicate keys (shared union-arm base fields): unions Literal annotations + and allowed_values across arms, and FAILS if two arms give a key a + different kind. Other duplicate annotations retain the first deterministic + representative. """ rows: list[dict[str, Any]] = [] @@ -339,6 +340,7 @@ def walk(cls: type, prefix: str, stack: tuple) -> None: rows.sort(key=lambda r: (r["key"], r["defining"])) first: dict[str, dict] = {} + union_annotations: dict[str, list[Any]] = {} union_allowed: dict[str, list[str]] = {} for r in rows: if r["key"] not in first: @@ -348,15 +350,27 @@ def walk(cls: type, prefix: str, stack: tuple) -> None: f"telemetry manifest: key '{r['key']}' has conflicting kinds " f"across union arms: {first[r['key']]['kind']} vs {r['kind']}" ) + annotations = union_annotations.setdefault(r["key"], []) + if r["annotation"] not in annotations: + annotations.append(r["annotation"]) seen = union_allowed.setdefault(r["key"], []) for v in r["allowed"]: if v not in seen: seen.append(v) + def merged_annotation(key: str) -> Any: + annotations = union_annotations[key] + if len(annotations) > 1 and all(_is_literal(ann) for ann in annotations): + values = tuple( + value for ann in annotations for value in get_args(ann) + ) + return Literal[values] + return annotations[0] + entries = [ _ManifestEntry( path=key, - annotation=r["annotation"], + annotation=merged_annotation(key), kind=r["kind"], converter=r["converter"], allowed_values=tuple(union_allowed[key]), diff --git a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py index 477e459b3e2c..530109421692 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py +++ b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py @@ -64,6 +64,7 @@ def _make_creator( c._speculative_config = None c._mapping = Mock() c._model_engine = Mock() + c._llm_args = SimpleNamespace(kv_cache_compression_config=None) c._kv_cache_manager_cls = Mock() c._kv_cache_manager_cls.get_cache_size_per_token = Mock( @@ -449,7 +450,6 @@ def test_no_disk_cache_leaves_none(self): assert target_config is c._kv_cache_config assert draft_config is None - class TestHostSplitIgnoresGpuFixedCost: """The fixed cost models GPU-resident state and is not host memory.""" diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index 4f628d6fbb94..d7fa9ffb4479 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -1,26 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Unit tests for the KV-cache compression manager framework -(``KVCacheCompressionManager`` in ``resource_manager.py``) — the -``BaseResourceManager``-based single-manager design. - -Covers: -- :class:`KVCacheCompressionManager` contract: the four lifecycle hooks - default to no-op, zero resource counts, and it inherits - :class:`BaseResourceManager` (so PyExecutor auto-drives it once registered). -- The resource-manager API -> lifecycle-hook translation, gated on PyExecutor's - own signals: ``prepare_resources`` fires ``on_request_init`` on each - request's first prefill chunk (``is_first_context_chunk``); - ``update_resources`` fires ``on_context_step_end`` once with the - ``context_requests_last_chunk`` list + one ``on_generation_step_end`` per - iteration; ``free_resources`` fires ``on_request_finish``. -- :func:`create_kv_cache_compression_manager` factory. - -The base class lives in ``resource_manager.py`` (it is a resource manager, not a -sparse-attention backend); the ``create_kv_cache_compression_manager`` factory -lives in ``_util.py`` next to ``_create_kv_cache_manager``. -""" +"""Tests for the KV-cache compression manager lifecycle and factory.""" from types import SimpleNamespace from typing import ClassVar @@ -48,7 +29,8 @@ class _RecordingMixin: translation without real algorithm side-effects.""" def __init__(self, kv_cache_manager, record_list, name="m"): - super().__init__(_compression_config(), kv_cache_manager) + super().__init__(_compression_config()) + self.bind_kv_cache_managers(kv_cache_manager) self._record_list = record_list self._name = name @@ -57,7 +39,7 @@ def _record(self, hook_name: str): class _MockCompressionManager(_RecordingMixin, KVCacheCompressionManager): - """Mock manager that records the four lifecycle hooks.""" + """Mock manager that records iteration lifecycle hooks.""" def on_request_init(self, request): self._record("on_request_init") @@ -126,8 +108,9 @@ def test_inherits_base_resource_manager(self): # So PyExecutor's main loop auto-invokes prepare/update/free_resources. assert issubclass(KVCacheCompressionManager, BaseResourceManager) - def test_four_hooks_default_noop(self, fake_kv_cache_manager): - m = KVCacheCompressionManager(_compression_config(), fake_kv_cache_manager) + def test_lifecycle_hooks_default_noop(self, fake_kv_cache_manager): + m = KVCacheCompressionManager(_compression_config()) + m.bind_kv_cache_managers(fake_kv_cache_manager) assert m.on_request_init(MagicMock()) is None assert m.on_context_step_end([MagicMock()]) is None assert m.on_generation_step_begin(MagicMock()) is None @@ -137,12 +120,14 @@ def test_four_hooks_default_noop(self, fake_kv_cache_manager): def test_hooks_accept_extra_kwargs(self, fake_kv_cache_manager): # **kwargs lets the framework pass new args later without breaking # existing overrides. - m = KVCacheCompressionManager(_compression_config(), fake_kv_cache_manager) + m = KVCacheCompressionManager(_compression_config()) + m.bind_kv_cache_managers(fake_kv_cache_manager) assert m.on_request_init(MagicMock(), future_arg=1) is None assert m.on_generation_step_end(MagicMock(), future_arg=1) is None def test_resource_counts_are_zero(self, fake_kv_cache_manager): - m = KVCacheCompressionManager(_compression_config(), fake_kv_cache_manager) + m = KVCacheCompressionManager(_compression_config()) + m.bind_kv_cache_managers(fake_kv_cache_manager) # The manager owns no physical resources (the V2 cache manager does), # so it must not gate the scheduler. assert m.get_max_resource_count() == 0 @@ -155,7 +140,8 @@ def test_physical_length_change_marks_target_and_draft_v2(self): draft = _v2_manager(is_draft=True) config = _PhysicalLengthChangingConfig(algorithm="test") - manager = KVCacheCompressionManager(config, target, draft) + manager = KVCacheCompressionManager(config) + manager.bind_kv_cache_managers(target, draft) assert manager.kv_cache_manager is target assert manager.draft_kv_cache_manager is draft @@ -166,9 +152,11 @@ def test_physical_length_change_marks_target_and_draft_v2(self): def test_rejects_non_v2_ownership(self): config = _compression_config() with pytest.raises(TypeError, match="requires KVCacheManagerV2"): - KVCacheCompressionManager(config, MagicMock()) + KVCacheCompressionManager(config).bind_kv_cache_managers(MagicMock()) with pytest.raises(TypeError, match="requires KVCacheManagerV2"): - KVCacheCompressionManager(config, _v2_manager(is_draft=False), MagicMock()) + KVCacheCompressionManager(config).bind_kv_cache_managers( + _v2_manager(is_draft=False), MagicMock() + ) def test_request_field_defaults_to_zero(self): """LlmRequest carries the compression count (the manager's only @@ -276,42 +264,26 @@ def test_free_fires_finish(self, fake_kv_cache_manager): class TestFactory: - def test_returns_none_when_no_algorithm_registered(self, fake_kv_cache_manager): - # Framework-only: no concrete algorithm ships, so any config -> None. + def test_returns_none_when_no_algorithm_registered(self): cfg = MagicMock() cfg.algorithm = "made_up_method" - assert create_kv_cache_compression_manager(cfg, fake_kv_cache_manager) is None + assert create_kv_cache_compression_manager(cfg) is None - def test_warns_for_unregistered_algorithm(self, fake_kv_cache_manager): + def test_warns_for_unregistered_algorithm(self): cfg = MagicMock() cfg.algorithm = "made_up_method" with patch.object(util_mod, "logger") as mock_logger: - create_kv_cache_compression_manager(cfg, fake_kv_cache_manager) + create_kv_cache_compression_manager(cfg) mock_logger.warning.assert_called_once() - def test_factory_accepts_independent_draft_manager(self): - cfg = MagicMock() - cfg.algorithm = "made_up_method" - target = _v2_manager(is_draft=False) - draft = _v2_manager(is_draft=True) - - assert ( - create_kv_cache_compression_manager( - cfg, - target, - draft_kv_cache_manager=draft, - ) - is None - ) - - def test_triattention_requires_sm100_family(self, fake_kv_cache_manager): + def test_triattention_requires_sm100_family(self): cfg = MagicMock() cfg.algorithm = "triattention" with ( patch.object(util_mod, "is_sm_100f", return_value=False), pytest.raises(RuntimeError, match="SM100-family"), ): - create_kv_cache_compression_manager(cfg, fake_kv_cache_manager) + create_kv_cache_compression_manager(cfg) def test_capabilities_default_false(self): config = KvCacheCompressionConfig(algorithm="offload") @@ -319,7 +291,8 @@ def test_capabilities_default_false(self): assert config.changes_physical_kv_length is False assert config.supports_block_reuse() is False assert config.supports_speculative_decoding() is False - m = KVCacheCompressionManager(config, target) + m = KVCacheCompressionManager(config) + m.bind_kv_cache_managers(target) assert target.kv_compression_manages_history is False assert not hasattr(m, "spec_config") @@ -343,6 +316,127 @@ def test_spec_gate_uses_config_capability(self): ) +class TestKvCacheCreatorLifecycle: + def test_estimation_still_creates_triattention_manager(self): + config = SimpleNamespace(algorithm="triattention") + pretrained_config = object() + expected_manager = SimpleNamespace(provides_cold_page_codec=False) + creator = object.__new__(util_mod.KvCacheCreator) + creator._skip_est = False + creator._max_seq_len = 1024 + creator._kv_cache_config = SimpleNamespace() + creator._llm_args = SimpleNamespace(kv_cache_compression_config=config) + creator._model_engine = SimpleNamespace( + model=SimpleNamespace(model_config=SimpleNamespace(pretrained_config=pretrained_config)) + ) + creator._draft_model_engine = None + creator._is_encoder_decoder = MagicMock(return_value=False) + creator._should_create_separate_draft_kv_cache = MagicMock(return_value=False) + target_manager = object() + creator._create_kv_cache_manager = MagicMock(return_value=target_manager) + + with patch.object( + util_mod, + "create_kv_cache_compression_manager", + return_value=expected_manager, + ) as factory: + resources = {} + creator.build_managers(resources, estimating_kv_cache=True) + + factory.assert_called_once_with( + config, + pretrained_config=pretrained_config, + ) + assert resources[ResourceManagerType.KV_CACHE_MANAGER] is target_manager + assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is expected_manager + + def test_teardown_pops_and_shuts_down_compression_manager(self): + creator = object.__new__(util_mod.KvCacheCreator) + compression_manager = MagicMock() + resources = { + ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER: compression_manager, + ResourceManagerType.KV_CACHE_MANAGER: MagicMock(), + ResourceManagerType.DRAFT_KV_CACHE_MANAGER: None, + } + + creator.teardown_managers(resources) + + compression_manager.shutdown.assert_called_once_with() + assert ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in resources + + +@pytest.mark.parametrize( + "provides_cold_page_codec", + (True, False), + ids=("cold-codec", "iteration-manager"), +) +def test_build_routes_compression_manager_by_capabilities(provides_cold_page_codec): + creator = object.__new__(util_mod.KvCacheCreator) + creator._skip_est = False + creator._max_seq_len = 1024 + creator._kv_cache_config = SimpleNamespace() + compression_config = SimpleNamespace(algorithm="triattention") + pretrained_config = object() + creator._llm_args = SimpleNamespace(kv_cache_compression_config=compression_config) + creator._model_engine = SimpleNamespace( + model=SimpleNamespace(model_config=SimpleNamespace(pretrained_config=pretrained_config)) + ) + creator._draft_model_engine = None + creator._kv_connector_manager = None + creator._is_kv_cache_manager_v2 = True + creator._fp8_ctx_mla_kv_len_cap = None + creator._is_encoder_decoder = MagicMock(return_value=False) + creator._should_create_separate_draft_kv_cache = MagicMock(return_value=True) + creator._needs_gpu_kv_cache_budget_split = MagicMock(return_value=False) + + build_order = [] + compression_manager = SimpleNamespace( + provides_cold_page_codec=provides_cold_page_codec, + ) + target_config = object() + draft_config = object() + target_manager = SimpleNamespace() + draft_manager = object() + creator._split_kv_cache_budget_for_draft = MagicMock( + side_effect=[(target_config, draft_config), (target_config, draft_config)] + ) + creator._create_kv_cache_manager = MagicMock( + side_effect=lambda *_args, **_kwargs: build_order.append("target") or target_manager + ) + creator._create_one_model_draft_kv_cache_manager = MagicMock( + side_effect=lambda *_args, **_kwargs: build_order.append("draft") or draft_manager + ) + + resources = {} + with patch.object( + util_mod, + "create_kv_cache_compression_manager", + side_effect=lambda *_args, **_kwargs: build_order.append("factory") + or compression_manager, + ) as factory: + creator.build_managers(resources) + + expected_codec_provider = compression_manager if provides_cold_page_codec else None + factory.assert_called_once_with( + compression_config, + pretrained_config=pretrained_config, + ) + assert ( + creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"] + is expected_codec_provider + ) + assert ( + creator._create_one_model_draft_kv_cache_manager.call_args.kwargs[ + "cold_page_codec_provider" + ] + is expected_codec_provider + ) + assert resources[ResourceManagerType.KV_CACHE_MANAGER] is target_manager + assert resources[ResourceManagerType.DRAFT_KV_CACHE_MANAGER] is draft_manager + assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is compression_manager + assert build_order == ["factory", "target", "draft"] + + # ---------------------------------------------------------------------- # # 4. Canonical names live in resource_manager, not in the sparse module # # ---------------------------------------------------------------------- # diff --git a/tests/unittest/_torch/executor/test_kv_cache_estimation.py b/tests/unittest/_torch/executor/test_kv_cache_estimation.py index 9f0a4b9ac0df..663fa065177c 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_estimation.py +++ b/tests/unittest/_torch/executor/test_kv_cache_estimation.py @@ -948,8 +948,13 @@ def test_separate_one_model_draft_normalizes_target_pool_ratio() -> None: return_value=Mock(), ) as create_manager, ): - creator._create_one_model_draft_kv_cache_manager(creator._max_seq_len) + codec_provider = object() + creator._create_one_model_draft_kv_cache_manager( + creator._max_seq_len, + cold_page_codec_provider=codec_provider, + ) draft_config = create_manager.call_args.kwargs["kv_cache_config"] assert draft_config.pool_ratio == [1.0] + assert create_manager.call_args.kwargs["cold_page_codec_provider"] is codec_provider assert creator._kv_cache_config.pool_ratio == target_pool_ratio diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index 3b6df54cb3bf..f86a30cad85c 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -110,6 +110,7 @@ def _make_manager_for_cache_tier_test( impl_side_effect: list[object], *, add_secondary_gpu_tier: bool = False, + cold_page_codec_provider: object | None = None, ) -> tuple[KVCacheManagerV2, Mock]: impl_constructor = Mock(side_effect=impl_side_effect) @@ -167,6 +168,7 @@ def build_cache_config( dtype=DataType.HALF, vocab_size=16, execution_stream=Mock(), + cold_page_codec_provider=cold_page_codec_provider, ) return manager, impl_constructor @@ -368,6 +370,34 @@ def test_host_init_fallback_drops_only_host_tier(tmp_path) -> None: ] +def test_host_init_fallback_recreates_cold_codec_and_keeps_disk(tmp_path) -> None: + impl = Mock() + codecs = [object(), object()] + codec_provider = Mock() + codec_provider.create_cold_page_codec.side_effect = codecs + manager, impl_constructor = _make_manager_for_cache_tier_test( + KvCacheConfig( + max_gpu_total_bytes=16 << 20, + host_cache_size=16 << 20, + disk_cache_size=16 << 20, + disk_cache_path=str(tmp_path), + ), + [_CacheTierInitError("host tier init failed"), impl], + cold_page_codec_provider=codec_provider, + ) + + assert manager.can_evict + assert codec_provider.create_cold_page_codec.call_count == 2 + assert impl_constructor.call_count == 2 + assert impl_constructor.call_args_list[0].kwargs["cold_page_codec"] is codecs[0] + assert impl_constructor.call_args_list[1].kwargs["cold_page_codec"] is codecs[1] + fallback_tiers = impl_constructor.call_args_list[1].args[0].cache_tiers + assert [type(tier) for tier in fallback_tiers] == [ + GpuCacheTierConfig, + DiskCacheTierConfig, + ] + + def test_extra_tokens_are_in_context_capacity() -> None: config = _make_cache_config_for_test( KvCacheConfig(avg_seq_len=264), diff --git a/tests/unittest/_torch/kv_cache_compression/conftest.py b/tests/unittest/_torch/kv_cache_compression/conftest.py index 1ab526d83d40..a823dfd02386 100644 --- a/tests/unittest/_torch/kv_cache_compression/conftest.py +++ b/tests/unittest/_torch/kv_cache_compression/conftest.py @@ -348,11 +348,12 @@ def make_triattention(**overrides): ) with mock.patch.object(TriAttentionCompressionManager, "_initialize_eviction_state"): - return TriAttentionCompressionManager( + manager = TriAttentionCompressionManager( make_tri_config(**overrides), - make_fake_v2(), pretrained_config=make_test_pretrained_config(), ) + manager.bind_kv_cache_managers(make_fake_v2()) + return manager def make_eviction_request( diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py new file mode 100644 index 000000000000..f07d0dee8646 --- /dev/null +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -0,0 +1,417 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Control-plane tests for NVFP4 cold-page compression.""" + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +import torch +from safetensors.torch import save_file + +from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page import ( + ColdPageQuantizationCompression, +) +from tensorrt_llm._torch.pyexecutor import _util as util_mod +from tensorrt_llm._torch.pyexecutor.resource_manager import DataType +from tensorrt_llm._torch.speculative.interface import SpeculativeDecodingMode +from tensorrt_llm._torch.speculative.utils import update_spec_config_from_model_config +from tensorrt_llm.llmapi.llm_args import ( + ColdPageQuantizationCompressionConfig, + MTPDecodingConfig, +) +from tensorrt_llm.runtime import kv_cache_manager_v2 as runtime_v2_mod +from tensorrt_llm.runtime.kv_cache_manager_v2 import ( + AttentionLayerConfig, + BufferConfig, + SsmLayerConfig, +) + + +def _manager(scale_checkpoint_path=None): + config = ColdPageQuantizationCompressionConfig( + scale_checkpoint_path=( + str(scale_checkpoint_path) if scale_checkpoint_path is not None else None + ) + ) + return ColdPageQuantizationCompression(config) + + +def _cache_config(*layers): + configs = [] + for layer_id, kind in layers: + layer_type = SsmLayerConfig if kind == "ssm" else AttentionLayerConfig + roles = ("ssm_state", "conv_state") if kind == "ssm" else ("key", "value") + configs.append( + layer_type( + layer_id=layer_id, + buffers=[BufferConfig(role=role, size=128) for role in roles], + ) + ) + return SimpleNamespace(tokens_per_block=64, layers=tuple(configs)) + + +def _native(): + def layer_config(): + return SimpleNamespace( + fp8_scale_orig_quant=(1.0, 1.0), + fp8_scale_quant_orig=(1.0, 1.0), + ) + + codec = MagicMock() + module = SimpleNamespace( + Nvfp4BoundaryRuntimeType=SimpleNamespace( + FLOAT16="native-fp16", + BFLOAT16="native-bf16", + FP8_E4M3="native-fp8", + ), + Nvfp4ColdPageLayerConfig=layer_config, + create_nvfp4_cold_page_codec=MagicMock(return_value=codec), + ) + return module, codec + + +def _write_quant_metadata(directory, algorithm="NVFP4"): + metadata = { + "producer": {"name": "modelopt"}, + "quantization": {"kv_cache_quant_algo": algorithm}, + } + (directory / "hf_quant_config.json").write_text(json.dumps(metadata)) + + +def _write_scales(directory, scales_by_layer, *, filename="model.safetensors", prefix="model"): + _write_quant_metadata(directory) + tensors = {} + for layer_id, (k_scale, v_scale) in scales_by_layer.items(): + base = f"{prefix}.layers.{layer_id}.self_attn" + tensors[f"{base}.k_proj.k_scale"] = torch.as_tensor(k_scale, dtype=torch.float32) + tensors[f"{base}.v_proj.v_scale"] = torch.as_tensor(v_scale, dtype=torch.float32) + save_file(tensors, str(directory / filename)) + + +def _validate_compression(mode=None): + spec_config = None if mode is None else SimpleNamespace(spec_dec_mode=mode) + util_mod.validate_kv_cache_compression_compatibility( + ColdPageQuantizationCompressionConfig(), + SimpleNamespace(enable_block_reuse=False), + spec_config, + ) + + +def test_optional_modelopt_scales_map_pp_layers_and_ignore_local_draft_id(tmp_path): + native, codec = _native() + _write_scales( + tmp_path, + {10: (0.5, 0.25)}, + filename="model-00001-of-00002.safetensors", + ) + _write_scales( + tmp_path, + {4: (0.125, 0.0625), 2: (0.75, 0.5)}, + filename="model-00002-of-00002.safetensors", + prefix="model.language_model", + ) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + result = _manager(tmp_path).create_cold_page_codec( + _cache_config((0, "attention"), (1, "attention"), (2, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(10, 4, 32), + num_kv_heads_per_layer=(8, 8, 8), + head_dim_per_layer=(128, 128, 128), + ) + + assert result is codec + configs = native.create_nvfp4_cold_page_codec.call_args.args[0] + assert [config.layer_id for config in configs] == [0, 1, 2] + assert [config.runtime_type for config in configs] == ["native-bf16"] * 3 + assert [ + (config.nvfp4_scale_orig_quant, config.nvfp4_scale_quant_orig) + for config in configs + ] == [ + ((2.0, 4.0), (0.5, 0.25)), + ((8.0, 16.0), (0.125, 0.0625)), + ((1.0, 1.0), (1.0, 1.0)), + ] + + +def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): + native, _ = _native() + cache_config = _cache_config((0, "attention")) + cache_config.tokens_per_block = 5 + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.HALF, + pp_layers=(10,), + num_kv_heads_per_layer=(4,), + head_dim_per_layer=(128,), + ) + + config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert config.runtime_type == "native-fp16" + assert config.num_kv_heads == 4 + assert config.tokens_per_page == 5 + assert config.head_dim == 128 + assert config.nvfp4_scale_orig_quant == config.nvfp4_scale_quant_orig == (1.0, 1.0) + + +def test_provider_creates_one_native_codec_per_kv_cache_manager(): + native, _ = _native() + codecs = (object(), object()) + native.create_nvfp4_cold_page_codec.side_effect = codecs + provider = _manager() + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + results = tuple( + provider.create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(layer_id,), + num_kv_heads_per_layer=(8,), + head_dim_per_layer=(128,), + ) + for layer_id in (0, 32) + ) + + assert results == codecs + assert native.create_nvfp4_cold_page_codec.call_count == 2 + + +def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): + _write_scales(tmp_path, {7: (0.5, 0.25)}, filename="model.safetensors") + _write_scales( + tmp_path, + {7: (0.125, 0.0625)}, + filename="consolidated.00.safetensors", + ) + assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + (2.0, 4.0), + (0.5, 0.25), + ) + + consolidated_only = tmp_path / "consolidated-only" + consolidated_only.mkdir() + _write_scales( + consolidated_only, + {9: (0.125, 0.0625)}, + filename="consolidated.00.safetensors", + ) + assert _manager(consolidated_only)._model_nvfp4_scales[9] == ( + (8.0, 16.0), + (0.125, 0.0625), + ) + + +def test_scale_loader_reduces_duplicate_shards_like_native_qkv_loader(tmp_path): + _write_scales(tmp_path, {7: (0.25, 0.125)}, filename="model-00001.safetensors") + _write_scales( + tmp_path, + {7: (0.5, 0.25)}, + filename="model-00002.safetensors", + prefix="model.language_model", + ) + assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + (2.0, 4.0), + (0.5, 0.25), + ) + + +def test_trtllm_load_kv_scales_zero_uses_identity(tmp_path, monkeypatch): + _write_scales(tmp_path, {7: (0.5, 0.25)}) + monkeypatch.setenv("TRTLLM_LOAD_KV_SCALES", "0") + assert _manager(tmp_path)._model_nvfp4_scales == {} + + +def test_non_nvfp4_checkpoint_scales_are_not_reused(tmp_path): + _write_scales(tmp_path, {7: (0.5, 0.25)}) + _write_quant_metadata(tmp_path, "FP8") + assert _manager(tmp_path)._model_nvfp4_scales == {} + + +def test_unquantized_checkpoint_uses_identity_scales(tmp_path): + save_file({"model.weight": torch.ones(1)}, str(tmp_path / "model.safetensors")) + assert _manager(tmp_path)._model_nvfp4_scales == {} + + +def test_explicit_scale_checkpoint_requires_safetensors(tmp_path): + with pytest.raises(FileNotFoundError, match="No safetensors files"): + _manager(tmp_path) + + +@pytest.mark.parametrize("present_kind", ["k", "v"]) +def test_scale_checkpoint_requires_kv_pair(tmp_path, present_kind): + _write_scales(tmp_path, {7: (0.5, 0.5)}) + base = "model.layers.7.self_attn" + name = f"{base}.{present_kind}_proj.{present_kind}_scale" + save_file({name: torch.tensor(0.5)}, str(tmp_path / "model.safetensors")) + with pytest.raises(ValueError, match="both K and V"): + _manager(tmp_path) + + +def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): + native, codec = _native() + _write_scales(tmp_path, {4: (0.5, 0.25)}) + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager(tmp_path).create_cold_page_codec( + _cache_config((0, "ssm"), (1, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(10, 4), + num_kv_heads_per_layer=(0, 8), + head_dim_per_layer=(128, 128), + ) + config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert config.layer_id == 1 + assert config.nvfp4_scale_orig_quant == (2.0, 4.0) + + result = _manager().create_cold_page_codec( + _cache_config((0, "ssm")), + runtime_dtype=DataType.INT8, + pp_layers=(10,), + num_kv_heads_per_layer=(0,), + head_dim_per_layer=(128,), + ) + + assert result is codec + native.create_nvfp4_cold_page_codec.assert_called_with([]) + + +def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): + native, codec = _native() + _write_scales(tmp_path, {10: (0.5, 0.25)}) + cache_config = SimpleNamespace( + tokens_per_block=64, + layers=( + AttentionLayerConfig( + layer_id=0, + buffers=[ + BufferConfig(role="key", size=64 * 576 * 2), + BufferConfig(role="index_key", size=64 * 132), + ], + ), + ), + ) + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + result = _manager(tmp_path).create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=(10,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(576,), + ) + + assert result is codec + config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert config.layer_id == 0 + assert config.runtime_type == "native-bf16" + assert config.num_kv_heads == 1 + assert config.tokens_per_page == 64 + assert config.head_dim == 576 + assert config.nvfp4_scale_orig_quant == config.nvfp4_scale_quant_orig == (1.0, 1.0) + + +def test_fp8_runtime_uses_native_unit_source_scale_default(tmp_path): + native, _ = _native() + _write_scales(tmp_path, {10: (0.5, 0.25)}) + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager(tmp_path).create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.FP8, + pp_layers=(10,), + num_kv_heads_per_layer=(8,), + head_dim_per_layer=(128,), + ) + + config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert config.fp8_scale_orig_quant == config.fp8_scale_quant_orig == (1.0, 1.0) + assert config.nvfp4_scale_quant_orig == (0.5, 0.25) + + +def test_runtime_admission_is_checked_in_utils_before_manager_creation(monkeypatch): + monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "python") + with pytest.raises(ValueError, match=r"require.*C\+\+ KVCacheManagerV2"): + _validate_compression() + + monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: False) + with pytest.raises(RuntimeError, match="requires an SM100-family device"): + _validate_compression() + + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) + _validate_compression() + + +def test_speculative_admission_accepts_verified_one_model_modes(monkeypatch): + monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) + + _validate_compression(SpeculativeDecodingMode.EAGLE3_ONE_MODEL) + _validate_compression(SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL) + for mode in ( + SpeculativeDecodingMode.MTP, + SpeculativeDecodingMode.MTP_EAGLE, + SpeculativeDecodingMode.EAGLE3, + SpeculativeDecodingMode.DFLASH, + ): + with pytest.raises(ValueError, match="one-model MTP-EAGLE or EAGLE3"): + _validate_compression(mode) + + +def test_qwen35_mtp3_resolves_to_supported_one_model_mode(monkeypatch): + monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) + + spec_config = MTPDecodingConfig(max_draft_len=3) + update_spec_config_from_model_config( + spec_config, + SimpleNamespace(mtp_num_hidden_layers=1), + ) + + assert spec_config.spec_dec_mode is SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL + assert spec_config.max_draft_len == 3 + util_mod.validate_kv_cache_compression_compatibility( + ColdPageQuantizationCompressionConfig(), + SimpleNamespace(enable_block_reuse=False), + spec_config, + ) + + +def test_cold_manager_is_disabled_for_estimation_and_active_nvfp4(): + def build(*, estimating=False, active_kv_quant=None): + creator = object.__new__(util_mod.KvCacheCreator) + creator._skip_est = False + creator._max_seq_len = 1024 + creator._kv_cache_config = SimpleNamespace() + creator._llm_args = SimpleNamespace( + kv_cache_compression_config=ColdPageQuantizationCompressionConfig() + ) + model_config = SimpleNamespace( + quant_config=active_kv_quant, pretrained_config=object() + ) + creator._model_engine = SimpleNamespace( + model=SimpleNamespace(model_config=model_config) + ) + creator._draft_model_engine = None + creator._kv_connector_manager = None + creator._fp8_ctx_mla_kv_len_cap = None + creator._is_encoder_decoder = MagicMock(return_value=False) + creator._should_create_separate_draft_kv_cache = MagicMock(return_value=False) + creator._create_kv_cache_manager = MagicMock(return_value=SimpleNamespace()) + resources = {} + creator.build_managers(resources, estimating_kv_cache=estimating) + return resources + + resources = build() + manager = resources[util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] + assert isinstance(manager, ColdPageQuantizationCompression) + assert manager.provides_cold_page_codec + assert not manager.uses_iteration_lifecycle + for kwargs in ( + {"estimating": True}, + {"active_kv_quant": SimpleNamespace(kv_cache_quant_algo="NVFP4")}, + ): + assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in build( + **kwargs + ) diff --git a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py index 92a8bda7ed4d..5c38c121a5e4 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py +++ b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py @@ -90,9 +90,9 @@ def test_factory_allows_block_reuse_and_propagates_config_fields(self): ): mgr = create_kv_cache_compression_manager( cfg, - kv_cache_manager=fake_v2, pretrained_config=_make_test_pretrained_config(), ) + mgr.bind_kv_cache_managers(fake_v2) assert isinstance(mgr, TriAttentionCompressionManager) assert mgr.budget == 32 assert mgr.beta == 16 @@ -286,7 +286,7 @@ def test_prepare_does_not_evict_and_update_runs_final_hook_once(self): @staticmethod def _make_due_decode_request(seq_len, *, num_extra_kv_tokens=0, kv_reserve_draft_tokens=0): # The growth and protected-tail capacity constants snapshot the - # manager at construction, so the reserve widths are set up front. + # manager at binding, so the reserve widths are set up front. request = _make_request( 7, py_prompt_len=1024, @@ -299,9 +299,9 @@ def _make_due_decode_request(seq_len, *, num_extra_kv_tokens=0, kv_reserve_draft with mock.patch.object(TriAttentionCompressionManager, "_initialize_eviction_state"): mgr = TriAttentionCompressionManager( _make_tri_config(budget=8), - fake_v2, pretrained_config=_make_test_pretrained_config(), ) + mgr.bind_kv_cache_managers(fake_v2) cache = SimpleNamespace( capacity=seq_len, history_length=1024, @@ -441,12 +441,11 @@ def test_confirmed_length_comes_from_capacity_ledger_not_logical_length(self): def test_one_model_draft_co_compression_is_accepted(self, spec_mode): draft_manager = _make_fake_v2(is_draft=True) with mock.patch.object(TriAttentionCompressionManager, "_initialize_eviction_state"): - TriAttentionCompressionManager( + manager = TriAttentionCompressionManager( _make_tri_config(budget=8), - _make_fake_v2(), - draft_kv_cache_manager=draft_manager, pretrained_config=_make_test_pretrained_config(), ) + manager.bind_kv_cache_managers(_make_fake_v2(), draft_manager) from tensorrt_llm._torch.pyexecutor._util import validate_kv_cache_compression_compatibility from tensorrt_llm.llmapi.llm_args import Eagle3DecodingConfig, MTPDecodingConfig diff --git a/tests/unittest/api_stability/references/llm.yaml b/tests/unittest/api_stability/references/llm.yaml index d66b9e72dafe..912875b4bfb8 100644 --- a/tests/unittest/api_stability/references/llm.yaml +++ b/tests/unittest/api_stability/references/llm.yaml @@ -280,7 +280,7 @@ methods: default: null status: prototype kv_cache_compression_config: - annotation: Union[tensorrt_llm.llmapi.llm_args.TriAttentionKvCacheCompressionConfig, NoneType] + annotation: Union[tensorrt_llm.llmapi.llm_args.ColdPageQuantizationCompressionConfig, tensorrt_llm.llmapi.llm_args.TriAttentionKvCacheCompressionConfig, NoneType] default: null status: prototype otlp_traces_endpoint: diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 529851947843..c47ceb93743f 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -3370,8 +3370,46 @@ def test_no_custom_init_methods(self): @pytest.mark.cpu_only def test_kv_cache_compression_config_dispatches_by_algorithm(): - from tensorrt_llm.llmapi.llm_args import \ - TriAttentionKvCacheCompressionConfig + from tensorrt_llm.llmapi.llm_args import ( + ColdPageQuantizationCompressionConfig, + TriAttentionKvCacheCompressionConfig) + + cold_config = TorchLlmArgs( + model="/tmp/dummy_model", + kv_cache_compression_config={ + "algorithm": "quantization_for_cold_page", + "quant": "nvfp4", + }, + ).kv_cache_compression_config + + assert isinstance(cold_config, ColdPageQuantizationCompressionConfig) + assert cold_config.model_dump() == { + "algorithm": "quantization_for_cold_page", + "quant": "nvfp4", + "scale_checkpoint_path": None, + } + assert not cold_config.changes_physical_kv_length + assert cold_config.supports_block_reuse() + assert cold_config.supports_speculative_decoding() + + cold_config_with_scales = TorchLlmArgs( + model="/tmp/dummy_model", + kv_cache_compression_config={ + "algorithm": "quantization_for_cold_page", + "quant": "nvfp4", + "scale_checkpoint_path": "/tmp/nvfp4-kv-scales", + }, + ).kv_cache_compression_config + assert cold_config_with_scales.scale_checkpoint_path == "/tmp/nvfp4-kv-scales" + + with pytest.raises(ValidationError): + TorchLlmArgs( + model="/tmp/dummy_model", + kv_cache_compression_config={ + "algorithm": "quantization_for_cold_page", + "quant": "fp8", + }, + ) config_dict = yaml.safe_load(""" kv_cache_compression_config: diff --git a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py index 059fad4224ef..ab9f77da6aed 100644 --- a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py +++ b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py @@ -286,6 +286,59 @@ def test_build_capture_manifest_matches_committed_golden(): _assert_committed_manifest_current(golden_manifest()) +def test_kv_cache_compression_discriminator_captures_both_algorithms(): + """Flattened telemetry preserves both discriminator Literal values.""" + from tensorrt_llm.llmapi.llm_args import ( + ColdPageQuantizationCompressionConfig, + TorchLlmArgs, + TriAttentionKvCacheCompressionConfig, + ) + from tensorrt_llm.usage.llmapi_config import ( + build_capture_manifest, + collect_llm_api_config_payloads, + ) + + entry = next( + item + for item in build_capture_manifest(TorchLlmArgs) + if item.path == "kv_cache_compression_config.algorithm" + ) + assert repr(entry.annotation) == ( + "typing.Literal['quantization_for_cold_page', 'triattention']" + ) + assert entry.converter == "" + assert set(entry.allowed_values) == { + "quantization_for_cold_page", + "triattention", + } + for config_cls, annotation in ( + ( + ColdPageQuantizationCompressionConfig, + "typing.Literal['quantization_for_cold_page']", + ), + (TriAttentionKvCacheCompressionConfig, "typing.Literal['triattention']"), + ): + field = config_cls.model_fields["algorithm"] + assert repr(field.annotation) == annotation + + configs = ( + ColdPageQuantizationCompressionConfig(), + TriAttentionKvCacheCompressionConfig( + calibration_path="/triattention.pt", + ), + ) + for config in configs: + args = TorchLlmArgs( + model="/model", + kv_cache_compression_config=config, + ) + config_json, metadata_json = collect_llm_api_config_payloads(args) + captured = json.loads(config_json) + metadata = json.loads(metadata_json) + assert captured["kv_cache_compression_config.algorithm"] == config.algorithm + assert metadata["capture_succeeded"] is True + + def test_load_generator_does_not_leak_sys_modules(): """_load_generator must not leak its temporary module into sys.modules. From 6b502cdbf83b84df96787fe80fe1cc7f67daa6e4 Mon Sep 17 00:00:00 2001 From: tianruih Date: Sat, 22 Aug 2026 07:14:26 -0700 Subject: [PATCH 02/29] [None][chore] Clarify KVCC lifecycle and test tiers --- tensorrt_llm/_torch/pyexecutor/resource_manager.py | 8 +++++++- .../_torch/executor/test_kv_cache_compression_manager.py | 2 ++ .../unittest/_torch/executor/test_kv_cache_manager_v2.py | 1 + .../test_quantization_for_cold_page.py | 2 ++ 4 files changed, 12 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 737300f52ed9..c8d189f82468 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -2715,7 +2715,13 @@ def _free_blocks(self, block_list: list): class KVCacheCompressionManager(BaseResourceManager): - """Base for iteration-driven and storage-boundary KV compression.""" + """Framework base for KV-cache compression methods in PyExecutor. + + Iteration-driven methods receive ResourceManager callbacks, while + storage-boundary methods provide a cold-page codec during cache construction. + Subclasses coordinate through KVCacheManagerV2 without owning its pools, + mappings, or migration lifecycle. + """ uses_iteration_lifecycle = True provides_cold_page_codec = False diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index d7fa9ffb4479..cc5d4619dc7d 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -19,6 +19,8 @@ ) from tensorrt_llm.llmapi.llm_args import KvCacheCompressionConfig +pytestmark = pytest.mark.cpu_only + # ---------------------------------------------------------------------- # # Mock infra: in-memory managers / requests (avoid touching V2 / model). # # ---------------------------------------------------------------------- # diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index f86a30cad85c..52231f4d7bb4 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -370,6 +370,7 @@ def test_host_init_fallback_drops_only_host_tier(tmp_path) -> None: ] +@pytest.mark.cpu_only def test_host_init_fallback_recreates_cold_codec_and_keeps_disk(tmp_path) -> None: impl = Mock() codecs = [object(), object()] diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index f07d0dee8646..3cbaa78086c0 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -28,6 +28,8 @@ SsmLayerConfig, ) +pytestmark = pytest.mark.cpu_only + def _manager(scale_checkpoint_path=None): config = ColdPageQuantizationCompressionConfig( From 8bb2c3feb3f18592b038ff1502ac89889faeb65e Mon Sep 17 00:00:00 2001 From: tianruih Date: Sun, 23 Aug 2026 23:55:11 -0700 Subject: [PATCH 03/29] [None][fix] Polish NVFP4 cold-page compression Signed-off-by: tianruih --- cpp/tensorrt_llm/CMakeLists.txt | 2 +- .../kernels/nvfp4BoundaryKernels.cu | 318 ++++++++-------- .../kernels/nvfp4BoundaryKernels.h | 5 +- .../kv_cache_compression/CMakeLists.txt | 2 +- .../nvfp4ColdPageCodec.cpp | 344 +++++++++--------- .../kv_cache_compression/nvfp4ColdPageCodec.h | 8 +- cpp/tests/unit_tests/CMakeLists.txt | 2 +- .../kernels/nvfp4BoundaryKernelsTest.cpp | 26 +- .../kv_cache_compression/CMakeLists.txt | 3 +- .../nvfp4ColdPageCodecTest.cpp | 36 +- .../quantization_for_cold_page/__init__.py | 4 - .../quantization_for_cold_page.py | 6 +- tensorrt_llm/_torch/pyexecutor/_util.py | 64 ++-- .../_torch/pyexecutor/kv_cache_manager_v2.py | 48 ++- tensorrt_llm/llmapi/llm_args.py | 18 +- .../_core/_kv_cache_manager.py | 3 + .../usage/llm_args_golden_manifest.json | 10 +- tensorrt_llm/usage/llmapi_config.py | 22 +- .../executor/test_kv_cache_budget_split.py | 1 + .../test_kv_cache_compression_manager.py | 37 +- .../test_quantization_for_cold_page.py | 56 +-- .../test_triattention_draft_cocompaction.py | 8 +- .../test_triattention_pipeline.py | 24 +- .../test_llmapi_config_telemetry_docs.py | 25 +- 24 files changed, 564 insertions(+), 508 deletions(-) diff --git a/cpp/tensorrt_llm/CMakeLists.txt b/cpp/tensorrt_llm/CMakeLists.txt index 096783b735c5..8be06c0bb516 100644 --- a/cpp/tensorrt_llm/CMakeLists.txt +++ b/cpp/tensorrt_llm/CMakeLists.txt @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2022-2024 NVIDIA CORPORATION & +# SPDX-FileCopyrightText: Copyright (c) 2022-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 diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu index 71147fd15db3..671f5243686a 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu @@ -48,34 +48,43 @@ constexpr std::uint32_t kThreadsPerBlock = 128; constexpr std::uint32_t kAsyncStages = 4; // Mapped-Host reads use eight cp.async stages; GPU-resident input uses four. constexpr std::uint32_t kHostLoadAsyncStages = 8; -constexpr std::uint32_t kHostMemorySplits = 1; +constexpr std::uint32_t kMappedHostGridSplits = 1; constexpr std::uint32_t kMaxTasksPerLaunch = 256; // Bound by-value buffer metadata to CUDA's kernel-parameter limit. constexpr std::uint32_t kMaxBuffersPerLaunch = kNvfp4BoundaryMaxBuffersPerLaunch; -constexpr std::uint32_t kElementsPerLane = 8; -constexpr std::uint32_t kElementsPerBlockScale = 16; +constexpr std::uint32_t kElementsPerHalfGroup = 8; +constexpr std::uint32_t kElementsPerScaleGroup = 16; +constexpr std::uint32_t kHalfGroupsPerScaleGroup = kElementsPerScaleGroup / kElementsPerHalfGroup; // Bound per-tile shared scale staging to 1 KiB. -constexpr std::uint32_t kTargetScaleTransferBytes = 1024; -constexpr std::uint32_t kHalfGroupsPerTransfer = 2U * kTargetScaleTransferBytes; -constexpr std::size_t kModernKernelParameterLimit = 32764; - -static_assert(kThreadsPerBlock % 2 == 0, "An NVFP4 scale group is shared by two lanes"); -static_assert(kTargetScaleTransferBytes > 0 && kTargetScaleTransferBytes % sizeof(uint4) == 0, - "Compact scale transfer must remain a positive 16-byte multiple"); -static_assert(kHalfGroupsPerTransfer % 2U == 0, "A transfer tile must not split an NVFP4 scale group"); -static_assert(std::is_trivially_copyable_v); -static_assert(std::is_trivially_copyable_v); -static_assert(std::is_trivially_copyable_v); +constexpr std::uint32_t kMaxScaleBytesPerTile = 1024; +constexpr std::uint32_t kMaxHalfGroupsPerTile = kHalfGroupsPerScaleGroup * kMaxScaleBytesPerTile; +constexpr std::size_t kKernelParameterLimitBytes = 32764; + +// Keep every CTA iteration and shared-memory tile on complete 16-value scale groups. +static_assert(kElementsPerScaleGroup % kElementsPerHalfGroup == 0, + "An NVFP4 scale group must contain a whole number of half-groups"); +static_assert(kHalfGroupsPerScaleGroup == 2U, "NVFP4 stores one scale for two eight-value half-groups"); +static_assert(kThreadsPerBlock % kHalfGroupsPerScaleGroup == 0, "A CTA iteration must not split an NVFP4 scale group"); +static_assert(kMaxScaleBytesPerTile > 0, "A transfer tile must make forward progress"); static_assert( - kMaxBuffersPerLaunch <= std::numeric_limits::max(), "Boundary buffer count must fit CUDA grid.y"); + kMaxScaleBytesPerTile % sizeof(uint4) == 0, "A full scale tile must preserve the 16-byte transfer fast path"); + +// Buffer plans are passed to CUDA kernels as raw by-value arguments. +static_assert(std::is_trivially_copyable_v, "Buffer plans must remain raw-copyable"); + +// Keep both kernel argument packs within CUDA's modern 32,764-byte limit. static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) - + 3U * sizeof(std::uint32_t) - <= kModernKernelParameterLimit); + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + + 3U * sizeof(std::uint32_t) + <= kKernelParameterLimitBytes, + "Offload kernel arguments exceed CUDA's parameter limit"); static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) - + 3U * sizeof(std::uint32_t) - <= kModernKernelParameterLimit); + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + + 3U * sizeof(std::uint32_t) + <= kKernelParameterLimitBytes, + "Onboard kernel arguments exceed CUDA's parameter limit"); + +// Device data path. // Issue one predicated 16-byte cp.async load. template @@ -124,39 +133,43 @@ __device__ OnboardBufferTask resolveTask(Nvfp4BoundaryOnboardPageTask const& pag reinterpret_cast(buffer.rawBase + gpuPage * buffer.rawSlotBytes)}; } -__host__ __device__ constexpr std::uint32_t packedBytesPerBuffer(std::uint32_t halfGroups) +// One eight-element half-group produces one uint32_t of packed E2M1 data. +__host__ __device__ constexpr std::uint32_t packedBytesForHalfGroups(std::uint32_t halfGroupCount) { - return halfGroups * sizeof(std::uint32_t); + return halfGroupCount * sizeof(std::uint32_t); } -__host__ __device__ constexpr std::uint32_t scaleBytesPerBuffer(std::uint32_t halfGroups) +// Store one E4M3 scale byte per pair of half-groups. +__host__ __device__ constexpr std::uint32_t scaleBytesForHalfGroups(std::uint32_t halfGroupCount) { - return halfGroups / 2U; + return halfGroupCount / kHalfGroupsPerScaleGroup; } -// Shared staging padding is not stored in the cold record. -__host__ __device__ constexpr std::uint32_t packedStagingBytesPerBuffer(std::uint32_t halfGroups) +// Align packed shared staging to uint4; this padding is not stored in the cold record. +__host__ __device__ constexpr std::uint32_t packedStageBytesForHalfGroups(std::uint32_t halfGroupCount) { - return (packedBytesPerBuffer(halfGroups) + sizeof(uint4) - 1U) / sizeof(uint4) * sizeof(uint4); + return (packedBytesForHalfGroups(halfGroupCount) + sizeof(uint4) - 1U) / sizeof(uint4) * sizeof(uint4); } -__host__ __device__ constexpr std::uint32_t compactStagingBytesPerBuffer(std::uint32_t halfGroups) +// Return shared-memory bytes for aligned packed data followed by scale bytes. +__host__ __device__ constexpr std::uint32_t compactStageBytesForHalfGroups(std::uint32_t halfGroupCount) { - return packedStagingBytesPerBuffer(halfGroups) + scaleBytesPerBuffer(halfGroups); + return packedStageBytesForHalfGroups(halfGroupCount) + scaleBytesForHalfGroups(halfGroupCount); } -__host__ __device__ constexpr std::uint32_t totalHalfGroupsPerBuffer(Nvfp4BoundaryKernelParams const& params) +// Flatten one buffer's [head, token, dim] geometry into eight-element half-groups. +__host__ __device__ constexpr std::uint32_t halfGroupCount(Nvfp4BoundaryKernelParams const& params) { return static_cast(params.numKvHeads) * static_cast(params.tokensPerPage) - * (static_cast(params.headDim) / kElementsPerLane); + * (static_cast(params.headDim) / kElementsPerHalfGroup); } -// Tile at most 2,048 half-groups; headDim % 16 prevents splitting a scale group. -__host__ __device__ constexpr std::uint32_t compressedTransferHalfGroups(Nvfp4BoundaryKernelParams const& params) +// Cap a buffer tile at shared-memory capacity without splitting a scale group. +__host__ __device__ constexpr std::uint32_t tileHalfGroupCount(Nvfp4BoundaryKernelParams const& params) { // Avoid std::min: its reference return ODR-uses this host/device constexpr. - auto const halfGroups = totalHalfGroupsPerBuffer(params); - return halfGroups < kHalfGroupsPerTransfer ? halfGroups : kHalfGroupsPerTransfer; + auto const halfGroups = halfGroupCount(params); + return halfGroups < kMaxHalfGroupsPerTile ? halfGroups : kMaxHalfGroupsPerTile; } // Flush vectorized packed values and scales, then copy remaining tails bytewise. @@ -371,7 +384,7 @@ __device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firs { float2 values[4]; unpackE2m1ToFloat(packedWords[laneInScale], values); - store16BitValues(output, (firstHalfGroup + laneInScale) * kElementsPerLane, values, dequantScale); + store16BitValues(output, (firstHalfGroup + laneInScale) * kElementsPerHalfGroup, values, dequantScale); } } else @@ -410,7 +423,7 @@ __global__ void offloadFrom16BitTiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) { copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); clearColdPadding(task, buffer); @@ -418,9 +431,9 @@ __global__ void offloadFrom16BitTiledKernel( } auto const params = buffer.params; - std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); - std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); - std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + std::uint32_t const halfGroupsPerBuffer = halfGroupCount(params); + std::uint32_t const tileHalfGroups = tileHalfGroupCount(params); + std::uint32_t const packedStageCapacityBytes = packedStageBytesForHalfGroups(tileHalfGroups); // Use a four-stage 16-byte cp.async ring and tile-bounded compact staging. __shared__ __align__(16) PackedVec rawStages[kAsyncStages][kThreadsPerBlock]; @@ -451,7 +464,7 @@ __global__ void offloadFrom16BitTiledKernel( std::uint32_t const localScaleOffset = localHalfGroup >> 1U; std::uint8_t* scale = laneInScale == 0 ? scaleStages + localScaleOffset : nullptr; PackedVec input = rawStages[stage][threadIdx.x]; - packedStages[localHalfGroup] = cvt_warp_fp16_to_fp4( + packedStages[localHalfGroup] = cvt_warp_fp16_to_fp4( input, params.nvfp4ScaleOrigQuant, scale); } } @@ -467,8 +480,9 @@ __global__ void offloadFrom16BitTiledKernel( // Publish quant results and finish the flush before reusing staging. cp_async_wait_group<0>(); __syncthreads(); - flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), - packedBytesPerBuffer(halfGroups), firstHalfGroup / 2U, halfGroups / 2U); + flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesForHalfGroups(firstHalfGroup), + packedBytesForHalfGroups(halfGroups), scaleBytesForHalfGroups(firstHalfGroup), + scaleBytesForHalfGroups(halfGroups)); __syncthreads(); } clearColdPadding(task, buffer); @@ -490,7 +504,7 @@ __global__ void offloadFromFp8TiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) { copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); clearColdPadding(task, buffer); @@ -498,9 +512,9 @@ __global__ void offloadFromFp8TiledKernel( } auto const params = buffer.params; - std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); - std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); - std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + std::uint32_t const halfGroupsPerBuffer = halfGroupCount(params); + std::uint32_t const tileHalfGroups = tileHalfGroupCount(params); + std::uint32_t const packedStageCapacityBytes = packedStageBytesForHalfGroups(tileHalfGroups); // Each cp.async moves one 16-byte grain in the production PackedVec layout. __shared__ __align__(16) PackedVec<__nv_fp8_e4m3> rawStages[kAsyncStages][kThreadsPerBlock]; @@ -543,8 +557,9 @@ __global__ void offloadFromFp8TiledKernel( cp_async_wait_group<0>(); __syncthreads(); - flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), - packedBytesPerBuffer(halfGroups), firstHalfGroup / 2U, halfGroups / 2U); + flushCompactRangeToHost(compactStages, task, packedStageCapacityBytes, packedBytesForHalfGroups(firstHalfGroup), + packedBytesForHalfGroups(halfGroups), scaleBytesForHalfGroups(firstHalfGroup), + scaleBytesForHalfGroups(halfGroups)); __syncthreads(); } clearColdPadding(task, buffer); @@ -633,16 +648,16 @@ __global__ void onboardTiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLossless) + if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) { copyLosslessBytes(task.coldData, task.raw, buffer.rawBytes); return; } auto const params = buffer.params; - std::uint32_t const halfGroupsPerBuffer = totalHalfGroupsPerBuffer(params); - std::uint32_t const tileHalfGroups = compressedTransferHalfGroups(params); - std::uint32_t const packedStageCapacityBytes = packedStagingBytesPerBuffer(tileHalfGroups); + std::uint32_t const halfGroupsPerBuffer = halfGroupCount(params); + std::uint32_t const tileHalfGroups = tileHalfGroupCount(params); + std::uint32_t const packedStageCapacityBytes = packedStageBytesForHalfGroups(tileHalfGroups); extern __shared__ __align__(16) std::uint8_t compactStages[]; auto* rawOutput = reinterpret_cast(task.raw); @@ -650,12 +665,12 @@ __global__ void onboardTiledKernel( firstHalfGroup += gridDim.x * tileHalfGroups) { std::uint32_t const halfGroups = std::min(tileHalfGroups, halfGroupsPerBuffer - firstHalfGroup); - std::uint32_t const packedBytes = packedBytesPerBuffer(halfGroups); - std::uint32_t const scaleBytes = halfGroups / 2U; + std::uint32_t const packedBytes = packedBytesForHalfGroups(halfGroups); + std::uint32_t const scaleBytes = scaleBytesForHalfGroups(halfGroups); // Stage packed data and scales before dequantization. - loadCompactRangeFromHost(compactStages, task, packedStageCapacityBytes, packedBytesPerBuffer(firstHalfGroup), - packedBytes, firstHalfGroup / 2U, scaleBytes); + loadCompactRangeFromHost(compactStages, task, packedStageCapacityBytes, + packedBytesForHalfGroups(firstHalfGroup), packedBytes, scaleBytesForHalfGroups(firstHalfGroup), scaleBytes); auto const* packedStages = reinterpret_cast(compactStages); auto const* scaleStages = compactStages + packedStageCapacityBytes; @@ -688,17 +703,20 @@ __global__ void onboardTiledKernel( #endif } -void validateParams(Nvfp4BoundaryKernelParams const& params, bool useFp8) +// Construction-time launch-plan validation. + +// Validate the geometry and scales used only by an NVFP4 transform. +void validateNvfp4Params(Nvfp4BoundaryKernelParams const& params, bool isFp8Runtime) { TLLM_CHECK_WITH_INFO(params.numKvHeads > 0, "numKvHeads must be positive"); TLLM_CHECK_WITH_INFO(params.tokensPerPage > 0, "tokensPerPage must be positive"); - TLLM_CHECK_WITH_INFO(params.headDim > 0 && params.headDim % kElementsPerBlockScale == 0, + TLLM_CHECK_WITH_INFO(params.headDim > 0 && params.headDim % kElementsPerScaleGroup == 0, "headDim must be positive and divisible by 16, got %d", params.headDim); std::uint64_t const rows = static_cast(params.numKvHeads) * static_cast(params.tokensPerPage); - constexpr std::uint64_t maxHalfGroups = std::numeric_limits::max() / kElementsPerLane; - std::uint64_t const halfGroupsPerRow = static_cast(params.headDim / kElementsPerLane); + constexpr std::uint64_t maxHalfGroups = std::numeric_limits::max() / kElementsPerHalfGroup; + std::uint64_t const halfGroupsPerRow = static_cast(params.headDim / kElementsPerHalfGroup); TLLM_CHECK_WITH_INFO(rows <= maxHalfGroups / halfGroupsPerRow, "Page geometry exceeds the 32-bit compact-offset range: " "heads=%d, tokens=%d, headDim=%d", @@ -708,7 +726,7 @@ void validateParams(Nvfp4BoundaryKernelParams const& params, bool useFp8) "NVFP4 original-to-quantized scale must be finite and positive"); TLLM_CHECK_WITH_INFO(std::isfinite(params.nvfp4ScaleQuantOrig) && params.nvfp4ScaleQuantOrig > 0.0F, "NVFP4 quantized-to-original scale must be finite and positive"); - if (useFp8) + if (isFp8Runtime) { TLLM_CHECK_WITH_INFO(std::isfinite(params.fp8ScaleOrigQuant) && params.fp8ScaleOrigQuant > 0.0F, "FP8 original-to-quantized scale must be finite and positive"); @@ -735,8 +753,9 @@ void addColdInterval(std::vector& intervals, std::size_t offset, s intervals.push_back({offset, offset + bytes}); } -void validateBufferPlan( - Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldPageBytes, bool useFp8, std::vector& intervals) +// Validate one raw/cold buffer mapping and append its occupied cold intervals. +void validateBufferPlan(Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldPageBytes, bool isFp8Runtime, + std::vector& intervals) { TLLM_CHECK_WITH_INFO(buffer.rawBase != 0U, "rawBase must not be null"); TLLM_CHECK_WITH_INFO(buffer.rawBytes > 0 && buffer.rawBytes <= buffer.rawSlotBytes, @@ -746,7 +765,7 @@ void validateBufferPlan( { case Nvfp4BoundaryTransform::kNvfp4: { - validateParams(buffer.params, useFp8); + validateNvfp4Params(buffer.params, isFp8Runtime); TLLM_CHECK_WITH_INFO( buffer.rawBase % alignof(uint4) == 0, "rawBase must be aligned to %zu bytes", alignof(uint4)); TLLM_CHECK_WITH_INFO(buffer.rawSlotBytes % alignof(uint4) == 0, @@ -754,7 +773,7 @@ void validateBufferPlan( std::uint64_t const elements = static_cast(buffer.params.numKvHeads) * static_cast(buffer.params.tokensPerPage) * static_cast(buffer.params.headDim); - std::uint64_t const expectedRawBytes = elements * (useFp8 ? 1U : 2U); + std::uint64_t const expectedRawBytes = elements * (isFp8Runtime ? 1U : 2U); TLLM_CHECK_WITH_INFO(buffer.rawBytes == static_cast(expectedRawBytes), "Raw buffer size does not match NVFP4 geometry and runtime type"); addColdInterval(intervals, buffer.coldDataOffset, static_cast(elements / 2U), coldPageBytes, @@ -763,7 +782,7 @@ void validateBufferPlan( "NVFP4 scale interval"); break; } - case Nvfp4BoundaryTransform::kLossless: + case Nvfp4BoundaryTransform::kLosslessCopy: addColdInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes, "Lossless-data interval"); break; default: TLLM_THROW("Unsupported NVFP4 boundary buffer transform"); @@ -772,111 +791,62 @@ void validateBufferPlan( intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes, "Cold-record padding interval"); } -// Launch the 256-descriptor raw-argument ABI with optional PDL. -// Only a partial chunk needs a zero-padded stack argument. -template -void launchBoundaryBatch(Kernel kernel, Task const* tasks, std::uint32_t count, dim3 grid, dim3 block, - std::uint32_t dynamicSmemBytes, std::array const& buffers, - ColdPointer coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers, cudaStream_t stream) -{ - static_assert(std::is_trivially_copyable_v>); - static_assert(sizeof(std::array) == sizeof(Task) * kMaxTasksPerLaunch); - TLLM_CHECK(count > 0 && count <= kMaxTasksPerLaunch); - - // Zero initialization keeps partial argument padding deterministic. - std::array tail{}; - void const* taskArgument = tasks; - if (count < kMaxTasksPerLaunch) - { - std::copy_n(tasks, count, tail.begin()); - taskArgument = tail.data(); - } - - cudaLaunchAttribute attribute{}; - attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; - attribute.val.programmaticStreamSerializationAllowed = common::getEnvEnablePDL() ? 1 : 0; - - cudaLaunchConfig_t config{}; - config.gridDim = grid; - config.blockDim = block; - config.dynamicSmemBytes = dynamicSmemBytes; - config.stream = stream; - config.attrs = &attribute; - config.numAttrs = 1; - - void* arguments[] = {const_cast(taskArgument), const_cast(buffers.data()), - &coldBase, &coldPageBytes, &numBuffers}; - TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); -} +// Host submission path. -// Submit Page descriptors in fixed-capacity chunks. -template -void launchTaskBatches(std::vector const& tasks, LaunchBatch const& launchBatch) +// Submit Page tasks through the fixed 256-descriptor kernel ABI. +template +void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4BoundaryPreparedPlan const& plan, + ColdPointer coldBase, cudaStream_t stream) { + static_assert(std::is_trivially_copyable_v>, + "Page tasks must remain raw-copyable kernel arguments"); + static_assert(sizeof(std::array) == sizeof(Task) * kMaxTasksPerLaunch, + "Page task arrays must not add ABI padding"); + dim3 const block(kThreadsPerBlock); std::size_t offset = 0; while (offset < tasks.size()) { - std::uint32_t const count + std::uint32_t const numChunkTasks = static_cast(std::min(tasks.size() - offset, kMaxTasksPerLaunch)); - launchBatch(tasks.data() + offset, count); - offset += count; - } -} - -// One tiled SM100 kernel family covers every runtime dtype, transform, and Page geometry. -// grid.y selects one independent buffer; geometry only changes the tile loop length. + Task const* chunkTasks = tasks.data() + offset; -template -void launchOffloadFrom16Bit(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, std::uint8_t* coldBase, cudaStream_t stream) -{ - dim3 const block(kThreadsPerBlock); - launchTaskBatches(pages, - [&](Nvfp4BoundaryOffloadPageTask const* taskData, std::uint32_t count) + // CUDA copies the full by-value array, so pad only the final partial chunk. + std::array paddedTasks{}; + if (numChunkTasks < kMaxTasksPerLaunch) { - dim3 const grid(kHostMemorySplits, plan.numBuffers, count); - launchBoundaryBatch(offloadFrom16BitTiledKernel, taskData, count, grid, block, - compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, - plan.numBuffers, stream); - }); -} + std::copy_n(chunkTasks, numChunkTasks, paddedTasks.begin()); + chunkTasks = paddedTasks.data(); + } -void launchOffloadFromFp8(std::vector const& pages, Nvfp4BoundaryPreparedPlan const& plan, - std::uint8_t* coldBase, cudaStream_t stream) -{ - dim3 const block(kThreadsPerBlock); - launchTaskBatches(pages, - [&](Nvfp4BoundaryOffloadPageTask const* taskData, std::uint32_t count) - { - dim3 const grid(kHostMemorySplits, plan.numBuffers, count); - launchBoundaryBatch(offloadFromFp8TiledKernel, taskData, count, grid, block, - compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, - plan.numBuffers, stream); - }); -} + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = common::getEnvEnablePDL() ? 1 : 0; -template -void launchOnboard(std::vector const& pages, Nvfp4BoundaryPreparedPlan const& plan, - std::uint8_t const* coldBase, cudaStream_t stream) -{ - dim3 const block(kThreadsPerBlock); - launchTaskBatches(pages, - [&](Nvfp4BoundaryOnboardPageTask const* taskData, std::uint32_t count) - { - dim3 const grid(kHostMemorySplits, plan.numBuffers, count); - launchBoundaryBatch(onboardTiledKernel, taskData, count, grid, block, - compactStagingBytesPerBuffer(plan.maxTileHalfGroups), plan.buffers, coldBase, plan.coldPageBytes, - plan.numBuffers, stream); - }); + cudaLaunchConfig_t config{}; + config.gridDim = dim3(kMappedHostGridSplits, plan.numBuffers, numChunkTasks); + config.blockDim = block; + config.dynamicSmemBytes = compactStageBytesForHalfGroups(plan.maxHalfGroupsPerTile); + config.stream = stream; + config.attrs = &attribute; + config.numAttrs = 1; + + auto coldPageBytes = plan.coldPageBytes; + auto numBuffers = plan.numBuffers; + void* arguments[] = {const_cast(chunkTasks), const_cast(plan.buffers.data()), + &coldBase, &coldPageBytes, &numBuffers}; + TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); + offset += numChunkTasks; + } } -// Drain earlier chunks after synchronous launch failure before Slots are recycled. -template -void launchAndDrainOnFailure(cudaStream_t stream, Launch const& launch) +// Drain earlier chunks after a later synchronous launch failure. +template +void submitBoundaryTasks(Kernel kernel, std::vector const& tasks, Nvfp4BoundaryPreparedPlan const& plan, + ColdPointer coldBase, cudaStream_t stream) { try { - launch(); + launchTaskChunks(kernel, tasks, plan, coldBase, stream); } catch (...) { @@ -884,7 +854,7 @@ void launchAndDrainOnFailure(cudaStream_t stream, Launch const& launch) if (drainStatus != cudaSuccess) { // An asynchronous drain failure leaves Slot ownership unknown; fail-stop. - TLLM_LOG_ERROR("NVFP4 boundary rollback drain failed: %s", cudaGetErrorString(drainStatus)); + TLLM_LOG_ERROR("NVFP4 boundary failure drain failed: %s", cudaGetErrorString(drainStatus)); std::terminate(); } throw; @@ -905,12 +875,12 @@ Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks( + offloadFrom16BitTiledKernel, pages, plan, static_cast(coldBase), stream); break; case Nvfp4BoundaryRuntimeType::kBfloat16: - launchAndDrainOnFailure(stream, - [&] { launchOffloadFrom16Bit<__nv_bfloat16>(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks( + offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); break; case Nvfp4BoundaryRuntimeType::kFp8E4m3: - launchAndDrainOnFailure( - stream, [&] { launchOffloadFromFp8(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks(offloadFromFp8TiledKernel, pages, plan, static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); } @@ -977,16 +946,15 @@ void invokeNvfp4BoundaryOnboardDecompress(std::vector(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks(onboardTiledKernel, pages, plan, static_cast(coldBase), stream); break; case Nvfp4BoundaryRuntimeType::kBfloat16: - launchAndDrainOnFailure(stream, - [&] { launchOnboard<__nv_bfloat16>(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks( + onboardTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); break; case Nvfp4BoundaryRuntimeType::kFp8E4m3: - launchAndDrainOnFailure(stream, - [&] { launchOnboard<__nv_fp8_e4m3>(pages, plan, static_cast(coldBase), stream); }); + submitBoundaryTasks( + onboardTiledKernel<__nv_fp8_e4m3>, pages, plan, static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); } diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h index 708e5835c2f0..5496bff9af92 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h +++ b/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h @@ -70,7 +70,8 @@ struct Nvfp4BoundaryKernelParams enum class Nvfp4BoundaryTransform : std::uint8_t { kNvfp4, - kLossless, + //! Byte-exact copy for an Attention side buffer such as DSA index_key. + kLosslessCopy, }; //! Immutable transform plan for one hot buffer and its fixed-offset cold record. @@ -94,7 +95,7 @@ struct Nvfp4BoundaryPreparedPlan { std::array buffers{}; std::uint32_t numBuffers = 0; - std::uint32_t maxTileHalfGroups = 0; + std::uint32_t maxHalfGroupsPerTile = 0; std::size_t coldPageBytes = 0; Nvfp4BoundaryRuntimeType runtimeType = Nvfp4BoundaryRuntimeType::kFloat16; }; diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt index 33e49bd5d927..e72249e383f0 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -6,4 +6,4 @@ target_include_directories( kv_cache_compression_src PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) set_property(TARGET kv_cache_compression_src PROPERTY POSITION_INDEPENDENT_CODE - ON) + ON) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp index 5ff5cba23636..9ad9f88f28a1 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp @@ -18,7 +18,6 @@ #include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" #include "tensorrt_llm/common/logger.h" -#include "tensorrt_llm/common/nvtxUtils.h" #include #include @@ -33,19 +32,20 @@ namespace tensorrt_llm::kv_cache_compression namespace { -constexpr std::size_t kCompactAlignment = 16U; -constexpr std::size_t kElementsPerBlockScale = 16U; +constexpr std::size_t kColdRecordAlignment = 16U; +constexpr std::size_t kElementsPerScaleGroup = 16U; constexpr std::size_t kPackedElementsPerByte = 2U; constexpr char const* kKeyRole = "key"; constexpr char const* kValueRole = "value"; -std::size_t alignUp(std::size_t value, std::size_t alignment) +// Descriptor-derived byte arithmetic must not wrap into an undersized cold record. +std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) { - if (value > std::numeric_limits::max() - (alignment - 1U)) + if (lhs > std::numeric_limits::max() - rhs) { - throw std::overflow_error("Cold Page size overflows size_t"); + throw std::overflow_error(label); } - return (value + alignment - 1U) / alignment * alignment; + return lhs + rhs; } std::size_t checkedMul(std::size_t lhs, std::size_t rhs, char const* label) @@ -57,22 +57,28 @@ std::size_t checkedMul(std::size_t lhs, std::size_t rhs, char const* label) return lhs * rhs; } -std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) +std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) { - if (lhs > std::numeric_limits::max() - rhs) + if (offset > std::numeric_limits::max() - base) { - throw std::overflow_error(label); + throw std::overflow_error("GPU buffer address overflows uintptr_t"); } - return lhs + rhs; + return base + offset; } -std::size_t scalarCount(Nvfp4ColdPageLayerConfig const& config) +std::size_t alignColdRecord(std::size_t bytes) +{ + return checkedAdd(bytes, kColdRecordAlignment - 1U, "Cold Page size overflows size_t") / kColdRecordAlignment + * kColdRecordAlignment; +} + +std::size_t checkedElementCount(Nvfp4ColdPageLayerConfig const& config) { if (config.numKvHeads <= 0 || config.tokensPerPage <= 0 || config.headDim <= 0) { throw std::invalid_argument("NVFP4 cold Page geometry must be positive"); } - if (config.headDim % static_cast(kElementsPerBlockScale) != 0) + if (config.headDim % static_cast(kElementsPerScaleGroup) != 0) { throw std::invalid_argument("NVFP4 cold Pages require headDim divisible by 16"); } @@ -95,32 +101,156 @@ kernels::Nvfp4BoundaryKernelParams makeKernelParams(Nvfp4ColdPageLayerConfig con return params; } -std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) +struct BufferLocation { - if (offset > std::numeric_limits::max() - base) + kv::PoolIndex poolIndex{0}; + std::size_t offset = 0; + std::size_t bytes = 0; + bool found = false; +}; + +struct AttentionLayerBuffers +{ + BufferLocation key; + BufferLocation value; + // Roles such as MLA index_key remain byte-exact. + std::vector sideBuffers; +}; + +using LayerConfigs = std::map; +using AttentionBufferMap = std::map; + +struct LifecycleBuffers +{ + AttentionBufferMap attention; + bool hasNonAttentionLayer = false; +}; + +LifecycleBuffers discoverLifecycleBuffers(kv::SlotDescVariant const& variant, LayerConfigs const& layerConfigs) +{ + LifecycleBuffers result; + for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) { - throw std::overflow_error("GPU buffer address overflows uintptr_t"); + auto const& coalesced = variant.coalescedBuffers.at(poolIndex); + std::size_t offset = 0U; + for (auto const& bufferId : coalesced.bufferIds) + { + auto const config = layerConfigs.find(bufferId.layerId); + if (config == layerConfigs.end()) + { + result.hasNonAttentionLayer = true; + } + else + { + auto& buffers = result.attention[bufferId.layerId]; + if (bufferId.role == kKeyRole || bufferId.role == kValueRole) + { + auto& location = bufferId.role == kKeyRole ? buffers.key : buffers.value; + if (location.found) + { + throw std::invalid_argument("GPU lifecycle contains a duplicate K/V buffer"); + } + location = BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}; + } + else + { + buffers.sideBuffers.push_back(BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}); + } + } + offset = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); + } } - return base + offset; + return result; +} + +struct AttentionPlan +{ + kernels::Nvfp4BoundaryPreparedPlan kernelPlan; + std::size_t coldPageBytes = 0; +}; + +AttentionPlan buildAttentionPlan( + kv::PoolGroupDesc const& gpuDesc, AttentionBufferMap const& layerBuffers, LayerConfigs const& layerConfigs) +{ + AttentionPlan result; + std::vector bufferPlans; + auto const runtimeType = layerConfigs.at(layerBuffers.begin()->first).runtimeType; + + for (auto const& [layerId, buffers] : layerBuffers) + { + // MHA/GQA has K and V; MLA exposes its latent KV as key only. + if (!buffers.key.found) + { + throw std::invalid_argument("Configured Attention layer is missing key"); + } + auto const& config = layerConfigs.at(layerId); + if (runtimeType != config.runtimeType) + { + throw std::invalid_argument("Attention lifecycle must use one runtime dtype"); + } + + auto const elements = checkedElementCount(config); + auto const layerOffset = result.coldPageBytes; + auto const packedBytes = elements / kPackedElementsPerByte; + auto const scaleBytes = elements / kElementsPerScaleGroup; + auto const compressedBufferCount = buffers.value.found ? 2U : 1U; + auto const packedRegionBytes = packedBytes * compressedBufferCount; + auto const scaleRegionBytes = scaleBytes * compressedBufferCount; + + // Per layer: [K packed | V packed? | K scale | V scale? | side buffers | padding]. + auto const scaleOffset = checkedAdd(layerOffset, packedRegionBytes, "NVFP4 cold Page size overflows size_t"); + auto coldOffset = checkedAdd(scaleOffset, scaleRegionBytes, "NVFP4 cold Page size overflows size_t"); + + auto appendNvfp4 = [&](BufferLocation const& location, std::size_t scaleIndex, std::size_t coldDataOffset, + std::size_t coldScaleOffset) + { + auto const& pool = gpuDesc.pools.at(location.poolIndex); + bufferPlans.push_back(kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), + pool.slotBytes, location.bytes, coldDataOffset, coldScaleOffset, 0U, 0U, + kernels::Nvfp4BoundaryTransform::kNvfp4, makeKernelParams(config, scaleIndex)}); + }; + + appendNvfp4(buffers.key, 0U, layerOffset, scaleOffset); + if (buffers.value.found) + { + appendNvfp4(buffers.value, 1U, layerOffset + packedBytes, scaleOffset + scaleBytes); + } + for (auto const& location : buffers.sideBuffers) + { + auto const& pool = gpuDesc.pools.at(location.poolIndex); + bufferPlans.push_back( + kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), pool.slotBytes, + location.bytes, coldOffset, 0U, 0U, 0U, kernels::Nvfp4BoundaryTransform::kLosslessCopy, {}}); + coldOffset = checkedAdd(coldOffset, location.bytes, "Lossless Attention side-buffer size overflows size_t"); + } + + auto const alignedEnd = alignColdRecord(coldOffset); + bufferPlans.back().coldPaddingOffset = coldOffset; + bufferPlans.back().coldPaddingBytes = static_cast(alignedEnd - coldOffset); + result.coldPageBytes = alignedEnd; + } + + result.kernelPlan = kernels::prepareNvfp4BoundaryPlan(bufferPlans, result.coldPageBytes, runtimeType); + return result; } } // namespace +// Validate and retain algorithm-owned metadata before KVCM creates its physical pools. Nvfp4ColdPageCodec::Nvfp4ColdPageCodec(std::vector layerConfigs) { + auto const validScales = [](auto const& scales) { + return std::all_of(scales.begin(), scales.end(), [](float value) { return std::isfinite(value) && value > 0; }); + }; for (auto& config : layerConfigs) { - static_cast(scalarCount(config)); if (config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFloat16 && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kBfloat16 && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3) { throw std::invalid_argument("Nvfp4ColdPageCodec received an unsupported runtime type"); } - auto const validScales = [](auto const& scales) { - return std::all_of( - scales.begin(), scales.end(), [](float value) { return std::isfinite(value) && value > 0; }); - }; + static_cast(checkedElementCount(config)); if (!validScales(config.nvfp4ScaleOrigQuant) || !validScales(config.nvfp4ScaleQuantOrig) || (config.runtimeType == kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3 && (!validScales(config.fp8ScaleOrigQuant) || !validScales(config.fp8ScaleQuantOrig)))) @@ -135,25 +265,11 @@ Nvfp4ColdPageCodec::Nvfp4ColdPageCodec(std::vector lay } } +// Bind the immutable metadata to KVCM's authoritative hot-pool layout exactly once. bool Nvfp4ColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept { try { - struct BufferLocation - { - kv::PoolIndex poolIndex{0}; - std::size_t offset = 0; - std::size_t bytes = 0; - bool found = false; - }; - - struct AttentionBuffers - { - BufferLocation key; - BufferLocation value; - std::vector losslessBuffers; - }; - // Use KVCM's default codec for non-Attention lifecycles. auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); if (!losslessCodec->configure(gpuDescs, numGpuDescs)) @@ -161,144 +277,48 @@ bool Nvfp4ColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGr throw std::invalid_argument("Default lossless codec rejected GPU layouts"); } - std::map pending; - std::set configuredAttentionLayers; + std::map pendingGroups; + std::set discoveredAttentionLayers; for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) { auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; for (auto const& variant : gpuDesc.slotDesc.variants) { - std::map attentionBuffers; - bool hasForeignLayerBuffer = false; - for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) - { - auto const& coalesced = variant.coalescedBuffers.at(poolIndex); - std::size_t offset = 0U; - for (auto const& bufferId : coalesced.bufferIds) - { - auto const config = mLayerConfigs.find(bufferId.layerId); - if (config == mLayerConfigs.end()) - { - hasForeignLayerBuffer = true; - } - else - { - configuredAttentionLayers.insert(bufferId.layerId); - auto& buffers = attentionBuffers[bufferId.layerId]; - if (bufferId.role == kKeyRole || bufferId.role == kValueRole) - { - auto& location = bufferId.role == kKeyRole ? buffers.key : buffers.value; - if (location.found) - { - throw std::invalid_argument("GPU lifecycle contains a duplicate K/V buffer"); - } - location = BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}; - } - else - { - buffers.losslessBuffers.push_back( - BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}); - } - } - offset - = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); - } - } + auto const buffers = discoverLifecycleBuffers(variant, mLayerConfigs); LayerGroupState state; - if (!attentionBuffers.empty()) + if (!buffers.attention.empty()) { - if (hasForeignLayerBuffer) + if (buffers.hasNonAttentionLayer) { - throw std::invalid_argument("Attention lifecycle mixes configured and unconfigured layers"); + throw std::invalid_argument("A lifecycle cannot mix Attention and non-Attention layers"); } - state.transform = Transform::kNvfp4Attention; - std::vector plans; - auto const runtimeType = mLayerConfigs.at(attentionBuffers.begin()->first).runtimeType; - for (auto const& [layerId, buffers] : attentionBuffers) + + auto plan = buildAttentionPlan(gpuDesc, buffers.attention, mLayerConfigs); + state.format = ColdPageFormat::kNvfp4Kv; + state.preparedPlan = std::move(plan.kernelPlan); + state.coldPageBytes = plan.coldPageBytes; + for (auto const& layer : buffers.attention) { - if (!buffers.key.found) - { - throw std::invalid_argument("Configured Attention layer is missing key"); - } - auto const& config = mLayerConfigs.at(layerId); - if (runtimeType != config.runtimeType) - { - throw std::invalid_argument("Attention lifecycle must use one runtime dtype"); - } - - auto const elements = scalarCount(config); - auto const rawElementBytes - = config.runtimeType == kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3 ? 1U : 2U; - auto const rawBytes - = checkedMul(elements, rawElementBytes, "Runtime KV Page size overflows size_t"); - if (buffers.key.bytes != rawBytes || (buffers.value.found && buffers.value.bytes != rawBytes)) - { - throw std::invalid_argument("GPU Attention buffer size does not match its geometry"); - } - - auto const layerOffset = state.coldPageBytes; - auto const packedBytesPerBuffer = elements / kPackedElementsPerByte; - auto const scaleBytesPerBuffer = elements / kElementsPerBlockScale; - auto const compressedBufferCount = buffers.value.found ? 2U : 1U; - auto const scaleOffset = checkedAdd(layerOffset, - checkedMul( - packedBytesPerBuffer, compressedBufferCount, "NVFP4 packed Page size overflows size_t"), - "NVFP4 cold Page size overflows size_t"); - - auto appendNvfp4 = [&](BufferLocation const& location, std::size_t scaleIndex, - std::size_t coldDataOffset, std::size_t coldScaleOffset) - { - auto const& pool = gpuDesc.pools.at(location.poolIndex); - plans.push_back( - kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), - pool.slotBytes, location.bytes, coldDataOffset, coldScaleOffset, 0U, 0U, - kernels::Nvfp4BoundaryTransform::kNvfp4, makeKernelParams(config, scaleIndex)}); - }; - - appendNvfp4(buffers.key, 0U, layerOffset, scaleOffset); - if (buffers.value.found) - { - appendNvfp4(buffers.value, 1U, - checkedAdd(layerOffset, packedBytesPerBuffer, "NVFP4 cold Page size overflows size_t"), - checkedAdd(scaleOffset, scaleBytesPerBuffer, "NVFP4 cold Page size overflows size_t")); - } - - auto coldOffset = checkedAdd(scaleOffset, - checkedMul( - scaleBytesPerBuffer, compressedBufferCount, "NVFP4 scale Page size overflows size_t"), - "NVFP4 cold Page size overflows size_t"); - for (auto const& location : buffers.losslessBuffers) - { - auto const& pool = gpuDesc.pools.at(location.poolIndex); - plans.push_back(kernels::Nvfp4BoundaryBufferPlan{ - checkedAddress(pool.baseAddress, location.offset), pool.slotBytes, location.bytes, - coldOffset, 0U, 0U, 0U, kernels::Nvfp4BoundaryTransform::kLossless, {}}); - coldOffset = checkedAdd( - coldOffset, location.bytes, "Lossless Attention side-buffer size overflows size_t"); - } - - auto const alignedEnd = alignUp(coldOffset, kCompactAlignment); - auto const paddingBytes = alignedEnd - coldOffset; - plans.back().coldPaddingOffset = coldOffset; - plans.back().coldPaddingBytes = static_cast(paddingBytes); - state.coldPageBytes = alignedEnd; + discoveredAttentionLayers.insert(layer.first); } - state.preparedPlan = kernels::prepareNvfp4BoundaryPlan(plans, state.coldPageBytes, runtimeType); } else { state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); } - pending.emplace(variant.lifeCycleId, std::move(state)); + if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) + { + throw std::invalid_argument("GPU lifecycle ID appears in multiple pool groups"); + } } } - if (configuredAttentionLayers.size() != mLayerConfigs.size()) + if (discoveredAttentionLayers.size() != mLayerConfigs.size()) { throw std::invalid_argument("A configured Attention layer is absent from all GPU descriptors"); } - mLayerGroups = std::move(pending); + mLayerGroups = std::move(pendingGroups); mLosslessCodec = std::move(losslessCodec); return true; } @@ -309,6 +329,13 @@ bool Nvfp4ColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGr } } +Nvfp4ColdPageCodec::LayerGroupState const* Nvfp4ColdPageCodec::findLayerGroup( + kv::LayerGroupId layerGroupId) const noexcept +{ + auto const found = mLayerGroups.find(layerGroupId); + return found == mLayerGroups.end() ? nullptr : &found->second; +} + std::size_t Nvfp4ColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept { auto const* state = findLayerGroup(layerGroupId); @@ -325,13 +352,6 @@ kv::PageIndexLocation Nvfp4ColdPageCodec::queryPageIndexLocation(kv::LayerGroupI return findLayerGroup(layerGroupId) == nullptr ? kv::PageIndexLocation::kBadLocation : kv::PageIndexLocation::kHost; } -Nvfp4ColdPageCodec::LayerGroupState const* Nvfp4ColdPageCodec::findLayerGroup( - kv::LayerGroupId layerGroupId) const noexcept -{ - auto const found = mLayerGroups.find(layerGroupId); - return found == mLayerGroups.end() ? nullptr : &found->second; -} - bool Nvfp4ColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { @@ -346,12 +366,11 @@ bool Nvfp4ColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, { return true; } - if (state->transform == Transform::kLosslessConcat) + if (state->format == ColdPageFormat::kLossless) { return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); } - NVTX3_SCOPED_RANGE(KVCC_OFFLOAD_COMPRESS_D2H); thread_local std::vector pages; pages.clear(); pages.reserve(numBasePages); @@ -383,12 +402,11 @@ bool Nvfp4ColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBa { return true; } - if (state->transform == Transform::kLosslessConcat) + if (state->format == ColdPageFormat::kLossless) { return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); } - NVTX3_SCOPED_RANGE(KVCC_ONBOARD_H2D_DECOMPRESS); thread_local std::vector pages; pages.clear(); pages.reserve(numBasePages); diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h index ef129f858207..93c318629868 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h @@ -67,15 +67,15 @@ class Nvfp4ColdPageCodec final : public kv::IKvCacheColdPageCodec std::size_t numBasePages, cudaStream_t stream) noexcept override; private: - enum class Transform + enum class ColdPageFormat { - kNvfp4Attention, - kLosslessConcat, + kNvfp4Kv, //!< NVFP4 K/V plus byte-exact Attention side buffers. + kLossless, //!< Entire lifecycle uses KVCM's default lossless concat. }; struct LayerGroupState { - Transform transform = Transform::kLosslessConcat; + ColdPageFormat format = ColdPageFormat::kLossless; kernels::Nvfp4BoundaryPreparedPlan preparedPlan; std::size_t coldPageBytes = 0; }; diff --git a/cpp/tests/unit_tests/CMakeLists.txt b/cpp/tests/unit_tests/CMakeLists.txt index 531bff2205cc..44291864657c 100644 --- a/cpp/tests/unit_tests/CMakeLists.txt +++ b/cpp/tests/unit_tests/CMakeLists.txt @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & +# SPDX-FileCopyrightText: Copyright (c) 2023-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 diff --git a/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp index a2130e87cd81..bd522cd80cff 100644 --- a/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp +++ b/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp @@ -113,6 +113,9 @@ class CudaStream return mStream; } + CudaStream(CudaStream const&) = delete; + CudaStream& operator=(CudaStream const&) = delete; + private: cudaStream_t mStream{}; }; @@ -337,24 +340,24 @@ std::uint32_t linearScaleOffset(std::uint32_t row, std::uint32_t scaleInRow, Pag return row * scalesPerRow + scaleInRow; } +constexpr std::array kE2m1Levels{0.0F, 0.5F, 1.0F, 1.5F, 2.0F, 3.0F, 4.0F, 6.0F}; + float e2m1Value(std::uint8_t nibble) { - constexpr std::array levels{0.0F, 0.5F, 1.0F, 1.5F, 2.0F, 3.0F, 4.0F, 6.0F}; - float const value = levels[nibble & 0x7U]; + float const value = kE2m1Levels[nibble & 0x7U]; return (nibble & 0x8U) != 0 ? -value : value; } //! Independent nearest-level oracle; fixtures avoid ties instead of duplicating production tie rules. std::uint8_t quantizeE2m1(float value) { - constexpr std::array levels{0.0F, 0.5F, 1.0F, 1.5F, 2.0F, 3.0F, 4.0F, 6.0F}; bool const negative = std::signbit(value); float const magnitude = std::abs(value); std::uint8_t best = 0; - float bestDistance = std::abs(magnitude - levels[0]); - for (std::uint8_t index = 1; index < levels.size(); ++index) + float bestDistance = std::abs(magnitude - kE2m1Levels[0]); + for (std::uint8_t index = 1; index < kE2m1Levels.size(); ++index) { - float const distance = std::abs(magnitude - levels[index]); + float const distance = std::abs(magnitude - kE2m1Levels[index]); if (distance < bestDistance) { best = index; @@ -397,6 +400,7 @@ std::vector makeRawPage(RawKind kind, std::size_t page, std::uint3 { normalizedValue = roundingMarginsPattern[index % roundingMarginsPattern.size()]; } + // kAllZero intentionally keeps the zero initializer. if (((page / 32) & 1U) != 0) { normalizedValue = -normalizedValue; @@ -919,8 +923,8 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) {reinterpret_cast(mla.data()), mlaRawBytes, mlaRawBytes, 0U, mlaPackedBytes, mlaPayloadBytes, static_cast(gapBeforeSide), Nvfp4BoundaryTransform::kNvfp4, params}, {reinterpret_cast(side.data()), sideSlotBytes, sideRawBytes, sideColdOffset, 0U, - sideColdEnd, static_cast(coldPageBytes - sideColdEnd), Nvfp4BoundaryTransform::kLossless, - {}}}; + sideColdEnd, static_cast(coldPageBytes - sideColdEnd), + Nvfp4BoundaryTransform::kLosslessCopy, {}}}; }; auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( makePlans(mlaInput, sideInput), coldPageBytes, runtimeType(kind)); @@ -1258,8 +1262,9 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) "zero FP8 quant scale", [](auto& params) { params.fp8ScaleOrigQuant = 0.0F; }, Nvfp4BoundaryRuntimeType::kFp8E4m3); expectInvalid( - "infinite FP8 dequant scale", [](auto& params) - { params.fp8ScaleQuantOrig = std::numeric_limits::infinity(); }, Nvfp4BoundaryRuntimeType::kFp8E4m3); + "infinite FP8 dequant scale", + [](auto& params) { params.fp8ScaleQuantOrig = std::numeric_limits::infinity(); }, + Nvfp4BoundaryRuntimeType::kFp8E4m3); } TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) @@ -1297,6 +1302,7 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({validOffload}, validPlan, nullptr, nullptr)); expectInvalid("unaligned raw stride", [](auto& buffer) { buffer.rawSlotBytes += alignof(uint4) / 2U; }); expectInvalid("raw bytes exceed stride", [](auto& buffer) { buffer.rawBytes = buffer.rawSlotBytes + 1U; }); + expectInvalid("raw bytes mismatch geometry", [](auto& buffer) { buffer.rawBytes -= alignof(uint4); }); expectInvalid("cold data interval exceeds page", [&](auto& buffer) { buffer.coldDataOffset = coldPageBytes; }); expectInvalid("cold intervals overlap", [](auto& buffer) { buffer.coldScaleOffset = buffer.coldDataOffset; }); EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes + alignof(uint4) / 2U))); diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index fdd5074c50ce..e442c2015239 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -16,7 +16,8 @@ set(NVFP4_COLD_PAGE_CODEC_TEST_SRC add_gtest(nvfp4ColdPageCodecTest "${NVFP4_COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) -target_link_libraries(nvfp4ColdPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) +target_link_libraries(nvfp4ColdPageCodecTest PRIVATE CUDA::cuda_driver + CUDA::cudart) target_include_directories( nvfp4ColdPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp index 339a2b277c04..e41ad149c26e 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp @@ -248,7 +248,7 @@ TEST(Nvfp4ColdPageCodecTest, KeyAndIndexAppendsLosslessIndexWithinTheLayerRecord EXPECT_EQ(index.coldDataOffset, 90U); EXPECT_EQ(index.coldPaddingOffset, 158U); EXPECT_EQ(index.coldPaddingBytes, 2U); - EXPECT_EQ(index.transform, kernels::Nvfp4BoundaryTransform::kLossless); + EXPECT_EQ(index.transform, kernels::Nvfp4BoundaryTransform::kLosslessCopy); } TEST(Nvfp4ColdPageCodecTest, FullAndSharedIndexerLayersHaveDistinctPerLayerRecords) @@ -371,6 +371,16 @@ TEST(Nvfp4ColdPageCodecTest, OnlyFp8RuntimeRequiresFp8Scales) EXPECT_THROW({ Nvfp4ColdPageCodec codec{layers}; }, std::invalid_argument); } +TEST(Nvfp4ColdPageCodecTest, Fp8SourceScalesDefaultToIdentity) +{ + auto layers = makeLayers(1); + layers.front().runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3; + + EXPECT_EQ(layers.front().fp8ScaleOrigQuant, (std::array{1.0F, 1.0F})); + EXPECT_EQ(layers.front().fp8ScaleQuantOrig, (std::array{1.0F, 1.0F})); + EXPECT_NO_THROW({ Nvfp4ColdPageCodec codec{layers}; }); +} + TEST(Nvfp4ColdPageCodecTest, DiscoversLifecycleMembershipAcrossPoolGroups) { auto layers = makeLayers(2); @@ -392,16 +402,6 @@ TEST(Nvfp4ColdPageCodecTest, RejectsConfiguredAttentionLayerAbsentFromAllGpuDesc EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1))); } -TEST(Nvfp4ColdPageCodecTest, RejectsAttentionBufferWithMismatchedGeometry) -{ - Nvfp4ColdPageCodec codec{makeLayers()}; - auto desc = makeAttentionDesc(); - desc.slotDesc.variants.front().coalescedBuffers[kv::PoolIndex{0}].singleBufferSize += 16U; - desc.pools[kv::PoolIndex{0}].slotBytes += 16U * kNumAttentionLayers; - - EXPECT_FALSE(configureOne(codec, desc)); -} - TEST(Nvfp4ColdPageCodecTest, CoalescedAttentionSideBufferUsesItsOwnBaseOffsetAndSlotStride) { resetLaunch(); @@ -423,7 +423,7 @@ TEST(Nvfp4ColdPageCodecTest, CoalescedAttentionSideBufferUsesItsOwnBaseOffsetAnd EXPECT_EQ(side.coldDataOffset, 180U); EXPECT_EQ(side.coldPaddingOffset, 500U); EXPECT_EQ(side.coldPaddingBytes, 12U); - EXPECT_EQ(side.transform, kernels::Nvfp4BoundaryTransform::kLossless); + EXPECT_EQ(side.transform, kernels::Nvfp4BoundaryTransform::kLosslessCopy); } TEST(Nvfp4ColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) @@ -526,6 +526,18 @@ TEST(Nvfp4ColdPageCodecTest, AttentionAndSsmSharingOneHotPoolGroupUseDifferentTr EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); } +TEST(Nvfp4ColdPageCodecTest, DuplicateLifecycleAcrossPoolGroupsIsRejected) +{ + Nvfp4ColdPageCodec codec{makeLayers(2)}; + std::array descs{ + makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1, 0), + makeAttentionDesc( + kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1, 1, kGpuKBase + kGpuSlotBytes, kGpuVBase + kGpuSlotBytes), + }; + + EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); +} + } // namespace } // namespace tensorrt_llm::kv_cache_compression diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py index 83cf2067bd62..c2d2e98e9de8 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/__init__.py @@ -1,6 +1,2 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -from .quantization_for_cold_page import ColdPageQuantizationCompression - -__all__ = ["ColdPageQuantizationCompression"] diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 86d03cf0bca1..3a23e157f8fa 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -29,8 +29,8 @@ LayerScales = tuple[ScalePair, ScalePair] _IDENTITY_NVFP4_SCALES: LayerScales = ((1.0, 1.0), (1.0, 1.0)) -_MODEL_OPT_KV_SCALE_KEY = re.compile( - r"(?:^|\.)layers\.(?P\d+)\.self_attn\." +_MODEL_OPT_LANGUAGE_KV_SCALE_KEY = re.compile( + r"^model(?:\.language_model)?\.layers\.(?P\d+)\.self_attn\." r"(?P[kv])_proj\.(?P=kind)_scale$" ) @@ -78,7 +78,7 @@ def _load_modelopt_nvfp4_scales( for file_path in weight_files: with safe_open(str(file_path), framework="pt", device="cpu") as checkpoint: for tensor_name in checkpoint.keys(): - match = _MODEL_OPT_KV_SCALE_KEY.search(tensor_name) + match = _MODEL_OPT_LANGUAGE_KV_SCALE_KEY.fullmatch(tensor_name) if match is None: continue value = float(checkpoint.get_tensor(tensor_name).reshape([]).item()) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 197f98ff992e..41eee837ca2b 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -41,6 +41,7 @@ supports_native_fp8_lora) from tensorrt_llm.logger import logger from tensorrt_llm.mapping import CpType, Mapping +from tensorrt_llm.quantization import QuantAlgo from ..attention_backend import get_sparse_attn_kv_cache_manager from ..hostfunc import set_low_latency_dispatch @@ -2057,22 +2058,17 @@ def build_managers(self, compression_manager = None compression_config = self._llm_args.kv_cache_compression_config - if compression_config is not None: - create_compression_manager = True - if compression_config.algorithm == "quantization_for_cold_page": - if estimating_kv_cache and not self._skip_est: - create_compression_manager = False - elif _uses_nvfp4_kv_cache(self._model_engine): - logger.info( - "Skipping cold-page NVFP4 quantization because the active " - "KV cache already uses NVFP4; KVCM will migrate its native " - "data and block-scale buffers losslessly.") - create_compression_manager = False - if create_compression_manager: - model_config = self._model_engine.model.model_config - compression_manager = create_kv_cache_compression_manager( - compression_config, - pretrained_config=model_config.pretrained_config) + is_cold_quantization = (compression_config is not None + and compression_config.algorithm + == "quantization_for_cold_page") + skip_compression_manager = is_cold_quantization and ( + (estimating_kv_cache and not self._skip_est) + or _uses_nvfp4_kv_cache(self._model_engine)) + if compression_config is not None and not skip_compression_manager: + model_config = self._model_engine.model.model_config + compression_manager = create_kv_cache_compression_manager( + compression_config, + pretrained_config=model_config.pretrained_config) cold_page_codec_provider = ( compression_manager if compression_manager is not None and compression_manager.provides_cold_page_codec else None) @@ -2795,8 +2791,13 @@ def validate_kv_cache_compression_compatibility( config: KvCacheCompressionConfig, kv_cache_config: KvCacheConfig, spec_config: Optional[SpeculativeConfig], + mapping: Optional[Mapping] = None, ) -> None: """Reject unsupported KV-cache compression feature combinations.""" + if mapping is not None and mapping.has_cp_helix(): + raise ValueError( + "KV-cache compression does not support HELIX context parallelism.") + if config.algorithm == "quantization_for_cold_page": from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND @@ -2808,6 +2809,9 @@ def validate_kv_cache_compression_compatibility( raise RuntimeError( "NVFP4 cold-page compression requires an SM100-family device " "(SM100 or SM103).") + elif config.algorithm == "triattention" and not is_sm_100f(): + raise RuntimeError( + "TriAttention requires an SM100-family device (SM100 or SM103).") if kv_cache_config.enable_block_reuse and not config.supports_block_reuse(): raise ValueError( @@ -2823,8 +2827,7 @@ def validate_kv_cache_compression_compatibility( "TriAttention requires eviction_mode='union'") mode = spec_config.spec_dec_mode if config.algorithm == "quantization_for_cold_page": - if not (mode.is_eagle3_one_model() - or mode.is_mtp_eagle_one_model()): + if not (mode.is_eagle3_one_model() or mode.is_mtp_eagle_one_model()): raise ValueError( "Cold-page quantization supports speculative decoding only " f"with one-model MTP-EAGLE or EAGLE3, not {mode.name}") @@ -2838,30 +2841,21 @@ def validate_kv_cache_compression_compatibility( def _uses_nvfp4_kv_cache(model_engine: PyTorchModelEngine) -> bool: quant_config = model_engine.model.model_config.quant_config return (quant_config is not None - and getattr(quant_config, "kv_cache_quant_algo", None) == "NVFP4") + and quant_config.kv_cache_quant_algo == QuantAlgo.NVFP4) def create_kv_cache_compression_manager( config: KvCacheCompressionConfig, pretrained_config: Optional["transformers.PretrainedConfig"] = None, ) -> Optional[KVCacheCompressionManager]: - """Build a KV-cache compression manager before KVCM construction. - - The caller may use the manager as a cold-page codec provider while building - KVCMs, then binds the completed target and optional draft managers. Only - iteration-driven implementations enter ``ResourceManager``. - """ + """Construct the configured compression manager before KVCM.""" if config.algorithm == "quantization_for_cold_page": - from ..kv_cache_compression.quantization_for_cold_page import \ - ColdPageQuantizationCompression + from ..kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import \ + ColdPageQuantizationCompression # noqa: E501 return ColdPageQuantizationCompression(config) if config.algorithm == "triattention": - if not is_sm_100f(): - raise RuntimeError( - "TriAttention requires an SM100-family device (SM100 or SM103)." - ) # TriAttention imports CuTe/CUTLASS; keep normal executor startup lazy. from ..kv_cache_compression.triattention.triattention import \ TriAttentionCompressionManager @@ -3653,11 +3647,17 @@ def validate_feature_combination(llm_args, model_engine, sampler_type): and compression_config.algorithm == "quantization_for_cold_page" and _uses_nvfp4_kv_cache(model_engine)) - if compression_config is not None and not cold_compression_is_redundant: + if cold_compression_is_redundant: + logger.info( + "Skipping cold-page NVFP4 quantization because the active KV cache " + "already uses NVFP4; KVCM will migrate its native data and " + "block-scale buffers losslessly.") + elif compression_config is not None: validate_kv_cache_compression_compatibility( compression_config, llm_args.kv_cache_config, model_engine.spec_config, + model_engine.mapping, ) def init_feature_status(llm_args) -> Dict[str, bool]: diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 635a33dd3d4a..2e0e391cb48e 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1120,26 +1120,23 @@ def append_to_kv_heads_per_layer( self.kv_cache_manager_py_config = config - def create_impl(cache_config): - manager_kwargs = {} - if cold_page_codec_provider is not None: - manager_kwargs["cold_page_codec"] = ( - cold_page_codec_provider.create_cold_page_codec( - cache_config, - runtime_dtype=self.dtype, - pp_layers=self.pp_layers, - num_kv_heads_per_layer=self.num_kv_heads_per_layer, - head_dim_per_layer=self.head_dim_per_layer, - ) + try: + cold_page_codec = ( + cold_page_codec_provider.create_cold_page_codec( + config, + runtime_dtype=self.dtype, + pp_layers=self.pp_layers, + num_kv_heads_per_layer=self.num_kv_heads_per_layer, + head_dim_per_layer=self.head_dim_per_layer, ) - return KVCacheManagerPy( - cache_config, + if cold_page_codec_provider is not None + else None + ) + self.impl = KVCacheManagerPy( + config, event_manager=self.event_manager, - **manager_kwargs, + cold_page_codec=cold_page_codec, ) - - try: - self.impl = create_impl(config) except (CuError, KVCacheOutOfMemoryError): if has_host_cache_tier: logger.warning( @@ -1152,7 +1149,22 @@ def create_impl(cache_config): ] config = replace(config, cache_tiers=cache_tiers_without_host) self.kv_cache_manager_py_config = config - self.impl = create_impl(config) + cold_page_codec = ( + cold_page_codec_provider.create_cold_page_codec( + config, + runtime_dtype=self.dtype, + pp_layers=self.pp_layers, + num_kv_heads_per_layer=self.num_kv_heads_per_layer, + head_dim_per_layer=self.head_dim_per_layer, + ) + if cold_page_codec_provider is not None + else None + ) + self.impl = KVCacheManagerPy( + config, + event_manager=self.event_manager, + cold_page_codec=cold_page_codec, + ) else: raise self.can_evict = len(config.cache_tiers) > 1 diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 3253c6c2cae0..116418e5ad72 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3706,10 +3706,17 @@ def supports_speculative_decoding(self) -> bool: return False +_KV_CACHE_COMPRESSION_ALGORITHM_TELEMETRY = TelemetryField.categorical( + "quantization_for_cold_page", "triattention") + + class ColdPageQuantizationCompressionConfig(KvCacheCompressionConfig): """Quantize Host and Disk KV pages without changing the active GPU cache.""" - algorithm: Literal["quantization_for_cold_page"] = "quantization_for_cold_page" + algorithm: Literal["quantization_for_cold_page"] = Field( + default="quantization_for_cold_page", + telemetry=False, + ) quant: Literal["nvfp4"] = Field( default="nvfp4", description="Quantization format stored in the compressed cache tier.") @@ -3722,11 +3729,11 @@ class ColdPageQuantizationCompressionConfig(KvCacheCompressionConfig): "K/V global scales. Omit it to use identity global scales.") def supports_block_reuse(self) -> bool: - # Compression changes representation and residency, not token identity. + """Block reuse is unchanged because token identity is preserved.""" return True def supports_speculative_decoding(self) -> bool: - # Each target or draft KVCM encodes its own pages at the storage boundary. + """Target and draft KVCMs encode their own cold pages independently.""" return True @@ -3740,7 +3747,10 @@ class TriAttentionKvCacheCompressionConfig(KvCacheCompressionConfig): changes_physical_kv_length: ClassVar[bool] = True - algorithm: Literal["triattention"] = "triattention" + algorithm: Literal["triattention"] = Field( + default="triattention", + telemetry=_KV_CACHE_COMPRESSION_ALGORITHM_TELEMETRY, + ) eviction_mode: Literal["union", "per_head", "per_layer_perhead"] = Field( default="union", description= diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py index 41c769c25cd9..3bafa3f368a1 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py @@ -255,7 +255,10 @@ def __init__( self, config: KVCacheManagerConfig, event_manager: "KVCacheEventManager | None" = None, + cold_page_codec: object | None = None, ) -> None: + if cold_page_codec is not None: + raise NotImplementedError("Cold-page codecs require the C++ KVCacheManagerV2 backend") init_cuda_once() config = deepcopy(config) self._init_config = config diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index b42702f438b8..8d54e1d9e742 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -299,7 +299,7 @@ "decode", "encode" ], - "annotation": "Literal['decode', 'encode']", + "annotation": "Literal['decode']", "converter": "", "kind": "categorical", "path": "cuda_graph_config.mode" @@ -634,8 +634,8 @@ "quantization_for_cold_page", "triattention" ], - "annotation": "Literal['quantization_for_cold_page', 'triattention']", - "converter": "", + "annotation": "Literal['triattention']", + "converter": "allowlist", "kind": "categorical", "path": "kv_cache_compression_config.algorithm" }, @@ -1558,7 +1558,7 @@ "rocket", "skip_softmax" ], - "annotation": "Literal['dsa', 'deepseek_v4', 'minimax_m3', 'rocket', 'skip_softmax']", + "annotation": "Literal['dsa']", "converter": "", "kind": "categorical", "path": "sparse_attention_config.algorithm" @@ -1882,7 +1882,7 @@ "SaveState", "User_Provided" ], - "annotation": "Literal['AUTO', 'DFlash', 'DSpark', 'Draft_Target', 'Eagle3', 'Eagle', 'Lookahead', 'MTP', 'Medusa', 'NGram', 'PARD', 'SA', 'SaveState', 'User_Provided']", + "annotation": "Literal['AUTO']", "converter": "", "kind": "categorical", "path": "speculative_config.decoding_type" diff --git a/tensorrt_llm/usage/llmapi_config.py b/tensorrt_llm/usage/llmapi_config.py index 7678a8308688..25c661af8d61 100644 --- a/tensorrt_llm/usage/llmapi_config.py +++ b/tensorrt_llm/usage/llmapi_config.py @@ -305,10 +305,9 @@ def build_capture_manifest(model_cls: type[BaseModel]) -> list[_ManifestEntry]: The single source of truth. Type-safe annotations auto-enroll; str/Any allowlist escape hatches opt in; telemetry=False opts out. Recurses into statically reachable nested BaseModels with a cycle guard. Collapses - duplicate keys (shared union-arm base fields): unions Literal annotations - and allowed_values across arms, and FAILS if two arms give a key a - different kind. Other duplicate annotations retain the first deterministic - representative. + duplicate keys (shared union-arm base fields): keeps the first by + (key, defining_class), unions allowed_values across arms, and FAILS if two + arms give a key a different kind. """ rows: list[dict[str, Any]] = [] @@ -340,7 +339,6 @@ def walk(cls: type, prefix: str, stack: tuple) -> None: rows.sort(key=lambda r: (r["key"], r["defining"])) first: dict[str, dict] = {} - union_annotations: dict[str, list[Any]] = {} union_allowed: dict[str, list[str]] = {} for r in rows: if r["key"] not in first: @@ -350,27 +348,15 @@ def walk(cls: type, prefix: str, stack: tuple) -> None: f"telemetry manifest: key '{r['key']}' has conflicting kinds " f"across union arms: {first[r['key']]['kind']} vs {r['kind']}" ) - annotations = union_annotations.setdefault(r["key"], []) - if r["annotation"] not in annotations: - annotations.append(r["annotation"]) seen = union_allowed.setdefault(r["key"], []) for v in r["allowed"]: if v not in seen: seen.append(v) - def merged_annotation(key: str) -> Any: - annotations = union_annotations[key] - if len(annotations) > 1 and all(_is_literal(ann) for ann in annotations): - values = tuple( - value for ann in annotations for value in get_args(ann) - ) - return Literal[values] - return annotations[0] - entries = [ _ManifestEntry( path=key, - annotation=merged_annotation(key), + annotation=r["annotation"], kind=r["kind"], converter=r["converter"], allowed_values=tuple(union_allowed[key]), diff --git a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py index 530109421692..f9675e56e375 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py +++ b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py @@ -450,6 +450,7 @@ def test_no_disk_cache_leaves_none(self): assert target_config is c._kv_cache_config assert draft_config is None + class TestHostSplitIgnoresGpuFixedCost: """The fixed cost models GPU-resident state and is not host memory.""" diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index cc5d4619dc7d..2f1609821631 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -17,7 +17,11 @@ ResourceManager, ResourceManagerType, ) -from tensorrt_llm.llmapi.llm_args import KvCacheCompressionConfig +from tensorrt_llm.llmapi.llm_args import ( + ColdPageQuantizationCompressionConfig, + KvCacheCompressionConfig, + TriAttentionKvCacheCompressionConfig, +) pytestmark = pytest.mark.cpu_only @@ -285,7 +289,11 @@ def test_triattention_requires_sm100_family(self): patch.object(util_mod, "is_sm_100f", return_value=False), pytest.raises(RuntimeError, match="SM100-family"), ): - create_kv_cache_compression_manager(cfg) + util_mod.validate_kv_cache_compression_compatibility( + cfg, + SimpleNamespace(enable_block_reuse=False), + None, + ) def test_capabilities_default_false(self): config = KvCacheCompressionConfig(algorithm="offload") @@ -326,7 +334,7 @@ def test_estimation_still_creates_triattention_manager(self): creator = object.__new__(util_mod.KvCacheCreator) creator._skip_est = False creator._max_seq_len = 1024 - creator._kv_cache_config = SimpleNamespace() + creator._kv_cache_config = SimpleNamespace(host_cache_size=None, disk_cache_size=None) creator._llm_args = SimpleNamespace(kv_cache_compression_config=config) creator._model_engine = SimpleNamespace( model=SimpleNamespace(model_config=SimpleNamespace(pretrained_config=pretrained_config)) @@ -413,8 +421,7 @@ def test_build_routes_compression_manager_by_capabilities(provides_cold_page_cod with patch.object( util_mod, "create_kv_cache_compression_manager", - side_effect=lambda *_args, **_kwargs: build_order.append("factory") - or compression_manager, + side_effect=lambda *_args, **_kwargs: build_order.append("factory") or compression_manager, ) as factory: creator.build_managers(resources) @@ -468,6 +475,26 @@ def test_names_not_in_sparse_module(self): class TestCompressionCompatibility: + @pytest.mark.parametrize( + "config", + ( + ColdPageQuantizationCompressionConfig(), + TriAttentionKvCacheCompressionConfig(), + ), + ) + def test_helix_is_rejected(self, config): + model_engine = SimpleNamespace( + mapping=SimpleNamespace(has_cp_helix=lambda: True), + spec_config=None, + model=SimpleNamespace(model_config=SimpleNamespace(quant_config=None)), + ) + llm_args = SimpleNamespace( + kv_cache_compression_config=config, + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) + with pytest.raises(ValueError, match="HELIX"): + util_mod.validate_feature_combination(llm_args, model_engine, None) + def test_raises_when_reuse_on(self): config = _compression_config() with pytest.raises(ValueError, match="block reuse"): diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 3cbaa78086c0..8e45b4da9c77 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -10,17 +10,16 @@ import torch from safetensors.torch import save_file -from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page import ( +from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import ( ColdPageQuantizationCompression, ) from tensorrt_llm._torch.pyexecutor import _util as util_mod from tensorrt_llm._torch.pyexecutor.resource_manager import DataType from tensorrt_llm._torch.speculative.interface import SpeculativeDecodingMode from tensorrt_llm._torch.speculative.utils import update_spec_config_from_model_config -from tensorrt_llm.llmapi.llm_args import ( - ColdPageQuantizationCompressionConfig, - MTPDecodingConfig, -) +from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig, MTPDecodingConfig +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization import QuantAlgo from tensorrt_llm.runtime import kv_cache_manager_v2 as runtime_v2_mod from tensorrt_llm.runtime.kv_cache_manager_v2 import ( AttentionLayerConfig, @@ -56,10 +55,7 @@ def _cache_config(*layers): def _native(): def layer_config(): - return SimpleNamespace( - fp8_scale_orig_quant=(1.0, 1.0), - fp8_scale_quant_orig=(1.0, 1.0), - ) + return SimpleNamespace() codec = MagicMock() module = SimpleNamespace( @@ -129,8 +125,7 @@ def test_optional_modelopt_scales_map_pp_layers_and_ignore_local_draft_id(tmp_pa assert [config.layer_id for config in configs] == [0, 1, 2] assert [config.runtime_type for config in configs] == ["native-bf16"] * 3 assert [ - (config.nvfp4_scale_orig_quant, config.nvfp4_scale_quant_orig) - for config in configs + (config.nvfp4_scale_orig_quant, config.nvfp4_scale_quant_orig) for config in configs ] == [ ((2.0, 4.0), (0.5, 0.25)), ((8.0, 16.0), (0.125, 0.0625)), @@ -221,6 +216,24 @@ def test_scale_loader_reduces_duplicate_shards_like_native_qkv_loader(tmp_path): ) +def test_scale_loader_ignores_multimodal_towers_with_the_same_layer_id(tmp_path): + _write_quant_metadata(tmp_path) + tensors = { + "model.language_model.layers.7.self_attn.k_proj.k_scale": torch.tensor(0.5), + "model.language_model.layers.7.self_attn.v_proj.v_scale": torch.tensor(0.25), + "model.vision_tower.encoder.layers.7.self_attn.k_proj.k_scale": torch.tensor(4.0), + "model.vision_tower.encoder.layers.7.self_attn.v_proj.v_scale": torch.tensor(2.0), + "model.audio_tower.layers.7.self_attn.k_proj.k_scale": torch.tensor(8.0), + "model.audio_tower.layers.7.self_attn.v_proj.v_scale": torch.tensor(4.0), + } + save_file(tensors, str(tmp_path / "model.safetensors")) + + assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + (2.0, 4.0), + (0.5, 0.25), + ) + + def test_trtllm_load_kv_scales_zero_uses_identity(tmp_path, monkeypatch): _write_scales(tmp_path, {7: (0.5, 0.25)}) monkeypatch.setenv("TRTLLM_LOAD_KV_SCALES", "0") @@ -314,7 +327,7 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): assert config.nvfp4_scale_orig_quant == config.nvfp4_scale_quant_orig == (1.0, 1.0) -def test_fp8_runtime_uses_native_unit_source_scale_default(tmp_path): +def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): native, _ = _native() _write_scales(tmp_path, {10: (0.5, 0.25)}) with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): @@ -327,7 +340,8 @@ def test_fp8_runtime_uses_native_unit_source_scale_default(tmp_path): ) config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert config.fp8_scale_orig_quant == config.fp8_scale_quant_orig == (1.0, 1.0) + assert config.runtime_type == "native-fp8" + assert config.nvfp4_scale_orig_quant == (2.0, 4.0) assert config.nvfp4_scale_quant_orig == (0.5, 0.25) @@ -385,16 +399,12 @@ def build(*, estimating=False, active_kv_quant=None): creator = object.__new__(util_mod.KvCacheCreator) creator._skip_est = False creator._max_seq_len = 1024 - creator._kv_cache_config = SimpleNamespace() + creator._kv_cache_config = SimpleNamespace(host_cache_size=None, disk_cache_size=None) creator._llm_args = SimpleNamespace( kv_cache_compression_config=ColdPageQuantizationCompressionConfig() ) - model_config = SimpleNamespace( - quant_config=active_kv_quant, pretrained_config=object() - ) - creator._model_engine = SimpleNamespace( - model=SimpleNamespace(model_config=model_config) - ) + model_config = SimpleNamespace(quant_config=active_kv_quant, pretrained_config=object()) + creator._model_engine = SimpleNamespace(model=SimpleNamespace(model_config=model_config)) creator._draft_model_engine = None creator._kv_connector_manager = None creator._fp8_ctx_mla_kv_len_cap = None @@ -412,8 +422,6 @@ def build(*, estimating=False, active_kv_quant=None): assert not manager.uses_iteration_lifecycle for kwargs in ( {"estimating": True}, - {"active_kv_quant": SimpleNamespace(kv_cache_quant_algo="NVFP4")}, + {"active_kv_quant": QuantConfig(kv_cache_quant_algo=QuantAlgo.NVFP4)}, ): - assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in build( - **kwargs - ) + assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in build(**kwargs) diff --git a/tests/unittest/_torch/kv_cache_compression/test_triattention_draft_cocompaction.py b/tests/unittest/_torch/kv_cache_compression/test_triattention_draft_cocompaction.py index 6824e5d1ce83..b67ab51b2b90 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_triattention_draft_cocompaction.py +++ b/tests/unittest/_torch/kv_cache_compression/test_triattention_draft_cocompaction.py @@ -91,7 +91,13 @@ def test_speculative_admission_gates_raise(gate, match): eviction_mode="per_head" if gate == "union_only_per_head" else "union", ) - with pytest.raises(ValueError, match=match): + with ( + mock.patch( + "tensorrt_llm._torch.pyexecutor._util.is_sm_100f", + return_value=True, + ), + pytest.raises(ValueError, match=match), + ): validate_kv_cache_compression_compatibility( config, SimpleNamespace(enable_block_reuse=False), diff --git a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py index 5c38c121a5e4..e2fd555a7723 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py +++ b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py @@ -74,11 +74,6 @@ def test_factory_allows_block_reuse_and_propagates_config_fields(self): # The factory contract is independent of GPU-owned persistent buffers. fake_v2 = _make_fake_v2(enable_block_reuse=True) cfg = _make_tri_config(budget=32, beta=16, eviction_mode="per_head") - validate_kv_cache_compression_compatibility( - cfg, - SimpleNamespace(enable_block_reuse=True), - None, - ) with ( mock.patch( "tensorrt_llm._torch.pyexecutor._util.is_sm_100f", @@ -88,6 +83,11 @@ def test_factory_allows_block_reuse_and_propagates_config_fields(self): TriAttentionCompressionManager, "_initialize_eviction_state" ) as initialize, ): + validate_kv_cache_compression_compatibility( + cfg, + SimpleNamespace(enable_block_reuse=True), + None, + ) mgr = create_kv_cache_compression_manager( cfg, pretrained_config=_make_test_pretrained_config(), @@ -460,11 +460,15 @@ def test_one_model_draft_co_compression_is_accepted(self, spec_mode): ) ) - validate_kv_cache_compression_compatibility( - _make_tri_config(budget=8), - SimpleNamespace(enable_block_reuse=False), - spec_config, - ) + with mock.patch( + "tensorrt_llm._torch.pyexecutor._util.is_sm_100f", + return_value=True, + ): + validate_kv_cache_compression_compatibility( + _make_tri_config(budget=8), + SimpleNamespace(enable_block_reuse=False), + spec_config, + ) class TestFixedScoreMetadata: diff --git a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py index ab9f77da6aed..7b58c5b52832 100644 --- a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py +++ b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py @@ -287,7 +287,7 @@ def test_build_capture_manifest_matches_committed_golden(): def test_kv_cache_compression_discriminator_captures_both_algorithms(): - """Flattened telemetry preserves both discriminator Literal values.""" + """The shared allowlist captures either compression discriminator.""" from tensorrt_llm.llmapi.llm_args import ( ColdPageQuantizationCompressionConfig, TorchLlmArgs, @@ -303,23 +303,20 @@ def test_kv_cache_compression_discriminator_captures_both_algorithms(): for item in build_capture_manifest(TorchLlmArgs) if item.path == "kv_cache_compression_config.algorithm" ) - assert repr(entry.annotation) == ( - "typing.Literal['quantization_for_cold_page', 'triattention']" - ) - assert entry.converter == "" + assert repr(entry.annotation) == "typing.Literal['triattention']" + assert entry.converter == "allowlist" assert set(entry.allowed_values) == { "quantization_for_cold_page", "triattention", } - for config_cls, annotation in ( - ( - ColdPageQuantizationCompressionConfig, - "typing.Literal['quantization_for_cold_page']", - ), - (TriAttentionKvCacheCompressionConfig, "typing.Literal['triattention']"), - ): - field = config_cls.model_fields["algorithm"] - assert repr(field.annotation) == annotation + cold_field = ColdPageQuantizationCompressionConfig.model_fields["algorithm"] + assert cold_field.json_schema_extra["telemetry"] == {"exclude": True} + tri_field = TriAttentionKvCacheCompressionConfig.model_fields["algorithm"] + assert tri_field.json_schema_extra["telemetry"] == { + "kind": "categorical", + "converter": "allowlist", + "allowed_values": ["quantization_for_cold_page", "triattention"], + } configs = ( ColdPageQuantizationCompressionConfig(), From f593902c25ede87c0b3f00168791322399672f70 Mon Sep 17 00:00:00 2001 From: tianruih Date: Mon, 24 Aug 2026 02:35:25 -0700 Subject: [PATCH 04/29] [None][test] Fix compression compatibility test setup --- .../_torch/executor/test_kv_cache_compression_manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index 2f1609821631..8383f8e3b1a6 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -479,7 +479,7 @@ class TestCompressionCompatibility: "config", ( ColdPageQuantizationCompressionConfig(), - TriAttentionKvCacheCompressionConfig(), + TriAttentionKvCacheCompressionConfig(calibration_path="triattention-calibration.pt"), ), ) def test_helix_is_rejected(self, config): From 2f8e95584a612103ce1bd166d6a228d50d15b9c6 Mon Sep 17 00:00:00 2001 From: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:23:47 -0700 Subject: [PATCH 05/29] [None][fix] Tighten KV cache compression admission Signed-off-by: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/_util.py | 10 +++--- .../test_kv_cache_compression_manager.py | 31 ++++++++++++++++--- .../test_llmapi_config_telemetry_docs.py | 17 ++++++++-- 3 files changed, 45 insertions(+), 13 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 41eee837ca2b..9e9a60fd0a4e 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -2791,13 +2791,8 @@ def validate_kv_cache_compression_compatibility( config: KvCacheCompressionConfig, kv_cache_config: KvCacheConfig, spec_config: Optional[SpeculativeConfig], - mapping: Optional[Mapping] = None, ) -> None: """Reject unsupported KV-cache compression feature combinations.""" - if mapping is not None and mapping.has_cp_helix(): - raise ValueError( - "KV-cache compression does not support HELIX context parallelism.") - if config.algorithm == "quantization_for_cold_page": from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND @@ -3643,6 +3638,10 @@ def _adjust_torch_mem_fraction(): def validate_feature_combination(llm_args, model_engine, sampler_type): # Validate the flags for features' combination compression_config = llm_args.kv_cache_compression_config + if (compression_config is not None and model_engine.mapping.has_cp_helix()): + # TODO: Revisit after KVCC validates HELIX-sharded Page ownership and migration. + raise ValueError( + "KV-cache compression does not support HELIX context parallelism.") cold_compression_is_redundant = (compression_config is not None and compression_config.algorithm == "quantization_for_cold_page" @@ -3657,7 +3656,6 @@ def validate_feature_combination(llm_args, model_engine, sampler_type): compression_config, llm_args.kv_cache_config, model_engine.spec_config, - model_engine.mapping, ) def init_feature_status(llm_args) -> Dict[str, bool]: diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index 8383f8e3b1a6..40e449cf5f2d 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -22,6 +22,7 @@ KvCacheCompressionConfig, TriAttentionKvCacheCompressionConfig, ) +from tensorrt_llm.quantization import QuantAlgo pytestmark = pytest.mark.cpu_only @@ -327,7 +328,7 @@ def test_spec_gate_uses_config_capability(self): class TestKvCacheCreatorLifecycle: - def test_estimation_still_creates_triattention_manager(self): + def test_estimation_still_creates_triattention_manager(self) -> None: config = SimpleNamespace(algorithm="triattention") pretrained_config = object() expected_manager = SimpleNamespace(provides_cold_page_codec=False) @@ -360,7 +361,7 @@ def test_estimation_still_creates_triattention_manager(self): assert resources[ResourceManagerType.KV_CACHE_MANAGER] is target_manager assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is expected_manager - def test_teardown_pops_and_shuts_down_compression_manager(self): + def test_teardown_pops_and_shuts_down_compression_manager(self) -> None: creator = object.__new__(util_mod.KvCacheCreator) compression_manager = MagicMock() resources = { @@ -380,7 +381,9 @@ def test_teardown_pops_and_shuts_down_compression_manager(self): (True, False), ids=("cold-codec", "iteration-manager"), ) -def test_build_routes_compression_manager_by_capabilities(provides_cold_page_codec): +def test_build_routes_compression_manager_by_capabilities( + provides_cold_page_codec: bool, +) -> None: creator = object.__new__(util_mod.KvCacheCreator) creator._skip_est = False creator._max_seq_len = 1024 @@ -482,7 +485,10 @@ class TestCompressionCompatibility: TriAttentionKvCacheCompressionConfig(calibration_path="triattention-calibration.pt"), ), ) - def test_helix_is_rejected(self, config): + def test_helix_is_rejected( + self, + config: KvCacheCompressionConfig, + ) -> None: model_engine = SimpleNamespace( mapping=SimpleNamespace(has_cp_helix=lambda: True), spec_config=None, @@ -495,6 +501,23 @@ def test_helix_is_rejected(self, config): with pytest.raises(ValueError, match="HELIX"): util_mod.validate_feature_combination(llm_args, model_engine, None) + def test_helix_is_rejected_before_redundant_cold_quantization(self) -> None: + model_engine = SimpleNamespace( + mapping=SimpleNamespace(has_cp_helix=lambda: True), + spec_config=None, + model=SimpleNamespace( + model_config=SimpleNamespace( + quant_config=SimpleNamespace(kv_cache_quant_algo=QuantAlgo.NVFP4) + ) + ), + ) + llm_args = SimpleNamespace( + kv_cache_compression_config=ColdPageQuantizationCompressionConfig(), + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) + with pytest.raises(ValueError, match="HELIX"): + util_mod.validate_feature_combination(llm_args, model_engine, None) + def test_raises_when_reuse_on(self): config = _compression_config() with pytest.raises(ValueError, match="block reuse"): diff --git a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py index 7b58c5b52832..e6e8f68b7c83 100644 --- a/tests/unittest/usage/test_llmapi_config_telemetry_docs.py +++ b/tests/unittest/usage/test_llmapi_config_telemetry_docs.py @@ -286,7 +286,7 @@ def test_build_capture_manifest_matches_committed_golden(): _assert_committed_manifest_current(golden_manifest()) -def test_kv_cache_compression_discriminator_captures_both_algorithms(): +def test_kv_cache_compression_discriminator_captures_both_algorithms() -> None: """The shared allowlist captures either compression discriminator.""" from tensorrt_llm.llmapi.llm_args import ( ColdPageQuantizationCompressionConfig, @@ -318,10 +318,16 @@ def test_kv_cache_compression_discriminator_captures_both_algorithms(): "allowed_values": ["quantization_for_cold_page", "triattention"], } + private_paths = ( + "/private/modelopt-scales", + "/private/triattention.pt", + ) configs = ( - ColdPageQuantizationCompressionConfig(), + ColdPageQuantizationCompressionConfig( + scale_checkpoint_path=private_paths[0], + ), TriAttentionKvCacheCompressionConfig( - calibration_path="/triattention.pt", + calibration_path=private_paths[1], ), ) for config in configs: @@ -333,6 +339,11 @@ def test_kv_cache_compression_discriminator_captures_both_algorithms(): captured = json.loads(config_json) metadata = json.loads(metadata_json) assert captured["kv_cache_compression_config.algorithm"] == config.algorithm + assert "kv_cache_compression_config.scale_checkpoint_path" not in captured + assert "kv_cache_compression_config.calibration_path" not in captured + for private_path in private_paths: + assert private_path not in config_json + assert private_path not in metadata_json assert metadata["capture_succeeded"] is True From 8bb2ed4bdffb9722144ef35a62769ff552868c88 Mon Sep 17 00:00:00 2001 From: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> Date: Mon, 24 Aug 2026 22:17:24 -0700 Subject: [PATCH 06/29] [None][test] Fix KV cache compression CI routing Signed-off-by: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> --- tests/integration/test_lists/test-db/l0_a10.yml | 1 - tests/unittest/_torch/executor/test_kv_cache_manager_v2.py | 1 + 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index 5b39eea23efa..da4fb68e1aa8 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -29,7 +29,6 @@ l0_a10: - unittest/_torch/sampler/test_trtllm_sampler.py - unittest/_torch/sampler/test_token_ban.py - unittest/_torch/executor/test_disagg_index_mapper_early_release.py - - unittest/_torch/executor/test_kv_cache_compression_manager.py - unittest/_torch/executor/test_kv_cache_v2_capacity_only.py - unittest/_torch/executor/test_error_classification.py - unittest/_torch/modules/moe/test_communication_factory.py diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index 52231f4d7bb4..c6785a2fe7f7 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -148,6 +148,7 @@ def build_cache_config( module = "tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2" with ( patch(f"{module}.CuError", _CacheTierInitError), + patch(f"{module}.IndexMapper"), patch(f"{module}.KVCacheManagerPy", impl_constructor), patch.object(KVCacheManagerV2, "_build_base_config", build_base_config), patch.object(KVCacheManagerV2, "_build_cache_config", build_cache_config), From d9fcc3c76d7419cfa23fc2795ee52dcd6df7429f Mon Sep 17 00:00:00 2001 From: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> Date: Mon, 24 Aug 2026 23:13:12 -0700 Subject: [PATCH 07/29] [None][test] Preserve existing compression test routing Signed-off-by: Tianrui 'Hudayday' Hu <32944717+Hudayday@users.noreply.github.com> --- tests/integration/test_lists/test-db/l0_a10.yml | 1 + .../_torch/executor/test_kv_cache_compression_manager.py | 6 ++++-- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index da4fb68e1aa8..5b39eea23efa 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -29,6 +29,7 @@ l0_a10: - unittest/_torch/sampler/test_trtllm_sampler.py - unittest/_torch/sampler/test_token_ban.py - unittest/_torch/executor/test_disagg_index_mapper_early_release.py + - unittest/_torch/executor/test_kv_cache_compression_manager.py - unittest/_torch/executor/test_kv_cache_v2_capacity_only.py - unittest/_torch/executor/test_error_classification.py - unittest/_torch/modules/moe/test_communication_factory.py diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index 40e449cf5f2d..df6f2abfea9d 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -24,8 +24,6 @@ ) from tensorrt_llm.quantization import QuantAlgo -pytestmark = pytest.mark.cpu_only - # ---------------------------------------------------------------------- # # Mock infra: in-memory managers / requests (avoid touching V2 / model). # # ---------------------------------------------------------------------- # @@ -327,6 +325,7 @@ def test_spec_gate_uses_config_capability(self): ) +@pytest.mark.cpu_only class TestKvCacheCreatorLifecycle: def test_estimation_still_creates_triattention_manager(self) -> None: config = SimpleNamespace(algorithm="triattention") @@ -376,6 +375,7 @@ def test_teardown_pops_and_shuts_down_compression_manager(self) -> None: assert ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in resources +@pytest.mark.cpu_only @pytest.mark.parametrize( "provides_cold_page_codec", (True, False), @@ -478,6 +478,7 @@ def test_names_not_in_sparse_module(self): class TestCompressionCompatibility: + @pytest.mark.cpu_only @pytest.mark.parametrize( "config", ( @@ -501,6 +502,7 @@ def test_helix_is_rejected( with pytest.raises(ValueError, match="HELIX"): util_mod.validate_feature_combination(llm_args, model_engine, None) + @pytest.mark.cpu_only def test_helix_is_rejected_before_redundant_cold_quantization(self) -> None: model_engine = SimpleNamespace( mapping=SimpleNamespace(has_cp_helix=lambda: True), From 50db0dcd1e278298aa3330974b997ff148d5a4a8 Mon Sep 17 00:00:00 2001 From: tianruih Date: Tue, 25 Aug 2026 15:02:48 -0700 Subject: [PATCH 08/29] [None][refactor] Move NVFP4 cold-page policy to Python Signed-off-by: tianruih --- ...daryKernels.cu => nvfp4ColdPageKernels.cu} | 126 ++-- ...undaryKernels.h => nvfp4ColdPageKernels.h} | 40 +- .../kv_cache_compression/CMakeLists.txt | 2 +- .../kv_cache_compression/coldPageCodec.cpp | 416 +++++++++++++ .../kv_cache_compression/coldPageCodec.h | 108 ++++ .../nvfp4ColdPageCodec.cpp | 427 ------------- .../kv_cache_compression/nvfp4ColdPageCodec.h | 90 --- .../nanobind/kvCacheCompression/bindings.cpp | 65 +- cpp/tests/unit_tests/kernels/CMakeLists.txt | 12 +- ...sTest.cpp => nvfp4ColdPageKernelsTest.cpp} | 241 ++++---- .../kv_cache_compression/CMakeLists.txt | 15 +- .../coldPageCodecTest.cpp | 391 ++++++++++++ .../nvfp4ColdPageCodecTest.cpp | 584 ------------------ .../quantization_for_cold_page/nvfp4.py | 230 +++++++ .../quantization_for_cold_page.py | 169 +---- tensorrt_llm/_torch/pyexecutor/_util.py | 2 +- .../_torch/pyexecutor/kv_cache_manager_v2.py | 2 + .../_torch/pyexecutor/resource_manager.py | 5 +- tensorrt_llm/llmapi/llm_args.py | 2 +- .../executor/test_kv_cache_manager_v2.py | 17 + .../test_quantization_for_cold_page.py | 378 ++++++++++-- 21 files changed, 1787 insertions(+), 1535 deletions(-) rename cpp/tensorrt_llm/kernels/{nvfp4BoundaryKernels.cu => nvfp4ColdPageKernels.cu} (91%) rename cpp/tensorrt_llm/kernels/{nvfp4BoundaryKernels.h => nvfp4ColdPageKernels.h} (68%) create mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp create mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h rename cpp/tests/unit_tests/kernels/{nvfp4BoundaryKernelsTest.cpp => nvfp4ColdPageKernelsTest.cpp} (86%) create mode 100644 cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp delete mode 100644 cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp create mode 100644 tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu similarity index 91% rename from cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu rename to cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index 671f5243686a..d24bc4144fe5 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -16,7 +16,7 @@ * limitations under the License. */ -#include "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" #include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/cudaUtils.h" @@ -51,7 +51,7 @@ constexpr std::uint32_t kHostLoadAsyncStages = 8; constexpr std::uint32_t kMappedHostGridSplits = 1; constexpr std::uint32_t kMaxTasksPerLaunch = 256; // Bound by-value buffer metadata to CUDA's kernel-parameter limit. -constexpr std::uint32_t kMaxBuffersPerLaunch = kNvfp4BoundaryMaxBuffersPerLaunch; +constexpr std::uint32_t kMaxBuffersPerLaunch = kNvfp4ColdPageMaxBuffersPerLaunch; constexpr std::uint32_t kElementsPerHalfGroup = 8; constexpr std::uint32_t kElementsPerScaleGroup = 16; constexpr std::uint32_t kHalfGroupsPerScaleGroup = kElementsPerScaleGroup / kElementsPerHalfGroup; @@ -70,16 +70,16 @@ static_assert( kMaxScaleBytesPerTile % sizeof(uint4) == 0, "A full scale tile must preserve the 16-byte transfer fast path"); // Buffer plans are passed to CUDA kernels as raw by-value arguments. -static_assert(std::is_trivially_copyable_v, "Buffer plans must remain raw-copyable"); +static_assert(std::is_trivially_copyable_v, "Buffer plans must remain raw-copyable"); // Keep both kernel argument packs within CUDA's modern 32,764-byte limit. -static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) +static_assert(sizeof(std::array) + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + 3U * sizeof(std::uint32_t) <= kKernelParameterLimitBytes, "Offload kernel arguments exceed CUDA's parameter limit"); -static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) +static_assert(sizeof(std::array) + + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + 3U * sizeof(std::uint32_t) <= kKernelParameterLimitBytes, "Onboard kernel arguments exceed CUDA's parameter limit"); @@ -90,7 +90,7 @@ static_assert(sizeof(std::array __device__ __forceinline__ void copyAsyncGlobalToShared(T* shared, T const* global, bool valid) { - static_assert(sizeof(T) == 16, "Boundary transfer grains must match batchedCopy's 16-byte width"); + static_assert(sizeof(T) == 16, "Cold-page transfer grains must match batchedCopy's 16-byte width"); if (valid) { asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" @@ -115,8 +115,8 @@ struct OnboardBufferTask std::uint8_t* raw; }; -__device__ OffloadBufferTask resolveTask(Nvfp4BoundaryOffloadPageTask const& page, - Nvfp4BoundaryBufferPlan const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) +__device__ OffloadBufferTask resolveTask(Nvfp4ColdPageOffloadPageTask const& page, + Nvfp4ColdPageBufferPlan const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.gpuPageIndex); auto* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; @@ -124,8 +124,8 @@ __device__ OffloadBufferTask resolveTask(Nvfp4BoundaryOffloadPageTask const& pag coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, coldPage + buffer.coldPaddingOffset}; } -__device__ OnboardBufferTask resolveTask(Nvfp4BoundaryOnboardPageTask const& page, - Nvfp4BoundaryBufferPlan const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) +__device__ OnboardBufferTask resolveTask(Nvfp4ColdPageOnboardPageTask const& page, + Nvfp4ColdPageBufferPlan const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.gpuPageIndex); auto const* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; @@ -158,14 +158,14 @@ __host__ __device__ constexpr std::uint32_t compactStageBytesForHalfGroups(std:: } // Flatten one buffer's [head, token, dim] geometry into eight-element half-groups. -__host__ __device__ constexpr std::uint32_t halfGroupCount(Nvfp4BoundaryKernelParams const& params) +__host__ __device__ constexpr std::uint32_t halfGroupCount(Nvfp4ColdPageKernelParams const& params) { return static_cast(params.numKvHeads) * static_cast(params.tokensPerPage) * (static_cast(params.headDim) / kElementsPerHalfGroup); } // Cap a buffer tile at shared-memory capacity without splitting a scale group. -__host__ __device__ constexpr std::uint32_t tileHalfGroupCount(Nvfp4BoundaryKernelParams const& params) +__host__ __device__ constexpr std::uint32_t tileHalfGroupCount(Nvfp4ColdPageKernelParams const& params) { // Avoid std::min: its reference return ODR-uses this host/device constexpr. auto const halfGroups = halfGroupCount(params); @@ -215,7 +215,7 @@ __device__ void flushCompactRangeToHost(std::uint8_t const* compactStages, Offlo } // Zero codec-specified record padding so persisted cold Slots are deterministic. -__device__ void clearColdPadding(OffloadBufferTask const& task, Nvfp4BoundaryBufferPlan const& buffer) +__device__ void clearColdPadding(OffloadBufferTask const& task, Nvfp4ColdPageBufferPlan const& buffer) { if (blockIdx.x != 0U) { @@ -361,7 +361,7 @@ __device__ uint2 quantizeFp8GrainToNvfp4(PackedVec<__nv_fp8_e4m3> const& grain, // Restore one natural 16-value NVFP4 scale group. template -__device__ float onboardDequantScale(std::uint8_t encodedScale, Nvfp4BoundaryKernelParams const& params) +__device__ float onboardDequantScale(std::uint8_t encodedScale, Nvfp4ColdPageKernelParams const& params) { __nv_fp8_e4m3 blockScale; blockScale.__x = encodedScale; @@ -410,8 +410,8 @@ __device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firs // FP16/BF16 GPU Page -> mapped-Host NVFP4 in bounded tiles. template __global__ void offloadFrom16BitTiledKernel( - std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, std::uint8_t* coldBase, + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, std::uint8_t* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) @@ -423,7 +423,7 @@ __global__ void offloadFrom16BitTiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) + if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) { copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); clearColdPadding(task, buffer); @@ -491,8 +491,8 @@ __global__ void offloadFrom16BitTiledKernel( // FP8 E4M3 GPU Page -> mapped-Host NVFP4 in bounded tiles. __global__ void offloadFromFp8TiledKernel( - std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, std::uint8_t* coldBase, + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, std::uint8_t* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) @@ -504,7 +504,7 @@ __global__ void offloadFromFp8TiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) + if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) { copyLosslessBytes(task.raw, task.coldData, buffer.rawBytes); clearColdPadding(task, buffer); @@ -635,8 +635,8 @@ __device__ void loadCompactRangeFromHost(std::uint8_t* compactStages, OnboardBuf // Mapped-Host NVFP4 -> runtime GPU Page in bounded tiles. template __global__ void onboardTiledKernel( - std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, + std::array const __grid_constant__ pages, + std::array const __grid_constant__ buffers, std::uint8_t const* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) @@ -648,7 +648,7 @@ __global__ void onboardTiledKernel( auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); - if (buffer.transform == Nvfp4BoundaryTransform::kLosslessCopy) + if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) { copyLosslessBytes(task.coldData, task.raw, buffer.rawBytes); return; @@ -706,7 +706,7 @@ __global__ void onboardTiledKernel( // Construction-time launch-plan validation. // Validate the geometry and scales used only by an NVFP4 transform. -void validateNvfp4Params(Nvfp4BoundaryKernelParams const& params, bool isFp8Runtime) +void validateNvfp4Params(Nvfp4ColdPageKernelParams const& params, bool isFp8Runtime) { TLLM_CHECK_WITH_INFO(params.numKvHeads > 0, "numKvHeads must be positive"); TLLM_CHECK_WITH_INFO(params.tokensPerPage > 0, "tokensPerPage must be positive"); @@ -754,7 +754,7 @@ void addColdInterval(std::vector& intervals, std::size_t offset, s } // Validate one raw/cold buffer mapping and append its occupied cold intervals. -void validateBufferPlan(Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldPageBytes, bool isFp8Runtime, +void validateBufferPlan(Nvfp4ColdPageBufferPlan const& buffer, std::size_t coldPageBytes, bool isFp8Runtime, std::vector& intervals) { TLLM_CHECK_WITH_INFO(buffer.rawBase != 0U, "rawBase must not be null"); @@ -763,7 +763,7 @@ void validateBufferPlan(Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldP switch (buffer.transform) { - case Nvfp4BoundaryTransform::kNvfp4: + case Nvfp4ColdPageTransform::kNvfp4: { validateNvfp4Params(buffer.params, isFp8Runtime); TLLM_CHECK_WITH_INFO( @@ -782,10 +782,10 @@ void validateBufferPlan(Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldP "NVFP4 scale interval"); break; } - case Nvfp4BoundaryTransform::kLosslessCopy: + case Nvfp4ColdPageTransform::kLosslessCopy: addColdInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes, "Lossless-data interval"); break; - default: TLLM_THROW("Unsupported NVFP4 boundary buffer transform"); + default: TLLM_THROW("Unsupported NVFP4 cold-page buffer transform"); } addColdInterval( intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes, "Cold-record padding interval"); @@ -795,7 +795,7 @@ void validateBufferPlan(Nvfp4BoundaryBufferPlan const& buffer, std::size_t coldP // Submit Page tasks through the fixed 256-descriptor kernel ABI. template -void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4BoundaryPreparedPlan const& plan, +void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPagePreparedPlan const& plan, ColdPointer coldBase, cudaStream_t stream) { static_assert(std::is_trivially_copyable_v>, @@ -832,7 +832,7 @@ void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4Bounda auto coldPageBytes = plan.coldPageBytes; auto numBuffers = plan.numBuffers; - void* arguments[] = {const_cast(chunkTasks), const_cast(plan.buffers.data()), + void* arguments[] = {const_cast(chunkTasks), const_cast(plan.buffers.data()), &coldBase, &coldPageBytes, &numBuffers}; TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); offset += numChunkTasks; @@ -841,7 +841,7 @@ void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4Bounda // Drain earlier chunks after a later synchronous launch failure. template -void submitBoundaryTasks(Kernel kernel, std::vector const& tasks, Nvfp4BoundaryPreparedPlan const& plan, +void submitColdPageTasks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPagePreparedPlan const& plan, ColdPointer coldBase, cudaStream_t stream) { try @@ -854,7 +854,7 @@ void submitBoundaryTasks(Kernel kernel, std::vector const& tasks, Nvfp4Bou if (drainStatus != cudaSuccess) { // An asynchronous drain failure leaves Slot ownership unknown; fail-stop. - TLLM_LOG_ERROR("NVFP4 boundary failure drain failed: %s", cudaGetErrorString(drainStatus)); + TLLM_LOG_ERROR("NVFP4 cold-page failure drain failed: %s", cudaGetErrorString(drainStatus)); std::terminate(); } throw; @@ -863,14 +863,14 @@ void submitBoundaryTasks(Kernel kernel, std::vector const& tasks, Nvfp4Bou } // namespace -Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType) +Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) { // TODO: Make this codec-private, then remove caller-guaranteed admission checks. - TLLM_CHECK_WITH_INFO(common::isSM100Family(), "NVFP4 boundary kernels require an SM100-family GPU"); - TLLM_CHECK_WITH_INFO(!buffers.empty(), "NVFP4 boundary launch requires at least one buffer"); + TLLM_CHECK_WITH_INFO(common::isSM100Family(), "NVFP4 cold-page kernels require an SM100-family GPU"); + TLLM_CHECK_WITH_INFO(!buffers.empty(), "NVFP4 cold-page launch requires at least one buffer"); TLLM_CHECK_WITH_INFO(buffers.size() <= kMaxBuffersPerLaunch, - "NVFP4 boundary launch supports at most %u local buffers, got %zu", kMaxBuffersPerLaunch, buffers.size()); + "NVFP4 cold-page launch supports at most %u local buffers, got %zu", kMaxBuffersPerLaunch, buffers.size()); TLLM_CHECK_WITH_INFO(coldPageBytes > 0, "Cold Page stride must be positive"); TLLM_CHECK_WITH_INFO( coldPageBytes % alignof(uint4) == 0, "Cold Page stride must be aligned to %zu bytes", alignof(uint4)); @@ -878,13 +878,13 @@ Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector(buffers.size()); plan.coldPageBytes = coldPageBytes; plan.runtimeType = runtimeType; @@ -894,7 +894,7 @@ Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageEncode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) { if (pages.empty()) { @@ -920,23 +920,23 @@ void invokeNvfp4BoundaryOffloadCompress(std::vector, pages, plan, static_cast(coldBase), stream); break; - case Nvfp4BoundaryRuntimeType::kBfloat16: - submitBoundaryTasks( + case Nvfp4ColdPageRuntimeType::kBfloat16: + submitColdPageTasks( offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); break; - case Nvfp4BoundaryRuntimeType::kFp8E4m3: - submitBoundaryTasks(offloadFromFp8TiledKernel, pages, plan, static_cast(coldBase), stream); + case Nvfp4ColdPageRuntimeType::kFp8E4m3: + submitColdPageTasks(offloadFromFp8TiledKernel, pages, plan, static_cast(coldBase), stream); break; - default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); + default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } } -void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageDecode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream) { if (pages.empty()) { @@ -945,18 +945,18 @@ void invokeNvfp4BoundaryOnboardDecompress(std::vector, pages, plan, static_cast(coldBase), stream); + case Nvfp4ColdPageRuntimeType::kFloat16: + submitColdPageTasks(onboardTiledKernel, pages, plan, static_cast(coldBase), stream); break; - case Nvfp4BoundaryRuntimeType::kBfloat16: - submitBoundaryTasks( + case Nvfp4ColdPageRuntimeType::kBfloat16: + submitColdPageTasks( onboardTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); break; - case Nvfp4BoundaryRuntimeType::kFp8E4m3: - submitBoundaryTasks( + case Nvfp4ColdPageRuntimeType::kFp8E4m3: + submitColdPageTasks( onboardTiledKernel<__nv_fp8_e4m3>, pages, plan, static_cast(coldBase), stream); break; - default: TLLM_THROW("Unsupported NVFP4 boundary runtime type"); + default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } } diff --git a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h similarity index 68% rename from cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h rename to cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h index 5496bff9af92..36af8c564489 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4BoundaryKernels.h +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h @@ -31,8 +31,8 @@ TRTLLM_NAMESPACE_BEGIN namespace kernels { -//! Active GPU representation at the cold-page boundary. -enum class Nvfp4BoundaryRuntimeType : std::uint8_t +//! Active GPU representation encoded into cold Pages. +enum class Nvfp4ColdPageRuntimeType : std::uint8_t { kFloat16, kBfloat16, @@ -40,14 +40,14 @@ enum class Nvfp4BoundaryRuntimeType : std::uint8_t }; //! One Base Page selected for GPU-to-Host transformation. -struct Nvfp4BoundaryOffloadPageTask +struct Nvfp4ColdPageOffloadPageTask { std::int32_t gpuPageIndex; std::int32_t coldPageIndex; }; //! One Base Page selected for Host-to-GPU transformation. -struct Nvfp4BoundaryOnboardPageTask +struct Nvfp4ColdPageOnboardPageTask { std::int32_t gpuPageIndex; std::int32_t coldPageIndex; @@ -55,7 +55,7 @@ struct Nvfp4BoundaryOnboardPageTask //! Per-buffer geometry and scales for one NVFP4 record in HND order. //! `headDim` is a multiple of 16; `*OrigQuant` encodes and `*QuantOrig` decodes this buffer. -struct Nvfp4BoundaryKernelParams +struct Nvfp4ColdPageKernelParams { std::int32_t numKvHeads; std::int32_t tokensPerPage; @@ -67,7 +67,7 @@ struct Nvfp4BoundaryKernelParams }; //! Transformation applied to one independently addressed hot buffer. -enum class Nvfp4BoundaryTransform : std::uint8_t +enum class Nvfp4ColdPageTransform : std::uint8_t { kNvfp4, //! Byte-exact copy for an Attention side buffer such as DSA index_key. @@ -75,7 +75,7 @@ enum class Nvfp4BoundaryTransform : std::uint8_t }; //! Immutable transform plan for one hot buffer and its fixed-offset cold record. -struct Nvfp4BoundaryBufferPlan +struct Nvfp4ColdPageBufferPlan { std::uintptr_t rawBase; std::size_t rawSlotBytes; @@ -84,33 +84,33 @@ struct Nvfp4BoundaryBufferPlan std::size_t coldScaleOffset; std::size_t coldPaddingOffset; std::uint32_t coldPaddingBytes; - Nvfp4BoundaryTransform transform; - Nvfp4BoundaryKernelParams params; + Nvfp4ColdPageTransform transform; + Nvfp4ColdPageKernelParams params; }; -inline constexpr std::uint32_t kNvfp4BoundaryMaxBuffersPerLaunch = 256; +inline constexpr std::uint32_t kNvfp4ColdPageMaxBuffersPerLaunch = 256; //! Configure-time launch plan for one Attention lifecycle. -struct Nvfp4BoundaryPreparedPlan +struct Nvfp4ColdPagePreparedPlan { - std::array buffers{}; + std::array buffers{}; std::uint32_t numBuffers = 0; std::uint32_t maxHalfGroupsPerTile = 0; std::size_t coldPageBytes = 0; - Nvfp4BoundaryRuntimeType runtimeType = Nvfp4BoundaryRuntimeType::kFloat16; + Nvfp4ColdPageRuntimeType runtimeType = Nvfp4ColdPageRuntimeType::kFloat16; }; -//! Validate and freeze one lifecycle's boundary-transform plan. -[[nodiscard]] Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType); +//! Validate and freeze one lifecycle's cold-page transform plan. +[[nodiscard]] Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType); //! Compress GPU Pages into mapped-Host NVFP4 records. -void invokeNvfp4BoundaryOffloadCompress(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream); +void invokeNvfp4ColdPageEncode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream); //! Restore mapped-Host NVFP4 records into GPU Pages. -void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream); +void invokeNvfp4ColdPageDecode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream); } // namespace kernels diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt index e72249e383f0..c2cf3d8d28eb 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -1,7 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. # All rights reserved. SPDX-License-Identifier: Apache-2.0 -add_library(kv_cache_compression_src OBJECT nvfp4ColdPageCodec.cpp) +add_library(kv_cache_compression_src OBJECT coldPageCodec.cpp) target_include_directories( kv_cache_compression_src PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp new file mode 100644 index 000000000000..810107ecb1a3 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp @@ -0,0 +1,416 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" + +#include "tensorrt_llm/common/logger.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) +{ + if (lhs > std::numeric_limits::max() - rhs) + { + throw std::overflow_error(label); + } + return lhs + rhs; +} + +std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) +{ + if (offset > std::numeric_limits::max() - base) + { + throw std::overflow_error("GPU buffer address overflows uintptr_t"); + } + return base + offset; +} + +struct BufferLocation +{ + kv::PoolIndex poolIndex{0}; + std::size_t offset = 0; + std::size_t bytes = 0; +}; + +using LayerBuffers = std::map; +using LifecycleBuffers = std::map; +using LayerPlans = std::map; + +LifecycleBuffers indexLifecycleBuffers(kv::SlotDescVariant const& variant) +{ + LifecycleBuffers result; + for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) + { + auto const& coalesced = variant.coalescedBuffers.at(poolIndex); + std::size_t offset = 0; + for (auto const& bufferId : coalesced.bufferIds) + { + auto& buffers = result[bufferId.layerId]; + if (!buffers.emplace(bufferId.role, BufferLocation{poolIndex, offset, coalesced.singleBufferSize}).second) + { + throw std::invalid_argument("GPU lifecycle contains a duplicate buffer role"); + } + offset = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); + } + } + return result; +} + +kernels::Nvfp4ColdPageKernelParams toKernelParams(Nvfp4ColdPageParams const& params) +{ + return kernels::Nvfp4ColdPageKernelParams{params.numKvHeads, params.tokensPerPage, params.headDim, + params.nvfp4ScaleOrigQuant, params.nvfp4ScaleQuantOrig, params.fp8ScaleOrigQuant, params.fp8ScaleQuantOrig}; +} + +struct CompiledLayerGroup +{ + kernels::Nvfp4ColdPagePreparedPlan preparedPlan; + std::size_t coldPageBytes = 0; + std::set consumedLayers; +}; + +CompiledLayerGroup compileLayerGroup( + kv::PoolGroupDesc const& gpuDesc, LifecycleBuffers const& physicalLayers, LayerPlans const& layerPlans) +{ + CompiledLayerGroup result; + std::vector groupBuffers; + std::optional groupRuntimeType; + + for (auto const& [layerId, physicalBuffers] : physicalLayers) + { + auto const planIt = layerPlans.find(layerId); + if (planIt == layerPlans.end()) + { + throw std::invalid_argument("A planned lifecycle contains a layer without a cold-page plan"); + } + auto const& layerPlan = planIt->second; + std::vector layerBuffers; + layerBuffers.reserve(layerPlan.buffers.size()); + std::optional layerRuntimeType; + + for (auto const& bufferPlan : layerPlan.buffers) + { + auto const locationIt = physicalBuffers.find(bufferPlan.role); + if (locationIt == physicalBuffers.end()) + { + throw std::invalid_argument( + "Cold-page plan references a buffer " + "absent from the GPU lifecycle"); + } + auto const& location = locationIt->second; + if (bufferPlan.rawBytes != location.bytes) + { + throw std::invalid_argument("Cold-page plan raw size does not match the GPU buffer size"); + } + + auto const& pool = gpuDesc.pools.at(location.poolIndex); + kernels::Nvfp4ColdPageBufferPlan nativePlan{checkedAddress(pool.baseAddress, location.offset), + pool.slotBytes, bufferPlan.rawBytes, bufferPlan.coldDataOffset, bufferPlan.coldScaleOffset, 0U, 0U, + kernels::Nvfp4ColdPageTransform::kLosslessCopy, {}}; + switch (bufferPlan.transform) + { + case ColdPageTransformKind::kLosslessCopy: break; + case ColdPageTransformKind::kNvfp4: + if (!bufferPlan.nvfp4Params) + { + throw std::invalid_argument("NVFP4 cold-page transform requires NVFP4 parameters"); + } + nativePlan.transform = kernels::Nvfp4ColdPageTransform::kNvfp4; + nativePlan.params = toKernelParams(*bufferPlan.nvfp4Params); + if (layerRuntimeType && *layerRuntimeType != bufferPlan.nvfp4Params->runtimeType) + { + throw std::invalid_argument("One layer cold-page plan must use one runtime type"); + } + layerRuntimeType = bufferPlan.nvfp4Params->runtimeType; + break; + default: throw std::invalid_argument("Unsupported cold-page transform kind"); + } + layerBuffers.push_back(nativePlan); + } + if (layerBuffers.size() != physicalBuffers.size()) + { + throw std::invalid_argument("Cold-page layer plan must cover every GPU buffer role"); + } + + layerBuffers.back().coldPaddingOffset = layerPlan.coldPaddingOffset; + layerBuffers.back().coldPaddingBytes = static_cast(layerPlan.coldPaddingBytes); + auto const runtimeType = layerRuntimeType.value_or(kernels::Nvfp4ColdPageRuntimeType::kFloat16); + auto const validatedLayer + = kernels::prepareNvfp4ColdPagePlan(layerBuffers, layerPlan.coldPageBytes, runtimeType); + + if (layerRuntimeType) + { + if (groupRuntimeType && *groupRuntimeType != *layerRuntimeType) + { + throw std::invalid_argument("One lifecycle cold-page plan must use one runtime type"); + } + groupRuntimeType = *layerRuntimeType; + } + for (std::uint32_t index = 0; index < validatedLayer.numBuffers; ++index) + { + auto nativePlan = validatedLayer.buffers[index]; + nativePlan.coldDataOffset + = checkedAdd(result.coldPageBytes, nativePlan.coldDataOffset, "Cold-page data offset overflows size_t"); + if (nativePlan.transform == kernels::Nvfp4ColdPageTransform::kNvfp4) + { + nativePlan.coldScaleOffset = checkedAdd( + result.coldPageBytes, nativePlan.coldScaleOffset, "Cold-page scale offset overflows size_t"); + } + if (nativePlan.coldPaddingBytes != 0U) + { + nativePlan.coldPaddingOffset = checkedAdd( + result.coldPageBytes, nativePlan.coldPaddingOffset, "Cold-page padding offset overflows size_t"); + } + groupBuffers.push_back(nativePlan); + } + result.coldPageBytes + = checkedAdd(result.coldPageBytes, layerPlan.coldPageBytes, "Lifecycle cold-page size overflows size_t"); + result.consumedLayers.insert(layerId); + } + + result.preparedPlan = kernels::prepareNvfp4ColdPagePlan( + groupBuffers, result.coldPageBytes, groupRuntimeType.value_or(kernels::Nvfp4ColdPageRuntimeType::kFloat16)); + return result; +} + +} // namespace + +PlannedColdPageCodec::PlannedColdPageCodec(std::vector layerPlans) +{ + for (auto& layerPlan : layerPlans) + { + if (layerPlan.coldPageBytes == 0U || layerPlan.buffers.empty()) + { + throw std::invalid_argument("A cold-page layer plan must contain a non-empty record"); + } + if (layerPlan.coldPaddingOffset > layerPlan.coldPageBytes + || layerPlan.coldPaddingBytes != layerPlan.coldPageBytes - layerPlan.coldPaddingOffset + || layerPlan.coldPaddingBytes > std::numeric_limits::max()) + { + throw std::invalid_argument("Cold-page layer padding exceeds its record"); + } + + std::set roles; + for (auto const& buffer : layerPlan.buffers) + { + if (buffer.role.empty() || buffer.rawBytes == 0U || !roles.emplace(buffer.role).second) + { + throw std::invalid_argument( + "Cold-page layer plans require unique " + "non-empty buffer roles and sizes"); + } + if (buffer.transform != ColdPageTransformKind::kLosslessCopy + && buffer.transform != ColdPageTransformKind::kNvfp4) + { + throw std::invalid_argument("Unsupported cold-page transform kind"); + } + } + auto const layerId = layerPlan.layerId; + if (!mLayerPlans.emplace(layerId, std::move(layerPlan)).second) + { + throw std::invalid_argument("Cold-page layer plan IDs must be unique"); + } + } +} + +bool PlannedColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept +{ + try + { + if (mLosslessCodec || !mLayerGroups.empty()) + { + throw std::invalid_argument("PlannedColdPageCodec can be configured only once"); + } + auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); + if (!losslessCodec->configure(gpuDescs, numGpuDescs)) + { + throw std::invalid_argument("Default lossless codec rejected GPU layouts"); + } + + std::map pendingGroups; + std::set consumedLayers; + for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) + { + auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; + for (auto const& variant : gpuDesc.slotDesc.variants) + { + auto const physicalLayers = indexLifecycleBuffers(variant); + auto const plannedLayerCount = std::count_if(physicalLayers.begin(), physicalLayers.end(), + [this](auto const& layer) { return mLayerPlans.count(layer.first) != 0U; }); + + LayerGroupState state; + if (plannedLayerCount == 0U) + { + state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); + } + else + { + if (plannedLayerCount != physicalLayers.size()) + { + throw std::invalid_argument("A lifecycle cannot mix planned and unplanned layers"); + } + auto compiled = compileLayerGroup(gpuDesc, physicalLayers, mLayerPlans); + state.execution = ExecutionKind::kPlanned; + state.preparedPlan = std::move(compiled.preparedPlan); + state.coldPageBytes = compiled.coldPageBytes; + for (auto const layerId : compiled.consumedLayers) + { + if (!consumedLayers.emplace(layerId).second) + { + throw std::invalid_argument("A cold-page layer plan appears in multiple lifecycles"); + } + } + } + if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) + { + throw std::invalid_argument("GPU lifecycle ID appears in multiple pool groups"); + } + } + } + if (consumedLayers.size() != mLayerPlans.size()) + { + throw std::invalid_argument("A cold-page layer plan is absent from all GPU descriptors"); + } + + mLayerGroups = std::move(pendingGroups); + mLosslessCodec = std::move(losslessCodec); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("PlannedColdPageCodec::configure rejected GPU layouts: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR( + "PlannedColdPageCodec::configure rejected GPU layouts: " + "unknown error"); + return false; + } +} + +PlannedColdPageCodec::LayerGroupState const* PlannedColdPageCodec::findLayerGroup( + kv::LayerGroupId layerGroupId) const noexcept +{ + auto const found = mLayerGroups.find(layerGroupId); + return found == mLayerGroups.end() ? nullptr : &found->second; +} + +std::size_t PlannedColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept +{ + auto const* state = findLayerGroup(layerGroupId); + return state == nullptr ? 0U : state->coldPageBytes; +} + +kv::LayerGroupId PlannedColdPageCodec::getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept +{ + return findLayerGroup(layerGroupId) == nullptr ? kv::LayerGroupId{-1} : layerGroupId; +} + +kv::PageIndexLocation PlannedColdPageCodec::queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept +{ + return findLayerGroup(layerGroupId) == nullptr ? kv::PageIndexLocation::kBadLocation : kv::PageIndexLocation::kHost; +} + +bool PlannedColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && (dstBasePtr == nullptr || pageIndices == nullptr))) + { + throw std::invalid_argument("encode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (state->execution == ExecutionKind::kLossless) + { + return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); + } + + thread_local std::vector pages; + pages.clear(); + pages.reserve(numBasePages); + for (std::size_t page = 0; page < numBasePages; ++page) + { + pages.push_back({pageIndices[page].src, pageIndices[page].dst}); + } + kernels::invokeNvfp4ColdPageEncode(pages, state->preparedPlan, dstBasePtr, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("PlannedColdPageCodec::encode failed before completion fencing: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR( + "PlannedColdPageCodec::encode failed before completion " + "fencing: unknown error"); + return false; + } +} + +bool PlannedColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, + kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && (srcBasePtr == nullptr || pageIndices == nullptr))) + { + throw std::invalid_argument("decode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (state->execution == ExecutionKind::kLossless) + { + return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); + } + + thread_local std::vector pages; + pages.clear(); + pages.reserve(numBasePages); + for (std::size_t page = 0; page < numBasePages; ++page) + { + pages.push_back({pageIndices[page].dst, pageIndices[page].src}); + } + kernels::invokeNvfp4ColdPageDecode(pages, state->preparedPlan, srcBasePtr, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("PlannedColdPageCodec::decode failed before completion fencing: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR( + "PlannedColdPageCodec::decode failed before completion " + "fencing: unknown error"); + return false; + } +} + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h new file mode 100644 index 000000000000..bbb725cba1e0 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h @@ -0,0 +1,108 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include "kv_cache_manager_v2/coldPageCodec.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +//! Native transform selected for one buffer in a planned cold-page record. +enum class ColdPageTransformKind : std::uint8_t +{ + kLosslessCopy, + kNvfp4, +}; + +//! NVFP4 parameters for one independently scaled buffer. +struct Nvfp4ColdPageParams +{ + kernels::Nvfp4ColdPageRuntimeType runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; + std::int32_t numKvHeads = 0; + std::int32_t tokensPerPage = 0; + std::int32_t headDim = 0; + float nvfp4ScaleOrigQuant = 1.0F; + float nvfp4ScaleQuantOrig = 1.0F; + float fp8ScaleOrigQuant = 1.0F; + float fp8ScaleQuantOrig = 1.0F; +}; + +//! Python-authored transform and layer-relative cold offsets for one buffer. +struct ColdPageBufferPlan +{ + kv::DataRole role; + ColdPageTransformKind transform = ColdPageTransformKind::kLosslessCopy; + std::size_t rawBytes = 0; + std::size_t coldDataOffset = 0; + std::size_t coldScaleOffset = 0; + std::optional nvfp4Params; +}; + +//! Python-authored fixed cold record for one layer. +struct ColdPageLayerPlan +{ + kv::LayerId layerId = 0; + std::size_t coldPageBytes = 0; + std::size_t coldPaddingOffset = 0; + std::size_t coldPaddingBytes = 0; + std::vector buffers; +}; + +//! Resolves declarative layer plans against KVCM's authoritative hot-pool +// descriptors. +class PlannedColdPageCodec final : public kv::IKvCacheColdPageCodec +{ +public: + explicit PlannedColdPageCodec(std::vector layerPlans); + + bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; + + [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; + + bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + + bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + +private: + enum class ExecutionKind : std::uint8_t + { + kLossless, + kPlanned, + }; + + struct LayerGroupState + { + ExecutionKind execution = ExecutionKind::kLossless; + kernels::Nvfp4ColdPagePreparedPlan preparedPlan; + std::size_t coldPageBytes = 0; + }; + + [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; + + std::map mLayerPlans; + std::map mLayerGroups; + std::unique_ptr mLosslessCodec; +}; + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp deleted file mode 100644 index 9ad9f88f28a1..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp +++ /dev/null @@ -1,427 +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 "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" - -#include "tensorrt_llm/common/logger.h" - -#include -#include -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ -namespace -{ - -constexpr std::size_t kColdRecordAlignment = 16U; -constexpr std::size_t kElementsPerScaleGroup = 16U; -constexpr std::size_t kPackedElementsPerByte = 2U; -constexpr char const* kKeyRole = "key"; -constexpr char const* kValueRole = "value"; - -// Descriptor-derived byte arithmetic must not wrap into an undersized cold record. -std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) -{ - if (lhs > std::numeric_limits::max() - rhs) - { - throw std::overflow_error(label); - } - return lhs + rhs; -} - -std::size_t checkedMul(std::size_t lhs, std::size_t rhs, char const* label) -{ - if (rhs != 0U && lhs > std::numeric_limits::max() / rhs) - { - throw std::overflow_error(label); - } - return lhs * rhs; -} - -std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) -{ - if (offset > std::numeric_limits::max() - base) - { - throw std::overflow_error("GPU buffer address overflows uintptr_t"); - } - return base + offset; -} - -std::size_t alignColdRecord(std::size_t bytes) -{ - return checkedAdd(bytes, kColdRecordAlignment - 1U, "Cold Page size overflows size_t") / kColdRecordAlignment - * kColdRecordAlignment; -} - -std::size_t checkedElementCount(Nvfp4ColdPageLayerConfig const& config) -{ - if (config.numKvHeads <= 0 || config.tokensPerPage <= 0 || config.headDim <= 0) - { - throw std::invalid_argument("NVFP4 cold Page geometry must be positive"); - } - if (config.headDim % static_cast(kElementsPerScaleGroup) != 0) - { - throw std::invalid_argument("NVFP4 cold Pages require headDim divisible by 16"); - } - auto const headsTimesTokens = checkedMul(static_cast(config.numKvHeads), - static_cast(config.tokensPerPage), "NVFP4 Page geometry overflows size_t"); - return checkedMul( - headsTimesTokens, static_cast(config.headDim), "NVFP4 Page geometry overflows size_t"); -} - -kernels::Nvfp4BoundaryKernelParams makeKernelParams(Nvfp4ColdPageLayerConfig const& config, std::size_t scaleIndex) -{ - kernels::Nvfp4BoundaryKernelParams params{}; - params.numKvHeads = config.numKvHeads; - params.tokensPerPage = config.tokensPerPage; - params.headDim = config.headDim; - params.nvfp4ScaleOrigQuant = config.nvfp4ScaleOrigQuant[scaleIndex]; - params.nvfp4ScaleQuantOrig = config.nvfp4ScaleQuantOrig[scaleIndex]; - params.fp8ScaleOrigQuant = config.fp8ScaleOrigQuant[scaleIndex]; - params.fp8ScaleQuantOrig = config.fp8ScaleQuantOrig[scaleIndex]; - return params; -} - -struct BufferLocation -{ - kv::PoolIndex poolIndex{0}; - std::size_t offset = 0; - std::size_t bytes = 0; - bool found = false; -}; - -struct AttentionLayerBuffers -{ - BufferLocation key; - BufferLocation value; - // Roles such as MLA index_key remain byte-exact. - std::vector sideBuffers; -}; - -using LayerConfigs = std::map; -using AttentionBufferMap = std::map; - -struct LifecycleBuffers -{ - AttentionBufferMap attention; - bool hasNonAttentionLayer = false; -}; - -LifecycleBuffers discoverLifecycleBuffers(kv::SlotDescVariant const& variant, LayerConfigs const& layerConfigs) -{ - LifecycleBuffers result; - for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) - { - auto const& coalesced = variant.coalescedBuffers.at(poolIndex); - std::size_t offset = 0U; - for (auto const& bufferId : coalesced.bufferIds) - { - auto const config = layerConfigs.find(bufferId.layerId); - if (config == layerConfigs.end()) - { - result.hasNonAttentionLayer = true; - } - else - { - auto& buffers = result.attention[bufferId.layerId]; - if (bufferId.role == kKeyRole || bufferId.role == kValueRole) - { - auto& location = bufferId.role == kKeyRole ? buffers.key : buffers.value; - if (location.found) - { - throw std::invalid_argument("GPU lifecycle contains a duplicate K/V buffer"); - } - location = BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}; - } - else - { - buffers.sideBuffers.push_back(BufferLocation{poolIndex, offset, coalesced.singleBufferSize, true}); - } - } - offset = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); - } - } - return result; -} - -struct AttentionPlan -{ - kernels::Nvfp4BoundaryPreparedPlan kernelPlan; - std::size_t coldPageBytes = 0; -}; - -AttentionPlan buildAttentionPlan( - kv::PoolGroupDesc const& gpuDesc, AttentionBufferMap const& layerBuffers, LayerConfigs const& layerConfigs) -{ - AttentionPlan result; - std::vector bufferPlans; - auto const runtimeType = layerConfigs.at(layerBuffers.begin()->first).runtimeType; - - for (auto const& [layerId, buffers] : layerBuffers) - { - // MHA/GQA has K and V; MLA exposes its latent KV as key only. - if (!buffers.key.found) - { - throw std::invalid_argument("Configured Attention layer is missing key"); - } - auto const& config = layerConfigs.at(layerId); - if (runtimeType != config.runtimeType) - { - throw std::invalid_argument("Attention lifecycle must use one runtime dtype"); - } - - auto const elements = checkedElementCount(config); - auto const layerOffset = result.coldPageBytes; - auto const packedBytes = elements / kPackedElementsPerByte; - auto const scaleBytes = elements / kElementsPerScaleGroup; - auto const compressedBufferCount = buffers.value.found ? 2U : 1U; - auto const packedRegionBytes = packedBytes * compressedBufferCount; - auto const scaleRegionBytes = scaleBytes * compressedBufferCount; - - // Per layer: [K packed | V packed? | K scale | V scale? | side buffers | padding]. - auto const scaleOffset = checkedAdd(layerOffset, packedRegionBytes, "NVFP4 cold Page size overflows size_t"); - auto coldOffset = checkedAdd(scaleOffset, scaleRegionBytes, "NVFP4 cold Page size overflows size_t"); - - auto appendNvfp4 = [&](BufferLocation const& location, std::size_t scaleIndex, std::size_t coldDataOffset, - std::size_t coldScaleOffset) - { - auto const& pool = gpuDesc.pools.at(location.poolIndex); - bufferPlans.push_back(kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), - pool.slotBytes, location.bytes, coldDataOffset, coldScaleOffset, 0U, 0U, - kernels::Nvfp4BoundaryTransform::kNvfp4, makeKernelParams(config, scaleIndex)}); - }; - - appendNvfp4(buffers.key, 0U, layerOffset, scaleOffset); - if (buffers.value.found) - { - appendNvfp4(buffers.value, 1U, layerOffset + packedBytes, scaleOffset + scaleBytes); - } - for (auto const& location : buffers.sideBuffers) - { - auto const& pool = gpuDesc.pools.at(location.poolIndex); - bufferPlans.push_back( - kernels::Nvfp4BoundaryBufferPlan{checkedAddress(pool.baseAddress, location.offset), pool.slotBytes, - location.bytes, coldOffset, 0U, 0U, 0U, kernels::Nvfp4BoundaryTransform::kLosslessCopy, {}}); - coldOffset = checkedAdd(coldOffset, location.bytes, "Lossless Attention side-buffer size overflows size_t"); - } - - auto const alignedEnd = alignColdRecord(coldOffset); - bufferPlans.back().coldPaddingOffset = coldOffset; - bufferPlans.back().coldPaddingBytes = static_cast(alignedEnd - coldOffset); - result.coldPageBytes = alignedEnd; - } - - result.kernelPlan = kernels::prepareNvfp4BoundaryPlan(bufferPlans, result.coldPageBytes, runtimeType); - return result; -} - -} // namespace - -// Validate and retain algorithm-owned metadata before KVCM creates its physical pools. -Nvfp4ColdPageCodec::Nvfp4ColdPageCodec(std::vector layerConfigs) -{ - auto const validScales = [](auto const& scales) { - return std::all_of(scales.begin(), scales.end(), [](float value) { return std::isfinite(value) && value > 0; }); - }; - for (auto& config : layerConfigs) - { - if (config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFloat16 - && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kBfloat16 - && config.runtimeType != kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3) - { - throw std::invalid_argument("Nvfp4ColdPageCodec received an unsupported runtime type"); - } - static_cast(checkedElementCount(config)); - if (!validScales(config.nvfp4ScaleOrigQuant) || !validScales(config.nvfp4ScaleQuantOrig) - || (config.runtimeType == kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3 - && (!validScales(config.fp8ScaleOrigQuant) || !validScales(config.fp8ScaleQuantOrig)))) - { - throw std::invalid_argument("NVFP4 runtime scales must be finite and positive"); - } - auto const layerId = config.layerId; - if (!mLayerConfigs.emplace(layerId, std::move(config)).second) - { - throw std::invalid_argument("Nvfp4ColdPageCodec layer IDs must be unique"); - } - } -} - -// Bind the immutable metadata to KVCM's authoritative hot-pool layout exactly once. -bool Nvfp4ColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept -{ - try - { - // Use KVCM's default codec for non-Attention lifecycles. - auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); - if (!losslessCodec->configure(gpuDescs, numGpuDescs)) - { - throw std::invalid_argument("Default lossless codec rejected GPU layouts"); - } - - std::map pendingGroups; - std::set discoveredAttentionLayers; - for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) - { - auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; - for (auto const& variant : gpuDesc.slotDesc.variants) - { - auto const buffers = discoverLifecycleBuffers(variant, mLayerConfigs); - - LayerGroupState state; - if (!buffers.attention.empty()) - { - if (buffers.hasNonAttentionLayer) - { - throw std::invalid_argument("A lifecycle cannot mix Attention and non-Attention layers"); - } - - auto plan = buildAttentionPlan(gpuDesc, buffers.attention, mLayerConfigs); - state.format = ColdPageFormat::kNvfp4Kv; - state.preparedPlan = std::move(plan.kernelPlan); - state.coldPageBytes = plan.coldPageBytes; - for (auto const& layer : buffers.attention) - { - discoveredAttentionLayers.insert(layer.first); - } - } - else - { - state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); - } - if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) - { - throw std::invalid_argument("GPU lifecycle ID appears in multiple pool groups"); - } - } - } - if (discoveredAttentionLayers.size() != mLayerConfigs.size()) - { - throw std::invalid_argument("A configured Attention layer is absent from all GPU descriptors"); - } - - mLayerGroups = std::move(pendingGroups); - mLosslessCodec = std::move(losslessCodec); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("Nvfp4ColdPageCodec::configure rejected GPU layouts: %s", error.what()); - return false; - } -} - -Nvfp4ColdPageCodec::LayerGroupState const* Nvfp4ColdPageCodec::findLayerGroup( - kv::LayerGroupId layerGroupId) const noexcept -{ - auto const found = mLayerGroups.find(layerGroupId); - return found == mLayerGroups.end() ? nullptr : &found->second; -} - -std::size_t Nvfp4ColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept -{ - auto const* state = findLayerGroup(layerGroupId); - return state == nullptr ? 0U : state->coldPageBytes; -} - -kv::LayerGroupId Nvfp4ColdPageCodec::getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept -{ - return findLayerGroup(layerGroupId) == nullptr ? kv::LayerGroupId{-1} : layerGroupId; -} - -kv::PageIndexLocation Nvfp4ColdPageCodec::queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept -{ - return findLayerGroup(layerGroupId) == nullptr ? kv::PageIndexLocation::kBadLocation : kv::PageIndexLocation::kHost; -} - -bool Nvfp4ColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept -{ - try - { - auto const* state = findLayerGroup(layerGroupId); - if (state == nullptr || (numBasePages != 0U && pageIndices == nullptr)) - { - throw std::invalid_argument("encode received an invalid lifecycle or Page batch"); - } - if (numBasePages == 0U) - { - return true; - } - if (state->format == ColdPageFormat::kLossless) - { - return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); - } - - thread_local std::vector pages; - pages.clear(); - pages.reserve(numBasePages); - for (std::size_t page = 0; page < numBasePages; ++page) - { - pages.push_back({pageIndices[page].src, pageIndices[page].dst}); - } - kernels::invokeNvfp4BoundaryOffloadCompress(pages, state->preparedPlan, dstBasePtr, stream); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("Nvfp4ColdPageCodec::encode failed before completion fencing: %s", error.what()); - return false; - } -} - -bool Nvfp4ColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, - kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept -{ - try - { - auto const* state = findLayerGroup(layerGroupId); - if (state == nullptr || (numBasePages != 0U && pageIndices == nullptr)) - { - throw std::invalid_argument("decode received an invalid lifecycle or Page batch"); - } - if (numBasePages == 0U) - { - return true; - } - if (state->format == ColdPageFormat::kLossless) - { - return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); - } - - thread_local std::vector pages; - pages.clear(); - pages.reserve(numBasePages); - for (std::size_t page = 0; page < numBasePages; ++page) - { - pages.push_back({pageIndices[page].dst, pageIndices[page].src}); - } - kernels::invokeNvfp4BoundaryOnboardDecompress(pages, state->preparedPlan, srcBasePtr, stream); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("Nvfp4ColdPageCodec::decode failed before completion fencing: %s", error.what()); - return false; - } -} - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h deleted file mode 100644 index 93c318629868..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h +++ /dev/null @@ -1,90 +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. - */ - -#pragma once - -#include "kv_cache_manager_v2/coldPageCodec.h" -#include "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" - -#include -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ - -namespace kv = batch_manager::kv_cache_manager_v2; - -//! Per-layer geometry and calibration for NVFP4 cold Pages. -struct Nvfp4ColdPageLayerConfig -{ - kv::LayerId layerId = 0; - kernels::Nvfp4BoundaryRuntimeType runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFloat16; - std::int32_t numKvHeads = 0; - std::int32_t tokensPerPage = 0; - std::int32_t headDim = 0; - std::array nvfp4ScaleOrigQuant{}; - std::array nvfp4ScaleQuantOrig{}; - std::array fp8ScaleOrigQuant{1.0F, 1.0F}; - std::array fp8ScaleQuantOrig{1.0F, 1.0F}; -}; - -//! NVFP4 codec for compact Attention records with lossless side-buffer spans. -class Nvfp4ColdPageCodec final : public kv::IKvCacheColdPageCodec -{ -public: - explicit Nvfp4ColdPageCodec(std::vector layerConfigs); - - bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; - - [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; - - bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - - bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - -private: - enum class ColdPageFormat - { - kNvfp4Kv, //!< NVFP4 K/V plus byte-exact Attention side buffers. - kLossless, //!< Entire lifecycle uses KVCM's default lossless concat. - }; - - struct LayerGroupState - { - ColdPageFormat format = ColdPageFormat::kLossless; - kernels::Nvfp4BoundaryPreparedPlan preparedPlan; - std::size_t coldPageBytes = 0; - }; - - [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; - - std::map mLayerConfigs; - std::map mLayerGroups; - std::unique_ptr mLosslessCodec; -}; - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 34eecd1fd38e..f60778ff175d 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -1,5 +1,6 @@ /* - * 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"); @@ -16,10 +17,11 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" +#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" #include -#include +#include +#include #include #include @@ -37,30 +39,49 @@ namespace tensorrt_llm::nanobind::kv_cache_compression void initBindings(nb::module_& module) { - nb::enum_(module, "Nvfp4BoundaryRuntimeType") - .value("FLOAT16", kernels::Nvfp4BoundaryRuntimeType::kFloat16) - .value("BFLOAT16", kernels::Nvfp4BoundaryRuntimeType::kBfloat16) - .value("FP8_E4M3", kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3); + nb::enum_(module, "ColdPageTransformKind") + .value("LOSSLESS_COPY", compression::ColdPageTransformKind::kLosslessCopy) + .value("NVFP4", compression::ColdPageTransformKind::kNvfp4); - nb::class_(module, "Nvfp4ColdPageLayerConfig") + nb::enum_(module, "Nvfp4ColdPageRuntimeType") + .value("FLOAT16", kernels::Nvfp4ColdPageRuntimeType::kFloat16) + .value("BFLOAT16", kernels::Nvfp4ColdPageRuntimeType::kBfloat16) + .value("FP8_E4M3", kernels::Nvfp4ColdPageRuntimeType::kFp8E4m3); + + nb::class_(module, "Nvfp4ColdPageParams") + .def(nb::init<>()) + .def_rw("runtime_type", &compression::Nvfp4ColdPageParams::runtimeType) + .def_rw("num_kv_heads", &compression::Nvfp4ColdPageParams::numKvHeads) + .def_rw("tokens_per_page", &compression::Nvfp4ColdPageParams::tokensPerPage) + .def_rw("head_dim", &compression::Nvfp4ColdPageParams::headDim) + .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageParams::nvfp4ScaleOrigQuant) + .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageParams::nvfp4ScaleQuantOrig) + .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageParams::fp8ScaleOrigQuant) + .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageParams::fp8ScaleQuantOrig); + + nb::class_(module, "ColdPageBufferPlan") + .def(nb::init<>()) + .def_rw("role", &compression::ColdPageBufferPlan::role) + .def_rw("transform", &compression::ColdPageBufferPlan::transform) + .def_rw("raw_bytes", &compression::ColdPageBufferPlan::rawBytes) + .def_rw("cold_data_offset", &compression::ColdPageBufferPlan::coldDataOffset) + .def_rw("cold_scale_offset", &compression::ColdPageBufferPlan::coldScaleOffset) + .def_rw("nvfp4_params", &compression::ColdPageBufferPlan::nvfp4Params); + + nb::class_(module, "ColdPageLayerPlan") .def(nb::init<>()) - .def_rw("layer_id", &compression::Nvfp4ColdPageLayerConfig::layerId) - .def_rw("runtime_type", &compression::Nvfp4ColdPageLayerConfig::runtimeType) - .def_rw("num_kv_heads", &compression::Nvfp4ColdPageLayerConfig::numKvHeads) - .def_rw("tokens_per_page", &compression::Nvfp4ColdPageLayerConfig::tokensPerPage) - .def_rw("head_dim", &compression::Nvfp4ColdPageLayerConfig::headDim) - .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageLayerConfig::nvfp4ScaleOrigQuant) - .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageLayerConfig::nvfp4ScaleQuantOrig) - .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageLayerConfig::fp8ScaleOrigQuant) - .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageLayerConfig::fp8ScaleQuantOrig); + .def_rw("layer_id", &compression::ColdPageLayerPlan::layerId) + .def_rw("cold_page_bytes", &compression::ColdPageLayerPlan::coldPageBytes) + .def_rw("cold_padding_offset", &compression::ColdPageLayerPlan::coldPaddingOffset) + .def_rw("cold_padding_bytes", &compression::ColdPageLayerPlan::coldPaddingBytes) + .def_rw("buffers", &compression::ColdPageLayerPlan::buffers); // Construct in C++ so ownership can transfer to KVCM as a unique_ptr codec. module.def( - "create_nvfp4_cold_page_codec", - [](std::vector layerConfigs) - -> std::unique_ptr - { return std::make_unique(std::move(layerConfigs)); }, - nb::arg("layer_configs"), "Create an owning NVFP4 cold-page codec for one-time transfer into KVCacheManager."); + "create_cold_page_codec", + [](std::vector layerPlans) -> std::unique_ptr + { return std::make_unique(std::move(layerPlans)); }, + nb::arg("layer_plans"), "Create an owning planned cold-page codec for transfer into KVCacheManager."); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tests/unit_tests/kernels/CMakeLists.txt b/cpp/tests/unit_tests/kernels/CMakeLists.txt index 643c98dae575..ef7d99072940 100644 --- a/cpp/tests/unit_tests/kernels/CMakeLists.txt +++ b/cpp/tests/unit_tests/kernels/CMakeLists.txt @@ -46,21 +46,21 @@ if(USING_OSS_CUTLASS_MOE_GEMM) endif() add_gtest(ropeTest ropeTest.cu) -set(NVFP4_BOUNDARY_KERNEL_TEST_SRC - nvfp4BoundaryKernelsTest.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4BoundaryKernels.cu +set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC + nvfp4ColdPageKernelsTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) -add_gtest(nvfp4BoundaryKernelsTest "${NVFP4_BOUNDARY_KERNEL_TEST_SRC}" +add_gtest(nvfp4ColdPageKernelsTest "${NVFP4_COLD_PAGE_KERNEL_TEST_SRC}" NO_TLLM_LINKAGE) target_include_directories( - nvfp4BoundaryKernelsTest + nvfp4ColdPageKernelsTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) -target_link_libraries(nvfp4BoundaryKernelsTest PRIVATE CUDA::cudart +target_link_libraries(nvfp4ColdPageKernelsTest PRIVATE CUDA::cudart CUDA::cuda_driver) add_gtest(shiftKCacheKernelTest shiftKCacheKernelTest.cu) add_gtest(smoothQuantKernelTest smoothQuant/smoothQuantKernelTest.cpp) diff --git a/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp similarity index 86% rename from cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp rename to cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp index bd522cd80cff..be09cca9c9b3 100644 --- a/cpp/tests/unit_tests/kernels/nvfp4BoundaryKernelsTest.cpp +++ b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp @@ -16,7 +16,7 @@ * limitations under the License. */ -#include "tensorrt_llm/kernels/nvfp4BoundaryKernels.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.h" #include "tensorrt_llm/common/cudaUtils.h" @@ -41,13 +41,13 @@ namespace using tensorrt_llm::batch_manager::kv_cache_manager_v2::HostMem; using tensorrt_llm::batch_manager::kv_cache_manager_v2::MemAddress; -using tensorrt_llm::kernels::Nvfp4BoundaryBufferPlan; -using tensorrt_llm::kernels::Nvfp4BoundaryKernelParams; -using tensorrt_llm::kernels::Nvfp4BoundaryOffloadPageTask; -using tensorrt_llm::kernels::Nvfp4BoundaryOnboardPageTask; -using tensorrt_llm::kernels::Nvfp4BoundaryPreparedPlan; -using tensorrt_llm::kernels::Nvfp4BoundaryRuntimeType; -using tensorrt_llm::kernels::Nvfp4BoundaryTransform; +using tensorrt_llm::kernels::Nvfp4ColdPageBufferPlan; +using tensorrt_llm::kernels::Nvfp4ColdPageKernelParams; +using tensorrt_llm::kernels::Nvfp4ColdPageOffloadPageTask; +using tensorrt_llm::kernels::Nvfp4ColdPageOnboardPageTask; +using tensorrt_llm::kernels::Nvfp4ColdPagePreparedPlan; +using tensorrt_llm::kernels::Nvfp4ColdPageRuntimeType; +using tensorrt_llm::kernels::Nvfp4ColdPageTransform; constexpr std::size_t kGuardBytes = 64; constexpr std::uint8_t kCanary = 0xA5; @@ -273,9 +273,9 @@ std::size_t scaleBytes(PageGeometry const& geometry) return numElements(geometry) / 16; } -Nvfp4BoundaryKernelParams makeParams(PageGeometry const& geometry = kDefaultGeometry, std::uint32_t role = 0U) +Nvfp4ColdPageKernelParams makeParams(PageGeometry const& geometry = kDefaultGeometry, std::uint32_t role = 0U) { - Nvfp4BoundaryKernelParams params{}; + Nvfp4ColdPageKernelParams params{}; params.numKvHeads = geometry.numHeads; params.tokensPerPage = geometry.tokensPerPage; params.headDim = geometry.headDim; @@ -286,15 +286,15 @@ Nvfp4BoundaryKernelParams makeParams(PageGeometry const& geometry = kDefaultGeom return params; } -Nvfp4BoundaryRuntimeType runtimeType(RawKind kind) +Nvfp4ColdPageRuntimeType runtimeType(RawKind kind) { switch (kind) { - case RawKind::kFloat16: return Nvfp4BoundaryRuntimeType::kFloat16; - case RawKind::kBfloat16: return Nvfp4BoundaryRuntimeType::kBfloat16; - case RawKind::kFp8: return Nvfp4BoundaryRuntimeType::kFp8E4m3; + case RawKind::kFloat16: return Nvfp4ColdPageRuntimeType::kFloat16; + case RawKind::kBfloat16: return Nvfp4ColdPageRuntimeType::kBfloat16; + case RawKind::kFp8: return Nvfp4ColdPageRuntimeType::kFp8E4m3; } - return Nvfp4BoundaryRuntimeType::kFloat16; + return Nvfp4ColdPageRuntimeType::kFloat16; } template @@ -312,7 +312,7 @@ T loadScalar(std::vector const& bytes, std::size_t index) } void storeRawValue(std::vector& bytes, RawKind kind, std::size_t index, float value, - Nvfp4BoundaryKernelParams const& params) + Nvfp4ColdPageKernelParams const& params) { switch (kind) { @@ -323,7 +323,7 @@ void storeRawValue(std::vector& bytes, RawKind kind, std::size_t i } float loadRawValue( - std::vector const& bytes, RawKind kind, std::size_t index, Nvfp4BoundaryKernelParams const& params) + std::vector const& bytes, RawKind kind, std::size_t index, Nvfp4ColdPageKernelParams const& params) { switch (kind) { @@ -369,7 +369,7 @@ std::uint8_t quantizeE2m1(float value) //! Exactly representable E2M1 values and E4M3 scales keep byte comparisons deterministic. std::vector makeRawPage(RawKind kind, std::size_t page, std::uint32_t role, - Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry, InputPattern inputPattern) + Nvfp4ColdPageKernelParams const& params, PageGeometry const& geometry, InputPattern inputPattern) { constexpr std::array densePattern{ 0.0F, 0.5F, -1.0F, 1.5F, -2.0F, 3.0F, -4.0F, 6.0F, -0.5F, 1.0F, -1.5F, 2.0F, -3.0F, 4.0F, -6.0F, 0.5F}; @@ -418,7 +418,7 @@ struct ReferenceNvfp4 }; ReferenceNvfp4 compressReference(std::vector const& raw, RawKind kind, - Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) + Nvfp4ColdPageKernelParams const& params, PageGeometry const& geometry) { ReferenceNvfp4 result{{}, {}}; result.packed.resize(packedBytes(geometry)); @@ -453,7 +453,7 @@ ReferenceNvfp4 compressReference(std::vector const& raw, RawKind k } std::vector decompressReference(ReferenceNvfp4 const& compressed, RawKind kind, - Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) + Nvfp4ColdPageKernelParams const& params, PageGeometry const& geometry) { std::vector raw(rawBytes(kind, geometry)); std::uint32_t const scalesPerRow = static_cast(geometry.headDim) / 16; @@ -477,16 +477,16 @@ std::vector decompressReference(ReferenceNvfp4 const& compressed, return raw; } -void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultGeometry, +void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultGeometry, std::size_t numPages = kDefaultNumPages, InputPattern inputPattern = InputPattern::kDense, bool synchronizeBetweenDirections = true, bool repeatRoundTrip = false, std::size_t coldBaseOffset = 0) { ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } - std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; + std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; CudaStream stream; std::size_t const rawSlotBytes = rawBytes(kind, geometry); // Align compact Slot strides while independently testing an arbitrary staging-base offset. @@ -500,7 +500,7 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG MappedHostRegion compactPages(coldBaseOffset + slotCapacity * compactSlotBytes); auto* compactBase = compactPages.bytes() + coldBaseOffset; std::vector, 2>> rawHost(numPages); - std::vector offloadTasks; + std::vector offloadTasks; offloadTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { @@ -516,14 +516,14 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG std::size_t const scale = scaleBytes(geometry); std::size_t const payloadBytes = 2U * (packed + scale); std::uint32_t const paddingBytes = static_cast(compactSlotBytes - payloadBytes); - std::vector const inputBuffers{ + std::vector const inputBuffers{ {reinterpret_cast(rawInputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, - Nvfp4BoundaryTransform::kNvfp4, params[0]}, + Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawInputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, - payloadBytes, paddingBytes, Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + payloadBytes, paddingBytes, Nvfp4ColdPageTransform::kNvfp4, params[1]}}; auto const inputPlan - = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan(inputBuffers, compactSlotBytes, runtimeType(kind)); + = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan(inputBuffers, compactSlotBytes, runtimeType(kind)); std::vector> references(numPages); for (std::size_t page = 0; page < numPages; ++page) { @@ -558,7 +558,7 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG } }; - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactBase, stream); if (synchronizeBetweenDirections) { @@ -579,28 +579,28 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG { std::memset(compactBase + 2U * page * compactSlotBytes, 0x5A, compactSlotBytes); } - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); EXPECT_EQ(compactPages.payload(), firstSerialization); verifyCompressedPages(); } } - std::vector onboardTasks; + std::vector onboardTasks; onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { std::size_t const slot = 2U * page; onboardTasks.push_back({static_cast(slot), static_cast(slot)}); } - std::vector const outputBuffers{ + std::vector const outputBuffers{ {reinterpret_cast(rawOutputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, - Nvfp4BoundaryTransform::kNvfp4, params[0]}, + Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawOutputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, - payloadBytes, paddingBytes, Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + payloadBytes, paddingBytes, Nvfp4ColdPageTransform::kNvfp4, params[1]}}; auto const outputPlan - = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan(outputBuffers, compactSlotBytes, runtimeType(kind)); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, compactBase, stream); + = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan(outputBuffers, compactSlotBytes, runtimeType(kind)); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); if (!synchronizeBetweenDirections) @@ -620,8 +620,8 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG references[page][role] = compressReference(restored, kind, params[role], geometry); } } - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, outputPlan, compactBase, stream); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, outputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, inputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); verifyCompressedPages(); } @@ -644,7 +644,7 @@ void runBoundaryRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG } std::vector makePartialRawPage(RawKind kind, std::int32_t validTokens, bool zeroTail, std::uint32_t role, - Nvfp4BoundaryKernelParams const& params, PageGeometry const& geometry) + Nvfp4ColdPageKernelParams const& params, PageGeometry const& geometry) { std::vector bytes(rawBytes(kind, geometry)); for (std::int32_t head = 0; head < geometry.numHeads; ++head) @@ -722,13 +722,13 @@ void runPartialPageTailIsolation(RawKind kind) ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } PageGeometry constexpr geometry = kModelLikeGeometry; std::size_t constexpr pageVariants = 2U; std::size_t const numPages = pageVariants * kValidTokenCounts.size(); - std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; + std::array const params{makeParams(geometry, 0U), makeParams(geometry, 1U)}; std::size_t const rawSlotBytes = rawBytes(kind, geometry); std::size_t const compactSlotBytes = roundUp(2U * (packedBytes(geometry) + scaleBytes(geometry)), alignof(uint4)); @@ -737,8 +737,8 @@ void runPartialPageTailIsolation(RawKind kind) DeviceRegion rawOutputK(numPages * rawSlotBytes); DeviceRegion rawOutputV(numPages * rawSlotBytes); MappedHostRegion compactPages(numPages * compactSlotBytes); - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; offloadTasks.reserve(numPages); onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) @@ -759,20 +759,20 @@ void runPartialPageTailIsolation(RawKind kind) std::size_t const payloadBytes = 2U * (packed + scale); auto const makeBuffers = [&](DeviceRegion const& rawK, DeviceRegion const& rawV) { - return std::vector{ + return std::vector{ {reinterpret_cast(rawK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, - Nvfp4BoundaryTransform::kNvfp4, params[0]}, + Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, payloadBytes, static_cast(compactSlotBytes - payloadBytes), - Nvfp4BoundaryTransform::kNvfp4, params[1]}}; + Nvfp4ColdPageTransform::kNvfp4, params[1]}}; }; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( makeBuffers(rawInputK, rawInputV), compactSlotBytes, runtimeType(kind)); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( makeBuffers(rawOutputK, rawOutputV), compactSlotBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, compactPages.data(), stream); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, compactPages.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactPages.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, compactPages.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); for (std::size_t pair = 0; pair < kValidTokenCounts.size(); ++pair) @@ -848,14 +848,14 @@ RoundTripCase constexpr kRoundTripCases[]{ {"PdlFp8", RawKind::kFp8, kSmallVectorGeometry, 65, InputPattern::kDense, false}, }; -class Nvfp4BoundaryRoundTripTest : public testing::TestWithParam +class Nvfp4ColdPageRoundTripTest : public testing::TestWithParam { }; -TEST_P(Nvfp4BoundaryRoundTripTest, MatchesReference) +TEST_P(Nvfp4ColdPageRoundTripTest, MatchesReference) { auto const& test = GetParam(); - runBoundaryRoundTrip(test.kind, test.geometry, test.numPages, test.inputPattern, test.synchronizeBetweenDirections, + runColdPageRoundTrip(test.kind, test.geometry, test.numPages, test.inputPattern, test.synchronizeBetweenDirections, test.repeatRoundTrip, test.coldBaseOffset); } @@ -864,14 +864,14 @@ std::string roundTripCaseName(testing::TestParamInfo const& info) return info.param.name; } -INSTANTIATE_TEST_SUITE_P(Scenarios, Nvfp4BoundaryRoundTripTest, testing::ValuesIn(kRoundTripCases), roundTripCaseName); +INSTANTIATE_TEST_SUITE_P(Scenarios, Nvfp4ColdPageRoundTripTest, testing::ValuesIn(kRoundTripCases), roundTripCaseName); void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) { ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } PageGeometry constexpr geometry{1, 64, 576}; @@ -899,8 +899,8 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) std::array, numPages> mlaHost; std::array, numPages> sideHost; std::array references; - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; for (std::size_t page = 0; page < numPages; ++page) { mlaHost[page] = makeRawPage(kind, page, 0U, params, geometry, InputPattern::kDense); @@ -919,21 +919,21 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) auto const makePlans = [&](DeviceRegion const& mla, DeviceRegion const& side) { - return std::vector{ + return std::vector{ {reinterpret_cast(mla.data()), mlaRawBytes, mlaRawBytes, 0U, mlaPackedBytes, - mlaPayloadBytes, static_cast(gapBeforeSide), Nvfp4BoundaryTransform::kNvfp4, params}, + mlaPayloadBytes, static_cast(gapBeforeSide), Nvfp4ColdPageTransform::kNvfp4, params}, {reinterpret_cast(side.data()), sideSlotBytes, sideRawBytes, sideColdOffset, 0U, sideColdEnd, static_cast(coldPageBytes - sideColdEnd), - Nvfp4BoundaryTransform::kLosslessCopy, {}}}; + Nvfp4ColdPageTransform::kLosslessCopy, {}}}; }; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( makePlans(mlaInput, sideInput), coldPageBytes, runtimeType(kind)); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( makePlans(mlaOutput, sideOutput), coldPageBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress(offloadTasks, inputPlan, coldBase, stream); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress(onboardTasks, outputPlan, coldBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, coldBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, coldBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const cold = coldStorage.payload(); @@ -973,24 +973,24 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) coldStorage.expectCanaries(); } -class Nvfp4BoundaryMlaSideTest : public testing::TestWithParam +class Nvfp4ColdPageMlaSideTest : public testing::TestWithParam { }; -TEST_P(Nvfp4BoundaryMlaSideTest, MlaPageAndDefaultDsaIndexKeyRoundTripExactly) +TEST_P(Nvfp4ColdPageMlaSideTest, MlaPageAndDefaultDsaIndexKeyRoundTripExactly) { runUnaryMlaWithLosslessSideRoundTrip(GetParam()); } INSTANTIATE_TEST_SUITE_P( - AllRuntimeTypes, Nvfp4BoundaryMlaSideTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); + AllRuntimeTypes, Nvfp4ColdPageMlaSideTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); -TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatch) +TEST(Nvfp4ColdPageWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatch) { ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } constexpr std::size_t numLayers = 2; @@ -1005,7 +1005,7 @@ TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc std::array, numLayers> rawInputV; std::array, numLayers> rawOutputK; std::array, numLayers> rawOutputV; - std::array, numLayers> params{ + std::array, numLayers> params{ std::array{makeParams(geometry, 0U), makeParams(geometry, 1U)}, std::array{makeParams(geometry, 0U), makeParams(geometry, 1U)}}; // Distinct per-layer K/V scales verify blockIdx.y selects immutable launch metadata. @@ -1016,8 +1016,8 @@ TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc std::array, 2>, numLayers> rawHost; std::array, numLayers> references; - std::vector inputPlans; - std::vector outputPlans; + std::vector inputPlans; + std::vector outputPlans; inputPlans.reserve(2U * numLayers); outputPlans.reserve(2U * numLayers); for (std::size_t layer = 0; layer < numLayers; ++layer) @@ -1039,10 +1039,10 @@ TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc auto const appendPlans = [&](auto& plans, DeviceRegion const& rawK, DeviceRegion const& rawV) { plans.push_back({reinterpret_cast(rawK.data()), rawSlotBytes, rawSlotBytes, base, - base + 2U * packed, 0U, 0U, Nvfp4BoundaryTransform::kNvfp4, params[layer][0]}); + base + 2U * packed, 0U, 0U, Nvfp4ColdPageTransform::kNvfp4, params[layer][0]}); plans.push_back({reinterpret_cast(rawV.data()), rawSlotBytes, rawSlotBytes, base + packed, base + 2U * packed + scale, base + layerRecordBytes, - static_cast(layerRecordStride - layerRecordBytes), Nvfp4BoundaryTransform::kNvfp4, + static_cast(layerRecordStride - layerRecordBytes), Nvfp4ColdPageTransform::kNvfp4, params[layer][1]}); }; appendPlans(inputPlans, *rawInputK[layer], *rawInputV[layer]); @@ -1051,12 +1051,12 @@ TEST(Nvfp4BoundaryWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc MappedHostRegion compactPage(coldPageBytes); CudaStream stream; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( - inputPlans, coldPageBytes, Nvfp4BoundaryRuntimeType::kBfloat16); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4BoundaryPlan( - outputPlans, coldPageBytes, Nvfp4BoundaryRuntimeType::kBfloat16); - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({{0, 0}}, inputPlan, compactPage.data(), stream); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress({{0, 0}}, outputPlan, compactPage.data(), stream); + auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( + inputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); + auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( + outputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({{0, 0}}, inputPlan, compactPage.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode({{0, 0}}, outputPlan, compactPage.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const compact = compactPage.payload(); @@ -1095,7 +1095,7 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector, numLayers> rawK; std::array, numLayers> rawV; - std::vector buffers; + std::vector buffers; buffers.reserve(2U * numLayers); for (std::size_t layer = 0; layer < numLayers; ++layer) { @@ -1109,15 +1109,15 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector(rawK[layer]->data()), rawSlotBytes, rawSlotBytes, base, - base + 2U * packed, 0U, 0U, Nvfp4BoundaryTransform::kNvfp4, kParams}); + base + 2U * packed, 0U, 0U, Nvfp4ColdPageTransform::kNvfp4, kParams}); buffers.push_back({reinterpret_cast(rawV[layer]->data()), rawSlotBytes, rawSlotBytes, base + packed, base + 2U * packed + scale, base + recordBytes, - static_cast(recordStride - recordBytes), Nvfp4BoundaryTransform::kNvfp4, vParams}); + static_cast(recordStride - recordBytes), Nvfp4ColdPageTransform::kNvfp4, vParams}); } MappedHostRegion coldPages(numPages * coldPageBytes); - std::vector offloadPages; - std::vector onboardPages; + std::vector offloadPages; + std::vector onboardPages; offloadPages.reserve(numPages); onboardPages.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) @@ -1129,7 +1129,7 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector +class Nvfp4ColdPageTailTest : public testing::TestWithParam { }; -TEST_P(Nvfp4BoundaryTailTest, InactiveRowsDoNotAffectTheValidPrefix) +TEST_P(Nvfp4ColdPageTailTest, InactiveRowsDoNotAffectTheValidPrefix) { runPartialPageTailIsolation(GetParam()); } INSTANTIATE_TEST_SUITE_P( - AllRuntimeTypes, Nvfp4BoundaryTailTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); + AllRuntimeTypes, Nvfp4ColdPageTailTest, testing::Values(RawKind::kFloat16, RawKind::kBfloat16, RawKind::kFp8)); -TEST(Nvfp4BoundaryValidationTest, EmptyBatchIsAnAsyncNoOp) +TEST(Nvfp4ColdPageValidationTest, EmptyBatchIsAnAsyncNoOp) { - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({}, Nvfp4BoundaryPreparedPlan{}, nullptr, nullptr); - tensorrt_llm::kernels::invokeNvfp4BoundaryOnboardDecompress({}, Nvfp4BoundaryPreparedPlan{}, nullptr, nullptr); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({}, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode({}, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); } -TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) +TEST(Nvfp4ColdPageValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) { ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); LayerBuffers buffers(rawSlotBytes); std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); auto const prepare = - [&](Nvfp4BoundaryKernelParams const& params, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + [&](Nvfp4ColdPageKernelParams const& params, Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) { std::size_t activeRawBytes = rawSlotBytes; std::size_t dataBytes = packedBytes(kDefaultGeometry); @@ -1219,7 +1219,7 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) { std::uint64_t const elements = static_cast(params.numKvHeads) * static_cast(params.tokensPerPage) * static_cast(params.headDim); - std::uint64_t const candidateRawBytes = elements * (type == Nvfp4BoundaryRuntimeType::kFp8E4m3 ? 1U : 2U); + std::uint64_t const candidateRawBytes = elements * (type == Nvfp4ColdPageRuntimeType::kFp8E4m3 ? 1U : 2U); if (candidateRawBytes > 0U && candidateRawBytes <= rawSlotBytes) { activeRawBytes = static_cast(candidateRawBytes); @@ -1227,13 +1227,13 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) scales = static_cast(elements / 16U); } } - Nvfp4BoundaryBufferPlan const buffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, + Nvfp4ColdPageBufferPlan const buffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, activeRawBytes, 0U, dataBytes, dataBytes + scales, - static_cast(coldPageBytes - dataBytes - scales), Nvfp4BoundaryTransform::kNvfp4, params}; - static_cast(tensorrt_llm::kernels::prepareNvfp4BoundaryPlan({buffer}, coldPageBytes, type)); + static_cast(coldPageBytes - dataBytes - scales), Nvfp4ColdPageTransform::kNvfp4, params}; + static_cast(tensorrt_llm::kernels::prepareNvfp4ColdPagePlan({buffer}, coldPageBytes, type)); }; auto const expectInvalid - = [&](char const* name, auto const& mutate, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + = [&](char const* name, auto const& mutate, Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) { SCOPED_TRACE(name); auto params = makeParams(); @@ -1260,35 +1260,35 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) [](auto& params) { params.nvfp4ScaleQuantOrig = std::numeric_limits::infinity(); }); expectInvalid( "zero FP8 quant scale", [](auto& params) { params.fp8ScaleOrigQuant = 0.0F; }, - Nvfp4BoundaryRuntimeType::kFp8E4m3); + Nvfp4ColdPageRuntimeType::kFp8E4m3); expectInvalid( "infinite FP8 dequant scale", [](auto& params) { params.fp8ScaleQuantOrig = std::numeric_limits::infinity(); }, - Nvfp4BoundaryRuntimeType::kFp8E4m3); + Nvfp4ColdPageRuntimeType::kFp8E4m3); } -TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) +TEST(Nvfp4ColdPageValidationTest, RejectsInvalidLaunchDescriptors) { ASSERT_EQ(cudaSetDevice(0), cudaSuccess); if (!tensorrt_llm::common::isSM100Family()) { - GTEST_SKIP() << "NVFP4 boundary kernels require an SM100-family GPU"; + GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; } std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); LayerBuffers buffers(rawSlotBytes); - Nvfp4BoundaryOffloadPageTask const validOffload{0, 0}; + Nvfp4ColdPageOffloadPageTask const validOffload{0, 0}; std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); std::size_t const packed = packedBytes(kDefaultGeometry); std::size_t const scale = scaleBytes(kDefaultGeometry); - Nvfp4BoundaryBufferPlan const validBuffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, + Nvfp4ColdPageBufferPlan const validBuffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, rawSlotBytes, 0U, packed, packed + scale, static_cast(coldPageBytes - packed - scale), - Nvfp4BoundaryTransform::kNvfp4, makeParams()}; - auto const prepare = [&](Nvfp4BoundaryBufferPlan const& buffer, std::size_t pageBytes, - Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) - { return tensorrt_llm::kernels::prepareNvfp4BoundaryPlan({buffer}, pageBytes, type); }; + Nvfp4ColdPageTransform::kNvfp4, makeParams()}; + auto const prepare = [&](Nvfp4ColdPageBufferPlan const& buffer, std::size_t pageBytes, + Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) + { return tensorrt_llm::kernels::prepareNvfp4ColdPagePlan({buffer}, pageBytes, type); }; auto const expectInvalid - = [&](char const* name, auto const& mutate, Nvfp4BoundaryRuntimeType type = Nvfp4BoundaryRuntimeType::kFloat16) + = [&](char const* name, auto const& mutate, Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) { SCOPED_TRACE(name); auto buffer = validBuffer; @@ -1298,8 +1298,7 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) auto const validPlan = prepare(validBuffer, coldPageBytes); expectInvalid("unaligned raw base", [](auto& buffer) { buffer.rawBase += 1U; }); - EXPECT_ANY_THROW( - tensorrt_llm::kernels::invokeNvfp4BoundaryOffloadCompress({validOffload}, validPlan, nullptr, nullptr)); + EXPECT_ANY_THROW(tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({validOffload}, validPlan, nullptr, nullptr)); expectInvalid("unaligned raw stride", [](auto& buffer) { buffer.rawSlotBytes += alignof(uint4) / 2U; }); expectInvalid("raw bytes exceed stride", [](auto& buffer) { buffer.rawBytes = buffer.rawSlotBytes + 1U; }); expectInvalid("raw bytes mismatch geometry", [](auto& buffer) { buffer.rawBytes -= alignof(uint4); }); @@ -1307,9 +1306,9 @@ TEST(Nvfp4BoundaryValidationTest, RejectsInvalidLaunchDescriptors) expectInvalid("cold intervals overlap", [](auto& buffer) { buffer.coldScaleOffset = buffer.coldDataOffset; }); EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes + alignof(uint4) / 2U))); expectInvalid( - "unaligned FP8 raw base", [](auto& buffer) { buffer.rawBase += 1U; }, Nvfp4BoundaryRuntimeType::kFp8E4m3); + "unaligned FP8 raw base", [](auto& buffer) { buffer.rawBase += 1U; }, Nvfp4ColdPageRuntimeType::kFp8E4m3); - auto const unsupportedType = static_cast(255); + auto const unsupportedType = static_cast(255); EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes, unsupportedType))); } diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index e442c2015239..650b8a19f1df 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -1,12 +1,12 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. # All rights reserved. SPDX-License-Identifier: Apache-2.0 -set(NVFP4_COLD_PAGE_CODEC_TEST_SRC - nvfp4ColdPageCodecTest.cpp +set(COLD_PAGE_CODEC_TEST_SRC + coldPageCodecTest.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp @@ -14,10 +14,7 @@ set(NVFP4_COLD_PAGE_CODEC_TEST_SRC ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) -add_gtest(nvfp4ColdPageCodecTest "${NVFP4_COLD_PAGE_CODEC_TEST_SRC}" - NO_TLLM_LINKAGE) -target_link_libraries(nvfp4ColdPageCodecTest PRIVATE CUDA::cuda_driver - CUDA::cudart) +add_gtest(coldPageCodecTest "${COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) +target_link_libraries(coldPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) target_include_directories( - nvfp4ColdPageCodecTest - PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) + coldPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp new file mode 100644 index 000000000000..0af91c512de2 --- /dev/null +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -0,0 +1,391 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +static_assert(std::is_base_of_v); + +struct RecordedLaunch +{ + int encodeCalls = 0; + int decodeCalls = 0; + std::vector encodePages; + std::vector decodePages; + kernels::Nvfp4ColdPagePreparedPlan plan; + void const* coldBase = nullptr; + cudaStream_t stream{}; +}; + +RecordedLaunch gLaunch; + +constexpr std::uintptr_t kGpuKBase = 0x100000; +constexpr std::uintptr_t kGpuVBase = 0x200000; +constexpr std::uintptr_t kColdBase = 0x300000; +constexpr std::uintptr_t kStreamValue = 0x7000; +constexpr std::size_t kRawBytes = 320; +constexpr std::size_t kColdBytes = 192; + +void resetLaunch() +{ + gLaunch = {}; +} + +Nvfp4ColdPageParams makeParams(float scale = 1.0F) +{ + Nvfp4ColdPageParams params; + params.runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; + params.numKvHeads = 1; + params.tokensPerPage = 5; + params.headDim = 32; + params.nvfp4ScaleOrigQuant = scale; + params.nvfp4ScaleQuantOrig = 1.0F / scale; + return params; +} + +ColdPageLayerPlan makeAttentionPlan(int layerId, float keyScale = 1.0F, float valueScale = 2.0F) +{ + return ColdPageLayerPlan{layerId, kColdBytes, 180U, 12U, + {ColdPageBufferPlan{"key", ColdPageTransformKind::kNvfp4, kRawBytes, 0U, 160U, makeParams(keyScale)}, + ColdPageBufferPlan{"value", ColdPageTransformKind::kNvfp4, kRawBytes, 80U, 170U, makeParams(valueScale)}}}; +} + +ColdPageLayerPlan makeMlaPlan(int layerId, bool hasIndex) +{ + std::vector buffers; + buffers.push_back(ColdPageBufferPlan{"key", ColdPageTransformKind::kNvfp4, kRawBytes, 0U, 80U, makeParams()}); + if (hasIndex) + { + buffers.push_back( + ColdPageBufferPlan{"index_key", ColdPageTransformKind::kLosslessCopy, 68U, 90U, 0U, std::nullopt}); + return ColdPageLayerPlan{layerId, 160U, 158U, 2U, std::move(buffers)}; + } + return ColdPageLayerPlan{layerId, 96U, 90U, 6U, std::move(buffers)}; +} + +kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::PoolGroupIndex{0}, + kv::LayerGroupId lifeCycle = kv::LayerGroupId{0}, std::size_t count = 1U, int firstLayer = 0, + std::uintptr_t keyBase = kGpuKBase, std::uintptr_t valueBase = kGpuVBase) +{ + kv::CoalescedBuffer keys{kRawBytes, {}}; + kv::CoalescedBuffer values{kRawBytes, {}}; + for (std::size_t index = 0; index < count; ++index) + { + auto const layerId = firstLayer + static_cast(index); + keys.bufferIds.push_back({layerId, "key"}); + values.bufferIds.push_back({layerId, "value"}); + } + auto const slotBytes = count * kRawBytes; + kv::SlotDescVariant variant{ + lifeCycle, kv::TypedVec{std::move(keys), std::move(values)}}; + return kv::PoolGroupDesc{poolGroupIndex, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, + kv::TypedVec{ + {kv::PoolIndex{0}, keyBase, slotBytes}, {kv::PoolIndex{1}, valueBase, slotBytes}}}; +} + +kv::PoolGroupDesc makeMlaDesc(bool hasIndex) +{ + kv::TypedVec buffers; + buffers.push_back(kv::CoalescedBuffer{kRawBytes, {{0, "key"}}}); + if (hasIndex) + { + buffers.push_back(kv::CoalescedBuffer{68U, {{0, "index_key"}}}); + } + kv::SlotDescVariant variant{kv::LayerGroupId{0}, std::move(buffers)}; + kv::TypedVec pools; + pools.push_back({kv::PoolIndex{0}, kGpuKBase, kRawBytes}); + if (hasIndex) + { + pools.push_back({kv::PoolIndex{1}, kGpuVBase, 68U}); + } + return {kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, std::move(pools)}; +} + +kv::PoolGroupDesc makeLosslessDesc(kv::PoolGroupIndex poolGroupIndex, kv::LayerGroupId lifeCycle) +{ + kv::SlotDescVariant variant; + variant.lifeCycleId = lifeCycle; + variant.coalescedBuffers = kv::TypedVec{ + kv::CoalescedBuffer{64U, {{10, "ssm_state"}}}, kv::CoalescedBuffer{32U, {{10, "conv_state"}}}}; + return {poolGroupIndex, kv::SlotCount{8}, kv::SlotDesc{{std::move(variant)}}, + kv::TypedVec{ + {kv::PoolIndex{0}, 0x400000, 64U}, {kv::PoolIndex{1}, 0x500000, 32U}}}; +} + +bool configureOne(PlannedColdPageCodec& codec, kv::PoolGroupDesc const& desc) +{ + return codec.configure(&desc, kv::PoolGroupIndex{1}); +} + +TEST(PlannedColdPageCodecTest, ResolvesPythonAuthoredLayoutAndScales) +{ + resetLaunch(); + PlannedColdPageCodec codec{{makeAttentionPlan(0, 2.0F, 3.0F), makeAttentionPlan(1, 4.0F, 5.0F)}}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{3}, 2U))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{3}), 2U * kColdBytes); + + kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; + auto const stream = reinterpret_cast(kStreamValue); + ASSERT_TRUE( + codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + ASSERT_EQ(gLaunch.encodeCalls, 1); + ASSERT_EQ(gLaunch.encodePages.size(), 2U); + EXPECT_EQ(gLaunch.encodePages[0].gpuPageIndex, 1); + EXPECT_EQ(gLaunch.encodePages[0].coldPageIndex, 2); + ASSERT_EQ(gLaunch.plan.numBuffers, 4U); + EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); + EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kRawBytes); + EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, kColdBytes); + EXPECT_EQ(gLaunch.plan.buffers[3].coldScaleOffset, kColdBytes + 170U); + EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingOffset, kColdBytes + 180U); + EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingBytes, 12U); + EXPECT_FLOAT_EQ(gLaunch.plan.buffers[0].params.nvfp4ScaleOrigQuant, 2.0F); + EXPECT_FLOAT_EQ(gLaunch.plan.buffers[3].params.nvfp4ScaleOrigQuant, 5.0F); + + ASSERT_TRUE(codec.decode( + kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + ASSERT_EQ(gLaunch.decodeCalls, 1); + EXPECT_EQ(gLaunch.decodePages[0].gpuPageIndex, 2); + EXPECT_EQ(gLaunch.decodePages[0].coldPageIndex, 1); +} + +TEST(PlannedColdPageCodecTest, PreservesExplicitMlaLosslessSideBuffer) +{ + resetLaunch(); + PlannedColdPageCodec codec{{makeMlaPlan(0, true)}}; + ASSERT_TRUE(configureOne(codec, makeMlaDesc(true))); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 160U); + + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 2U); + EXPECT_EQ(gLaunch.plan.buffers[0].transform, kernels::Nvfp4ColdPageTransform::kNvfp4); + EXPECT_EQ(gLaunch.plan.buffers[1].transform, kernels::Nvfp4ColdPageTransform::kLosslessCopy); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBytes, 68U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingOffset, 158U); +} + +TEST(PlannedColdPageCodecTest, UnplannedLifecycleDelegatesToDefaultCodec) +{ + PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; + std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), kColdBytes); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 96U); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); +} + +TEST(PlannedColdPageCodecTest, RejectsDuplicateLayerAndRolePlans) +{ + std::vector duplicateLayers{makeAttentionPlan(0), makeAttentionPlan(0)}; + EXPECT_THROW({ PlannedColdPageCodec codec{duplicateLayers}; }, std::invalid_argument); + + auto plan = makeAttentionPlan(0); + plan.buffers.push_back(plan.buffers.front()); + EXPECT_THROW({ PlannedColdPageCodec codec{{plan}}; }, std::invalid_argument); +} + +TEST(PlannedColdPageCodecTest, RejectsMissingNvfp4Parameters) +{ + auto plan = makeAttentionPlan(0); + plan.buffers.front().nvfp4Params.reset(); + PlannedColdPageCodec codec{{plan}}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); +} + +TEST(PlannedColdPageCodecTest, RejectsRawSizeAndRoleMismatches) +{ + auto wrongSize = makeAttentionPlan(0); + wrongSize.buffers.front().rawBytes += 1U; + PlannedColdPageCodec sizeCodec{{wrongSize}}; + EXPECT_FALSE(configureOne(sizeCodec, makeAttentionDesc())); + + auto missingRole = makeAttentionPlan(0); + missingRole.buffers.pop_back(); + PlannedColdPageCodec roleCodec{{missingRole}}; + EXPECT_FALSE(configureOne(roleCodec, makeAttentionDesc())); +} + +TEST(PlannedColdPageCodecTest, RejectsPlannedAndUnplannedLayersInOneLifecycle) +{ + PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2U))); +} + +TEST(PlannedColdPageCodecTest, RejectsPlanAbsentFromGpuDescriptors) +{ + PlannedColdPageCodec codec{{makeAttentionPlan(0), makeAttentionPlan(1)}}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); +} + +TEST(PlannedColdPageCodecTest, RejectsDuplicateLifecycleAcrossPoolGroups) +{ + PlannedColdPageCodec codec{{makeAttentionPlan(0), makeAttentionPlan(1)}}; + std::array descs{ + makeAttentionDesc(), makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1U, 1, 0x600000, 0x700000)}; + EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); +} + +TEST(PlannedColdPageCodecTest, RejectsInvalidLayerRelativeIntervals) +{ + auto plan = makeAttentionPlan(0); + plan.buffers.back().coldDataOffset = plan.coldPageBytes; + PlannedColdPageCodec codec{{plan}}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); +} + +TEST(PlannedColdPageCodecTest, EmptyBatchIsValidAndInvalidBatchesFailBeforeLaunch) +{ + resetLaunch(); + PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); + EXPECT_TRUE(codec.encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + EXPECT_TRUE(codec.decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + + kv::PageIndexPair const valid[]{{0, 0}}; + EXPECT_FALSE(codec.encode(kv::LayerGroupId{0}, nullptr, valid, 1U, nullptr)); + EXPECT_FALSE(codec.decode(kv::LayerGroupId{0}, nullptr, valid, 1U, nullptr)); + EXPECT_EQ(gLaunch.encodeCalls, 0); + EXPECT_EQ(gLaunch.decodeCalls, 0); +} + +TEST(PlannedColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) +{ + PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); + EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); +} + +} // namespace +} // namespace tensorrt_llm::kv_cache_compression + +namespace tensorrt_llm::kernels +{ +namespace +{ + +struct Interval +{ + std::size_t begin; + std::size_t end; +}; + +void appendInterval(std::vector& intervals, std::size_t begin, std::size_t bytes, std::size_t pageBytes) +{ + if (bytes == 0U) + { + return; + } + if (begin > pageBytes || bytes > pageBytes - begin) + { + throw std::invalid_argument("test interval exceeds cold Page"); + } + intervals.push_back({begin, begin + bytes}); +} + +} // namespace + +Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) +{ + if (buffers.empty() || buffers.size() > kNvfp4ColdPageMaxBuffersPerLaunch || coldPageBytes == 0U) + { + throw std::invalid_argument("invalid test launch plan"); + } + std::vector intervals; + for (auto const& buffer : buffers) + { + if (buffer.rawBytes == 0U || buffer.rawBytes > buffer.rawSlotBytes) + { + throw std::invalid_argument("invalid raw buffer"); + } + if (buffer.transform == Nvfp4ColdPageTransform::kNvfp4) + { + auto const& params = buffer.params; + if (params.numKvHeads <= 0 || params.tokensPerPage <= 0 || params.headDim <= 0 || params.headDim % 16 != 0) + { + throw std::invalid_argument("invalid NVFP4 geometry"); + } + std::size_t const elements = static_cast(params.numKvHeads) + * static_cast(params.tokensPerPage) * static_cast(params.headDim); + auto const elementBytes = runtimeType == Nvfp4ColdPageRuntimeType::kFp8E4m3 ? 1U : 2U; + if (buffer.rawBytes != elements * elementBytes) + { + throw std::invalid_argument("raw size mismatch"); + } + appendInterval(intervals, buffer.coldDataOffset, elements / 2U, coldPageBytes); + appendInterval(intervals, buffer.coldScaleOffset, elements / 16U, coldPageBytes); + } + else + { + appendInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes); + } + appendInterval(intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes); + } + std::sort( + intervals.begin(), intervals.end(), [](auto const& lhs, auto const& rhs) { return lhs.begin < rhs.begin; }); + for (std::size_t index = 1; index < intervals.size(); ++index) + { + if (intervals[index - 1U].end > intervals[index].begin) + { + throw std::invalid_argument("test intervals overlap"); + } + } + + Nvfp4ColdPagePreparedPlan plan; + std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); + plan.numBuffers = static_cast(buffers.size()); + plan.coldPageBytes = coldPageBytes; + plan.runtimeType = runtimeType; + return plan; +} + +void invokeNvfp4ColdPageEncode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) +{ + auto& launch = kv_cache_compression::gLaunch; + ++launch.encodeCalls; + launch.encodePages = pages; + launch.plan = plan; + launch.coldBase = coldBase; + launch.stream = stream; +} + +void invokeNvfp4ColdPageDecode(std::vector const& pages, + Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream) +{ + auto& launch = kv_cache_compression::gLaunch; + ++launch.decodeCalls; + launch.decodePages = pages; + launch.plan = plan; + launch.coldBase = coldBase; + launch.stream = stream; +} + +} // namespace tensorrt_llm::kernels diff --git a/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp deleted file mode 100644 index e41ad149c26e..000000000000 --- a/cpp/tests/unit_tests/kv_cache_compression/nvfp4ColdPageCodecTest.cpp +++ /dev/null @@ -1,584 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" - -#include - -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ -namespace -{ - -namespace kv = batch_manager::kv_cache_manager_v2; - -static_assert(std::is_base_of_v); - -struct RecordedLaunch -{ - int offloadCalls = 0; - int onboardCalls = 0; - std::vector offloadPages; - std::vector onboardPages; - kernels::Nvfp4BoundaryPreparedPlan plan; - void const* coldBase = nullptr; - cudaStream_t stream{}; -}; - -RecordedLaunch gLaunch; - -constexpr std::uintptr_t kGpuKBase = 0x100000; -constexpr std::uintptr_t kGpuVBase = 0x200000; -constexpr std::uintptr_t kColdBase = 0x300000; -constexpr std::size_t kLayerRawBytes = 320; -constexpr std::size_t kLayerColdBytesAligned = 192; -constexpr std::size_t kNumAttentionLayers = 8; -constexpr std::size_t kGpuSlotBytes = kNumAttentionLayers * kLayerRawBytes; -constexpr std::size_t kColdSlotBytes = kNumAttentionLayers * kLayerColdBytesAligned; -constexpr std::size_t kMlaRawBytes = 64U * 576U * 2U; -constexpr std::size_t kMlaColdBytes = 64U * 576U / 2U + 64U * 576U / 16U; -constexpr std::size_t kDsaIndexKeyBytes = 64U * (128U + 4U); -constexpr std::uintptr_t kStreamValue = 0x7000; - -void resetLaunch() -{ - gLaunch = {}; -} - -std::vector makeLayers(std::size_t count = kNumAttentionLayers, int firstLayer = 0) -{ - std::vector layers; - layers.reserve(count); - for (std::size_t index = 0; index < count; ++index) - { - auto const scale = static_cast(index + 2U); - layers.push_back({firstLayer + static_cast(index), kernels::Nvfp4BoundaryRuntimeType::kFloat16, 1, 5, 32, - {scale, scale + 0.5F}, {1.0F / scale, 1.0F / (scale + 0.5F)}}); - } - return layers; -} - -std::vector makeMlaLayers(std::size_t count) -{ - std::vector layers; - layers.reserve(count); - for (std::size_t layer = 0; layer < count; ++layer) - { - layers.push_back({static_cast(layer), kernels::Nvfp4BoundaryRuntimeType::kFloat16, 1, 64, 576, - {1.0F, 1.0F}, {1.0F, 1.0F}}); - } - return layers; -} - -kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::PoolGroupIndex{0}, - kv::LayerGroupId lifeCycle = kv::LayerGroupId{3}, std::size_t count = kNumAttentionLayers, int firstLayer = 0, - std::uintptr_t keyBase = kGpuKBase, std::uintptr_t valueBase = kGpuVBase) -{ - kv::CoalescedBuffer keys{kLayerRawBytes, {}}; - kv::CoalescedBuffer values{kLayerRawBytes, {}}; - for (std::size_t index = 0; index < count; ++index) - { - auto const layerId = firstLayer + static_cast(index); - keys.bufferIds.push_back({layerId, "key"}); - values.bufferIds.push_back({layerId, "value"}); - } - - kv::SlotDescVariant variant{ - lifeCycle, kv::TypedVec{std::move(keys), std::move(values)}}; - auto const slotBytes = count * kLayerRawBytes; - return kv::PoolGroupDesc{poolGroupIndex, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, - kv::TypedVec{ - kv::PoolDesc{kv::PoolIndex{0}, keyBase, slotBytes}, kv::PoolDesc{kv::PoolIndex{1}, valueBase, slotBytes}}}; -} - -kv::PoolGroupDesc makeMlaDesc(std::vector const& ownsIndexer, kv::LayerGroupId lifeCycle = kv::LayerGroupId{0}, - int firstLayer = 0, std::size_t keyBytes = kLayerRawBytes, std::size_t indexBytes = 68U, - std::uintptr_t keyBase = kGpuKBase, std::uintptr_t indexBase = kGpuVBase) -{ - kv::CoalescedBuffer keys{keyBytes, {}}; - kv::CoalescedBuffer indexes{indexBytes, {}}; - for (std::size_t index = 0; index < ownsIndexer.size(); ++index) - { - auto const layerId = firstLayer + static_cast(index); - keys.bufferIds.push_back({layerId, "key"}); - if (ownsIndexer[index]) - { - indexes.bufferIds.push_back({layerId, "index_key"}); - } - } - - kv::TypedVec buffers; - buffers.push_back(std::move(keys)); - if (!indexes.bufferIds.empty()) - { - buffers.push_back(std::move(indexes)); - } - kv::SlotDescVariant variant{lifeCycle, std::move(buffers)}; - - kv::TypedVec pools; - pools.push_back(kv::PoolDesc{kv::PoolIndex{0}, keyBase, ownsIndexer.size() * keyBytes}); - auto const indexCount = static_cast(std::count(ownsIndexer.begin(), ownsIndexer.end(), true)); - if (indexCount != 0U) - { - pools.push_back(kv::PoolDesc{kv::PoolIndex{1}, indexBase, indexCount * indexBytes}); - } - return kv::PoolGroupDesc{ - kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, std::move(pools)}; -} - -bool configureOne(Nvfp4ColdPageCodec& codec, kv::PoolGroupDesc const& desc) -{ - return codec.configure(&desc, kv::PoolGroupIndex{1}); -} - -std::unique_ptr makeConfiguredAttentionCodec( - std::size_t count = kNumAttentionLayers, kv::LayerGroupId lifeCycle = kv::LayerGroupId{3}) -{ - auto codec = std::make_unique(makeLayers(count)); - EXPECT_TRUE(configureOne(*codec, makeAttentionDesc(kv::PoolGroupIndex{0}, lifeCycle, count))); - return codec; -} - -TEST(Nvfp4ColdPageCodecTest, OneCompletePageTaskCoversAllLayersWithDistinctScales) -{ - resetLaunch(); - auto codec = makeConfiguredAttentionCodec(); - EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{3}), kColdSlotBytes); - EXPECT_EQ(codec->queryPageIndexLocation(kv::LayerGroupId{3}), kv::PageIndexLocation::kHost); - - kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; - auto const stream = reinterpret_cast(kStreamValue); - ASSERT_TRUE( - codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - - EXPECT_EQ(gLaunch.offloadCalls, 1); - ASSERT_EQ(gLaunch.offloadPages.size(), 2U); - EXPECT_EQ(gLaunch.offloadPages[0].gpuPageIndex, 1); - EXPECT_EQ(gLaunch.offloadPages[0].coldPageIndex, 2); - EXPECT_EQ(gLaunch.offloadPages[1].gpuPageIndex, 3); - EXPECT_EQ(gLaunch.offloadPages[1].coldPageIndex, 5); - ASSERT_EQ(gLaunch.plan.numBuffers, 2U * kNumAttentionLayers); - for (std::size_t layer = 0; layer < kNumAttentionLayers; ++layer) - { - auto const& key = gLaunch.plan.buffers[2U * layer]; - auto const& value = gLaunch.plan.buffers[2U * layer + 1U]; - auto const layerOffset = layer * kLayerColdBytesAligned; - EXPECT_EQ(key.rawBase, kGpuKBase + layer * kLayerRawBytes); - EXPECT_EQ(value.rawBase, kGpuVBase + layer * kLayerRawBytes); - EXPECT_EQ(key.rawSlotBytes, kGpuSlotBytes); - EXPECT_EQ(value.rawSlotBytes, kGpuSlotBytes); - EXPECT_EQ(key.rawBytes, kLayerRawBytes); - EXPECT_EQ(value.rawBytes, kLayerRawBytes); - EXPECT_EQ(key.coldDataOffset, layerOffset); - EXPECT_EQ(value.coldDataOffset, layerOffset + 80U); - EXPECT_EQ(key.coldScaleOffset, layerOffset + 160U); - EXPECT_EQ(value.coldScaleOffset, layerOffset + 170U); - EXPECT_EQ(value.coldPaddingOffset, layerOffset + 180U); - EXPECT_EQ(value.coldPaddingBytes, 12U); - EXPECT_EQ(key.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); - EXPECT_EQ(value.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); - EXPECT_EQ(key.params.tokensPerPage, 5); - EXPECT_EQ(key.params.headDim, 32); - EXPECT_FLOAT_EQ(key.params.nvfp4ScaleOrigQuant, static_cast(layer + 2U)); - EXPECT_FLOAT_EQ(value.params.nvfp4ScaleOrigQuant, static_cast(layer + 2U) + 0.5F); - } - EXPECT_EQ(gLaunch.coldBase, reinterpret_cast(kColdBase)); - EXPECT_EQ(gLaunch.plan.coldPageBytes, kColdSlotBytes); - EXPECT_EQ(gLaunch.stream, stream); - - ASSERT_TRUE(codec->decode( - kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - EXPECT_EQ(gLaunch.onboardCalls, 1); - ASSERT_EQ(gLaunch.onboardPages.size(), 2U); - EXPECT_EQ(gLaunch.onboardPages[0].gpuPageIndex, 2); - EXPECT_EQ(gLaunch.onboardPages[0].coldPageIndex, 1); -} - -TEST(Nvfp4ColdPageCodecTest, KeyOnlyMlaUsesLatentPackedThenScaleLayout) -{ - resetLaunch(); - Nvfp4ColdPageCodec codec{makeLayers(1)}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc({false}))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 96U); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 1U); - auto const& latent = gLaunch.plan.buffers[0]; - EXPECT_EQ(latent.rawBase, kGpuKBase); - EXPECT_EQ(latent.rawSlotBytes, kLayerRawBytes); - EXPECT_EQ(latent.rawBytes, kLayerRawBytes); - EXPECT_EQ(latent.coldDataOffset, 0U); - EXPECT_EQ(latent.coldScaleOffset, 80U); - EXPECT_EQ(latent.coldPaddingOffset, 90U); - EXPECT_EQ(latent.coldPaddingBytes, 6U); - EXPECT_EQ(latent.transform, kernels::Nvfp4BoundaryTransform::kNvfp4); -} - -TEST(Nvfp4ColdPageCodecTest, KeyAndIndexAppendsLosslessIndexWithinTheLayerRecord) -{ - resetLaunch(); - Nvfp4ColdPageCodec codec{makeLayers(1)}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc({true}))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 160U); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 2U); - auto const& latent = gLaunch.plan.buffers[0]; - auto const& index = gLaunch.plan.buffers[1]; - EXPECT_EQ(latent.coldDataOffset, 0U); - EXPECT_EQ(latent.coldScaleOffset, 80U); - EXPECT_EQ(index.rawBase, kGpuVBase); - EXPECT_EQ(index.rawSlotBytes, 68U); - EXPECT_EQ(index.rawBytes, 68U); - EXPECT_EQ(index.coldDataOffset, 90U); - EXPECT_EQ(index.coldPaddingOffset, 158U); - EXPECT_EQ(index.coldPaddingBytes, 2U); - EXPECT_EQ(index.transform, kernels::Nvfp4BoundaryTransform::kLosslessCopy); -} - -TEST(Nvfp4ColdPageCodecTest, FullAndSharedIndexerLayersHaveDistinctPerLayerRecords) -{ - resetLaunch(); - Nvfp4ColdPageCodec codec{makeLayers(3)}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc({true, false, true}))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 416U); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 5U); - EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); - EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kLayerRawBytes); - EXPECT_EQ(gLaunch.plan.buffers[3].rawBase, kGpuKBase + 2U * kLayerRawBytes); - EXPECT_EQ(gLaunch.plan.buffers[4].rawBase, kGpuVBase + 68U); - EXPECT_EQ(gLaunch.plan.buffers[0].rawSlotBytes, 3U * kLayerRawBytes); - EXPECT_EQ(gLaunch.plan.buffers[1].rawSlotBytes, 2U * 68U); - EXPECT_EQ(gLaunch.plan.buffers[0].coldDataOffset, 0U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); - EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, 160U); - EXPECT_EQ(gLaunch.plan.buffers[3].coldDataOffset, 256U); - EXPECT_EQ(gLaunch.plan.buffers[4].coldDataOffset, 346U); - EXPECT_EQ(gLaunch.plan.buffers[4].coldPaddingOffset, 414U); - EXPECT_EQ(gLaunch.plan.buffers[4].coldPaddingBytes, 2U); -} - -TEST(Nvfp4ColdPageCodecTest, DeepSeekV32AllIndexerLayoutFitsOneColdPagePlan) -{ - resetLaunch(); - std::vector const ownsIndexer(61, true); - Nvfp4ColdPageCodec codec{makeMlaLayers(ownsIndexer.size())}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc(ownsIndexer, kv::LayerGroupId{0}, 0, kMlaRawBytes, kDsaIndexKeyBytes))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 61U * (kMlaColdBytes + kDsaIndexKeyBytes)); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - EXPECT_EQ(gLaunch.plan.numBuffers, 122U); - EXPECT_EQ(gLaunch.plan.coldPageBytes, 1780224U); -} - -TEST(Nvfp4ColdPageCodecTest, Glm52MixedIndexerLayoutFitsOneColdPagePlan) -{ - resetLaunch(); - std::vector ownsIndexer(78, false); - for (std::size_t layer = 0; layer < ownsIndexer.size(); ++layer) - { - ownsIndexer[layer] = layer < 3U || (layer >= 6U && layer % 4U == 2U); - } - ASSERT_EQ(std::count(ownsIndexer.begin(), ownsIndexer.end(), true), 21); - - Nvfp4ColdPageCodec codec{makeMlaLayers(ownsIndexer.size())}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc(ownsIndexer, kv::LayerGroupId{0}, 0, kMlaRawBytes, kDsaIndexKeyBytes))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 78U * kMlaColdBytes + 21U * kDsaIndexKeyBytes); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - EXPECT_EQ(gLaunch.plan.numBuffers, 99U); - EXPECT_EQ(gLaunch.plan.coldPageBytes, 1794816U); -} - -TEST(Nvfp4ColdPageCodecTest, PreservesOneCodecSubmissionAcrossThe256PageKernelBoundary) -{ - resetLaunch(); - auto codec = makeConfiguredAttentionCodec(); - - std::vector indices(257); - for (std::size_t page = 0; page < indices.size(); ++page) - { - indices[page] = {static_cast(500U - page), static_cast(page * 2U)}; - } - ASSERT_TRUE(codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices.data(), indices.size(), - reinterpret_cast(kStreamValue))); - EXPECT_EQ(gLaunch.offloadCalls, 1); - EXPECT_EQ(gLaunch.offloadPages.size(), 257U); - EXPECT_EQ(gLaunch.plan.numBuffers, 2U * kNumAttentionLayers); -} - -TEST(Nvfp4ColdPageCodecTest, EmptyAttentionBatchIsValidAndDoesNotLaunch) -{ - resetLaunch(); - auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); - - EXPECT_TRUE(codec->encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); - EXPECT_TRUE(codec->decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); - EXPECT_EQ(gLaunch.offloadCalls, 0); - EXPECT_EQ(gLaunch.onboardCalls, 0); -} - -TEST(Nvfp4ColdPageCodecTest, DefaultStreamIsAccepted) -{ - resetLaunch(); - auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); - - kv::PageIndexPair const indices[]{{0, 0}}; - EXPECT_TRUE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - EXPECT_TRUE(codec->decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - EXPECT_EQ(gLaunch.offloadCalls, 1); - EXPECT_EQ(gLaunch.onboardCalls, 1); -} - -TEST(Nvfp4ColdPageCodecTest, NonEmptyAttentionBatchRequiresPageIndices) -{ - auto codec = makeConfiguredAttentionCodec(1, kv::LayerGroupId{0}); - - auto const stream = reinterpret_cast(kStreamValue); - EXPECT_FALSE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), nullptr, 1U, stream)); - EXPECT_FALSE(codec->decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), nullptr, 1U, stream)); -} - -TEST(Nvfp4ColdPageCodecTest, OnlyFp8RuntimeRequiresFp8Scales) -{ - auto layers = makeLayers(1); - layers.front().fp8ScaleOrigQuant = {0.0F, 0.0F}; - layers.front().fp8ScaleQuantOrig = {0.0F, 0.0F}; - EXPECT_NO_THROW({ Nvfp4ColdPageCodec codec{layers}; }); - - layers.front().runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3; - EXPECT_THROW({ Nvfp4ColdPageCodec codec{layers}; }, std::invalid_argument); -} - -TEST(Nvfp4ColdPageCodecTest, Fp8SourceScalesDefaultToIdentity) -{ - auto layers = makeLayers(1); - layers.front().runtimeType = kernels::Nvfp4BoundaryRuntimeType::kFp8E4m3; - - EXPECT_EQ(layers.front().fp8ScaleOrigQuant, (std::array{1.0F, 1.0F})); - EXPECT_EQ(layers.front().fp8ScaleQuantOrig, (std::array{1.0F, 1.0F})); - EXPECT_NO_THROW({ Nvfp4ColdPageCodec codec{layers}; }); -} - -TEST(Nvfp4ColdPageCodecTest, DiscoversLifecycleMembershipAcrossPoolGroups) -{ - auto layers = makeLayers(2); - auto secondGroupLayers = makeLayers(2, 2); - layers.insert(layers.end(), secondGroupLayers.begin(), secondGroupLayers.end()); - Nvfp4ColdPageCodec codec{layers}; - std::array descs{makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2), - makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1}, 2, 2, 0x400000, 0x500000)}; - ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 2U * kLayerColdBytesAligned); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kLayerColdBytesAligned); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{0}), kv::LayerGroupId{0}); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); -} - -TEST(Nvfp4ColdPageCodecTest, RejectsConfiguredAttentionLayerAbsentFromAllGpuDescriptors) -{ - Nvfp4ColdPageCodec codec{makeLayers(2)}; - EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1))); -} - -TEST(Nvfp4ColdPageCodecTest, CoalescedAttentionSideBufferUsesItsOwnBaseOffsetAndSlotStride) -{ - resetLaunch(); - Nvfp4ColdPageCodec codec{makeLayers(1)}; - auto desc = makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1); - auto& keys = desc.slotDesc.variants.front().coalescedBuffers[kv::PoolIndex{0}]; - keys.bufferIds.push_back({0, "index_key"}); - desc.pools[kv::PoolIndex{0}].slotBytes += keys.singleBufferSize; - - ASSERT_TRUE(configureOne(codec, desc)); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 512U); - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 3U); - auto const& side = gLaunch.plan.buffers[2]; - EXPECT_EQ(side.rawBase, kGpuKBase + kLayerRawBytes); - EXPECT_EQ(side.rawSlotBytes, 2U * kLayerRawBytes); - EXPECT_EQ(side.rawBytes, kLayerRawBytes); - EXPECT_EQ(side.coldDataOffset, 180U); - EXPECT_EQ(side.coldPaddingOffset, 500U); - EXPECT_EQ(side.coldPaddingBytes, 12U); - EXPECT_EQ(side.transform, kernels::Nvfp4BoundaryTransform::kLosslessCopy); -} - -TEST(Nvfp4ColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) -{ - auto codec = makeConfiguredAttentionCodec(); - EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{99}), 0U); - EXPECT_EQ(codec->getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); - EXPECT_EQ(codec->queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); -} - -TEST(Nvfp4ColdPageCodecTest, NonAttentionLifecycleUsesLosslessSingleBlob) -{ - resetLaunch(); - int deviceCount = 0; - if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) - { - GTEST_SKIP() << "CUDA device is required for the lossless copy data-plane test"; - } - - constexpr std::size_t kPoolBytes = 64; - constexpr std::size_t kSlots = 2; - std::byte* statePool = nullptr; - std::byte* convPool = nullptr; - std::byte* coldPool = nullptr; - cudaStream_t stream = nullptr; - ASSERT_EQ(cudaMalloc(reinterpret_cast(&statePool), kPoolBytes * kSlots), cudaSuccess); - ASSERT_EQ(cudaMalloc(reinterpret_cast(&convPool), kPoolBytes * kSlots), cudaSuccess); - ASSERT_EQ(cudaMalloc(reinterpret_cast(&coldPool), 2U * kPoolBytes * kSlots), cudaSuccess); - ASSERT_EQ(cudaStreamCreate(&stream), cudaSuccess); - - kv::SlotDescVariant variant; - variant.lifeCycleId = kv::LayerGroupId{1}; - variant.coalescedBuffers = kv::TypedVec{ - kv::CoalescedBuffer{kPoolBytes, {{10, "ssm_state"}}}, - kv::CoalescedBuffer{kPoolBytes, {{10, "conv_state"}}}, - }; - kv::PoolGroupDesc desc; - desc.poolGroupIndex = kv::PoolGroupIndex{1}; - desc.numSlots = kSlots; - desc.slotDesc.variants = {variant}; - desc.pools = kv::TypedVec{ - kv::PoolDesc{kv::PoolIndex{0}, reinterpret_cast(statePool), kPoolBytes}, - kv::PoolDesc{kv::PoolIndex{1}, reinterpret_cast(convPool), kPoolBytes}, - }; - - Nvfp4ColdPageCodec codec{makeLayers(1)}; - std::array descs{makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1), desc}; - ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kPoolBytes); - - std::vector state(kPoolBytes); - std::vector conv(kPoolBytes); - for (std::size_t index = 0; index < kPoolBytes; ++index) - { - state[index] = static_cast(index + 1U); - conv[index] = static_cast(index + 65U); - } - ASSERT_EQ(cudaMemcpy(statePool + kPoolBytes, state.data(), kPoolBytes, cudaMemcpyHostToDevice), cudaSuccess); - ASSERT_EQ(cudaMemcpy(convPool + kPoolBytes, conv.data(), kPoolBytes, cudaMemcpyHostToDevice), cudaSuccess); - - kv::PageIndexPair const encodePair{0, 1}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{1}, coldPool, &encodePair, 1U, stream)); - ASSERT_EQ(cudaMemsetAsync(statePool, 0, kPoolBytes, stream), cudaSuccess); - ASSERT_EQ(cudaMemsetAsync(convPool, 0, kPoolBytes, stream), cudaSuccess); - kv::PageIndexPair const decodePair{0, 0}; - ASSERT_TRUE(codec.decode(kv::LayerGroupId{1}, coldPool, &decodePair, 1U, stream)); - ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); - - std::vector restoredState(kPoolBytes); - std::vector restoredConv(kPoolBytes); - ASSERT_EQ(cudaMemcpy(restoredState.data(), statePool, kPoolBytes, cudaMemcpyDeviceToHost), cudaSuccess); - ASSERT_EQ(cudaMemcpy(restoredConv.data(), convPool, kPoolBytes, cudaMemcpyDeviceToHost), cudaSuccess); - EXPECT_EQ(restoredState, state); - EXPECT_EQ(restoredConv, conv); - EXPECT_EQ(gLaunch.offloadCalls, 0); - - EXPECT_EQ(cudaStreamDestroy(stream), cudaSuccess); - EXPECT_EQ(cudaFree(coldPool), cudaSuccess); - EXPECT_EQ(cudaFree(convPool), cudaSuccess); - EXPECT_EQ(cudaFree(statePool), cudaSuccess); -} - -TEST(Nvfp4ColdPageCodecTest, AttentionAndSsmSharingOneHotPoolGroupUseDifferentTransforms) -{ - Nvfp4ColdPageCodec codec{makeLayers(1)}; - auto desc = makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1); - - kv::SlotDescVariant ssm; - ssm.lifeCycleId = kv::LayerGroupId{1}; - ssm.coalescedBuffers = kv::TypedVec{ - kv::CoalescedBuffer{kLayerRawBytes, {{10, "ssm_state"}}}, - kv::CoalescedBuffer{kLayerRawBytes, {{10, "conv_state"}}}, - }; - desc.slotDesc.variants.push_back(std::move(ssm)); - - ASSERT_TRUE(configureOne(codec, desc)); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), kLayerColdBytesAligned); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 2U * kLayerRawBytes); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{0}), kv::LayerGroupId{0}); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); -} - -TEST(Nvfp4ColdPageCodecTest, DuplicateLifecycleAcrossPoolGroupsIsRejected) -{ - Nvfp4ColdPageCodec codec{makeLayers(2)}; - std::array descs{ - makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 1, 0), - makeAttentionDesc( - kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1, 1, kGpuKBase + kGpuSlotBytes, kGpuVBase + kGpuSlotBytes), - }; - - EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); -} - -} // namespace -} // namespace tensorrt_llm::kv_cache_compression - -namespace tensorrt_llm::kernels -{ - -Nvfp4BoundaryPreparedPlan prepareNvfp4BoundaryPlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4BoundaryRuntimeType runtimeType) -{ - if (buffers.empty() || buffers.size() > kNvfp4BoundaryMaxBuffersPerLaunch || coldPageBytes == 0U) - { - throw std::invalid_argument("invalid test launch plan"); - } - Nvfp4BoundaryPreparedPlan plan; - std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); - plan.numBuffers = static_cast(buffers.size()); - plan.coldPageBytes = coldPageBytes; - plan.runtimeType = runtimeType; - return plan; -} - -void invokeNvfp4BoundaryOffloadCompress(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void* coldBase, cudaStream_t stream) -{ - auto& launch = kv_cache_compression::gLaunch; - ++launch.offloadCalls; - launch.offloadPages = pages; - launch.plan = plan; - launch.coldBase = coldBase; - launch.stream = stream; -} - -void invokeNvfp4BoundaryOnboardDecompress(std::vector const& pages, - Nvfp4BoundaryPreparedPlan const& plan, void const* coldBase, cudaStream_t stream) -{ - auto& launch = kv_cache_compression::gLaunch; - ++launch.onboardCalls; - launch.onboardPages = pages; - launch.plan = plan; - launch.coldBase = coldBase; - launch.stream = stream; -} - -} // namespace tensorrt_llm::kernels diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py new file mode 100644 index 000000000000..49ac51185822 --- /dev/null +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py @@ -0,0 +1,230 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""NVFP4 policy and cold-page layout construction.""" + +import json +import math +import os +import re +from pathlib import Path +from typing import Sequence + +from tensorrt_llm.quantization.modelopt_config import ( + is_modelopt_quant_config, + read_modelopt_quant_config, +) + +from ...pyexecutor.resource_manager import DataType + +ScalePair = tuple[float, float] +LayerScales = tuple[ScalePair, ScalePair] + +_IDENTITY_NVFP4_SCALES: LayerScales = ((1.0, 1.0), (1.0, 1.0)) +_MODEL_OPT_LANGUAGE_KV_SCALE_KEY = re.compile( + r"^model(?:\.language_model)?\.layers\.(?P\d+)\.self_attn\." + r"(?P[kv])_proj\.(?P=kind)_scale$" +) +_COLD_PAGE_ALIGNMENT = 16 +_ELEMENTS_PER_BYTE = 2 +_ELEMENTS_PER_SCALE = 16 + + +def _load_modelopt_nvfp4_scales( + checkpoint_path: str | None, +) -> dict[int, LayerScales]: + """Load optional ModelOpt NVFP4 K/V global scales by model layer.""" + + if checkpoint_path is None or os.environ.get("TRTLLM_LOAD_KV_SCALES", "1") != "1": + return {} + + checkpoint_dir = Path(checkpoint_path) + weight_files = sorted(checkpoint_dir.glob("*.safetensors")) + ordinary_files = [path for path in weight_files if "consolidated" not in path.name] + weight_files = ordinary_files or weight_files + if not weight_files: + raise FileNotFoundError( + f"No safetensors files in ModelOpt scale checkpoint {checkpoint_dir}" + ) + + metadata_path = checkpoint_dir / "hf_quant_config.json" + if metadata_path.exists(): + metadata = json.loads(metadata_path.read_text()) + else: + config_path = checkpoint_dir / "config.json" + metadata = ( + json.loads(config_path.read_text()).get("quantization_config") + if config_path.exists() + else None + ) + if not is_modelopt_quant_config(metadata): + return {} + if read_modelopt_quant_config(metadata).get("kv_cache_quant_algo") != "NVFP4": + return {} + + from safetensors import safe_open + + values: dict[int, dict[str, list[float]]] = {} + for file_path in weight_files: + with safe_open(str(file_path), framework="pt", device="cpu") as checkpoint: + for tensor_name in checkpoint.keys(): + match = _MODEL_OPT_LANGUAGE_KV_SCALE_KEY.fullmatch(tensor_name) + if match is None: + continue + value = float(checkpoint.get_tensor(tensor_name).reshape([]).item()) + if not math.isfinite(value) or value <= 0.0: + raise ValueError( + f"ModelOpt KV scale {file_path}:{tensor_name} must be finite and positive" + ) + layer_values = values.setdefault(int(match.group("layer_id")), {"k": [], "v": []}) + layer_values[match.group("kind")].append(value) + + result: dict[int, LayerScales] = {} + for layer_id, layer_values in values.items(): + k_values, v_values = layer_values["k"], layer_values["v"] + if not k_values or not v_values: + raise ValueError(f"ModelOpt KV scales for layer {layer_id} must contain both K and V") + quant_orig = (max(k_values), max(v_values)) + result[layer_id] = ( + (1.0 / quant_orig[0], 1.0 / quant_orig[1]), + quant_orig, + ) + return result + + +def _align_up(value: int, alignment: int = _COLD_PAGE_ALIGNMENT) -> int: + return (value + alignment - 1) // alignment * alignment + + +def _buffer_bytes(buffer: object, tokens_per_page: int) -> int: + buffer_tokens = buffer.tokens_per_block_override or tokens_per_page + if buffer_tokens <= 0 or tokens_per_page % buffer_tokens != 0: + raise ValueError("tokens_per_block_override must be a positive divisor of tokens_per_block") + return int(buffer.size) * (tokens_per_page // buffer_tokens) + + +class Nvfp4ColdPagePolicy: + """Build native NVFP4 plans without retaining Python in the data path.""" + + def __init__(self, checkpoint_path: str | None) -> None: + self._model_scales = _load_modelopt_nvfp4_scales(checkpoint_path) + + def create_cold_page_codec( + self, + cache_config: object, + *, + runtime_dtype: DataType, + pp_layers: Sequence[int], + num_kv_heads_per_layer: Sequence[int], + head_dim_per_layer: Sequence[int], + is_draft: bool = False, + ) -> object: + """Create the native generic codec from explicit per-buffer plans.""" + + from tensorrt_llm.bindings.internal import kv_cache_compression as native + from tensorrt_llm.runtime.kv_cache_manager_v2 import AttentionLayerConfig + + attention_layers = [ + layer for layer in cache_config.layers if isinstance(layer, AttentionLayerConfig) + ] + if not attention_layers: + return native.create_cold_page_codec([]) + + runtime_type = { + DataType.HALF: native.Nvfp4ColdPageRuntimeType.FLOAT16, + DataType.BF16: native.Nvfp4ColdPageRuntimeType.BFLOAT16, + DataType.FP8: native.Nvfp4ColdPageRuntimeType.FP8_E4M3, + }.get(runtime_dtype) + if runtime_type is None: + raise RuntimeError( + "NVFP4 cold-page compression supports FP16, BF16, or FP8 " + f"Attention KV, not {runtime_dtype}" + ) + + layer_plans = [] + for layer in attention_layers: + layer_id = int(layer.layer_id) + buffers_by_role = {str(buffer.role): buffer for buffer in layer.buffers} + if "key" not in buffers_by_role: + raise NotImplementedError( + "NVFP4 cold-page compression requires an Attention key buffer" + ) + + compressed_roles = ("key", "value") if "value" in buffers_by_role else ("key",) + if len(compressed_roles) == 2 and not is_draft: + orig_quant, quant_orig = self._model_scales.get( + int(pp_layers[layer_id]), _IDENTITY_NVFP4_SCALES + ) + else: + # Target projection scales describe neither MLA latents nor a + # separately numbered draft model. + orig_quant, quant_orig = _IDENTITY_NVFP4_SCALES + + num_kv_heads = int(num_kv_heads_per_layer[layer_id]) + tokens_per_page = int(cache_config.tokens_per_block) + head_dim = int(head_dim_per_layer[layer_id]) + if head_dim <= 0 or head_dim % _ELEMENTS_PER_SCALE != 0: + raise ValueError( + f"NVFP4 cold pages require head_dim divisible by 16, got {head_dim}" + ) + elements = num_kv_heads * tokens_per_page * head_dim + packed_bytes = elements // _ELEMENTS_PER_BYTE + scale_bytes = elements // _ELEMENTS_PER_SCALE + raw_bytes_by_role = { + role: _buffer_bytes(buffer, tokens_per_page) + for role, buffer in buffers_by_role.items() + } + + data_offsets = { + role: index * packed_bytes for index, role in enumerate(compressed_roles) + } + scale_base = len(compressed_roles) * packed_bytes + scale_offsets = { + role: scale_base + index * scale_bytes + for index, role in enumerate(compressed_roles) + } + cursor = scale_base + len(compressed_roles) * scale_bytes + + buffer_plans = [] + for scale_index, role in enumerate(compressed_roles): + params = native.Nvfp4ColdPageParams() + params.runtime_type = runtime_type + params.num_kv_heads = num_kv_heads + params.tokens_per_page = tokens_per_page + params.head_dim = head_dim + params.nvfp4_scale_orig_quant = orig_quant[scale_index] + params.nvfp4_scale_quant_orig = quant_orig[scale_index] + params.fp8_scale_orig_quant = 1.0 + params.fp8_scale_quant_orig = 1.0 + + plan = native.ColdPageBufferPlan() + plan.role = role + plan.transform = native.ColdPageTransformKind.NVFP4 + plan.raw_bytes = raw_bytes_by_role[role] + plan.cold_data_offset = data_offsets[role] + plan.cold_scale_offset = scale_offsets[role] + plan.nvfp4_params = params + buffer_plans.append(plan) + + for buffer in layer.buffers: + role = str(buffer.role) + if role in compressed_roles: + continue + plan = native.ColdPageBufferPlan() + plan.role = role + plan.transform = native.ColdPageTransformKind.LOSSLESS_COPY + plan.raw_bytes = raw_bytes_by_role[role] + plan.cold_data_offset = cursor + plan.cold_scale_offset = 0 + buffer_plans.append(plan) + cursor += raw_bytes_by_role[role] + + cold_page_bytes = _align_up(cursor) + layer_plan = native.ColdPageLayerPlan() + layer_plan.layer_id = layer_id + layer_plan.cold_page_bytes = cold_page_bytes + layer_plan.cold_padding_offset = cursor + layer_plan.cold_padding_bytes = cold_page_bytes - cursor + layer_plan.buffers = buffer_plans + layer_plans.append(layer_plan) + + return native.create_cold_page_codec(layer_plans) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 3a23e157f8fa..3a58b7a73f69 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -1,116 +1,32 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""NVFP4 compression for KVCM V2 cold pages. +"""Quantization policies for KVCM V2 cold pages.""" -The compression manager owns optional ModelOpt K/V global scales and creates -the native storage-boundary codec before KVCM allocates cold Slots. Attention, -the active KV-cache format, and the normal model-loading path remain unaware of -the cold representation. -""" - -import json -import math -import os -import re -from pathlib import Path from typing import TYPE_CHECKING, Sequence -from tensorrt_llm.quantization.modelopt_config import ( - is_modelopt_quant_config, - read_modelopt_quant_config, -) - from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager +from .nvfp4 import Nvfp4ColdPagePolicy if TYPE_CHECKING: from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig -ScalePair = tuple[float, float] -LayerScales = tuple[ScalePair, ScalePair] - -_IDENTITY_NVFP4_SCALES: LayerScales = ((1.0, 1.0), (1.0, 1.0)) -_MODEL_OPT_LANGUAGE_KV_SCALE_KEY = re.compile( - r"^model(?:\.language_model)?\.layers\.(?P\d+)\.self_attn\." - r"(?P[kv])_proj\.(?P=kind)_scale$" -) - - -def _load_modelopt_nvfp4_scales( - checkpoint_path: str | None, -) -> dict[int, LayerScales]: - """Load optional ModelOpt NVFP4 K/V global scales by model layer. - - This mirrors the native NVFP4 KV loader's contract: no checkpoint means - identity scales, ``TRTLLM_LOAD_KV_SCALES`` can disable loading, and K/V - scalars from standard safetensors shards are reduced with ``max``. - """ - - if checkpoint_path is None or os.environ.get("TRTLLM_LOAD_KV_SCALES", "1") != "1": - return {} - - checkpoint_dir = Path(checkpoint_path) - weight_files = sorted(checkpoint_dir.glob("*.safetensors")) - ordinary_files = [path for path in weight_files if "consolidated" not in path.name] - weight_files = ordinary_files or weight_files - if not weight_files: - raise FileNotFoundError( - f"No safetensors files in ModelOpt scale checkpoint {checkpoint_dir}" - ) - - metadata_path = checkpoint_dir / "hf_quant_config.json" - if metadata_path.exists(): - metadata = json.loads(metadata_path.read_text()) - else: - config_path = checkpoint_dir / "config.json" - metadata = ( - json.loads(config_path.read_text()).get("quantization_config") - if config_path.exists() - else None - ) - if not is_modelopt_quant_config(metadata): - return {} - if read_modelopt_quant_config(metadata).get("kv_cache_quant_algo") != "NVFP4": - return {} - - from safetensors import safe_open - - values: dict[int, dict[str, list[float]]] = {} - for file_path in weight_files: - with safe_open(str(file_path), framework="pt", device="cpu") as checkpoint: - for tensor_name in checkpoint.keys(): - match = _MODEL_OPT_LANGUAGE_KV_SCALE_KEY.fullmatch(tensor_name) - if match is None: - continue - value = float(checkpoint.get_tensor(tensor_name).reshape([]).item()) - if not math.isfinite(value) or value <= 0.0: - raise ValueError( - f"ModelOpt KV scale {file_path}:{tensor_name} must be finite and positive" - ) - layer_values = values.setdefault(int(match.group("layer_id")), {"k": [], "v": []}) - layer_values[match.group("kind")].append(value) - - result: dict[int, LayerScales] = {} - for layer_id, layer_values in values.items(): - k_values, v_values = layer_values["k"], layer_values["v"] - if not k_values or not v_values: - raise ValueError(f"ModelOpt KV scales for layer {layer_id} must contain both K and V") - quant_orig = (max(k_values), max(v_values)) - result[layer_id] = ( - (1.0 / quant_orig[0], 1.0 / quant_orig[1]), - quant_orig, - ) - return result - class ColdPageQuantizationCompression(KVCacheCompressionManager): - """NVFP4 cold-page manager and owner of its optional model scales.""" + """Select and own the configured cold-page quantization policy.""" uses_iteration_lifecycle = False provides_cold_page_codec = True def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) - self._model_nvfp4_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) + policies = {"nvfp4": Nvfp4ColdPagePolicy} + try: + policy = policies[config.quant] + except KeyError as error: + raise NotImplementedError( + f"Unsupported cold-page quantization format {config.quant!r}" + ) from error + self._policy = policy(config.scale_checkpoint_path) def create_cold_page_codec( self, @@ -120,56 +36,15 @@ def create_cold_page_codec( pp_layers: Sequence[int], num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], + is_draft: bool = False, ) -> object: - """Create the native codec that KVCM consumes exactly once.""" - - from tensorrt_llm.bindings.internal import kv_cache_compression as native - from tensorrt_llm.runtime.kv_cache_manager_v2 import SsmLayerConfig - - attention_layers = [] - for layer in cache_config.layers: - if isinstance(layer, SsmLayerConfig): - continue - roles = {buffer.role for buffer in layer.buffers} - if "key" not in roles: - raise NotImplementedError( - "NVFP4 cold-page compression requires an Attention key buffer" - ) - has_value = "value" in roles - attention_layers.append((layer, has_value)) - if not attention_layers: - return native.create_nvfp4_cold_page_codec([]) - - runtime_type = { - DataType.HALF: native.Nvfp4BoundaryRuntimeType.FLOAT16, - DataType.BF16: native.Nvfp4BoundaryRuntimeType.BFLOAT16, - DataType.FP8: native.Nvfp4BoundaryRuntimeType.FP8_E4M3, - }.get(runtime_dtype) - if runtime_type is None: - raise RuntimeError( - "NVFP4 cold-page compression supports FP16, BF16, or FP8 " - f"Attention KV, not {runtime_dtype}" - ) - - native_configs = [] - for layer, has_value in attention_layers: - layer_id = int(layer.layer_id) - if has_value: - orig_quant, quant_orig = self._model_nvfp4_scales.get( - int(pp_layers[layer_id]), _IDENTITY_NVFP4_SCALES - ) - else: - # ModelOpt K/V projection scales do not apply to MLA latent buffers. - orig_quant, quant_orig = _IDENTITY_NVFP4_SCALES - native_config = native.Nvfp4ColdPageLayerConfig() - native_config.layer_id = layer_id - native_config.runtime_type = runtime_type - native_config.num_kv_heads = int(num_kv_heads_per_layer[layer_id]) - native_config.tokens_per_page = int(cache_config.tokens_per_block) - native_config.head_dim = int(head_dim_per_layer[layer_id]) - native_config.nvfp4_scale_orig_quant = orig_quant - native_config.nvfp4_scale_quant_orig = quant_orig - native_configs.append(native_config) - - # The C++ factory is required for unique_ptr ownership transfer into KVCM. - return native.create_nvfp4_cold_page_codec(native_configs) + """Create the native codec selected by the quantization policy.""" + + return self._policy.create_cold_page_codec( + cache_config, + runtime_dtype=runtime_dtype, + pp_layers=pp_layers, + num_kv_heads_per_layer=num_kv_heads_per_layer, + head_dim_per_layer=head_dim_per_layer, + is_draft=is_draft, + ) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 9e9a60fd0a4e..4c912b75a96f 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -3149,7 +3149,7 @@ def create_py_executor_instance( resource_manager.resource_managers.move_to_end( ResourceManagerType.CROSS_KV_CACHE_MANAGER, last=True) # Iteration-driven compression is the final reconciler after every native - # KV manager. Boundary quantization runs only at native storage migration. + # KV manager. Cold-page quantization runs only at native storage migration. if (ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER in resource_manager.resource_managers): resource_manager.resource_managers.move_to_end( diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 2e0e391cb48e..90650698f3d0 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1128,6 +1128,7 @@ def append_to_kv_heads_per_layer( pp_layers=self.pp_layers, num_kv_heads_per_layer=self.num_kv_heads_per_layer, head_dim_per_layer=self.head_dim_per_layer, + is_draft=self.is_draft, ) if cold_page_codec_provider is not None else None @@ -1156,6 +1157,7 @@ def append_to_kv_heads_per_layer( pp_layers=self.pp_layers, num_kv_heads_per_layer=self.num_kv_heads_per_layer, head_dim_per_layer=self.head_dim_per_layer, + is_draft=self.is_draft, ) if cold_page_codec_provider is not None else None diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index c8d189f82468..fc31cd9dd56f 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -2718,7 +2718,7 @@ class KVCacheCompressionManager(BaseResourceManager): """Framework base for KV-cache compression methods in PyExecutor. Iteration-driven methods receive ResourceManager callbacks, while - storage-boundary methods provide a cold-page codec during cache construction. + storage-bound methods provide a cold-page codec during cache construction. Subclasses coordinate through KVCacheManagerV2 without owning its pools, mappings, or migration lifecycle. """ @@ -2765,8 +2765,9 @@ def create_cold_page_codec( pp_layers: Sequence[int], num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], + is_draft: bool = False, ) -> Optional[object]: - """Create a native storage-boundary codec when the algorithm provides one.""" + """Create a native cold-page codec when the algorithm provides one.""" return None # ================================================================== # diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 116418e5ad72..62d9b1a162f0 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3685,7 +3685,7 @@ class KvCacheCompressionConfig(StrictBaseModel): Kept separate from SparseAttentionConfig by design -- compression changes which KV is stored, not the attention computation. Iteration-driven methods - use the resource-manager cycle; storage-boundary managers provide a native + use the resource-manager cycle; storage-bound managers provide a native codec that KVCacheManagerV2 retains and invokes. Concrete algorithms subclass this and add their parameters. """ diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index c6785a2fe7f7..bd82f07dc94c 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -111,6 +111,7 @@ def _make_manager_for_cache_tier_test( *, add_secondary_gpu_tier: bool = False, cold_page_codec_provider: object | None = None, + is_draft: bool = False, ) -> tuple[KVCacheManagerV2, Mock]: impl_constructor = Mock(side_effect=impl_side_effect) @@ -168,6 +169,7 @@ def build_cache_config( mapping=Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), dtype=DataType.HALF, vocab_size=16, + is_draft=is_draft, execution_stream=Mock(), cold_page_codec_provider=cold_page_codec_provider, ) @@ -400,6 +402,21 @@ def test_host_init_fallback_recreates_cold_codec_and_keeps_disk(tmp_path) -> Non ] +@pytest.mark.cpu_only +def test_cold_codec_provider_receives_draft_role() -> None: + impl = Mock() + codec_provider = Mock() + codec_provider.create_cold_page_codec.return_value = object() + _make_manager_for_cache_tier_test( + KvCacheConfig(max_gpu_total_bytes=16 << 20), + [impl], + cold_page_codec_provider=codec_provider, + is_draft=True, + ) + + assert codec_provider.create_cold_page_codec.call_args.kwargs["is_draft"] is True + + def test_extra_tokens_are_in_context_capacity() -> None: config = _make_cache_config_for_test( KvCacheConfig(avg_seq_len=264), diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 8e45b4da9c77..155258902f73 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -10,6 +10,9 @@ import torch from safetensors.torch import save_file +from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.nvfp4 import ( + _load_modelopt_nvfp4_scales, +) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import ( ColdPageQuantizationCompression, ) @@ -54,18 +57,27 @@ def _cache_config(*layers): def _native(): - def layer_config(): + def config(): return SimpleNamespace() + def buffer_plan(): + return SimpleNamespace(nvfp4_params=None) + codec = MagicMock() module = SimpleNamespace( - Nvfp4BoundaryRuntimeType=SimpleNamespace( + ColdPageTransformKind=SimpleNamespace( + NVFP4="native-nvfp4", + LOSSLESS_COPY="native-lossless", + ), + Nvfp4ColdPageRuntimeType=SimpleNamespace( FLOAT16="native-fp16", BFLOAT16="native-bf16", FP8_E4M3="native-fp8", ), - Nvfp4ColdPageLayerConfig=layer_config, - create_nvfp4_cold_page_codec=MagicMock(return_value=codec), + Nvfp4ColdPageParams=config, + ColdPageBufferPlan=buffer_plan, + ColdPageLayerPlan=config, + create_cold_page_codec=MagicMock(return_value=codec), ) return module, codec @@ -97,7 +109,7 @@ def _validate_compression(mode=None): ) -def test_optional_modelopt_scales_map_pp_layers_and_ignore_local_draft_id(tmp_path): +def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_path): native, codec = _native() _write_scales( tmp_path, @@ -121,11 +133,17 @@ def test_optional_modelopt_scales_map_pp_layers_and_ignore_local_draft_id(tmp_pa ) assert result is codec - configs = native.create_nvfp4_cold_page_codec.call_args.args[0] - assert [config.layer_id for config in configs] == [0, 1, 2] - assert [config.runtime_type for config in configs] == ["native-bf16"] * 3 + plans = native.create_cold_page_codec.call_args.args[0] + assert [plan.layer_id for plan in plans] == [0, 1, 2] + assert [buffer.nvfp4_params.runtime_type for plan in plans for buffer in plan.buffers] == [ + "native-bf16" + ] * 6 assert [ - (config.nvfp4_scale_orig_quant, config.nvfp4_scale_quant_orig) for config in configs + ( + tuple(buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers), + tuple(buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers), + ) + for plan in plans ] == [ ((2.0, 4.0), (0.5, 0.25)), ((8.0, 16.0), (0.125, 0.0625)), @@ -133,6 +151,44 @@ def test_optional_modelopt_scales_map_pp_layers_and_ignore_local_draft_id(tmp_pa ] +def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: + native, _ = _native() + _write_scales(tmp_path, {10: (0.5, 0.25)}) + manager = _manager(tmp_path) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + manager.create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(10,), + num_kv_heads_per_layer=(8,), + head_dim_per_layer=(128,), + ) + target_plan = native.create_cold_page_codec.call_args.args[0][0] + manager.create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(10,), + num_kv_heads_per_layer=(8,), + head_dim_per_layer=(128,), + is_draft=True, + ) + + draft_plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in target_plan.buffers] == [ + 2.0, + 4.0, + ] + assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in draft_plan.buffers] == [ + 1.0, + 1.0, + ] + assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in draft_plan.buffers] == [ + 1.0, + 1.0, + ] + + def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): native, _ = _native() cache_config = _cache_config((0, "attention")) @@ -147,18 +203,64 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): head_dim_per_layer=(128,), ) - config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert config.runtime_type == "native-fp16" - assert config.num_kv_heads == 4 - assert config.tokens_per_page == 5 - assert config.head_dim == 128 - assert config.nvfp4_scale_orig_quant == config.nvfp4_scale_quant_orig == (1.0, 1.0) + plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.role for buffer in plan.buffers] == ["key", "value"] + assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 1280] + assert [buffer.cold_scale_offset for buffer in plan.buffers] == [2560, 2720] + assert plan.cold_page_bytes == 2880 + assert plan.cold_padding_bytes == 0 + params = plan.buffers[0].nvfp4_params + assert params.runtime_type == "native-fp16" + assert params.num_kv_heads == 4 + assert params.tokens_per_page == 5 + assert params.head_dim == 128 + assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + 1.0, + 1.0, + ] + assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + 1.0, + 1.0, + ] + + +def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: + native, _ = _native() + cache_config = SimpleNamespace( + tokens_per_block=5, + layers=( + AttentionLayerConfig( + layer_id=0, + buffers=[ + BufferConfig(role="key", size=320), + BufferConfig(role="value", size=320), + ], + ), + ), + ) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.HALF, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(32,), + ) + + plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.role for buffer in plan.buffers] == ["key", "value"] + assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 80] + assert [buffer.cold_scale_offset for buffer in plan.buffers] == [160, 170] + assert plan.cold_padding_offset == 180 + assert plan.cold_padding_bytes == 12 + assert plan.cold_page_bytes == 192 def test_provider_creates_one_native_codec_per_kv_cache_manager(): native, _ = _native() codecs = (object(), object()) - native.create_nvfp4_cold_page_codec.side_effect = codecs + native.create_cold_page_codec.side_effect = codecs provider = _manager() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): @@ -174,7 +276,21 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): ) assert results == codecs - assert native.create_nvfp4_cold_page_codec.call_count == 2 + assert native.create_cold_page_codec.call_count == 2 + + +def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: + config = SimpleNamespace( + quant="future-format", + scale_checkpoint_path="/not/a/checkpoint", + ) + with patch( + "tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page." + "quantization_for_cold_page.Nvfp4ColdPagePolicy" + ) as policy: + with pytest.raises(NotImplementedError, match="future-format"): + ColdPageQuantizationCompression(config) + policy.assert_not_called() def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): @@ -184,7 +300,7 @@ def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): {7: (0.125, 0.0625)}, filename="consolidated.00.safetensors", ) - assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + assert _load_modelopt_nvfp4_scales(str(tmp_path))[7] == ( (2.0, 4.0), (0.5, 0.25), ) @@ -196,7 +312,7 @@ def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): {9: (0.125, 0.0625)}, filename="consolidated.00.safetensors", ) - assert _manager(consolidated_only)._model_nvfp4_scales[9] == ( + assert _load_modelopt_nvfp4_scales(str(consolidated_only))[9] == ( (8.0, 16.0), (0.125, 0.0625), ) @@ -210,7 +326,7 @@ def test_scale_loader_reduces_duplicate_shards_like_native_qkv_loader(tmp_path): filename="model-00002.safetensors", prefix="model.language_model", ) - assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + assert _load_modelopt_nvfp4_scales(str(tmp_path))[7] == ( (2.0, 4.0), (0.5, 0.25), ) @@ -228,7 +344,7 @@ def test_scale_loader_ignores_multimodal_towers_with_the_same_layer_id(tmp_path) } save_file(tensors, str(tmp_path / "model.safetensors")) - assert _manager(tmp_path)._model_nvfp4_scales[7] == ( + assert _load_modelopt_nvfp4_scales(str(tmp_path))[7] == ( (2.0, 4.0), (0.5, 0.25), ) @@ -237,23 +353,23 @@ def test_scale_loader_ignores_multimodal_towers_with_the_same_layer_id(tmp_path) def test_trtllm_load_kv_scales_zero_uses_identity(tmp_path, monkeypatch): _write_scales(tmp_path, {7: (0.5, 0.25)}) monkeypatch.setenv("TRTLLM_LOAD_KV_SCALES", "0") - assert _manager(tmp_path)._model_nvfp4_scales == {} + assert _load_modelopt_nvfp4_scales(str(tmp_path)) == {} def test_non_nvfp4_checkpoint_scales_are_not_reused(tmp_path): _write_scales(tmp_path, {7: (0.5, 0.25)}) _write_quant_metadata(tmp_path, "FP8") - assert _manager(tmp_path)._model_nvfp4_scales == {} + assert _load_modelopt_nvfp4_scales(str(tmp_path)) == {} def test_unquantized_checkpoint_uses_identity_scales(tmp_path): save_file({"model.weight": torch.ones(1)}, str(tmp_path / "model.safetensors")) - assert _manager(tmp_path)._model_nvfp4_scales == {} + assert _load_modelopt_nvfp4_scales(str(tmp_path)) == {} def test_explicit_scale_checkpoint_requires_safetensors(tmp_path): with pytest.raises(FileNotFoundError, match="No safetensors files"): - _manager(tmp_path) + _load_modelopt_nvfp4_scales(str(tmp_path)) @pytest.mark.parametrize("present_kind", ["k", "v"]) @@ -263,7 +379,7 @@ def test_scale_checkpoint_requires_kv_pair(tmp_path, present_kind): name = f"{base}.{present_kind}_proj.{present_kind}_scale" save_file({name: torch.tensor(0.5)}, str(tmp_path / "model.safetensors")) with pytest.raises(ValueError, match="both K and V"): - _manager(tmp_path) + _load_modelopt_nvfp4_scales(str(tmp_path)) def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): @@ -277,9 +393,12 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): num_kv_heads_per_layer=(0, 8), head_dim_per_layer=(128, 128), ) - config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert config.layer_id == 1 - assert config.nvfp4_scale_orig_quant == (2.0, 4.0) + plan = native.create_cold_page_codec.call_args.args[0][0] + assert plan.layer_id == 1 + assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + 2.0, + 4.0, + ] result = _manager().create_cold_page_codec( _cache_config((0, "ssm")), @@ -290,7 +409,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - native.create_nvfp4_cold_page_codec.assert_called_with([]) + native.create_cold_page_codec.assert_called_with([]) def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -318,13 +437,177 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): ) assert result is codec - config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert config.layer_id == 0 - assert config.runtime_type == "native-bf16" - assert config.num_kv_heads == 1 - assert config.tokens_per_page == 64 - assert config.head_dim == 576 - assert config.nvfp4_scale_orig_quant == config.nvfp4_scale_quant_orig == (1.0, 1.0) + plan = native.create_cold_page_codec.call_args.args[0][0] + assert plan.layer_id == 0 + assert plan.cold_page_bytes == 29184 + assert plan.cold_padding_offset == 29184 + assert plan.cold_padding_bytes == 0 + assert [buffer.role for buffer in plan.buffers] == ["key", "index_key"] + assert [buffer.transform for buffer in plan.buffers] == [ + "native-nvfp4", + "native-lossless", + ] + assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 20736] + assert [buffer.cold_scale_offset for buffer in plan.buffers] == [18432, 0] + params = plan.buffers[0].nvfp4_params + assert params.runtime_type == "native-bf16" + assert params.num_kv_heads == 1 + assert params.tokens_per_page == 64 + assert params.head_dim == 576 + assert params.nvfp4_scale_orig_quant == params.nvfp4_scale_quant_orig == 1.0 + assert plan.buffers[1].nvfp4_params is None + + +def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: + native, _ = _native() + cache_config = SimpleNamespace( + tokens_per_block=5, + layers=( + AttentionLayerConfig( + layer_id=0, + buffers=[ + BufferConfig(role="key", size=320), + BufferConfig(role="index_key", size=68), + BufferConfig(role="rope_state", size=7), + ], + ), + ), + ) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(32,), + ) + + plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.role for buffer in plan.buffers] == [ + "key", + "index_key", + "rope_state", + ] + assert [buffer.transform for buffer in plan.buffers] == [ + "native-nvfp4", + "native-lossless", + "native-lossless", + ] + assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 90, 158] + assert plan.cold_padding_offset == 165 + assert plan.cold_padding_bytes == 11 + assert plan.cold_page_bytes == 176 + + +def test_tokens_per_block_override_expands_raw_and_lossless_bytes() -> None: + native, _ = _native() + cache_config = SimpleNamespace( + tokens_per_block=4, + layers=( + AttentionLayerConfig( + layer_id=0, + buffers=[ + BufferConfig(role="key", size=128), + BufferConfig( + role="index_key", + size=3, + tokens_per_block_override=2, + ), + ], + ), + ), + ) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(16,), + ) + + plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.raw_bytes for buffer in plan.buffers] == [128, 6] + assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 36] + assert plan.cold_padding_offset == 42 + assert plan.cold_padding_bytes == 6 + assert plan.cold_page_bytes == 48 + + +def test_tokens_per_block_override_must_divide_page_size() -> None: + native, _ = _native() + cache_config = SimpleNamespace( + tokens_per_block=5, + layers=( + AttentionLayerConfig( + layer_id=0, + buffers=[ + BufferConfig( + role="key", + size=64, + tokens_per_block_override=2, + ), + ], + ), + ), + ) + + with ( + patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native), + pytest.raises(ValueError, match="positive divisor"), + ): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(16,), + ) + + +@pytest.mark.parametrize( + ("owns_index", "expected_layers", "expected_buffers", "expected_bytes"), + [ + ([True] * 61, 61, 122, 1_780_224), + ( + [layer < 3 or (layer >= 6 and layer % 4 == 2) for layer in range(78)], + 78, + 99, + 1_794_816, + ), + ], + ids=("deepseek-v3.2", "glm-5.2"), +) +def test_mla_model_layouts_are_built_in_python( + owns_index: list[bool], + expected_layers: int, + expected_buffers: int, + expected_bytes: int, +) -> None: + native, _ = _native() + layers = [] + for layer_id, has_index in enumerate(owns_index): + buffers = [BufferConfig(role="key", size=64 * 576 * 2)] + if has_index: + buffers.append(BufferConfig(role="index_key", size=64 * (128 + 4))) + layers.append(AttentionLayerConfig(layer_id=layer_id, buffers=buffers)) + cache_config = SimpleNamespace(tokens_per_block=64, layers=tuple(layers)) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=tuple(range(expected_layers)), + num_kv_heads_per_layer=(1,) * expected_layers, + head_dim_per_layer=(576,) * expected_layers, + ) + + plans = native.create_cold_page_codec.call_args.args[0] + assert len(plans) == expected_layers + assert sum(len(plan.buffers) for plan in plans) == expected_buffers + assert sum(plan.cold_page_bytes for plan in plans) == expected_bytes def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): @@ -339,10 +622,23 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): head_dim_per_layer=(128,), ) - config = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert config.runtime_type == "native-fp8" - assert config.nvfp4_scale_orig_quant == (2.0, 4.0) - assert config.nvfp4_scale_quant_orig == (0.5, 0.25) + plan = native.create_cold_page_codec.call_args.args[0][0] + assert [buffer.nvfp4_params.runtime_type for buffer in plan.buffers] == [ + "native-fp8", + "native-fp8", + ] + assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + 2.0, + 4.0, + ] + assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + 0.5, + 0.25, + ] + assert all( + buffer.nvfp4_params.fp8_scale_orig_quant == buffer.nvfp4_params.fp8_scale_quant_orig == 1.0 + for buffer in plan.buffers + ) def test_runtime_admission_is_checked_in_utils_before_manager_creation(monkeypatch): From 1971d2f03f22e7c2b16f5e48e88ff26ab934e896 Mon Sep 17 00:00:00 2001 From: tianruih Date: Tue, 25 Aug 2026 16:08:26 -0700 Subject: [PATCH 09/29] [None][refactor] Split generic cold-page codec backend Signed-off-by: tianruih --- .../kv_cache_compression/CMakeLists.txt | 3 +- .../kv_cache_compression/coldPageCodec.cpp | 416 ---------------- .../kv_cache_compression/coldPageCodec.h | 108 ---- .../nativeColdPageCodec.cpp | 260 ++++++++++ .../nativeColdPageCodec.h | 97 ++++ .../nvfp4ColdPageCodecBackend.cpp | 181 +++++++ .../nvfp4ColdPageCodecBackend.h | 58 +++ .../nanobind/kvCacheCompression/bindings.cpp | 59 +-- .../kv_cache_compression/CMakeLists.txt | 3 +- .../coldPageCodecTest.cpp | 469 +++++++++--------- .../quantization_for_cold_page/nvfp4.py | 75 ++- .../test_quantization_for_cold_page.py | 116 ++--- 12 files changed, 927 insertions(+), 918 deletions(-) delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp create mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt index c2cf3d8d28eb..756e80176dac 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -1,7 +1,8 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. # All rights reserved. SPDX-License-Identifier: Apache-2.0 -add_library(kv_cache_compression_src OBJECT coldPageCodec.cpp) +add_library(kv_cache_compression_src OBJECT nativeColdPageCodec.cpp + nvfp4ColdPageCodecBackend.cpp) target_include_directories( kv_cache_compression_src PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp deleted file mode 100644 index 810107ecb1a3..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp +++ /dev/null @@ -1,416 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" - -#include "tensorrt_llm/common/logger.h" - -#include -#include -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ -namespace -{ - -std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) -{ - if (lhs > std::numeric_limits::max() - rhs) - { - throw std::overflow_error(label); - } - return lhs + rhs; -} - -std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) -{ - if (offset > std::numeric_limits::max() - base) - { - throw std::overflow_error("GPU buffer address overflows uintptr_t"); - } - return base + offset; -} - -struct BufferLocation -{ - kv::PoolIndex poolIndex{0}; - std::size_t offset = 0; - std::size_t bytes = 0; -}; - -using LayerBuffers = std::map; -using LifecycleBuffers = std::map; -using LayerPlans = std::map; - -LifecycleBuffers indexLifecycleBuffers(kv::SlotDescVariant const& variant) -{ - LifecycleBuffers result; - for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) - { - auto const& coalesced = variant.coalescedBuffers.at(poolIndex); - std::size_t offset = 0; - for (auto const& bufferId : coalesced.bufferIds) - { - auto& buffers = result[bufferId.layerId]; - if (!buffers.emplace(bufferId.role, BufferLocation{poolIndex, offset, coalesced.singleBufferSize}).second) - { - throw std::invalid_argument("GPU lifecycle contains a duplicate buffer role"); - } - offset = checkedAdd(offset, coalesced.singleBufferSize, "GPU Slot buffer offsets overflow size_t"); - } - } - return result; -} - -kernels::Nvfp4ColdPageKernelParams toKernelParams(Nvfp4ColdPageParams const& params) -{ - return kernels::Nvfp4ColdPageKernelParams{params.numKvHeads, params.tokensPerPage, params.headDim, - params.nvfp4ScaleOrigQuant, params.nvfp4ScaleQuantOrig, params.fp8ScaleOrigQuant, params.fp8ScaleQuantOrig}; -} - -struct CompiledLayerGroup -{ - kernels::Nvfp4ColdPagePreparedPlan preparedPlan; - std::size_t coldPageBytes = 0; - std::set consumedLayers; -}; - -CompiledLayerGroup compileLayerGroup( - kv::PoolGroupDesc const& gpuDesc, LifecycleBuffers const& physicalLayers, LayerPlans const& layerPlans) -{ - CompiledLayerGroup result; - std::vector groupBuffers; - std::optional groupRuntimeType; - - for (auto const& [layerId, physicalBuffers] : physicalLayers) - { - auto const planIt = layerPlans.find(layerId); - if (planIt == layerPlans.end()) - { - throw std::invalid_argument("A planned lifecycle contains a layer without a cold-page plan"); - } - auto const& layerPlan = planIt->second; - std::vector layerBuffers; - layerBuffers.reserve(layerPlan.buffers.size()); - std::optional layerRuntimeType; - - for (auto const& bufferPlan : layerPlan.buffers) - { - auto const locationIt = physicalBuffers.find(bufferPlan.role); - if (locationIt == physicalBuffers.end()) - { - throw std::invalid_argument( - "Cold-page plan references a buffer " - "absent from the GPU lifecycle"); - } - auto const& location = locationIt->second; - if (bufferPlan.rawBytes != location.bytes) - { - throw std::invalid_argument("Cold-page plan raw size does not match the GPU buffer size"); - } - - auto const& pool = gpuDesc.pools.at(location.poolIndex); - kernels::Nvfp4ColdPageBufferPlan nativePlan{checkedAddress(pool.baseAddress, location.offset), - pool.slotBytes, bufferPlan.rawBytes, bufferPlan.coldDataOffset, bufferPlan.coldScaleOffset, 0U, 0U, - kernels::Nvfp4ColdPageTransform::kLosslessCopy, {}}; - switch (bufferPlan.transform) - { - case ColdPageTransformKind::kLosslessCopy: break; - case ColdPageTransformKind::kNvfp4: - if (!bufferPlan.nvfp4Params) - { - throw std::invalid_argument("NVFP4 cold-page transform requires NVFP4 parameters"); - } - nativePlan.transform = kernels::Nvfp4ColdPageTransform::kNvfp4; - nativePlan.params = toKernelParams(*bufferPlan.nvfp4Params); - if (layerRuntimeType && *layerRuntimeType != bufferPlan.nvfp4Params->runtimeType) - { - throw std::invalid_argument("One layer cold-page plan must use one runtime type"); - } - layerRuntimeType = bufferPlan.nvfp4Params->runtimeType; - break; - default: throw std::invalid_argument("Unsupported cold-page transform kind"); - } - layerBuffers.push_back(nativePlan); - } - if (layerBuffers.size() != physicalBuffers.size()) - { - throw std::invalid_argument("Cold-page layer plan must cover every GPU buffer role"); - } - - layerBuffers.back().coldPaddingOffset = layerPlan.coldPaddingOffset; - layerBuffers.back().coldPaddingBytes = static_cast(layerPlan.coldPaddingBytes); - auto const runtimeType = layerRuntimeType.value_or(kernels::Nvfp4ColdPageRuntimeType::kFloat16); - auto const validatedLayer - = kernels::prepareNvfp4ColdPagePlan(layerBuffers, layerPlan.coldPageBytes, runtimeType); - - if (layerRuntimeType) - { - if (groupRuntimeType && *groupRuntimeType != *layerRuntimeType) - { - throw std::invalid_argument("One lifecycle cold-page plan must use one runtime type"); - } - groupRuntimeType = *layerRuntimeType; - } - for (std::uint32_t index = 0; index < validatedLayer.numBuffers; ++index) - { - auto nativePlan = validatedLayer.buffers[index]; - nativePlan.coldDataOffset - = checkedAdd(result.coldPageBytes, nativePlan.coldDataOffset, "Cold-page data offset overflows size_t"); - if (nativePlan.transform == kernels::Nvfp4ColdPageTransform::kNvfp4) - { - nativePlan.coldScaleOffset = checkedAdd( - result.coldPageBytes, nativePlan.coldScaleOffset, "Cold-page scale offset overflows size_t"); - } - if (nativePlan.coldPaddingBytes != 0U) - { - nativePlan.coldPaddingOffset = checkedAdd( - result.coldPageBytes, nativePlan.coldPaddingOffset, "Cold-page padding offset overflows size_t"); - } - groupBuffers.push_back(nativePlan); - } - result.coldPageBytes - = checkedAdd(result.coldPageBytes, layerPlan.coldPageBytes, "Lifecycle cold-page size overflows size_t"); - result.consumedLayers.insert(layerId); - } - - result.preparedPlan = kernels::prepareNvfp4ColdPagePlan( - groupBuffers, result.coldPageBytes, groupRuntimeType.value_or(kernels::Nvfp4ColdPageRuntimeType::kFloat16)); - return result; -} - -} // namespace - -PlannedColdPageCodec::PlannedColdPageCodec(std::vector layerPlans) -{ - for (auto& layerPlan : layerPlans) - { - if (layerPlan.coldPageBytes == 0U || layerPlan.buffers.empty()) - { - throw std::invalid_argument("A cold-page layer plan must contain a non-empty record"); - } - if (layerPlan.coldPaddingOffset > layerPlan.coldPageBytes - || layerPlan.coldPaddingBytes != layerPlan.coldPageBytes - layerPlan.coldPaddingOffset - || layerPlan.coldPaddingBytes > std::numeric_limits::max()) - { - throw std::invalid_argument("Cold-page layer padding exceeds its record"); - } - - std::set roles; - for (auto const& buffer : layerPlan.buffers) - { - if (buffer.role.empty() || buffer.rawBytes == 0U || !roles.emplace(buffer.role).second) - { - throw std::invalid_argument( - "Cold-page layer plans require unique " - "non-empty buffer roles and sizes"); - } - if (buffer.transform != ColdPageTransformKind::kLosslessCopy - && buffer.transform != ColdPageTransformKind::kNvfp4) - { - throw std::invalid_argument("Unsupported cold-page transform kind"); - } - } - auto const layerId = layerPlan.layerId; - if (!mLayerPlans.emplace(layerId, std::move(layerPlan)).second) - { - throw std::invalid_argument("Cold-page layer plan IDs must be unique"); - } - } -} - -bool PlannedColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept -{ - try - { - if (mLosslessCodec || !mLayerGroups.empty()) - { - throw std::invalid_argument("PlannedColdPageCodec can be configured only once"); - } - auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); - if (!losslessCodec->configure(gpuDescs, numGpuDescs)) - { - throw std::invalid_argument("Default lossless codec rejected GPU layouts"); - } - - std::map pendingGroups; - std::set consumedLayers; - for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) - { - auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; - for (auto const& variant : gpuDesc.slotDesc.variants) - { - auto const physicalLayers = indexLifecycleBuffers(variant); - auto const plannedLayerCount = std::count_if(physicalLayers.begin(), physicalLayers.end(), - [this](auto const& layer) { return mLayerPlans.count(layer.first) != 0U; }); - - LayerGroupState state; - if (plannedLayerCount == 0U) - { - state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); - } - else - { - if (plannedLayerCount != physicalLayers.size()) - { - throw std::invalid_argument("A lifecycle cannot mix planned and unplanned layers"); - } - auto compiled = compileLayerGroup(gpuDesc, physicalLayers, mLayerPlans); - state.execution = ExecutionKind::kPlanned; - state.preparedPlan = std::move(compiled.preparedPlan); - state.coldPageBytes = compiled.coldPageBytes; - for (auto const layerId : compiled.consumedLayers) - { - if (!consumedLayers.emplace(layerId).second) - { - throw std::invalid_argument("A cold-page layer plan appears in multiple lifecycles"); - } - } - } - if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) - { - throw std::invalid_argument("GPU lifecycle ID appears in multiple pool groups"); - } - } - } - if (consumedLayers.size() != mLayerPlans.size()) - { - throw std::invalid_argument("A cold-page layer plan is absent from all GPU descriptors"); - } - - mLayerGroups = std::move(pendingGroups); - mLosslessCodec = std::move(losslessCodec); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("PlannedColdPageCodec::configure rejected GPU layouts: %s", error.what()); - return false; - } - catch (...) - { - TLLM_LOG_ERROR( - "PlannedColdPageCodec::configure rejected GPU layouts: " - "unknown error"); - return false; - } -} - -PlannedColdPageCodec::LayerGroupState const* PlannedColdPageCodec::findLayerGroup( - kv::LayerGroupId layerGroupId) const noexcept -{ - auto const found = mLayerGroups.find(layerGroupId); - return found == mLayerGroups.end() ? nullptr : &found->second; -} - -std::size_t PlannedColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept -{ - auto const* state = findLayerGroup(layerGroupId); - return state == nullptr ? 0U : state->coldPageBytes; -} - -kv::LayerGroupId PlannedColdPageCodec::getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept -{ - return findLayerGroup(layerGroupId) == nullptr ? kv::LayerGroupId{-1} : layerGroupId; -} - -kv::PageIndexLocation PlannedColdPageCodec::queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept -{ - return findLayerGroup(layerGroupId) == nullptr ? kv::PageIndexLocation::kBadLocation : kv::PageIndexLocation::kHost; -} - -bool PlannedColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept -{ - try - { - auto const* state = findLayerGroup(layerGroupId); - if (state == nullptr || (numBasePages != 0U && (dstBasePtr == nullptr || pageIndices == nullptr))) - { - throw std::invalid_argument("encode received an invalid lifecycle or Page batch"); - } - if (numBasePages == 0U) - { - return true; - } - if (state->execution == ExecutionKind::kLossless) - { - return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); - } - - thread_local std::vector pages; - pages.clear(); - pages.reserve(numBasePages); - for (std::size_t page = 0; page < numBasePages; ++page) - { - pages.push_back({pageIndices[page].src, pageIndices[page].dst}); - } - kernels::invokeNvfp4ColdPageEncode(pages, state->preparedPlan, dstBasePtr, stream); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("PlannedColdPageCodec::encode failed before completion fencing: %s", error.what()); - return false; - } - catch (...) - { - TLLM_LOG_ERROR( - "PlannedColdPageCodec::encode failed before completion " - "fencing: unknown error"); - return false; - } -} - -bool PlannedColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, - kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept -{ - try - { - auto const* state = findLayerGroup(layerGroupId); - if (state == nullptr || (numBasePages != 0U && (srcBasePtr == nullptr || pageIndices == nullptr))) - { - throw std::invalid_argument("decode received an invalid lifecycle or Page batch"); - } - if (numBasePages == 0U) - { - return true; - } - if (state->execution == ExecutionKind::kLossless) - { - return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); - } - - thread_local std::vector pages; - pages.clear(); - pages.reserve(numBasePages); - for (std::size_t page = 0; page < numBasePages; ++page) - { - pages.push_back({pageIndices[page].dst, pageIndices[page].src}); - } - kernels::invokeNvfp4ColdPageDecode(pages, state->preparedPlan, srcBasePtr, stream); - return true; - } - catch (std::exception const& error) - { - TLLM_LOG_ERROR("PlannedColdPageCodec::decode failed before completion fencing: %s", error.what()); - return false; - } - catch (...) - { - TLLM_LOG_ERROR( - "PlannedColdPageCodec::decode failed before completion " - "fencing: unknown error"); - return false; - } -} - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h deleted file mode 100644 index bbb725cba1e0..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/coldPageCodec.h +++ /dev/null @@ -1,108 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#pragma once - -#include "kv_cache_manager_v2/coldPageCodec.h" -#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" - -#include -#include -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ - -namespace kv = batch_manager::kv_cache_manager_v2; - -//! Native transform selected for one buffer in a planned cold-page record. -enum class ColdPageTransformKind : std::uint8_t -{ - kLosslessCopy, - kNvfp4, -}; - -//! NVFP4 parameters for one independently scaled buffer. -struct Nvfp4ColdPageParams -{ - kernels::Nvfp4ColdPageRuntimeType runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; - std::int32_t numKvHeads = 0; - std::int32_t tokensPerPage = 0; - std::int32_t headDim = 0; - float nvfp4ScaleOrigQuant = 1.0F; - float nvfp4ScaleQuantOrig = 1.0F; - float fp8ScaleOrigQuant = 1.0F; - float fp8ScaleQuantOrig = 1.0F; -}; - -//! Python-authored transform and layer-relative cold offsets for one buffer. -struct ColdPageBufferPlan -{ - kv::DataRole role; - ColdPageTransformKind transform = ColdPageTransformKind::kLosslessCopy; - std::size_t rawBytes = 0; - std::size_t coldDataOffset = 0; - std::size_t coldScaleOffset = 0; - std::optional nvfp4Params; -}; - -//! Python-authored fixed cold record for one layer. -struct ColdPageLayerPlan -{ - kv::LayerId layerId = 0; - std::size_t coldPageBytes = 0; - std::size_t coldPaddingOffset = 0; - std::size_t coldPaddingBytes = 0; - std::vector buffers; -}; - -//! Resolves declarative layer plans against KVCM's authoritative hot-pool -// descriptors. -class PlannedColdPageCodec final : public kv::IKvCacheColdPageCodec -{ -public: - explicit PlannedColdPageCodec(std::vector layerPlans); - - bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; - - [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; - - bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - - bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - -private: - enum class ExecutionKind : std::uint8_t - { - kLossless, - kPlanned, - }; - - struct LayerGroupState - { - ExecutionKind execution = ExecutionKind::kLossless; - kernels::Nvfp4ColdPagePreparedPlan preparedPlan; - std::size_t coldPageBytes = 0; - }; - - [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; - - std::map mLayerPlans; - std::map mLayerGroups; - std::unique_ptr mLosslessCodec; -}; - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp new file mode 100644 index 000000000000..e30d4922efb4 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -0,0 +1,260 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" + +#include "tensorrt_llm/common/logger.h" + +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +std::size_t checkedAdd(std::size_t lhs, std::size_t rhs) +{ + if (lhs > std::numeric_limits::max() - rhs) + { + throw std::overflow_error("GPU Slot buffer offsets overflow size_t"); + } + return lhs + rhs; +} + +std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) +{ + if (offset > std::numeric_limits::max() - base) + { + throw std::overflow_error("GPU buffer address overflows uintptr_t"); + } + return base + offset; +} + +ResolvedColdPageLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::SlotDescVariant const& variant) +{ + ResolvedColdPageLifecycle result; + for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) + { + auto const& coalesced = variant.coalescedBuffers.at(poolIndex); + auto const& pool = gpuDesc.pools.at(poolIndex); + std::size_t offset = 0; + for (auto const& bufferId : coalesced.bufferIds) + { + auto& layer = result[bufferId.layerId]; + if (!layer + .emplace(bufferId.role, + ResolvedColdPageBuffer{ + checkedAddress(pool.baseAddress, offset), pool.slotBytes, coalesced.singleBufferSize}) + .second) + { + throw std::invalid_argument("GPU lifecycle contains a duplicate buffer role"); + } + offset = checkedAdd(offset, coalesced.singleBufferSize); + } + } + return result; +} + +} // namespace + +NativeColdPageCodec::NativeColdPageCodec(std::unique_ptr backend) + : mBackend(std::move(backend)) +{ + if (!mBackend) + { + throw std::invalid_argument("NativeColdPageCodec requires a backend"); + } +} + +bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept +{ + try + { + auto losslessCodec = kv::createDefaultKvCacheColdPageCodec(); + if (!losslessCodec->configure(gpuDescs, numGpuDescs)) + { + throw std::invalid_argument("Default lossless codec rejected GPU layouts"); + } + + auto const& backendLayerIds = mBackend->getLayerIds(); + std::map pendingGroups; + std::vector backendLifecycles; + std::vector backendLayerGroups; + std::set consumedLayers; + + for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) + { + auto const& gpuDesc = gpuDescs[kv::toSizeT(poolGroupIndex)]; + for (auto const& variant : gpuDesc.slotDesc.variants) + { + auto resolved = resolveLifecycle(gpuDesc, variant); + auto const backendLayerCount = std::count_if(resolved.begin(), resolved.end(), + [&backendLayerIds](auto const& layer) { return backendLayerIds.count(layer.first) != 0U; }); + + LayerGroupState state; + if (backendLayerCount == 0U) + { + state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); + state.pageIndexLocation = losslessCodec->queryPageIndexLocation(variant.lifeCycleId); + } + else + { + if (backendLayerCount != resolved.size()) + { + throw std::invalid_argument("A lifecycle cannot mix backend-owned and fallback layers"); + } + for (auto const& [layerId, buffers] : resolved) + { + static_cast(buffers); + if (!consumedLayers.emplace(layerId).second) + { + throw std::invalid_argument("A backend layer appears in multiple lifecycles"); + } + } + state.backendIndex = backendLifecycles.size(); + backendLayerGroups.push_back(variant.lifeCycleId); + backendLifecycles.push_back(std::move(resolved)); + } + + if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) + { + throw std::invalid_argument("GPU lifecycle ID appears in multiple pool groups"); + } + } + } + if (consumedLayers != backendLayerIds) + { + throw std::invalid_argument("A backend layer is absent from all GPU descriptors"); + } + + auto const backendConfigs = mBackend->configure(backendLifecycles); + if (backendConfigs.size() != backendLifecycles.size()) + { + throw std::invalid_argument("Cold-page backend returned an unexpected lifecycle count"); + } + for (std::size_t index = 0; index < backendConfigs.size(); ++index) + { + auto const& config = backendConfigs[index]; + if (config.coldPageBytes == 0U || config.pageIndexLocation == kv::PageIndexLocation::kBadLocation) + { + throw std::invalid_argument("Cold-page backend returned invalid storage properties"); + } + auto& state = pendingGroups.at(backendLayerGroups[index]); + state.coldPageBytes = config.coldPageBytes; + state.pageIndexLocation = config.pageIndexLocation; + } + + mLayerGroups = std::move(pendingGroups); + mLosslessCodec = std::move(losslessCodec); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("NativeColdPageCodec::configure rejected GPU layouts: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR("NativeColdPageCodec::configure rejected GPU layouts: unknown error"); + return false; + } +} + +NativeColdPageCodec::LayerGroupState const* NativeColdPageCodec::findLayerGroup( + kv::LayerGroupId layerGroupId) const noexcept +{ + auto const found = mLayerGroups.find(layerGroupId); + return found == mLayerGroups.end() ? nullptr : &found->second; +} + +std::size_t NativeColdPageCodec::queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept +{ + auto const* state = findLayerGroup(layerGroupId); + return state == nullptr ? 0U : state->coldPageBytes; +} + +kv::LayerGroupId NativeColdPageCodec::getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept +{ + return findLayerGroup(layerGroupId) == nullptr ? kv::LayerGroupId{-1} : layerGroupId; +} + +kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept +{ + auto const* state = findLayerGroup(layerGroupId); + return state == nullptr ? kv::PageIndexLocation::kBadLocation : state->pageIndexLocation; +} + +bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && (dstBasePtr == nullptr || pageIndices == nullptr))) + { + throw std::invalid_argument("encode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (!state->backendIndex) + { + return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); + } + mBackend->encode(*state->backendIndex, dstBasePtr, pageIndices, numBasePages, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: unknown error"); + return false; + } +} + +bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, + kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept +{ + try + { + auto const* state = findLayerGroup(layerGroupId); + if (state == nullptr || (numBasePages != 0U && (srcBasePtr == nullptr || pageIndices == nullptr))) + { + throw std::invalid_argument("decode received an invalid lifecycle or Page batch"); + } + if (numBasePages == 0U) + { + return true; + } + if (!state->backendIndex) + { + return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); + } + mBackend->decode(*state->backendIndex, srcBasePtr, pageIndices, numBasePages, stream); + return true; + } + catch (std::exception const& error) + { + TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: %s", error.what()); + return false; + } + catch (...) + { + TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: unknown error"); + return false; + } +} + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h new file mode 100644 index 000000000000..c2e110d24634 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -0,0 +1,97 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include "kv_cache_manager_v2/coldPageCodec.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +//! One hot buffer resolved from KVCM's authoritative pool descriptors. +struct ResolvedColdPageBuffer +{ + std::uintptr_t rawBase = 0; + std::size_t rawSlotBytes = 0; + std::size_t rawBytes = 0; +}; + +using ResolvedColdPageLayer = std::map; +using ResolvedColdPageLifecycle = std::map; + +//! Storage properties produced while a backend prepares one lifecycle. +struct ColdPageLifecycleConfig +{ + std::size_t coldPageBytes = 0; + kv::PageIndexLocation pageIndexLocation = kv::PageIndexLocation::kBadLocation; +}; + +//! Native algorithm backend retained by the generic KVCM codec adapter. +class IColdPageCodecBackend +{ +public: + virtual ~IColdPageCodecBackend() = default; + + [[nodiscard]] virtual std::set const& getLayerIds() const noexcept = 0; + + virtual std::vector configure(std::vector const& lifecycles) + = 0; + + virtual void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) + = 0; + + virtual void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) + = 0; +}; + +//! Resolves KVCM layouts, routes lifecycles, and owns one native backend. +class NativeColdPageCodec final : public kv::IKvCacheColdPageCodec +{ +public: + explicit NativeColdPageCodec(std::unique_ptr backend); + + bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; + + [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; + + [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; + + bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + + bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept override; + +private: + struct LayerGroupState + { + std::optional backendIndex; + std::size_t coldPageBytes = 0; + kv::PageIndexLocation pageIndexLocation = kv::PageIndexLocation::kBadLocation; + }; + + [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; + + std::unique_ptr mBackend; + std::map mLayerGroups; + std::unique_ptr mLosslessCodec; +}; + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp new file mode 100644 index 000000000000..0611e78f03ab --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp @@ -0,0 +1,181 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" + +#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" + +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ +namespace +{ + +std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) +{ + if (lhs > std::numeric_limits::max() - rhs) + { + throw std::overflow_error(label); + } + return lhs + rhs; +} + +kernels::Nvfp4ColdPageKernelParams makeKernelParams( + Nvfp4ColdPageLayerLayout const& layout, Nvfp4ColdPageScales const& scales) +{ + return kernels::Nvfp4ColdPageKernelParams{layout.numKvHeads, layout.tokensPerPage, layout.headDim, + scales.nvfp4ScaleOrigQuant, scales.nvfp4ScaleQuantOrig, scales.fp8ScaleOrigQuant, scales.fp8ScaleQuantOrig}; +} + +class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend +{ +public: + explicit Nvfp4ColdPageCodecBackend(std::vector layerLayouts) + { + for (auto& layout : layerLayouts) + { + if (layout.coldPageBytes == 0U || layout.buffers.empty() || layout.coldPaddingOffset > layout.coldPageBytes) + { + throw std::invalid_argument("An NVFP4 layer layout must contain a valid cold-page record"); + } + std::set roles; + for (auto const& buffer : layout.buffers) + { + if (buffer.role.empty() || !roles.emplace(buffer.role).second) + { + throw std::invalid_argument("NVFP4 layer layouts require unique non-empty buffer roles"); + } + } + auto const layerId = layout.layerId; + if (!mLayerLayouts.emplace(layerId, std::move(layout)).second) + { + throw std::invalid_argument("NVFP4 layer layout IDs must be unique"); + } + mLayerIds.emplace(layerId); + } + } + + [[nodiscard]] std::set const& getLayerIds() const noexcept override + { + return mLayerIds; + } + + std::vector configure(std::vector const& lifecycles) override + { + std::vector preparedPlans; + std::vector configs; + preparedPlans.reserve(lifecycles.size()); + configs.reserve(lifecycles.size()); + + for (auto const& lifecycle : lifecycles) + { + std::vector buffers; + std::optional runtimeType; + std::size_t coldPageBytes = 0; + + for (auto const& [layerId, resolvedBuffers] : lifecycle) + { + auto const& layout = mLayerLayouts.at(layerId); + if (layout.buffers.size() != resolvedBuffers.size()) + { + throw std::invalid_argument("NVFP4 layer layout must cover every GPU buffer role"); + } + if (runtimeType && *runtimeType != layout.runtimeType) + { + throw std::invalid_argument("One lifecycle must use one NVFP4 runtime type"); + } + runtimeType = layout.runtimeType; + + for (auto const& bufferLayout : layout.buffers) + { + auto const resolved = resolvedBuffers.find(bufferLayout.role); + if (resolved == resolvedBuffers.end()) + { + throw std::invalid_argument("NVFP4 layer layout references an absent GPU buffer role"); + } + auto const& raw = resolved->second; + auto const transform = bufferLayout.scales ? kernels::Nvfp4ColdPageTransform::kNvfp4 + : kernels::Nvfp4ColdPageTransform::kLosslessCopy; + buffers.push_back(kernels::Nvfp4ColdPageBufferPlan{raw.rawBase, raw.rawSlotBytes, raw.rawBytes, + checkedAdd( + coldPageBytes, bufferLayout.coldDataOffset, "Cold-page data offset overflows size_t"), + bufferLayout.scales ? checkedAdd( + coldPageBytes, bufferLayout.coldScaleOffset, "Cold-page scale offset overflows size_t") + : 0U, + 0U, 0U, transform, + bufferLayout.scales ? makeKernelParams(layout, *bufferLayout.scales) + : kernels::Nvfp4ColdPageKernelParams{}}); + } + + auto const paddingBytes = layout.coldPageBytes - layout.coldPaddingOffset; + if (paddingBytes > std::numeric_limits::max()) + { + throw std::overflow_error("Cold-page padding exceeds the kernel ABI"); + } + buffers.back().coldPaddingOffset + = checkedAdd(coldPageBytes, layout.coldPaddingOffset, "Cold-page padding offset overflows size_t"); + buffers.back().coldPaddingBytes = static_cast(paddingBytes); + coldPageBytes + = checkedAdd(coldPageBytes, layout.coldPageBytes, "Lifecycle cold-page size overflows size_t"); + } + + if (!runtimeType) + { + throw std::invalid_argument("NVFP4 backend received an empty lifecycle"); + } + preparedPlans.push_back(kernels::prepareNvfp4ColdPagePlan(buffers, coldPageBytes, *runtimeType)); + configs.push_back({coldPageBytes, kv::PageIndexLocation::kHost}); + } + + mPreparedPlans = std::move(preparedPlans); + return configs; + } + + void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, + cudaStream_t stream) override + { + thread_local std::vector pages; + pages.clear(); + pages.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + pages.push_back({pageIndices[page].src, pageIndices[page].dst}); + } + kernels::invokeNvfp4ColdPageEncode(pages, mPreparedPlans.at(lifecycleIndex), coldBase, stream); + } + + void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override + { + thread_local std::vector pages; + pages.clear(); + pages.reserve(numPages); + for (std::size_t page = 0; page < numPages; ++page) + { + pages.push_back({pageIndices[page].dst, pageIndices[page].src}); + } + kernels::invokeNvfp4ColdPageDecode(pages, mPreparedPlans.at(lifecycleIndex), coldBase, stream); + } + +private: + std::map mLayerLayouts; + std::set mLayerIds; + std::vector mPreparedPlans; +}; + +} // namespace + +std::unique_ptr createNvfp4ColdPageCodec(std::vector layerLayouts) +{ + return std::make_unique(std::make_unique(std::move(layerLayouts))); +} + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h new file mode 100644 index 000000000000..7cd0d0cb1972 --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h @@ -0,0 +1,58 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include "kv_cache_manager_v2/coldPageCodec.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" + +#include +#include +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ + +namespace kv = batch_manager::kv_cache_manager_v2; + +//! Per-buffer global scales used by the NVFP4 cold-page kernels. +struct Nvfp4ColdPageScales +{ + float nvfp4ScaleOrigQuant = 1.0F; + float nvfp4ScaleQuantOrig = 1.0F; + float fp8ScaleOrigQuant = 1.0F; + float fp8ScaleQuantOrig = 1.0F; +}; + +//! Algorithm and layer-relative cold offsets for one buffer. +struct Nvfp4ColdPageBufferLayout +{ + kv::DataRole role; + std::size_t coldDataOffset = 0; + std::size_t coldScaleOffset = 0; + std::optional scales; +}; + +//! Python-authored NVFP4 record layout for one Attention layer. +struct Nvfp4ColdPageLayerLayout +{ + kv::LayerId layerId = 0; + kernels::Nvfp4ColdPageRuntimeType runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; + std::int32_t numKvHeads = 0; + std::int32_t tokensPerPage = 0; + std::int32_t headDim = 0; + std::size_t coldPageBytes = 0; + std::size_t coldPaddingOffset = 0; + std::vector buffers; +}; + +//! Create an owning native codec configured by NVFP4 layer layouts. +[[nodiscard]] std::unique_ptr createNvfp4ColdPageCodec( + std::vector layerLayouts); + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index f60778ff175d..e5b4816d2d7d 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,7 +17,7 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" #include #include @@ -25,63 +25,46 @@ #include #include -#include -#include -#include - namespace nb = nanobind; namespace compression = tensorrt_llm::kv_cache_compression; namespace kernels = tensorrt_llm::kernels; -namespace kv = tensorrt_llm::batch_manager::kv_cache_manager_v2; namespace tensorrt_llm::nanobind::kv_cache_compression { void initBindings(nb::module_& module) { - nb::enum_(module, "ColdPageTransformKind") - .value("LOSSLESS_COPY", compression::ColdPageTransformKind::kLosslessCopy) - .value("NVFP4", compression::ColdPageTransformKind::kNvfp4); - nb::enum_(module, "Nvfp4ColdPageRuntimeType") .value("FLOAT16", kernels::Nvfp4ColdPageRuntimeType::kFloat16) .value("BFLOAT16", kernels::Nvfp4ColdPageRuntimeType::kBfloat16) .value("FP8_E4M3", kernels::Nvfp4ColdPageRuntimeType::kFp8E4m3); - nb::class_(module, "Nvfp4ColdPageParams") + nb::class_(module, "Nvfp4ColdPageScales") .def(nb::init<>()) - .def_rw("runtime_type", &compression::Nvfp4ColdPageParams::runtimeType) - .def_rw("num_kv_heads", &compression::Nvfp4ColdPageParams::numKvHeads) - .def_rw("tokens_per_page", &compression::Nvfp4ColdPageParams::tokensPerPage) - .def_rw("head_dim", &compression::Nvfp4ColdPageParams::headDim) - .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageParams::nvfp4ScaleOrigQuant) - .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageParams::nvfp4ScaleQuantOrig) - .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageParams::fp8ScaleOrigQuant) - .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageParams::fp8ScaleQuantOrig); + .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageScales::nvfp4ScaleOrigQuant) + .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageScales::nvfp4ScaleQuantOrig) + .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageScales::fp8ScaleOrigQuant) + .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageScales::fp8ScaleQuantOrig); - nb::class_(module, "ColdPageBufferPlan") + nb::class_(module, "Nvfp4ColdPageBufferLayout") .def(nb::init<>()) - .def_rw("role", &compression::ColdPageBufferPlan::role) - .def_rw("transform", &compression::ColdPageBufferPlan::transform) - .def_rw("raw_bytes", &compression::ColdPageBufferPlan::rawBytes) - .def_rw("cold_data_offset", &compression::ColdPageBufferPlan::coldDataOffset) - .def_rw("cold_scale_offset", &compression::ColdPageBufferPlan::coldScaleOffset) - .def_rw("nvfp4_params", &compression::ColdPageBufferPlan::nvfp4Params); + .def_rw("role", &compression::Nvfp4ColdPageBufferLayout::role) + .def_rw("cold_data_offset", &compression::Nvfp4ColdPageBufferLayout::coldDataOffset) + .def_rw("cold_scale_offset", &compression::Nvfp4ColdPageBufferLayout::coldScaleOffset) + .def_rw("scales", &compression::Nvfp4ColdPageBufferLayout::scales); - nb::class_(module, "ColdPageLayerPlan") + nb::class_(module, "Nvfp4ColdPageLayerLayout") .def(nb::init<>()) - .def_rw("layer_id", &compression::ColdPageLayerPlan::layerId) - .def_rw("cold_page_bytes", &compression::ColdPageLayerPlan::coldPageBytes) - .def_rw("cold_padding_offset", &compression::ColdPageLayerPlan::coldPaddingOffset) - .def_rw("cold_padding_bytes", &compression::ColdPageLayerPlan::coldPaddingBytes) - .def_rw("buffers", &compression::ColdPageLayerPlan::buffers); + .def_rw("layer_id", &compression::Nvfp4ColdPageLayerLayout::layerId) + .def_rw("runtime_type", &compression::Nvfp4ColdPageLayerLayout::runtimeType) + .def_rw("num_kv_heads", &compression::Nvfp4ColdPageLayerLayout::numKvHeads) + .def_rw("tokens_per_page", &compression::Nvfp4ColdPageLayerLayout::tokensPerPage) + .def_rw("head_dim", &compression::Nvfp4ColdPageLayerLayout::headDim) + .def_rw("cold_page_bytes", &compression::Nvfp4ColdPageLayerLayout::coldPageBytes) + .def_rw("cold_padding_offset", &compression::Nvfp4ColdPageLayerLayout::coldPaddingOffset) + .def_rw("buffers", &compression::Nvfp4ColdPageLayerLayout::buffers); - // Construct in C++ so ownership can transfer to KVCM as a unique_ptr codec. - module.def( - "create_cold_page_codec", - [](std::vector layerPlans) -> std::unique_ptr - { return std::make_unique(std::move(layerPlans)); }, - nb::arg("layer_plans"), "Create an owning planned cold-page codec for transfer into KVCacheManager."); + module.def("create_nvfp4_cold_page_codec", &compression::createNvfp4ColdPageCodec, nb::arg("layer_layouts")); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index 650b8a19f1df..9962aa05df62 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -6,7 +6,8 @@ set(COLD_PAGE_CODEC_TEST_SRC ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/coldPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp index 0af91c512de2..ca692381bc43 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -4,7 +4,8 @@ * SPDX-License-Identifier: Apache-2.0 */ -#include "tensorrt_llm/kv_cache_compression/coldPageCodec.h" +#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" #include @@ -12,10 +13,11 @@ #include #include #include -#include #include +#include #include #include +#include #include namespace tensorrt_llm::kv_cache_compression @@ -25,20 +27,7 @@ namespace namespace kv = batch_manager::kv_cache_manager_v2; -static_assert(std::is_base_of_v); - -struct RecordedLaunch -{ - int encodeCalls = 0; - int decodeCalls = 0; - std::vector encodePages; - std::vector decodePages; - kernels::Nvfp4ColdPagePreparedPlan plan; - void const* coldBase = nullptr; - cudaStream_t stream{}; -}; - -RecordedLaunch gLaunch; +static_assert(std::is_base_of_v); constexpr std::uintptr_t kGpuKBase = 0x100000; constexpr std::uintptr_t kGpuVBase = 0x200000; @@ -47,43 +36,6 @@ constexpr std::uintptr_t kStreamValue = 0x7000; constexpr std::size_t kRawBytes = 320; constexpr std::size_t kColdBytes = 192; -void resetLaunch() -{ - gLaunch = {}; -} - -Nvfp4ColdPageParams makeParams(float scale = 1.0F) -{ - Nvfp4ColdPageParams params; - params.runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; - params.numKvHeads = 1; - params.tokensPerPage = 5; - params.headDim = 32; - params.nvfp4ScaleOrigQuant = scale; - params.nvfp4ScaleQuantOrig = 1.0F / scale; - return params; -} - -ColdPageLayerPlan makeAttentionPlan(int layerId, float keyScale = 1.0F, float valueScale = 2.0F) -{ - return ColdPageLayerPlan{layerId, kColdBytes, 180U, 12U, - {ColdPageBufferPlan{"key", ColdPageTransformKind::kNvfp4, kRawBytes, 0U, 160U, makeParams(keyScale)}, - ColdPageBufferPlan{"value", ColdPageTransformKind::kNvfp4, kRawBytes, 80U, 170U, makeParams(valueScale)}}}; -} - -ColdPageLayerPlan makeMlaPlan(int layerId, bool hasIndex) -{ - std::vector buffers; - buffers.push_back(ColdPageBufferPlan{"key", ColdPageTransformKind::kNvfp4, kRawBytes, 0U, 80U, makeParams()}); - if (hasIndex) - { - buffers.push_back( - ColdPageBufferPlan{"index_key", ColdPageTransformKind::kLosslessCopy, 68U, 90U, 0U, std::nullopt}); - return ColdPageLayerPlan{layerId, 160U, 158U, 2U, std::move(buffers)}; - } - return ColdPageLayerPlan{layerId, 96U, 90U, 6U, std::move(buffers)}; -} - kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::PoolGroupIndex{0}, kv::LayerGroupId lifeCycle = kv::LayerGroupId{0}, std::size_t count = 1U, int firstLayer = 0, std::uintptr_t keyBase = kGpuKBase, std::uintptr_t valueBase = kGpuVBase) @@ -104,260 +56,297 @@ kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::Pool {kv::PoolIndex{0}, keyBase, slotBytes}, {kv::PoolIndex{1}, valueBase, slotBytes}}}; } -kv::PoolGroupDesc makeMlaDesc(bool hasIndex) +kv::PoolGroupDesc makeMlaDesc() { - kv::TypedVec buffers; - buffers.push_back(kv::CoalescedBuffer{kRawBytes, {{0, "key"}}}); - if (hasIndex) - { - buffers.push_back(kv::CoalescedBuffer{68U, {{0, "index_key"}}}); - } - kv::SlotDescVariant variant{kv::LayerGroupId{0}, std::move(buffers)}; - kv::TypedVec pools; - pools.push_back({kv::PoolIndex{0}, kGpuKBase, kRawBytes}); - if (hasIndex) - { - pools.push_back({kv::PoolIndex{1}, kGpuVBase, 68U}); - } - return {kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, std::move(pools)}; + kv::SlotDescVariant variant{kv::LayerGroupId{0}, + kv::TypedVec{ + kv::CoalescedBuffer{kRawBytes, {{0, "key"}}}, kv::CoalescedBuffer{68U, {{0, "index_key"}}}}}; + return {kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, + kv::TypedVec{ + {kv::PoolIndex{0}, kGpuKBase, kRawBytes}, {kv::PoolIndex{1}, kGpuVBase, 68U}}}; } kv::PoolGroupDesc makeLosslessDesc(kv::PoolGroupIndex poolGroupIndex, kv::LayerGroupId lifeCycle) { - kv::SlotDescVariant variant; - variant.lifeCycleId = lifeCycle; - variant.coalescedBuffers = kv::TypedVec{ - kv::CoalescedBuffer{64U, {{10, "ssm_state"}}}, kv::CoalescedBuffer{32U, {{10, "conv_state"}}}}; + kv::SlotDescVariant variant{lifeCycle, + kv::TypedVec{ + kv::CoalescedBuffer{64U, {{10, "ssm_state"}}}, kv::CoalescedBuffer{32U, {{10, "conv_state"}}}}}; return {poolGroupIndex, kv::SlotCount{8}, kv::SlotDesc{{std::move(variant)}}, kv::TypedVec{ {kv::PoolIndex{0}, 0x400000, 64U}, {kv::PoolIndex{1}, 0x500000, 32U}}}; } -bool configureOne(PlannedColdPageCodec& codec, kv::PoolGroupDesc const& desc) +bool configureOne(kv::IKvCacheColdPageCodec& codec, kv::PoolGroupDesc const& desc) { return codec.configure(&desc, kv::PoolGroupIndex{1}); } -TEST(PlannedColdPageCodecTest, ResolvesPythonAuthoredLayoutAndScales) +class RecordingBackend final : public IColdPageCodecBackend { - resetLaunch(); - PlannedColdPageCodec codec{{makeAttentionPlan(0, 2.0F, 3.0F), makeAttentionPlan(1, 4.0F, 5.0F)}}; +public: + explicit RecordingBackend(std::set layerIds) + : mLayerIds(std::move(layerIds)) + { + } + + [[nodiscard]] std::set const& getLayerIds() const noexcept override + { + return mLayerIds; + } + + std::vector configure(std::vector const& lifecycles) override + { + resolved = lifecycles; + if (failConfigure) + { + throw std::runtime_error("requested configure failure"); + } + return std::vector( + lifecycles.size(), ColdPageLifecycleConfig{777U, kv::PageIndexLocation::kHost}); + } + + void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, + cudaStream_t stream) override + { + ++encodeCalls; + lastLifecycleIndex = lifecycleIndex; + lastColdBase = coldBase; + lastIndices.assign(pageIndices, pageIndices + numPages); + lastStream = stream; + } + + void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override + { + ++decodeCalls; + lastLifecycleIndex = lifecycleIndex; + lastColdBase = coldBase; + lastIndices.assign(pageIndices, pageIndices + numPages); + lastStream = stream; + } + + bool failConfigure = false; + int encodeCalls = 0; + int decodeCalls = 0; + std::size_t lastLifecycleIndex = 0; + void const* lastColdBase = nullptr; + cudaStream_t lastStream{}; + std::vector lastIndices; + std::vector resolved; + +private: + std::set mLayerIds; +}; + +TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsBatches) +{ + auto backend = std::make_unique(std::set{0, 1}); + auto* recorder = backend.get(); + NativeColdPageCodec codec{std::move(backend)}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{3}, 2U))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{3}), 2U * kColdBytes); + ASSERT_EQ(recorder->resolved.size(), 1U); + auto const& layers = recorder->resolved.front(); + EXPECT_EQ(layers.at(0).at("key").rawBase, kGpuKBase); + EXPECT_EQ(layers.at(0).at("value").rawBase, kGpuVBase); + EXPECT_EQ(layers.at(1).at("key").rawBase, kGpuKBase + kRawBytes); + EXPECT_EQ(layers.at(1).at("key").rawSlotBytes, 2U * kRawBytes); + EXPECT_EQ(layers.at(1).at("key").rawBytes, kRawBytes); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{3}), 777U); kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; auto const stream = reinterpret_cast(kStreamValue); ASSERT_TRUE( codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - ASSERT_EQ(gLaunch.encodeCalls, 1); - ASSERT_EQ(gLaunch.encodePages.size(), 2U); - EXPECT_EQ(gLaunch.encodePages[0].gpuPageIndex, 1); - EXPECT_EQ(gLaunch.encodePages[0].coldPageIndex, 2); - ASSERT_EQ(gLaunch.plan.numBuffers, 4U); - EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); - EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kRawBytes); - EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, kColdBytes); - EXPECT_EQ(gLaunch.plan.buffers[3].coldScaleOffset, kColdBytes + 170U); - EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingOffset, kColdBytes + 180U); - EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingBytes, 12U); - EXPECT_FLOAT_EQ(gLaunch.plan.buffers[0].params.nvfp4ScaleOrigQuant, 2.0F); - EXPECT_FLOAT_EQ(gLaunch.plan.buffers[3].params.nvfp4ScaleOrigQuant, 5.0F); + EXPECT_EQ(recorder->encodeCalls, 1); + EXPECT_EQ(recorder->lastLifecycleIndex, 0U); + EXPECT_EQ(recorder->lastIndices[1].src, 3); + EXPECT_EQ(recorder->lastColdBase, reinterpret_cast(kColdBase)); + EXPECT_EQ(recorder->lastStream, stream); ASSERT_TRUE(codec.decode( kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - ASSERT_EQ(gLaunch.decodeCalls, 1); - EXPECT_EQ(gLaunch.decodePages[0].gpuPageIndex, 2); - EXPECT_EQ(gLaunch.decodePages[0].coldPageIndex, 1); -} - -TEST(PlannedColdPageCodecTest, PreservesExplicitMlaLosslessSideBuffer) -{ - resetLaunch(); - PlannedColdPageCodec codec{{makeMlaPlan(0, true)}}; - ASSERT_TRUE(configureOne(codec, makeMlaDesc(true))); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 160U); - - kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 2U); - EXPECT_EQ(gLaunch.plan.buffers[0].transform, kernels::Nvfp4ColdPageTransform::kNvfp4); - EXPECT_EQ(gLaunch.plan.buffers[1].transform, kernels::Nvfp4ColdPageTransform::kLosslessCopy); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBytes, 68U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingOffset, 158U); + EXPECT_EQ(recorder->decodeCalls, 1); } -TEST(PlannedColdPageCodecTest, UnplannedLifecycleDelegatesToDefaultCodec) +TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) { - PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; + auto backend = std::make_unique(std::set{0}); + auto* recorder = backend.get(); + NativeColdPageCodec codec{std::move(backend)}; std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), kColdBytes); + EXPECT_EQ(recorder->resolved.size(), 1U); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 777U); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 96U); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{1}), kv::LayerGroupId{1}); + EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{1}), kv::PageIndexLocation::kHost); } -TEST(PlannedColdPageCodecTest, RejectsDuplicateLayerAndRolePlans) +TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) { - std::vector duplicateLayers{makeAttentionPlan(0), makeAttentionPlan(0)}; - EXPECT_THROW({ PlannedColdPageCodec codec{duplicateLayers}; }, std::invalid_argument); - - auto plan = makeAttentionPlan(0); - plan.buffers.push_back(plan.buffers.front()); - EXPECT_THROW({ PlannedColdPageCodec codec{{plan}}; }, std::invalid_argument); + { + NativeColdPageCodec codec{std::make_unique(std::set{0})}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2U))); + } + { + NativeColdPageCodec codec{std::make_unique(std::set{0, 1})}; + EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); + } + { + NativeColdPageCodec codec{std::make_unique(std::set{0, 1})}; + std::array descs{makeAttentionDesc(), + makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1U, 1, 0x600000, 0x700000)}; + EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + } } -TEST(PlannedColdPageCodecTest, RejectsMissingNvfp4Parameters) +TEST(NativeColdPageCodecTest, CatchesBackendConfigureAndBatchFailures) { - auto plan = makeAttentionPlan(0); - plan.buffers.front().nvfp4Params.reset(); - PlannedColdPageCodec codec{{plan}}; + auto backend = std::make_unique(std::set{0}); + backend->failConfigure = true; + NativeColdPageCodec codec{std::move(backend)}; EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); + + auto validBackend = std::make_unique(std::set{0}); + auto* recorder = validBackend.get(); + NativeColdPageCodec validCodec{std::move(validBackend)}; + ASSERT_TRUE(configureOne(validCodec, makeAttentionDesc())); + EXPECT_TRUE(validCodec.encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + EXPECT_TRUE(validCodec.decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); + kv::PageIndexPair const indices[]{{0, 0}}; + EXPECT_FALSE(validCodec.encode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); + EXPECT_FALSE(validCodec.decode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); + EXPECT_EQ(recorder->encodeCalls, 0); + EXPECT_EQ(recorder->decodeCalls, 0); } -TEST(PlannedColdPageCodecTest, RejectsRawSizeAndRoleMismatches) +TEST(NativeColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) { - auto wrongSize = makeAttentionPlan(0); - wrongSize.buffers.front().rawBytes += 1U; - PlannedColdPageCodec sizeCodec{{wrongSize}}; - EXPECT_FALSE(configureOne(sizeCodec, makeAttentionDesc())); - - auto missingRole = makeAttentionPlan(0); - missingRole.buffers.pop_back(); - PlannedColdPageCodec roleCodec{{missingRole}}; - EXPECT_FALSE(configureOne(roleCodec, makeAttentionDesc())); + NativeColdPageCodec codec{std::make_unique(std::set{0})}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); + EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); } -TEST(PlannedColdPageCodecTest, RejectsPlannedAndUnplannedLayersInOneLifecycle) +Nvfp4ColdPageScales makeScales(float scale) { - PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; - EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2U))); + Nvfp4ColdPageScales scales; + scales.nvfp4ScaleOrigQuant = scale; + scales.nvfp4ScaleQuantOrig = 1.0F / scale; + return scales; } -TEST(PlannedColdPageCodecTest, RejectsPlanAbsentFromGpuDescriptors) +Nvfp4ColdPageLayerLayout makeAttentionLayout(int layerId, float keyScale = 1.0F, float valueScale = 2.0F) { - PlannedColdPageCodec codec{{makeAttentionPlan(0), makeAttentionPlan(1)}}; - EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); + return Nvfp4ColdPageLayerLayout{layerId, kernels::Nvfp4ColdPageRuntimeType::kFloat16, 1, 5, 32, kColdBytes, 180U, + {Nvfp4ColdPageBufferLayout{"key", 0U, 160U, makeScales(keyScale)}, + Nvfp4ColdPageBufferLayout{"value", 80U, 170U, makeScales(valueScale)}}}; } -TEST(PlannedColdPageCodecTest, RejectsDuplicateLifecycleAcrossPoolGroups) +Nvfp4ColdPageLayerLayout makeMlaLayout() { - PlannedColdPageCodec codec{{makeAttentionPlan(0), makeAttentionPlan(1)}}; - std::array descs{ - makeAttentionDesc(), makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1U, 1, 0x600000, 0x700000)}; - EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + return Nvfp4ColdPageLayerLayout{0, kernels::Nvfp4ColdPageRuntimeType::kFloat16, 1, 5, 32, 160U, 158U, + {Nvfp4ColdPageBufferLayout{"key", 0U, 80U, makeScales(1.0F)}, + Nvfp4ColdPageBufferLayout{"index_key", 90U, 0U, std::nullopt}}}; } -TEST(PlannedColdPageCodecTest, RejectsInvalidLayerRelativeIntervals) +struct RecordedLaunch { - auto plan = makeAttentionPlan(0); - plan.buffers.back().coldDataOffset = plan.coldPageBytes; - PlannedColdPageCodec codec{{plan}}; - EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); -} + int prepareCalls = 0; + int encodeCalls = 0; + int decodeCalls = 0; + std::vector encodePages; + std::vector decodePages; + kernels::Nvfp4ColdPagePreparedPlan plan; +}; -TEST(PlannedColdPageCodecTest, EmptyBatchIsValidAndInvalidBatchesFailBeforeLaunch) +RecordedLaunch gLaunch; + +TEST(Nvfp4ColdPageCodecBackendTest, LowersMhaLayoutOnceAndDispatchesEncodeDecode) { - resetLaunch(); - PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; - ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); - EXPECT_TRUE(codec.encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); - EXPECT_TRUE(codec.decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); - - kv::PageIndexPair const valid[]{{0, 0}}; - EXPECT_FALSE(codec.encode(kv::LayerGroupId{0}, nullptr, valid, 1U, nullptr)); - EXPECT_FALSE(codec.decode(kv::LayerGroupId{0}, nullptr, valid, 1U, nullptr)); - EXPECT_EQ(gLaunch.encodeCalls, 0); - EXPECT_EQ(gLaunch.decodeCalls, 0); + gLaunch = {}; + auto codec = createNvfp4ColdPageCodec({makeAttentionLayout(0, 2.0F, 3.0F), makeAttentionLayout(1, 4.0F, 5.0F)}); + ASSERT_TRUE(configureOne(*codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{3}, 2U))); + EXPECT_EQ(gLaunch.prepareCalls, 1); + EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{3}), 2U * kColdBytes); + + kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; + auto const stream = reinterpret_cast(kStreamValue); + ASSERT_TRUE( + codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + ASSERT_EQ(gLaunch.plan.numBuffers, 4U); + EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); + EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kRawBytes); + EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, kColdBytes); + EXPECT_EQ(gLaunch.plan.buffers[3].coldScaleOffset, kColdBytes + 170U); + EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingOffset, kColdBytes + 180U); + EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingBytes, 12U); + EXPECT_FLOAT_EQ(gLaunch.plan.buffers[0].params.nvfp4ScaleOrigQuant, 2.0F); + EXPECT_FLOAT_EQ(gLaunch.plan.buffers[3].params.nvfp4ScaleOrigQuant, 5.0F); + ASSERT_EQ(gLaunch.encodePages.size(), 2U); + EXPECT_EQ(gLaunch.encodePages[0].gpuPageIndex, 1); + EXPECT_EQ(gLaunch.encodePages[0].coldPageIndex, 2); + + ASSERT_TRUE(codec->decode( + kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + ASSERT_EQ(gLaunch.decodePages.size(), 2U); + EXPECT_EQ(gLaunch.decodePages[0].gpuPageIndex, 2); + EXPECT_EQ(gLaunch.decodePages[0].coldPageIndex, 1); } -TEST(PlannedColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) +TEST(Nvfp4ColdPageCodecBackendTest, PreservesMlaSideBufferLosslessly) { - PlannedColdPageCodec codec{{makeAttentionPlan(0)}}; - ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); - EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); -} + gLaunch = {}; + auto codec = createNvfp4ColdPageCodec({makeMlaLayout()}); + ASSERT_TRUE(configureOne(*codec, makeMlaDesc())); -} // namespace -} // namespace tensorrt_llm::kv_cache_compression + kv::PageIndexPair const indices[]{{0, 0}}; + ASSERT_TRUE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + ASSERT_EQ(gLaunch.plan.numBuffers, 2U); + EXPECT_EQ(gLaunch.plan.buffers[0].transform, kernels::Nvfp4ColdPageTransform::kNvfp4); + EXPECT_EQ(gLaunch.plan.buffers[1].transform, kernels::Nvfp4ColdPageTransform::kLosslessCopy); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); + EXPECT_EQ(gLaunch.plan.buffers[1].rawBytes, 68U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingOffset, 158U); + EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingBytes, 2U); +} -namespace tensorrt_llm::kernels -{ -namespace +TEST(Nvfp4ColdPageCodecBackendTest, RejectsDuplicateLayoutsAndMissingRoles) { + EXPECT_THROW( + { + auto codec = createNvfp4ColdPageCodec({makeAttentionLayout(0), makeAttentionLayout(0)}); + }, + std::invalid_argument); -struct Interval -{ - std::size_t begin; - std::size_t end; -}; + auto duplicateRole = makeAttentionLayout(0); + duplicateRole.buffers.push_back(duplicateRole.buffers.front()); + EXPECT_THROW({ auto codec = createNvfp4ColdPageCodec({duplicateRole}); }, std::invalid_argument); -void appendInterval(std::vector& intervals, std::size_t begin, std::size_t bytes, std::size_t pageBytes) -{ - if (bytes == 0U) - { - return; - } - if (begin > pageBytes || bytes > pageBytes - begin) - { - throw std::invalid_argument("test interval exceeds cold Page"); - } - intervals.push_back({begin, begin + bytes}); + auto missingRole = makeAttentionLayout(0); + missingRole.buffers.pop_back(); + auto codec = createNvfp4ColdPageCodec({missingRole}); + EXPECT_FALSE(configureOne(*codec, makeAttentionDesc())); } } // namespace +} // namespace tensorrt_llm::kv_cache_compression + +namespace tensorrt_llm::kernels +{ Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) { - if (buffers.empty() || buffers.size() > kNvfp4ColdPageMaxBuffersPerLaunch || coldPageBytes == 0U) + auto& launch = kv_cache_compression::gLaunch; + ++launch.prepareCalls; + if (buffers.empty() || buffers.size() > kNvfp4ColdPageMaxBuffersPerLaunch) { throw std::invalid_argument("invalid test launch plan"); } - std::vector intervals; - for (auto const& buffer : buffers) - { - if (buffer.rawBytes == 0U || buffer.rawBytes > buffer.rawSlotBytes) - { - throw std::invalid_argument("invalid raw buffer"); - } - if (buffer.transform == Nvfp4ColdPageTransform::kNvfp4) - { - auto const& params = buffer.params; - if (params.numKvHeads <= 0 || params.tokensPerPage <= 0 || params.headDim <= 0 || params.headDim % 16 != 0) - { - throw std::invalid_argument("invalid NVFP4 geometry"); - } - std::size_t const elements = static_cast(params.numKvHeads) - * static_cast(params.tokensPerPage) * static_cast(params.headDim); - auto const elementBytes = runtimeType == Nvfp4ColdPageRuntimeType::kFp8E4m3 ? 1U : 2U; - if (buffer.rawBytes != elements * elementBytes) - { - throw std::invalid_argument("raw size mismatch"); - } - appendInterval(intervals, buffer.coldDataOffset, elements / 2U, coldPageBytes); - appendInterval(intervals, buffer.coldScaleOffset, elements / 16U, coldPageBytes); - } - else - { - appendInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes); - } - appendInterval(intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes); - } - std::sort( - intervals.begin(), intervals.end(), [](auto const& lhs, auto const& rhs) { return lhs.begin < rhs.begin; }); - for (std::size_t index = 1; index < intervals.size(); ++index) - { - if (intervals[index - 1U].end > intervals[index].begin) - { - throw std::invalid_argument("test intervals overlap"); - } - } - Nvfp4ColdPagePreparedPlan plan; std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); plan.numBuffers = static_cast(buffers.size()); @@ -366,26 +355,22 @@ Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageEncode( + std::vector const& pages, Nvfp4ColdPagePreparedPlan const& plan, void*, cudaStream_t) { auto& launch = kv_cache_compression::gLaunch; ++launch.encodeCalls; launch.encodePages = pages; launch.plan = plan; - launch.coldBase = coldBase; - launch.stream = stream; } void invokeNvfp4ColdPageDecode(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream) + Nvfp4ColdPagePreparedPlan const& plan, void const*, cudaStream_t) { auto& launch = kv_cache_compression::gLaunch; ++launch.decodeCalls; launch.decodePages = pages; launch.plan = plan; - launch.coldBase = coldBase; - launch.stream = stream; } } // namespace tensorrt_llm::kernels diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py index 49ac51185822..4dbb3cf0002f 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py @@ -127,7 +127,7 @@ def create_cold_page_codec( layer for layer in cache_config.layers if isinstance(layer, AttentionLayerConfig) ] if not attention_layers: - return native.create_cold_page_codec([]) + return native.create_nvfp4_cold_page_codec([]) runtime_type = { DataType.HALF: native.Nvfp4ColdPageRuntimeType.FLOAT16, @@ -140,7 +140,7 @@ def create_cold_page_codec( f"Attention KV, not {runtime_dtype}" ) - layer_plans = [] + layer_layouts = [] for layer in attention_layers: layer_id = int(layer.layer_id) buffers_by_role = {str(buffer.role): buffer for buffer in layer.buffers} @@ -169,11 +169,6 @@ def create_cold_page_codec( elements = num_kv_heads * tokens_per_page * head_dim packed_bytes = elements // _ELEMENTS_PER_BYTE scale_bytes = elements // _ELEMENTS_PER_SCALE - raw_bytes_by_role = { - role: _buffer_bytes(buffer, tokens_per_page) - for role, buffer in buffers_by_role.items() - } - data_offsets = { role: index * packed_bytes for index, role in enumerate(compressed_roles) } @@ -184,47 +179,41 @@ def create_cold_page_codec( } cursor = scale_base + len(compressed_roles) * scale_bytes - buffer_plans = [] + buffer_layouts = [] for scale_index, role in enumerate(compressed_roles): - params = native.Nvfp4ColdPageParams() - params.runtime_type = runtime_type - params.num_kv_heads = num_kv_heads - params.tokens_per_page = tokens_per_page - params.head_dim = head_dim - params.nvfp4_scale_orig_quant = orig_quant[scale_index] - params.nvfp4_scale_quant_orig = quant_orig[scale_index] - params.fp8_scale_orig_quant = 1.0 - params.fp8_scale_quant_orig = 1.0 - - plan = native.ColdPageBufferPlan() - plan.role = role - plan.transform = native.ColdPageTransformKind.NVFP4 - plan.raw_bytes = raw_bytes_by_role[role] - plan.cold_data_offset = data_offsets[role] - plan.cold_scale_offset = scale_offsets[role] - plan.nvfp4_params = params - buffer_plans.append(plan) + scales = native.Nvfp4ColdPageScales() + scales.nvfp4_scale_orig_quant = orig_quant[scale_index] + scales.nvfp4_scale_quant_orig = quant_orig[scale_index] + scales.fp8_scale_orig_quant = 1.0 + scales.fp8_scale_quant_orig = 1.0 + + buffer_layout = native.Nvfp4ColdPageBufferLayout() + buffer_layout.role = role + buffer_layout.cold_data_offset = data_offsets[role] + buffer_layout.cold_scale_offset = scale_offsets[role] + buffer_layout.scales = scales + buffer_layouts.append(buffer_layout) for buffer in layer.buffers: role = str(buffer.role) if role in compressed_roles: continue - plan = native.ColdPageBufferPlan() - plan.role = role - plan.transform = native.ColdPageTransformKind.LOSSLESS_COPY - plan.raw_bytes = raw_bytes_by_role[role] - plan.cold_data_offset = cursor - plan.cold_scale_offset = 0 - buffer_plans.append(plan) - cursor += raw_bytes_by_role[role] + buffer_layout = native.Nvfp4ColdPageBufferLayout() + buffer_layout.role = role + buffer_layout.cold_data_offset = cursor + buffer_layouts.append(buffer_layout) + cursor += _buffer_bytes(buffer, tokens_per_page) cold_page_bytes = _align_up(cursor) - layer_plan = native.ColdPageLayerPlan() - layer_plan.layer_id = layer_id - layer_plan.cold_page_bytes = cold_page_bytes - layer_plan.cold_padding_offset = cursor - layer_plan.cold_padding_bytes = cold_page_bytes - cursor - layer_plan.buffers = buffer_plans - layer_plans.append(layer_plan) - - return native.create_cold_page_codec(layer_plans) + layer_layout = native.Nvfp4ColdPageLayerLayout() + layer_layout.layer_id = layer_id + layer_layout.runtime_type = runtime_type + layer_layout.num_kv_heads = num_kv_heads + layer_layout.tokens_per_page = tokens_per_page + layer_layout.head_dim = head_dim + layer_layout.cold_page_bytes = cold_page_bytes + layer_layout.cold_padding_offset = cursor + layer_layout.buffers = buffer_layouts + layer_layouts.append(layer_layout) + + return native.create_nvfp4_cold_page_codec(layer_layouts) diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 155258902f73..2c98a353617b 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -60,24 +60,20 @@ def _native(): def config(): return SimpleNamespace() - def buffer_plan(): - return SimpleNamespace(nvfp4_params=None) + def buffer_layout(): + return SimpleNamespace(scales=None, cold_scale_offset=0) codec = MagicMock() module = SimpleNamespace( - ColdPageTransformKind=SimpleNamespace( - NVFP4="native-nvfp4", - LOSSLESS_COPY="native-lossless", - ), Nvfp4ColdPageRuntimeType=SimpleNamespace( FLOAT16="native-fp16", BFLOAT16="native-bf16", FP8_E4M3="native-fp8", ), - Nvfp4ColdPageParams=config, - ColdPageBufferPlan=buffer_plan, - ColdPageLayerPlan=config, - create_cold_page_codec=MagicMock(return_value=codec), + Nvfp4ColdPageScales=config, + Nvfp4ColdPageBufferLayout=buffer_layout, + Nvfp4ColdPageLayerLayout=config, + create_nvfp4_cold_page_codec=MagicMock(return_value=codec), ) return module, codec @@ -133,15 +129,13 @@ def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_p ) assert result is codec - plans = native.create_cold_page_codec.call_args.args[0] + plans = native.create_nvfp4_cold_page_codec.call_args.args[0] assert [plan.layer_id for plan in plans] == [0, 1, 2] - assert [buffer.nvfp4_params.runtime_type for plan in plans for buffer in plan.buffers] == [ - "native-bf16" - ] * 6 + assert [plan.runtime_type for plan in plans] == ["native-bf16"] * 3 assert [ ( - tuple(buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers), - tuple(buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers), + tuple(buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers), + tuple(buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers), ) for plan in plans ] == [ @@ -164,7 +158,7 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: num_kv_heads_per_layer=(8,), head_dim_per_layer=(128,), ) - target_plan = native.create_cold_page_codec.call_args.args[0][0] + target_plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] manager.create_cold_page_codec( _cache_config((0, "attention")), runtime_dtype=DataType.BF16, @@ -174,16 +168,16 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: is_draft=True, ) - draft_plan = native.create_cold_page_codec.call_args.args[0][0] - assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in target_plan.buffers] == [ + draft_plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert [buffer.scales.nvfp4_scale_orig_quant for buffer in target_plan.buffers] == [ 2.0, 4.0, ] - assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_scale_orig_quant for buffer in draft_plan.buffers] == [ 1.0, 1.0, ] - assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_scale_quant_orig for buffer in draft_plan.buffers] == [ 1.0, 1.0, ] @@ -203,22 +197,21 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): head_dim_per_layer=(128,), ) - plan = native.create_cold_page_codec.call_args.args[0][0] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert [buffer.role for buffer in plan.buffers] == ["key", "value"] assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 1280] assert [buffer.cold_scale_offset for buffer in plan.buffers] == [2560, 2720] assert plan.cold_page_bytes == 2880 - assert plan.cold_padding_bytes == 0 - params = plan.buffers[0].nvfp4_params - assert params.runtime_type == "native-fp16" - assert params.num_kv_heads == 4 - assert params.tokens_per_page == 5 - assert params.head_dim == 128 - assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + assert plan.cold_padding_offset == plan.cold_page_bytes + assert plan.runtime_type == "native-fp16" + assert plan.num_kv_heads == 4 + assert plan.tokens_per_page == 5 + assert plan.head_dim == 128 + assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ 1.0, 1.0, ] - assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ 1.0, 1.0, ] @@ -248,19 +241,18 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: head_dim_per_layer=(32,), ) - plan = native.create_cold_page_codec.call_args.args[0][0] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert [buffer.role for buffer in plan.buffers] == ["key", "value"] assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 80] assert [buffer.cold_scale_offset for buffer in plan.buffers] == [160, 170] assert plan.cold_padding_offset == 180 - assert plan.cold_padding_bytes == 12 assert plan.cold_page_bytes == 192 def test_provider_creates_one_native_codec_per_kv_cache_manager(): native, _ = _native() codecs = (object(), object()) - native.create_cold_page_codec.side_effect = codecs + native.create_nvfp4_cold_page_codec.side_effect = codecs provider = _manager() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): @@ -276,7 +268,7 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): ) assert results == codecs - assert native.create_cold_page_codec.call_count == 2 + assert native.create_nvfp4_cold_page_codec.call_count == 2 def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: @@ -393,9 +385,9 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): num_kv_heads_per_layer=(0, 8), head_dim_per_layer=(128, 128), ) - plan = native.create_cold_page_codec.call_args.args[0][0] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert plan.layer_id == 1 - assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ 2.0, 4.0, ] @@ -409,7 +401,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - native.create_cold_page_codec.assert_called_with([]) + native.create_nvfp4_cold_page_codec.assert_called_with([]) def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -437,25 +429,21 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): ) assert result is codec - plan = native.create_cold_page_codec.call_args.args[0][0] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert plan.layer_id == 0 assert plan.cold_page_bytes == 29184 assert plan.cold_padding_offset == 29184 - assert plan.cold_padding_bytes == 0 assert [buffer.role for buffer in plan.buffers] == ["key", "index_key"] - assert [buffer.transform for buffer in plan.buffers] == [ - "native-nvfp4", - "native-lossless", - ] + assert [buffer.scales is not None for buffer in plan.buffers] == [True, False] assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 20736] assert [buffer.cold_scale_offset for buffer in plan.buffers] == [18432, 0] - params = plan.buffers[0].nvfp4_params - assert params.runtime_type == "native-bf16" - assert params.num_kv_heads == 1 - assert params.tokens_per_page == 64 - assert params.head_dim == 576 - assert params.nvfp4_scale_orig_quant == params.nvfp4_scale_quant_orig == 1.0 - assert plan.buffers[1].nvfp4_params is None + assert plan.runtime_type == "native-bf16" + assert plan.num_kv_heads == 1 + assert plan.tokens_per_page == 64 + assert plan.head_dim == 576 + scales = plan.buffers[0].scales + assert scales.nvfp4_scale_orig_quant == scales.nvfp4_scale_quant_orig == 1.0 + assert plan.buffers[1].scales is None def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: @@ -483,20 +471,15 @@ def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: head_dim_per_layer=(32,), ) - plan = native.create_cold_page_codec.call_args.args[0][0] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert [buffer.role for buffer in plan.buffers] == [ "key", "index_key", "rope_state", ] - assert [buffer.transform for buffer in plan.buffers] == [ - "native-nvfp4", - "native-lossless", - "native-lossless", - ] + assert [buffer.scales is not None for buffer in plan.buffers] == [True, False, False] assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 90, 158] assert plan.cold_padding_offset == 165 - assert plan.cold_padding_bytes == 11 assert plan.cold_page_bytes == 176 @@ -528,11 +511,9 @@ def test_tokens_per_block_override_expands_raw_and_lossless_bytes() -> None: head_dim_per_layer=(16,), ) - plan = native.create_cold_page_codec.call_args.args[0][0] - assert [buffer.raw_bytes for buffer in plan.buffers] == [128, 6] + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 36] assert plan.cold_padding_offset == 42 - assert plan.cold_padding_bytes == 6 assert plan.cold_page_bytes == 48 @@ -604,7 +585,7 @@ def test_mla_model_layouts_are_built_in_python( head_dim_per_layer=(576,) * expected_layers, ) - plans = native.create_cold_page_codec.call_args.args[0] + plans = native.create_nvfp4_cold_page_codec.call_args.args[0] assert len(plans) == expected_layers assert sum(len(plan.buffers) for plan in plans) == expected_buffers assert sum(plan.cold_page_bytes for plan in plans) == expected_bytes @@ -622,21 +603,18 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): head_dim_per_layer=(128,), ) - plan = native.create_cold_page_codec.call_args.args[0][0] - assert [buffer.nvfp4_params.runtime_type for buffer in plan.buffers] == [ - "native-fp8", - "native-fp8", - ] - assert [buffer.nvfp4_params.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + assert plan.runtime_type == "native-fp8" + assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ 2.0, 4.0, ] - assert [buffer.nvfp4_params.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ 0.5, 0.25, ] assert all( - buffer.nvfp4_params.fp8_scale_orig_quant == buffer.nvfp4_params.fp8_scale_quant_orig == 1.0 + buffer.scales.fp8_scale_orig_quant == buffer.scales.fp8_scale_quant_orig == 1.0 for buffer in plan.buffers ) From a2e177e828cd8e33fef9a4d8f4dab604b94c5323 Mon Sep 17 00:00:00 2001 From: tianruih Date: Tue, 25 Aug 2026 16:33:45 -0700 Subject: [PATCH 10/29] [None][test] Rely on KVCM buffer validation Signed-off-by: tianruih --- .../quantization_for_cold_page/nvfp4.py | 2 -- .../test_quantization_for_cold_page.py | 33 +------------------ 2 files changed, 1 insertion(+), 34 deletions(-) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py index 4dbb3cf0002f..7b7d95f321e0 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py @@ -97,8 +97,6 @@ def _align_up(value: int, alignment: int = _COLD_PAGE_ALIGNMENT) -> int: def _buffer_bytes(buffer: object, tokens_per_page: int) -> int: buffer_tokens = buffer.tokens_per_block_override or tokens_per_page - if buffer_tokens <= 0 or tokens_per_page % buffer_tokens != 0: - raise ValueError("tokens_per_block_override must be a positive divisor of tokens_per_block") return int(buffer.size) * (tokens_per_page // buffer_tokens) diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 2c98a353617b..dadc7e9af538 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -483,7 +483,7 @@ def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: assert plan.cold_page_bytes == 176 -def test_tokens_per_block_override_expands_raw_and_lossless_bytes() -> None: +def test_tokens_per_block_override_expands_lossless_bytes() -> None: native, _ = _native() cache_config = SimpleNamespace( tokens_per_block=4, @@ -517,37 +517,6 @@ def test_tokens_per_block_override_expands_raw_and_lossless_bytes() -> None: assert plan.cold_page_bytes == 48 -def test_tokens_per_block_override_must_divide_page_size() -> None: - native, _ = _native() - cache_config = SimpleNamespace( - tokens_per_block=5, - layers=( - AttentionLayerConfig( - layer_id=0, - buffers=[ - BufferConfig( - role="key", - size=64, - tokens_per_block_override=2, - ), - ], - ), - ), - ) - - with ( - patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native), - pytest.raises(ValueError, match="positive divisor"), - ): - _manager().create_cold_page_codec( - cache_config, - runtime_dtype=DataType.BF16, - pp_layers=(0,), - num_kv_heads_per_layer=(1,), - head_dim_per_layer=(16,), - ) - - @pytest.mark.parametrize( ("owns_index", "expected_layers", "expected_buffers", "expected_bytes"), [ From ba58b114c05cf773d30c403435ad80fd16f785a5 Mon Sep 17 00:00:00 2001 From: tianruih Date: Tue, 25 Aug 2026 17:21:49 -0700 Subject: [PATCH 11/29] [None][refactor] Simplify native cold-page codec Signed-off-by: tianruih --- .../kv_cache_compression/CMakeLists.txt | 2 +- .../nativeColdPageCodec.cpp | 95 ++++++------------ .../nativeColdPageCodec.h | 68 +++++++------ ...odecBackend.cpp => nvfp4ColdPageCodec.cpp} | 33 ++++--- ...ageCodecBackend.h => nvfp4ColdPageCodec.h} | 0 .../nanobind/kvCacheCompression/bindings.cpp | 2 +- .../kv_cache_compression/CMakeLists.txt | 2 +- .../coldPageCodecTest.cpp | 98 ++++++++++--------- 8 files changed, 138 insertions(+), 162 deletions(-) rename cpp/tensorrt_llm/kv_cache_compression/{nvfp4ColdPageCodecBackend.cpp => nvfp4ColdPageCodec.cpp} (85%) rename cpp/tensorrt_llm/kv_cache_compression/{nvfp4ColdPageCodecBackend.h => nvfp4ColdPageCodec.h} (100%) diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt index 756e80176dac..eb6c284f6759 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -2,7 +2,7 @@ # All rights reserved. SPDX-License-Identifier: Apache-2.0 add_library(kv_cache_compression_src OBJECT nativeColdPageCodec.cpp - nvfp4ColdPageCodecBackend.cpp) + nvfp4ColdPageCodec.cpp) target_include_directories( kv_cache_compression_src PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp index e30d4922efb4..98d2e1a056ae 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -9,7 +9,6 @@ #include "tensorrt_llm/common/logger.h" #include -#include #include #include #include @@ -19,27 +18,9 @@ namespace tensorrt_llm::kv_cache_compression namespace { -std::size_t checkedAdd(std::size_t lhs, std::size_t rhs) +ResolvedHotLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::SlotDescVariant const& variant) { - if (lhs > std::numeric_limits::max() - rhs) - { - throw std::overflow_error("GPU Slot buffer offsets overflow size_t"); - } - return lhs + rhs; -} - -std::uintptr_t checkedAddress(std::uintptr_t base, std::size_t offset) -{ - if (offset > std::numeric_limits::max() - base) - { - throw std::overflow_error("GPU buffer address overflows uintptr_t"); - } - return base + offset; -} - -ResolvedColdPageLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::SlotDescVariant const& variant) -{ - ResolvedColdPageLifecycle result; + ResolvedHotLifecycle result{variant.lifeCycleId, {}}; for (kv::PoolIndex poolIndex{0}; poolIndex < variant.coalescedBuffers.size(); ++poolIndex) { auto const& coalesced = variant.coalescedBuffers.at(poolIndex); @@ -47,16 +28,15 @@ ResolvedColdPageLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv: std::size_t offset = 0; for (auto const& bufferId : coalesced.bufferIds) { - auto& layer = result[bufferId.layerId]; + auto& layer = result.layers[bufferId.layerId]; if (!layer .emplace(bufferId.role, - ResolvedColdPageBuffer{ - checkedAddress(pool.baseAddress, offset), pool.slotBytes, coalesced.singleBufferSize}) + ResolvedHotBuffer{pool.baseAddress + offset, pool.slotBytes, coalesced.singleBufferSize}) .second) { throw std::invalid_argument("GPU lifecycle contains a duplicate buffer role"); } - offset = checkedAdd(offset, coalesced.singleBufferSize); + offset += coalesced.singleBufferSize; } } return result; @@ -64,15 +44,6 @@ ResolvedColdPageLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv: } // namespace -NativeColdPageCodec::NativeColdPageCodec(std::unique_ptr backend) - : mBackend(std::move(backend)) -{ - if (!mBackend) - { - throw std::invalid_argument("NativeColdPageCodec requires a backend"); - } -} - bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept { try @@ -83,10 +54,9 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG throw std::invalid_argument("Default lossless codec rejected GPU layouts"); } - auto const& backendLayerIds = mBackend->getLayerIds(); + auto const& algorithmLayerIds = getLayerIds(); std::map pendingGroups; - std::vector backendLifecycles; - std::vector backendLayerGroups; + std::vector algorithmLifecycles; std::set consumedLayers; for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) @@ -95,32 +65,31 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG for (auto const& variant : gpuDesc.slotDesc.variants) { auto resolved = resolveLifecycle(gpuDesc, variant); - auto const backendLayerCount = std::count_if(resolved.begin(), resolved.end(), - [&backendLayerIds](auto const& layer) { return backendLayerIds.count(layer.first) != 0U; }); + auto const algorithmLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), + [&algorithmLayerIds](auto const& layer) { return algorithmLayerIds.count(layer.first) != 0U; }); LayerGroupState state; - if (backendLayerCount == 0U) + if (algorithmLayerCount == 0U) { state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); state.pageIndexLocation = losslessCodec->queryPageIndexLocation(variant.lifeCycleId); } else { - if (backendLayerCount != resolved.size()) + if (algorithmLayerCount != resolved.layers.size()) { - throw std::invalid_argument("A lifecycle cannot mix backend-owned and fallback layers"); + throw std::invalid_argument("A lifecycle cannot mix algorithm-owned and fallback layers"); } - for (auto const& [layerId, buffers] : resolved) + for (auto const& [layerId, buffers] : resolved.layers) { static_cast(buffers); if (!consumedLayers.emplace(layerId).second) { - throw std::invalid_argument("A backend layer appears in multiple lifecycles"); + throw std::invalid_argument("An algorithm layer appears in multiple lifecycles"); } } - state.backendIndex = backendLifecycles.size(); - backendLayerGroups.push_back(variant.lifeCycleId); - backendLifecycles.push_back(std::move(resolved)); + state.planIndex = algorithmLifecycles.size(); + algorithmLifecycles.push_back(std::move(resolved)); } if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) @@ -129,26 +98,26 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG } } } - if (consumedLayers != backendLayerIds) + if (consumedLayers != algorithmLayerIds) { - throw std::invalid_argument("A backend layer is absent from all GPU descriptors"); + throw std::invalid_argument("An algorithm layer is absent from all GPU descriptors"); } - auto const backendConfigs = mBackend->configure(backendLifecycles); - if (backendConfigs.size() != backendLifecycles.size()) + auto const properties = configureAlgorithm(algorithmLifecycles); + if (properties.size() != algorithmLifecycles.size()) { - throw std::invalid_argument("Cold-page backend returned an unexpected lifecycle count"); + throw std::invalid_argument("Cold-page algorithm returned an unexpected lifecycle count"); } - for (std::size_t index = 0; index < backendConfigs.size(); ++index) + for (std::size_t index = 0; index < properties.size(); ++index) { - auto const& config = backendConfigs[index]; - if (config.coldPageBytes == 0U || config.pageIndexLocation == kv::PageIndexLocation::kBadLocation) + auto const& lifecycle = properties[index]; + if (lifecycle.coldPageBytes == 0U || lifecycle.pageIndexLocation == kv::PageIndexLocation::kBadLocation) { - throw std::invalid_argument("Cold-page backend returned invalid storage properties"); + throw std::invalid_argument("Cold-page algorithm returned invalid storage properties"); } - auto& state = pendingGroups.at(backendLayerGroups[index]); - state.coldPageBytes = config.coldPageBytes; - state.pageIndexLocation = config.pageIndexLocation; + auto& state = pendingGroups.at(algorithmLifecycles[index].lifeCycleId); + state.coldPageBytes = lifecycle.coldPageBytes; + state.pageIndexLocation = lifecycle.pageIndexLocation; } mLayerGroups = std::move(pendingGroups); @@ -205,11 +174,11 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr { return true; } - if (!state->backendIndex) + if (!state->planIndex) { return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); } - mBackend->encode(*state->backendIndex, dstBasePtr, pageIndices, numBasePages, stream); + encodeAlgorithm(*state->planIndex, dstBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) @@ -238,11 +207,11 @@ bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcB { return true; } - if (!state->backendIndex) + if (!state->planIndex) { return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); } - mBackend->decode(*state->backendIndex, srcBasePtr, pageIndices, numBasePages, stream); + decodeAlgorithm(*state->planIndex, srcBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h index c2e110d24634..4ea28a9b3dfe 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -22,74 +22,72 @@ namespace tensorrt_llm::kv_cache_compression namespace kv = batch_manager::kv_cache_manager_v2; //! One hot buffer resolved from KVCM's authoritative pool descriptors. -struct ResolvedColdPageBuffer +struct ResolvedHotBuffer { std::uintptr_t rawBase = 0; std::size_t rawSlotBytes = 0; std::size_t rawBytes = 0; }; -using ResolvedColdPageLayer = std::map; -using ResolvedColdPageLifecycle = std::map; +using ResolvedHotLayer = std::map; -//! Storage properties produced while a backend prepares one lifecycle. -struct ColdPageLifecycleConfig +//! One KVCM lifecycle resolved into its hot buffers. +struct ResolvedHotLifecycle +{ + kv::LifeCycleId lifeCycleId{-1}; + std::map layers; +}; + +//! Storage properties produced while an algorithm prepares one lifecycle. +struct ColdPageLifecycleProperties { std::size_t coldPageBytes = 0; kv::PageIndexLocation pageIndexLocation = kv::PageIndexLocation::kBadLocation; }; -//! Native algorithm backend retained by the generic KVCM codec adapter. -class IColdPageCodecBackend +//! Resolves KVCM layouts and routes lifecycles for one native compression method. +class NativeColdPageCodec : public kv::IKvCacheColdPageCodec { public: - virtual ~IColdPageCodecBackend() = default; + bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept final; + + [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept final; + + [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept final; + + [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept final; + bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept final; + + bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, + std::size_t numBasePages, cudaStream_t stream) noexcept final; + +private: [[nodiscard]] virtual std::set const& getLayerIds() const noexcept = 0; - virtual std::vector configure(std::vector const& lifecycles) + virtual std::vector configureAlgorithm( + std::vector const& lifecycles) = 0; - virtual void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + //! Enqueue only on stream; drain earlier partial submissions before throwing. + virtual void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; - virtual void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + virtual void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; -}; -//! Resolves KVCM layouts, routes lifecycles, and owns one native backend. -class NativeColdPageCodec final : public kv::IKvCacheColdPageCodec -{ -public: - explicit NativeColdPageCodec(std::unique_ptr backend); - - bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept override; - - [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::LayerGroupId getBatchingLayerGroupId(kv::LayerGroupId layerGroupId) const noexcept override; - - [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept override; - - bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - - bool decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, - std::size_t numBasePages, cudaStream_t stream) noexcept override; - -private: struct LayerGroupState { - std::optional backendIndex; + std::optional planIndex; std::size_t coldPageBytes = 0; kv::PageIndexLocation pageIndexLocation = kv::PageIndexLocation::kBadLocation; }; [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; - std::unique_ptr mBackend; std::map mLayerGroups; std::unique_ptr mLosslessCodec; }; diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp similarity index 85% rename from cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp rename to cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp index 0611e78f03ab..c7af9cb279c4 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp @@ -4,7 +4,7 @@ * SPDX-License-Identifier: Apache-2.0 */ -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" @@ -35,10 +35,10 @@ kernels::Nvfp4ColdPageKernelParams makeKernelParams( scales.nvfp4ScaleOrigQuant, scales.nvfp4ScaleQuantOrig, scales.fp8ScaleOrigQuant, scales.fp8ScaleQuantOrig}; } -class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend +class Nvfp4ColdPageCodec final : public NativeColdPageCodec { public: - explicit Nvfp4ColdPageCodecBackend(std::vector layerLayouts) + explicit Nvfp4ColdPageCodec(std::vector layerLayouts) { for (auto& layout : layerLayouts) { @@ -68,12 +68,13 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend return mLayerIds; } - std::vector configure(std::vector const& lifecycles) override + std::vector configureAlgorithm( + std::vector const& lifecycles) override { std::vector preparedPlans; - std::vector configs; + std::vector properties; preparedPlans.reserve(lifecycles.size()); - configs.reserve(lifecycles.size()); + properties.reserve(lifecycles.size()); for (auto const& lifecycle : lifecycles) { @@ -81,7 +82,7 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend std::optional runtimeType; std::size_t coldPageBytes = 0; - for (auto const& [layerId, resolvedBuffers] : lifecycle) + for (auto const& [layerId, resolvedBuffers] : lifecycle.layers) { auto const& layout = mLayerLayouts.at(layerId); if (layout.buffers.size() != resolvedBuffers.size()) @@ -129,18 +130,18 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend if (!runtimeType) { - throw std::invalid_argument("NVFP4 backend received an empty lifecycle"); + throw std::invalid_argument("NVFP4 received an empty lifecycle"); } preparedPlans.push_back(kernels::prepareNvfp4ColdPagePlan(buffers, coldPageBytes, *runtimeType)); - configs.push_back({coldPageBytes, kv::PageIndexLocation::kHost}); + properties.push_back({coldPageBytes, kv::PageIndexLocation::kHost}); } mPreparedPlans = std::move(preparedPlans); - return configs; + return properties; } - void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, - cudaStream_t stream) override + void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override { thread_local std::vector pages; pages.clear(); @@ -149,10 +150,10 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend { pages.push_back({pageIndices[page].src, pageIndices[page].dst}); } - kernels::invokeNvfp4ColdPageEncode(pages, mPreparedPlans.at(lifecycleIndex), coldBase, stream); + kernels::invokeNvfp4ColdPageEncode(pages, mPreparedPlans.at(planIndex), coldBase, stream); } - void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { thread_local std::vector pages; @@ -162,7 +163,7 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend { pages.push_back({pageIndices[page].dst, pageIndices[page].src}); } - kernels::invokeNvfp4ColdPageDecode(pages, mPreparedPlans.at(lifecycleIndex), coldBase, stream); + kernels::invokeNvfp4ColdPageDecode(pages, mPreparedPlans.at(planIndex), coldBase, stream); } private: @@ -175,7 +176,7 @@ class Nvfp4ColdPageCodecBackend final : public IColdPageCodecBackend std::unique_ptr createNvfp4ColdPageCodec(std::vector layerLayouts) { - return std::make_unique(std::make_unique(std::move(layerLayouts))); + return std::make_unique(std::move(layerLayouts)); } } // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h similarity index 100% rename from cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h rename to cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index e5b4816d2d7d..07a92305270b 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,7 +17,7 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" #include #include diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index 9962aa05df62..bada5d6e6578 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -7,7 +7,7 @@ set(COLD_PAGE_CODEC_TEST_SRC ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp index ca692381bc43..5786189eac58 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -5,7 +5,7 @@ */ #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodecBackend.h" +#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" #include @@ -13,7 +13,6 @@ #include #include #include -#include #include #include #include @@ -27,6 +26,7 @@ namespace namespace kv = batch_manager::kv_cache_manager_v2; +static_assert(std::is_abstract_v); static_assert(std::is_base_of_v); constexpr std::uintptr_t kGpuKBase = 0x100000; @@ -81,10 +81,10 @@ bool configureOne(kv::IKvCacheColdPageCodec& codec, kv::PoolGroupDesc const& des return codec.configure(&desc, kv::PoolGroupIndex{1}); } -class RecordingBackend final : public IColdPageCodecBackend +class RecordingCodec final : public NativeColdPageCodec { public: - explicit RecordingBackend(std::set layerIds) + explicit RecordingCodec(std::set layerIds) : mLayerIds(std::move(layerIds)) { } @@ -94,45 +94,55 @@ class RecordingBackend final : public IColdPageCodecBackend return mLayerIds; } - std::vector configure(std::vector const& lifecycles) override + std::vector configureAlgorithm( + std::vector const& lifecycles) override { resolved = lifecycles; if (failConfigure) { throw std::runtime_error("requested configure failure"); } - return std::vector( - lifecycles.size(), ColdPageLifecycleConfig{777U, kv::PageIndexLocation::kHost}); + return std::vector( + lifecycles.size(), ColdPageLifecycleProperties{777U, kv::PageIndexLocation::kHost}); } - void encode(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, - cudaStream_t stream) override + void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override { + if (failBatches) + { + throw std::runtime_error("requested batch failure"); + } ++encodeCalls; - lastLifecycleIndex = lifecycleIndex; + lastPlanIndex = planIndex; lastColdBase = coldBase; lastIndices.assign(pageIndices, pageIndices + numPages); lastStream = stream; } - void decode(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { + if (failBatches) + { + throw std::runtime_error("requested batch failure"); + } ++decodeCalls; - lastLifecycleIndex = lifecycleIndex; + lastPlanIndex = planIndex; lastColdBase = coldBase; lastIndices.assign(pageIndices, pageIndices + numPages); lastStream = stream; } bool failConfigure = false; + bool failBatches = false; int encodeCalls = 0; int decodeCalls = 0; - std::size_t lastLifecycleIndex = 0; + std::size_t lastPlanIndex = 0; void const* lastColdBase = nullptr; cudaStream_t lastStream{}; std::vector lastIndices; - std::vector resolved; + std::vector resolved; private: std::set mLayerIds; @@ -140,13 +150,12 @@ class RecordingBackend final : public IColdPageCodecBackend TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsBatches) { - auto backend = std::make_unique(std::set{0, 1}); - auto* recorder = backend.get(); - NativeColdPageCodec codec{std::move(backend)}; + RecordingCodec codec{{0, 1}}; ASSERT_TRUE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{3}, 2U))); - ASSERT_EQ(recorder->resolved.size(), 1U); - auto const& layers = recorder->resolved.front(); + ASSERT_EQ(codec.resolved.size(), 1U); + EXPECT_EQ(codec.resolved.front().lifeCycleId, kv::LifeCycleId{3}); + auto const& layers = codec.resolved.front().layers; EXPECT_EQ(layers.at(0).at("key").rawBase, kGpuKBase); EXPECT_EQ(layers.at(0).at("value").rawBase, kGpuVBase); EXPECT_EQ(layers.at(1).at("key").rawBase, kGpuKBase + kRawBytes); @@ -158,26 +167,24 @@ TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsBatches) auto const stream = reinterpret_cast(kStreamValue); ASSERT_TRUE( codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - EXPECT_EQ(recorder->encodeCalls, 1); - EXPECT_EQ(recorder->lastLifecycleIndex, 0U); - EXPECT_EQ(recorder->lastIndices[1].src, 3); - EXPECT_EQ(recorder->lastColdBase, reinterpret_cast(kColdBase)); - EXPECT_EQ(recorder->lastStream, stream); + EXPECT_EQ(codec.encodeCalls, 1); + EXPECT_EQ(codec.lastPlanIndex, 0U); + EXPECT_EQ(codec.lastIndices[1].src, 3); + EXPECT_EQ(codec.lastColdBase, reinterpret_cast(kColdBase)); + EXPECT_EQ(codec.lastStream, stream); ASSERT_TRUE(codec.decode( kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - EXPECT_EQ(recorder->decodeCalls, 1); + EXPECT_EQ(codec.decodeCalls, 1); } TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) { - auto backend = std::make_unique(std::set{0}); - auto* recorder = backend.get(); - NativeColdPageCodec codec{std::move(backend)}; + RecordingCodec codec{{0}}; std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - EXPECT_EQ(recorder->resolved.size(), 1U); + EXPECT_EQ(codec.resolved.size(), 1U); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 777U); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{1}), 96U); EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{1}), kv::PageIndexLocation::kHost); @@ -186,44 +193,45 @@ TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) { { - NativeColdPageCodec codec{std::make_unique(std::set{0})}; + RecordingCodec codec{{0}}; EXPECT_FALSE(configureOne(codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{0}, 2U))); } { - NativeColdPageCodec codec{std::make_unique(std::set{0, 1})}; + RecordingCodec codec{{0, 1}}; EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); } { - NativeColdPageCodec codec{std::make_unique(std::set{0, 1})}; + RecordingCodec codec{{0, 1}}; std::array descs{makeAttentionDesc(), makeAttentionDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{0}, 1U, 1, 0x600000, 0x700000)}; EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); } } -TEST(NativeColdPageCodecTest, CatchesBackendConfigureAndBatchFailures) +TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureAndBatchFailures) { - auto backend = std::make_unique(std::set{0}); - backend->failConfigure = true; - NativeColdPageCodec codec{std::move(backend)}; + RecordingCodec codec{{0}}; + codec.failConfigure = true; EXPECT_FALSE(configureOne(codec, makeAttentionDesc())); - auto validBackend = std::make_unique(std::set{0}); - auto* recorder = validBackend.get(); - NativeColdPageCodec validCodec{std::move(validBackend)}; + RecordingCodec validCodec{{0}}; ASSERT_TRUE(configureOne(validCodec, makeAttentionDesc())); EXPECT_TRUE(validCodec.encode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); EXPECT_TRUE(validCodec.decode(kv::LayerGroupId{0}, nullptr, nullptr, 0U, nullptr)); kv::PageIndexPair const indices[]{{0, 0}}; EXPECT_FALSE(validCodec.encode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); EXPECT_FALSE(validCodec.decode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); - EXPECT_EQ(recorder->encodeCalls, 0); - EXPECT_EQ(recorder->decodeCalls, 0); + validCodec.failBatches = true; + EXPECT_FALSE(validCodec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_FALSE( + validCodec.decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); + EXPECT_EQ(validCodec.encodeCalls, 0); + EXPECT_EQ(validCodec.decodeCalls, 0); } TEST(NativeColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) { - NativeColdPageCodec codec{std::make_unique(std::set{0})}; + RecordingCodec codec{{0}}; ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); @@ -264,7 +272,7 @@ struct RecordedLaunch RecordedLaunch gLaunch; -TEST(Nvfp4ColdPageCodecBackendTest, LowersMhaLayoutOnceAndDispatchesEncodeDecode) +TEST(Nvfp4ColdPageCodecTest, LowersMhaLayoutOnceAndDispatchesEncodeDecode) { gLaunch = {}; auto codec = createNvfp4ColdPageCodec({makeAttentionLayout(0, 2.0F, 3.0F), makeAttentionLayout(1, 4.0F, 5.0F)}); @@ -296,7 +304,7 @@ TEST(Nvfp4ColdPageCodecBackendTest, LowersMhaLayoutOnceAndDispatchesEncodeDecode EXPECT_EQ(gLaunch.decodePages[0].coldPageIndex, 1); } -TEST(Nvfp4ColdPageCodecBackendTest, PreservesMlaSideBufferLosslessly) +TEST(Nvfp4ColdPageCodecTest, PreservesMlaSideBufferLosslessly) { gLaunch = {}; auto codec = createNvfp4ColdPageCodec({makeMlaLayout()}); @@ -314,7 +322,7 @@ TEST(Nvfp4ColdPageCodecBackendTest, PreservesMlaSideBufferLosslessly) EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingBytes, 2U); } -TEST(Nvfp4ColdPageCodecBackendTest, RejectsDuplicateLayoutsAndMissingRoles) +TEST(Nvfp4ColdPageCodecTest, RejectsDuplicateLayoutsAndMissingRoles) { EXPECT_THROW( { From c0d6c43d155683c3f17f9f76284700cb6c8b859d Mon Sep 17 00:00:00 2001 From: tianruih Date: Tue, 25 Aug 2026 20:29:32 -0700 Subject: [PATCH 12/29] [None][refactor] Dispatch cold-page compression through Python --- .../kernels/nvfp4ColdPageKernels.cu | 138 +++++------ .../kernels/nvfp4ColdPageKernels.h | 41 ++-- .../kv_cache_compression/CMakeLists.txt | 3 +- .../coldPageCallbackAbi.h | 29 +++ .../nativeColdPageCodec.cpp | 31 +++ .../nativeColdPageCodec.h | 2 +- .../nvfp4ColdPageCodec.cpp | 182 -------------- .../kv_cache_compression/nvfp4ColdPageCodec.h | 58 ----- .../nanobind/kvCacheCompression/bindings.cpp | 154 +++++++++--- cpp/tensorrt_llm/thop/CMakeLists.txt | 1 + cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp | 229 ++++++++++++++++++ .../kernels/nvfp4ColdPageKernelsTest.cpp | 101 +++++--- .../kv_cache_compression/CMakeLists.txt | 1 - .../coldPageCodecTest.cpp | 215 ++++------------ .../quantization_for_cold_page/nvfp4.py | 224 +++++++++++++---- .../quantization_for_cold_page.py | 68 +++++- .../test_quantization_for_cold_page.py | 170 ++++++++----- 17 files changed, 947 insertions(+), 700 deletions(-) create mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h create mode 100644 cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index d24bc4144fe5..0651c8560e8a 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -28,10 +28,10 @@ #include #include #include +#include #include #include #include -#include #include #include #include @@ -73,16 +73,11 @@ static_assert( static_assert(std::is_trivially_copyable_v, "Buffer plans must remain raw-copyable"); // Keep both kernel argument packs within CUDA's modern 32,764-byte limit. -static_assert(sizeof(std::array) +static_assert(sizeof(std::array) + sizeof(std::array) + 2U * sizeof(std::uintptr_t) + 3U * sizeof(std::uint32_t) <= kKernelParameterLimitBytes, - "Offload kernel arguments exceed CUDA's parameter limit"); -static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) - + 3U * sizeof(std::uint32_t) - <= kKernelParameterLimitBytes, - "Onboard kernel arguments exceed CUDA's parameter limit"); + "Cold-page kernel arguments exceed CUDA's parameter limit"); // Device data path. @@ -115,20 +110,20 @@ struct OnboardBufferTask std::uint8_t* raw; }; -__device__ OffloadBufferTask resolveTask(Nvfp4ColdPageOffloadPageTask const& page, - Nvfp4ColdPageBufferPlan const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) +__device__ OffloadBufferTask resolveOffloadTask(ColdPageIndexPair const& page, Nvfp4ColdPageBufferPlan const& buffer, + std::uint8_t* coldBase, std::size_t coldPageBytes) { - std::size_t const gpuPage = static_cast(page.gpuPageIndex); - auto* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; + std::size_t const gpuPage = static_cast(page.src); + auto* coldPage = coldBase + static_cast(page.dst) * coldPageBytes; return {reinterpret_cast(buffer.rawBase + gpuPage * buffer.rawSlotBytes), coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, coldPage + buffer.coldPaddingOffset}; } -__device__ OnboardBufferTask resolveTask(Nvfp4ColdPageOnboardPageTask const& page, - Nvfp4ColdPageBufferPlan const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) +__device__ OnboardBufferTask resolveOnboardTask(ColdPageIndexPair const& page, Nvfp4ColdPageBufferPlan const& buffer, + std::uint8_t const* coldBase, std::size_t coldPageBytes) { - std::size_t const gpuPage = static_cast(page.gpuPageIndex); - auto const* coldPage = coldBase + static_cast(page.coldPageIndex) * coldPageBytes; + std::size_t const gpuPage = static_cast(page.dst); + auto const* coldPage = coldBase + static_cast(page.src) * coldPageBytes; return {coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, reinterpret_cast(buffer.rawBase + gpuPage * buffer.rawSlotBytes)}; } @@ -410,7 +405,7 @@ __device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firs // FP16/BF16 GPU Page -> mapped-Host NVFP4 in bounded tiles. template __global__ void offloadFrom16BitTiledKernel( - std::array const __grid_constant__ pages, + std::array const __grid_constant__ pages, std::array const __grid_constant__ buffers, std::uint8_t* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { @@ -420,7 +415,7 @@ __global__ void offloadFrom16BitTiledKernel( std::uint32_t const bufferIndex = blockIdx.y; assert(bufferIndex < numBuffers); auto const& buffer = buffers[bufferIndex]; - auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + auto const task = resolveOffloadTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) @@ -491,7 +486,7 @@ __global__ void offloadFrom16BitTiledKernel( // FP8 E4M3 GPU Page -> mapped-Host NVFP4 in bounded tiles. __global__ void offloadFromFp8TiledKernel( - std::array const __grid_constant__ pages, + std::array const __grid_constant__ pages, std::array const __grid_constant__ buffers, std::uint8_t* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { @@ -501,7 +496,7 @@ __global__ void offloadFromFp8TiledKernel( std::uint32_t const bufferIndex = blockIdx.y; assert(bufferIndex < numBuffers); auto const& buffer = buffers[bufferIndex]; - auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + auto const task = resolveOffloadTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) @@ -634,8 +629,7 @@ __device__ void loadCompactRangeFromHost(std::uint8_t* compactStages, OnboardBuf // Mapped-Host NVFP4 -> runtime GPU Page in bounded tiles. template -__global__ void onboardTiledKernel( - std::array const __grid_constant__ pages, +__global__ void onboardTiledKernel(std::array const __grid_constant__ pages, std::array const __grid_constant__ buffers, std::uint8_t const* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) { @@ -645,7 +639,7 @@ __global__ void onboardTiledKernel( std::uint32_t const bufferIndex = blockIdx.y; assert(bufferIndex < numBuffers); auto const& buffer = buffers[bufferIndex]; - auto const task = resolveTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); + auto const task = resolveOnboardTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); if (buffer.transform == Nvfp4ColdPageTransform::kLosslessCopy) @@ -793,29 +787,31 @@ void validateBufferPlan(Nvfp4ColdPageBufferPlan const& buffer, std::size_t coldP // Host submission path. -// Submit Page tasks through the fixed 256-descriptor kernel ABI. -template -void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPagePreparedPlan const& plan, +// Submit one whole KVCM Page batch through the fixed 256-descriptor kernel ABI. +template +void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, ColdPointer coldBase, cudaStream_t stream) { - static_assert(std::is_trivially_copyable_v>, - "Page tasks must remain raw-copyable kernel arguments"); - static_assert(sizeof(std::array) == sizeof(Task) * kMaxTasksPerLaunch, - "Page task arrays must not add ABI padding"); + static_assert(std::is_trivially_copyable_v>, + "Page-index pairs must remain raw-copyable kernel arguments"); + static_assert( + sizeof(std::array) == sizeof(ColdPageIndexPair) * kMaxTasksPerLaunch, + "Page-index pair arrays must not add ABI padding"); dim3 const block(kThreadsPerBlock); + auto const* pageBytes = static_cast(pages); std::size_t offset = 0; - while (offset < tasks.size()) + while (offset < numPages) { - std::uint32_t const numChunkTasks - = static_cast(std::min(tasks.size() - offset, kMaxTasksPerLaunch)); - Task const* chunkTasks = tasks.data() + offset; + std::uint32_t const numChunkPages + = static_cast(std::min(numPages - offset, kMaxTasksPerLaunch)); + auto const* chunkPages = pageBytes + offset * sizeof(ColdPageIndexPair); // CUDA copies the full by-value array, so pad only the final partial chunk. - std::array paddedTasks{}; - if (numChunkTasks < kMaxTasksPerLaunch) + std::array paddedPages{}; + if (numChunkPages < kMaxTasksPerLaunch) { - std::copy_n(chunkTasks, numChunkTasks, paddedTasks.begin()); - chunkTasks = paddedTasks.data(); + std::memcpy(paddedPages.data(), chunkPages, numChunkPages * sizeof(ColdPageIndexPair)); + chunkPages = reinterpret_cast(paddedPages.data()); } cudaLaunchAttribute attribute{}; @@ -823,7 +819,7 @@ void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPa attribute.val.programmaticStreamSerializationAllowed = common::getEnvEnablePDL() ? 1 : 0; cudaLaunchConfig_t config{}; - config.gridDim = dim3(kMappedHostGridSplits, plan.numBuffers, numChunkTasks); + config.gridDim = dim3(kMappedHostGridSplits, plan.numBuffers, numChunkPages); config.blockDim = block; config.dynamicSmemBytes = compactStageBytesForHalfGroups(plan.maxHalfGroupsPerTile); config.stream = stream; @@ -832,32 +828,10 @@ void launchTaskChunks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPa auto coldPageBytes = plan.coldPageBytes; auto numBuffers = plan.numBuffers; - void* arguments[] = {const_cast(chunkTasks), const_cast(plan.buffers.data()), - &coldBase, &coldPageBytes, &numBuffers}; + void* arguments[] = {const_cast(chunkPages), + const_cast(plan.buffers.data()), &coldBase, &coldPageBytes, &numBuffers}; TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); - offset += numChunkTasks; - } -} - -// Drain earlier chunks after a later synchronous launch failure. -template -void submitColdPageTasks(Kernel kernel, std::vector const& tasks, Nvfp4ColdPagePreparedPlan const& plan, - ColdPointer coldBase, cudaStream_t stream) -{ - try - { - launchTaskChunks(kernel, tasks, plan, coldBase, stream); - } - catch (...) - { - cudaError_t const drainStatus = cudaStreamSynchronize(stream); - if (drainStatus != cudaSuccess) - { - // An asynchronous drain failure leaves Slot ownership unknown; fail-stop. - TLLM_LOG_ERROR("NVFP4 cold-page failure drain failed: %s", cudaGetErrorString(drainStatus)); - std::terminate(); - } - throw; + offset += numChunkPages; } } @@ -910,51 +884,55 @@ Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageEncode( + void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) { - if (pages.empty()) + if (numPages == 0) { return; } + TLLM_CHECK_WITH_INFO(pages != nullptr, "pages must not be null"); TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); switch (plan.runtimeType) { case Nvfp4ColdPageRuntimeType::kFloat16: - submitColdPageTasks( - offloadFrom16BitTiledKernel, pages, plan, static_cast(coldBase), stream); + launchPageChunks( + offloadFrom16BitTiledKernel, pages, numPages, plan, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kBfloat16: - submitColdPageTasks( - offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); + launchPageChunks(offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, numPages, plan, + static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kFp8E4m3: - submitColdPageTasks(offloadFromFp8TiledKernel, pages, plan, static_cast(coldBase), stream); + launchPageChunks( + offloadFromFp8TiledKernel, pages, numPages, plan, static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } } -void invokeNvfp4ColdPageDecode(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, + void const* coldBase, cudaStream_t stream) { - if (pages.empty()) + if (numPages == 0) { return; } + TLLM_CHECK_WITH_INFO(pages != nullptr, "pages must not be null"); TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); switch (plan.runtimeType) { case Nvfp4ColdPageRuntimeType::kFloat16: - submitColdPageTasks(onboardTiledKernel, pages, plan, static_cast(coldBase), stream); + launchPageChunks( + onboardTiledKernel, pages, numPages, plan, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kBfloat16: - submitColdPageTasks( - onboardTiledKernel<__nv_bfloat16>, pages, plan, static_cast(coldBase), stream); + launchPageChunks(onboardTiledKernel<__nv_bfloat16>, pages, numPages, plan, + static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kFp8E4m3: - submitColdPageTasks( - onboardTiledKernel<__nv_fp8_e4m3>, pages, plan, static_cast(coldBase), stream); + launchPageChunks(onboardTiledKernel<__nv_fp8_e4m3>, pages, numPages, plan, + static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h index 36af8c564489..a3eb70d7116a 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h @@ -19,6 +19,7 @@ #pragma once #include "tensorrt_llm/common/config.h" +#include "tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h" #include #include @@ -31,27 +32,15 @@ TRTLLM_NAMESPACE_BEGIN namespace kernels { -//! Active GPU representation encoded into cold Pages. +//! Active GPU representation encoded into cold Pages; values are part of the Python custom-op ABI. enum class Nvfp4ColdPageRuntimeType : std::uint8_t { - kFloat16, - kBfloat16, - kFp8E4m3, + kFloat16 = 0, + kBfloat16 = 1, + kFp8E4m3 = 2, }; -//! One Base Page selected for GPU-to-Host transformation. -struct Nvfp4ColdPageOffloadPageTask -{ - std::int32_t gpuPageIndex; - std::int32_t coldPageIndex; -}; - -//! One Base Page selected for Host-to-GPU transformation. -struct Nvfp4ColdPageOnboardPageTask -{ - std::int32_t gpuPageIndex; - std::int32_t coldPageIndex; -}; +using ColdPageIndexPair = ::tensorrt_llm::kv_cache_compression::ColdPageIndexPair; //! Per-buffer geometry and scales for one NVFP4 record in HND order. //! `headDim` is a multiple of 16; `*OrigQuant` encodes and `*QuantOrig` decodes this buffer. @@ -66,12 +55,12 @@ struct Nvfp4ColdPageKernelParams float fp8ScaleQuantOrig; }; -//! Transformation applied to one independently addressed hot buffer. +//! Per-buffer transform selected by the Python custom-op metadata ABI. enum class Nvfp4ColdPageTransform : std::uint8_t { - kNvfp4, + kNvfp4 = 0, //! Byte-exact copy for an Attention side buffer such as DSA index_key. - kLosslessCopy, + kLosslessCopy = 1, }; //! Immutable transform plan for one hot buffer and its fixed-offset cold record. @@ -104,13 +93,13 @@ struct Nvfp4ColdPagePreparedPlan [[nodiscard]] Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType); -//! Compress GPU Pages into mapped-Host NVFP4 records. -void invokeNvfp4ColdPageEncode(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream); +//! Compress one whole KVCM Page-index batch; the launcher performs 256-Page chunking internally. +void invokeNvfp4ColdPageEncode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, + void* coldBase, cudaStream_t stream); -//! Restore mapped-Host NVFP4 records into GPU Pages. -void invokeNvfp4ColdPageDecode(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void const* coldBase, cudaStream_t stream); +//! Restore one whole KVCM Page-index batch; the launcher performs 256-Page chunking internally. +void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, + void const* coldBase, cudaStream_t stream); } // namespace kernels diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt index eb6c284f6759..2a95aba2807c 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt @@ -1,8 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. # All rights reserved. SPDX-License-Identifier: Apache-2.0 -add_library(kv_cache_compression_src OBJECT nativeColdPageCodec.cpp - nvfp4ColdPageCodec.cpp) +add_library(kv_cache_compression_src OBJECT nativeColdPageCodec.cpp) target_include_directories( kv_cache_compression_src PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h b/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h new file mode 100644 index 000000000000..8958961f50ae --- /dev/null +++ b/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h @@ -0,0 +1,29 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include +#include +#include + +namespace tensorrt_llm::kv_cache_compression +{ + +//! Algorithm-neutral descriptor ABI borrowed by one cold-page callback. +struct alignas(8) ColdPageIndexPair +{ + std::int32_t dst; + std::int32_t src; +}; + +static_assert(sizeof(ColdPageIndexPair) == 8); +static_assert(alignof(ColdPageIndexPair) == 8); +static_assert(offsetof(ColdPageIndexPair, dst) == 0); +static_assert(offsetof(ColdPageIndexPair, src) == 4); +static_assert(std::is_trivially_copyable_v); + +} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp index 98d2e1a056ae..2e7222d7675d 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -9,6 +9,7 @@ #include "tensorrt_llm/common/logger.h" #include +#include #include #include #include @@ -42,6 +43,16 @@ ResolvedHotLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::Slot return result; } +void drainAfterAlgorithmFailure(cudaStream_t stream) noexcept +{ + auto const status = cudaStreamSynchronize(stream); + if (status != cudaSuccess) + { + TLLM_LOG_ERROR("Cold-page algorithm rollback drain failed: %s", cudaGetErrorString(status)); + std::terminate(); + } +} + } // namespace bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept @@ -163,6 +174,7 @@ kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroup bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { + bool algorithmStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -178,16 +190,25 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr { return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); } + algorithmStarted = true; encodeAlgorithm(*state->planIndex, dstBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { + if (algorithmStarted) + { + drainAfterAlgorithmFailure(stream); + } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: %s", error.what()); return false; } catch (...) { + if (algorithmStarted) + { + drainAfterAlgorithmFailure(stream); + } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: unknown error"); return false; } @@ -196,6 +217,7 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { + bool algorithmStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -211,16 +233,25 @@ bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcB { return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); } + algorithmStarted = true; decodeAlgorithm(*state->planIndex, srcBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { + if (algorithmStarted) + { + drainAfterAlgorithmFailure(stream); + } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: %s", error.what()); return false; } catch (...) { + if (algorithmStarted) + { + drainAfterAlgorithmFailure(stream); + } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: unknown error"); return false; } diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h index 4ea28a9b3dfe..fc0d8810365e 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -70,7 +70,7 @@ class NativeColdPageCodec : public kv::IKvCacheColdPageCodec std::vector const& lifecycles) = 0; - //! Enqueue only on stream; drain earlier partial submissions before throwing. + //! Enqueue only on stream; this codec drains partial submissions after a throw. virtual void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp deleted file mode 100644 index c7af9cb279c4..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp +++ /dev/null @@ -1,182 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" - -#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" - -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ -namespace -{ - -std::size_t checkedAdd(std::size_t lhs, std::size_t rhs, char const* label) -{ - if (lhs > std::numeric_limits::max() - rhs) - { - throw std::overflow_error(label); - } - return lhs + rhs; -} - -kernels::Nvfp4ColdPageKernelParams makeKernelParams( - Nvfp4ColdPageLayerLayout const& layout, Nvfp4ColdPageScales const& scales) -{ - return kernels::Nvfp4ColdPageKernelParams{layout.numKvHeads, layout.tokensPerPage, layout.headDim, - scales.nvfp4ScaleOrigQuant, scales.nvfp4ScaleQuantOrig, scales.fp8ScaleOrigQuant, scales.fp8ScaleQuantOrig}; -} - -class Nvfp4ColdPageCodec final : public NativeColdPageCodec -{ -public: - explicit Nvfp4ColdPageCodec(std::vector layerLayouts) - { - for (auto& layout : layerLayouts) - { - if (layout.coldPageBytes == 0U || layout.buffers.empty() || layout.coldPaddingOffset > layout.coldPageBytes) - { - throw std::invalid_argument("An NVFP4 layer layout must contain a valid cold-page record"); - } - std::set roles; - for (auto const& buffer : layout.buffers) - { - if (buffer.role.empty() || !roles.emplace(buffer.role).second) - { - throw std::invalid_argument("NVFP4 layer layouts require unique non-empty buffer roles"); - } - } - auto const layerId = layout.layerId; - if (!mLayerLayouts.emplace(layerId, std::move(layout)).second) - { - throw std::invalid_argument("NVFP4 layer layout IDs must be unique"); - } - mLayerIds.emplace(layerId); - } - } - - [[nodiscard]] std::set const& getLayerIds() const noexcept override - { - return mLayerIds; - } - - std::vector configureAlgorithm( - std::vector const& lifecycles) override - { - std::vector preparedPlans; - std::vector properties; - preparedPlans.reserve(lifecycles.size()); - properties.reserve(lifecycles.size()); - - for (auto const& lifecycle : lifecycles) - { - std::vector buffers; - std::optional runtimeType; - std::size_t coldPageBytes = 0; - - for (auto const& [layerId, resolvedBuffers] : lifecycle.layers) - { - auto const& layout = mLayerLayouts.at(layerId); - if (layout.buffers.size() != resolvedBuffers.size()) - { - throw std::invalid_argument("NVFP4 layer layout must cover every GPU buffer role"); - } - if (runtimeType && *runtimeType != layout.runtimeType) - { - throw std::invalid_argument("One lifecycle must use one NVFP4 runtime type"); - } - runtimeType = layout.runtimeType; - - for (auto const& bufferLayout : layout.buffers) - { - auto const resolved = resolvedBuffers.find(bufferLayout.role); - if (resolved == resolvedBuffers.end()) - { - throw std::invalid_argument("NVFP4 layer layout references an absent GPU buffer role"); - } - auto const& raw = resolved->second; - auto const transform = bufferLayout.scales ? kernels::Nvfp4ColdPageTransform::kNvfp4 - : kernels::Nvfp4ColdPageTransform::kLosslessCopy; - buffers.push_back(kernels::Nvfp4ColdPageBufferPlan{raw.rawBase, raw.rawSlotBytes, raw.rawBytes, - checkedAdd( - coldPageBytes, bufferLayout.coldDataOffset, "Cold-page data offset overflows size_t"), - bufferLayout.scales ? checkedAdd( - coldPageBytes, bufferLayout.coldScaleOffset, "Cold-page scale offset overflows size_t") - : 0U, - 0U, 0U, transform, - bufferLayout.scales ? makeKernelParams(layout, *bufferLayout.scales) - : kernels::Nvfp4ColdPageKernelParams{}}); - } - - auto const paddingBytes = layout.coldPageBytes - layout.coldPaddingOffset; - if (paddingBytes > std::numeric_limits::max()) - { - throw std::overflow_error("Cold-page padding exceeds the kernel ABI"); - } - buffers.back().coldPaddingOffset - = checkedAdd(coldPageBytes, layout.coldPaddingOffset, "Cold-page padding offset overflows size_t"); - buffers.back().coldPaddingBytes = static_cast(paddingBytes); - coldPageBytes - = checkedAdd(coldPageBytes, layout.coldPageBytes, "Lifecycle cold-page size overflows size_t"); - } - - if (!runtimeType) - { - throw std::invalid_argument("NVFP4 received an empty lifecycle"); - } - preparedPlans.push_back(kernels::prepareNvfp4ColdPagePlan(buffers, coldPageBytes, *runtimeType)); - properties.push_back({coldPageBytes, kv::PageIndexLocation::kHost}); - } - - mPreparedPlans = std::move(preparedPlans); - return properties; - } - - void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, - std::size_t numPages, cudaStream_t stream) override - { - thread_local std::vector pages; - pages.clear(); - pages.reserve(numPages); - for (std::size_t page = 0; page < numPages; ++page) - { - pages.push_back({pageIndices[page].src, pageIndices[page].dst}); - } - kernels::invokeNvfp4ColdPageEncode(pages, mPreparedPlans.at(planIndex), coldBase, stream); - } - - void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, - std::size_t numPages, cudaStream_t stream) override - { - thread_local std::vector pages; - pages.clear(); - pages.reserve(numPages); - for (std::size_t page = 0; page < numPages; ++page) - { - pages.push_back({pageIndices[page].dst, pageIndices[page].src}); - } - kernels::invokeNvfp4ColdPageDecode(pages, mPreparedPlans.at(planIndex), coldBase, stream); - } - -private: - std::map mLayerLayouts; - std::set mLayerIds; - std::vector mPreparedPlans; -}; - -} // namespace - -std::unique_ptr createNvfp4ColdPageCodec(std::vector layerLayouts) -{ - return std::make_unique(std::move(layerLayouts)); -} - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h deleted file mode 100644 index 7cd0d0cb1972..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h +++ /dev/null @@ -1,58 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#pragma once - -#include "kv_cache_manager_v2/coldPageCodec.h" -#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" - -#include -#include -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ - -namespace kv = batch_manager::kv_cache_manager_v2; - -//! Per-buffer global scales used by the NVFP4 cold-page kernels. -struct Nvfp4ColdPageScales -{ - float nvfp4ScaleOrigQuant = 1.0F; - float nvfp4ScaleQuantOrig = 1.0F; - float fp8ScaleOrigQuant = 1.0F; - float fp8ScaleQuantOrig = 1.0F; -}; - -//! Algorithm and layer-relative cold offsets for one buffer. -struct Nvfp4ColdPageBufferLayout -{ - kv::DataRole role; - std::size_t coldDataOffset = 0; - std::size_t coldScaleOffset = 0; - std::optional scales; -}; - -//! Python-authored NVFP4 record layout for one Attention layer. -struct Nvfp4ColdPageLayerLayout -{ - kv::LayerId layerId = 0; - kernels::Nvfp4ColdPageRuntimeType runtimeType = kernels::Nvfp4ColdPageRuntimeType::kFloat16; - std::int32_t numKvHeads = 0; - std::int32_t tokensPerPage = 0; - std::int32_t headDim = 0; - std::size_t coldPageBytes = 0; - std::size_t coldPaddingOffset = 0; - std::vector buffers; -}; - -//! Create an owning native codec configured by NVFP4 layer layouts. -[[nodiscard]] std::unique_ptr createNvfp4ColdPageCodec( - std::vector layerLayouts); - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 07a92305270b..e4d27412467e 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,54 +17,148 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" +#include "tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h" +#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" #include -#include +#include #include #include #include +#include +#include +#include +#include +#include + namespace nb = nanobind; namespace compression = tensorrt_llm::kv_cache_compression; -namespace kernels = tensorrt_llm::kernels; +namespace kv = tensorrt_llm::batch_manager::kv_cache_manager_v2; namespace tensorrt_llm::nanobind::kv_cache_compression { +namespace +{ + +static_assert(sizeof(kv::PageIndexPair) == sizeof(compression::ColdPageIndexPair)); +static_assert(alignof(kv::PageIndexPair) == alignof(compression::ColdPageIndexPair)); +static_assert(offsetof(kv::PageIndexPair, dst) == offsetof(compression::ColdPageIndexPair, dst)); +static_assert(offsetof(kv::PageIndexPair, src) == offsetof(compression::ColdPageIndexPair, src)); + +//! Algorithm-neutral adapter from KVCM's native codec calls to one Python policy. +class PythonColdPageCodec final : public compression::NativeColdPageCodec +{ +public: + PythonColdPageCodec(std::vector layerIds, nb::handle policy) + : mLayerIds(layerIds.begin(), layerIds.end()) + , mPolicy(policy.ptr()) + { + if (mLayerIds.size() != layerIds.size()) + { + throw std::invalid_argument("Cold-page policy layer IDs must be unique"); + } + if (policy.is_none()) + { + throw std::invalid_argument("Cold-page policy must not be None"); + } + Py_INCREF(mPolicy); + } + + ~PythonColdPageCodec() override + { + if (Py_IsInitialized()) + { + nb::gil_scoped_acquire acquire; + Py_DECREF(mPolicy); + } + } + +private: + [[nodiscard]] std::set const& getLayerIds() const noexcept override + { + return mLayerIds; + } + + std::vector configureAlgorithm( + std::vector const& lifecycles) override + { + nb::gil_scoped_acquire acquire; + try + { + return nb::cast>( + nb::borrow(mPolicy).attr("configure")(lifecycles)); + } + catch (nb::python_error const& error) + { + throw std::runtime_error(error.what()); + } + } + + void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override + { + invoke("encode", planIndex, coldBase, pageIndices, numPages, stream); + } + + void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) override + { + invoke("decode", planIndex, coldBase, pageIndices, numPages, stream); + } + + template + void invoke(char const* method, std::size_t planIndex, ColdPointer coldBase, kv::PageIndexPair const* pageIndices, + std::size_t numPages, cudaStream_t stream) + { + // Forward the complete KVCM batch once. The method custom op owns any launch chunking. + nb::gil_scoped_acquire acquire; + try + { + nb::borrow(mPolicy).attr(method)(planIndex, reinterpret_cast(coldBase), + reinterpret_cast(pageIndices), numPages, reinterpret_cast(stream)); + } + catch (nb::python_error const& error) + { + throw std::runtime_error(error.what()); + } + } + + std::set mLayerIds; + PyObject* mPolicy; +}; + +std::unique_ptr createPythonColdPageCodec( + std::vector layerIds, nb::handle policy) +{ + return std::make_unique(std::move(layerIds), policy); +} + +} // namespace void initBindings(nb::module_& module) { - nb::enum_(module, "Nvfp4ColdPageRuntimeType") - .value("FLOAT16", kernels::Nvfp4ColdPageRuntimeType::kFloat16) - .value("BFLOAT16", kernels::Nvfp4ColdPageRuntimeType::kBfloat16) - .value("FP8_E4M3", kernels::Nvfp4ColdPageRuntimeType::kFp8E4m3); + nb::enum_(module, "ColdPageIndexLocation") + .value("BAD_LOCATION", kv::PageIndexLocation::kBadLocation) + .value("HOST", kv::PageIndexLocation::kHost) + .value("DEVICE", kv::PageIndexLocation::kDevice); - nb::class_(module, "Nvfp4ColdPageScales") - .def(nb::init<>()) - .def_rw("nvfp4_scale_orig_quant", &compression::Nvfp4ColdPageScales::nvfp4ScaleOrigQuant) - .def_rw("nvfp4_scale_quant_orig", &compression::Nvfp4ColdPageScales::nvfp4ScaleQuantOrig) - .def_rw("fp8_scale_orig_quant", &compression::Nvfp4ColdPageScales::fp8ScaleOrigQuant) - .def_rw("fp8_scale_quant_orig", &compression::Nvfp4ColdPageScales::fp8ScaleQuantOrig); + nb::class_(module, "ResolvedHotBuffer") + .def_ro("raw_base", &compression::ResolvedHotBuffer::rawBase) + .def_ro("raw_slot_bytes", &compression::ResolvedHotBuffer::rawSlotBytes) + .def_ro("raw_bytes", &compression::ResolvedHotBuffer::rawBytes); - nb::class_(module, "Nvfp4ColdPageBufferLayout") - .def(nb::init<>()) - .def_rw("role", &compression::Nvfp4ColdPageBufferLayout::role) - .def_rw("cold_data_offset", &compression::Nvfp4ColdPageBufferLayout::coldDataOffset) - .def_rw("cold_scale_offset", &compression::Nvfp4ColdPageBufferLayout::coldScaleOffset) - .def_rw("scales", &compression::Nvfp4ColdPageBufferLayout::scales); + nb::class_(module, "ResolvedHotLifecycle") + .def_prop_ro("life_cycle_id", + [](compression::ResolvedHotLifecycle const& lifecycle) { return lifecycle.lifeCycleId.value(); }) + .def_ro("layers", &compression::ResolvedHotLifecycle::layers); - nb::class_(module, "Nvfp4ColdPageLayerLayout") + nb::class_(module, "ColdPageLifecycleProperties") .def(nb::init<>()) - .def_rw("layer_id", &compression::Nvfp4ColdPageLayerLayout::layerId) - .def_rw("runtime_type", &compression::Nvfp4ColdPageLayerLayout::runtimeType) - .def_rw("num_kv_heads", &compression::Nvfp4ColdPageLayerLayout::numKvHeads) - .def_rw("tokens_per_page", &compression::Nvfp4ColdPageLayerLayout::tokensPerPage) - .def_rw("head_dim", &compression::Nvfp4ColdPageLayerLayout::headDim) - .def_rw("cold_page_bytes", &compression::Nvfp4ColdPageLayerLayout::coldPageBytes) - .def_rw("cold_padding_offset", &compression::Nvfp4ColdPageLayerLayout::coldPaddingOffset) - .def_rw("buffers", &compression::Nvfp4ColdPageLayerLayout::buffers); - - module.def("create_nvfp4_cold_page_codec", &compression::createNvfp4ColdPageCodec, nb::arg("layer_layouts")); + .def_rw("cold_page_bytes", &compression::ColdPageLifecycleProperties::coldPageBytes) + .def_rw("page_index_location", &compression::ColdPageLifecycleProperties::pageIndexLocation); + + module.def("create_python_cold_page_codec", &createPythonColdPageCodec, nb::arg("layer_ids"), nb::arg("policy")); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index cdf0a0228fac..4215a1e0d724 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -111,6 +111,7 @@ add_library( fp8PerTensorScaleMoe.cpp fp4BlockScaleMoe.cpp noAuxTcOp.cpp + nvfp4ColdPageOp.cpp fusedCatFp4Op.cpp fusedCatFp8Op.cpp IndexerKCacheGatherOp.cpp diff --git a/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp b/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp new file mode 100644 index 000000000000..c9e46f358bbc --- /dev/null +++ b/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp @@ -0,0 +1,229 @@ +/* + * 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 "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" + +#include +#include +#include +#include +#include +#include +#include + +TRTLLM_NAMESPACE_BEGIN + +namespace torch_ext +{ +namespace +{ + +using kernels::Nvfp4ColdPageBufferPlan; +using kernels::ColdPageIndexPair; +using kernels::Nvfp4ColdPageKernelParams; +using kernels::Nvfp4ColdPagePreparedPlan; +using kernels::Nvfp4ColdPageRuntimeType; +using kernels::Nvfp4ColdPageTransform; + +enum BufferIntegerField : std::size_t +{ + kRawBase, + kRawSlotBytes, + kRawBytes, + kColdDataOffset, + kColdScaleOffset, + kColdPaddingOffset, + kColdPaddingBytes, + kTransform, + kNumKvHeads, + kTokensPerPage, + kHeadDim, + kNumBufferIntegerFields, +}; + +enum BufferScaleField : std::size_t +{ + kNvfp4ScaleOrigQuant, + kNvfp4ScaleQuantOrig, + kFp8ScaleOrigQuant, + kFp8ScaleQuantOrig, + kNumBufferScaleFields, +}; + +template +T checkedNonNegativeCast(std::int64_t value, char const* field) +{ + TORCH_CHECK(value >= 0, field, " must be non-negative"); + TORCH_CHECK(static_cast(value) <= static_cast(std::numeric_limits::max()), field, + " exceeds its native range"); + return static_cast(value); +} + +Nvfp4ColdPageRuntimeType parseRuntimeType(std::int64_t value) +{ + switch (value) + { + case 0: return Nvfp4ColdPageRuntimeType::kFloat16; + case 1: return Nvfp4ColdPageRuntimeType::kBfloat16; + case 2: return Nvfp4ColdPageRuntimeType::kFp8E4m3; + default: TORCH_CHECK(false, "Unsupported NVFP4 cold-page runtime type ", value); + } + return Nvfp4ColdPageRuntimeType::kFloat16; +} + +Nvfp4ColdPageTransform parseTransform(std::int64_t value) +{ + switch (value) + { + case 0: return Nvfp4ColdPageTransform::kNvfp4; + case 1: return Nvfp4ColdPageTransform::kLosslessCopy; + default: TORCH_CHECK(false, "Unsupported NVFP4 cold-page transform ", value); + } + return Nvfp4ColdPageTransform::kNvfp4; +} + +struct ColdPageInvocation +{ + std::uintptr_t coldBase; + void const* pagePairs; + std::size_t numPages; + cudaStream_t stream; +}; + +ColdPageInvocation parseInvocation( + std::int64_t coldBase, std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) +{ + auto const coldAddress = checkedNonNegativeCast(coldBase, "cold_base"); + auto const pairsAddress = checkedNonNegativeCast(pagePairs, "page_pairs"); + auto const numPages = checkedNonNegativeCast(pageCount, "page_count"); + auto const streamAddress = checkedNonNegativeCast(stream, "stream"); + if (numPages != 0) + { + TORCH_CHECK(coldAddress != 0, "cold_base must not be null"); + TORCH_CHECK(pairsAddress != 0, "page_pairs must not be null"); + TORCH_CHECK(pairsAddress % alignof(ColdPageIndexPair) == 0, "page_pairs is misaligned"); + } + return {coldAddress, reinterpret_cast(pairsAddress), numPages, + reinterpret_cast(streamAddress)}; +} + +} // namespace + +//! Opaque, immutable configure-time plan consumed by the runtime custom ops. +class Nvfp4ColdPageProgram : public torch::CustomClassHolder +{ +public: + explicit Nvfp4ColdPageProgram(Nvfp4ColdPagePreparedPlan plan) + : mPlan(std::move(plan)) + { + } + + [[nodiscard]] Nvfp4ColdPagePreparedPlan const& getPlan() const noexcept + { + return mPlan; + } + +private: + Nvfp4ColdPagePreparedPlan const mPlan; +}; + +c10::intrusive_ptr prepareNvfp4ColdPageProgram( + c10::List> const& bufferIntegers, c10::List> const& bufferScales, + std::int64_t coldPageBytes, std::int64_t runtimeType) +{ + TORCH_CHECK(bufferIntegers.size() == bufferScales.size(), "Each cold-page buffer needs one scale row"); + + std::vector buffers; + buffers.reserve(bufferIntegers.size()); + for (std::size_t index = 0; index < bufferIntegers.size(); ++index) + { + auto const integers = bufferIntegers.get(index); + auto const scales = bufferScales.get(index); + TORCH_CHECK(integers.size() == kNumBufferIntegerFields, "NVFP4 buffer ", index, " requires ", + kNumBufferIntegerFields, " integer fields, got ", integers.size()); + TORCH_CHECK(scales.size() == kNumBufferScaleFields, "NVFP4 buffer ", index, " requires ", kNumBufferScaleFields, + " scale fields, got ", scales.size()); + + Nvfp4ColdPageBufferPlan buffer{}; + buffer.rawBase = checkedNonNegativeCast(integers.get(kRawBase), "raw_base"); + buffer.rawSlotBytes = checkedNonNegativeCast(integers.get(kRawSlotBytes), "raw_slot_bytes"); + buffer.rawBytes = checkedNonNegativeCast(integers.get(kRawBytes), "raw_bytes"); + buffer.coldDataOffset = checkedNonNegativeCast(integers.get(kColdDataOffset), "cold_data_offset"); + buffer.coldScaleOffset + = checkedNonNegativeCast(integers.get(kColdScaleOffset), "cold_scale_offset"); + buffer.coldPaddingOffset + = checkedNonNegativeCast(integers.get(kColdPaddingOffset), "cold_padding_offset"); + buffer.coldPaddingBytes + = checkedNonNegativeCast(integers.get(kColdPaddingBytes), "cold_padding_bytes"); + buffer.transform = parseTransform(integers.get(kTransform)); + buffer.params = Nvfp4ColdPageKernelParams{ + checkedNonNegativeCast(integers.get(kNumKvHeads), "num_kv_heads"), + checkedNonNegativeCast(integers.get(kTokensPerPage), "tokens_per_page"), + checkedNonNegativeCast(integers.get(kHeadDim), "head_dim"), + static_cast(scales.get(kNvfp4ScaleOrigQuant)), static_cast(scales.get(kNvfp4ScaleQuantOrig)), + static_cast(scales.get(kFp8ScaleOrigQuant)), static_cast(scales.get(kFp8ScaleQuantOrig))}; + buffers.push_back(buffer); + } + + auto plan = kernels::prepareNvfp4ColdPagePlan( + buffers, checkedNonNegativeCast(coldPageBytes, "cold_page_bytes"), parseRuntimeType(runtimeType)); + return c10::make_intrusive(std::move(plan)); +} + +void nvfp4ColdPageEncode(c10::intrusive_ptr const& program, std::int64_t coldBase, + std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) +{ + TORCH_CHECK(program, "NVFP4 cold-page program must not be null"); + auto const batch = parseInvocation(coldBase, pagePairs, pageCount, stream); + kernels::invokeNvfp4ColdPageEncode( + batch.pagePairs, batch.numPages, program->getPlan(), reinterpret_cast(batch.coldBase), batch.stream); +} + +void nvfp4ColdPageDecode(c10::intrusive_ptr const& program, std::int64_t coldBase, + std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) +{ + TORCH_CHECK(program, "NVFP4 cold-page program must not be null"); + auto const batch = parseInvocation(coldBase, pagePairs, pageCount, stream); + kernels::invokeNvfp4ColdPageDecode(batch.pagePairs, batch.numPages, program->getPlan(), + reinterpret_cast(batch.coldBase), batch.stream); +} + +} // namespace torch_ext + +TRTLLM_NAMESPACE_END + +TORCH_LIBRARY_FRAGMENT(trtllm, m) +{ + m.class_("Nvfp4ColdPageProgram"); + m.def( + "prepare_nvfp4_cold_page_program(int[][] buffer_ints, float[][] buffer_scales, int cold_page_bytes, " + "int runtime_type) -> __torch__.torch.classes.trtllm.Nvfp4ColdPageProgram"); + m.def( + "nvfp4_cold_page_encode(__torch__.torch.classes.trtllm.Nvfp4ColdPageProgram program, int cold_base, " + "int page_pairs, int page_count, int stream) -> ()"); + m.def( + "nvfp4_cold_page_decode(__torch__.torch.classes.trtllm.Nvfp4ColdPageProgram program, int cold_base, " + "int page_pairs, int page_count, int stream) -> ()"); +} + +TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, m) +{ + m.impl("prepare_nvfp4_cold_page_program", &tensorrt_llm::torch_ext::prepareNvfp4ColdPageProgram); + m.impl("nvfp4_cold_page_encode", &tensorrt_llm::torch_ext::nvfp4ColdPageEncode); + m.impl("nvfp4_cold_page_decode", &tensorrt_llm::torch_ext::nvfp4ColdPageDecode); +} diff --git a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp index be09cca9c9b3..88db07644510 100644 --- a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp +++ b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp @@ -42,9 +42,8 @@ namespace using tensorrt_llm::batch_manager::kv_cache_manager_v2::HostMem; using tensorrt_llm::batch_manager::kv_cache_manager_v2::MemAddress; using tensorrt_llm::kernels::Nvfp4ColdPageBufferPlan; +using tensorrt_llm::kernels::ColdPageIndexPair; using tensorrt_llm::kernels::Nvfp4ColdPageKernelParams; -using tensorrt_llm::kernels::Nvfp4ColdPageOffloadPageTask; -using tensorrt_llm::kernels::Nvfp4ColdPageOnboardPageTask; using tensorrt_llm::kernels::Nvfp4ColdPagePreparedPlan; using tensorrt_llm::kernels::Nvfp4ColdPageRuntimeType; using tensorrt_llm::kernels::Nvfp4ColdPageTransform; @@ -500,16 +499,17 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG MappedHostRegion compactPages(coldBaseOffset + slotCapacity * compactSlotBytes); auto* compactBase = compactPages.bytes() + coldBaseOffset; std::vector, 2>> rawHost(numPages); - std::vector offloadTasks; + std::vector offloadTasks; offloadTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { - std::size_t const slot = 2U * page; + std::size_t const rawSlot = 2U * page; + std::size_t const coldSlot = rawSlot + 1U; rawHost[page][0] = makeRawPage(kind, page, 0, params[0], geometry, inputPattern); rawHost[page][1] = makeRawPage(kind, page, 1, params[1], geometry, inputPattern); - rawInputK.copyFrom(slot * rawSlotBytes, rawHost[page][0]); - rawInputV.copyFrom(slot * rawSlotBytes, rawHost[page][1]); - offloadTasks.push_back({static_cast(slot), static_cast(slot)}); + rawInputK.copyFrom(rawSlot * rawSlotBytes, rawHost[page][0]); + rawInputV.copyFrom(rawSlot * rawSlotBytes, rawHost[page][1]); + offloadTasks.push_back({static_cast(coldSlot), static_cast(rawSlot)}); } std::size_t const packed = packedBytes(geometry); @@ -538,7 +538,7 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG [](std::uint8_t value) { return value == kCanary; })); for (std::size_t page = 0; page < numPages; ++page) { - std::size_t const base = coldBaseOffset + 2U * page * compactSlotBytes; + std::size_t const base = coldBaseOffset + (2U * page + 1U) * compactSlotBytes; auto const region = [&](std::size_t offset, std::size_t bytes) { return std::vector(payload.begin() + static_cast(base + offset), @@ -551,14 +551,15 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG auto const padding = region(2U * (packed + scale), compactSlotBytes - 2U * (packed + scale)); EXPECT_TRUE(std::all_of(padding.begin(), padding.end(), [](std::uint8_t value) { return value == 0U; })); - std::size_t const unusedBase = coldBaseOffset + (2U * page + 1U) * compactSlotBytes; + std::size_t const unusedBase = coldBaseOffset + 2U * page * compactSlotBytes; EXPECT_TRUE(std::all_of(payload.begin() + static_cast(unusedBase), payload.begin() + static_cast(unusedBase + compactSlotBytes), [](std::uint8_t value) { return value == kCanary; })); } }; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + offloadTasks.data(), offloadTasks.size(), inputPlan, compactBase, stream); if (synchronizeBetweenDirections) { @@ -577,21 +578,23 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG auto const firstSerialization = compactPages.payload(); for (std::size_t page = 0; page < numPages; ++page) { - std::memset(compactBase + 2U * page * compactSlotBytes, 0x5A, compactSlotBytes); + std::memset(compactBase + (2U * page + 1U) * compactSlotBytes, 0x5A, compactSlotBytes); } - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + offloadTasks.data(), offloadTasks.size(), inputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); EXPECT_EQ(compactPages.payload(), firstSerialization); verifyCompressedPages(); } } - std::vector onboardTasks; + std::vector onboardTasks; onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { - std::size_t const slot = 2U * page; - onboardTasks.push_back({static_cast(slot), static_cast(slot)}); + std::size_t const rawSlot = 2U * page; + std::size_t const coldSlot = rawSlot + 1U; + onboardTasks.push_back({static_cast(rawSlot), static_cast(coldSlot)}); } std::vector const outputBuffers{ {reinterpret_cast(rawOutputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, @@ -600,7 +603,8 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG payloadBytes, paddingBytes, Nvfp4ColdPageTransform::kNvfp4, params[1]}}; auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan(outputBuffers, compactSlotBytes, runtimeType(kind)); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + onboardTasks.data(), onboardTasks.size(), outputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); if (!synchronizeBetweenDirections) @@ -620,8 +624,10 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG references[page][role] = compressReference(restored, kind, params[role], geometry); } } - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, outputPlan, compactBase, stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, inputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + offloadTasks.data(), offloadTasks.size(), outputPlan, compactBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + onboardTasks.data(), onboardTasks.size(), inputPlan, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); verifyCompressedPages(); } @@ -737,8 +743,8 @@ void runPartialPageTailIsolation(RawKind kind) DeviceRegion rawOutputK(numPages * rawSlotBytes); DeviceRegion rawOutputV(numPages * rawSlotBytes); MappedHostRegion compactPages(numPages * compactSlotBytes); - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; offloadTasks.reserve(numPages); onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) @@ -771,8 +777,10 @@ void runPartialPageTailIsolation(RawKind kind) auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( makeBuffers(rawOutputK, rawOutputV), compactSlotBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, compactPages.data(), stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, compactPages.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + offloadTasks.data(), offloadTasks.size(), inputPlan, compactPages.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + onboardTasks.data(), onboardTasks.size(), outputPlan, compactPages.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); for (std::size_t pair = 0; pair < kValidTokenCounts.size(); ++pair) @@ -899,8 +907,8 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) std::array, numPages> mlaHost; std::array, numPages> sideHost; std::array references; - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; for (std::size_t page = 0; page < numPages; ++page) { mlaHost[page] = makeRawPage(kind, page, 0U, params, geometry, InputPattern::kDense); @@ -932,8 +940,10 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) makePlans(mlaOutput, sideOutput), coldPageBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(offloadTasks, inputPlan, coldBase, stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(onboardTasks, outputPlan, coldBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + offloadTasks.data(), offloadTasks.size(), inputPlan, coldBase, stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + onboardTasks.data(), onboardTasks.size(), outputPlan, coldBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const cold = coldStorage.payload(); @@ -1055,8 +1065,9 @@ TEST(Nvfp4ColdPageWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc inputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( outputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({{0, 0}}, inputPlan, compactPage.data(), stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode({{0, 0}}, outputPlan, compactPage.data(), stream); + ColdPageIndexPair const page{0, 0}; + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(&page, 1U, inputPlan, compactPage.data(), stream); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(&page, 1U, outputPlan, compactPage.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const compact = compactPage.payload(); @@ -1116,8 +1127,8 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector offloadPages; - std::vector onboardPages; + std::vector offloadPages; + std::vector onboardPages; offloadPages.reserve(numPages); onboardPages.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) @@ -1165,9 +1176,17 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector(16, 256)); +} + class Nvfp4ColdPageTailTest : public testing::TestWithParam { }; @@ -1194,8 +1223,8 @@ INSTANTIATE_TEST_SUITE_P( TEST(Nvfp4ColdPageValidationTest, EmptyBatchIsAnAsyncNoOp) { - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({}, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode({}, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(nullptr, 0U, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(nullptr, 0U, Nvfp4ColdPagePreparedPlan{}, nullptr, nullptr); } TEST(Nvfp4ColdPageValidationTest, RejectsInvalidGeometryAndScalesBeforeLaunch) @@ -1277,7 +1306,7 @@ TEST(Nvfp4ColdPageValidationTest, RejectsInvalidLaunchDescriptors) std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); LayerBuffers buffers(rawSlotBytes); - Nvfp4ColdPageOffloadPageTask const validOffload{0, 0}; + ColdPageIndexPair const validPage{0, 0}; std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); std::size_t const packed = packedBytes(kDefaultGeometry); std::size_t const scale = scaleBytes(kDefaultGeometry); @@ -1298,7 +1327,7 @@ TEST(Nvfp4ColdPageValidationTest, RejectsInvalidLaunchDescriptors) auto const validPlan = prepare(validBuffer, coldPageBytes); expectInvalid("unaligned raw base", [](auto& buffer) { buffer.rawBase += 1U; }); - EXPECT_ANY_THROW(tensorrt_llm::kernels::invokeNvfp4ColdPageEncode({validOffload}, validPlan, nullptr, nullptr)); + EXPECT_ANY_THROW(tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(&validPage, 1U, validPlan, nullptr, nullptr)); expectInvalid("unaligned raw stride", [](auto& buffer) { buffer.rawSlotBytes += alignof(uint4) / 2U; }); expectInvalid("raw bytes exceed stride", [](auto& buffer) { buffer.rawBytes = buffer.rawSlotBytes + 1U; }); expectInvalid("raw bytes mismatch geometry", [](auto& buffer) { buffer.rawBytes -= alignof(uint4); }); diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index bada5d6e6578..ee19f85157e9 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -7,7 +7,6 @@ set(COLD_PAGE_CODEC_TEST_SRC ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp index 5786189eac58..8ac907f0a094 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -5,14 +5,12 @@ */ #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" -#include "tensorrt_llm/kv_cache_compression/nvfp4ColdPageCodec.h" #include -#include #include +#include #include -#include #include #include #include @@ -34,7 +32,6 @@ constexpr std::uintptr_t kGpuVBase = 0x200000; constexpr std::uintptr_t kColdBase = 0x300000; constexpr std::uintptr_t kStreamValue = 0x7000; constexpr std::size_t kRawBytes = 320; -constexpr std::size_t kColdBytes = 192; kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::PoolGroupIndex{0}, kv::LayerGroupId lifeCycle = kv::LayerGroupId{0}, std::size_t count = 1U, int firstLayer = 0, @@ -56,16 +53,6 @@ kv::PoolGroupDesc makeAttentionDesc(kv::PoolGroupIndex poolGroupIndex = kv::Pool {kv::PoolIndex{0}, keyBase, slotBytes}, {kv::PoolIndex{1}, valueBase, slotBytes}}}; } -kv::PoolGroupDesc makeMlaDesc() -{ - kv::SlotDescVariant variant{kv::LayerGroupId{0}, - kv::TypedVec{ - kv::CoalescedBuffer{kRawBytes, {{0, "key"}}}, kv::CoalescedBuffer{68U, {{0, "index_key"}}}}}; - return {kv::PoolGroupIndex{0}, kv::SlotCount{512}, kv::SlotDesc{{std::move(variant)}}, - kv::TypedVec{ - {kv::PoolIndex{0}, kGpuKBase, kRawBytes}, {kv::PoolIndex{1}, kGpuVBase, 68U}}}; -} - kv::PoolGroupDesc makeLosslessDesc(kv::PoolGroupIndex poolGroupIndex, kv::LayerGroupId lifeCycle) { kv::SlotDescVariant variant{lifeCycle, @@ -111,6 +98,7 @@ class RecordingCodec final : public NativeColdPageCodec { if (failBatches) { + enqueueFailureMarker(stream); throw std::runtime_error("requested batch failure"); } ++encodeCalls; @@ -125,6 +113,7 @@ class RecordingCodec final : public NativeColdPageCodec { if (failBatches) { + enqueueFailureMarker(stream); throw std::runtime_error("requested batch failure"); } ++decodeCalls; @@ -136,6 +125,7 @@ class RecordingCodec final : public NativeColdPageCodec bool failConfigure = false; bool failBatches = false; + std::atomic_bool* failureMarker = nullptr; int encodeCalls = 0; int decodeCalls = 0; std::size_t lastPlanIndex = 0; @@ -145,10 +135,21 @@ class RecordingCodec final : public NativeColdPageCodec std::vector resolved; private: + void enqueueFailureMarker(cudaStream_t stream) + { + if (failureMarker != nullptr + && cudaLaunchHostFunc( + stream, [](void* marker) { static_cast(marker)->store(true); }, failureMarker) + != cudaSuccess) + { + throw std::runtime_error("failed to enqueue the requested batch failure marker"); + } + } + std::set mLayerIds; }; -TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsBatches) +TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsWholeBatchOnce) { RecordingCodec codec{{0, 1}}; @@ -163,18 +164,23 @@ TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsBatches) EXPECT_EQ(layers.at(1).at("key").rawBytes, kRawBytes); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{3}), 777U); - kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; + std::vector indices(4096); + for (std::size_t index = 0; index < indices.size(); ++index) + { + indices[index] = {static_cast(index + 1U), static_cast(index)}; + } auto const stream = reinterpret_cast(kStreamValue); ASSERT_TRUE( - codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices.data(), indices.size(), stream)); EXPECT_EQ(codec.encodeCalls, 1); EXPECT_EQ(codec.lastPlanIndex, 0U); - EXPECT_EQ(codec.lastIndices[1].src, 3); + EXPECT_EQ(codec.lastIndices.size(), 4096U); + EXPECT_EQ(codec.lastIndices.back().src, 4095); EXPECT_EQ(codec.lastColdBase, reinterpret_cast(kColdBase)); EXPECT_EQ(codec.lastStream, stream); ASSERT_TRUE(codec.decode( - kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); + kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices.data(), indices.size(), stream)); EXPECT_EQ(codec.decodeCalls, 1); } @@ -208,7 +214,7 @@ TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) } } -TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureAndBatchFailures) +TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureFailuresAndInvalidBatches) { RecordingCodec codec{{0}}; codec.failConfigure = true; @@ -221,164 +227,41 @@ TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureAndBatchFailures) kv::PageIndexPair const indices[]{{0, 0}}; EXPECT_FALSE(validCodec.encode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); EXPECT_FALSE(validCodec.decode(kv::LayerGroupId{0}, nullptr, indices, 1U, nullptr)); - validCodec.failBatches = true; - EXPECT_FALSE(validCodec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - EXPECT_FALSE( - validCodec.decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); EXPECT_EQ(validCodec.encodeCalls, 0); EXPECT_EQ(validCodec.decodeCalls, 0); } -TEST(NativeColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) +TEST(NativeColdPageCodecTest, AlgorithmFailureUsesTheSuppliedCudaStreamForRollback) { + int deviceCount = 0; + if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) + { + GTEST_SKIP() << "Failure draining requires a CUDA device"; + } + RecordingCodec codec{{0}}; ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); - EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); - EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); - EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); -} - -Nvfp4ColdPageScales makeScales(float scale) -{ - Nvfp4ColdPageScales scales; - scales.nvfp4ScaleOrigQuant = scale; - scales.nvfp4ScaleQuantOrig = 1.0F / scale; - return scales; -} - -Nvfp4ColdPageLayerLayout makeAttentionLayout(int layerId, float keyScale = 1.0F, float valueScale = 2.0F) -{ - return Nvfp4ColdPageLayerLayout{layerId, kernels::Nvfp4ColdPageRuntimeType::kFloat16, 1, 5, 32, kColdBytes, 180U, - {Nvfp4ColdPageBufferLayout{"key", 0U, 160U, makeScales(keyScale)}, - Nvfp4ColdPageBufferLayout{"value", 80U, 170U, makeScales(valueScale)}}}; -} - -Nvfp4ColdPageLayerLayout makeMlaLayout() -{ - return Nvfp4ColdPageLayerLayout{0, kernels::Nvfp4ColdPageRuntimeType::kFloat16, 1, 5, 32, 160U, 158U, - {Nvfp4ColdPageBufferLayout{"key", 0U, 80U, makeScales(1.0F)}, - Nvfp4ColdPageBufferLayout{"index_key", 90U, 0U, std::nullopt}}}; -} - -struct RecordedLaunch -{ - int prepareCalls = 0; - int encodeCalls = 0; - int decodeCalls = 0; - std::vector encodePages; - std::vector decodePages; - kernels::Nvfp4ColdPagePreparedPlan plan; -}; - -RecordedLaunch gLaunch; - -TEST(Nvfp4ColdPageCodecTest, LowersMhaLayoutOnceAndDispatchesEncodeDecode) -{ - gLaunch = {}; - auto codec = createNvfp4ColdPageCodec({makeAttentionLayout(0, 2.0F, 3.0F), makeAttentionLayout(1, 4.0F, 5.0F)}); - ASSERT_TRUE(configureOne(*codec, makeAttentionDesc(kv::PoolGroupIndex{0}, kv::LayerGroupId{3}, 2U))); - EXPECT_EQ(gLaunch.prepareCalls, 1); - EXPECT_EQ(codec->queryColdPageBytes(kv::LayerGroupId{3}), 2U * kColdBytes); - - kv::PageIndexPair const indices[]{{2, 1}, {5, 3}}; - auto const stream = reinterpret_cast(kStreamValue); - ASSERT_TRUE( - codec->encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - ASSERT_EQ(gLaunch.plan.numBuffers, 4U); - EXPECT_EQ(gLaunch.plan.buffers[0].rawBase, kGpuKBase); - EXPECT_EQ(gLaunch.plan.buffers[2].rawBase, kGpuKBase + kRawBytes); - EXPECT_EQ(gLaunch.plan.buffers[2].coldDataOffset, kColdBytes); - EXPECT_EQ(gLaunch.plan.buffers[3].coldScaleOffset, kColdBytes + 170U); - EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingOffset, kColdBytes + 180U); - EXPECT_EQ(gLaunch.plan.buffers[3].coldPaddingBytes, 12U); - EXPECT_FLOAT_EQ(gLaunch.plan.buffers[0].params.nvfp4ScaleOrigQuant, 2.0F); - EXPECT_FLOAT_EQ(gLaunch.plan.buffers[3].params.nvfp4ScaleOrigQuant, 5.0F); - ASSERT_EQ(gLaunch.encodePages.size(), 2U); - EXPECT_EQ(gLaunch.encodePages[0].gpuPageIndex, 1); - EXPECT_EQ(gLaunch.encodePages[0].coldPageIndex, 2); - - ASSERT_TRUE(codec->decode( - kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices, std::size(indices), stream)); - ASSERT_EQ(gLaunch.decodePages.size(), 2U); - EXPECT_EQ(gLaunch.decodePages[0].gpuPageIndex, 2); - EXPECT_EQ(gLaunch.decodePages[0].coldPageIndex, 1); -} - -TEST(Nvfp4ColdPageCodecTest, PreservesMlaSideBufferLosslessly) -{ - gLaunch = {}; - auto codec = createNvfp4ColdPageCodec({makeMlaLayout()}); - ASSERT_TRUE(configureOne(*codec, makeMlaDesc())); - + codec.failBatches = true; + std::atomic_bool completed = false; + codec.failureMarker = &completed; + cudaStream_t stream{}; + ASSERT_EQ(cudaStreamCreate(&stream), cudaSuccess); kv::PageIndexPair const indices[]{{0, 0}}; - ASSERT_TRUE(codec->encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, nullptr)); - ASSERT_EQ(gLaunch.plan.numBuffers, 2U); - EXPECT_EQ(gLaunch.plan.buffers[0].transform, kernels::Nvfp4ColdPageTransform::kNvfp4); - EXPECT_EQ(gLaunch.plan.buffers[1].transform, kernels::Nvfp4ColdPageTransform::kLosslessCopy); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBase, kGpuVBase); - EXPECT_EQ(gLaunch.plan.buffers[1].rawBytes, 68U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldDataOffset, 90U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingOffset, 158U); - EXPECT_EQ(gLaunch.plan.buffers[1].coldPaddingBytes, 2U); + EXPECT_FALSE(codec.encode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, stream)); + EXPECT_TRUE(completed.exchange(false)); + EXPECT_FALSE(codec.decode(kv::LayerGroupId{0}, reinterpret_cast(kColdBase), indices, 1U, stream)); + EXPECT_TRUE(completed.load()); + EXPECT_EQ(cudaStreamDestroy(stream), cudaSuccess); } -TEST(Nvfp4ColdPageCodecTest, RejectsDuplicateLayoutsAndMissingRoles) +TEST(NativeColdPageCodecTest, UnknownLifecycleUsesFailureSentinels) { - EXPECT_THROW( - { - auto codec = createNvfp4ColdPageCodec({makeAttentionLayout(0), makeAttentionLayout(0)}); - }, - std::invalid_argument); - - auto duplicateRole = makeAttentionLayout(0); - duplicateRole.buffers.push_back(duplicateRole.buffers.front()); - EXPECT_THROW({ auto codec = createNvfp4ColdPageCodec({duplicateRole}); }, std::invalid_argument); - - auto missingRole = makeAttentionLayout(0); - missingRole.buffers.pop_back(); - auto codec = createNvfp4ColdPageCodec({missingRole}); - EXPECT_FALSE(configureOne(*codec, makeAttentionDesc())); + RecordingCodec codec{{0}}; + ASSERT_TRUE(configureOne(codec, makeAttentionDesc())); + EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{99}), 0U); + EXPECT_EQ(codec.getBatchingLayerGroupId(kv::LayerGroupId{99}), kv::LayerGroupId{-1}); + EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{99}), kv::PageIndexLocation::kBadLocation); } } // namespace } // namespace tensorrt_llm::kv_cache_compression - -namespace tensorrt_llm::kernels -{ - -Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) -{ - auto& launch = kv_cache_compression::gLaunch; - ++launch.prepareCalls; - if (buffers.empty() || buffers.size() > kNvfp4ColdPageMaxBuffersPerLaunch) - { - throw std::invalid_argument("invalid test launch plan"); - } - Nvfp4ColdPagePreparedPlan plan; - std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); - plan.numBuffers = static_cast(buffers.size()); - plan.coldPageBytes = coldPageBytes; - plan.runtimeType = runtimeType; - return plan; -} - -void invokeNvfp4ColdPageEncode( - std::vector const& pages, Nvfp4ColdPagePreparedPlan const& plan, void*, cudaStream_t) -{ - auto& launch = kv_cache_compression::gLaunch; - ++launch.encodeCalls; - launch.encodePages = pages; - launch.plan = plan; -} - -void invokeNvfp4ColdPageDecode(std::vector const& pages, - Nvfp4ColdPagePreparedPlan const& plan, void const*, cudaStream_t) -{ - auto& launch = kv_cache_compression::gLaunch; - ++launch.decodeCalls; - launch.decodePages = pages; - launch.plan = plan; -} - -} // namespace tensorrt_llm::kernels diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py index 7b7d95f321e0..028c28b04813 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py @@ -6,15 +6,19 @@ import math import os import re +from dataclasses import dataclass from pathlib import Path from typing import Sequence +import torch + from tensorrt_llm.quantization.modelopt_config import ( is_modelopt_quant_config, read_modelopt_quant_config, ) from ...pyexecutor.resource_manager import DataType +from .quantization_for_cold_page import ColdPageCodecPolicy, ColdPageQuantizationMethod ScalePair = tuple[float, float] LayerScales = tuple[ScalePair, ScalePair] @@ -27,6 +31,36 @@ _COLD_PAGE_ALIGNMENT = 16 _ELEMENTS_PER_BYTE = 2 _ELEMENTS_PER_SCALE = 16 +_NVFP4_TRANSFORM = 0 +_LOSSLESS_TRANSFORM = 1 + + +@dataclass(frozen=True) +class _Nvfp4Scales: + nvfp4_orig_quant: float + nvfp4_quant_orig: float + fp8_orig_quant: float = 1.0 + fp8_quant_orig: float = 1.0 + + +@dataclass(frozen=True) +class _Nvfp4BufferLayout: + role: str + data_offset: int + scale_offset: int = 0 + scales: _Nvfp4Scales | None = None + + +@dataclass(frozen=True) +class _Nvfp4LayerLayout: + layer_id: int + runtime_type: int + num_kv_heads: int + tokens_per_page: int + head_dim: int + cold_page_bytes: int + padding_offset: int + buffers: tuple[_Nvfp4BufferLayout, ...] def _load_modelopt_nvfp4_scales( @@ -100,8 +134,111 @@ def _buffer_bytes(buffer: object, tokens_per_page: int) -> int: return int(buffer.size) * (tokens_per_page // buffer_tokens) -class Nvfp4ColdPagePolicy: - """Build native NVFP4 plans without retaining Python in the data path.""" +class Nvfp4ColdPagePolicy(ColdPageCodecPolicy): + """Own NVFP4 lifecycle programs and dispatch one operation per codec batch.""" + + def __init__(self, layer_layouts: Sequence[_Nvfp4LayerLayout]) -> None: + self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} + self._programs: list[object] = [] + + @property + def layer_ids(self) -> tuple[int, ...]: + return tuple(sorted(self._layer_layouts)) + + def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: + """Resolve Python layouts against hot buffers and prepare method programs.""" + + from tensorrt_llm.bindings.internal import kv_cache_compression as native + + programs = [] + properties = [] + for lifecycle in lifecycles: + metadata: list[list[int]] = [] + scales: list[list[float]] = [] + cold_page_bytes = 0 + runtime_type = None + + for layer_id, hot_buffers in lifecycle.layers.items(): + layout = self._layer_layouts[int(layer_id)] + if runtime_type is not None and runtime_type != layout.runtime_type: + raise ValueError("One cold-page lifecycle must use one runtime dtype") + runtime_type = layout.runtime_type + + for index, buffer in enumerate(layout.buffers): + hot = hot_buffers[buffer.role] + padding_offset = 0 + padding_bytes = 0 + if index + 1 == len(layout.buffers): + padding_offset = cold_page_bytes + layout.padding_offset + padding_bytes = layout.cold_page_bytes - layout.padding_offset + metadata.append( + [ + int(hot.raw_base), + int(hot.raw_slot_bytes), + int(hot.raw_bytes), + cold_page_bytes + buffer.data_offset, + cold_page_bytes + buffer.scale_offset if buffer.scales else 0, + padding_offset, + padding_bytes, + _NVFP4_TRANSFORM if buffer.scales else _LOSSLESS_TRANSFORM, + layout.num_kv_heads if buffer.scales else 0, + layout.tokens_per_page if buffer.scales else 0, + layout.head_dim if buffer.scales else 0, + ] + ) + buffer_scales = buffer.scales or _Nvfp4Scales(1.0, 1.0) + scales.append( + [ + buffer_scales.nvfp4_orig_quant, + buffer_scales.nvfp4_quant_orig, + buffer_scales.fp8_orig_quant, + buffer_scales.fp8_quant_orig, + ] + ) + cold_page_bytes += layout.cold_page_bytes + + if runtime_type is None: + raise ValueError("NVFP4 received an empty cold-page lifecycle") + programs.append( + torch.ops.trtllm.prepare_nvfp4_cold_page_program( + metadata, scales, cold_page_bytes, runtime_type + ) + ) + lifecycle_properties = native.ColdPageLifecycleProperties() + lifecycle_properties.cold_page_bytes = cold_page_bytes + lifecycle_properties.page_index_location = native.ColdPageIndexLocation.HOST + properties.append(lifecycle_properties) + + self._programs = programs + return properties + + def encode( + self, + program_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + torch.ops.trtllm.nvfp4_cold_page_encode( + self._programs[program_index], cold_base, page_indices, num_pages, stream + ) + + def decode( + self, + program_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + torch.ops.trtllm.nvfp4_cold_page_decode( + self._programs[program_index], cold_base, page_indices, num_pages, stream + ) + + +class Nvfp4ColdPageQuantization(ColdPageQuantizationMethod): + """Build one fresh NVFP4 callback policy for each KVCM construction.""" def __init__(self, checkpoint_path: str | None) -> None: self._model_scales = _load_modelopt_nvfp4_scales(checkpoint_path) @@ -116,23 +253,18 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create the native generic codec from explicit per-buffer plans.""" - from tensorrt_llm.bindings.internal import kv_cache_compression as native from tensorrt_llm.runtime.kv_cache_manager_v2 import AttentionLayerConfig + runtime_type = { + DataType.HALF: 0, + DataType.BF16: 1, + DataType.FP8: 2, + }.get(runtime_dtype) attention_layers = [ layer for layer in cache_config.layers if isinstance(layer, AttentionLayerConfig) ] - if not attention_layers: - return native.create_nvfp4_cold_page_codec([]) - - runtime_type = { - DataType.HALF: native.Nvfp4ColdPageRuntimeType.FLOAT16, - DataType.BF16: native.Nvfp4ColdPageRuntimeType.BFLOAT16, - DataType.FP8: native.Nvfp4ColdPageRuntimeType.FP8_E4M3, - }.get(runtime_dtype) - if runtime_type is None: + if attention_layers and runtime_type is None: raise RuntimeError( "NVFP4 cold-page compression supports FP16, BF16, or FP8 " f"Attention KV, not {runtime_dtype}" @@ -153,8 +285,6 @@ def create_cold_page_codec( int(pp_layers[layer_id]), _IDENTITY_NVFP4_SCALES ) else: - # Target projection scales describe neither MLA latents nor a - # separately numbered draft model. orig_quant, quant_orig = _IDENTITY_NVFP4_SCALES num_kv_heads = int(num_kv_heads_per_layer[layer_id]) @@ -177,41 +307,33 @@ def create_cold_page_codec( } cursor = scale_base + len(compressed_roles) * scale_bytes - buffer_layouts = [] - for scale_index, role in enumerate(compressed_roles): - scales = native.Nvfp4ColdPageScales() - scales.nvfp4_scale_orig_quant = orig_quant[scale_index] - scales.nvfp4_scale_quant_orig = quant_orig[scale_index] - scales.fp8_scale_orig_quant = 1.0 - scales.fp8_scale_quant_orig = 1.0 - - buffer_layout = native.Nvfp4ColdPageBufferLayout() - buffer_layout.role = role - buffer_layout.cold_data_offset = data_offsets[role] - buffer_layout.cold_scale_offset = scale_offsets[role] - buffer_layout.scales = scales - buffer_layouts.append(buffer_layout) - + buffer_layouts = [ + _Nvfp4BufferLayout( + role=role, + data_offset=data_offsets[role], + scale_offset=scale_offsets[role], + scales=_Nvfp4Scales(orig_quant[index], quant_orig[index]), + ) + for index, role in enumerate(compressed_roles) + ] for buffer in layer.buffers: role = str(buffer.role) - if role in compressed_roles: - continue - buffer_layout = native.Nvfp4ColdPageBufferLayout() - buffer_layout.role = role - buffer_layout.cold_data_offset = cursor - buffer_layouts.append(buffer_layout) - cursor += _buffer_bytes(buffer, tokens_per_page) - - cold_page_bytes = _align_up(cursor) - layer_layout = native.Nvfp4ColdPageLayerLayout() - layer_layout.layer_id = layer_id - layer_layout.runtime_type = runtime_type - layer_layout.num_kv_heads = num_kv_heads - layer_layout.tokens_per_page = tokens_per_page - layer_layout.head_dim = head_dim - layer_layout.cold_page_bytes = cold_page_bytes - layer_layout.cold_padding_offset = cursor - layer_layout.buffers = buffer_layouts - layer_layouts.append(layer_layout) - - return native.create_nvfp4_cold_page_codec(layer_layouts) + if role not in compressed_roles: + buffer_layouts.append(_Nvfp4BufferLayout(role=role, data_offset=cursor)) + cursor += _buffer_bytes(buffer, tokens_per_page) + + layer_layouts.append( + _Nvfp4LayerLayout( + layer_id=layer_id, + runtime_type=runtime_type, + num_kv_heads=num_kv_heads, + tokens_per_page=tokens_per_page, + head_dim=head_dim, + cold_page_bytes=_align_up(cursor), + padding_offset=cursor, + buffers=tuple(buffer_layouts), + ) + ) + + policy = Nvfp4ColdPagePolicy(layer_layouts) + return native.create_python_cold_page_codec(policy.layer_ids, policy) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 3a58b7a73f69..d08c2e93936a 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -2,31 +2,85 @@ # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. """Quantization policies for KVCM V2 cold pages.""" +from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Sequence from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager -from .nvfp4 import Nvfp4ColdPagePolicy if TYPE_CHECKING: from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig +class ColdPageQuantizationMethod(ABC): + """Configured quantization method that creates one codec per KVCM.""" + + @abstractmethod + def create_cold_page_codec( + self, + cache_config: object, + *, + runtime_dtype: DataType, + pp_layers: Sequence[int], + num_kv_heads_per_layer: Sequence[int], + head_dim_per_layer: Sequence[int], + is_draft: bool = False, + ) -> object: + """Create a codec using this method's immutable calibration.""" + + +class ColdPageCodecPolicy(ABC): + """Per-KVCM Python callback contract used by the generic native codec.""" + + @property + @abstractmethod + def layer_ids(self) -> tuple[int, ...]: + """Layers transformed by this policy; other lifecycles stay lossless.""" + + @abstractmethod + def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: + """Prepare one immutable method program per owned lifecycle.""" + + @abstractmethod + def encode( + self, + program_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + """Submit one complete hot-to-cold Page batch.""" + + @abstractmethod + def decode( + self, + program_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + """Submit one complete cold-to-hot Page batch.""" + + class ColdPageQuantizationCompression(KVCacheCompressionManager): - """Select and own the configured cold-page quantization policy.""" + """Select and own the configured cold-page quantization method.""" uses_iteration_lifecycle = False provides_cold_page_codec = True def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) - policies = {"nvfp4": Nvfp4ColdPagePolicy} + from .nvfp4 import Nvfp4ColdPageQuantization + + methods = {"nvfp4": Nvfp4ColdPageQuantization} try: - policy = policies[config.quant] + method = methods[config.quant] except KeyError as error: raise NotImplementedError( f"Unsupported cold-page quantization format {config.quant!r}" ) from error - self._policy = policy(config.scale_checkpoint_path) + self._method: ColdPageQuantizationMethod = method(config.scale_checkpoint_path) def create_cold_page_codec( self, @@ -38,9 +92,9 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create the native codec selected by the quantization policy.""" + """Create the native codec selected by the quantization method.""" - return self._policy.create_cold_page_codec( + return self._method.create_cold_page_codec( cache_config, runtime_dtype=runtime_dtype, pp_layers=pp_layers, diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index dadc7e9af538..bda086325da8 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -56,28 +56,24 @@ def _cache_config(*layers): return SimpleNamespace(tokens_per_block=64, layers=tuple(configs)) -def _native(): - def config(): - return SimpleNamespace() - - def buffer_layout(): - return SimpleNamespace(scales=None, cold_scale_offset=0) - +def _native() -> tuple[SimpleNamespace, MagicMock]: codec = MagicMock() module = SimpleNamespace( - Nvfp4ColdPageRuntimeType=SimpleNamespace( - FLOAT16="native-fp16", - BFLOAT16="native-bf16", - FP8_E4M3="native-fp8", - ), - Nvfp4ColdPageScales=config, - Nvfp4ColdPageBufferLayout=buffer_layout, - Nvfp4ColdPageLayerLayout=config, - create_nvfp4_cold_page_codec=MagicMock(return_value=codec), + ColdPageLifecycleProperties=lambda: SimpleNamespace(), + ColdPageIndexLocation=SimpleNamespace(HOST="host"), + create_python_cold_page_codec=MagicMock(return_value=codec), ) return module, codec +def _policy(native: SimpleNamespace) -> object: + return native.create_python_cold_page_codec.call_args.args[1] + + +def _plans(native: SimpleNamespace) -> list[object]: + return list(_policy(native)._layer_layouts.values()) + + def _write_quant_metadata(directory, algorithm="NVFP4"): metadata = { "producer": {"name": "modelopt"}, @@ -129,13 +125,13 @@ def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_p ) assert result is codec - plans = native.create_nvfp4_cold_page_codec.call_args.args[0] + plans = _plans(native) assert [plan.layer_id for plan in plans] == [0, 1, 2] - assert [plan.runtime_type for plan in plans] == ["native-bf16"] * 3 + assert [plan.runtime_type for plan in plans] == [1] * 3 assert [ ( - tuple(buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers), - tuple(buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers), + tuple(buffer.scales.nvfp4_orig_quant for buffer in plan.buffers), + tuple(buffer.scales.nvfp4_quant_orig for buffer in plan.buffers), ) for plan in plans ] == [ @@ -158,7 +154,7 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: num_kv_heads_per_layer=(8,), head_dim_per_layer=(128,), ) - target_plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + target_plan = _plans(native)[0] manager.create_cold_page_codec( _cache_config((0, "attention")), runtime_dtype=DataType.BF16, @@ -168,16 +164,16 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: is_draft=True, ) - draft_plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert [buffer.scales.nvfp4_scale_orig_quant for buffer in target_plan.buffers] == [ + draft_plan = _plans(native)[0] + assert [buffer.scales.nvfp4_orig_quant for buffer in target_plan.buffers] == [ 2.0, 4.0, ] - assert [buffer.scales.nvfp4_scale_orig_quant for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_orig_quant for buffer in draft_plan.buffers] == [ 1.0, 1.0, ] - assert [buffer.scales.nvfp4_scale_quant_orig for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in draft_plan.buffers] == [ 1.0, 1.0, ] @@ -197,21 +193,21 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): head_dim_per_layer=(128,), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + plan = _plans(native)[0] assert [buffer.role for buffer in plan.buffers] == ["key", "value"] - assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 1280] - assert [buffer.cold_scale_offset for buffer in plan.buffers] == [2560, 2720] + assert [buffer.data_offset for buffer in plan.buffers] == [0, 1280] + assert [buffer.scale_offset for buffer in plan.buffers] == [2560, 2720] assert plan.cold_page_bytes == 2880 - assert plan.cold_padding_offset == plan.cold_page_bytes - assert plan.runtime_type == "native-fp16" + assert plan.padding_offset == plan.cold_page_bytes + assert plan.runtime_type == 0 assert plan.num_kv_heads == 4 assert plan.tokens_per_page == 5 assert plan.head_dim == 128 - assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ 1.0, 1.0, ] - assert [buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in plan.buffers] == [ 1.0, 1.0, ] @@ -241,18 +237,18 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: head_dim_per_layer=(32,), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + plan = _plans(native)[0] assert [buffer.role for buffer in plan.buffers] == ["key", "value"] - assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 80] - assert [buffer.cold_scale_offset for buffer in plan.buffers] == [160, 170] - assert plan.cold_padding_offset == 180 + assert [buffer.data_offset for buffer in plan.buffers] == [0, 80] + assert [buffer.scale_offset for buffer in plan.buffers] == [160, 170] + assert plan.padding_offset == 180 assert plan.cold_page_bytes == 192 def test_provider_creates_one_native_codec_per_kv_cache_manager(): native, _ = _native() codecs = (object(), object()) - native.create_nvfp4_cold_page_codec.side_effect = codecs + native.create_python_cold_page_codec.side_effect = codecs provider = _manager() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): @@ -268,7 +264,61 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): ) assert results == codecs - assert native.create_nvfp4_cold_page_codec.call_count == 2 + assert native.create_python_cold_page_codec.call_count == 2 + + +def test_policy_forwards_a_4096_page_batch_through_one_custom_op(monkeypatch) -> None: + native, _ = _native() + program = object() + prepare = MagicMock(return_value=program) + encode = MagicMock() + decode = MagicMock() + monkeypatch.setattr( + torch.ops.trtllm, + "prepare_nvfp4_cold_page_program", + prepare, + raising=False, + ) + monkeypatch.setattr( + torch.ops.trtllm, + "nvfp4_cold_page_encode", + encode, + raising=False, + ) + monkeypatch.setattr( + torch.ops.trtllm, + "nvfp4_cold_page_decode", + decode, + raising=False, + ) + + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(16,), + ) + policy = _policy(native) + hot = { + role: SimpleNamespace( + raw_base=0x1000 + index * 0x1000, + raw_slot_bytes=4096, + raw_bytes=2048, + ) + for index, role in enumerate(("key", "value")) + } + properties = policy.configure([SimpleNamespace(layers={0: hot})]) + + assert prepare.call_count == 1 + assert properties[0].cold_page_bytes == 1152 + assert properties[0].page_index_location == "host" + + policy.encode(0, 0x3000, 0x4000, 4096, 0x5000) + policy.decode(0, 0x3000, 0x4000, 4096, 0x5000) + encode.assert_called_once_with(program, 0x3000, 0x4000, 4096, 0x5000) + decode.assert_called_once_with(program, 0x3000, 0x4000, 4096, 0x5000) def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: @@ -278,7 +328,7 @@ def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: ) with patch( "tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page." - "quantization_for_cold_page.Nvfp4ColdPagePolicy" + "nvfp4.Nvfp4ColdPageQuantization" ) as policy: with pytest.raises(NotImplementedError, match="future-format"): ColdPageQuantizationCompression(config) @@ -385,9 +435,9 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): num_kv_heads_per_layer=(0, 8), head_dim_per_layer=(128, 128), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + plan = _plans(native)[0] assert plan.layer_id == 1 - assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ 2.0, 4.0, ] @@ -401,7 +451,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - native.create_nvfp4_cold_page_codec.assert_called_with([]) + assert native.create_python_cold_page_codec.call_args.args[0] == () def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -429,20 +479,20 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): ) assert result is codec - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + plan = _plans(native)[0] assert plan.layer_id == 0 assert plan.cold_page_bytes == 29184 - assert plan.cold_padding_offset == 29184 + assert plan.padding_offset == 29184 assert [buffer.role for buffer in plan.buffers] == ["key", "index_key"] assert [buffer.scales is not None for buffer in plan.buffers] == [True, False] - assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 20736] - assert [buffer.cold_scale_offset for buffer in plan.buffers] == [18432, 0] - assert plan.runtime_type == "native-bf16" + assert [buffer.data_offset for buffer in plan.buffers] == [0, 20736] + assert [buffer.scale_offset for buffer in plan.buffers] == [18432, 0] + assert plan.runtime_type == 1 assert plan.num_kv_heads == 1 assert plan.tokens_per_page == 64 assert plan.head_dim == 576 scales = plan.buffers[0].scales - assert scales.nvfp4_scale_orig_quant == scales.nvfp4_scale_quant_orig == 1.0 + assert scales.nvfp4_orig_quant == scales.nvfp4_quant_orig == 1.0 assert plan.buffers[1].scales is None @@ -471,15 +521,15 @@ def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: head_dim_per_layer=(32,), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] + plan = _plans(native)[0] assert [buffer.role for buffer in plan.buffers] == [ "key", "index_key", "rope_state", ] assert [buffer.scales is not None for buffer in plan.buffers] == [True, False, False] - assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 90, 158] - assert plan.cold_padding_offset == 165 + assert [buffer.data_offset for buffer in plan.buffers] == [0, 90, 158] + assert plan.padding_offset == 165 assert plan.cold_page_bytes == 176 @@ -511,9 +561,9 @@ def test_tokens_per_block_override_expands_lossless_bytes() -> None: head_dim_per_layer=(16,), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert [buffer.cold_data_offset for buffer in plan.buffers] == [0, 36] - assert plan.cold_padding_offset == 42 + plan = _plans(native)[0] + assert [buffer.data_offset for buffer in plan.buffers] == [0, 36] + assert plan.padding_offset == 42 assert plan.cold_page_bytes == 48 @@ -554,7 +604,7 @@ def test_mla_model_layouts_are_built_in_python( head_dim_per_layer=(576,) * expected_layers, ) - plans = native.create_nvfp4_cold_page_codec.call_args.args[0] + plans = _plans(native) assert len(plans) == expected_layers assert sum(len(plan.buffers) for plan in plans) == expected_buffers assert sum(plan.cold_page_bytes for plan in plans) == expected_bytes @@ -572,18 +622,18 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): head_dim_per_layer=(128,), ) - plan = native.create_nvfp4_cold_page_codec.call_args.args[0][0] - assert plan.runtime_type == "native-fp8" - assert [buffer.scales.nvfp4_scale_orig_quant for buffer in plan.buffers] == [ + plan = _plans(native)[0] + assert plan.runtime_type == 2 + assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ 2.0, 4.0, ] - assert [buffer.scales.nvfp4_scale_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in plan.buffers] == [ 0.5, 0.25, ] assert all( - buffer.scales.fp8_scale_orig_quant == buffer.scales.fp8_scale_quant_orig == 1.0 + buffer.scales.fp8_orig_quant == buffer.scales.fp8_quant_orig == 1.0 for buffer in plan.buffers ) From af6f49db13be592929de6a9e74f4257d22e190a3 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 00:30:38 -0700 Subject: [PATCH 13/29] [None][refactor] Run cold-page methods through Python policy Signed-off-by: tianruih --- .../kernels/nvfp4ColdPageKernels.cu | 311 +++++++----------- .../kernels/nvfp4ColdPageKernels.h | 70 +--- .../nativeColdPageCodec.cpp | 74 +++-- .../nativeColdPageCodec.h | 13 +- .../nanobind/kvCacheCompression/bindings.cpp | 50 +-- cpp/tensorrt_llm/thop/CMakeLists.txt | 5 +- .../thop/coldPageMethods/nvfp4ColdPageOp.cu | 84 +++++ cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp | 229 ------------- .../kernels/nvfp4ColdPageKernelsTest.cpp | 310 ++++++++--------- .../coldPageCodecTest.cpp | 27 +- .../quantization_for_cold_page/nvfp4.py | 241 +++++++++----- .../quantization_for_cold_page.py | 6 +- .../test_quantization_for_cold_page.py | 293 ++++++++++++----- .../test_nvfp4_cold_page_op.py | 48 +++ 14 files changed, 856 insertions(+), 905 deletions(-) create mode 100644 cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu delete mode 100644 cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp create mode 100644 tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index 0651c8560e8a..577a9778632a 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -26,15 +26,12 @@ #include #include -#include #include #include #include #include #include -#include #include -#include TRTLLM_NAMESPACE_BEGIN @@ -50,8 +47,6 @@ constexpr std::uint32_t kAsyncStages = 4; constexpr std::uint32_t kHostLoadAsyncStages = 8; constexpr std::uint32_t kMappedHostGridSplits = 1; constexpr std::uint32_t kMaxTasksPerLaunch = 256; -// Bound by-value buffer metadata to CUDA's kernel-parameter limit. -constexpr std::uint32_t kMaxBuffersPerLaunch = kNvfp4ColdPageMaxBuffersPerLaunch; constexpr std::uint32_t kElementsPerHalfGroup = 8; constexpr std::uint32_t kElementsPerScaleGroup = 16; constexpr std::uint32_t kHalfGroupsPerScaleGroup = kElementsPerScaleGroup / kElementsPerHalfGroup; @@ -60,22 +55,74 @@ constexpr std::uint32_t kMaxScaleBytesPerTile = 1024; constexpr std::uint32_t kMaxHalfGroupsPerTile = kHalfGroupsPerScaleGroup * kMaxScaleBytesPerTile; constexpr std::size_t kKernelParameterLimitBytes = 32764; +enum WideField : std::uint32_t +{ + kRawBase, + kRawSlotBytes, + kRawBytes, + kColdDataOffset, + kColdScaleOffset, + kColdPaddingOffset, +}; + +enum IntegerField : std::uint32_t +{ + kColdPaddingBytes, + kTransform, + kNumKvHeads, + kTokensPerPage, + kHeadDim, +}; + +enum ScaleField : std::uint32_t +{ + kNvfp4ScaleOrigQuant, + kNvfp4ScaleQuantOrig, + kFp8ScaleOrigQuant, + kFp8ScaleQuantOrig, +}; + +struct Nvfp4ColdPageKernelParams +{ + std::int32_t numKvHeads; + std::int32_t tokensPerPage; + std::int32_t headDim; + float nvfp4ScaleOrigQuant; + float nvfp4ScaleQuantOrig; + float fp8ScaleOrigQuant; + float fp8ScaleQuantOrig; +}; + +enum class Nvfp4ColdPageTransform : std::int32_t +{ + kNvfp4 = 0, + kLosslessCopy = 1, +}; + +struct Nvfp4ColdPageBuffer +{ + std::uintptr_t rawBase; + std::size_t rawSlotBytes; + std::size_t rawBytes; + std::size_t coldDataOffset; + std::size_t coldScaleOffset; + std::size_t coldPaddingOffset; + std::uint32_t coldPaddingBytes; + Nvfp4ColdPageTransform transform; + Nvfp4ColdPageKernelParams params; +}; + // Keep every CTA iteration and shared-memory tile on complete 16-value scale groups. static_assert(kElementsPerScaleGroup % kElementsPerHalfGroup == 0, "An NVFP4 scale group must contain a whole number of half-groups"); static_assert(kHalfGroupsPerScaleGroup == 2U, "NVFP4 stores one scale for two eight-value half-groups"); static_assert(kThreadsPerBlock % kHalfGroupsPerScaleGroup == 0, "A CTA iteration must not split an NVFP4 scale group"); -static_assert(kMaxScaleBytesPerTile > 0, "A transfer tile must make forward progress"); static_assert( kMaxScaleBytesPerTile % sizeof(uint4) == 0, "A full scale tile must preserve the 16-byte transfer fast path"); -// Buffer plans are passed to CUDA kernels as raw by-value arguments. -static_assert(std::is_trivially_copyable_v, "Buffer plans must remain raw-copyable"); - // Keep both kernel argument packs within CUDA's modern 32,764-byte limit. -static_assert(sizeof(std::array) - + sizeof(std::array) + 2U * sizeof(std::uintptr_t) - + 3U * sizeof(std::uint32_t) +static_assert(sizeof(std::array) + sizeof(Nvfp4ColdPageWideTable) + + sizeof(Nvfp4ColdPageIntegerTable) + sizeof(Nvfp4ColdPageScaleTable) + 2U * sizeof(std::uintptr_t) <= kKernelParameterLimitBytes, "Cold-page kernel arguments exceed CUDA's parameter limit"); @@ -110,8 +157,22 @@ struct OnboardBufferTask std::uint8_t* raw; }; -__device__ OffloadBufferTask resolveOffloadTask(ColdPageIndexPair const& page, Nvfp4ColdPageBufferPlan const& buffer, - std::uint8_t* coldBase, std::size_t coldPageBytes) +__device__ Nvfp4ColdPageBuffer loadBuffer(std::uint32_t index, Nvfp4ColdPageWideTable const& wide, + Nvfp4ColdPageIntegerTable const& integers, Nvfp4ColdPageScaleTable const& scales) +{ + auto const& w = wide[index]; + auto const& i = integers[index]; + auto const& s = scales[index]; + return {static_cast(w[kRawBase]), static_cast(w[kRawSlotBytes]), + static_cast(w[kRawBytes]), static_cast(w[kColdDataOffset]), + static_cast(w[kColdScaleOffset]), static_cast(w[kColdPaddingOffset]), + static_cast(i[kColdPaddingBytes]), static_cast(i[kTransform]), + {i[kNumKvHeads], i[kTokensPerPage], i[kHeadDim], s[kNvfp4ScaleOrigQuant], s[kNvfp4ScaleQuantOrig], + s[kFp8ScaleOrigQuant], s[kFp8ScaleQuantOrig]}}; +} + +__device__ OffloadBufferTask resolveOffloadTask( + ColdPageIndexPair const& page, Nvfp4ColdPageBuffer const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.src); auto* coldPage = coldBase + static_cast(page.dst) * coldPageBytes; @@ -119,7 +180,7 @@ __device__ OffloadBufferTask resolveOffloadTask(ColdPageIndexPair const& page, N coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, coldPage + buffer.coldPaddingOffset}; } -__device__ OnboardBufferTask resolveOnboardTask(ColdPageIndexPair const& page, Nvfp4ColdPageBufferPlan const& buffer, +__device__ OnboardBufferTask resolveOnboardTask(ColdPageIndexPair const& page, Nvfp4ColdPageBuffer const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.dst); @@ -210,7 +271,7 @@ __device__ void flushCompactRangeToHost(std::uint8_t const* compactStages, Offlo } // Zero codec-specified record padding so persisted cold Slots are deterministic. -__device__ void clearColdPadding(OffloadBufferTask const& task, Nvfp4ColdPageBufferPlan const& buffer) +__device__ void clearColdPadding(OffloadBufferTask const& task, Nvfp4ColdPageBuffer const& buffer) { if (blockIdx.x != 0U) { @@ -406,15 +467,14 @@ __device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firs template __global__ void offloadFrom16BitTiledKernel( std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, std::uint8_t* coldBase, - std::size_t coldPageBytes, std::uint32_t numBuffers) + Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, + Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t* coldBase, std::size_t coldPageBytes) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) asm volatile("griddepcontrol.launch_dependents;\n"); std::uint32_t const bufferIndex = blockIdx.y; - assert(bufferIndex < numBuffers); - auto const& buffer = buffers[bufferIndex]; + auto const buffer = loadBuffer(bufferIndex, wide, integers, scales); auto const task = resolveOffloadTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); @@ -487,15 +547,14 @@ __global__ void offloadFrom16BitTiledKernel( // FP8 E4M3 GPU Page -> mapped-Host NVFP4 in bounded tiles. __global__ void offloadFromFp8TiledKernel( std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, std::uint8_t* coldBase, - std::size_t coldPageBytes, std::uint32_t numBuffers) + Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, + Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t* coldBase, std::size_t coldPageBytes) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) asm volatile("griddepcontrol.launch_dependents;\n"); std::uint32_t const bufferIndex = blockIdx.y; - assert(bufferIndex < numBuffers); - auto const& buffer = buffers[bufferIndex]; + auto const buffer = loadBuffer(bufferIndex, wide, integers, scales); auto const task = resolveOffloadTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); @@ -630,15 +689,14 @@ __device__ void loadCompactRangeFromHost(std::uint8_t* compactStages, OnboardBuf // Mapped-Host NVFP4 -> runtime GPU Page in bounded tiles. template __global__ void onboardTiledKernel(std::array const __grid_constant__ pages, - std::array const __grid_constant__ buffers, - std::uint8_t const* coldBase, std::size_t coldPageBytes, std::uint32_t numBuffers) + Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, + Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t const* coldBase, std::size_t coldPageBytes) { #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) asm volatile("griddepcontrol.launch_dependents;\n"); std::uint32_t const bufferIndex = blockIdx.y; - assert(bufferIndex < numBuffers); - auto const& buffer = buffers[bufferIndex]; + auto const buffer = loadBuffer(bufferIndex, wide, integers, scales); auto const task = resolveOnboardTask(pages[blockIdx.z], buffer, coldBase, coldPageBytes); asm volatile("griddepcontrol.wait;\n" : : : "memory"); @@ -697,106 +755,12 @@ __global__ void onboardTiledKernel(std::array 0, "numKvHeads must be positive"); - TLLM_CHECK_WITH_INFO(params.tokensPerPage > 0, "tokensPerPage must be positive"); - TLLM_CHECK_WITH_INFO(params.headDim > 0 && params.headDim % kElementsPerScaleGroup == 0, - "headDim must be positive and divisible by 16, got %d", params.headDim); - - std::uint64_t const rows - = static_cast(params.numKvHeads) * static_cast(params.tokensPerPage); - constexpr std::uint64_t maxHalfGroups = std::numeric_limits::max() / kElementsPerHalfGroup; - std::uint64_t const halfGroupsPerRow = static_cast(params.headDim / kElementsPerHalfGroup); - TLLM_CHECK_WITH_INFO(rows <= maxHalfGroups / halfGroupsPerRow, - "Page geometry exceeds the 32-bit compact-offset range: " - "heads=%d, tokens=%d, headDim=%d", - params.numKvHeads, params.tokensPerPage, params.headDim); - - TLLM_CHECK_WITH_INFO(std::isfinite(params.nvfp4ScaleOrigQuant) && params.nvfp4ScaleOrigQuant > 0.0F, - "NVFP4 original-to-quantized scale must be finite and positive"); - TLLM_CHECK_WITH_INFO(std::isfinite(params.nvfp4ScaleQuantOrig) && params.nvfp4ScaleQuantOrig > 0.0F, - "NVFP4 quantized-to-original scale must be finite and positive"); - if (isFp8Runtime) - { - TLLM_CHECK_WITH_INFO(std::isfinite(params.fp8ScaleOrigQuant) && params.fp8ScaleOrigQuant > 0.0F, - "FP8 original-to-quantized scale must be finite and positive"); - TLLM_CHECK_WITH_INFO(std::isfinite(params.fp8ScaleQuantOrig) && params.fp8ScaleQuantOrig > 0.0F, - "FP8 quantized-to-original scale must be finite and positive"); - } -} - -struct ColdInterval -{ - std::size_t begin; - std::size_t end; -}; - -void addColdInterval(std::vector& intervals, std::size_t offset, std::size_t bytes, - std::size_t coldPageBytes, char const* label) -{ - if (bytes == 0U) - { - return; - } - TLLM_CHECK_WITH_INFO( - offset <= coldPageBytes && bytes <= coldPageBytes - offset, "%s exceeds the cold Page stride", label); - intervals.push_back({offset, offset + bytes}); -} - -// Validate one raw/cold buffer mapping and append its occupied cold intervals. -void validateBufferPlan(Nvfp4ColdPageBufferPlan const& buffer, std::size_t coldPageBytes, bool isFp8Runtime, - std::vector& intervals) -{ - TLLM_CHECK_WITH_INFO(buffer.rawBase != 0U, "rawBase must not be null"); - TLLM_CHECK_WITH_INFO(buffer.rawBytes > 0 && buffer.rawBytes <= buffer.rawSlotBytes, - "Raw buffer bytes must be positive and fit within the GPU Slot stride"); - - switch (buffer.transform) - { - case Nvfp4ColdPageTransform::kNvfp4: - { - validateNvfp4Params(buffer.params, isFp8Runtime); - TLLM_CHECK_WITH_INFO( - buffer.rawBase % alignof(uint4) == 0, "rawBase must be aligned to %zu bytes", alignof(uint4)); - TLLM_CHECK_WITH_INFO(buffer.rawSlotBytes % alignof(uint4) == 0, - "GPU raw Slot stride must be aligned to %zu bytes for NVFP4", alignof(uint4)); - std::uint64_t const elements = static_cast(buffer.params.numKvHeads) - * static_cast(buffer.params.tokensPerPage) - * static_cast(buffer.params.headDim); - std::uint64_t const expectedRawBytes = elements * (isFp8Runtime ? 1U : 2U); - TLLM_CHECK_WITH_INFO(buffer.rawBytes == static_cast(expectedRawBytes), - "Raw buffer size does not match NVFP4 geometry and runtime type"); - addColdInterval(intervals, buffer.coldDataOffset, static_cast(elements / 2U), coldPageBytes, - "NVFP4 packed-data interval"); - addColdInterval(intervals, buffer.coldScaleOffset, static_cast(elements / 16U), coldPageBytes, - "NVFP4 scale interval"); - break; - } - case Nvfp4ColdPageTransform::kLosslessCopy: - addColdInterval(intervals, buffer.coldDataOffset, buffer.rawBytes, coldPageBytes, "Lossless-data interval"); - break; - default: TLLM_THROW("Unsupported NVFP4 cold-page buffer transform"); - } - addColdInterval( - intervals, buffer.coldPaddingOffset, buffer.coldPaddingBytes, coldPageBytes, "Cold-record padding interval"); -} - -// Host submission path. - // Submit one whole KVCM Page batch through the fixed 256-descriptor kernel ABI. template -void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, - ColdPointer coldBase, cudaStream_t stream) -{ - static_assert(std::is_trivially_copyable_v>, - "Page-index pairs must remain raw-copyable kernel arguments"); - static_assert( - sizeof(std::array) == sizeof(ColdPageIndexPair) * kMaxTasksPerLaunch, - "Page-index pair arrays must not add ABI padding"); +void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, std::int64_t const* wide, + std::int32_t const* integers, float const* scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, ColdPointer coldBase, cudaStream_t stream) +{ dim3 const block(kThreadsPerBlock); auto const* pageBytes = static_cast(pages); std::size_t offset = 0; @@ -819,17 +783,15 @@ void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, Nv attribute.val.programmaticStreamSerializationAllowed = common::getEnvEnablePDL() ? 1 : 0; cudaLaunchConfig_t config{}; - config.gridDim = dim3(kMappedHostGridSplits, plan.numBuffers, numChunkPages); + config.gridDim = dim3(kMappedHostGridSplits, numBuffers, numChunkPages); config.blockDim = block; - config.dynamicSmemBytes = compactStageBytesForHalfGroups(plan.maxHalfGroupsPerTile); + config.dynamicSmemBytes = compactStageBytesForHalfGroups(maxHalfGroupsPerTile); config.stream = stream; config.attrs = &attribute; config.numAttrs = 1; - auto coldPageBytes = plan.coldPageBytes; - auto numBuffers = plan.numBuffers; - void* arguments[] = {const_cast(chunkPages), - const_cast(plan.buffers.data()), &coldBase, &coldPageBytes, &numBuffers}; + void* arguments[] = {const_cast(chunkPages), const_cast(wide), + const_cast(integers), const_cast(scales), &coldBase, &coldPageBytes}; TLLM_CUDA_CHECK(cudaLaunchKernelExC(&config, reinterpret_cast(kernel), arguments)); offset += numChunkPages; } @@ -837,55 +799,9 @@ void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, Nv } // namespace -Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) -{ - // TODO: Make this codec-private, then remove caller-guaranteed admission checks. - TLLM_CHECK_WITH_INFO(common::isSM100Family(), "NVFP4 cold-page kernels require an SM100-family GPU"); - TLLM_CHECK_WITH_INFO(!buffers.empty(), "NVFP4 cold-page launch requires at least one buffer"); - TLLM_CHECK_WITH_INFO(buffers.size() <= kMaxBuffersPerLaunch, - "NVFP4 cold-page launch supports at most %u local buffers, got %zu", kMaxBuffersPerLaunch, buffers.size()); - TLLM_CHECK_WITH_INFO(coldPageBytes > 0, "Cold Page stride must be positive"); - TLLM_CHECK_WITH_INFO( - coldPageBytes % alignof(uint4) == 0, "Cold Page stride must be aligned to %zu bytes", alignof(uint4)); - - bool isFp8Runtime = false; - switch (runtimeType) - { - case Nvfp4ColdPageRuntimeType::kFloat16: - case Nvfp4ColdPageRuntimeType::kBfloat16: break; - case Nvfp4ColdPageRuntimeType::kFp8E4m3: isFp8Runtime = true; break; - default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); - } - - Nvfp4ColdPagePreparedPlan plan; - plan.numBuffers = static_cast(buffers.size()); - plan.coldPageBytes = coldPageBytes; - plan.runtimeType = runtimeType; - std::copy(buffers.begin(), buffers.end(), plan.buffers.begin()); - std::vector intervals; - intervals.reserve(3U * buffers.size()); - for (auto const& buffer : buffers) - { - validateBufferPlan(buffer, coldPageBytes, isFp8Runtime, intervals); - if (buffer.transform == Nvfp4ColdPageTransform::kNvfp4) - { - plan.maxHalfGroupsPerTile = std::max(plan.maxHalfGroupsPerTile, tileHalfGroupCount(buffer.params)); - } - } - std::sort(intervals.begin(), intervals.end(), - [](ColdInterval const& lhs, ColdInterval const& rhs) - { return lhs.begin < rhs.begin || (lhs.begin == rhs.begin && lhs.end < rhs.end); }); - for (std::size_t index = 1; index < intervals.size(); ++index) - { - TLLM_CHECK_WITH_INFO( - intervals[index - 1U].end <= intervals[index].begin, "Cold record intervals must not overlap"); - } - return plan; -} - -void invokeNvfp4ColdPageEncode( - void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, void* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageEncode(void const* pages, std::size_t numPages, std::int64_t const* wide, + std::int32_t const* integers, float const* scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType, void* coldBase, cudaStream_t stream) { if (numPages == 0) { @@ -893,26 +809,27 @@ void invokeNvfp4ColdPageEncode( } TLLM_CHECK_WITH_INFO(pages != nullptr, "pages must not be null"); TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); - switch (plan.runtimeType) + switch (runtimeType) { case Nvfp4ColdPageRuntimeType::kFloat16: - launchPageChunks( - offloadFrom16BitTiledKernel, pages, numPages, plan, static_cast(coldBase), stream); + launchPageChunks(offloadFrom16BitTiledKernel, pages, numPages, wide, integers, scales, numBuffers, + maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kBfloat16: - launchPageChunks(offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, numPages, plan, - static_cast(coldBase), stream); + launchPageChunks(offloadFrom16BitTiledKernel<__nv_bfloat16>, pages, numPages, wide, integers, scales, + numBuffers, maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kFp8E4m3: - launchPageChunks( - offloadFromFp8TiledKernel, pages, numPages, plan, static_cast(coldBase), stream); + launchPageChunks(offloadFromFp8TiledKernel, pages, numPages, wide, integers, scales, numBuffers, + maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } } -void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, - void const* coldBase, cudaStream_t stream) +void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, std::int64_t const* wide, + std::int32_t const* integers, float const* scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType, void const* coldBase, cudaStream_t stream) { if (numPages == 0) { @@ -920,19 +837,19 @@ void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4Col } TLLM_CHECK_WITH_INFO(pages != nullptr, "pages must not be null"); TLLM_CHECK_WITH_INFO(coldBase != nullptr, "coldBase must not be null"); - switch (plan.runtimeType) + switch (runtimeType) { case Nvfp4ColdPageRuntimeType::kFloat16: - launchPageChunks( - onboardTiledKernel, pages, numPages, plan, static_cast(coldBase), stream); + launchPageChunks(onboardTiledKernel, pages, numPages, wide, integers, scales, numBuffers, + maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kBfloat16: - launchPageChunks(onboardTiledKernel<__nv_bfloat16>, pages, numPages, plan, - static_cast(coldBase), stream); + launchPageChunks(onboardTiledKernel<__nv_bfloat16>, pages, numPages, wide, integers, scales, numBuffers, + maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; case Nvfp4ColdPageRuntimeType::kFp8E4m3: - launchPageChunks(onboardTiledKernel<__nv_fp8_e4m3>, pages, numPages, plan, - static_cast(coldBase), stream); + launchPageChunks(onboardTiledKernel<__nv_fp8_e4m3>, pages, numPages, wide, integers, scales, numBuffers, + maxHalfGroupsPerTile, coldPageBytes, static_cast(coldBase), stream); break; default: TLLM_THROW("Unsupported NVFP4 cold-page runtime type"); } diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h index a3eb70d7116a..95bce5177ed8 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h @@ -25,14 +25,13 @@ #include #include #include -#include TRTLLM_NAMESPACE_BEGIN namespace kernels { -//! Active GPU representation encoded into cold Pages; values are part of the Python custom-op ABI. +//! Active GPU representation encoded into cold Pages. enum class Nvfp4ColdPageRuntimeType : std::uint8_t { kFloat16 = 0, @@ -42,64 +41,27 @@ enum class Nvfp4ColdPageRuntimeType : std::uint8_t using ColdPageIndexPair = ::tensorrt_llm::kv_cache_compression::ColdPageIndexPair; -//! Per-buffer geometry and scales for one NVFP4 record in HND order. -//! `headDim` is a multiple of 16; `*OrigQuant` encodes and `*QuantOrig` decodes this buffer. -struct Nvfp4ColdPageKernelParams -{ - std::int32_t numKvHeads; - std::int32_t tokensPerPage; - std::int32_t headDim; - float nvfp4ScaleOrigQuant; - float nvfp4ScaleQuantOrig; - float fp8ScaleOrigQuant; - float fp8ScaleQuantOrig; -}; - -//! Per-buffer transform selected by the Python custom-op metadata ABI. -enum class Nvfp4ColdPageTransform : std::uint8_t -{ - kNvfp4 = 0, - //! Byte-exact copy for an Attention side buffer such as DSA index_key. - kLosslessCopy = 1, -}; - -//! Immutable transform plan for one hot buffer and its fixed-offset cold record. -struct Nvfp4ColdPageBufferPlan -{ - std::uintptr_t rawBase; - std::size_t rawSlotBytes; - std::size_t rawBytes; - std::size_t coldDataOffset; - std::size_t coldScaleOffset; - std::size_t coldPaddingOffset; - std::uint32_t coldPaddingBytes; - Nvfp4ColdPageTransform transform; - Nvfp4ColdPageKernelParams params; -}; - inline constexpr std::uint32_t kNvfp4ColdPageMaxBuffersPerLaunch = 256; +inline constexpr std::uint32_t kNvfp4ColdPageWideFields = 6; +inline constexpr std::uint32_t kNvfp4ColdPageIntegerFields = 5; +inline constexpr std::uint32_t kNvfp4ColdPageScaleFields = 4; -//! Configure-time launch plan for one Attention lifecycle. -struct Nvfp4ColdPagePreparedPlan -{ - std::array buffers{}; - std::uint32_t numBuffers = 0; - std::uint32_t maxHalfGroupsPerTile = 0; - std::size_t coldPageBytes = 0; - Nvfp4ColdPageRuntimeType runtimeType = Nvfp4ColdPageRuntimeType::kFloat16; -}; - -//! Validate and freeze one lifecycle's cold-page transform plan. -[[nodiscard]] Nvfp4ColdPagePreparedPlan prepareNvfp4ColdPagePlan(std::vector const& buffers, - std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType); +using Nvfp4ColdPageWideTable + = std::array, kNvfp4ColdPageMaxBuffersPerLaunch>; +using Nvfp4ColdPageIntegerTable + = std::array, kNvfp4ColdPageMaxBuffersPerLaunch>; +using Nvfp4ColdPageScaleTable + = std::array, kNvfp4ColdPageMaxBuffersPerLaunch>; //! Compress one whole KVCM Page-index batch; the launcher performs 256-Page chunking internally. -void invokeNvfp4ColdPageEncode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, - void* coldBase, cudaStream_t stream); +void invokeNvfp4ColdPageEncode(void const* pages, std::size_t numPages, std::int64_t const* wide, + std::int32_t const* integers, float const* scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType, void* coldBase, cudaStream_t stream); //! Restore one whole KVCM Page-index batch; the launcher performs 256-Page chunking internally. -void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4ColdPagePreparedPlan const& plan, - void const* coldBase, cudaStream_t stream); +void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, std::int64_t const* wide, + std::int32_t const* integers, float const* scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType, void const* coldBase, cudaStream_t stream); } // namespace kernels diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp index 2e7222d7675d..0ce52227457e 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -43,18 +43,23 @@ ResolvedHotLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::Slot return result; } -void drainAfterAlgorithmFailure(cudaStream_t stream) noexcept +void drainAfterPolicyFailure(cudaStream_t stream) noexcept { auto const status = cudaStreamSynchronize(stream); if (status != cudaSuccess) { - TLLM_LOG_ERROR("Cold-page algorithm rollback drain failed: %s", cudaGetErrorString(status)); + TLLM_LOG_ERROR("Cold-page policy rollback drain failed: %s", cudaGetErrorString(status)); std::terminate(); } } } // namespace +NativeColdPageCodec::NativeColdPageCodec(std::set layerIds) + : mLayerIds(std::move(layerIds)) +{ +} + bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept { try @@ -65,9 +70,8 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG throw std::invalid_argument("Default lossless codec rejected GPU layouts"); } - auto const& algorithmLayerIds = getLayerIds(); std::map pendingGroups; - std::vector algorithmLifecycles; + std::vector policyLifecycles; std::set consumedLayers; for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) @@ -76,31 +80,31 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG for (auto const& variant : gpuDesc.slotDesc.variants) { auto resolved = resolveLifecycle(gpuDesc, variant); - auto const algorithmLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), - [&algorithmLayerIds](auto const& layer) { return algorithmLayerIds.count(layer.first) != 0U; }); + auto const policyLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), + [this](auto const& layer) { return mLayerIds.count(layer.first) != 0U; }); LayerGroupState state; - if (algorithmLayerCount == 0U) + if (policyLayerCount == 0U) { state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); state.pageIndexLocation = losslessCodec->queryPageIndexLocation(variant.lifeCycleId); } else { - if (algorithmLayerCount != resolved.layers.size()) + if (policyLayerCount != resolved.layers.size()) { - throw std::invalid_argument("A lifecycle cannot mix algorithm-owned and fallback layers"); + throw std::invalid_argument("A lifecycle cannot mix policy-owned and fallback layers"); } for (auto const& [layerId, buffers] : resolved.layers) { static_cast(buffers); if (!consumedLayers.emplace(layerId).second) { - throw std::invalid_argument("An algorithm layer appears in multiple lifecycles"); + throw std::invalid_argument("A policy layer appears in multiple lifecycles"); } } - state.planIndex = algorithmLifecycles.size(); - algorithmLifecycles.push_back(std::move(resolved)); + state.lifecycleIndex = policyLifecycles.size(); + policyLifecycles.push_back(std::move(resolved)); } if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) @@ -109,24 +113,24 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG } } } - if (consumedLayers != algorithmLayerIds) + if (consumedLayers != mLayerIds) { - throw std::invalid_argument("An algorithm layer is absent from all GPU descriptors"); + throw std::invalid_argument("A policy layer is absent from all GPU descriptors"); } - auto const properties = configureAlgorithm(algorithmLifecycles); - if (properties.size() != algorithmLifecycles.size()) + auto const properties = configurePolicy(policyLifecycles); + if (properties.size() != policyLifecycles.size()) { - throw std::invalid_argument("Cold-page algorithm returned an unexpected lifecycle count"); + throw std::invalid_argument("Cold-page policy returned an unexpected lifecycle count"); } for (std::size_t index = 0; index < properties.size(); ++index) { auto const& lifecycle = properties[index]; if (lifecycle.coldPageBytes == 0U || lifecycle.pageIndexLocation == kv::PageIndexLocation::kBadLocation) { - throw std::invalid_argument("Cold-page algorithm returned invalid storage properties"); + throw std::invalid_argument("Cold-page policy returned invalid storage properties"); } - auto& state = pendingGroups.at(algorithmLifecycles[index].lifeCycleId); + auto& state = pendingGroups.at(policyLifecycles[index].lifeCycleId); state.coldPageBytes = lifecycle.coldPageBytes; state.pageIndexLocation = lifecycle.pageIndexLocation; } @@ -174,7 +178,7 @@ kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroup bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { - bool algorithmStarted = false; + bool policyStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -186,28 +190,28 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr { return true; } - if (!state->planIndex) + if (!state->lifecycleIndex) { return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); } - algorithmStarted = true; - encodeAlgorithm(*state->planIndex, dstBasePtr, pageIndices, numBasePages, stream); + policyStarted = true; + encodePolicy(*state->lifecycleIndex, dstBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { - if (algorithmStarted) + if (policyStarted) { - drainAfterAlgorithmFailure(stream); + drainAfterPolicyFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: %s", error.what()); return false; } catch (...) { - if (algorithmStarted) + if (policyStarted) { - drainAfterAlgorithmFailure(stream); + drainAfterPolicyFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: unknown error"); return false; @@ -217,7 +221,7 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { - bool algorithmStarted = false; + bool policyStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -229,28 +233,28 @@ bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcB { return true; } - if (!state->planIndex) + if (!state->lifecycleIndex) { return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); } - algorithmStarted = true; - decodeAlgorithm(*state->planIndex, srcBasePtr, pageIndices, numBasePages, stream); + policyStarted = true; + decodePolicy(*state->lifecycleIndex, srcBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { - if (algorithmStarted) + if (policyStarted) { - drainAfterAlgorithmFailure(stream); + drainAfterPolicyFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: %s", error.what()); return false; } catch (...) { - if (algorithmStarted) + if (policyStarted) { - drainAfterAlgorithmFailure(stream); + drainAfterPolicyFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: unknown error"); return false; diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h index fc0d8810365e..16b6494c837f 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -49,6 +49,8 @@ struct ColdPageLifecycleProperties class NativeColdPageCodec : public kv::IKvCacheColdPageCodec { public: + explicit NativeColdPageCodec(std::set layerIds); + bool configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolGroupIndex numGpuDescs) noexcept final; [[nodiscard]] std::size_t queryColdPageBytes(kv::LayerGroupId layerGroupId) const noexcept final; @@ -64,30 +66,29 @@ class NativeColdPageCodec : public kv::IKvCacheColdPageCodec std::size_t numBasePages, cudaStream_t stream) noexcept final; private: - [[nodiscard]] virtual std::set const& getLayerIds() const noexcept = 0; - - virtual std::vector configureAlgorithm( + virtual std::vector configurePolicy( std::vector const& lifecycles) = 0; //! Enqueue only on stream; this codec drains partial submissions after a throw. - virtual void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + virtual void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; - virtual void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + virtual void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; struct LayerGroupState { - std::optional planIndex; + std::optional lifecycleIndex; std::size_t coldPageBytes = 0; kv::PageIndexLocation pageIndexLocation = kv::PageIndexLocation::kBadLocation; }; [[nodiscard]] LayerGroupState const* findLayerGroup(kv::LayerGroupId layerGroupId) const noexcept; + std::set mLayerIds; std::map mLayerGroups; std::unique_ptr mLosslessCodec; }; diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index e4d27412467e..5a28fbf63fbb 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -50,18 +50,10 @@ static_assert(offsetof(kv::PageIndexPair, src) == offsetof(compression::ColdPage class PythonColdPageCodec final : public compression::NativeColdPageCodec { public: - PythonColdPageCodec(std::vector layerIds, nb::handle policy) - : mLayerIds(layerIds.begin(), layerIds.end()) + explicit PythonColdPageCodec(nb::handle policy) + : NativeColdPageCodec(readLayerIds(policy)) , mPolicy(policy.ptr()) { - if (mLayerIds.size() != layerIds.size()) - { - throw std::invalid_argument("Cold-page policy layer IDs must be unique"); - } - if (policy.is_none()) - { - throw std::invalid_argument("Cold-page policy must not be None"); - } Py_INCREF(mPolicy); } @@ -75,12 +67,22 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec } private: - [[nodiscard]] std::set const& getLayerIds() const noexcept override + static std::set readLayerIds(nb::handle policy) { - return mLayerIds; + if (policy.is_none()) + { + throw std::invalid_argument("Cold-page policy must not be None"); + } + auto const layerIds = nb::cast>(policy.attr("layer_ids")); + std::set result(layerIds.begin(), layerIds.end()); + if (result.size() != layerIds.size()) + { + throw std::invalid_argument("Cold-page policy layer IDs must be unique"); + } + return result; } - std::vector configureAlgorithm( + std::vector configurePolicy( std::vector const& lifecycles) override { nb::gil_scoped_acquire acquire; @@ -95,27 +97,27 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec } } - void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { - invoke("encode", planIndex, coldBase, pageIndices, numPages, stream); + invoke("encode", lifecycleIndex, coldBase, pageIndices, numPages, stream); } - void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { - invoke("decode", planIndex, coldBase, pageIndices, numPages, stream); + invoke("decode", lifecycleIndex, coldBase, pageIndices, numPages, stream); } template - void invoke(char const* method, std::size_t planIndex, ColdPointer coldBase, kv::PageIndexPair const* pageIndices, - std::size_t numPages, cudaStream_t stream) + void invoke(char const* method, std::size_t lifecycleIndex, ColdPointer coldBase, + kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) { // Forward the complete KVCM batch once. The method custom op owns any launch chunking. nb::gil_scoped_acquire acquire; try { - nb::borrow(mPolicy).attr(method)(planIndex, reinterpret_cast(coldBase), + nb::borrow(mPolicy).attr(method)(lifecycleIndex, reinterpret_cast(coldBase), reinterpret_cast(pageIndices), numPages, reinterpret_cast(stream)); } catch (nb::python_error const& error) @@ -124,14 +126,12 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec } } - std::set mLayerIds; PyObject* mPolicy; }; -std::unique_ptr createPythonColdPageCodec( - std::vector layerIds, nb::handle policy) +std::unique_ptr createPythonColdPageCodec(nb::handle policy) { - return std::make_unique(std::move(layerIds), policy); + return std::make_unique(policy); } } // namespace @@ -158,7 +158,7 @@ void initBindings(nb::module_& module) .def_rw("cold_page_bytes", &compression::ColdPageLifecycleProperties::coldPageBytes) .def_rw("page_index_location", &compression::ColdPageLifecycleProperties::pageIndexLocation); - module.def("create_python_cold_page_codec", &createPythonColdPageCodec, nb::arg("layer_ids"), nb::arg("policy")); + module.def("create_python_cold_page_codec", &createPythonColdPageCodec, nb::arg("policy")); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index 4215a1e0d724..939513c53563 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -38,6 +38,9 @@ target_link_libraries(th_utils PUBLIC ${SHARED_TARGET} ${TORCH_LIBRARIES} ${CUBLAS_LIB} ${CURAND_LIB}) # TODO This does not compile with internal cutlass MOE gemm +file(GLOB COLD_PAGE_METHOD_SOURCES CONFIGURE_DEPENDS + "${CMAKE_CURRENT_SOURCE_DIR}/coldPageMethods/*.cu") + add_library( th_common SHARED mlaPreprocessOp.cpp @@ -111,7 +114,7 @@ add_library( fp8PerTensorScaleMoe.cpp fp4BlockScaleMoe.cpp noAuxTcOp.cpp - nvfp4ColdPageOp.cpp + ${COLD_PAGE_METHOD_SOURCES} fusedCatFp4Op.cpp fusedCatFp8Op.cpp IndexerKCacheGatherOp.cpp diff --git a/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu b/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu new file mode 100644 index 000000000000..2f1a473d9b28 --- /dev/null +++ b/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu @@ -0,0 +1,84 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. + * All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" + +#include +#include + +namespace +{ + +using namespace tensorrt_llm::kernels; + +void checkMetadata(at::Tensor const& wide, at::Tensor const& integers, at::Tensor const& scales) +{ + TORCH_CHECK(wide.device().is_cpu() && wide.scalar_type() == at::kLong && wide.is_contiguous() + && wide.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageWideFields}), + "NVFP4 wide metadata must be a contiguous CPU int64 [256, 6] tensor"); + TORCH_CHECK(integers.device().is_cpu() && integers.scalar_type() == at::kInt && integers.is_contiguous() + && integers.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageIntegerFields}), + "NVFP4 integer metadata must be a contiguous CPU int32 [256, 5] tensor"); + TORCH_CHECK(scales.device().is_cpu() && scales.scalar_type() == at::kFloat && scales.is_contiguous() + && scales.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageScaleFields}), + "NVFP4 scale metadata must be a contiguous CPU float32 [256, 4] tensor"); +} + +template +void runNvfp4ColdPage(at::Tensor const& wide, at::Tensor const& integers, at::Tensor const& scales, + std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, std::int64_t runtimeType, + std::int64_t coldBase, std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) +{ + // This internal op consumes metadata prepared and validated by its Python policy. + checkMetadata(wide, integers, scales); + TORCH_CHECK(numBuffers > 0 && numBuffers <= kNvfp4ColdPageMaxBuffersPerLaunch, "Invalid NVFP4 buffer count"); + TORCH_CHECK(maxHalfGroupsPerTile > 0 && maxHalfGroupsPerTile <= 2048, "Invalid NVFP4 tile geometry"); + TORCH_CHECK(coldPageBytes > 0 && pageCount >= 0, "Invalid NVFP4 cold-page size or Page count"); + TORCH_CHECK(runtimeType >= 0 && runtimeType <= 2, "Invalid NVFP4 runtime type"); + TORCH_CHECK(coldBase >= 0 && pagePairs >= 0 && stream >= 0, "NVFP4 pointer arguments must be non-negative"); + + auto const* wideData = wide.const_data_ptr(); + auto const* integerData = integers.const_data_ptr(); + auto const* scaleData = scales.const_data_ptr(); + auto const type = static_cast(runtimeType); + auto const* pages = reinterpret_cast(static_cast(pagePairs)); + auto const cudaStream = reinterpret_cast(static_cast(stream)); + + if constexpr (Encode) + { + invokeNvfp4ColdPageEncode(pages, static_cast(pageCount), wideData, integerData, scaleData, + static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), + static_cast(coldPageBytes), type, + reinterpret_cast(static_cast(coldBase)), cudaStream); + } + else + { + invokeNvfp4ColdPageDecode(pages, static_cast(pageCount), wideData, integerData, scaleData, + static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), + static_cast(coldPageBytes), type, + reinterpret_cast(static_cast(coldBase)), cudaStream); + } +} + +} // namespace + +TORCH_LIBRARY_FRAGMENT(trtllm, m) +{ + m.def( + "nvfp4_cold_page_encode(Tensor wide, Tensor integers, Tensor scales, int num_buffers, " + "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int page_pairs, " + "int page_count, int stream) -> ()"); + m.def( + "nvfp4_cold_page_decode(Tensor wide, Tensor integers, Tensor scales, int num_buffers, " + "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int page_pairs, " + "int page_count, int stream) -> ()"); +} + +TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, m) +{ + m.impl("nvfp4_cold_page_encode", &runNvfp4ColdPage); + m.impl("nvfp4_cold_page_decode", &runNvfp4ColdPage); +} diff --git a/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp b/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp deleted file mode 100644 index c9e46f358bbc..000000000000 --- a/cpp/tensorrt_llm/thop/nvfp4ColdPageOp.cpp +++ /dev/null @@ -1,229 +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 "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" - -#include -#include -#include -#include -#include -#include -#include - -TRTLLM_NAMESPACE_BEGIN - -namespace torch_ext -{ -namespace -{ - -using kernels::Nvfp4ColdPageBufferPlan; -using kernels::ColdPageIndexPair; -using kernels::Nvfp4ColdPageKernelParams; -using kernels::Nvfp4ColdPagePreparedPlan; -using kernels::Nvfp4ColdPageRuntimeType; -using kernels::Nvfp4ColdPageTransform; - -enum BufferIntegerField : std::size_t -{ - kRawBase, - kRawSlotBytes, - kRawBytes, - kColdDataOffset, - kColdScaleOffset, - kColdPaddingOffset, - kColdPaddingBytes, - kTransform, - kNumKvHeads, - kTokensPerPage, - kHeadDim, - kNumBufferIntegerFields, -}; - -enum BufferScaleField : std::size_t -{ - kNvfp4ScaleOrigQuant, - kNvfp4ScaleQuantOrig, - kFp8ScaleOrigQuant, - kFp8ScaleQuantOrig, - kNumBufferScaleFields, -}; - -template -T checkedNonNegativeCast(std::int64_t value, char const* field) -{ - TORCH_CHECK(value >= 0, field, " must be non-negative"); - TORCH_CHECK(static_cast(value) <= static_cast(std::numeric_limits::max()), field, - " exceeds its native range"); - return static_cast(value); -} - -Nvfp4ColdPageRuntimeType parseRuntimeType(std::int64_t value) -{ - switch (value) - { - case 0: return Nvfp4ColdPageRuntimeType::kFloat16; - case 1: return Nvfp4ColdPageRuntimeType::kBfloat16; - case 2: return Nvfp4ColdPageRuntimeType::kFp8E4m3; - default: TORCH_CHECK(false, "Unsupported NVFP4 cold-page runtime type ", value); - } - return Nvfp4ColdPageRuntimeType::kFloat16; -} - -Nvfp4ColdPageTransform parseTransform(std::int64_t value) -{ - switch (value) - { - case 0: return Nvfp4ColdPageTransform::kNvfp4; - case 1: return Nvfp4ColdPageTransform::kLosslessCopy; - default: TORCH_CHECK(false, "Unsupported NVFP4 cold-page transform ", value); - } - return Nvfp4ColdPageTransform::kNvfp4; -} - -struct ColdPageInvocation -{ - std::uintptr_t coldBase; - void const* pagePairs; - std::size_t numPages; - cudaStream_t stream; -}; - -ColdPageInvocation parseInvocation( - std::int64_t coldBase, std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) -{ - auto const coldAddress = checkedNonNegativeCast(coldBase, "cold_base"); - auto const pairsAddress = checkedNonNegativeCast(pagePairs, "page_pairs"); - auto const numPages = checkedNonNegativeCast(pageCount, "page_count"); - auto const streamAddress = checkedNonNegativeCast(stream, "stream"); - if (numPages != 0) - { - TORCH_CHECK(coldAddress != 0, "cold_base must not be null"); - TORCH_CHECK(pairsAddress != 0, "page_pairs must not be null"); - TORCH_CHECK(pairsAddress % alignof(ColdPageIndexPair) == 0, "page_pairs is misaligned"); - } - return {coldAddress, reinterpret_cast(pairsAddress), numPages, - reinterpret_cast(streamAddress)}; -} - -} // namespace - -//! Opaque, immutable configure-time plan consumed by the runtime custom ops. -class Nvfp4ColdPageProgram : public torch::CustomClassHolder -{ -public: - explicit Nvfp4ColdPageProgram(Nvfp4ColdPagePreparedPlan plan) - : mPlan(std::move(plan)) - { - } - - [[nodiscard]] Nvfp4ColdPagePreparedPlan const& getPlan() const noexcept - { - return mPlan; - } - -private: - Nvfp4ColdPagePreparedPlan const mPlan; -}; - -c10::intrusive_ptr prepareNvfp4ColdPageProgram( - c10::List> const& bufferIntegers, c10::List> const& bufferScales, - std::int64_t coldPageBytes, std::int64_t runtimeType) -{ - TORCH_CHECK(bufferIntegers.size() == bufferScales.size(), "Each cold-page buffer needs one scale row"); - - std::vector buffers; - buffers.reserve(bufferIntegers.size()); - for (std::size_t index = 0; index < bufferIntegers.size(); ++index) - { - auto const integers = bufferIntegers.get(index); - auto const scales = bufferScales.get(index); - TORCH_CHECK(integers.size() == kNumBufferIntegerFields, "NVFP4 buffer ", index, " requires ", - kNumBufferIntegerFields, " integer fields, got ", integers.size()); - TORCH_CHECK(scales.size() == kNumBufferScaleFields, "NVFP4 buffer ", index, " requires ", kNumBufferScaleFields, - " scale fields, got ", scales.size()); - - Nvfp4ColdPageBufferPlan buffer{}; - buffer.rawBase = checkedNonNegativeCast(integers.get(kRawBase), "raw_base"); - buffer.rawSlotBytes = checkedNonNegativeCast(integers.get(kRawSlotBytes), "raw_slot_bytes"); - buffer.rawBytes = checkedNonNegativeCast(integers.get(kRawBytes), "raw_bytes"); - buffer.coldDataOffset = checkedNonNegativeCast(integers.get(kColdDataOffset), "cold_data_offset"); - buffer.coldScaleOffset - = checkedNonNegativeCast(integers.get(kColdScaleOffset), "cold_scale_offset"); - buffer.coldPaddingOffset - = checkedNonNegativeCast(integers.get(kColdPaddingOffset), "cold_padding_offset"); - buffer.coldPaddingBytes - = checkedNonNegativeCast(integers.get(kColdPaddingBytes), "cold_padding_bytes"); - buffer.transform = parseTransform(integers.get(kTransform)); - buffer.params = Nvfp4ColdPageKernelParams{ - checkedNonNegativeCast(integers.get(kNumKvHeads), "num_kv_heads"), - checkedNonNegativeCast(integers.get(kTokensPerPage), "tokens_per_page"), - checkedNonNegativeCast(integers.get(kHeadDim), "head_dim"), - static_cast(scales.get(kNvfp4ScaleOrigQuant)), static_cast(scales.get(kNvfp4ScaleQuantOrig)), - static_cast(scales.get(kFp8ScaleOrigQuant)), static_cast(scales.get(kFp8ScaleQuantOrig))}; - buffers.push_back(buffer); - } - - auto plan = kernels::prepareNvfp4ColdPagePlan( - buffers, checkedNonNegativeCast(coldPageBytes, "cold_page_bytes"), parseRuntimeType(runtimeType)); - return c10::make_intrusive(std::move(plan)); -} - -void nvfp4ColdPageEncode(c10::intrusive_ptr const& program, std::int64_t coldBase, - std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) -{ - TORCH_CHECK(program, "NVFP4 cold-page program must not be null"); - auto const batch = parseInvocation(coldBase, pagePairs, pageCount, stream); - kernels::invokeNvfp4ColdPageEncode( - batch.pagePairs, batch.numPages, program->getPlan(), reinterpret_cast(batch.coldBase), batch.stream); -} - -void nvfp4ColdPageDecode(c10::intrusive_ptr const& program, std::int64_t coldBase, - std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) -{ - TORCH_CHECK(program, "NVFP4 cold-page program must not be null"); - auto const batch = parseInvocation(coldBase, pagePairs, pageCount, stream); - kernels::invokeNvfp4ColdPageDecode(batch.pagePairs, batch.numPages, program->getPlan(), - reinterpret_cast(batch.coldBase), batch.stream); -} - -} // namespace torch_ext - -TRTLLM_NAMESPACE_END - -TORCH_LIBRARY_FRAGMENT(trtllm, m) -{ - m.class_("Nvfp4ColdPageProgram"); - m.def( - "prepare_nvfp4_cold_page_program(int[][] buffer_ints, float[][] buffer_scales, int cold_page_bytes, " - "int runtime_type) -> __torch__.torch.classes.trtllm.Nvfp4ColdPageProgram"); - m.def( - "nvfp4_cold_page_encode(__torch__.torch.classes.trtllm.Nvfp4ColdPageProgram program, int cold_base, " - "int page_pairs, int page_count, int stream) -> ()"); - m.def( - "nvfp4_cold_page_decode(__torch__.torch.classes.trtllm.Nvfp4ColdPageProgram program, int cold_base, " - "int page_pairs, int page_count, int stream) -> ()"); -} - -TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, m) -{ - m.impl("prepare_nvfp4_cold_page_program", &tensorrt_llm::torch_ext::prepareNvfp4ColdPageProgram); - m.impl("nvfp4_cold_page_encode", &tensorrt_llm::torch_ext::nvfp4ColdPageEncode); - m.impl("nvfp4_cold_page_decode", &tensorrt_llm::torch_ext::nvfp4ColdPageDecode); -} diff --git a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp index 88db07644510..030d8030dd9e 100644 --- a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp +++ b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp @@ -41,12 +41,98 @@ namespace using tensorrt_llm::batch_manager::kv_cache_manager_v2::HostMem; using tensorrt_llm::batch_manager::kv_cache_manager_v2::MemAddress; -using tensorrt_llm::kernels::Nvfp4ColdPageBufferPlan; using tensorrt_llm::kernels::ColdPageIndexPair; -using tensorrt_llm::kernels::Nvfp4ColdPageKernelParams; -using tensorrt_llm::kernels::Nvfp4ColdPagePreparedPlan; +using tensorrt_llm::kernels::Nvfp4ColdPageIntegerTable; using tensorrt_llm::kernels::Nvfp4ColdPageRuntimeType; -using tensorrt_llm::kernels::Nvfp4ColdPageTransform; +using tensorrt_llm::kernels::Nvfp4ColdPageScaleTable; +using tensorrt_llm::kernels::Nvfp4ColdPageWideTable; + +struct Nvfp4ColdPageKernelParams +{ + std::int32_t numKvHeads; + std::int32_t tokensPerPage; + std::int32_t headDim; + float nvfp4ScaleOrigQuant; + float nvfp4ScaleQuantOrig; + float fp8ScaleOrigQuant; + float fp8ScaleQuantOrig; +}; + +enum class Nvfp4ColdPageTransform : std::int32_t +{ + kNvfp4 = 0, + kLosslessCopy = 1, +}; + +struct Nvfp4ColdPageTestBuffer +{ + std::uintptr_t rawBase; + std::size_t rawSlotBytes; + std::size_t rawBytes; + std::size_t coldDataOffset; + std::size_t coldScaleOffset; + std::size_t coldPaddingOffset; + std::uint32_t coldPaddingBytes; + Nvfp4ColdPageTransform transform; + Nvfp4ColdPageKernelParams params; +}; + +struct Nvfp4ColdPageTestMetadata +{ + Nvfp4ColdPageWideTable wide{}; + Nvfp4ColdPageIntegerTable integers{}; + Nvfp4ColdPageScaleTable scales{}; + std::uint32_t numBuffers{}; + std::uint32_t maxHalfGroupsPerTile{}; + std::size_t coldPageBytes{}; + Nvfp4ColdPageRuntimeType runtimeType{}; +}; + +Nvfp4ColdPageTestMetadata makeNvfp4ColdPageTestMetadata(std::vector const& buffers, + std::size_t coldPageBytes, Nvfp4ColdPageRuntimeType runtimeType) +{ + Nvfp4ColdPageTestMetadata metadata; + metadata.numBuffers = static_cast(buffers.size()); + metadata.coldPageBytes = coldPageBytes; + metadata.runtimeType = runtimeType; + for (std::size_t index = 0; index < buffers.size(); ++index) + { + auto const& buffer = buffers[index]; + metadata.wide[index] + = {static_cast(buffer.rawBase), static_cast(buffer.rawSlotBytes), + static_cast(buffer.rawBytes), static_cast(buffer.coldDataOffset), + static_cast(buffer.coldScaleOffset), static_cast(buffer.coldPaddingOffset)}; + metadata.integers[index] + = {static_cast(buffer.coldPaddingBytes), static_cast(buffer.transform), + buffer.params.numKvHeads, buffer.params.tokensPerPage, buffer.params.headDim}; + metadata.scales[index] = {buffer.params.nvfp4ScaleOrigQuant, buffer.params.nvfp4ScaleQuantOrig, + buffer.params.fp8ScaleOrigQuant, buffer.params.fp8ScaleQuantOrig}; + if (buffer.transform == Nvfp4ColdPageTransform::kNvfp4) + { + auto const halfGroups = static_cast(buffer.params.numKvHeads) + * static_cast(buffer.params.tokensPerPage) + * (static_cast(buffer.params.headDim) / 8U); + metadata.maxHalfGroupsPerTile = std::max(metadata.maxHalfGroupsPerTile, std::min(halfGroups, 2048U)); + } + } + return metadata; +} + +void invokeNvfp4ColdPageEncode(void const* pages, std::size_t numPages, Nvfp4ColdPageTestMetadata const& metadata, + void* coldBase, cudaStream_t stream) +{ + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(pages, numPages, metadata.wide.front().data(), + metadata.integers.front().data(), metadata.scales.front().data(), metadata.numBuffers, + metadata.maxHalfGroupsPerTile, metadata.coldPageBytes, metadata.runtimeType, coldBase, stream); +} + +void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, Nvfp4ColdPageTestMetadata const& metadata, + void const* coldBase, cudaStream_t stream) +{ + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(pages, numPages, metadata.wide.front().data(), + metadata.integers.front().data(), metadata.scales.front().data(), metadata.numBuffers, + metadata.maxHalfGroupsPerTile, metadata.coldPageBytes, metadata.runtimeType, coldBase, stream); +} constexpr std::size_t kGuardBytes = 64; constexpr std::uint8_t kCanary = 0xA5; @@ -516,14 +602,13 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG std::size_t const scale = scaleBytes(geometry); std::size_t const payloadBytes = 2U * (packed + scale); std::uint32_t const paddingBytes = static_cast(compactSlotBytes - payloadBytes); - std::vector const inputBuffers{ + std::vector const inputBuffers{ {reinterpret_cast(rawInputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawInputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, payloadBytes, paddingBytes, Nvfp4ColdPageTransform::kNvfp4, params[1]}}; - auto const inputPlan - = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan(inputBuffers, compactSlotBytes, runtimeType(kind)); + auto const inputMetadata = makeNvfp4ColdPageTestMetadata(inputBuffers, compactSlotBytes, runtimeType(kind)); std::vector> references(numPages); for (std::size_t page = 0; page < numPages; ++page) { @@ -558,8 +643,7 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG } }; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - offloadTasks.data(), offloadTasks.size(), inputPlan, compactBase, stream); + invokeNvfp4ColdPageEncode(offloadTasks.data(), offloadTasks.size(), inputMetadata, compactBase, stream); if (synchronizeBetweenDirections) { @@ -580,8 +664,7 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG { std::memset(compactBase + (2U * page + 1U) * compactSlotBytes, 0x5A, compactSlotBytes); } - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - offloadTasks.data(), offloadTasks.size(), inputPlan, compactBase, stream); + invokeNvfp4ColdPageEncode(offloadTasks.data(), offloadTasks.size(), inputMetadata, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); EXPECT_EQ(compactPages.payload(), firstSerialization); verifyCompressedPages(); @@ -596,15 +679,13 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG std::size_t const coldSlot = rawSlot + 1U; onboardTasks.push_back({static_cast(rawSlot), static_cast(coldSlot)}); } - std::vector const outputBuffers{ + std::vector const outputBuffers{ {reinterpret_cast(rawOutputK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawOutputV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, payloadBytes, paddingBytes, Nvfp4ColdPageTransform::kNvfp4, params[1]}}; - auto const outputPlan - = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan(outputBuffers, compactSlotBytes, runtimeType(kind)); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( - onboardTasks.data(), onboardTasks.size(), outputPlan, compactBase, stream); + auto const outputMetadata = makeNvfp4ColdPageTestMetadata(outputBuffers, compactSlotBytes, runtimeType(kind)); + invokeNvfp4ColdPageDecode(onboardTasks.data(), onboardTasks.size(), outputMetadata, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); if (!synchronizeBetweenDirections) @@ -624,10 +705,8 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG references[page][role] = compressReference(restored, kind, params[role], geometry); } } - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - offloadTasks.data(), offloadTasks.size(), outputPlan, compactBase, stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( - onboardTasks.data(), onboardTasks.size(), inputPlan, compactBase, stream); + invokeNvfp4ColdPageEncode(offloadTasks.data(), offloadTasks.size(), outputMetadata, compactBase, stream); + invokeNvfp4ColdPageDecode(onboardTasks.data(), onboardTasks.size(), inputMetadata, compactBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); verifyCompressedPages(); } @@ -765,22 +844,20 @@ void runPartialPageTailIsolation(RawKind kind) std::size_t const payloadBytes = 2U * (packed + scale); auto const makeBuffers = [&](DeviceRegion const& rawK, DeviceRegion const& rawV) { - return std::vector{ + return std::vector{ {reinterpret_cast(rawK.data()), rawSlotBytes, rawSlotBytes, 0U, 2U * packed, 0U, 0U, Nvfp4ColdPageTransform::kNvfp4, params[0]}, {reinterpret_cast(rawV.data()), rawSlotBytes, rawSlotBytes, packed, 2U * packed + scale, payloadBytes, static_cast(compactSlotBytes - payloadBytes), Nvfp4ColdPageTransform::kNvfp4, params[1]}}; }; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - makeBuffers(rawInputK, rawInputV), compactSlotBytes, runtimeType(kind)); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - makeBuffers(rawOutputK, rawOutputV), compactSlotBytes, runtimeType(kind)); + auto const inputMetadata + = makeNvfp4ColdPageTestMetadata(makeBuffers(rawInputK, rawInputV), compactSlotBytes, runtimeType(kind)); + auto const outputMetadata + = makeNvfp4ColdPageTestMetadata(makeBuffers(rawOutputK, rawOutputV), compactSlotBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - offloadTasks.data(), offloadTasks.size(), inputPlan, compactPages.data(), stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( - onboardTasks.data(), onboardTasks.size(), outputPlan, compactPages.data(), stream); + invokeNvfp4ColdPageEncode(offloadTasks.data(), offloadTasks.size(), inputMetadata, compactPages.data(), stream); + invokeNvfp4ColdPageDecode(onboardTasks.data(), onboardTasks.size(), outputMetadata, compactPages.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); for (std::size_t pair = 0; pair < kValidTokenCounts.size(); ++pair) @@ -927,23 +1004,21 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) auto const makePlans = [&](DeviceRegion const& mla, DeviceRegion const& side) { - return std::vector{ + return std::vector{ {reinterpret_cast(mla.data()), mlaRawBytes, mlaRawBytes, 0U, mlaPackedBytes, mlaPayloadBytes, static_cast(gapBeforeSide), Nvfp4ColdPageTransform::kNvfp4, params}, {reinterpret_cast(side.data()), sideSlotBytes, sideRawBytes, sideColdOffset, 0U, sideColdEnd, static_cast(coldPageBytes - sideColdEnd), Nvfp4ColdPageTransform::kLosslessCopy, {}}}; }; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - makePlans(mlaInput, sideInput), coldPageBytes, runtimeType(kind)); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - makePlans(mlaOutput, sideOutput), coldPageBytes, runtimeType(kind)); + auto const inputMetadata + = makeNvfp4ColdPageTestMetadata(makePlans(mlaInput, sideInput), coldPageBytes, runtimeType(kind)); + auto const outputMetadata + = makeNvfp4ColdPageTestMetadata(makePlans(mlaOutput, sideOutput), coldPageBytes, runtimeType(kind)); CudaStream stream; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - offloadTasks.data(), offloadTasks.size(), inputPlan, coldBase, stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( - onboardTasks.data(), onboardTasks.size(), outputPlan, coldBase, stream); + invokeNvfp4ColdPageEncode(offloadTasks.data(), offloadTasks.size(), inputMetadata, coldBase, stream); + invokeNvfp4ColdPageDecode(onboardTasks.data(), onboardTasks.size(), outputMetadata, coldBase, stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const cold = coldStorage.payload(); @@ -1026,10 +1101,10 @@ TEST(Nvfp4ColdPageWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc std::array, 2>, numLayers> rawHost; std::array, numLayers> references; - std::vector inputPlans; - std::vector outputPlans; - inputPlans.reserve(2U * numLayers); - outputPlans.reserve(2U * numLayers); + std::vector inputMetadatas; + std::vector outputMetadatas; + inputMetadatas.reserve(2U * numLayers); + outputMetadatas.reserve(2U * numLayers); for (std::size_t layer = 0; layer < numLayers; ++layer) { rawInputK[layer] = std::make_unique(rawSlotBytes); @@ -1055,19 +1130,19 @@ TEST(Nvfp4ColdPageWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc static_cast(layerRecordStride - layerRecordBytes), Nvfp4ColdPageTransform::kNvfp4, params[layer][1]}); }; - appendPlans(inputPlans, *rawInputK[layer], *rawInputV[layer]); - appendPlans(outputPlans, *rawOutputK[layer], *rawOutputV[layer]); + appendPlans(inputMetadatas, *rawInputK[layer], *rawInputV[layer]); + appendPlans(outputMetadatas, *rawOutputK[layer], *rawOutputV[layer]); } MappedHostRegion compactPage(coldPageBytes); CudaStream stream; - auto const inputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - inputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); - auto const outputPlan = tensorrt_llm::kernels::prepareNvfp4ColdPagePlan( - outputPlans, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); + auto const inputMetadata + = makeNvfp4ColdPageTestMetadata(inputMetadatas, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); + auto const outputMetadata + = makeNvfp4ColdPageTestMetadata(outputMetadatas, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); ColdPageIndexPair const page{0, 0}; - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(&page, 1U, inputPlan, compactPage.data(), stream); - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode(&page, 1U, outputPlan, compactPage.data(), stream); + invokeNvfp4ColdPageEncode(&page, 1U, inputMetadata, compactPage.data(), stream); + invokeNvfp4ColdPageDecode(&page, 1U, outputMetadata, compactPage.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); auto const compact = compactPage.payload(); @@ -1106,7 +1181,7 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector, numLayers> rawK; std::array, numLayers> rawV; - std::vector buffers; + std::vector buffers; buffers.reserve(2U * numLayers); for (std::size_t layer = 0; layer < numLayers; ++layer) { @@ -1139,8 +1214,7 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector 0 && params.tokensPerPage > 0 && params.headDim > 0) - { - std::uint64_t const elements = static_cast(params.numKvHeads) - * static_cast(params.tokensPerPage) * static_cast(params.headDim); - std::uint64_t const candidateRawBytes = elements * (type == Nvfp4ColdPageRuntimeType::kFp8E4m3 ? 1U : 2U); - if (candidateRawBytes > 0U && candidateRawBytes <= rawSlotBytes) - { - activeRawBytes = static_cast(candidateRawBytes); - dataBytes = static_cast(elements / 2U); - scales = static_cast(elements / 16U); - } - } - Nvfp4ColdPageBufferPlan const buffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, - activeRawBytes, 0U, dataBytes, dataBytes + scales, - static_cast(coldPageBytes - dataBytes - scales), Nvfp4ColdPageTransform::kNvfp4, params}; - static_cast(tensorrt_llm::kernels::prepareNvfp4ColdPagePlan({buffer}, coldPageBytes, type)); - }; - auto const expectInvalid - = [&](char const* name, auto const& mutate, Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) - { - SCOPED_TRACE(name); - auto params = makeParams(); - mutate(params); - EXPECT_ANY_THROW(prepare(params, type)); - }; - - auto valid = makeParams(); - valid.tokensPerPage = 6; - EXPECT_NO_THROW(prepare(valid)); - EXPECT_NO_THROW(prepare(makeParams(PageGeometry{1, 1, 16}))); - - expectInvalid("zero heads", [](auto& params) { params.numKvHeads = 0; }); - expectInvalid("zero tokens", [](auto& params) { params.tokensPerPage = 0; }); - expectInvalid("zero head dimension", [](auto& params) { params.headDim = 0; }); - expectInvalid("unaligned head dimension", [](auto& params) { params.headDim = 24; }); - expectInvalid( - "element count overflow", [](auto& params) { params.numKvHeads = std::numeric_limits::max(); }); - expectInvalid("zero NVFP4 quant scale", [](auto& params) { params.nvfp4ScaleOrigQuant = 0.0F; }); - expectInvalid("negative NVFP4 dequant scale", [](auto& params) { params.nvfp4ScaleQuantOrig = -1.0F; }); - expectInvalid("NaN NVFP4 quant scale", - [](auto& params) { params.nvfp4ScaleOrigQuant = std::numeric_limits::quiet_NaN(); }); - expectInvalid("infinite NVFP4 dequant scale", - [](auto& params) { params.nvfp4ScaleQuantOrig = std::numeric_limits::infinity(); }); - expectInvalid( - "zero FP8 quant scale", [](auto& params) { params.fp8ScaleOrigQuant = 0.0F; }, - Nvfp4ColdPageRuntimeType::kFp8E4m3); - expectInvalid( - "infinite FP8 dequant scale", - [](auto& params) { params.fp8ScaleQuantOrig = std::numeric_limits::infinity(); }, - Nvfp4ColdPageRuntimeType::kFp8E4m3); -} - -TEST(Nvfp4ColdPageValidationTest, RejectsInvalidLaunchDescriptors) -{ - ASSERT_EQ(cudaSetDevice(0), cudaSuccess); - if (!tensorrt_llm::common::isSM100Family()) - { - GTEST_SKIP() << "NVFP4 cold-page kernels require an SM100-family GPU"; - } - - std::size_t const rawSlotBytes = rawBytes(RawKind::kFloat16, kDefaultGeometry); - LayerBuffers buffers(rawSlotBytes); - ColdPageIndexPair const validPage{0, 0}; - std::size_t const coldPageBytes = 2U * (packedBytes(kDefaultGeometry) + scaleBytes(kDefaultGeometry)); - std::size_t const packed = packedBytes(kDefaultGeometry); - std::size_t const scale = scaleBytes(kDefaultGeometry); - Nvfp4ColdPageBufferPlan const validBuffer{reinterpret_cast(buffers.rawK.data()), rawSlotBytes, - rawSlotBytes, 0U, packed, packed + scale, static_cast(coldPageBytes - packed - scale), - Nvfp4ColdPageTransform::kNvfp4, makeParams()}; - auto const prepare = [&](Nvfp4ColdPageBufferPlan const& buffer, std::size_t pageBytes, - Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) - { return tensorrt_llm::kernels::prepareNvfp4ColdPagePlan({buffer}, pageBytes, type); }; - auto const expectInvalid - = [&](char const* name, auto const& mutate, Nvfp4ColdPageRuntimeType type = Nvfp4ColdPageRuntimeType::kFloat16) - { - SCOPED_TRACE(name); - auto buffer = validBuffer; - mutate(buffer); - EXPECT_ANY_THROW(static_cast(prepare(buffer, coldPageBytes, type))); - }; - auto const validPlan = prepare(validBuffer, coldPageBytes); - - expectInvalid("unaligned raw base", [](auto& buffer) { buffer.rawBase += 1U; }); - EXPECT_ANY_THROW(tensorrt_llm::kernels::invokeNvfp4ColdPageEncode(&validPage, 1U, validPlan, nullptr, nullptr)); - expectInvalid("unaligned raw stride", [](auto& buffer) { buffer.rawSlotBytes += alignof(uint4) / 2U; }); - expectInvalid("raw bytes exceed stride", [](auto& buffer) { buffer.rawBytes = buffer.rawSlotBytes + 1U; }); - expectInvalid("raw bytes mismatch geometry", [](auto& buffer) { buffer.rawBytes -= alignof(uint4); }); - expectInvalid("cold data interval exceeds page", [&](auto& buffer) { buffer.coldDataOffset = coldPageBytes; }); - expectInvalid("cold intervals overlap", [](auto& buffer) { buffer.coldScaleOffset = buffer.coldDataOffset; }); - EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes + alignof(uint4) / 2U))); - expectInvalid( - "unaligned FP8 raw base", [](auto& buffer) { buffer.rawBase += 1U; }, Nvfp4ColdPageRuntimeType::kFp8E4m3); - - auto const unsupportedType = static_cast(255); - EXPECT_ANY_THROW(static_cast(prepare(validBuffer, coldPageBytes, unsupportedType))); + invokeNvfp4ColdPageEncode(nullptr, 0U, Nvfp4ColdPageTestMetadata{}, nullptr, nullptr); + invokeNvfp4ColdPageDecode(nullptr, 0U, Nvfp4ColdPageTestMetadata{}, nullptr, nullptr); } } // namespace diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp index 8ac907f0a094..bdb27f11825e 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -72,16 +72,11 @@ class RecordingCodec final : public NativeColdPageCodec { public: explicit RecordingCodec(std::set layerIds) - : mLayerIds(std::move(layerIds)) + : NativeColdPageCodec(std::move(layerIds)) { } - [[nodiscard]] std::set const& getLayerIds() const noexcept override - { - return mLayerIds; - } - - std::vector configureAlgorithm( + std::vector configurePolicy( std::vector const& lifecycles) override { resolved = lifecycles; @@ -93,7 +88,7 @@ class RecordingCodec final : public NativeColdPageCodec lifecycles.size(), ColdPageLifecycleProperties{777U, kv::PageIndexLocation::kHost}); } - void encodeAlgorithm(std::size_t planIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { if (failBatches) @@ -102,13 +97,13 @@ class RecordingCodec final : public NativeColdPageCodec throw std::runtime_error("requested batch failure"); } ++encodeCalls; - lastPlanIndex = planIndex; + lastLifecycleIndex = lifecycleIndex; lastColdBase = coldBase; lastIndices.assign(pageIndices, pageIndices + numPages); lastStream = stream; } - void decodeAlgorithm(std::size_t planIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { if (failBatches) @@ -117,7 +112,7 @@ class RecordingCodec final : public NativeColdPageCodec throw std::runtime_error("requested batch failure"); } ++decodeCalls; - lastPlanIndex = planIndex; + lastLifecycleIndex = lifecycleIndex; lastColdBase = coldBase; lastIndices.assign(pageIndices, pageIndices + numPages); lastStream = stream; @@ -128,7 +123,7 @@ class RecordingCodec final : public NativeColdPageCodec std::atomic_bool* failureMarker = nullptr; int encodeCalls = 0; int decodeCalls = 0; - std::size_t lastPlanIndex = 0; + std::size_t lastLifecycleIndex = 0; void const* lastColdBase = nullptr; cudaStream_t lastStream{}; std::vector lastIndices; @@ -145,8 +140,6 @@ class RecordingCodec final : public NativeColdPageCodec throw std::runtime_error("failed to enqueue the requested batch failure marker"); } } - - std::set mLayerIds; }; TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsWholeBatchOnce) @@ -173,7 +166,7 @@ TEST(NativeColdPageCodecTest, ResolvesKvcManagerLayoutAndForwardsWholeBatchOnce) ASSERT_TRUE( codec.encode(kv::LayerGroupId{3}, reinterpret_cast(kColdBase), indices.data(), indices.size(), stream)); EXPECT_EQ(codec.encodeCalls, 1); - EXPECT_EQ(codec.lastPlanIndex, 0U); + EXPECT_EQ(codec.lastLifecycleIndex, 0U); EXPECT_EQ(codec.lastIndices.size(), 4096U); EXPECT_EQ(codec.lastIndices.back().src, 4095); EXPECT_EQ(codec.lastColdBase, reinterpret_cast(kColdBase)); @@ -214,7 +207,7 @@ TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) } } -TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureFailuresAndInvalidBatches) +TEST(NativeColdPageCodecTest, CatchesPolicyConfigureFailuresAndInvalidBatches) { RecordingCodec codec{{0}}; codec.failConfigure = true; @@ -231,7 +224,7 @@ TEST(NativeColdPageCodecTest, CatchesAlgorithmConfigureFailuresAndInvalidBatches EXPECT_EQ(validCodec.decodeCalls, 0); } -TEST(NativeColdPageCodecTest, AlgorithmFailureUsesTheSuppliedCudaStreamForRollback) +TEST(NativeColdPageCodecTest, PolicyFailureUsesTheSuppliedCudaStreamForRollback) { int deviceCount = 0; if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py index 028c28b04813..afe159e96a40 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py @@ -31,6 +31,12 @@ _COLD_PAGE_ALIGNMENT = 16 _ELEMENTS_PER_BYTE = 2 _ELEMENTS_PER_SCALE = 16 +_ELEMENTS_PER_HALF_GROUP = 8 +_MAX_HALF_GROUPS_PER_TILE = 2048 +_MAX_BUFFERS_PER_LAUNCH = 256 +_WIDE_FIELDS = 6 +_INTEGER_FIELDS = 5 +_SCALE_FIELDS = 4 _NVFP4_TRANSFORM = 0 _LOSSLESS_TRANSFORM = 1 @@ -46,23 +52,31 @@ class _Nvfp4Scales: @dataclass(frozen=True) class _Nvfp4BufferLayout: role: str - data_offset: int - scale_offset: int = 0 scales: _Nvfp4Scales | None = None @dataclass(frozen=True) class _Nvfp4LayerLayout: layer_id: int - runtime_type: int num_kv_heads: int tokens_per_page: int head_dim: int - cold_page_bytes: int - padding_offset: int buffers: tuple[_Nvfp4BufferLayout, ...] +@dataclass(frozen=True) +class _Nvfp4ColdPageMetadata: + """Python-owned launch metadata for one KVCM lifecycle.""" + + wide: torch.Tensor + integers: torch.Tensor + scales: torch.Tensor + num_buffers: int + max_half_groups_per_tile: int + cold_page_bytes: int + runtime_type: int + + def _load_modelopt_nvfp4_scales( checkpoint_path: str | None, ) -> dict[int, LayerScales]: @@ -118,10 +132,15 @@ def _load_modelopt_nvfp4_scales( if not k_values or not v_values: raise ValueError(f"ModelOpt KV scales for layer {layer_id} must contain both K and V") quant_orig = (max(k_values), max(v_values)) - result[layer_id] = ( - (1.0 / quant_orig[0], 1.0 / quant_orig[1]), - quant_orig, - ) + orig_quant = (1.0 / quant_orig[0], 1.0 / quant_orig[1]) + stored_scales = torch.tensor( + (*orig_quant, *quant_orig), dtype=torch.float32, device="cpu" + ).tolist() + if any(not math.isfinite(value) or value <= 0.0 for value in stored_scales): + raise ValueError( + f"ModelOpt KV scales for layer {layer_id} are not representable as float32" + ) + result[layer_id] = (orig_quant, quant_orig) return result @@ -129,57 +148,102 @@ def _align_up(value: int, alignment: int = _COLD_PAGE_ALIGNMENT) -> int: return (value + alignment - 1) // alignment * alignment -def _buffer_bytes(buffer: object, tokens_per_page: int) -> int: - buffer_tokens = buffer.tokens_per_block_override or tokens_per_page - return int(buffer.size) * (tokens_per_page // buffer_tokens) - - class Nvfp4ColdPagePolicy(ColdPageCodecPolicy): - """Own NVFP4 lifecycle programs and dispatch one operation per codec batch.""" + """Resolve NVFP4 metadata and submit one CUDA op per codec batch.""" - def __init__(self, layer_layouts: Sequence[_Nvfp4LayerLayout]) -> None: + def __init__(self, layer_layouts: Sequence[_Nvfp4LayerLayout], runtime_type: int) -> None: self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} - self._programs: list[object] = [] + self._runtime_type = runtime_type + self._lifecycle_metadata: list[_Nvfp4ColdPageMetadata] = [] @property def layer_ids(self) -> tuple[int, ...]: return tuple(sorted(self._layer_layouts)) def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: - """Resolve Python layouts against hot buffers and prepare method programs.""" + """Resolve hot buffers into immutable Python-owned launch metadata.""" from tensorrt_llm.bindings.internal import kv_cache_compression as native - programs = [] + lifecycle_metadata = [] properties = [] for lifecycle in lifecycles: - metadata: list[list[int]] = [] - scales: list[list[float]] = [] + wide_rows: list[list[int]] = [] + integer_rows: list[list[int]] = [] + scale_rows: list[list[float]] = [] cold_page_bytes = 0 - runtime_type = None + max_half_groups_per_tile = 0 for layer_id, hot_buffers in lifecycle.layers.items(): layout = self._layer_layouts[int(layer_id)] - if runtime_type is not None and runtime_type != layout.runtime_type: - raise ValueError("One cold-page lifecycle must use one runtime dtype") - runtime_type = layout.runtime_type - - for index, buffer in enumerate(layout.buffers): + expected_roles = {buffer.role for buffer in layout.buffers} + if set(hot_buffers) != expected_roles: + raise ValueError( + f"Cold-page layer {layer_id} roles do not match its KVCM layout" + ) + compressed = [buffer for buffer in layout.buffers if buffer.scales] + packed_bytes = ( + layout.num_kv_heads * layout.tokens_per_page * layout.head_dim + ) // _ELEMENTS_PER_BYTE + scale_bytes = ( + layout.num_kv_heads * layout.tokens_per_page * layout.head_dim + ) // _ELEMENTS_PER_SCALE + layer_start = cold_page_bytes + scale_start = layer_start + len(compressed) * packed_bytes + cursor = scale_start + len(compressed) * scale_bytes + + compressed_index = 0 + for buffer in layout.buffers: hot = hot_buffers[buffer.role] - padding_offset = 0 - padding_bytes = 0 - if index + 1 == len(layout.buffers): - padding_offset = cold_page_bytes + layout.padding_offset - padding_bytes = layout.cold_page_bytes - layout.padding_offset - metadata.append( + raw_base = int(hot.raw_base) + raw_slot_bytes = int(hot.raw_slot_bytes) + raw_bytes = int(hot.raw_bytes) + if raw_base <= 0 or raw_bytes <= 0 or raw_bytes > raw_slot_bytes: + raise ValueError("Cold-page hot buffer has invalid address or size") + + if buffer.scales: + data_offset = layer_start + compressed_index * packed_bytes + scale_offset = scale_start + compressed_index * scale_bytes + compressed_index += 1 + expected_raw_bytes = ( + layout.num_kv_heads + * layout.tokens_per_page + * layout.head_dim + * (1 if self._runtime_type == 2 else 2) + ) + if raw_bytes != expected_raw_bytes: + raise ValueError("Hot buffer size does not match NVFP4 geometry") + if raw_base % 16 or raw_slot_bytes % 16: + raise ValueError( + "NVFP4 hot address and Slot stride must be 16-byte aligned" + ) + half_groups = ( + expected_raw_bytes + // (1 if self._runtime_type == 2 else 2) + // _ELEMENTS_PER_HALF_GROUP + ) + max_half_groups_per_tile = max( + max_half_groups_per_tile, + min(half_groups, _MAX_HALF_GROUPS_PER_TILE), + ) + else: + data_offset = cursor + scale_offset = 0 + cursor += raw_bytes + + wide_rows.append( + [ + raw_base, + raw_slot_bytes, + raw_bytes, + data_offset, + scale_offset, + 0, + ] + ) + integer_rows.append( [ - int(hot.raw_base), - int(hot.raw_slot_bytes), - int(hot.raw_bytes), - cold_page_bytes + buffer.data_offset, - cold_page_bytes + buffer.scale_offset if buffer.scales else 0, - padding_offset, - padding_bytes, + 0, _NVFP4_TRANSFORM if buffer.scales else _LOSSLESS_TRANSFORM, layout.num_kv_heads if buffer.scales else 0, layout.tokens_per_page if buffer.scales else 0, @@ -187,7 +251,7 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: ] ) buffer_scales = buffer.scales or _Nvfp4Scales(1.0, 1.0) - scales.append( + scale_rows.append( [ buffer_scales.nvfp4_orig_quant, buffer_scales.nvfp4_quant_orig, @@ -195,13 +259,39 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: buffer_scales.fp8_quant_orig, ] ) - cold_page_bytes += layout.cold_page_bytes + layer_end = _align_up(cursor) + wide_rows[-1][5] = cursor + integer_rows[-1][0] = layer_end - cursor + cold_page_bytes = layer_end - if runtime_type is None: - raise ValueError("NVFP4 received an empty cold-page lifecycle") - programs.append( - torch.ops.trtllm.prepare_nvfp4_cold_page_program( - metadata, scales, cold_page_bytes, runtime_type + num_buffers = len(wide_rows) + if not 0 < num_buffers <= _MAX_BUFFERS_PER_LAUNCH: + raise ValueError( + f"NVFP4 cold-page lifecycle has {num_buffers} buffers; " + f"the maximum is {_MAX_BUFFERS_PER_LAUNCH}" + ) + padding = _MAX_BUFFERS_PER_LAUNCH - num_buffers + lifecycle_metadata.append( + _Nvfp4ColdPageMetadata( + wide=torch.tensor( + wide_rows + [[0] * _WIDE_FIELDS for _ in range(padding)], + dtype=torch.int64, + device="cpu", + ), + integers=torch.tensor( + integer_rows + [[0] * _INTEGER_FIELDS for _ in range(padding)], + dtype=torch.int32, + device="cpu", + ), + scales=torch.tensor( + scale_rows + [[0.0] * _SCALE_FIELDS for _ in range(padding)], + dtype=torch.float32, + device="cpu", + ), + num_buffers=num_buffers, + max_half_groups_per_tile=max_half_groups_per_tile, + cold_page_bytes=cold_page_bytes, + runtime_type=self._runtime_type, ) ) lifecycle_properties = native.ColdPageLifecycleProperties() @@ -209,31 +299,53 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: lifecycle_properties.page_index_location = native.ColdPageIndexLocation.HOST properties.append(lifecycle_properties) - self._programs = programs + self._lifecycle_metadata = lifecycle_metadata return properties def encode( self, - program_index: int, + lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, stream: int, ) -> None: + metadata = self._lifecycle_metadata[lifecycle_index] torch.ops.trtllm.nvfp4_cold_page_encode( - self._programs[program_index], cold_base, page_indices, num_pages, stream + metadata.wide, + metadata.integers, + metadata.scales, + metadata.num_buffers, + metadata.max_half_groups_per_tile, + metadata.cold_page_bytes, + metadata.runtime_type, + cold_base, + page_indices, + num_pages, + stream, ) def decode( self, - program_index: int, + lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, stream: int, ) -> None: + metadata = self._lifecycle_metadata[lifecycle_index] torch.ops.trtllm.nvfp4_cold_page_decode( - self._programs[program_index], cold_base, page_indices, num_pages, stream + metadata.wide, + metadata.integers, + metadata.scales, + metadata.num_buffers, + metadata.max_half_groups_per_tile, + metadata.cold_page_bytes, + metadata.runtime_type, + cold_base, + page_indices, + num_pages, + stream, ) @@ -294,24 +406,9 @@ def create_cold_page_codec( raise ValueError( f"NVFP4 cold pages require head_dim divisible by 16, got {head_dim}" ) - elements = num_kv_heads * tokens_per_page * head_dim - packed_bytes = elements // _ELEMENTS_PER_BYTE - scale_bytes = elements // _ELEMENTS_PER_SCALE - data_offsets = { - role: index * packed_bytes for index, role in enumerate(compressed_roles) - } - scale_base = len(compressed_roles) * packed_bytes - scale_offsets = { - role: scale_base + index * scale_bytes - for index, role in enumerate(compressed_roles) - } - cursor = scale_base + len(compressed_roles) * scale_bytes - buffer_layouts = [ _Nvfp4BufferLayout( role=role, - data_offset=data_offsets[role], - scale_offset=scale_offsets[role], scales=_Nvfp4Scales(orig_quant[index], quant_orig[index]), ) for index, role in enumerate(compressed_roles) @@ -319,21 +416,17 @@ def create_cold_page_codec( for buffer in layer.buffers: role = str(buffer.role) if role not in compressed_roles: - buffer_layouts.append(_Nvfp4BufferLayout(role=role, data_offset=cursor)) - cursor += _buffer_bytes(buffer, tokens_per_page) + buffer_layouts.append(_Nvfp4BufferLayout(role=role)) layer_layouts.append( _Nvfp4LayerLayout( layer_id=layer_id, - runtime_type=runtime_type, num_kv_heads=num_kv_heads, tokens_per_page=tokens_per_page, head_dim=head_dim, - cold_page_bytes=_align_up(cursor), - padding_offset=cursor, buffers=tuple(buffer_layouts), ) ) - policy = Nvfp4ColdPagePolicy(layer_layouts) - return native.create_python_cold_page_codec(policy.layer_ids, policy) + policy = Nvfp4ColdPagePolicy(layer_layouts, runtime_type or 0) + return native.create_python_cold_page_codec(policy) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index d08c2e93936a..f775e3d44896 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -38,12 +38,12 @@ def layer_ids(self) -> tuple[int, ...]: @abstractmethod def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: - """Prepare one immutable method program per owned lifecycle.""" + """Resolve each owned lifecycle into immutable method metadata.""" @abstractmethod def encode( self, - program_index: int, + lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, @@ -54,7 +54,7 @@ def encode( @abstractmethod def decode( self, - program_index: int, + lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index bda086325da8..2ec1bef82553 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -67,13 +67,34 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: def _policy(native: SimpleNamespace) -> object: - return native.create_python_cold_page_codec.call_args.args[1] + return native.create_python_cold_page_codec.call_args.args[0] -def _plans(native: SimpleNamespace) -> list[object]: +def _layouts(native: SimpleNamespace) -> list[object]: return list(_policy(native)._layer_layouts.values()) +def _configure_lifecycle(native: SimpleNamespace, layer_bytes: dict[int, dict[str, int]]) -> object: + address = 0x10000 + layers = {} + for layer_id, roles in layer_bytes.items(): + hot = {} + for role, raw_bytes in roles.items(): + hot[role] = SimpleNamespace( + raw_base=address, + raw_slot_bytes=(raw_bytes + 15) // 16 * 16, + raw_bytes=raw_bytes, + ) + address += 0x10000 + layers[layer_id] = hot + _policy(native).configure([SimpleNamespace(layers=layers)]) + return _policy(native)._lifecycle_metadata[0] + + +def _configure_policy(native: SimpleNamespace, raw_bytes: int) -> object: + return _configure_lifecycle(native, {0: {"key": raw_bytes, "value": raw_bytes}}) + + def _write_quant_metadata(directory, algorithm="NVFP4"): metadata = { "producer": {"name": "modelopt"}, @@ -125,20 +146,32 @@ def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_p ) assert result is codec - plans = _plans(native) - assert [plan.layer_id for plan in plans] == [0, 1, 2] - assert [plan.runtime_type for plan in plans] == [1] * 3 + layouts = _layouts(native) + assert [layout.layer_id for layout in layouts] == [0, 1, 2] + assert _policy(native)._runtime_type == 1 assert [ ( - tuple(buffer.scales.nvfp4_orig_quant for buffer in plan.buffers), - tuple(buffer.scales.nvfp4_quant_orig for buffer in plan.buffers), + tuple(buffer.scales.nvfp4_orig_quant for buffer in layout.buffers), + tuple(buffer.scales.nvfp4_quant_orig for buffer in layout.buffers), ) - for plan in plans + for layout in layouts ] == [ ((2.0, 4.0), (0.5, 0.25)), ((8.0, 16.0), (0.125, 0.0625)), ((1.0, 1.0), (1.0, 1.0)), ] + metadata = _configure_lifecycle( + native, + {layer_id: {"key": 131072, "value": 131072} for layer_id in range(3)}, + ) + assert metadata.scales[:6].tolist() == [ + [2.0, 0.5, 1.0, 1.0], + [4.0, 0.25, 1.0, 1.0], + [8.0, 0.125, 1.0, 1.0], + [16.0, 0.0625, 1.0, 1.0], + [1.0, 1.0, 1.0, 1.0], + [1.0, 1.0, 1.0, 1.0], + ] def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: @@ -154,7 +187,7 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: num_kv_heads_per_layer=(8,), head_dim_per_layer=(128,), ) - target_plan = _plans(native)[0] + target_layout = _layouts(native)[0] manager.create_cold_page_codec( _cache_config((0, "attention")), runtime_dtype=DataType.BF16, @@ -164,16 +197,16 @@ def test_draft_codec_does_not_reuse_target_modelopt_scales(tmp_path) -> None: is_draft=True, ) - draft_plan = _plans(native)[0] - assert [buffer.scales.nvfp4_orig_quant for buffer in target_plan.buffers] == [ + draft_layout = _layouts(native)[0] + assert [buffer.scales.nvfp4_orig_quant for buffer in target_layout.buffers] == [ 2.0, 4.0, ] - assert [buffer.scales.nvfp4_orig_quant for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_orig_quant for buffer in draft_layout.buffers] == [ 1.0, 1.0, ] - assert [buffer.scales.nvfp4_quant_orig for buffer in draft_plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in draft_layout.buffers] == [ 1.0, 1.0, ] @@ -193,24 +226,25 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): head_dim_per_layer=(128,), ) - plan = _plans(native)[0] - assert [buffer.role for buffer in plan.buffers] == ["key", "value"] - assert [buffer.data_offset for buffer in plan.buffers] == [0, 1280] - assert [buffer.scale_offset for buffer in plan.buffers] == [2560, 2720] - assert plan.cold_page_bytes == 2880 - assert plan.padding_offset == plan.cold_page_bytes - assert plan.runtime_type == 0 - assert plan.num_kv_heads == 4 - assert plan.tokens_per_page == 5 - assert plan.head_dim == 128 - assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ + layout = _layouts(native)[0] + assert [buffer.role for buffer in layout.buffers] == ["key", "value"] + assert layout.num_kv_heads == 4 + assert layout.tokens_per_page == 5 + assert layout.head_dim == 128 + assert [buffer.scales.nvfp4_orig_quant for buffer in layout.buffers] == [ 1.0, 1.0, ] - assert [buffer.scales.nvfp4_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in layout.buffers] == [ 1.0, 1.0, ] + metadata = _configure_policy(native, raw_bytes=5120) + assert metadata.runtime_type == 0 + assert metadata.cold_page_bytes == 2880 + assert metadata.wide[:2, 3].tolist() == [0, 1280] + assert metadata.wide[:2, 4].tolist() == [2560, 2720] + assert metadata.integers[:2, 0].tolist() == [0, 0] def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: @@ -237,12 +271,12 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: head_dim_per_layer=(32,), ) - plan = _plans(native)[0] - assert [buffer.role for buffer in plan.buffers] == ["key", "value"] - assert [buffer.data_offset for buffer in plan.buffers] == [0, 80] - assert [buffer.scale_offset for buffer in plan.buffers] == [160, 170] - assert plan.padding_offset == 180 - assert plan.cold_page_bytes == 192 + metadata = _configure_policy(native, raw_bytes=320) + assert metadata.cold_page_bytes == 192 + assert metadata.wide[:2, 3].tolist() == [0, 80] + assert metadata.wide[:2, 4].tolist() == [160, 170] + assert metadata.wide[1, 5].item() == 180 + assert metadata.integers[:2, 0].tolist() == [0, 12] def test_provider_creates_one_native_codec_per_kv_cache_manager(): @@ -269,16 +303,8 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): def test_policy_forwards_a_4096_page_batch_through_one_custom_op(monkeypatch) -> None: native, _ = _native() - program = object() - prepare = MagicMock(return_value=program) encode = MagicMock() decode = MagicMock() - monkeypatch.setattr( - torch.ops.trtllm, - "prepare_nvfp4_cold_page_program", - prepare, - raising=False, - ) monkeypatch.setattr( torch.ops.trtllm, "nvfp4_cold_page_encode", @@ -311,14 +337,98 @@ def test_policy_forwards_a_4096_page_batch_through_one_custom_op(monkeypatch) -> } properties = policy.configure([SimpleNamespace(layers={0: hot})]) - assert prepare.call_count == 1 assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" policy.encode(0, 0x3000, 0x4000, 4096, 0x5000) policy.decode(0, 0x3000, 0x4000, 4096, 0x5000) - encode.assert_called_once_with(program, 0x3000, 0x4000, 4096, 0x5000) - decode.assert_called_once_with(program, 0x3000, 0x4000, 4096, 0x5000) + metadata = policy._lifecycle_metadata[0] + for operation in (encode, decode): + operation.assert_called_once() + arguments = operation.call_args.args + assert arguments[0] is metadata.wide + assert arguments[1] is metadata.integers + assert arguments[2] is metadata.scales + assert arguments[3:] == (2, 128, 1152, 1, 0x3000, 0x4000, 4096, 0x5000) + + +def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: + native, _ = _native() + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(16,), + ) + + with torch.device("meta"): + metadata = _configure_policy(native, raw_bytes=2048) + assert metadata.wide.device.type == "cpu" + assert metadata.integers.device.type == "cpu" + assert metadata.scales.device.type == "cpu" + + +def test_policy_rejects_invalid_resolved_hot_buffers() -> None: + native, _ = _native() + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + _cache_config((0, "attention")), + runtime_dtype=DataType.BF16, + pp_layers=(0,), + num_kv_heads_per_layer=(1,), + head_dim_per_layer=(16,), + ) + + policy = _policy(native) + + def hot(raw_base: int = 0x1000, raw_bytes: int = 2048) -> SimpleNamespace: + return SimpleNamespace( + raw_base=raw_base, + raw_slot_bytes=2048, + raw_bytes=raw_bytes, + ) + + with pytest.raises(ValueError, match="roles do not match"): + policy.configure( + [SimpleNamespace(layers={0: {"key": hot(), "value": hot(), "extra": hot()}})] + ) + with pytest.raises(ValueError, match="size does not match"): + policy.configure([SimpleNamespace(layers={0: {"key": hot(raw_bytes=32), "value": hot()}})]) + with pytest.raises(ValueError, match="16-byte aligned"): + policy.configure( + [SimpleNamespace(layers={0: {"key": hot(raw_base=0x1001), "value": hot()}})] + ) + + +def test_policy_rejects_more_than_256_lifecycle_buffers() -> None: + native, _ = _native() + layers = tuple( + AttentionLayerConfig( + layer_id=layer_id, + buffers=[ + BufferConfig(role="key", size=32), + BufferConfig(role="value", size=32), + ], + ) + for layer_id in range(129) + ) + cache_config = SimpleNamespace(tokens_per_block=1, layers=layers) + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + _manager().create_cold_page_codec( + cache_config, + runtime_dtype=DataType.BF16, + pp_layers=tuple(range(129)), + num_kv_heads_per_layer=(1,) * 129, + head_dim_per_layer=(16,) * 129, + ) + + with pytest.raises(ValueError, match="maximum is 256"): + _configure_lifecycle( + native, + {layer_id: {"key": 32, "value": 32} for layer_id in range(129)}, + ) def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: @@ -424,6 +534,13 @@ def test_scale_checkpoint_requires_kv_pair(tmp_path, present_kind): _load_modelopt_nvfp4_scales(str(tmp_path)) +def test_scale_checkpoint_requires_float32_reciprocals(tmp_path) -> None: + smallest_subnormal = torch.tensor(1e-45, dtype=torch.float32).item() + _write_scales(tmp_path, {7: (smallest_subnormal, smallest_subnormal)}) + with pytest.raises(ValueError, match="representable as float32"): + _load_modelopt_nvfp4_scales(str(tmp_path)) + + def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): native, codec = _native() _write_scales(tmp_path, {4: (0.5, 0.25)}) @@ -435,9 +552,9 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): num_kv_heads_per_layer=(0, 8), head_dim_per_layer=(128, 128), ) - plan = _plans(native)[0] - assert plan.layer_id == 1 - assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ + layout = _layouts(native)[0] + assert layout.layer_id == 1 + assert [buffer.scales.nvfp4_orig_quant for buffer in layout.buffers] == [ 2.0, 4.0, ] @@ -451,7 +568,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - assert native.create_python_cold_page_codec.call_args.args[0] == () + assert _policy(native).layer_ids == () def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -479,21 +596,21 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): ) assert result is codec - plan = _plans(native)[0] - assert plan.layer_id == 0 - assert plan.cold_page_bytes == 29184 - assert plan.padding_offset == 29184 - assert [buffer.role for buffer in plan.buffers] == ["key", "index_key"] - assert [buffer.scales is not None for buffer in plan.buffers] == [True, False] - assert [buffer.data_offset for buffer in plan.buffers] == [0, 20736] - assert [buffer.scale_offset for buffer in plan.buffers] == [18432, 0] - assert plan.runtime_type == 1 - assert plan.num_kv_heads == 1 - assert plan.tokens_per_page == 64 - assert plan.head_dim == 576 - scales = plan.buffers[0].scales + layout = _layouts(native)[0] + assert layout.layer_id == 0 + assert [buffer.role for buffer in layout.buffers] == ["key", "index_key"] + assert [buffer.scales is not None for buffer in layout.buffers] == [True, False] + assert _policy(native)._runtime_type == 1 + assert layout.num_kv_heads == 1 + assert layout.tokens_per_page == 64 + assert layout.head_dim == 576 + scales = layout.buffers[0].scales assert scales.nvfp4_orig_quant == scales.nvfp4_quant_orig == 1.0 - assert plan.buffers[1].scales is None + assert layout.buffers[1].scales is None + metadata = _configure_lifecycle(native, {0: {"key": 64 * 576 * 2, "index_key": 64 * 132}}) + assert metadata.cold_page_bytes == 29184 + assert metadata.wide[:2, 3].tolist() == [0, 20736] + assert metadata.wide[:2, 4].tolist() == [18432, 0] def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: @@ -521,19 +638,21 @@ def test_mla_all_non_latent_roles_are_explicit_lossless_spans() -> None: head_dim_per_layer=(32,), ) - plan = _plans(native)[0] - assert [buffer.role for buffer in plan.buffers] == [ + layout = _layouts(native)[0] + assert [buffer.role for buffer in layout.buffers] == [ "key", "index_key", "rope_state", ] - assert [buffer.scales is not None for buffer in plan.buffers] == [True, False, False] - assert [buffer.data_offset for buffer in plan.buffers] == [0, 90, 158] - assert plan.padding_offset == 165 - assert plan.cold_page_bytes == 176 + assert [buffer.scales is not None for buffer in layout.buffers] == [True, False, False] + metadata = _configure_lifecycle(native, {0: {"key": 320, "index_key": 68, "rope_state": 7}}) + assert metadata.wide[:3, 3].tolist() == [0, 90, 158] + assert metadata.wide[2, 5].item() == 165 + assert metadata.integers[2, 0].item() == 11 + assert metadata.cold_page_bytes == 176 -def test_tokens_per_block_override_expands_lossless_bytes() -> None: +def test_lossless_layout_uses_resolved_hot_buffer_bytes() -> None: native, _ = _native() cache_config = SimpleNamespace( tokens_per_block=4, @@ -542,11 +661,7 @@ def test_tokens_per_block_override_expands_lossless_bytes() -> None: layer_id=0, buffers=[ BufferConfig(role="key", size=128), - BufferConfig( - role="index_key", - size=3, - tokens_per_block_override=2, - ), + BufferConfig(role="index_key", size=3), ], ), ), @@ -561,10 +676,11 @@ def test_tokens_per_block_override_expands_lossless_bytes() -> None: head_dim_per_layer=(16,), ) - plan = _plans(native)[0] - assert [buffer.data_offset for buffer in plan.buffers] == [0, 36] - assert plan.padding_offset == 42 - assert plan.cold_page_bytes == 48 + metadata = _configure_lifecycle(native, {0: {"key": 128, "index_key": 6}}) + assert metadata.wide[:2, 3].tolist() == [0, 36] + assert metadata.wide[1, 5].item() == 42 + assert metadata.integers[1, 0].item() == 6 + assert metadata.cold_page_bytes == 48 @pytest.mark.parametrize( @@ -604,10 +720,17 @@ def test_mla_model_layouts_are_built_in_python( head_dim_per_layer=(576,) * expected_layers, ) - plans = _plans(native) - assert len(plans) == expected_layers - assert sum(len(plan.buffers) for plan in plans) == expected_buffers - assert sum(plan.cold_page_bytes for plan in plans) == expected_bytes + layouts = _layouts(native) + assert len(layouts) == expected_layers + assert sum(len(layout.buffers) for layout in layouts) == expected_buffers + layer_bytes = { + layer_id: { + "key": 64 * 576 * 2, + **({"index_key": 64 * (128 + 4)} if has_index else {}), + } + for layer_id, has_index in enumerate(owns_index) + } + assert _configure_lifecycle(native, layer_bytes).cold_page_bytes == expected_bytes def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): @@ -622,19 +745,19 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): head_dim_per_layer=(128,), ) - plan = _plans(native)[0] - assert plan.runtime_type == 2 - assert [buffer.scales.nvfp4_orig_quant for buffer in plan.buffers] == [ + layout = _layouts(native)[0] + assert _policy(native)._runtime_type == 2 + assert [buffer.scales.nvfp4_orig_quant for buffer in layout.buffers] == [ 2.0, 4.0, ] - assert [buffer.scales.nvfp4_quant_orig for buffer in plan.buffers] == [ + assert [buffer.scales.nvfp4_quant_orig for buffer in layout.buffers] == [ 0.5, 0.25, ] assert all( buffer.scales.fp8_orig_quant == buffer.scales.fp8_quant_orig == 1.0 - for buffer in plan.buffers + for buffer in layout.buffers ) diff --git a/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py b/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py new file mode 100644 index 000000000000..9b0ea56ec293 --- /dev/null +++ b/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""CPU-only ABI checks for NVFP4 cold-page custom ops.""" + +import pytest +import torch + +pytestmark = pytest.mark.cpu_only + + +def _encode(wide: torch.Tensor, integers: torch.Tensor, scales: torch.Tensor) -> None: + torch.ops.trtllm.nvfp4_cold_page_encode( + wide, + integers, + scales, + 1, + 2, + 16, + 0, + 0, + 0, + 0, + 0, + ) + + +def test_rejects_malformed_metadata_tensors() -> None: + wide = torch.zeros((256, 6), dtype=torch.int64) + integers = torch.zeros((256, 5), dtype=torch.int32) + scales = torch.zeros((256, 4), dtype=torch.float32) + + invalid = ( + (wide.to(torch.int32), integers, scales, "wide metadata"), + (torch.zeros((256, 12), dtype=torch.int64)[:, ::2], integers, scales, "wide metadata"), + (wide, integers[:, :4], scales, "integer metadata"), + (wide, integers, scales.to(torch.float64), "scale metadata"), + ) + for bad_wide, bad_integers, bad_scales, message in invalid: + with pytest.raises(RuntimeError, match=message): + _encode(bad_wide, bad_integers, bad_scales) + + +def test_rejects_invalid_launch_scalars() -> None: + wide = torch.zeros((256, 6), dtype=torch.int64) + integers = torch.zeros((256, 5), dtype=torch.int32) + scales = torch.zeros((256, 4), dtype=torch.float32) + with pytest.raises(RuntimeError, match="buffer count"): + torch.ops.trtllm.nvfp4_cold_page_encode(wide, integers, scales, 0, 2, 16, 0, 0, 0, 0, 0) From 74b788cd8de98f26029b82618b110bcc477130fd Mon Sep 17 00:00:00 2001 From: Tianrui Hu <32944717+Hudayday@users.noreply.github.com> Date: Wed, 26 Aug 2026 09:56:38 -0700 Subject: [PATCH 14/29] [None][refactor] Remove duplicate cold-page callback ABI Signed-off-by: Tianrui Hu <32944717+Hudayday@users.noreply.github.com> --- .../kernels/nvfp4ColdPageKernels.cu | 31 +++++++++++++------ .../kernels/nvfp4ColdPageKernels.h | 3 -- .../coldPageCallbackAbi.h | 29 ----------------- .../nanobind/kvCacheCompression/bindings.cpp | 11 ++++--- .../kernels/nvfp4ColdPageKernelsTest.cpp | 21 +++++++------ 5 files changed, 39 insertions(+), 56 deletions(-) delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index 577a9778632a..57d2d3b9433f 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -112,6 +112,19 @@ struct Nvfp4ColdPageBuffer Nvfp4ColdPageKernelParams params; }; +// Private KVCM Page-index view; its matching layout is asserted here and at the callback boundary. +struct alignas(8) PageIndexPairView +{ + std::int32_t dst; + std::int32_t src; +}; + +static_assert(sizeof(PageIndexPairView) == 8); +static_assert(alignof(PageIndexPairView) == 8); +static_assert(offsetof(PageIndexPairView, dst) == 0); +static_assert(offsetof(PageIndexPairView, src) == 4); +static_assert(std::is_trivially_copyable_v); + // Keep every CTA iteration and shared-memory tile on complete 16-value scale groups. static_assert(kElementsPerScaleGroup % kElementsPerHalfGroup == 0, "An NVFP4 scale group must contain a whole number of half-groups"); @@ -121,7 +134,7 @@ static_assert( kMaxScaleBytesPerTile % sizeof(uint4) == 0, "A full scale tile must preserve the 16-byte transfer fast path"); // Keep both kernel argument packs within CUDA's modern 32,764-byte limit. -static_assert(sizeof(std::array) + sizeof(Nvfp4ColdPageWideTable) +static_assert(sizeof(std::array) + sizeof(Nvfp4ColdPageWideTable) + sizeof(Nvfp4ColdPageIntegerTable) + sizeof(Nvfp4ColdPageScaleTable) + 2U * sizeof(std::uintptr_t) <= kKernelParameterLimitBytes, "Cold-page kernel arguments exceed CUDA's parameter limit"); @@ -172,7 +185,7 @@ __device__ Nvfp4ColdPageBuffer loadBuffer(std::uint32_t index, Nvfp4ColdPageWide } __device__ OffloadBufferTask resolveOffloadTask( - ColdPageIndexPair const& page, Nvfp4ColdPageBuffer const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) + PageIndexPairView const& page, Nvfp4ColdPageBuffer const& buffer, std::uint8_t* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.src); auto* coldPage = coldBase + static_cast(page.dst) * coldPageBytes; @@ -180,7 +193,7 @@ __device__ OffloadBufferTask resolveOffloadTask( coldPage + buffer.coldDataOffset, coldPage + buffer.coldScaleOffset, coldPage + buffer.coldPaddingOffset}; } -__device__ OnboardBufferTask resolveOnboardTask(ColdPageIndexPair const& page, Nvfp4ColdPageBuffer const& buffer, +__device__ OnboardBufferTask resolveOnboardTask(PageIndexPairView const& page, Nvfp4ColdPageBuffer const& buffer, std::uint8_t const* coldBase, std::size_t coldPageBytes) { std::size_t const gpuPage = static_cast(page.dst); @@ -466,7 +479,7 @@ __device__ void restoreNvfp4Pair(uint2 packedPair, T* output, std::uint32_t firs // FP16/BF16 GPU Page -> mapped-Host NVFP4 in bounded tiles. template __global__ void offloadFrom16BitTiledKernel( - std::array const __grid_constant__ pages, + std::array const __grid_constant__ pages, Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t* coldBase, std::size_t coldPageBytes) { @@ -546,7 +559,7 @@ __global__ void offloadFrom16BitTiledKernel( // FP8 E4M3 GPU Page -> mapped-Host NVFP4 in bounded tiles. __global__ void offloadFromFp8TiledKernel( - std::array const __grid_constant__ pages, + std::array const __grid_constant__ pages, Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t* coldBase, std::size_t coldPageBytes) { @@ -688,7 +701,7 @@ __device__ void loadCompactRangeFromHost(std::uint8_t* compactStages, OnboardBuf // Mapped-Host NVFP4 -> runtime GPU Page in bounded tiles. template -__global__ void onboardTiledKernel(std::array const __grid_constant__ pages, +__global__ void onboardTiledKernel(std::array const __grid_constant__ pages, Nvfp4ColdPageWideTable const __grid_constant__ wide, Nvfp4ColdPageIntegerTable const __grid_constant__ integers, Nvfp4ColdPageScaleTable const __grid_constant__ scales, std::uint8_t const* coldBase, std::size_t coldPageBytes) { @@ -768,13 +781,13 @@ void launchPageChunks(Kernel kernel, void const* pages, std::size_t numPages, st { std::uint32_t const numChunkPages = static_cast(std::min(numPages - offset, kMaxTasksPerLaunch)); - auto const* chunkPages = pageBytes + offset * sizeof(ColdPageIndexPair); + auto const* chunkPages = pageBytes + offset * sizeof(PageIndexPairView); // CUDA copies the full by-value array, so pad only the final partial chunk. - std::array paddedPages{}; + std::array paddedPages{}; if (numChunkPages < kMaxTasksPerLaunch) { - std::memcpy(paddedPages.data(), chunkPages, numChunkPages * sizeof(ColdPageIndexPair)); + std::memcpy(paddedPages.data(), chunkPages, numChunkPages * sizeof(PageIndexPairView)); chunkPages = reinterpret_cast(paddedPages.data()); } diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h index 95bce5177ed8..6a249b246cd9 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.h @@ -19,7 +19,6 @@ #pragma once #include "tensorrt_llm/common/config.h" -#include "tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h" #include #include @@ -39,8 +38,6 @@ enum class Nvfp4ColdPageRuntimeType : std::uint8_t kFp8E4m3 = 2, }; -using ColdPageIndexPair = ::tensorrt_llm::kv_cache_compression::ColdPageIndexPair; - inline constexpr std::uint32_t kNvfp4ColdPageMaxBuffersPerLaunch = 256; inline constexpr std::uint32_t kNvfp4ColdPageWideFields = 6; inline constexpr std::uint32_t kNvfp4ColdPageIntegerFields = 5; diff --git a/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h b/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h deleted file mode 100644 index 8958961f50ae..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h +++ /dev/null @@ -1,29 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#pragma once - -#include -#include -#include - -namespace tensorrt_llm::kv_cache_compression -{ - -//! Algorithm-neutral descriptor ABI borrowed by one cold-page callback. -struct alignas(8) ColdPageIndexPair -{ - std::int32_t dst; - std::int32_t src; -}; - -static_assert(sizeof(ColdPageIndexPair) == 8); -static_assert(alignof(ColdPageIndexPair) == 8); -static_assert(offsetof(ColdPageIndexPair, dst) == 0); -static_assert(offsetof(ColdPageIndexPair, src) == 4); -static_assert(std::is_trivially_copyable_v); - -} // namespace tensorrt_llm::kv_cache_compression diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 5a28fbf63fbb..6bb6a54a2d69 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,7 +17,6 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/coldPageCallbackAbi.h" #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" #include @@ -30,6 +29,7 @@ #include #include #include +#include #include namespace nb = nanobind; @@ -41,10 +41,11 @@ namespace tensorrt_llm::nanobind::kv_cache_compression namespace { -static_assert(sizeof(kv::PageIndexPair) == sizeof(compression::ColdPageIndexPair)); -static_assert(alignof(kv::PageIndexPair) == alignof(compression::ColdPageIndexPair)); -static_assert(offsetof(kv::PageIndexPair, dst) == offsetof(compression::ColdPageIndexPair, dst)); -static_assert(offsetof(kv::PageIndexPair, src) == offsetof(compression::ColdPageIndexPair, src)); +static_assert(sizeof(kv::PageIndexPair) == 8); +static_assert(alignof(kv::PageIndexPair) == 8); +static_assert(offsetof(kv::PageIndexPair, dst) == 0); +static_assert(offsetof(kv::PageIndexPair, src) == 4); +static_assert(std::is_trivially_copyable_v); //! Algorithm-neutral adapter from KVCM's native codec calls to one Python policy. class PythonColdPageCodec final : public compression::NativeColdPageCodec diff --git a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp index 030d8030dd9e..9a552d1e83af 100644 --- a/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp +++ b/cpp/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp @@ -17,6 +17,7 @@ */ #include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" +#include "kv_cache_manager_v2/coldPageCodec.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.h" #include "tensorrt_llm/common/cudaUtils.h" @@ -41,7 +42,7 @@ namespace using tensorrt_llm::batch_manager::kv_cache_manager_v2::HostMem; using tensorrt_llm::batch_manager::kv_cache_manager_v2::MemAddress; -using tensorrt_llm::kernels::ColdPageIndexPair; +using tensorrt_llm::batch_manager::kv_cache_manager_v2::PageIndexPair; using tensorrt_llm::kernels::Nvfp4ColdPageIntegerTable; using tensorrt_llm::kernels::Nvfp4ColdPageRuntimeType; using tensorrt_llm::kernels::Nvfp4ColdPageScaleTable; @@ -585,7 +586,7 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG MappedHostRegion compactPages(coldBaseOffset + slotCapacity * compactSlotBytes); auto* compactBase = compactPages.bytes() + coldBaseOffset; std::vector, 2>> rawHost(numPages); - std::vector offloadTasks; + std::vector offloadTasks; offloadTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { @@ -671,7 +672,7 @@ void runColdPageRoundTrip(RawKind kind, PageGeometry const& geometry = kDefaultG } } - std::vector onboardTasks; + std::vector onboardTasks; onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) { @@ -822,8 +823,8 @@ void runPartialPageTailIsolation(RawKind kind) DeviceRegion rawOutputK(numPages * rawSlotBytes); DeviceRegion rawOutputV(numPages * rawSlotBytes); MappedHostRegion compactPages(numPages * compactSlotBytes); - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; offloadTasks.reserve(numPages); onboardTasks.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) @@ -984,8 +985,8 @@ void runUnaryMlaWithLosslessSideRoundTrip(RawKind kind) std::array, numPages> mlaHost; std::array, numPages> sideHost; std::array references; - std::vector offloadTasks; - std::vector onboardTasks; + std::vector offloadTasks; + std::vector onboardTasks; for (std::size_t page = 0; page < numPages; ++page) { mlaHost[page] = makeRawPage(kind, page, 0U, params, geometry, InputPattern::kDense); @@ -1140,7 +1141,7 @@ TEST(Nvfp4ColdPageWholePageTest, DifferentLayerScalesRemainInOneCompletePageBatc = makeNvfp4ColdPageTestMetadata(inputMetadatas, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); auto const outputMetadata = makeNvfp4ColdPageTestMetadata(outputMetadatas, coldPageBytes, Nvfp4ColdPageRuntimeType::kBfloat16); - ColdPageIndexPair const page{0, 0}; + PageIndexPair const page{0, 0}; invokeNvfp4ColdPageEncode(&page, 1U, inputMetadata, compactPage.data(), stream); invokeNvfp4ColdPageDecode(&page, 1U, outputMetadata, compactPage.data(), stream); ASSERT_EQ(cudaStreamSynchronize(stream), cudaSuccess); @@ -1202,8 +1203,8 @@ void expectWholePageLaunchTopology(std::size_t numPages, std::vector offloadPages; - std::vector onboardPages; + std::vector offloadPages; + std::vector onboardPages; offloadPages.reserve(numPages); onboardPages.reserve(numPages); for (std::size_t page = 0; page < numPages; ++page) From e55e82941022cd18de3fed57f6b5884df17354b3 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 12:11:04 -0700 Subject: [PATCH 15/29] [None][refactor] Simplify cold-page quantization dispatch Signed-off-by: tianruih --- .../nanobind/kvCacheCompression/bindings.cpp | 46 ++++-- cpp/tensorrt_llm/thop/CMakeLists.txt | 4 - .../thop/coldPageMethods/nvfp4ColdPageOp.cu | 84 ----------- .../{nvfp4.py => nvfp4_quantization.py} | 132 +++++++++--------- .../quantization_for_cold_page.py | 88 +----------- tensorrt_llm/_torch/pyexecutor/_util.py | 9 +- .../test_quantization_for_cold_page.py | 76 +++++----- .../test_nvfp4_cold_page_op.py | 48 ------- 8 files changed, 152 insertions(+), 335 deletions(-) delete mode 100644 cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu rename tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/{nvfp4.py => nvfp4_quantization.py} (81%) delete mode 100644 tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 6bb6a54a2d69..724f86ccc7f9 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,6 +17,7 @@ */ #include "bindings.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" #include @@ -30,10 +31,10 @@ #include #include #include -#include namespace nb = nanobind; namespace compression = tensorrt_llm::kv_cache_compression; +namespace kernels = tensorrt_llm::kernels; namespace kv = tensorrt_llm::batch_manager::kv_cache_manager_v2; namespace tensorrt_llm::nanobind::kv_cache_compression @@ -114,7 +115,7 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec void invoke(char const* method, std::size_t lifecycleIndex, ColdPointer coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) { - // Forward the complete KVCM batch once. The method custom op owns any launch chunking. + // Forward the complete KVCM batch once. The native launcher owns any chunking. nb::gil_scoped_acquire acquire; try { @@ -130,11 +131,6 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec PyObject* mPolicy; }; -std::unique_ptr createPythonColdPageCodec(nb::handle policy) -{ - return std::make_unique(policy); -} - } // namespace void initBindings(nb::module_& module) @@ -159,7 +155,41 @@ void initBindings(nb::module_& module) .def_rw("cold_page_bytes", &compression::ColdPageLifecycleProperties::coldPageBytes) .def_rw("page_index_location", &compression::ColdPageLifecycleProperties::pageIndexLocation); - module.def("create_python_cold_page_codec", &createPythonColdPageCodec, nb::arg("policy")); + module.def( + "create_python_cold_page_codec", + [](nb::handle policy) -> std::unique_ptr + { return std::make_unique(policy); }, + nb::arg("policy")); + module.def( + "invoke_nvfp4_cold_page_encode", + [](std::uintptr_t pageIndices, std::size_t numPages, std::uintptr_t wide, std::uintptr_t integers, + std::uintptr_t scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, int runtimeType, std::uintptr_t coldBase, std::uintptr_t stream) + { + kernels::invokeNvfp4ColdPageEncode(reinterpret_cast(pageIndices), numPages, + reinterpret_cast(wide), reinterpret_cast(integers), + reinterpret_cast(scales), numBuffers, maxHalfGroupsPerTile, coldPageBytes, + static_cast(runtimeType), reinterpret_cast(coldBase), + reinterpret_cast(stream)); + }, + nb::arg("page_indices"), nb::arg("num_pages"), nb::arg("wide"), nb::arg("integers"), nb::arg("scales"), + nb::arg("num_buffers"), nb::arg("max_half_groups_per_tile"), nb::arg("cold_page_bytes"), + nb::arg("runtime_type"), nb::arg("cold_base"), nb::arg("stream"), nb::call_guard()); + module.def( + "invoke_nvfp4_cold_page_decode", + [](std::uintptr_t pageIndices, std::size_t numPages, std::uintptr_t wide, std::uintptr_t integers, + std::uintptr_t scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, + std::size_t coldPageBytes, int runtimeType, std::uintptr_t coldBase, std::uintptr_t stream) + { + kernels::invokeNvfp4ColdPageDecode(reinterpret_cast(pageIndices), numPages, + reinterpret_cast(wide), reinterpret_cast(integers), + reinterpret_cast(scales), numBuffers, maxHalfGroupsPerTile, coldPageBytes, + static_cast(runtimeType), reinterpret_cast(coldBase), + reinterpret_cast(stream)); + }, + nb::arg("page_indices"), nb::arg("num_pages"), nb::arg("wide"), nb::arg("integers"), nb::arg("scales"), + nb::arg("num_buffers"), nb::arg("max_half_groups_per_tile"), nb::arg("cold_page_bytes"), + nb::arg("runtime_type"), nb::arg("cold_base"), nb::arg("stream"), nb::call_guard()); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index 612aadbe5297..cdf0a0228fac 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -37,9 +37,6 @@ set_property(TARGET th_utils PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) target_link_libraries(th_utils PUBLIC ${SHARED_TARGET} ${TORCH_LIBRARIES} ${CUBLAS_LIB} ${CURAND_LIB}) -file(GLOB COLD_PAGE_METHOD_SOURCES CONFIGURE_DEPENDS - "${CMAKE_CURRENT_SOURCE_DIR}/coldPageMethods/*.cu") - # TODO This does not compile with internal cutlass MOE gemm add_library( th_common SHARED @@ -114,7 +111,6 @@ add_library( fp8PerTensorScaleMoe.cpp fp4BlockScaleMoe.cpp noAuxTcOp.cpp - ${COLD_PAGE_METHOD_SOURCES} fusedCatFp4Op.cpp fusedCatFp8Op.cpp IndexerKCacheGatherOp.cpp diff --git a/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu b/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu deleted file mode 100644 index 2f1a473d9b28..000000000000 --- a/cpp/tensorrt_llm/thop/coldPageMethods/nvfp4ColdPageOp.cu +++ /dev/null @@ -1,84 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. - * All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - */ - -#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" - -#include -#include - -namespace -{ - -using namespace tensorrt_llm::kernels; - -void checkMetadata(at::Tensor const& wide, at::Tensor const& integers, at::Tensor const& scales) -{ - TORCH_CHECK(wide.device().is_cpu() && wide.scalar_type() == at::kLong && wide.is_contiguous() - && wide.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageWideFields}), - "NVFP4 wide metadata must be a contiguous CPU int64 [256, 6] tensor"); - TORCH_CHECK(integers.device().is_cpu() && integers.scalar_type() == at::kInt && integers.is_contiguous() - && integers.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageIntegerFields}), - "NVFP4 integer metadata must be a contiguous CPU int32 [256, 5] tensor"); - TORCH_CHECK(scales.device().is_cpu() && scales.scalar_type() == at::kFloat && scales.is_contiguous() - && scales.sizes() == at::IntArrayRef({kNvfp4ColdPageMaxBuffersPerLaunch, kNvfp4ColdPageScaleFields}), - "NVFP4 scale metadata must be a contiguous CPU float32 [256, 4] tensor"); -} - -template -void runNvfp4ColdPage(at::Tensor const& wide, at::Tensor const& integers, at::Tensor const& scales, - std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, std::int64_t runtimeType, - std::int64_t coldBase, std::int64_t pagePairs, std::int64_t pageCount, std::int64_t stream) -{ - // This internal op consumes metadata prepared and validated by its Python policy. - checkMetadata(wide, integers, scales); - TORCH_CHECK(numBuffers > 0 && numBuffers <= kNvfp4ColdPageMaxBuffersPerLaunch, "Invalid NVFP4 buffer count"); - TORCH_CHECK(maxHalfGroupsPerTile > 0 && maxHalfGroupsPerTile <= 2048, "Invalid NVFP4 tile geometry"); - TORCH_CHECK(coldPageBytes > 0 && pageCount >= 0, "Invalid NVFP4 cold-page size or Page count"); - TORCH_CHECK(runtimeType >= 0 && runtimeType <= 2, "Invalid NVFP4 runtime type"); - TORCH_CHECK(coldBase >= 0 && pagePairs >= 0 && stream >= 0, "NVFP4 pointer arguments must be non-negative"); - - auto const* wideData = wide.const_data_ptr(); - auto const* integerData = integers.const_data_ptr(); - auto const* scaleData = scales.const_data_ptr(); - auto const type = static_cast(runtimeType); - auto const* pages = reinterpret_cast(static_cast(pagePairs)); - auto const cudaStream = reinterpret_cast(static_cast(stream)); - - if constexpr (Encode) - { - invokeNvfp4ColdPageEncode(pages, static_cast(pageCount), wideData, integerData, scaleData, - static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), - static_cast(coldPageBytes), type, - reinterpret_cast(static_cast(coldBase)), cudaStream); - } - else - { - invokeNvfp4ColdPageDecode(pages, static_cast(pageCount), wideData, integerData, scaleData, - static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), - static_cast(coldPageBytes), type, - reinterpret_cast(static_cast(coldBase)), cudaStream); - } -} - -} // namespace - -TORCH_LIBRARY_FRAGMENT(trtllm, m) -{ - m.def( - "nvfp4_cold_page_encode(Tensor wide, Tensor integers, Tensor scales, int num_buffers, " - "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int page_pairs, " - "int page_count, int stream) -> ()"); - m.def( - "nvfp4_cold_page_decode(Tensor wide, Tensor integers, Tensor scales, int num_buffers, " - "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int page_pairs, " - "int page_count, int stream) -> ()"); -} - -TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, m) -{ - m.impl("nvfp4_cold_page_encode", &runNvfp4ColdPage); - m.impl("nvfp4_cold_page_decode", &runNvfp4ColdPage); -} diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py similarity index 81% rename from tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py rename to tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index 66ed0688e567..d143e9d1650f 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""NVFP4 policy and cold-page layout construction.""" +"""NVFP4 quantization policy and cold-page layout construction.""" import json import math @@ -8,7 +8,7 @@ import re from dataclasses import dataclass from pathlib import Path -from typing import Sequence +from typing import TYPE_CHECKING, Sequence import torch @@ -18,12 +18,14 @@ ) from ...pyexecutor.resource_manager import DataType -from .quantization_for_cold_page import ColdPageCodecPolicy, ColdPageQuantizationMethod +from .quantization_for_cold_page import ColdPageQuantizationCompression -ScalePair = tuple[float, float] -LayerScales = tuple[ScalePair, ScalePair] +if TYPE_CHECKING: + from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig -_IDENTITY_NVFP4_SCALES: LayerScales = ((1.0, 1.0), (1.0, 1.0)) +_LayerScales = tuple[tuple[float, float], tuple[float, float]] + +_IDENTITY_NVFP4_SCALES: _LayerScales = ((1.0, 1.0), (1.0, 1.0)) _MODEL_OPT_LANGUAGE_KV_SCALE_KEY = re.compile( r"^model(?:\.language_model)?\.layers\.(?P\d+)\.self_attn\." r"(?P[kv])_proj\.(?P=kind)_scale$" @@ -74,12 +76,11 @@ class _Nvfp4ColdPageMetadata: num_buffers: int max_half_groups_per_tile: int cold_page_bytes: int - runtime_type: int def _load_modelopt_nvfp4_scales( checkpoint_path: str | None, -) -> dict[int, LayerScales]: +) -> dict[int, _LayerScales]: """Load optional ModelOpt NVFP4 K/V global scales by model layer.""" if checkpoint_path is None or os.environ.get("TRTLLM_LOAD_KV_SCALES", "1") != "1": @@ -126,7 +127,7 @@ def _load_modelopt_nvfp4_scales( layer_values = values.setdefault(int(match.group("layer_id")), {"k": [], "v": []}) layer_values[match.group("kind")].append(value) - result: dict[int, LayerScales] = {} + result: dict[int, _LayerScales] = {} for layer_id, layer_values in values.items(): k_values, v_values = layer_values["k"], layer_values["v"] if not k_values or not v_values: @@ -147,22 +148,15 @@ def _load_modelopt_nvfp4_scales( return result -def _align_up(value: int, alignment: int = _COLD_PAGE_ALIGNMENT) -> int: - return (value + alignment - 1) // alignment * alignment - - -class Nvfp4ColdPagePolicy(ColdPageCodecPolicy): - """Resolve NVFP4 metadata and submit one CUDA op per codec batch.""" +class _Nvfp4ColdPagePolicy: + """Resolve NVFP4 metadata and submit one native launch per codec batch.""" def __init__(self, layer_layouts: Sequence[_Nvfp4LayerLayout], runtime_type: int) -> None: self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} + self.layer_ids = tuple(sorted(self._layer_layouts)) self._runtime_type = runtime_type self._lifecycle_metadata: list[_Nvfp4ColdPageMetadata] = [] - @property - def layer_ids(self) -> tuple[int, ...]: - return tuple(sorted(self._layer_layouts)) - def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: """Resolve hot buffers into immutable Python-owned launch metadata.""" @@ -184,19 +178,20 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: raise ValueError( f"Cold-page layer {layer_id} roles do not match its KVCM layout" ) - compressed = [buffer for buffer in layout.buffers if buffer.scales] - packed_bytes = ( - layout.num_kv_heads * layout.tokens_per_page * layout.head_dim - ) // _ELEMENTS_PER_BYTE - scale_bytes = ( - layout.num_kv_heads * layout.tokens_per_page * layout.head_dim - ) // _ELEMENTS_PER_SCALE + elements = layout.num_kv_heads * layout.tokens_per_page * layout.head_dim + element_bytes = 1 if self._runtime_type == 2 else 2 + expected_raw_bytes = elements * element_bytes + half_groups = elements // _ELEMENTS_PER_HALF_GROUP + compressed_count = sum(buffer.scales is not None for buffer in layout.buffers) + packed_bytes = elements // _ELEMENTS_PER_BYTE + scale_bytes = elements // _ELEMENTS_PER_SCALE layer_start = cold_page_bytes - scale_start = layer_start + len(compressed) * packed_bytes - cursor = scale_start + len(compressed) * scale_bytes + scale_start = layer_start + compressed_count * packed_bytes + cursor = scale_start + compressed_count * scale_bytes compressed_index = 0 for buffer in layout.buffers: + is_compressed = buffer.scales is not None hot = hot_buffers[buffer.role] raw_base = int(hot.raw_base) raw_slot_bytes = int(hot.raw_slot_bytes) @@ -204,27 +199,16 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: if raw_base <= 0 or raw_bytes <= 0 or raw_bytes > raw_slot_bytes: raise ValueError("Cold-page hot buffer has invalid address or size") - if buffer.scales: + if is_compressed: data_offset = layer_start + compressed_index * packed_bytes scale_offset = scale_start + compressed_index * scale_bytes compressed_index += 1 - expected_raw_bytes = ( - layout.num_kv_heads - * layout.tokens_per_page - * layout.head_dim - * (1 if self._runtime_type == 2 else 2) - ) if raw_bytes != expected_raw_bytes: raise ValueError("Hot buffer size does not match NVFP4 geometry") if raw_base % 16 or raw_slot_bytes % 16: raise ValueError( "NVFP4 hot address and Slot stride must be 16-byte aligned" ) - half_groups = ( - expected_raw_bytes - // (1 if self._runtime_type == 2 else 2) - // _ELEMENTS_PER_HALF_GROUP - ) max_half_groups_per_tile = max( max_half_groups_per_tile, min(half_groups, _MAX_HALF_GROUPS_PER_TILE), @@ -247,13 +231,13 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: integer_rows.append( [ 0, - _NVFP4_TRANSFORM if buffer.scales else _LOSSLESS_TRANSFORM, - layout.num_kv_heads if buffer.scales else 0, - layout.tokens_per_page if buffer.scales else 0, - layout.head_dim if buffer.scales else 0, + _NVFP4_TRANSFORM if is_compressed else _LOSSLESS_TRANSFORM, + layout.num_kv_heads if is_compressed else 0, + layout.tokens_per_page if is_compressed else 0, + layout.head_dim if is_compressed else 0, ] ) - buffer_scales = buffer.scales or _Nvfp4Scales(1.0, 1.0) + buffer_scales = buffer.scales if is_compressed else _Nvfp4Scales(1.0, 1.0) scale_rows.append( [ buffer_scales.nvfp4_orig_quant, @@ -262,7 +246,11 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: buffer_scales.fp8_quant_orig, ] ) - layer_end = _align_up(cursor) + layer_end = ( + (cursor + _COLD_PAGE_ALIGNMENT - 1) + // _COLD_PAGE_ALIGNMENT + * _COLD_PAGE_ALIGNMENT + ) wide_rows[-1][5] = cursor integer_rows[-1][0] = layer_end - cursor cold_page_bytes = layer_end @@ -294,7 +282,6 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: num_buffers=num_buffers, max_half_groups_per_tile=max_half_groups_per_tile, cold_page_bytes=cold_page_bytes, - runtime_type=self._runtime_type, ) ) lifecycle_properties = native.ColdPageLifecycleProperties() @@ -313,18 +300,20 @@ def encode( num_pages: int, stream: int, ) -> None: + from tensorrt_llm.bindings.internal import kv_cache_compression as native + metadata = self._lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_encode( - metadata.wide, - metadata.integers, - metadata.scales, + native.invoke_nvfp4_cold_page_encode( + page_indices, + num_pages, + metadata.wide.data_ptr(), + metadata.integers.data_ptr(), + metadata.scales.data_ptr(), metadata.num_buffers, metadata.max_half_groups_per_tile, metadata.cold_page_bytes, - metadata.runtime_type, + self._runtime_type, cold_base, - page_indices, - num_pages, stream, ) @@ -336,27 +325,30 @@ def decode( num_pages: int, stream: int, ) -> None: + from tensorrt_llm.bindings.internal import kv_cache_compression as native + metadata = self._lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_decode( - metadata.wide, - metadata.integers, - metadata.scales, + native.invoke_nvfp4_cold_page_decode( + page_indices, + num_pages, + metadata.wide.data_ptr(), + metadata.integers.data_ptr(), + metadata.scales.data_ptr(), metadata.num_buffers, metadata.max_half_groups_per_tile, metadata.cold_page_bytes, - metadata.runtime_type, + self._runtime_type, cold_base, - page_indices, - num_pages, stream, ) -class Nvfp4ColdPageQuantization(ColdPageQuantizationMethod): - """Build one fresh NVFP4 callback policy for each KVCM construction.""" +class Nvfp4ColdPageQuantizationCompression(ColdPageQuantizationCompression): + """NVFP4 cold-page quantization and per-KVCM codec construction.""" - def __init__(self, checkpoint_path: str | None) -> None: - self._model_scales = _load_modelopt_nvfp4_scales(checkpoint_path) + def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: + super().__init__(config) + self._model_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) def create_cold_page_codec( self, @@ -388,13 +380,13 @@ def create_cold_page_codec( layer_layouts = [] for layer in attention_layers: layer_id = int(layer.layer_id) - buffers_by_role = {str(buffer.role): buffer for buffer in layer.buffers} - if "key" not in buffers_by_role: + buffer_roles = {str(buffer.role) for buffer in layer.buffers} + if "key" not in buffer_roles: raise NotImplementedError( "NVFP4 cold-page compression requires an Attention key buffer" ) - compressed_roles = ("key", "value") if "value" in buffers_by_role else ("key",) + compressed_roles = ("key", "value") if "value" in buffer_roles else ("key",) if len(compressed_roles) == 2 and not is_draft: orig_quant, quant_orig = self._model_scales.get( int(pp_layers[layer_id]), _IDENTITY_NVFP4_SCALES @@ -431,5 +423,7 @@ def create_cold_page_codec( ) ) - policy = Nvfp4ColdPagePolicy(layer_layouts, runtime_type or 0) + policy = _Nvfp4ColdPagePolicy( + layer_layouts, runtime_type if runtime_type is not None else 0 + ) return native.create_python_cold_page_codec(policy) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index f775e3d44896..84bc1a7b5e0f 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -1,87 +1,20 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Quantization policies for KVCM V2 cold pages.""" +"""Base class for KVCM V2 cold-page quantization.""" from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Sequence +from typing import Sequence from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager -if TYPE_CHECKING: - from tensorrt_llm.llmapi.llm_args import ColdPageQuantizationCompressionConfig - -class ColdPageQuantizationMethod(ABC): - """Configured quantization method that creates one codec per KVCM.""" - - @abstractmethod - def create_cold_page_codec( - self, - cache_config: object, - *, - runtime_dtype: DataType, - pp_layers: Sequence[int], - num_kv_heads_per_layer: Sequence[int], - head_dim_per_layer: Sequence[int], - is_draft: bool = False, - ) -> object: - """Create a codec using this method's immutable calibration.""" - - -class ColdPageCodecPolicy(ABC): - """Per-KVCM Python callback contract used by the generic native codec.""" - - @property - @abstractmethod - def layer_ids(self) -> tuple[int, ...]: - """Layers transformed by this policy; other lifecycles stay lossless.""" - - @abstractmethod - def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: - """Resolve each owned lifecycle into immutable method metadata.""" - - @abstractmethod - def encode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - """Submit one complete hot-to-cold Page batch.""" - - @abstractmethod - def decode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - """Submit one complete cold-to-hot Page batch.""" - - -class ColdPageQuantizationCompression(KVCacheCompressionManager): - """Select and own the configured cold-page quantization method.""" +class ColdPageQuantizationCompression(KVCacheCompressionManager, ABC): + """Base for storage-bound cold-page quantization implementations.""" uses_iteration_lifecycle = False provides_cold_page_codec = True - def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: - super().__init__(config) - from .nvfp4 import Nvfp4ColdPageQuantization - - methods = {"nvfp4": Nvfp4ColdPageQuantization} - try: - method = methods[config.quant] - except KeyError as error: - raise NotImplementedError( - f"Unsupported cold-page quantization format {config.quant!r}" - ) from error - self._method: ColdPageQuantizationMethod = method(config.scale_checkpoint_path) - + @abstractmethod def create_cold_page_codec( self, cache_config: object, @@ -92,13 +25,4 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create the native codec selected by the quantization method.""" - - return self._method.create_cold_page_codec( - cache_config, - runtime_dtype=runtime_dtype, - pp_layers=pp_layers, - num_kv_heads_per_layer=num_kv_heads_per_layer, - head_dim_per_layer=head_dim_per_layer, - is_draft=is_draft, - ) + """Create a codec using the implementation's immutable calibration.""" diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 2806ebbd8887..9e5d7e56d61a 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -2847,10 +2847,13 @@ def create_kv_cache_compression_manager( ) -> Optional[KVCacheCompressionManager]: """Construct the configured compression manager before KVCM.""" if config.algorithm == "quantization_for_cold_page": - from ..kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import \ - ColdPageQuantizationCompression # noqa: E501 + if config.quant == "nvfp4": + from ..kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import \ + Nvfp4ColdPageQuantizationCompression # noqa: E501 - return ColdPageQuantizationCompression(config) + return Nvfp4ColdPageQuantizationCompression(config) + raise NotImplementedError( + f"Unsupported cold-page quantization format {config.quant!r}") if config.algorithm == "triattention": # TriAttention imports CuTe/CUTLASS; keep normal executor startup lazy. diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 2ec1bef82553..37212522c122 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -10,7 +10,8 @@ import torch from safetensors.torch import save_file -from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.nvfp4 import ( +from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import ( + Nvfp4ColdPageQuantizationCompression, _load_modelopt_nvfp4_scales, ) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import ( @@ -39,7 +40,7 @@ def _manager(scale_checkpoint_path=None): str(scale_checkpoint_path) if scale_checkpoint_path is not None else None ) ) - return ColdPageQuantizationCompression(config) + return Nvfp4ColdPageQuantizationCompression(config) def _cache_config(*layers): @@ -62,6 +63,8 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: ColdPageLifecycleProperties=lambda: SimpleNamespace(), ColdPageIndexLocation=SimpleNamespace(HOST="host"), create_python_cold_page_codec=MagicMock(return_value=codec), + invoke_nvfp4_cold_page_encode=MagicMock(), + invoke_nvfp4_cold_page_decode=MagicMock(), ) return module, codec @@ -240,7 +243,6 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): 1.0, ] metadata = _configure_policy(native, raw_bytes=5120) - assert metadata.runtime_type == 0 assert metadata.cold_page_bytes == 2880 assert metadata.wide[:2, 3].tolist() == [0, 1280] assert metadata.wide[:2, 4].tolist() == [2560, 2720] @@ -301,22 +303,8 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): assert native.create_python_cold_page_codec.call_count == 2 -def test_policy_forwards_a_4096_page_batch_through_one_custom_op(monkeypatch) -> None: +def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: native, _ = _native() - encode = MagicMock() - decode = MagicMock() - monkeypatch.setattr( - torch.ops.trtllm, - "nvfp4_cold_page_encode", - encode, - raising=False, - ) - monkeypatch.setattr( - torch.ops.trtllm, - "nvfp4_cold_page_decode", - decode, - raising=False, - ) with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( @@ -336,20 +324,31 @@ def test_policy_forwards_a_4096_page_batch_through_one_custom_op(monkeypatch) -> for index, role in enumerate(("key", "value")) } properties = policy.configure([SimpleNamespace(layers={0: hot})]) + policy.encode(0, 0x3000, 0x4000, 4096, 0x5000) + policy.decode(0, 0x3000, 0x4000, 4096, 0x5000) assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" - - policy.encode(0, 0x3000, 0x4000, 4096, 0x5000) - policy.decode(0, 0x3000, 0x4000, 4096, 0x5000) metadata = policy._lifecycle_metadata[0] - for operation in (encode, decode): + for operation in ( + native.invoke_nvfp4_cold_page_encode, + native.invoke_nvfp4_cold_page_decode, + ): operation.assert_called_once() arguments = operation.call_args.args - assert arguments[0] is metadata.wide - assert arguments[1] is metadata.integers - assert arguments[2] is metadata.scales - assert arguments[3:] == (2, 128, 1152, 1, 0x3000, 0x4000, 4096, 0x5000) + assert arguments == ( + 0x4000, + 4096, + metadata.wide.data_ptr(), + metadata.integers.data_ptr(), + metadata.scales.data_ptr(), + 2, + 128, + 1152, + 1, + 0x3000, + 0x5000, + ) def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: @@ -365,9 +364,15 @@ def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: with torch.device("meta"): metadata = _configure_policy(native, raw_bytes=2048) - assert metadata.wide.device.type == "cpu" - assert metadata.integers.device.type == "cpu" - assert metadata.scales.device.type == "cpu" + for tensor, dtype, shape in ( + (metadata.wide, torch.int64, (256, 6)), + (metadata.integers, torch.int32, (256, 5)), + (metadata.scales, torch.float32, (256, 4)), + ): + assert tensor.device.type == "cpu" + assert tensor.dtype == dtype + assert tensor.shape == shape + assert tensor.is_contiguous() def test_policy_rejects_invalid_resolved_hot_buffers() -> None: @@ -431,18 +436,14 @@ def test_policy_rejects_more_than_256_lifecycle_buffers() -> None: ) -def test_unsupported_quant_does_not_construct_nvfp4_policy() -> None: +def test_unsupported_quant_is_rejected_before_manager_construction() -> None: config = SimpleNamespace( + algorithm="quantization_for_cold_page", quant="future-format", scale_checkpoint_path="/not/a/checkpoint", ) - with patch( - "tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page." - "nvfp4.Nvfp4ColdPageQuantization" - ) as policy: - with pytest.raises(NotImplementedError, match="future-format"): - ColdPageQuantizationCompression(config) - policy.assert_not_called() + with pytest.raises(NotImplementedError, match="future-format"): + util_mod.create_kv_cache_compression_manager(config) def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): @@ -833,6 +834,7 @@ def build(*, estimating=False, active_kv_quant=None): resources = build() manager = resources[util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] + assert isinstance(manager, Nvfp4ColdPageQuantizationCompression) assert isinstance(manager, ColdPageQuantizationCompression) assert manager.provides_cold_page_codec assert not manager.uses_iteration_lifecycle diff --git a/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py b/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py deleted file mode 100644 index 9b0ea56ec293..000000000000 --- a/tests/unittest/_torch/thop/parallel_hw_agnostic/test_nvfp4_cold_page_op.py +++ /dev/null @@ -1,48 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""CPU-only ABI checks for NVFP4 cold-page custom ops.""" - -import pytest -import torch - -pytestmark = pytest.mark.cpu_only - - -def _encode(wide: torch.Tensor, integers: torch.Tensor, scales: torch.Tensor) -> None: - torch.ops.trtllm.nvfp4_cold_page_encode( - wide, - integers, - scales, - 1, - 2, - 16, - 0, - 0, - 0, - 0, - 0, - ) - - -def test_rejects_malformed_metadata_tensors() -> None: - wide = torch.zeros((256, 6), dtype=torch.int64) - integers = torch.zeros((256, 5), dtype=torch.int32) - scales = torch.zeros((256, 4), dtype=torch.float32) - - invalid = ( - (wide.to(torch.int32), integers, scales, "wide metadata"), - (torch.zeros((256, 12), dtype=torch.int64)[:, ::2], integers, scales, "wide metadata"), - (wide, integers[:, :4], scales, "integer metadata"), - (wide, integers, scales.to(torch.float64), "scale metadata"), - ) - for bad_wide, bad_integers, bad_scales, message in invalid: - with pytest.raises(RuntimeError, match=message): - _encode(bad_wide, bad_integers, bad_scales) - - -def test_rejects_invalid_launch_scalars() -> None: - wide = torch.zeros((256, 6), dtype=torch.int64) - integers = torch.zeros((256, 5), dtype=torch.int32) - scales = torch.zeros((256, 4), dtype=torch.float32) - with pytest.raises(RuntimeError, match="buffer count"): - torch.ops.trtllm.nvfp4_cold_page_encode(wide, integers, scales, 0, 2, 16, 0, 0, 0, 0, 0) From 37509a7a3132962a03b05d337111430e063a40fd Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 15:38:26 -0700 Subject: [PATCH 16/29] [None][refactor] Dispatch cold-page methods through Python Signed-off-by: tianruih --- cpp/tensorrt_llm/kernels/CMakeLists.txt | 3 + .../kernels/nvfp4ColdPageKernels.cu | 57 ++++++++ .../nanobind/kvCacheCompression/bindings.cpp | 32 ----- cpp/tensorrt_llm/thop/CMakeLists.txt | 5 +- .../nvfp4_quantization.py | 16 +-- .../quantization_for_cold_page.py | 119 ++++++++++++++- tensorrt_llm/_torch/pyexecutor/_util.py | 110 +++++--------- .../test_kv_cache_compression_manager.py | 135 ++++++++++++++---- .../test_quantization_for_cold_page.py | 122 +++++++++++----- .../test_triattention_pipeline.py | 24 ++-- 10 files changed, 432 insertions(+), 191 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/CMakeLists.txt b/cpp/tensorrt_llm/kernels/CMakeLists.txt index c6c6fe4a292a..34bc62fcfa31 100644 --- a/cpp/tensorrt_llm/kernels/CMakeLists.txt +++ b/cpp/tensorrt_llm/kernels/CMakeLists.txt @@ -67,6 +67,9 @@ list(FILTER SRC_CPP EXCLUDE REGEX "mhcKernels/.*") list(FILTER SRC_CU EXCLUDE REGEX "mhcKernels/.*") list(FILTER SRC_CPP EXCLUDE REGEX "compressorKernels/.*") list(FILTER SRC_CU EXCLUDE REGEX "compressorKernels/.*") +# This TU registers its Torch op in th_common; exclude it here to avoid +# duplicate symbols. +list(FILTER SRC_CU EXCLUDE REGEX "nvfp4ColdPageKernels\\.cu$") # Marlin is built as its own architecture-scoped OBJECT library below. list(FILTER SRC_CPP EXCLUDE REGEX "marlin/.*") list(FILTER SRC_CU EXCLUDE REGEX "marlin/.*") diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index 57d2d3b9433f..90a4a28fe182 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -33,6 +33,10 @@ #include #include +#if defined(TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) +#include +#endif + TRTLLM_NAMESPACE_BEGIN namespace kernels @@ -871,3 +875,56 @@ void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, std::int } // namespace kernels TRTLLM_NAMESPACE_END + +#if defined(TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) +namespace +{ + +void nvfp4ColdPageEncode(std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, + std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, + std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) +{ + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + reinterpret_cast(static_cast(pageIndices)), static_cast(numPages), + reinterpret_cast(static_cast(wide)), + reinterpret_cast(static_cast(integers)), + reinterpret_cast(static_cast(scales)), static_cast(numBuffers), + static_cast(maxHalfGroupsPerTile), static_cast(coldPageBytes), + static_cast(runtimeType), + reinterpret_cast(static_cast(coldBase)), + reinterpret_cast(static_cast(stream))); +} + +void nvfp4ColdPageDecode(std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, + std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, + std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) +{ + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + reinterpret_cast(static_cast(pageIndices)), static_cast(numPages), + reinterpret_cast(static_cast(wide)), + reinterpret_cast(static_cast(integers)), + reinterpret_cast(static_cast(scales)), static_cast(numBuffers), + static_cast(maxHalfGroupsPerTile), static_cast(coldPageBytes), + static_cast(runtimeType), + reinterpret_cast(static_cast(coldBase)), + reinterpret_cast(static_cast(stream))); +} + +} // namespace + +TORCH_LIBRARY_FRAGMENT(trtllm, module) +{ + module.def( + "nvfp4_cold_page_encode(int page_indices, int num_pages, int wide, int integers, int scales, int num_buffers, " + "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int stream) -> ()"); + module.def( + "nvfp4_cold_page_decode(int page_indices, int num_pages, int wide, int integers, int scales, int num_buffers, " + "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int stream) -> ()"); +} + +TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, module) +{ + module.impl("nvfp4_cold_page_encode", &nvfp4ColdPageEncode); + module.impl("nvfp4_cold_page_decode", &nvfp4ColdPageDecode); +} +#endif diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 724f86ccc7f9..4332203d80d3 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,7 +17,6 @@ */ #include "bindings.h" -#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" #include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" #include @@ -34,7 +33,6 @@ namespace nb = nanobind; namespace compression = tensorrt_llm::kv_cache_compression; -namespace kernels = tensorrt_llm::kernels; namespace kv = tensorrt_llm::batch_manager::kv_cache_manager_v2; namespace tensorrt_llm::nanobind::kv_cache_compression @@ -160,36 +158,6 @@ void initBindings(nb::module_& module) [](nb::handle policy) -> std::unique_ptr { return std::make_unique(policy); }, nb::arg("policy")); - module.def( - "invoke_nvfp4_cold_page_encode", - [](std::uintptr_t pageIndices, std::size_t numPages, std::uintptr_t wide, std::uintptr_t integers, - std::uintptr_t scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, - std::size_t coldPageBytes, int runtimeType, std::uintptr_t coldBase, std::uintptr_t stream) - { - kernels::invokeNvfp4ColdPageEncode(reinterpret_cast(pageIndices), numPages, - reinterpret_cast(wide), reinterpret_cast(integers), - reinterpret_cast(scales), numBuffers, maxHalfGroupsPerTile, coldPageBytes, - static_cast(runtimeType), reinterpret_cast(coldBase), - reinterpret_cast(stream)); - }, - nb::arg("page_indices"), nb::arg("num_pages"), nb::arg("wide"), nb::arg("integers"), nb::arg("scales"), - nb::arg("num_buffers"), nb::arg("max_half_groups_per_tile"), nb::arg("cold_page_bytes"), - nb::arg("runtime_type"), nb::arg("cold_base"), nb::arg("stream"), nb::call_guard()); - module.def( - "invoke_nvfp4_cold_page_decode", - [](std::uintptr_t pageIndices, std::size_t numPages, std::uintptr_t wide, std::uintptr_t integers, - std::uintptr_t scales, std::uint32_t numBuffers, std::uint32_t maxHalfGroupsPerTile, - std::size_t coldPageBytes, int runtimeType, std::uintptr_t coldBase, std::uintptr_t stream) - { - kernels::invokeNvfp4ColdPageDecode(reinterpret_cast(pageIndices), numPages, - reinterpret_cast(wide), reinterpret_cast(integers), - reinterpret_cast(scales), numBuffers, maxHalfGroupsPerTile, coldPageBytes, - static_cast(runtimeType), reinterpret_cast(coldBase), - reinterpret_cast(stream)); - }, - nb::arg("page_indices"), nb::arg("num_pages"), nb::arg("wide"), nb::arg("integers"), nb::arg("scales"), - nb::arg("num_buffers"), nb::arg("max_half_groups_per_tile"), nb::arg("cold_page_bytes"), - nb::arg("runtime_type"), nb::arg("cold_base"), nb::arg("stream"), nb::call_guard()); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index cdf0a0228fac..6b427970593c 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -146,8 +146,11 @@ add_library( trtllmGenQKVProcessOp.cpp inplaceSliceCopyOp.cpp mhcOp.cpp - compressorOp.cpp) + compressorOp.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu) set_property(TARGET th_common PROPERTY POSITION_INDEPENDENT_CODE ON) +target_compile_definitions(th_common + PRIVATE TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) target_link_libraries( th_common PRIVATE ${TORCH_LIBRARIES} th_utils ${Python3_LIBRARIES} ${SHARED_TARGET} pg_utils) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index d143e9d1650f..a33ebe838154 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -300,10 +300,8 @@ def encode( num_pages: int, stream: int, ) -> None: - from tensorrt_llm.bindings.internal import kv_cache_compression as native - metadata = self._lifecycle_metadata[lifecycle_index] - native.invoke_nvfp4_cold_page_encode( + torch.ops.trtllm.nvfp4_cold_page_encode( page_indices, num_pages, metadata.wide.data_ptr(), @@ -325,10 +323,8 @@ def decode( num_pages: int, stream: int, ) -> None: - from tensorrt_llm.bindings.internal import kv_cache_compression as native - metadata = self._lifecycle_metadata[lifecycle_index] - native.invoke_nvfp4_cold_page_decode( + torch.ops.trtllm.nvfp4_cold_page_decode( page_indices, num_pages, metadata.wide.data_ptr(), @@ -350,7 +346,7 @@ def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) self._model_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) - def create_cold_page_codec( + def _create_cold_page_policy( self, cache_config: object, *, @@ -360,7 +356,6 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - from tensorrt_llm.bindings.internal import kv_cache_compression as native from tensorrt_llm.runtime.kv_cache_manager_v2 import AttentionLayerConfig runtime_type = { @@ -423,7 +418,4 @@ def create_cold_page_codec( ) ) - policy = _Nvfp4ColdPagePolicy( - layer_layouts, runtime_type if runtime_type is not None else 0 - ) - return native.create_python_cold_page_codec(policy) + return _Nvfp4ColdPagePolicy(layer_layouts, runtime_type if runtime_type is not None else 0) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 84bc1a7b5e0f..20b353f1ffeb 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -1,20 +1,104 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Base class for KVCM V2 cold-page quantization.""" +"""Common construction pipeline for cold-page quantization.""" from abc import ABC, abstractmethod -from typing import Sequence +from typing import TYPE_CHECKING, Optional, Sequence + +from tensorrt_llm._utils import is_sm_100f +from tensorrt_llm.logger import logger +from tensorrt_llm.quantization import QuantAlgo from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager +if TYPE_CHECKING: + from tensorrt_llm._torch.model_config import ModelConfig + from tensorrt_llm.llmapi.llm_args import ( + ColdPageQuantizationCompressionConfig, + KvCacheConfig, + SpeculativeConfig, + ) + +_NATIVE_KV_CACHE_EQUIVALENTS = {"nvfp4": QuantAlgo.NVFP4} + + +def validate_cold_page_quantization_compatibility( + config: "ColdPageQuantizationCompressionConfig", + kv_cache_config: "KvCacheConfig", + spec_config: Optional["SpeculativeConfig"], +) -> None: + """Validate a cold-page method that will participate in this executor.""" + + from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND + + if _BACKEND == "python": + raise ValueError("Cold-page quantization requires the C++ KVCacheManagerV2 backend") + if kv_cache_config.enable_block_reuse and not config.supports_block_reuse(): + raise ValueError( + f"KV-cache compression algorithm {config.algorithm!r} does not " + "support KV-cache block reuse. Set " + "KvCacheConfig.enable_block_reuse=False." + ) + if spec_config is None: + return + if not config.supports_speculative_decoding(): + raise ValueError( + f"KV-cache compression algorithm {config.algorithm!r} does not " + "support speculative decoding with its current configuration" + ) + mode = spec_config.spec_dec_mode + if not (mode.is_eagle3_one_model() or mode.is_mtp_eagle_one_model()): + raise ValueError( + "Cold-page quantization supports speculative decoding only " + f"with one-model MTP-EAGLE or EAGLE3, not {mode.name}" + ) + + +def create_cold_page_quantization_manager( + config: "ColdPageQuantizationCompressionConfig", + *, + model_config: "ModelConfig", + kv_cache_config: "KvCacheConfig", + spec_config: Optional["SpeculativeConfig"], + estimating_kv_cache: bool = False, +) -> Optional["ColdPageQuantizationCompression"]: + """Select, validate, and construct one cold-page quantization method.""" + + active_quant_algo = _NATIVE_KV_CACHE_EQUIVALENTS.get(config.quant) + if active_quant_algo is None: + raise NotImplementedError(f"Unsupported cold-page quantization format {config.quant!r}") + if estimating_kv_cache: + return None + + quant_config = model_config.quant_config + if ( + quant_config is not None + and getattr(quant_config, "kv_cache_quant_algo", None) == active_quant_algo + ): + logger.info( + "Skipping cold-page %s quantization because the active KV cache " + "already uses the same format; KVCM will migrate it losslessly.", + config.quant.upper(), + ) + return None + + validate_cold_page_quantization_compatibility(config, kv_cache_config, spec_config) + + from .nvfp4_quantization import Nvfp4ColdPageQuantizationCompression + + if not is_sm_100f(): + raise RuntimeError( + "NVFP4 cold-page quantization requires an SM100-family device (SM100 or SM103)." + ) + return Nvfp4ColdPageQuantizationCompression(config) + class ColdPageQuantizationCompression(KVCacheCompressionManager, ABC): - """Base for storage-bound cold-page quantization implementations.""" + """Base pipeline shared by storage-bound cold-page quantizers.""" uses_iteration_lifecycle = False provides_cold_page_codec = True - @abstractmethod def create_cold_page_codec( self, cache_config: object, @@ -25,4 +109,29 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create a codec using the implementation's immutable calibration.""" + """Create one generic native codec around a format-specific policy.""" + + from tensorrt_llm.bindings.internal import kv_cache_compression as native + + policy = self._create_cold_page_policy( + cache_config, + runtime_dtype=runtime_dtype, + pp_layers=pp_layers, + num_kv_heads_per_layer=num_kv_heads_per_layer, + head_dim_per_layer=head_dim_per_layer, + is_draft=is_draft, + ) + return native.create_python_cold_page_codec(policy) + + @abstractmethod + def _create_cold_page_policy( + self, + cache_config: object, + *, + runtime_dtype: DataType, + pp_layers: Sequence[int], + num_kv_heads_per_layer: Sequence[int], + head_dim_per_layer: Sequence[int], + is_draft: bool = False, + ) -> object: + """Build the format-specific layout and callback policy.""" diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 9e5d7e56d61a..db351e43dc2f 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -15,7 +15,7 @@ import copy import dataclasses import os -from typing import TYPE_CHECKING, Dict, List, Optional, Union +from typing import Dict, List, Optional, Union import torch @@ -41,7 +41,6 @@ supports_native_fp8_lora) from tensorrt_llm.logger import logger from tensorrt_llm.mapping import CpType, Mapping -from tensorrt_llm.quantization import QuantAlgo from ..attention_backend import get_sparse_attn_kv_cache_manager from ..hostfunc import set_low_latency_dispatch @@ -81,9 +80,6 @@ SimpleUnifiedScheduler) from .seq_slot_manager import SeqSlotManager -if TYPE_CHECKING: - import transformers - GB = 1 << 30 @@ -2055,19 +2051,13 @@ def build_managers(self, budget_attr, self_kv_cache_config, draft_kv_cache_config)) - compression_manager = None compression_config = self._llm_args.kv_cache_compression_config - is_cold_quantization = (compression_config is not None - and compression_config.algorithm - == "quantization_for_cold_page") - skip_compression_manager = is_cold_quantization and ( - (estimating_kv_cache and not self._skip_est) - or _uses_nvfp4_kv_cache(self._model_engine)) - if compression_config is not None and not skip_compression_manager: - model_config = self._model_engine.model.model_config - compression_manager = create_kv_cache_compression_manager( - compression_config, - pretrained_config=model_config.pretrained_config) + compression_manager = create_kv_cache_compression_manager( + compression_config, + model_engine=self._model_engine, + kv_cache_config=self_kv_cache_config, + estimating_kv_cache=estimating_kv_cache and not self._skip_est, + ) cold_page_codec_provider = ( compression_manager if compression_manager is not None and compression_manager.provides_cold_page_codec else None) @@ -2130,7 +2120,8 @@ def build_managers(self, ResourceManagerType.DRAFT_KV_CACHE_MANAGER] = draft_kv_cache_manager resources[ ResourceManagerType.CROSS_KV_CACHE_MANAGER] = cross_kv_cache_manager - if compression_manager is not None: + if (compression_manager is not None + and compression_manager.uses_iteration_lifecycle): resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] = ( compression_manager) @@ -2795,18 +2786,7 @@ def validate_kv_cache_compression_compatibility( spec_config: Optional[SpeculativeConfig], ) -> None: """Reject unsupported KV-cache compression feature combinations.""" - if config.algorithm == "quantization_for_cold_page": - from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND - - if _BACKEND == "python": - raise ValueError( - "Cold-page quantization requires the C++ KVCacheManagerV2 backend" - ) - if not is_sm_100f(): - raise RuntimeError( - "NVFP4 cold-page compression requires an SM100-family device " - "(SM100 or SM103).") - elif config.algorithm == "triattention" and not is_sm_100f(): + if config.algorithm == "triattention" and not is_sm_100f(): raise RuntimeError( "TriAttention requires an SM100-family device (SM100 or SM103).") @@ -2823,46 +2803,49 @@ def validate_kv_cache_compression_compatibility( "support speculative decoding with its current configuration; " "TriAttention requires eviction_mode='union'") mode = spec_config.spec_dec_mode - if config.algorithm == "quantization_for_cold_page": - if not (mode.is_eagle3_one_model() or mode.is_mtp_eagle_one_model()): - raise ValueError( - "Cold-page quantization supports speculative decoding only " - f"with one-model MTP-EAGLE or EAGLE3, not {mode.name}") - return if not (mode.is_mtp_one_model() or mode.is_eagle3_one_model()): raise ValueError( f"KV-cache compression does not support speculative decoding " f"mode {mode.name}; use one-model MTP or EAGLE3") -def _uses_nvfp4_kv_cache(model_engine: PyTorchModelEngine) -> bool: - quant_config = model_engine.model.model_config.quant_config - return (quant_config is not None - and quant_config.kv_cache_quant_algo == QuantAlgo.NVFP4) - - def create_kv_cache_compression_manager( - config: KvCacheCompressionConfig, - pretrained_config: Optional["transformers.PretrainedConfig"] = None, + config: Optional[KvCacheCompressionConfig], + *, + model_engine: PyTorchModelEngine, + kv_cache_config: KvCacheConfig, + estimating_kv_cache: bool = False, ) -> Optional[KVCacheCompressionManager]: - """Construct the configured compression manager before KVCM.""" + """Validate, select, and construct the configured manager before KVCM.""" + if config is None: + return None + if model_engine.mapping.has_cp_helix(): + # TODO: Revisit after KVCC validates HELIX-sharded Page ownership and migration. + raise ValueError( + "KV-cache compression does not support HELIX context parallelism.") + if config.algorithm == "quantization_for_cold_page": - if config.quant == "nvfp4": - from ..kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import \ - Nvfp4ColdPageQuantizationCompression # noqa: E501 + from ..kv_cache_compression.quantization_for_cold_page import \ + quantization_for_cold_page as cold_page_quantization - return Nvfp4ColdPageQuantizationCompression(config) - raise NotImplementedError( - f"Unsupported cold-page quantization format {config.quant!r}") + return cold_page_quantization.create_cold_page_quantization_manager( + config, + model_config=model_engine.model.model_config, + kv_cache_config=kv_cache_config, + spec_config=model_engine.spec_config, + estimating_kv_cache=estimating_kv_cache, + ) if config.algorithm == "triattention": + validate_kv_cache_compression_compatibility(config, kv_cache_config, + model_engine.spec_config) # TriAttention imports CuTe/CUTLASS; keep normal executor startup lazy. from ..kv_cache_compression.triattention.triattention import \ TriAttentionCompressionManager return TriAttentionCompressionManager( config, - pretrained_config=pretrained_config, + pretrained_config=model_engine.model.model_config.pretrained_config, ) logger.warning( @@ -3139,8 +3122,6 @@ def create_py_executor_instance( resources[ResourceManagerType.KV_CACHE_MANAGER], resources.get(ResourceManagerType.DRAFT_KV_CACHE_MANAGER), ) - if not compression_manager.uses_iteration_lifecycle: - resources.pop(ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER) resource_manager = ResourceManager(resources) # KV cache manager runs last (others may depend on it), except the @@ -3639,27 +3620,6 @@ def _adjust_torch_mem_fraction(): def validate_feature_combination(llm_args, model_engine, sampler_type): # Validate the flags for features' combination - compression_config = llm_args.kv_cache_compression_config - if (compression_config is not None and model_engine.mapping.has_cp_helix()): - # TODO: Revisit after KVCC validates HELIX-sharded Page ownership and migration. - raise ValueError( - "KV-cache compression does not support HELIX context parallelism.") - cold_compression_is_redundant = (compression_config is not None - and compression_config.algorithm - == "quantization_for_cold_page" - and _uses_nvfp4_kv_cache(model_engine)) - if cold_compression_is_redundant: - logger.info( - "Skipping cold-page NVFP4 quantization because the active KV cache " - "already uses NVFP4; KVCM will migrate its native data and " - "block-scale buffers losslessly.") - elif compression_config is not None: - validate_kv_cache_compression_compatibility( - compression_config, - llm_args.kv_cache_config, - model_engine.spec_config, - ) - def init_feature_status(llm_args) -> Dict[str, bool]: assert isinstance( llm_args, TorchLlmArgs diff --git a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py index df6f2abfea9d..644fe34c0508 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py +++ b/tests/unittest/_torch/executor/test_kv_cache_compression_manager.py @@ -72,6 +72,25 @@ def _compression_config() -> KvCacheCompressionConfig: return KvCacheCompressionConfig(algorithm="test") +def _factory_model_engine( + *, + pretrained_config: object | None = None, + quant_config: object | None = None, + spec_config: object | None = None, + helix: bool = False, +) -> SimpleNamespace: + return SimpleNamespace( + mapping=SimpleNamespace(has_cp_helix=lambda: helix), + spec_config=spec_config, + model=SimpleNamespace( + model_config=SimpleNamespace( + pretrained_config=pretrained_config, + quant_config=quant_config, + ) + ), + ) + + def _v2_manager(*, is_draft: bool): from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 @@ -269,16 +288,27 @@ def test_free_fires_finish(self, fake_kv_cache_manager): class TestFactory: - def test_returns_none_when_no_algorithm_registered(self): + def test_returns_none_when_no_algorithm_registered(self) -> None: cfg = MagicMock() cfg.algorithm = "made_up_method" - assert create_kv_cache_compression_manager(cfg) is None + assert ( + create_kv_cache_compression_manager( + cfg, + model_engine=_factory_model_engine(), + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) + is None + ) - def test_warns_for_unregistered_algorithm(self): + def test_warns_for_unregistered_algorithm(self) -> None: cfg = MagicMock() cfg.algorithm = "made_up_method" with patch.object(util_mod, "logger") as mock_logger: - create_kv_cache_compression_manager(cfg) + create_kv_cache_compression_manager( + cfg, + model_engine=_factory_model_engine(), + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) mock_logger.warning.assert_called_once() def test_triattention_requires_sm100_family(self): @@ -330,15 +360,16 @@ class TestKvCacheCreatorLifecycle: def test_estimation_still_creates_triattention_manager(self) -> None: config = SimpleNamespace(algorithm="triattention") pretrained_config = object() - expected_manager = SimpleNamespace(provides_cold_page_codec=False) + expected_manager = MagicMock( + provides_cold_page_codec=False, + uses_iteration_lifecycle=True, + ) creator = object.__new__(util_mod.KvCacheCreator) creator._skip_est = False creator._max_seq_len = 1024 creator._kv_cache_config = SimpleNamespace(host_cache_size=None, disk_cache_size=None) creator._llm_args = SimpleNamespace(kv_cache_compression_config=config) - creator._model_engine = SimpleNamespace( - model=SimpleNamespace(model_config=SimpleNamespace(pretrained_config=pretrained_config)) - ) + creator._model_engine = _factory_model_engine(pretrained_config=pretrained_config) creator._draft_model_engine = None creator._is_encoder_decoder = MagicMock(return_value=False) creator._should_create_separate_draft_kv_cache = MagicMock(return_value=False) @@ -355,10 +386,13 @@ def test_estimation_still_creates_triattention_manager(self) -> None: factory.assert_called_once_with( config, - pretrained_config=pretrained_config, + model_engine=creator._model_engine, + kv_cache_config=creator._kv_cache_config, + estimating_kv_cache=True, ) assert resources[ResourceManagerType.KV_CACHE_MANAGER] is target_manager assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is expected_manager + expected_manager.bind_kv_cache_managers.assert_not_called() def test_teardown_pops_and_shuts_down_compression_manager(self) -> None: creator = object.__new__(util_mod.KvCacheCreator) @@ -375,6 +409,49 @@ def test_teardown_pops_and_shuts_down_compression_manager(self) -> None: assert ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in resources +@pytest.mark.cpu_only +def test_executor_binds_iteration_manager_after_extra_resource_registration() -> None: + class _StopAfterBind(Exception): + pass + + target_manager = object() + draft_manager = object() + compression_manager = MagicMock() + resources = { + ResourceManagerType.KV_CACHE_MANAGER: target_manager, + ResourceManagerType.DRAFT_KV_CACHE_MANAGER: draft_manager, + } + llm_args = SimpleNamespace( + enable_low_latency_host_dispatch=False, + extra_resource_managers={ + ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER: compression_manager, + }, + ) + model_engine = SimpleNamespace(spec_config=None) + + with ( + patch.object(util_mod, "set_low_latency_dispatch"), + patch.object(util_mod, "ResourceManager", side_effect=_StopAfterBind), + pytest.raises(_StopAfterBind), + ): + util_mod.create_py_executor_instance( + dist=None, + resources=resources, + mapping=SimpleNamespace(), + llm_args=llm_args, + ctx_chunk_config=None, + model_engine=model_engine, + start_worker=False, + sampler=None, + drafter=None, + max_num_sequences=1, + ) + + compression_manager.bind_kv_cache_managers.assert_called_once_with( + target_manager, draft_manager + ) + + @pytest.mark.cpu_only @pytest.mark.parametrize( "provides_cold_page_codec", @@ -391,9 +468,7 @@ def test_build_routes_compression_manager_by_capabilities( compression_config = SimpleNamespace(algorithm="triattention") pretrained_config = object() creator._llm_args = SimpleNamespace(kv_cache_compression_config=compression_config) - creator._model_engine = SimpleNamespace( - model=SimpleNamespace(model_config=SimpleNamespace(pretrained_config=pretrained_config)) - ) + creator._model_engine = _factory_model_engine(pretrained_config=pretrained_config) creator._draft_model_engine = None creator._kv_connector_manager = None creator._is_kv_cache_manager_v2 = True @@ -405,6 +480,8 @@ def test_build_routes_compression_manager_by_capabilities( build_order = [] compression_manager = SimpleNamespace( provides_cold_page_codec=provides_cold_page_codec, + uses_iteration_lifecycle=not provides_cold_page_codec, + bind_kv_cache_managers=MagicMock(side_effect=lambda *_args: build_order.append("bind")), ) target_config = object() draft_config = object() @@ -431,7 +508,9 @@ def test_build_routes_compression_manager_by_capabilities( expected_codec_provider = compression_manager if provides_cold_page_codec else None factory.assert_called_once_with( compression_config, - pretrained_config=pretrained_config, + model_engine=creator._model_engine, + kv_cache_config=target_config, + estimating_kv_cache=False, ) assert ( creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"] @@ -445,8 +524,14 @@ def test_build_routes_compression_manager_by_capabilities( ) assert resources[ResourceManagerType.KV_CACHE_MANAGER] is target_manager assert resources[ResourceManagerType.DRAFT_KV_CACHE_MANAGER] is draft_manager - assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is compression_manager - assert build_order == ["factory", "target", "draft"] + if provides_cold_page_codec: + assert ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in resources + compression_manager.bind_kv_cache_managers.assert_not_called() + assert build_order == ["factory", "target", "draft"] + else: + assert resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] is compression_manager + compression_manager.bind_kv_cache_managers.assert_not_called() + assert build_order == ["factory", "target", "draft"] # ---------------------------------------------------------------------- # @@ -495,12 +580,12 @@ def test_helix_is_rejected( spec_config=None, model=SimpleNamespace(model_config=SimpleNamespace(quant_config=None)), ) - llm_args = SimpleNamespace( - kv_cache_compression_config=config, - kv_cache_config=SimpleNamespace(enable_block_reuse=False), - ) with pytest.raises(ValueError, match="HELIX"): - util_mod.validate_feature_combination(llm_args, model_engine, None) + create_kv_cache_compression_manager( + config, + model_engine=model_engine, + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) @pytest.mark.cpu_only def test_helix_is_rejected_before_redundant_cold_quantization(self) -> None: @@ -513,12 +598,12 @@ def test_helix_is_rejected_before_redundant_cold_quantization(self) -> None: ) ), ) - llm_args = SimpleNamespace( - kv_cache_compression_config=ColdPageQuantizationCompressionConfig(), - kv_cache_config=SimpleNamespace(enable_block_reuse=False), - ) with pytest.raises(ValueError, match="HELIX"): - util_mod.validate_feature_combination(llm_args, model_engine, None) + create_kv_cache_compression_manager( + ColdPageQuantizationCompressionConfig(), + model_engine=model_engine, + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) def test_raises_when_reuse_on(self): config = _compression_config() diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 37212522c122..7e7e96a96f38 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -10,12 +10,16 @@ import torch from safetensors.torch import save_file +from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page import ( + quantization_for_cold_page as cold_quant_mod, +) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import ( Nvfp4ColdPageQuantizationCompression, _load_modelopt_nvfp4_scales, ) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import ( ColdPageQuantizationCompression, + validate_cold_page_quantization_compatibility, ) from tensorrt_llm._torch.pyexecutor import _util as util_mod from tensorrt_llm._torch.pyexecutor.resource_manager import DataType @@ -43,6 +47,21 @@ def _manager(scale_checkpoint_path=None): return Nvfp4ColdPageQuantizationCompression(config) +def _factory_model_engine( + *, active_kv_quant: object | None = None, helix: bool = False +) -> SimpleNamespace: + return SimpleNamespace( + mapping=SimpleNamespace(has_cp_helix=lambda: helix), + spec_config=None, + model=SimpleNamespace( + model_config=SimpleNamespace( + quant_config=active_kv_quant, + pretrained_config=object(), + ) + ), + ) + + def _cache_config(*layers): configs = [] for layer_id, kind in layers: @@ -63,8 +82,6 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: ColdPageLifecycleProperties=lambda: SimpleNamespace(), ColdPageIndexLocation=SimpleNamespace(HOST="host"), create_python_cold_page_codec=MagicMock(return_value=codec), - invoke_nvfp4_cold_page_encode=MagicMock(), - invoke_nvfp4_cold_page_decode=MagicMock(), ) return module, codec @@ -116,9 +133,9 @@ def _write_scales(directory, scales_by_layer, *, filename="model.safetensors", p save_file(tensors, str(directory / filename)) -def _validate_compression(mode=None): +def _validate_compression(mode: object | None = None) -> None: spec_config = None if mode is None else SimpleNamespace(spec_dec_mode=mode) - util_mod.validate_kv_cache_compression_compatibility( + validate_cold_page_quantization_compatibility( ColdPageQuantizationCompressionConfig(), SimpleNamespace(enable_block_reuse=False), spec_config, @@ -306,7 +323,11 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: native, _ = _native() - with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): + with ( + patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native), + patch.object(torch.ops.trtllm, "nvfp4_cold_page_encode", create=True) as encode, + patch.object(torch.ops.trtllm, "nvfp4_cold_page_decode", create=True) as decode, + ): _manager().create_cold_page_codec( _cache_config((0, "attention")), runtime_dtype=DataType.BF16, @@ -330,10 +351,7 @@ def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" metadata = policy._lifecycle_metadata[0] - for operation in ( - native.invoke_nvfp4_cold_page_encode, - native.invoke_nvfp4_cold_page_decode, - ): + for operation in (encode, decode): operation.assert_called_once() arguments = operation.call_args.args assert arguments == ( @@ -443,7 +461,11 @@ def test_unsupported_quant_is_rejected_before_manager_construction() -> None: scale_checkpoint_path="/not/a/checkpoint", ) with pytest.raises(NotImplementedError, match="future-format"): - util_mod.create_kv_cache_compression_manager(config) + util_mod.create_kv_cache_compression_manager( + config, + model_engine=_factory_model_engine(), + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + ) def test_scale_loader_matches_hf_shard_and_consolidated_policy(tmp_path): @@ -762,23 +784,35 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): ) -def test_runtime_admission_is_checked_in_utils_before_manager_creation(monkeypatch): +def test_runtime_admission_is_checked_before_manager_creation(monkeypatch) -> None: monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "python") with pytest.raises(ValueError, match=r"require.*C\+\+ KVCacheManagerV2"): _validate_compression() monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") - monkeypatch.setattr(util_mod, "is_sm_100f", lambda: False) + monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: False) with pytest.raises(RuntimeError, match="requires an SM100-family device"): - _validate_compression() + cold_quant_mod.create_cold_page_quantization_manager( + ColdPageQuantizationCompressionConfig(), + model_config=_factory_model_engine().model.model_config, + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + spec_config=None, + ) - monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) - _validate_compression() + monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: True) + assert isinstance( + cold_quant_mod.create_cold_page_quantization_manager( + ColdPageQuantizationCompressionConfig(), + model_config=_factory_model_engine().model.model_config, + kv_cache_config=SimpleNamespace(enable_block_reuse=False), + spec_config=None, + ), + Nvfp4ColdPageQuantizationCompression, + ) -def test_speculative_admission_accepts_verified_one_model_modes(monkeypatch): +def test_speculative_admission_accepts_verified_one_model_modes(monkeypatch) -> None: monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") - monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) _validate_compression(SpeculativeDecodingMode.EAGLE3_ONE_MODEL) _validate_compression(SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL) @@ -792,9 +826,8 @@ def test_speculative_admission_accepts_verified_one_model_modes(monkeypatch): _validate_compression(mode) -def test_qwen35_mtp3_resolves_to_supported_one_model_mode(monkeypatch): +def test_qwen35_mtp3_resolves_to_supported_one_model_mode(monkeypatch) -> None: monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") - monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) spec_config = MTPDecodingConfig(max_draft_len=3) update_spec_config_from_model_config( @@ -804,42 +837,65 @@ def test_qwen35_mtp3_resolves_to_supported_one_model_mode(monkeypatch): assert spec_config.spec_dec_mode is SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL assert spec_config.max_draft_len == 3 - util_mod.validate_kv_cache_compression_compatibility( + validate_cold_page_quantization_compatibility( ColdPageQuantizationCompressionConfig(), SimpleNamespace(enable_block_reuse=False), spec_config, ) -def test_cold_manager_is_disabled_for_estimation_and_active_nvfp4(): - def build(*, estimating=False, active_kv_quant=None): +def test_cold_manager_is_disabled_for_estimation_and_active_nvfp4(monkeypatch) -> None: + monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") + monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: True) + + def build( + *, + estimating: bool = False, + skip_est: bool = False, + active_kv_quant: object | None = None, + ) -> tuple[dict, util_mod.KvCacheCreator]: creator = object.__new__(util_mod.KvCacheCreator) - creator._skip_est = False + creator._skip_est = skip_est creator._max_seq_len = 1024 creator._kv_cache_config = SimpleNamespace(host_cache_size=None, disk_cache_size=None) creator._llm_args = SimpleNamespace( kv_cache_compression_config=ColdPageQuantizationCompressionConfig() ) - model_config = SimpleNamespace(quant_config=active_kv_quant, pretrained_config=object()) - creator._model_engine = SimpleNamespace(model=SimpleNamespace(model_config=model_config)) + creator._model_engine = _factory_model_engine(active_kv_quant=active_kv_quant) creator._draft_model_engine = None creator._kv_connector_manager = None creator._fp8_ctx_mla_kv_len_cap = None creator._is_encoder_decoder = MagicMock(return_value=False) creator._should_create_separate_draft_kv_cache = MagicMock(return_value=False) creator._create_kv_cache_manager = MagicMock(return_value=SimpleNamespace()) + creator.configure_kv_cache_capacity = MagicMock() resources = {} creator.build_managers(resources, estimating_kv_cache=estimating) - return resources + return resources, creator - resources = build() - manager = resources[util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] + resources, creator = build() + manager = creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"] assert isinstance(manager, Nvfp4ColdPageQuantizationCompression) assert isinstance(manager, ColdPageQuantizationCompression) assert manager.provides_cold_page_codec assert not manager.uses_iteration_lifecycle - for kwargs in ( - {"estimating": True}, - {"active_kv_quant": QuantConfig(kv_cache_quant_algo=QuantAlgo.NVFP4)}, - ): - assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in build(**kwargs) + assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in resources + _, estimation_creator = build(estimating=True) + assert ( + estimation_creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"] + is None + ) + _, skip_est_creator = build(estimating=True, skip_est=True) + assert isinstance( + skip_est_creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"], + Nvfp4ColdPageQuantizationCompression, + ) + with patch.object(cold_quant_mod.logger, "info") as log: + active_resources, active_creator = build( + active_kv_quant=QuantConfig(kv_cache_quant_algo=QuantAlgo.NVFP4) + ) + assert util_mod.ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER not in active_resources + assert ( + active_creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"] is None + ) + log.assert_called_once() diff --git a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py index e2fd555a7723..0dcc4430b04b 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py +++ b/tests/unittest/_torch/kv_cache_compression/test_triattention_pipeline.py @@ -69,8 +69,21 @@ def _make_hf_config(**values): return SimpleNamespace(get_text_config=lambda: text_config) +def _factory_model_engine(pretrained_config: object) -> SimpleNamespace: + return SimpleNamespace( + mapping=SimpleNamespace(has_cp_helix=lambda: False), + spec_config=None, + model=SimpleNamespace( + model_config=SimpleNamespace( + pretrained_config=pretrained_config, + quant_config=None, + ) + ), + ) + + class TestConfigAndFactory: - def test_factory_allows_block_reuse_and_propagates_config_fields(self): + def test_factory_allows_block_reuse_and_propagates_config_fields(self) -> None: # The factory contract is independent of GPU-owned persistent buffers. fake_v2 = _make_fake_v2(enable_block_reuse=True) cfg = _make_tri_config(budget=32, beta=16, eviction_mode="per_head") @@ -83,14 +96,10 @@ def test_factory_allows_block_reuse_and_propagates_config_fields(self): TriAttentionCompressionManager, "_initialize_eviction_state" ) as initialize, ): - validate_kv_cache_compression_compatibility( - cfg, - SimpleNamespace(enable_block_reuse=True), - None, - ) mgr = create_kv_cache_compression_manager( cfg, - pretrained_config=_make_test_pretrained_config(), + model_engine=_factory_model_engine(_make_test_pretrained_config()), + kv_cache_config=SimpleNamespace(enable_block_reuse=True), ) mgr.bind_kv_cache_managers(fake_v2) assert isinstance(mgr, TriAttentionCompressionManager) @@ -447,7 +456,6 @@ def test_one_model_draft_co_compression_is_accepted(self, spec_mode): ) manager.bind_kv_cache_managers(_make_fake_v2(), draft_manager) - from tensorrt_llm._torch.pyexecutor._util import validate_kv_cache_compression_compatibility from tensorrt_llm.llmapi.llm_args import Eagle3DecodingConfig, MTPDecodingConfig spec_config = ( From 9a12ce89b8b9fc9644c4c0b40bfbb4cecb1a5cac Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 15:48:48 -0700 Subject: [PATCH 17/29] [None][refactor] Keep cold-page quantization base concrete Signed-off-by: tianruih --- .../quantization_for_cold_page/quantization_for_cold_page.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 20b353f1ffeb..27432ea5456f 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -2,7 +2,6 @@ # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. """Common construction pipeline for cold-page quantization.""" -from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Optional, Sequence from tensorrt_llm._utils import is_sm_100f @@ -93,7 +92,7 @@ def create_cold_page_quantization_manager( return Nvfp4ColdPageQuantizationCompression(config) -class ColdPageQuantizationCompression(KVCacheCompressionManager, ABC): +class ColdPageQuantizationCompression(KVCacheCompressionManager): """Base pipeline shared by storage-bound cold-page quantizers.""" uses_iteration_lifecycle = False @@ -123,7 +122,6 @@ def create_cold_page_codec( ) return native.create_python_cold_page_codec(policy) - @abstractmethod def _create_cold_page_policy( self, cache_config: object, @@ -135,3 +133,4 @@ def _create_cold_page_policy( is_draft: bool = False, ) -> object: """Build the format-specific layout and callback policy.""" + raise NotImplementedError From ff0f40639b7430972d24a42eda10eeab0d5230b4 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 16:04:11 -0700 Subject: [PATCH 18/29] [None][build] Scope cold-page Torch op definition to its source --- cpp/tensorrt_llm/thop/CMakeLists.txt | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index 6b427970593c..13432ae7e6bc 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -149,8 +149,9 @@ add_library( compressorOp.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu) set_property(TARGET th_common PROPERTY POSITION_INDEPENDENT_CODE ON) -target_compile_definitions(th_common - PRIVATE TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) +set_source_files_properties( + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu + PROPERTIES COMPILE_DEFINITIONS TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) target_link_libraries( th_common PRIVATE ${TORCH_LIBRARIES} th_utils ${Python3_LIBRARIES} ${SHARED_TARGET} pg_utils) From 3512a2ecc2bed4e01d8d8cf9e69e20e40ae66f78 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 16:40:54 -0700 Subject: [PATCH 19/29] [None][refactor] Simplify cold-page quantization callbacks Signed-off-by: tianruih --- .../nvfp4_quantization.py | 346 ++++++++---------- .../quantization_for_cold_page.py | 176 +++++---- tensorrt_llm/_torch/pyexecutor/_util.py | 54 ++- .../test_quantization_for_cold_page.py | 108 +++--- 4 files changed, 330 insertions(+), 354 deletions(-) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index a33ebe838154..dba13733bad3 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""NVFP4 quantization policy and cold-page layout construction.""" +"""NVFP4 cold-page layout, scales, metadata, and kernel dispatch.""" import json import math @@ -148,205 +148,14 @@ def _load_modelopt_nvfp4_scales( return result -class _Nvfp4ColdPagePolicy: - """Resolve NVFP4 metadata and submit one native launch per codec batch.""" - - def __init__(self, layer_layouts: Sequence[_Nvfp4LayerLayout], runtime_type: int) -> None: - self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} - self.layer_ids = tuple(sorted(self._layer_layouts)) - self._runtime_type = runtime_type - self._lifecycle_metadata: list[_Nvfp4ColdPageMetadata] = [] - - def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: - """Resolve hot buffers into immutable Python-owned launch metadata.""" - - from tensorrt_llm.bindings.internal import kv_cache_compression as native - - lifecycle_metadata = [] - properties = [] - for lifecycle in lifecycles: - wide_rows: list[list[int]] = [] - integer_rows: list[list[int]] = [] - scale_rows: list[list[float]] = [] - cold_page_bytes = 0 - max_half_groups_per_tile = 0 - - for layer_id, hot_buffers in lifecycle.layers.items(): - layout = self._layer_layouts[int(layer_id)] - expected_roles = {buffer.role for buffer in layout.buffers} - if set(hot_buffers) != expected_roles: - raise ValueError( - f"Cold-page layer {layer_id} roles do not match its KVCM layout" - ) - elements = layout.num_kv_heads * layout.tokens_per_page * layout.head_dim - element_bytes = 1 if self._runtime_type == 2 else 2 - expected_raw_bytes = elements * element_bytes - half_groups = elements // _ELEMENTS_PER_HALF_GROUP - compressed_count = sum(buffer.scales is not None for buffer in layout.buffers) - packed_bytes = elements // _ELEMENTS_PER_BYTE - scale_bytes = elements // _ELEMENTS_PER_SCALE - layer_start = cold_page_bytes - scale_start = layer_start + compressed_count * packed_bytes - cursor = scale_start + compressed_count * scale_bytes - - compressed_index = 0 - for buffer in layout.buffers: - is_compressed = buffer.scales is not None - hot = hot_buffers[buffer.role] - raw_base = int(hot.raw_base) - raw_slot_bytes = int(hot.raw_slot_bytes) - raw_bytes = int(hot.raw_bytes) - if raw_base <= 0 or raw_bytes <= 0 or raw_bytes > raw_slot_bytes: - raise ValueError("Cold-page hot buffer has invalid address or size") - - if is_compressed: - data_offset = layer_start + compressed_index * packed_bytes - scale_offset = scale_start + compressed_index * scale_bytes - compressed_index += 1 - if raw_bytes != expected_raw_bytes: - raise ValueError("Hot buffer size does not match NVFP4 geometry") - if raw_base % 16 or raw_slot_bytes % 16: - raise ValueError( - "NVFP4 hot address and Slot stride must be 16-byte aligned" - ) - max_half_groups_per_tile = max( - max_half_groups_per_tile, - min(half_groups, _MAX_HALF_GROUPS_PER_TILE), - ) - else: - data_offset = cursor - scale_offset = 0 - cursor += raw_bytes - - wide_rows.append( - [ - raw_base, - raw_slot_bytes, - raw_bytes, - data_offset, - scale_offset, - 0, - ] - ) - integer_rows.append( - [ - 0, - _NVFP4_TRANSFORM if is_compressed else _LOSSLESS_TRANSFORM, - layout.num_kv_heads if is_compressed else 0, - layout.tokens_per_page if is_compressed else 0, - layout.head_dim if is_compressed else 0, - ] - ) - buffer_scales = buffer.scales if is_compressed else _Nvfp4Scales(1.0, 1.0) - scale_rows.append( - [ - buffer_scales.nvfp4_orig_quant, - buffer_scales.nvfp4_quant_orig, - buffer_scales.fp8_orig_quant, - buffer_scales.fp8_quant_orig, - ] - ) - layer_end = ( - (cursor + _COLD_PAGE_ALIGNMENT - 1) - // _COLD_PAGE_ALIGNMENT - * _COLD_PAGE_ALIGNMENT - ) - wide_rows[-1][5] = cursor - integer_rows[-1][0] = layer_end - cursor - cold_page_bytes = layer_end - - num_buffers = len(wide_rows) - if not 0 < num_buffers <= _MAX_BUFFERS_PER_LAUNCH: - raise ValueError( - f"NVFP4 cold-page lifecycle has {num_buffers} buffers; " - f"the maximum is {_MAX_BUFFERS_PER_LAUNCH}" - ) - padding = _MAX_BUFFERS_PER_LAUNCH - num_buffers - lifecycle_metadata.append( - _Nvfp4ColdPageMetadata( - wide=torch.tensor( - wide_rows + [[0] * _WIDE_FIELDS for _ in range(padding)], - dtype=torch.int64, - device="cpu", - ), - integers=torch.tensor( - integer_rows + [[0] * _INTEGER_FIELDS for _ in range(padding)], - dtype=torch.int32, - device="cpu", - ), - scales=torch.tensor( - scale_rows + [[0.0] * _SCALE_FIELDS for _ in range(padding)], - dtype=torch.float32, - device="cpu", - ), - num_buffers=num_buffers, - max_half_groups_per_tile=max_half_groups_per_tile, - cold_page_bytes=cold_page_bytes, - ) - ) - lifecycle_properties = native.ColdPageLifecycleProperties() - lifecycle_properties.cold_page_bytes = cold_page_bytes - lifecycle_properties.page_index_location = native.ColdPageIndexLocation.HOST - properties.append(lifecycle_properties) - - self._lifecycle_metadata = lifecycle_metadata - return properties - - def encode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - metadata = self._lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_encode( - page_indices, - num_pages, - metadata.wide.data_ptr(), - metadata.integers.data_ptr(), - metadata.scales.data_ptr(), - metadata.num_buffers, - metadata.max_half_groups_per_tile, - metadata.cold_page_bytes, - self._runtime_type, - cold_base, - stream, - ) - - def decode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - metadata = self._lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_decode( - page_indices, - num_pages, - metadata.wide.data_ptr(), - metadata.integers.data_ptr(), - metadata.scales.data_ptr(), - metadata.num_buffers, - metadata.max_half_groups_per_tile, - metadata.cold_page_bytes, - self._runtime_type, - cold_base, - stream, - ) - - class Nvfp4ColdPageQuantizationCompression(ColdPageQuantizationCompression): - """NVFP4 cold-page quantization and per-KVCM codec construction.""" + """NVFP4 layout, calibration metadata, and CUDA dispatch.""" def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) self._model_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) - def _create_cold_page_policy( + def _initialize_codec( self, cache_config: object, *, @@ -355,7 +164,7 @@ def _create_cold_page_policy( num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], is_draft: bool = False, - ) -> object: + ) -> None: from tensorrt_llm.runtime.kv_cache_manager_v2 import AttentionLayerConfig runtime_type = { @@ -418,4 +227,149 @@ def _create_cold_page_policy( ) ) - return _Nvfp4ColdPagePolicy(layer_layouts, runtime_type if runtime_type is not None else 0) + self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} + self._layer_ids = tuple(sorted(self._layer_layouts)) + self._runtime_type = runtime_type if runtime_type is not None else 0 + + def _build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: + wide_rows: list[list[int]] = [] + integer_rows: list[list[int]] = [] + scale_rows: list[list[float]] = [] + cold_page_bytes = 0 + max_half_groups_per_tile = 0 + + for layer_id, hot_buffers in lifecycle.layers.items(): + layout = self._layer_layouts[int(layer_id)] + expected_roles = {buffer.role for buffer in layout.buffers} + if set(hot_buffers) != expected_roles: + raise ValueError(f"Cold-page layer {layer_id} roles do not match its KVCM layout") + elements = layout.num_kv_heads * layout.tokens_per_page * layout.head_dim + element_bytes = 1 if self._runtime_type == 2 else 2 + expected_raw_bytes = elements * element_bytes + half_groups = elements // _ELEMENTS_PER_HALF_GROUP + compressed_count = sum(buffer.scales is not None for buffer in layout.buffers) + packed_bytes = elements // _ELEMENTS_PER_BYTE + scale_bytes = elements // _ELEMENTS_PER_SCALE + layer_start = cold_page_bytes + scale_start = layer_start + compressed_count * packed_bytes + cursor = scale_start + compressed_count * scale_bytes + + compressed_index = 0 + for buffer in layout.buffers: + is_compressed = buffer.scales is not None + hot = hot_buffers[buffer.role] + raw_base = int(hot.raw_base) + raw_slot_bytes = int(hot.raw_slot_bytes) + raw_bytes = int(hot.raw_bytes) + if raw_base <= 0 or raw_bytes <= 0 or raw_bytes > raw_slot_bytes: + raise ValueError("Cold-page hot buffer has invalid address or size") + + if is_compressed: + data_offset = layer_start + compressed_index * packed_bytes + scale_offset = scale_start + compressed_index * scale_bytes + compressed_index += 1 + if raw_bytes != expected_raw_bytes: + raise ValueError("Hot buffer size does not match NVFP4 geometry") + if raw_base % 16 or raw_slot_bytes % 16: + raise ValueError( + "NVFP4 hot address and Slot stride must be 16-byte aligned" + ) + max_half_groups_per_tile = max( + max_half_groups_per_tile, + min(half_groups, _MAX_HALF_GROUPS_PER_TILE), + ) + else: + data_offset = cursor + scale_offset = 0 + cursor += raw_bytes + + wide_rows.append( + [ + raw_base, + raw_slot_bytes, + raw_bytes, + data_offset, + scale_offset, + 0, + ] + ) + integer_rows.append( + [ + 0, + _NVFP4_TRANSFORM if is_compressed else _LOSSLESS_TRANSFORM, + layout.num_kv_heads if is_compressed else 0, + layout.tokens_per_page if is_compressed else 0, + layout.head_dim if is_compressed else 0, + ] + ) + buffer_scales = buffer.scales if is_compressed else _Nvfp4Scales(1.0, 1.0) + scale_rows.append( + [ + buffer_scales.nvfp4_orig_quant, + buffer_scales.nvfp4_quant_orig, + buffer_scales.fp8_orig_quant, + buffer_scales.fp8_quant_orig, + ] + ) + layer_end = ( + (cursor + _COLD_PAGE_ALIGNMENT - 1) // _COLD_PAGE_ALIGNMENT * _COLD_PAGE_ALIGNMENT + ) + wide_rows[-1][5] = cursor + integer_rows[-1][0] = layer_end - cursor + cold_page_bytes = layer_end + + num_buffers = len(wide_rows) + if not 0 < num_buffers <= _MAX_BUFFERS_PER_LAUNCH: + raise ValueError( + f"NVFP4 cold-page lifecycle has {num_buffers} buffers; " + f"the maximum is {_MAX_BUFFERS_PER_LAUNCH}" + ) + padding = _MAX_BUFFERS_PER_LAUNCH - num_buffers + return _Nvfp4ColdPageMetadata( + wide=torch.tensor( + wide_rows + [[0] * _WIDE_FIELDS for _ in range(padding)], + dtype=torch.int64, + device="cpu", + ), + integers=torch.tensor( + integer_rows + [[0] * _INTEGER_FIELDS for _ in range(padding)], + dtype=torch.int32, + device="cpu", + ), + scales=torch.tensor( + scale_rows + [[0.0] * _SCALE_FIELDS for _ in range(padding)], + dtype=torch.float32, + device="cpu", + ), + num_buffers=num_buffers, + max_half_groups_per_tile=max_half_groups_per_tile, + cold_page_bytes=cold_page_bytes, + ) + + def _invoke_kernel( + self, + operation: str, + metadata: _Nvfp4ColdPageMetadata, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + op = ( + torch.ops.trtllm.nvfp4_cold_page_encode + if operation == "encode" + else torch.ops.trtllm.nvfp4_cold_page_decode + ) + op( + page_indices, + num_pages, + metadata.wide.data_ptr(), + metadata.integers.data_ptr(), + metadata.scales.data_ptr(), + metadata.num_buffers, + metadata.max_half_groups_per_tile, + metadata.cold_page_bytes, + self._runtime_type, + cold_base, + stream, + ) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 27432ea5456f..9dfdc21be7e9 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -1,99 +1,15 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Common construction pipeline for cold-page quantization.""" +"""Common runtime pipeline for cold-page quantization.""" -from typing import TYPE_CHECKING, Optional, Sequence - -from tensorrt_llm._utils import is_sm_100f -from tensorrt_llm.logger import logger -from tensorrt_llm.quantization import QuantAlgo +import copy +from typing import Sequence from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager -if TYPE_CHECKING: - from tensorrt_llm._torch.model_config import ModelConfig - from tensorrt_llm.llmapi.llm_args import ( - ColdPageQuantizationCompressionConfig, - KvCacheConfig, - SpeculativeConfig, - ) - -_NATIVE_KV_CACHE_EQUIVALENTS = {"nvfp4": QuantAlgo.NVFP4} - - -def validate_cold_page_quantization_compatibility( - config: "ColdPageQuantizationCompressionConfig", - kv_cache_config: "KvCacheConfig", - spec_config: Optional["SpeculativeConfig"], -) -> None: - """Validate a cold-page method that will participate in this executor.""" - - from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND - - if _BACKEND == "python": - raise ValueError("Cold-page quantization requires the C++ KVCacheManagerV2 backend") - if kv_cache_config.enable_block_reuse and not config.supports_block_reuse(): - raise ValueError( - f"KV-cache compression algorithm {config.algorithm!r} does not " - "support KV-cache block reuse. Set " - "KvCacheConfig.enable_block_reuse=False." - ) - if spec_config is None: - return - if not config.supports_speculative_decoding(): - raise ValueError( - f"KV-cache compression algorithm {config.algorithm!r} does not " - "support speculative decoding with its current configuration" - ) - mode = spec_config.spec_dec_mode - if not (mode.is_eagle3_one_model() or mode.is_mtp_eagle_one_model()): - raise ValueError( - "Cold-page quantization supports speculative decoding only " - f"with one-model MTP-EAGLE or EAGLE3, not {mode.name}" - ) - - -def create_cold_page_quantization_manager( - config: "ColdPageQuantizationCompressionConfig", - *, - model_config: "ModelConfig", - kv_cache_config: "KvCacheConfig", - spec_config: Optional["SpeculativeConfig"], - estimating_kv_cache: bool = False, -) -> Optional["ColdPageQuantizationCompression"]: - """Select, validate, and construct one cold-page quantization method.""" - - active_quant_algo = _NATIVE_KV_CACHE_EQUIVALENTS.get(config.quant) - if active_quant_algo is None: - raise NotImplementedError(f"Unsupported cold-page quantization format {config.quant!r}") - if estimating_kv_cache: - return None - - quant_config = model_config.quant_config - if ( - quant_config is not None - and getattr(quant_config, "kv_cache_quant_algo", None) == active_quant_algo - ): - logger.info( - "Skipping cold-page %s quantization because the active KV cache " - "already uses the same format; KVCM will migrate it losslessly.", - config.quant.upper(), - ) - return None - - validate_cold_page_quantization_compatibility(config, kv_cache_config, spec_config) - - from .nvfp4_quantization import Nvfp4ColdPageQuantizationCompression - - if not is_sm_100f(): - raise RuntimeError( - "NVFP4 cold-page quantization requires an SM100-family device (SM100 or SM103)." - ) - return Nvfp4ColdPageQuantizationCompression(config) - class ColdPageQuantizationCompression(KVCacheCompressionManager): - """Base pipeline shared by storage-bound cold-page quantizers.""" + """Common codec registration and callbacks for cold-page quantizers.""" uses_iteration_lifecycle = False provides_cold_page_codec = True @@ -108,11 +24,12 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create one generic native codec around a format-specific policy.""" + """Create one callback instance with state isolated to this KVCM.""" from tensorrt_llm.bindings.internal import kv_cache_compression as native - policy = self._create_cold_page_policy( + callback = copy.copy(self) + callback._initialize_codec( cache_config, runtime_dtype=runtime_dtype, pp_layers=pp_layers, @@ -120,9 +37,64 @@ def create_cold_page_codec( head_dim_per_layer=head_dim_per_layer, is_draft=is_draft, ) - return native.create_python_cold_page_codec(policy) + callback._lifecycle_metadata = [] + return native.create_python_cold_page_codec(callback) + + @property + def layer_ids(self) -> tuple[int, ...]: + return self._layer_ids + + def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: + """Resolve hot buffers and publish each lifecycle's cold-page size.""" + + from tensorrt_llm.bindings.internal import kv_cache_compression as native + + self._lifecycle_metadata = [ + self._build_lifecycle_metadata(lifecycle) for lifecycle in lifecycles + ] + properties = [] + for metadata in self._lifecycle_metadata: + lifecycle = native.ColdPageLifecycleProperties() + lifecycle.cold_page_bytes = metadata.cold_page_bytes + lifecycle.page_index_location = native.ColdPageIndexLocation.HOST + properties.append(lifecycle) + return properties + + def encode( + self, + lifecycle_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + self._invoke_kernel( + "encode", + self._lifecycle_metadata[lifecycle_index], + cold_base, + page_indices, + num_pages, + stream, + ) + + def decode( + self, + lifecycle_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + self._invoke_kernel( + "decode", + self._lifecycle_metadata[lifecycle_index], + cold_base, + page_indices, + num_pages, + stream, + ) - def _create_cold_page_policy( + def _initialize_codec( self, cache_config: object, *, @@ -131,6 +103,22 @@ def _create_cold_page_policy( num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], is_draft: bool = False, - ) -> object: - """Build the format-specific layout and callback policy.""" + ) -> None: + """Build format-specific immutable state for one KVCM.""" + raise NotImplementedError + + def _build_lifecycle_metadata(self, lifecycle: object) -> object: + """Resolve one KVCM lifecycle into format-specific launch metadata.""" + raise NotImplementedError + + def _invoke_kernel( + self, + operation: str, + metadata: object, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + """Submit one complete codec batch to the format-specific kernel.""" raise NotImplementedError diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index db351e43dc2f..5ad066a88231 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -41,6 +41,7 @@ supports_native_fp8_lora) from tensorrt_llm.logger import logger from tensorrt_llm.mapping import CpType, Mapping +from tensorrt_llm.quantization import QuantAlgo from ..attention_backend import get_sparse_attn_kv_cache_manager from ..hostfunc import set_low_latency_dispatch @@ -2786,7 +2787,18 @@ def validate_kv_cache_compression_compatibility( spec_config: Optional[SpeculativeConfig], ) -> None: """Reject unsupported KV-cache compression feature combinations.""" - if config.algorithm == "triattention" and not is_sm_100f(): + if config.algorithm == "quantization_for_cold_page": + from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND + + if _BACKEND == "python": + raise ValueError( + "Cold-page quantization requires the C++ KVCacheManagerV2 backend" + ) + if config.quant == "nvfp4" and not is_sm_100f(): + raise RuntimeError( + "NVFP4 cold-page quantization requires an SM100-family device " + "(SM100 or SM103).") + elif config.algorithm == "triattention" and not is_sm_100f(): raise RuntimeError( "TriAttention requires an SM100-family device (SM100 or SM103).") @@ -2800,13 +2812,18 @@ def validate_kv_cache_compression_compatibility( if not config.supports_speculative_decoding(): raise ValueError( f"KV-cache compression algorithm {config.algorithm!r} does not " - "support speculative decoding with its current configuration; " - "TriAttention requires eviction_mode='union'") + "support speculative decoding with its current configuration") mode = spec_config.spec_dec_mode - if not (mode.is_mtp_one_model() or mode.is_eagle3_one_model()): + if config.algorithm == "quantization_for_cold_page": + supported = mode.is_mtp_eagle_one_model() or mode.is_eagle3_one_model() + guidance = "one-model MTP-EAGLE or EAGLE3" + else: + supported = mode.is_mtp_one_model() or mode.is_eagle3_one_model() + guidance = "one-model MTP or EAGLE3" + if not supported: raise ValueError( f"KV-cache compression does not support speculative decoding " - f"mode {mode.name}; use one-model MTP or EAGLE3") + f"mode {mode.name}; use {guidance}") def create_kv_cache_compression_manager( @@ -2825,16 +2842,25 @@ def create_kv_cache_compression_manager( "KV-cache compression does not support HELIX context parallelism.") if config.algorithm == "quantization_for_cold_page": - from ..kv_cache_compression.quantization_for_cold_page import \ - quantization_for_cold_page as cold_page_quantization + if config.quant != "nvfp4": + raise NotImplementedError( + f"Unsupported cold-page quantization format {config.quant!r}") + if estimating_kv_cache: + return None + quant_config = model_engine.model.model_config.quant_config + if (quant_config is not None and getattr( + quant_config, "kv_cache_quant_algo", None) == QuantAlgo.NVFP4): + logger.info( + "Skipping cold-page NVFP4 quantization because the active KV " + "cache already uses NVFP4; KVCM will migrate it losslessly.") + return None - return cold_page_quantization.create_cold_page_quantization_manager( - config, - model_config=model_engine.model.model_config, - kv_cache_config=kv_cache_config, - spec_config=model_engine.spec_config, - estimating_kv_cache=estimating_kv_cache, - ) + validate_kv_cache_compression_compatibility(config, kv_cache_config, + model_engine.spec_config) + from ..kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import \ + Nvfp4ColdPageQuantizationCompression + + return Nvfp4ColdPageQuantizationCompression(config) if config.algorithm == "triattention": validate_kv_cache_compression_compatibility(config, kv_cache_config, diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 7e7e96a96f38..f23596f8ee65 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -10,16 +10,12 @@ import torch from safetensors.torch import save_file -from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page import ( - quantization_for_cold_page as cold_quant_mod, -) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import ( Nvfp4ColdPageQuantizationCompression, _load_modelopt_nvfp4_scales, ) from tensorrt_llm._torch.kv_cache_compression.quantization_for_cold_page.quantization_for_cold_page import ( ColdPageQuantizationCompression, - validate_cold_page_quantization_compatibility, ) from tensorrt_llm._torch.pyexecutor import _util as util_mod from tensorrt_llm._torch.pyexecutor.resource_manager import DataType @@ -86,12 +82,12 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: return module, codec -def _policy(native: SimpleNamespace) -> object: +def _callback(native: SimpleNamespace) -> object: return native.create_python_cold_page_codec.call_args.args[0] def _layouts(native: SimpleNamespace) -> list[object]: - return list(_policy(native)._layer_layouts.values()) + return list(_callback(native)._layer_layouts.values()) def _configure_lifecycle(native: SimpleNamespace, layer_bytes: dict[int, dict[str, int]]) -> object: @@ -107,11 +103,11 @@ def _configure_lifecycle(native: SimpleNamespace, layer_bytes: dict[int, dict[st ) address += 0x10000 layers[layer_id] = hot - _policy(native).configure([SimpleNamespace(layers=layers)]) - return _policy(native)._lifecycle_metadata[0] + _callback(native).configure([SimpleNamespace(layers=layers)]) + return _callback(native)._lifecycle_metadata[0] -def _configure_policy(native: SimpleNamespace, raw_bytes: int) -> object: +def _configure_callback(native: SimpleNamespace, raw_bytes: int) -> object: return _configure_lifecycle(native, {0: {"key": raw_bytes, "value": raw_bytes}}) @@ -135,11 +131,12 @@ def _write_scales(directory, scales_by_layer, *, filename="model.safetensors", p def _validate_compression(mode: object | None = None) -> None: spec_config = None if mode is None else SimpleNamespace(spec_dec_mode=mode) - validate_cold_page_quantization_compatibility( - ColdPageQuantizationCompressionConfig(), - SimpleNamespace(enable_block_reuse=False), - spec_config, - ) + with patch.object(util_mod, "is_sm_100f", return_value=True): + util_mod.validate_kv_cache_compression_compatibility( + ColdPageQuantizationCompressionConfig(), + SimpleNamespace(enable_block_reuse=False), + spec_config, + ) def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_path): @@ -168,7 +165,7 @@ def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_p assert result is codec layouts = _layouts(native) assert [layout.layer_id for layout in layouts] == [0, 1, 2] - assert _policy(native)._runtime_type == 1 + assert _callback(native)._runtime_type == 1 assert [ ( tuple(buffer.scales.nvfp4_orig_quant for buffer in layout.buffers), @@ -259,7 +256,7 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): 1.0, 1.0, ] - metadata = _configure_policy(native, raw_bytes=5120) + metadata = _configure_callback(native, raw_bytes=5120) assert metadata.cold_page_bytes == 2880 assert metadata.wide[:2, 3].tolist() == [0, 1280] assert metadata.wide[:2, 4].tolist() == [2560, 2720] @@ -290,7 +287,7 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: head_dim_per_layer=(32,), ) - metadata = _configure_policy(native, raw_bytes=320) + metadata = _configure_callback(native, raw_bytes=320) assert metadata.cold_page_bytes == 192 assert metadata.wide[:2, 3].tolist() == [0, 80] assert metadata.wide[:2, 4].tolist() == [160, 170] @@ -318,9 +315,15 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): assert results == codecs assert native.create_python_cold_page_codec.call_count == 2 + target_callback, draft_callback = ( + call.args[0] for call in native.create_python_cold_page_codec.call_args_list + ) + assert target_callback is not draft_callback + assert target_callback is not provider + assert draft_callback is not provider -def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: +def test_callback_forwards_a_4096_page_batch_through_one_native_launch() -> None: native, _ = _native() with ( @@ -335,7 +338,7 @@ def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: num_kv_heads_per_layer=(1,), head_dim_per_layer=(16,), ) - policy = _policy(native) + callback = _callback(native) hot = { role: SimpleNamespace( raw_base=0x1000 + index * 0x1000, @@ -344,13 +347,13 @@ def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: ) for index, role in enumerate(("key", "value")) } - properties = policy.configure([SimpleNamespace(layers={0: hot})]) - policy.encode(0, 0x3000, 0x4000, 4096, 0x5000) - policy.decode(0, 0x3000, 0x4000, 4096, 0x5000) + properties = callback.configure([SimpleNamespace(layers={0: hot})]) + callback.encode(0, 0x3000, 0x4000, 4096, 0x5000) + callback.decode(0, 0x3000, 0x4000, 4096, 0x5000) assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" - metadata = policy._lifecycle_metadata[0] + metadata = callback._lifecycle_metadata[0] for operation in (encode, decode): operation.assert_called_once() arguments = operation.call_args.args @@ -369,7 +372,7 @@ def test_policy_forwards_a_4096_page_batch_through_one_native_launch() -> None: ) -def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: +def test_callback_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: native, _ = _native() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( @@ -381,7 +384,7 @@ def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: ) with torch.device("meta"): - metadata = _configure_policy(native, raw_bytes=2048) + metadata = _configure_callback(native, raw_bytes=2048) for tensor, dtype, shape in ( (metadata.wide, torch.int64, (256, 6)), (metadata.integers, torch.int32, (256, 5)), @@ -393,7 +396,7 @@ def test_policy_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: assert tensor.is_contiguous() -def test_policy_rejects_invalid_resolved_hot_buffers() -> None: +def test_callback_rejects_invalid_resolved_hot_buffers() -> None: native, _ = _native() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( @@ -404,7 +407,7 @@ def test_policy_rejects_invalid_resolved_hot_buffers() -> None: head_dim_per_layer=(16,), ) - policy = _policy(native) + callback = _callback(native) def hot(raw_base: int = 0x1000, raw_bytes: int = 2048) -> SimpleNamespace: return SimpleNamespace( @@ -414,18 +417,20 @@ def hot(raw_base: int = 0x1000, raw_bytes: int = 2048) -> SimpleNamespace: ) with pytest.raises(ValueError, match="roles do not match"): - policy.configure( + callback.configure( [SimpleNamespace(layers={0: {"key": hot(), "value": hot(), "extra": hot()}})] ) with pytest.raises(ValueError, match="size does not match"): - policy.configure([SimpleNamespace(layers={0: {"key": hot(raw_bytes=32), "value": hot()}})]) + callback.configure( + [SimpleNamespace(layers={0: {"key": hot(raw_bytes=32), "value": hot()}})] + ) with pytest.raises(ValueError, match="16-byte aligned"): - policy.configure( + callback.configure( [SimpleNamespace(layers={0: {"key": hot(raw_base=0x1001), "value": hot()}})] ) -def test_policy_rejects_more_than_256_lifecycle_buffers() -> None: +def test_callback_rejects_more_than_256_lifecycle_buffers() -> None: native, _ = _native() layers = tuple( AttentionLayerConfig( @@ -591,7 +596,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - assert _policy(native).layer_ids == () + assert _callback(native).layer_ids == () def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -623,7 +628,7 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): assert layout.layer_id == 0 assert [buffer.role for buffer in layout.buffers] == ["key", "index_key"] assert [buffer.scales is not None for buffer in layout.buffers] == [True, False] - assert _policy(native)._runtime_type == 1 + assert _callback(native)._runtime_type == 1 assert layout.num_kv_heads == 1 assert layout.tokens_per_page == 64 assert layout.head_dim == 576 @@ -769,7 +774,7 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): ) layout = _layouts(native)[0] - assert _policy(native)._runtime_type == 2 + assert _callback(native)._runtime_type == 2 assert [buffer.scales.nvfp4_orig_quant for buffer in layout.buffers] == [ 2.0, 4.0, @@ -790,22 +795,20 @@ def test_runtime_admission_is_checked_before_manager_creation(monkeypatch) -> No _validate_compression() monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") - monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: False) + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: False) with pytest.raises(RuntimeError, match="requires an SM100-family device"): - cold_quant_mod.create_cold_page_quantization_manager( + util_mod.create_kv_cache_compression_manager( ColdPageQuantizationCompressionConfig(), - model_config=_factory_model_engine().model.model_config, + model_engine=_factory_model_engine(), kv_cache_config=SimpleNamespace(enable_block_reuse=False), - spec_config=None, ) - monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: True) + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) assert isinstance( - cold_quant_mod.create_cold_page_quantization_manager( + util_mod.create_kv_cache_compression_manager( ColdPageQuantizationCompressionConfig(), - model_config=_factory_model_engine().model.model_config, + model_engine=_factory_model_engine(), kv_cache_config=SimpleNamespace(enable_block_reuse=False), - spec_config=None, ), Nvfp4ColdPageQuantizationCompression, ) @@ -837,16 +840,17 @@ def test_qwen35_mtp3_resolves_to_supported_one_model_mode(monkeypatch) -> None: assert spec_config.spec_dec_mode is SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL assert spec_config.max_draft_len == 3 - validate_cold_page_quantization_compatibility( - ColdPageQuantizationCompressionConfig(), - SimpleNamespace(enable_block_reuse=False), - spec_config, - ) + with patch.object(util_mod, "is_sm_100f", return_value=True): + util_mod.validate_kv_cache_compression_compatibility( + ColdPageQuantizationCompressionConfig(), + SimpleNamespace(enable_block_reuse=False), + spec_config, + ) def test_cold_manager_is_disabled_for_estimation_and_active_nvfp4(monkeypatch) -> None: monkeypatch.setattr(runtime_v2_mod, "_BACKEND", "cpp") - monkeypatch.setattr(cold_quant_mod, "is_sm_100f", lambda: True) + monkeypatch.setattr(util_mod, "is_sm_100f", lambda: True) def build( *, @@ -857,7 +861,11 @@ def build( creator = object.__new__(util_mod.KvCacheCreator) creator._skip_est = skip_est creator._max_seq_len = 1024 - creator._kv_cache_config = SimpleNamespace(host_cache_size=None, disk_cache_size=None) + creator._kv_cache_config = SimpleNamespace( + host_cache_size=None, + disk_cache_size=None, + enable_block_reuse=False, + ) creator._llm_args = SimpleNamespace( kv_cache_compression_config=ColdPageQuantizationCompressionConfig() ) @@ -890,7 +898,7 @@ def build( skip_est_creator._create_kv_cache_manager.call_args.kwargs["cold_page_codec_provider"], Nvfp4ColdPageQuantizationCompression, ) - with patch.object(cold_quant_mod.logger, "info") as log: + with patch.object(util_mod.logger, "info") as log: active_resources, active_creator = build( active_kv_quant=QuantConfig(kv_cache_quant_algo=QuantAlgo.NVFP4) ) From f651c081ee7efe3738b2a0cdd72529f84004f297 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 19:21:43 -0700 Subject: [PATCH 20/29] [None][fix] Finalize cold-page Python callback contract Signed-off-by: tianruih --- .../nanobind/kvCacheCompression/bindings.cpp | 4 +- cpp/tests/unit_tests/kernels/CMakeLists.txt | 16 ------- .../kv_cache_compression/CMakeLists.txt | 17 +++++++ .../nvfp4_quantization.py | 34 ++++++++++---- .../quantization_for_cold_page.py | 46 ------------------- tensorrt_llm/_torch/pyexecutor/_util.py | 5 +- .../_torch/pyexecutor/resource_manager.py | 22 +++++++++ tests/integration/defs/cpp/test_unit_tests.py | 6 +-- .../test_lists/test-db/l0_b200.yml | 15 ++++++ .../test_quantization_for_cold_page.py | 6 +-- 10 files changed, 92 insertions(+), 79 deletions(-) diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 4332203d80d3..18bd4ba221c4 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -100,13 +100,13 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { - invoke("encode", lifecycleIndex, coldBase, pageIndices, numPages, stream); + invoke("encode_cold_pages", lifecycleIndex, coldBase, pageIndices, numPages, stream); } void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { - invoke("decode", lifecycleIndex, coldBase, pageIndices, numPages, stream); + invoke("decode_cold_pages", lifecycleIndex, coldBase, pageIndices, numPages, stream); } template diff --git a/cpp/tests/unit_tests/kernels/CMakeLists.txt b/cpp/tests/unit_tests/kernels/CMakeLists.txt index ef7d99072940..6b0e5a118211 100644 --- a/cpp/tests/unit_tests/kernels/CMakeLists.txt +++ b/cpp/tests/unit_tests/kernels/CMakeLists.txt @@ -46,22 +46,6 @@ if(USING_OSS_CUTLASS_MOE_GEMM) endif() add_gtest(ropeTest ropeTest.cu) -set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC - nvfp4ColdPageKernelsTest.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) -add_gtest(nvfp4ColdPageKernelsTest "${NVFP4_COLD_PAGE_KERNEL_TEST_SRC}" - NO_TLLM_LINKAGE) -target_include_directories( - nvfp4ColdPageKernelsTest - PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) -target_link_libraries(nvfp4ColdPageKernelsTest PRIVATE CUDA::cudart - CUDA::cuda_driver) add_gtest(shiftKCacheKernelTest shiftKCacheKernelTest.cu) add_gtest(smoothQuantKernelTest smoothQuant/smoothQuantKernelTest.cpp) add_gtest(stopCriteriaKernelsTest stopCriteriaKernelsTest.cpp) diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index ee19f85157e9..dbfe8cc5e023 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -18,3 +18,20 @@ add_gtest(coldPageCodecTest "${COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) target_link_libraries(coldPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) target_include_directories( coldPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) + +set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC + ${PROJECT_SOURCE_DIR}/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) +add_gtest(nvfp4ColdPageKernelsTest "${NVFP4_COLD_PAGE_KERNEL_TEST_SRC}" + NO_TLLM_LINKAGE) +target_include_directories( + nvfp4ColdPageKernelsTest + PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) +target_link_libraries(nvfp4ColdPageKernelsTest PRIVATE CUDA::cudart + CUDA::cuda_driver) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index dba13733bad3..34600078c13b 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -346,21 +346,39 @@ def _build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata cold_page_bytes=cold_page_bytes, ) - def _invoke_kernel( + def encode_cold_pages( self, - operation: str, - metadata: _Nvfp4ColdPageMetadata, + lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, stream: int, ) -> None: - op = ( - torch.ops.trtllm.nvfp4_cold_page_encode - if operation == "encode" - else torch.ops.trtllm.nvfp4_cold_page_decode + metadata = self._lifecycle_metadata[lifecycle_index] + torch.ops.trtllm.nvfp4_cold_page_encode( + page_indices, + num_pages, + metadata.wide.data_ptr(), + metadata.integers.data_ptr(), + metadata.scales.data_ptr(), + metadata.num_buffers, + metadata.max_half_groups_per_tile, + metadata.cold_page_bytes, + self._runtime_type, + cold_base, + stream, ) - op( + + def decode_cold_pages( + self, + lifecycle_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + metadata = self._lifecycle_metadata[lifecycle_index] + torch.ops.trtllm.nvfp4_cold_page_decode( page_indices, num_pages, metadata.wide.data_ptr(), diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 9dfdc21be7e9..5865dc00b4ce 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -60,40 +60,6 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: properties.append(lifecycle) return properties - def encode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - self._invoke_kernel( - "encode", - self._lifecycle_metadata[lifecycle_index], - cold_base, - page_indices, - num_pages, - stream, - ) - - def decode( - self, - lifecycle_index: int, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - self._invoke_kernel( - "decode", - self._lifecycle_metadata[lifecycle_index], - cold_base, - page_indices, - num_pages, - stream, - ) - def _initialize_codec( self, cache_config: object, @@ -110,15 +76,3 @@ def _initialize_codec( def _build_lifecycle_metadata(self, lifecycle: object) -> object: """Resolve one KVCM lifecycle into format-specific launch metadata.""" raise NotImplementedError - - def _invoke_kernel( - self, - operation: str, - metadata: object, - cold_base: int, - page_indices: int, - num_pages: int, - stream: int, - ) -> None: - """Submit one complete codec batch to the format-specific kernel.""" - raise NotImplementedError diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 5ad066a88231..b04fa4d46ea3 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -2810,9 +2810,12 @@ def validate_kv_cache_compression_compatibility( if spec_config is None: return if not config.supports_speculative_decoding(): + guidance = ("; TriAttention requires eviction_mode='union'" + if config.algorithm == "triattention" else "") raise ValueError( f"KV-cache compression algorithm {config.algorithm!r} does not " - "support speculative decoding with its current configuration") + "support speculative decoding with its current configuration" + f"{guidance}") mode = spec_config.spec_dec_mode if config.algorithm == "quantization_for_cold_page": supported = mode.is_mtp_eagle_one_model() or mode.is_eagle3_one_model() diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 754827a5c459..94e0b53fa522 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -2819,6 +2819,28 @@ def create_cold_page_codec( """Create a native cold-page codec when the algorithm provides one.""" return None + def encode_cold_pages( + self, + lifecycle_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + """Encode one complete KVCM migration batch into cold storage.""" + raise NotImplementedError + + def decode_cold_pages( + self, + lifecycle_index: int, + cold_base: int, + page_indices: int, + num_pages: int, + stream: int, + ) -> None: + """Decode one complete KVCM migration batch from cold storage.""" + raise NotImplementedError + # ================================================================== # # KV-cache lifecycle hooks (5, in temporal order). # # Subclasses override what they need; all default to no-op. # diff --git a/tests/integration/defs/cpp/test_unit_tests.py b/tests/integration/defs/cpp/test_unit_tests.py index 730ebbf389ed..9989e70a355c 100644 --- a/tests/integration/defs/cpp/test_unit_tests.py +++ b/tests/integration/defs/cpp/test_unit_tests.py @@ -4,11 +4,11 @@ import pytest -@pytest.mark.parametrize("build_google_tests", ["80", "86", "89", "90"], +@pytest.mark.parametrize("build_google_tests", ["80", "86", "89", "90", "100"], indirect=True) @pytest.mark.parametrize("test_group", [ - "batch_manager", "common", "executor", "kernels", "layers", "runtime", - "thop" + "batch_manager", "common", "executor", "kernels", "kv_cache_compression", + "layers", "runtime", "thop" ]) def test_unit_tests(build_google_tests, test_group, build_dir, lora_setup): diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index eb98bc2a7b27..faa08fb6ea9a 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -1,5 +1,20 @@ version: 0.0.1 l0_b200: +- condition: + ranges: + system_gpu_count: + gte: 1 + lte: 1 + wildcards: + gpu: + - '*b100*' + - '*b200*' + linux_distribution_name: ubuntu* + terms: + stage: pre_merge + backend: cpp + tests: + - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-100] - condition: ranges: system_gpu_count: diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index f23596f8ee65..e7e7e8dd7bc3 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -323,7 +323,7 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): assert draft_callback is not provider -def test_callback_forwards_a_4096_page_batch_through_one_native_launch() -> None: +def test_callback_forwards_a_4096_page_batch_through_one_custom_op_call() -> None: native, _ = _native() with ( @@ -348,8 +348,8 @@ def test_callback_forwards_a_4096_page_batch_through_one_native_launch() -> None for index, role in enumerate(("key", "value")) } properties = callback.configure([SimpleNamespace(layers={0: hot})]) - callback.encode(0, 0x3000, 0x4000, 4096, 0x5000) - callback.decode(0, 0x3000, 0x4000, 4096, 0x5000) + callback.encode_cold_pages(0, 0x3000, 0x4000, 4096, 0x5000) + callback.decode_cold_pages(0, 0x3000, 0x4000, 4096, 0x5000) assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" From 4f2070d1f524231b4788870d4522d874b28679bc Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 19:39:23 -0700 Subject: [PATCH 21/29] [None][refactor] Expose cold-page quantization hooks Signed-off-by: tianruih --- .../quantization_for_cold_page/nvfp4_quantization.py | 4 ++-- .../quantization_for_cold_page.py | 8 ++++---- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index 34600078c13b..c1cbc3981cc1 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -155,7 +155,7 @@ def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) self._model_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) - def _initialize_codec( + def initialize_codec( self, cache_config: object, *, @@ -231,7 +231,7 @@ def _initialize_codec( self._layer_ids = tuple(sorted(self._layer_layouts)) self._runtime_type = runtime_type if runtime_type is not None else 0 - def _build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: + def build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: wide_rows: list[list[int]] = [] integer_rows: list[list[int]] = [] scale_rows: list[list[float]] = [] diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index 5865dc00b4ce..c2bff20ed6bd 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -29,7 +29,7 @@ def create_cold_page_codec( from tensorrt_llm.bindings.internal import kv_cache_compression as native callback = copy.copy(self) - callback._initialize_codec( + callback.initialize_codec( cache_config, runtime_dtype=runtime_dtype, pp_layers=pp_layers, @@ -50,7 +50,7 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: from tensorrt_llm.bindings.internal import kv_cache_compression as native self._lifecycle_metadata = [ - self._build_lifecycle_metadata(lifecycle) for lifecycle in lifecycles + self.build_lifecycle_metadata(lifecycle) for lifecycle in lifecycles ] properties = [] for metadata in self._lifecycle_metadata: @@ -60,7 +60,7 @@ def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: properties.append(lifecycle) return properties - def _initialize_codec( + def initialize_codec( self, cache_config: object, *, @@ -73,6 +73,6 @@ def _initialize_codec( """Build format-specific immutable state for one KVCM.""" raise NotImplementedError - def _build_lifecycle_metadata(self, lifecycle: object) -> object: + def build_lifecycle_metadata(self, lifecycle: object) -> object: """Resolve one KVCM lifecycle into format-specific launch metadata.""" raise NotImplementedError From 0e4d4e157725f035fb0ecf930a7cea0a0c95ab05 Mon Sep 17 00:00:00 2001 From: tianruih Date: Wed, 26 Aug 2026 20:12:08 -0700 Subject: [PATCH 22/29] [None][refactor] Separate cold-page codec state Signed-off-by: tianruih --- .../nativeColdPageCodec.cpp | 60 +++++++------- .../nativeColdPageCodec.h | 6 +- .../nanobind/kvCacheCompression/bindings.cpp | 47 ++++++----- .../coldPageCodecTest.cpp | 10 +-- .../nvfp4_quantization.py | 43 +++++++--- .../quantization_for_cold_page.py | 33 +++----- .../_torch/pyexecutor/resource_manager.py | 2 + .../integration/test_lists/test-db/l0_a30.yml | 1 + .../test_lists/test-db/l0_b200.yml | 15 ---- .../test_quantization_for_cold_page.py | 79 +++++++++++-------- 10 files changed, 154 insertions(+), 142 deletions(-) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp index 0ce52227457e..aae08e5e2d82 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -43,12 +43,12 @@ ResolvedHotLifecycle resolveLifecycle(kv::PoolGroupDesc const& gpuDesc, kv::Slot return result; } -void drainAfterPolicyFailure(cudaStream_t stream) noexcept +void drainAfterProviderFailure(cudaStream_t stream) noexcept { auto const status = cudaStreamSynchronize(stream); if (status != cudaSuccess) { - TLLM_LOG_ERROR("Cold-page policy rollback drain failed: %s", cudaGetErrorString(status)); + TLLM_LOG_ERROR("Cold-page provider rollback drain failed: %s", cudaGetErrorString(status)); std::terminate(); } } @@ -71,7 +71,7 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG } std::map pendingGroups; - std::vector policyLifecycles; + std::vector providerLifecycles; std::set consumedLayers; for (kv::PoolGroupIndex poolGroupIndex{0}; poolGroupIndex < numGpuDescs; ++poolGroupIndex) @@ -80,31 +80,31 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG for (auto const& variant : gpuDesc.slotDesc.variants) { auto resolved = resolveLifecycle(gpuDesc, variant); - auto const policyLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), + auto const providerLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), [this](auto const& layer) { return mLayerIds.count(layer.first) != 0U; }); LayerGroupState state; - if (policyLayerCount == 0U) + if (providerLayerCount == 0U) { state.coldPageBytes = losslessCodec->queryColdPageBytes(variant.lifeCycleId); state.pageIndexLocation = losslessCodec->queryPageIndexLocation(variant.lifeCycleId); } else { - if (policyLayerCount != resolved.layers.size()) + if (providerLayerCount != resolved.layers.size()) { - throw std::invalid_argument("A lifecycle cannot mix policy-owned and fallback layers"); + throw std::invalid_argument("A lifecycle cannot mix provider-owned and fallback layers"); } for (auto const& [layerId, buffers] : resolved.layers) { static_cast(buffers); if (!consumedLayers.emplace(layerId).second) { - throw std::invalid_argument("A policy layer appears in multiple lifecycles"); + throw std::invalid_argument("A provider layer appears in multiple lifecycles"); } } - state.lifecycleIndex = policyLifecycles.size(); - policyLifecycles.push_back(std::move(resolved)); + state.lifecycleIndex = providerLifecycles.size(); + providerLifecycles.push_back(std::move(resolved)); } if (!pendingGroups.emplace(variant.lifeCycleId, std::move(state)).second) @@ -115,22 +115,22 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG } if (consumedLayers != mLayerIds) { - throw std::invalid_argument("A policy layer is absent from all GPU descriptors"); + throw std::invalid_argument("A provider layer is absent from all GPU descriptors"); } - auto const properties = configurePolicy(policyLifecycles); - if (properties.size() != policyLifecycles.size()) + auto const properties = configureProvider(providerLifecycles); + if (properties.size() != providerLifecycles.size()) { - throw std::invalid_argument("Cold-page policy returned an unexpected lifecycle count"); + throw std::invalid_argument("Cold-page provider returned an unexpected lifecycle count"); } for (std::size_t index = 0; index < properties.size(); ++index) { auto const& lifecycle = properties[index]; if (lifecycle.coldPageBytes == 0U || lifecycle.pageIndexLocation == kv::PageIndexLocation::kBadLocation) { - throw std::invalid_argument("Cold-page policy returned invalid storage properties"); + throw std::invalid_argument("Cold-page provider returned invalid storage properties"); } - auto& state = pendingGroups.at(policyLifecycles[index].lifeCycleId); + auto& state = pendingGroups.at(providerLifecycles[index].lifeCycleId); state.coldPageBytes = lifecycle.coldPageBytes; state.pageIndexLocation = lifecycle.pageIndexLocation; } @@ -178,7 +178,7 @@ kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroup bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { - bool policyStarted = false; + bool providerStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -194,24 +194,24 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr { return mLosslessCodec->encode(layerGroupId, dstBasePtr, pageIndices, numBasePages, stream); } - policyStarted = true; - encodePolicy(*state->lifecycleIndex, dstBasePtr, pageIndices, numBasePages, stream); + providerStarted = true; + encodeProvider(*state->lifecycleIndex, dstBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { - if (policyStarted) + if (providerStarted) { - drainAfterPolicyFailure(stream); + drainAfterProviderFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: %s", error.what()); return false; } catch (...) { - if (policyStarted) + if (providerStarted) { - drainAfterPolicyFailure(stream); + drainAfterProviderFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::encode failed before completion fencing: unknown error"); return false; @@ -221,7 +221,7 @@ bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { - bool policyStarted = false; + bool providerStarted = false; try { auto const* state = findLayerGroup(layerGroupId); @@ -237,24 +237,24 @@ bool NativeColdPageCodec::decode(kv::LayerGroupId layerGroupId, void const* srcB { return mLosslessCodec->decode(layerGroupId, srcBasePtr, pageIndices, numBasePages, stream); } - policyStarted = true; - decodePolicy(*state->lifecycleIndex, srcBasePtr, pageIndices, numBasePages, stream); + providerStarted = true; + decodeProvider(*state->lifecycleIndex, srcBasePtr, pageIndices, numBasePages, stream); return true; } catch (std::exception const& error) { - if (policyStarted) + if (providerStarted) { - drainAfterPolicyFailure(stream); + drainAfterProviderFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: %s", error.what()); return false; } catch (...) { - if (policyStarted) + if (providerStarted) { - drainAfterPolicyFailure(stream); + drainAfterProviderFailure(stream); } TLLM_LOG_ERROR("NativeColdPageCodec::decode failed before completion fencing: unknown error"); return false; diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h index 16b6494c837f..6c09a3ffb93d 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -66,16 +66,16 @@ class NativeColdPageCodec : public kv::IKvCacheColdPageCodec std::size_t numBasePages, cudaStream_t stream) noexcept final; private: - virtual std::vector configurePolicy( + virtual std::vector configureProvider( std::vector const& lifecycles) = 0; //! Enqueue only on stream; this codec drains partial submissions after a throw. - virtual void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + virtual void encodeProvider(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; - virtual void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + virtual void decodeProvider(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) = 0; diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 18bd4ba221c4..8d4f7f42568c 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -46,15 +46,17 @@ static_assert(offsetof(kv::PageIndexPair, dst) == 0); static_assert(offsetof(kv::PageIndexPair, src) == 4); static_assert(std::is_trivially_copyable_v); -//! Algorithm-neutral adapter from KVCM's native codec calls to one Python policy. +//! Algorithm-neutral adapter from KVCM migration calls to a Python provider. class PythonColdPageCodec final : public compression::NativeColdPageCodec { public: - explicit PythonColdPageCodec(nb::handle policy) - : NativeColdPageCodec(readLayerIds(policy)) - , mPolicy(policy.ptr()) + PythonColdPageCodec(nb::handle provider, nb::handle codecState) + : NativeColdPageCodec(readLayerIds(codecState)) + , mProvider(provider.ptr()) + , mCodecState(codecState.ptr()) { - Py_INCREF(mPolicy); + Py_INCREF(mProvider); + Py_INCREF(mCodecState); } ~PythonColdPageCodec() override @@ -62,34 +64,35 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec if (Py_IsInitialized()) { nb::gil_scoped_acquire acquire; - Py_DECREF(mPolicy); + Py_DECREF(mCodecState); + Py_DECREF(mProvider); } } private: - static std::set readLayerIds(nb::handle policy) + static std::set readLayerIds(nb::handle codecState) { - if (policy.is_none()) + if (codecState.is_none()) { - throw std::invalid_argument("Cold-page policy must not be None"); + throw std::invalid_argument("Cold-page codec state must not be None"); } - auto const layerIds = nb::cast>(policy.attr("layer_ids")); + auto const layerIds = nb::cast>(codecState.attr("layer_ids")); std::set result(layerIds.begin(), layerIds.end()); if (result.size() != layerIds.size()) { - throw std::invalid_argument("Cold-page policy layer IDs must be unique"); + throw std::invalid_argument("Cold-page codec state layer IDs must be unique"); } return result; } - std::vector configurePolicy( + std::vector configureProvider( std::vector const& lifecycles) override { nb::gil_scoped_acquire acquire; try { return nb::cast>( - nb::borrow(mPolicy).attr("configure")(lifecycles)); + nb::borrow(mProvider).attr("configure")(nb::borrow(mCodecState), lifecycles)); } catch (nb::python_error const& error) { @@ -97,13 +100,13 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec } } - void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + void encodeProvider(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { invoke("encode_cold_pages", lifecycleIndex, coldBase, pageIndices, numPages, stream); } - void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodeProvider(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { invoke("decode_cold_pages", lifecycleIndex, coldBase, pageIndices, numPages, stream); @@ -117,8 +120,9 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec nb::gil_scoped_acquire acquire; try { - nb::borrow(mPolicy).attr(method)(lifecycleIndex, reinterpret_cast(coldBase), - reinterpret_cast(pageIndices), numPages, reinterpret_cast(stream)); + nb::borrow(mProvider).attr(method)(nb::borrow(mCodecState), lifecycleIndex, + reinterpret_cast(coldBase), reinterpret_cast(pageIndices), numPages, + reinterpret_cast(stream)); } catch (nb::python_error const& error) { @@ -126,7 +130,8 @@ class PythonColdPageCodec final : public compression::NativeColdPageCodec } } - PyObject* mPolicy; + PyObject* mProvider; + PyObject* mCodecState; }; } // namespace @@ -155,9 +160,9 @@ void initBindings(nb::module_& module) module.def( "create_python_cold_page_codec", - [](nb::handle policy) -> std::unique_ptr - { return std::make_unique(policy); }, - nb::arg("policy")); + [](nb::handle provider, nb::handle codecState) -> std::unique_ptr + { return std::make_unique(provider, codecState); }, + nb::arg("provider"), nb::arg("codec_state")); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp index bdb27f11825e..b853e7bc34be 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp @@ -76,7 +76,7 @@ class RecordingCodec final : public NativeColdPageCodec { } - std::vector configurePolicy( + std::vector configureProvider( std::vector const& lifecycles) override { resolved = lifecycles; @@ -88,7 +88,7 @@ class RecordingCodec final : public NativeColdPageCodec lifecycles.size(), ColdPageLifecycleProperties{777U, kv::PageIndexLocation::kHost}); } - void encodePolicy(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, + void encodeProvider(std::size_t lifecycleIndex, void* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { if (failBatches) @@ -103,7 +103,7 @@ class RecordingCodec final : public NativeColdPageCodec lastStream = stream; } - void decodePolicy(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, + void decodeProvider(std::size_t lifecycleIndex, void const* coldBase, kv::PageIndexPair const* pageIndices, std::size_t numPages, cudaStream_t stream) override { if (failBatches) @@ -207,7 +207,7 @@ TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) } } -TEST(NativeColdPageCodecTest, CatchesPolicyConfigureFailuresAndInvalidBatches) +TEST(NativeColdPageCodecTest, CatchesProviderConfigureFailuresAndInvalidBatches) { RecordingCodec codec{{0}}; codec.failConfigure = true; @@ -224,7 +224,7 @@ TEST(NativeColdPageCodecTest, CatchesPolicyConfigureFailuresAndInvalidBatches) EXPECT_EQ(validCodec.decodeCalls, 0); } -TEST(NativeColdPageCodecTest, PolicyFailureUsesTheSuppliedCudaStreamForRollback) +TEST(NativeColdPageCodecTest, ProviderFailureUsesTheSuppliedCudaStreamForRollback) { int deviceCount = 0; if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index c1cbc3981cc1..44853ef736e0 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -6,7 +6,7 @@ import math import os import re -from dataclasses import dataclass +from dataclasses import dataclass, field from pathlib import Path from typing import TYPE_CHECKING, Sequence @@ -78,6 +78,16 @@ class _Nvfp4ColdPageMetadata: cold_page_bytes: int +@dataclass +class _Nvfp4ColdPageCodecState: + """NVFP4 state owned by one target, draft, or retry codec.""" + + layer_layouts: dict[int, _Nvfp4LayerLayout] + layer_ids: tuple[int, ...] + runtime_type: int + lifecycle_metadata: tuple[_Nvfp4ColdPageMetadata, ...] = field(init=False) + + def _load_modelopt_nvfp4_scales( checkpoint_path: str | None, ) -> dict[int, _LayerScales]: @@ -155,7 +165,7 @@ def __init__(self, config: "ColdPageQuantizationCompressionConfig") -> None: super().__init__(config) self._model_scales = _load_modelopt_nvfp4_scales(config.scale_checkpoint_path) - def initialize_codec( + def build_codec_state( self, cache_config: object, *, @@ -164,7 +174,7 @@ def initialize_codec( num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], is_draft: bool = False, - ) -> None: + ) -> _Nvfp4ColdPageCodecState: from tensorrt_llm.runtime.kv_cache_manager_v2 import AttentionLayerConfig runtime_type = { @@ -227,11 +237,16 @@ def initialize_codec( ) ) - self._layer_layouts = {layout.layer_id: layout for layout in layer_layouts} - self._layer_ids = tuple(sorted(self._layer_layouts)) - self._runtime_type = runtime_type if runtime_type is not None else 0 + layouts_by_layer = {layout.layer_id: layout for layout in layer_layouts} + return _Nvfp4ColdPageCodecState( + layer_layouts=layouts_by_layer, + layer_ids=tuple(sorted(layouts_by_layer)), + runtime_type=runtime_type if runtime_type is not None else 0, + ) - def build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: + def build_lifecycle_metadata( + self, codec_state: _Nvfp4ColdPageCodecState, lifecycle: object + ) -> _Nvfp4ColdPageMetadata: wide_rows: list[list[int]] = [] integer_rows: list[list[int]] = [] scale_rows: list[list[float]] = [] @@ -239,12 +254,12 @@ def build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: max_half_groups_per_tile = 0 for layer_id, hot_buffers in lifecycle.layers.items(): - layout = self._layer_layouts[int(layer_id)] + layout = codec_state.layer_layouts[int(layer_id)] expected_roles = {buffer.role for buffer in layout.buffers} if set(hot_buffers) != expected_roles: raise ValueError(f"Cold-page layer {layer_id} roles do not match its KVCM layout") elements = layout.num_kv_heads * layout.tokens_per_page * layout.head_dim - element_bytes = 1 if self._runtime_type == 2 else 2 + element_bytes = 1 if codec_state.runtime_type == 2 else 2 expected_raw_bytes = elements * element_bytes half_groups = elements // _ELEMENTS_PER_HALF_GROUP compressed_count = sum(buffer.scales is not None for buffer in layout.buffers) @@ -348,13 +363,14 @@ def build_lifecycle_metadata(self, lifecycle: object) -> _Nvfp4ColdPageMetadata: def encode_cold_pages( self, + codec_state: _Nvfp4ColdPageCodecState, lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, stream: int, ) -> None: - metadata = self._lifecycle_metadata[lifecycle_index] + metadata = codec_state.lifecycle_metadata[lifecycle_index] torch.ops.trtllm.nvfp4_cold_page_encode( page_indices, num_pages, @@ -364,20 +380,21 @@ def encode_cold_pages( metadata.num_buffers, metadata.max_half_groups_per_tile, metadata.cold_page_bytes, - self._runtime_type, + codec_state.runtime_type, cold_base, stream, ) def decode_cold_pages( self, + codec_state: _Nvfp4ColdPageCodecState, lifecycle_index: int, cold_base: int, page_indices: int, num_pages: int, stream: int, ) -> None: - metadata = self._lifecycle_metadata[lifecycle_index] + metadata = codec_state.lifecycle_metadata[lifecycle_index] torch.ops.trtllm.nvfp4_cold_page_decode( page_indices, num_pages, @@ -387,7 +404,7 @@ def decode_cold_pages( metadata.num_buffers, metadata.max_half_groups_per_tile, metadata.cold_page_bytes, - self._runtime_type, + codec_state.runtime_type, cold_base, stream, ) diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py index c2bff20ed6bd..d1c49f2bf0e8 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/quantization_for_cold_page.py @@ -2,8 +2,7 @@ # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. """Common runtime pipeline for cold-page quantization.""" -import copy -from typing import Sequence +from typing import Any, Sequence from ...pyexecutor.resource_manager import DataType, KVCacheCompressionManager @@ -24,12 +23,11 @@ def create_cold_page_codec( head_dim_per_layer: Sequence[int], is_draft: bool = False, ) -> object: - """Create one callback instance with state isolated to this KVCM.""" + """Create one native codec with state isolated to this KVCM.""" from tensorrt_llm.bindings.internal import kv_cache_compression as native - callback = copy.copy(self) - callback.initialize_codec( + codec_state = self.build_codec_state( cache_config, runtime_dtype=runtime_dtype, pp_layers=pp_layers, @@ -37,30 +35,25 @@ def create_cold_page_codec( head_dim_per_layer=head_dim_per_layer, is_draft=is_draft, ) - callback._lifecycle_metadata = [] - return native.create_python_cold_page_codec(callback) + return native.create_python_cold_page_codec(self, codec_state) - @property - def layer_ids(self) -> tuple[int, ...]: - return self._layer_ids - - def configure(self, lifecycles: Sequence[object]) -> Sequence[object]: + def configure(self, codec_state: Any, lifecycles: Sequence[object]) -> Sequence[object]: """Resolve hot buffers and publish each lifecycle's cold-page size.""" from tensorrt_llm.bindings.internal import kv_cache_compression as native - self._lifecycle_metadata = [ - self.build_lifecycle_metadata(lifecycle) for lifecycle in lifecycles - ] + codec_state.lifecycle_metadata = tuple( + self.build_lifecycle_metadata(codec_state, lifecycle) for lifecycle in lifecycles + ) properties = [] - for metadata in self._lifecycle_metadata: + for metadata in codec_state.lifecycle_metadata: lifecycle = native.ColdPageLifecycleProperties() lifecycle.cold_page_bytes = metadata.cold_page_bytes lifecycle.page_index_location = native.ColdPageIndexLocation.HOST properties.append(lifecycle) return properties - def initialize_codec( + def build_codec_state( self, cache_config: object, *, @@ -69,10 +62,10 @@ def initialize_codec( num_kv_heads_per_layer: Sequence[int], head_dim_per_layer: Sequence[int], is_draft: bool = False, - ) -> None: - """Build format-specific immutable state for one KVCM.""" + ) -> object: + """Build the format-specific state owned by one native codec.""" raise NotImplementedError - def build_lifecycle_metadata(self, lifecycle: object) -> object: + def build_lifecycle_metadata(self, codec_state: object, lifecycle: object) -> object: """Resolve one KVCM lifecycle into format-specific launch metadata.""" raise NotImplementedError diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 94e0b53fa522..b02e82f5d995 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -2821,6 +2821,7 @@ def create_cold_page_codec( def encode_cold_pages( self, + codec_state: object, lifecycle_index: int, cold_base: int, page_indices: int, @@ -2832,6 +2833,7 @@ def encode_cold_pages( def decode_cold_pages( self, + codec_state: object, lifecycle_index: int, cold_base: int, page_indices: int, diff --git a/tests/integration/test_lists/test-db/l0_a30.yml b/tests/integration/test_lists/test-db/l0_a30.yml index b165611714e6..5c415adcca16 100644 --- a/tests/integration/test_lists/test-db/l0_a30.yml +++ b/tests/integration/test_lists/test-db/l0_a30.yml @@ -43,6 +43,7 @@ l0_a30: tests: # ------------- CPP tests --------------- - cpp/test_unit_tests.py::test_unit_tests[batch_manager-80] + - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-80] - condition: ranges: system_gpu_count: diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index faa08fb6ea9a..eb98bc2a7b27 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -1,20 +1,5 @@ version: 0.0.1 l0_b200: -- condition: - ranges: - system_gpu_count: - gte: 1 - lte: 1 - wildcards: - gpu: - - '*b100*' - - '*b200*' - linux_distribution_name: ubuntu* - terms: - stage: pre_merge - backend: cpp - tests: - - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-100] - condition: ranges: system_gpu_count: diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index e7e7e8dd7bc3..17374cd034f9 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -82,12 +82,16 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: return module, codec -def _callback(native: SimpleNamespace) -> object: +def _provider(native: SimpleNamespace) -> object: return native.create_python_cold_page_codec.call_args.args[0] +def _codec_state(native: SimpleNamespace) -> object: + return native.create_python_cold_page_codec.call_args.args[1] + + def _layouts(native: SimpleNamespace) -> list[object]: - return list(_callback(native)._layer_layouts.values()) + return list(_codec_state(native).layer_layouts.values()) def _configure_lifecycle(native: SimpleNamespace, layer_bytes: dict[int, dict[str, int]]) -> object: @@ -103,11 +107,13 @@ def _configure_lifecycle(native: SimpleNamespace, layer_bytes: dict[int, dict[st ) address += 0x10000 layers[layer_id] = hot - _callback(native).configure([SimpleNamespace(layers=layers)]) - return _callback(native)._lifecycle_metadata[0] + provider = _provider(native) + codec_state = _codec_state(native) + provider.configure(codec_state, [SimpleNamespace(layers=layers)]) + return codec_state.lifecycle_metadata[0] -def _configure_callback(native: SimpleNamespace, raw_bytes: int) -> object: +def _configure_default_lifecycle(native: SimpleNamespace, raw_bytes: int) -> object: return _configure_lifecycle(native, {0: {"key": raw_bytes, "value": raw_bytes}}) @@ -165,7 +171,7 @@ def test_optional_modelopt_scales_map_pp_layers_and_default_missing_layers(tmp_p assert result is codec layouts = _layouts(native) assert [layout.layer_id for layout in layouts] == [0, 1, 2] - assert _callback(native)._runtime_type == 1 + assert _codec_state(native).runtime_type == 1 assert [ ( tuple(buffer.scales.nvfp4_orig_quant for buffer in layout.buffers), @@ -256,7 +262,7 @@ def test_omitted_scale_checkpoint_uses_identity_and_keeps_kv_geometry(): 1.0, 1.0, ] - metadata = _configure_callback(native, raw_bytes=5120) + metadata = _configure_default_lifecycle(native, raw_bytes=5120) assert metadata.cold_page_bytes == 2880 assert metadata.wide[:2, 3].tolist() == [0, 1280] assert metadata.wide[:2, 4].tolist() == [2560, 2720] @@ -287,7 +293,7 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: head_dim_per_layer=(32,), ) - metadata = _configure_callback(native, raw_bytes=320) + metadata = _configure_default_lifecycle(native, raw_bytes=320) assert metadata.cold_page_bytes == 192 assert metadata.wide[:2, 3].tolist() == [0, 80] assert metadata.wide[:2, 4].tolist() == [160, 170] @@ -295,7 +301,7 @@ def test_mha_layout_is_k_v_then_scales_and_layer_padding() -> None: assert metadata.integers[:2, 0].tolist() == [0, 12] -def test_provider_creates_one_native_codec_per_kv_cache_manager(): +def test_provider_creates_independent_state_per_kv_cache_manager() -> None: native, _ = _native() codecs = (object(), object()) native.create_python_cold_page_codec.side_effect = codecs @@ -315,15 +321,14 @@ def test_provider_creates_one_native_codec_per_kv_cache_manager(): assert results == codecs assert native.create_python_cold_page_codec.call_count == 2 - target_callback, draft_callback = ( - call.args[0] for call in native.create_python_cold_page_codec.call_args_list - ) - assert target_callback is not draft_callback - assert target_callback is not provider - assert draft_callback is not provider + calls = native.create_python_cold_page_codec.call_args_list + assert all(call.args[0] is provider for call in calls) + target_state, draft_state = (call.args[1] for call in calls) + assert target_state is not draft_state + assert target_state.layer_ids == draft_state.layer_ids == (0,) -def test_callback_forwards_a_4096_page_batch_through_one_custom_op_call() -> None: +def test_provider_forwards_a_4096_page_batch_through_one_custom_op_call() -> None: native, _ = _native() with ( @@ -338,7 +343,8 @@ def test_callback_forwards_a_4096_page_batch_through_one_custom_op_call() -> Non num_kv_heads_per_layer=(1,), head_dim_per_layer=(16,), ) - callback = _callback(native) + provider = _provider(native) + codec_state = _codec_state(native) hot = { role: SimpleNamespace( raw_base=0x1000 + index * 0x1000, @@ -347,13 +353,13 @@ def test_callback_forwards_a_4096_page_batch_through_one_custom_op_call() -> Non ) for index, role in enumerate(("key", "value")) } - properties = callback.configure([SimpleNamespace(layers={0: hot})]) - callback.encode_cold_pages(0, 0x3000, 0x4000, 4096, 0x5000) - callback.decode_cold_pages(0, 0x3000, 0x4000, 4096, 0x5000) + properties = provider.configure(codec_state, [SimpleNamespace(layers={0: hot})]) + provider.encode_cold_pages(codec_state, 0, 0x3000, 0x4000, 4096, 0x5000) + provider.decode_cold_pages(codec_state, 0, 0x3000, 0x4000, 4096, 0x5000) assert properties[0].cold_page_bytes == 1152 assert properties[0].page_index_location == "host" - metadata = callback._lifecycle_metadata[0] + metadata = codec_state.lifecycle_metadata[0] for operation in (encode, decode): operation.assert_called_once() arguments = operation.call_args.args @@ -372,7 +378,7 @@ def test_callback_forwards_a_4096_page_batch_through_one_custom_op_call() -> Non ) -def test_callback_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: +def test_codec_state_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: native, _ = _native() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( @@ -384,7 +390,7 @@ def test_callback_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: ) with torch.device("meta"): - metadata = _configure_callback(native, raw_bytes=2048) + metadata = _configure_default_lifecycle(native, raw_bytes=2048) for tensor, dtype, shape in ( (metadata.wide, torch.int64, (256, 6)), (metadata.integers, torch.int32, (256, 5)), @@ -396,7 +402,7 @@ def test_callback_metadata_stays_on_cpu_with_non_cpu_default_device() -> None: assert tensor.is_contiguous() -def test_callback_rejects_invalid_resolved_hot_buffers() -> None: +def test_provider_rejects_invalid_resolved_hot_buffers() -> None: native, _ = _native() with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( @@ -407,7 +413,8 @@ def test_callback_rejects_invalid_resolved_hot_buffers() -> None: head_dim_per_layer=(16,), ) - callback = _callback(native) + provider = _provider(native) + codec_state = _codec_state(native) def hot(raw_base: int = 0x1000, raw_bytes: int = 2048) -> SimpleNamespace: return SimpleNamespace( @@ -417,20 +424,22 @@ def hot(raw_base: int = 0x1000, raw_bytes: int = 2048) -> SimpleNamespace: ) with pytest.raises(ValueError, match="roles do not match"): - callback.configure( - [SimpleNamespace(layers={0: {"key": hot(), "value": hot(), "extra": hot()}})] + provider.configure( + codec_state, + [SimpleNamespace(layers={0: {"key": hot(), "value": hot(), "extra": hot()}})], ) with pytest.raises(ValueError, match="size does not match"): - callback.configure( - [SimpleNamespace(layers={0: {"key": hot(raw_bytes=32), "value": hot()}})] + provider.configure( + codec_state, [SimpleNamespace(layers={0: {"key": hot(raw_bytes=32), "value": hot()}})] ) with pytest.raises(ValueError, match="16-byte aligned"): - callback.configure( - [SimpleNamespace(layers={0: {"key": hot(raw_base=0x1001), "value": hot()}})] + provider.configure( + codec_state, + [SimpleNamespace(layers={0: {"key": hot(raw_base=0x1001), "value": hot()}})], ) -def test_callback_rejects_more_than_256_lifecycle_buffers() -> None: +def test_provider_rejects_more_than_256_lifecycle_buffers() -> None: native, _ = _native() layers = tuple( AttentionLayerConfig( @@ -596,7 +605,7 @@ def test_hybrid_codec_skips_ssm_layers_and_ssm_only_rank_is_lossless(tmp_path): ) assert result is codec - assert _callback(native).layer_ids == () + assert _codec_state(native).layer_ids == () def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): @@ -628,7 +637,7 @@ def test_mla_key_only_layout_with_index_key_uses_identity_scales(tmp_path): assert layout.layer_id == 0 assert [buffer.role for buffer in layout.buffers] == ["key", "index_key"] assert [buffer.scales is not None for buffer in layout.buffers] == [True, False] - assert _callback(native)._runtime_type == 1 + assert _codec_state(native).runtime_type == 1 assert layout.num_kv_heads == 1 assert layout.tokens_per_page == 64 assert layout.head_dim == 576 @@ -774,7 +783,7 @@ def test_fp8_runtime_uses_modelopt_nvfp4_scales(tmp_path): ) layout = _layouts(native)[0] - assert _callback(native)._runtime_type == 2 + assert _codec_state(native).runtime_type == 2 assert [buffer.scales.nvfp4_orig_quant for buffer in layout.buffers] == [ 2.0, 4.0, From 9a9bf5773e95860e7daed348486803358b32615f Mon Sep 17 00:00:00 2001 From: Hudayday <32944717+Hudayday@users.noreply.github.com> Date: Thu, 27 Aug 2026 15:51:38 -0700 Subject: [PATCH 23/29] [None][fix] Add complete cold-page codec license headers Signed-off-by: Hudayday <32944717+Hudayday@users.noreply.github.com> --- .../kv_cache_compression/nativeColdPageCodec.cpp | 12 ++++++++++++ .../kv_cache_compression/nativeColdPageCodec.h | 12 ++++++++++++ 2 files changed, 24 insertions(+) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp index aae08e5e2d82..1274e9bb6094 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp @@ -2,6 +2,18 @@ * 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 "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h index 6c09a3ffb93d..34e6515cd295 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h @@ -2,6 +2,18 @@ * 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. */ #pragma once From 559b6bc41cd8bbf426ba0cfbc1520ff28325f779 Mon Sep 17 00:00:00 2001 From: tianruih Date: Fri, 28 Aug 2026 21:43:38 -0700 Subject: [PATCH 24/29] [None][refactor] Address cold-page review feedback on build layout and HostMem forwarding - Register the NVFP4 cold-page launchers through nanobind next to the cold-page codec bindings and drop the Torch-op layer, so the kernel .cu compiles once into the kernels library. - Move cpp/tensorrt_llm/kv_cache_compression under batch_manager next to kv_cache_manager_v2 and fold it into the batch_manager static library; move coldPageCodecTest into the batch_manager test group. - Make HostMem registration a virtual codec capability and forward it through NativeColdPageCodec to the embedded lossless codec, so SSM/GDN fallback lifecycles keep the registration-boundary copy split; add forwarding regression tests. - Run the NVFP4 kernel unit test on B200: new DGX_B200-CPP-1 pre-merge stage and l0_b200.yml entry. Signed-off-by: tianruih --- cpp/tensorrt_llm/CMakeLists.txt | 2 - cpp/tensorrt_llm/batch_manager/CMakeLists.txt | 4 ++ .../kv_cache_compression/CMakeLists.txt | 18 ++++++ .../nativeColdPageCodec.cpp | 14 ++++- .../nativeColdPageCodec.h | 6 ++ .../kv_cache_manager_v2/coldPageCodec.cpp | 24 ++++++-- .../kv_cache_manager_v2/coldPageCodec.h | 10 +++- cpp/tensorrt_llm/kernels/CMakeLists.txt | 3 - .../kernels/nvfp4ColdPageKernels.cu | 57 ------------------- .../kv_cache_compression/CMakeLists.txt | 9 --- .../nanobind/kvCacheCompression/bindings.cpp | 40 ++++++++++++- cpp/tensorrt_llm/thop/CMakeLists.txt | 6 +- .../unit_tests/batch_manager/CMakeLists.txt | 16 ++++++ .../coldPageCodecTest.cpp | 32 ++++++++++- .../kv_cache_compression/CMakeLists.txt | 18 ------ jenkins/L0_Test.groovy | 1 + .../nvfp4_quantization.py | 8 ++- .../test_lists/test-db/l0_b200.yml | 16 ++++++ .../test_quantization_for_cold_page.py | 12 ++-- 19 files changed, 184 insertions(+), 112 deletions(-) create mode 100644 cpp/tensorrt_llm/batch_manager/kv_cache_compression/CMakeLists.txt rename cpp/tensorrt_llm/{ => batch_manager}/kv_cache_compression/nativeColdPageCodec.cpp (95%) rename cpp/tensorrt_llm/{ => batch_manager}/kv_cache_compression/nativeColdPageCodec.h (93%) delete mode 100644 cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt rename cpp/tests/unit_tests/{kv_cache_compression => batch_manager}/coldPageCodecTest.cpp (88%) diff --git a/cpp/tensorrt_llm/CMakeLists.txt b/cpp/tensorrt_llm/CMakeLists.txt index 71a88308be36..348c60fa787f 100644 --- a/cpp/tensorrt_llm/CMakeLists.txt +++ b/cpp/tensorrt_llm/CMakeLists.txt @@ -144,7 +144,6 @@ add_subdirectory(common) add_subdirectory(kernels) add_subdirectory(layers) add_subdirectory(runtime) -add_subdirectory(kv_cache_compression) set(BATCH_MANAGER_TARGET tensorrt_llm_batch_manager_static) set(BATCH_MANAGER_TARGET_ARCH ${TARGET_ARCH}) @@ -200,7 +199,6 @@ set(TRTLLM_LINK_LIBS cute_dsl_src layers_src runtime_src - kv_cache_compression_src compressorKernels_src mhcKernels_src userbuffers_src diff --git a/cpp/tensorrt_llm/batch_manager/CMakeLists.txt b/cpp/tensorrt_llm/batch_manager/CMakeLists.txt index f61e58e16b28..f04647f9d059 100644 --- a/cpp/tensorrt_llm/batch_manager/CMakeLists.txt +++ b/cpp/tensorrt_llm/batch_manager/CMakeLists.txt @@ -80,6 +80,10 @@ endif() include(${CMAKE_CURRENT_SOURCE_DIR}/kv_cache_manager_v2/CMakeLists.txt) list(APPEND SRCS ${KV_CACHE_MANAGER_V2_SRCS}) +# Include the KVCM cold-page compression codec bridge sources. +include(${CMAKE_CURRENT_SOURCE_DIR}/kv_cache_compression/CMakeLists.txt) +list(APPEND SRCS ${KV_CACHE_COMPRESSION_SRCS}) + add_library(${BATCH_MANAGER_STATIC_TARGET} STATIC ${SRCS}) target_include_directories( ${BATCH_MANAGER_STATIC_TARGET} diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/CMakeLists.txt new file mode 100644 index 000000000000..12018b632383 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/CMakeLists.txt @@ -0,0 +1,18 @@ +# 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. + +# Sources for the KVCM cold-page compression codec bridge. These are added to +# the tensorrt_llm_batch_manager_static target by the parent CMakeLists.txt. +set(KV_CACHE_COMPRESSION_SRCS kv_cache_compression/nativeColdPageCodec.cpp) diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp similarity index 95% rename from cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp rename to cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp index 1274e9bb6094..e0dccd18f418 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp @@ -16,8 +16,9 @@ * limitations under the License. */ -#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/logger.h" #include @@ -187,6 +188,17 @@ kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroup return state == nullptr ? kv::PageIndexLocation::kBadLocation : state->pageIndexLocation; } +bool NativeColdPageCodec::needsHostMemRegistration() const noexcept +{ + return mLosslessCodec != nullptr && mLosslessCodec->needsHostMemRegistration(); +} + +void NativeColdPageCodec::registerHostMem(kv::HostMem const* memory) +{ + TLLM_CHECK(mLosslessCodec != nullptr); + mLosslessCodec->registerHostMem(memory); +} + bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { diff --git a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h similarity index 93% rename from cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h rename to cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h index 34e6515cd295..39d72950242e 100644 --- a/cpp/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h @@ -71,6 +71,12 @@ class NativeColdPageCodec : public kv::IKvCacheColdPageCodec [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept final; + //! Both forward to the embedded lossless codec so fallback lifecycles keep the batched-copy + //! registration-boundary workaround. + [[nodiscard]] bool needsHostMemRegistration() const noexcept final; + + void registerHostMem(kv::HostMem const* memory) final; + bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept final; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp index aa93eb6d1374..0232009c4ea2 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp @@ -55,7 +55,12 @@ class ConcatKvCacheColdPageCodec final : public IKvCacheColdPageCodec TypedVec copyPlans; }; - void registerHostMem(HostMem const* memory) + [[nodiscard]] bool needsHostMemRegistration() const noexcept override + { + return HostMem::shouldUseChunkedRegistration(); + } + + void registerHostMem(HostMem const* memory) override { TLLM_CHECK( memory != nullptr && std::find(mHostMemories.begin(), mHostMemories.end(), memory) == mHostMemories.end()); @@ -335,15 +340,12 @@ namespace detail bool needsHostMemRegistration(IKvCacheColdPageCodec const& codec) noexcept { - return HostMem::shouldUseChunkedRegistration() - && dynamic_cast(&codec) != nullptr; + return codec.needsHostMemRegistration(); } void registerHostMem(IKvCacheColdPageCodec& codec, HostMem const* memory) { - auto* concatCodec = dynamic_cast(&codec); - TLLM_CHECK(concatCodec != nullptr); - concatCodec->registerHostMem(memory); + codec.registerHostMem(memory); } } // namespace detail @@ -356,6 +358,16 @@ LayerGroupId IKvCacheColdPageCodec::getBatchingLayerGroupId(LayerGroupId layerGr return layerGroupId; } +bool IKvCacheColdPageCodec::needsHostMemRegistration() const noexcept +{ + return false; +} + +void IKvCacheColdPageCodec::registerHostMem(HostMem const* /*memory*/) +{ + TLLM_CHECK_WITH_INFO(false, "This cold-page codec does not accept HostMem registration"); +} + std::unique_ptr createDefaultKvCacheColdPageCodec() { return std::make_unique(); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h index 0b3cb1c6a533..84b2483e9a5e 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h @@ -106,6 +106,14 @@ class IKvCacheColdPageCodec //! PageIndexLocation::kBadLocation on failure or for an unknown layer group. [[nodiscard]] virtual PageIndexLocation queryPageIndexLocation(LayerGroupId layerGroupId) const noexcept = 0; + //! Returns whether this codec needs HostMem spans for the batched-copy registration-boundary workaround. + //! Defaults to false. + [[nodiscard]] virtual bool needsHostMemRegistration() const noexcept; + + //! Registers one KVCM-owned HostMem span (non-owning). Wrapper codecs must forward both methods to their + //! inner codec. + virtual void registerHostMem(HostMem const* memory); + //! Encodes hot pages into cold pages. //! //! The cold base pointer is GPU-accessible. The index-array location is selected by queryPageIndexLocation(). Host @@ -131,7 +139,7 @@ class IKvCacheColdPageCodec namespace detail { -//! Returns whether the default codec needs HostMem spans for the batched-copy registration-boundary workaround. +//! Returns whether the codec needs HostMem spans for the batched-copy registration-boundary workaround. [[nodiscard]] bool needsHostMemRegistration(IKvCacheColdPageCodec const& codec) noexcept; //! Registers KVCM-owned pinned memory after needsHostMemRegistration() returns true. diff --git a/cpp/tensorrt_llm/kernels/CMakeLists.txt b/cpp/tensorrt_llm/kernels/CMakeLists.txt index 34bc62fcfa31..c6c6fe4a292a 100644 --- a/cpp/tensorrt_llm/kernels/CMakeLists.txt +++ b/cpp/tensorrt_llm/kernels/CMakeLists.txt @@ -67,9 +67,6 @@ list(FILTER SRC_CPP EXCLUDE REGEX "mhcKernels/.*") list(FILTER SRC_CU EXCLUDE REGEX "mhcKernels/.*") list(FILTER SRC_CPP EXCLUDE REGEX "compressorKernels/.*") list(FILTER SRC_CU EXCLUDE REGEX "compressorKernels/.*") -# This TU registers its Torch op in th_common; exclude it here to avoid -# duplicate symbols. -list(FILTER SRC_CU EXCLUDE REGEX "nvfp4ColdPageKernels\\.cu$") # Marlin is built as its own architecture-scoped OBJECT library below. list(FILTER SRC_CPP EXCLUDE REGEX "marlin/.*") list(FILTER SRC_CU EXCLUDE REGEX "marlin/.*") diff --git a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu index 90a4a28fe182..57d2d3b9433f 100644 --- a/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu +++ b/cpp/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu @@ -33,10 +33,6 @@ #include #include -#if defined(TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) -#include -#endif - TRTLLM_NAMESPACE_BEGIN namespace kernels @@ -875,56 +871,3 @@ void invokeNvfp4ColdPageDecode(void const* pages, std::size_t numPages, std::int } // namespace kernels TRTLLM_NAMESPACE_END - -#if defined(TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) -namespace -{ - -void nvfp4ColdPageEncode(std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, - std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, - std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) -{ - tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( - reinterpret_cast(static_cast(pageIndices)), static_cast(numPages), - reinterpret_cast(static_cast(wide)), - reinterpret_cast(static_cast(integers)), - reinterpret_cast(static_cast(scales)), static_cast(numBuffers), - static_cast(maxHalfGroupsPerTile), static_cast(coldPageBytes), - static_cast(runtimeType), - reinterpret_cast(static_cast(coldBase)), - reinterpret_cast(static_cast(stream))); -} - -void nvfp4ColdPageDecode(std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, - std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, - std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) -{ - tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( - reinterpret_cast(static_cast(pageIndices)), static_cast(numPages), - reinterpret_cast(static_cast(wide)), - reinterpret_cast(static_cast(integers)), - reinterpret_cast(static_cast(scales)), static_cast(numBuffers), - static_cast(maxHalfGroupsPerTile), static_cast(coldPageBytes), - static_cast(runtimeType), - reinterpret_cast(static_cast(coldBase)), - reinterpret_cast(static_cast(stream))); -} - -} // namespace - -TORCH_LIBRARY_FRAGMENT(trtllm, module) -{ - module.def( - "nvfp4_cold_page_encode(int page_indices, int num_pages, int wide, int integers, int scales, int num_buffers, " - "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int stream) -> ()"); - module.def( - "nvfp4_cold_page_decode(int page_indices, int num_pages, int wide, int integers, int scales, int num_buffers, " - "int max_half_groups_per_tile, int cold_page_bytes, int runtime_type, int cold_base, int stream) -> ()"); -} - -TORCH_LIBRARY_IMPL(trtllm, CompositeExplicitAutograd, module) -{ - module.impl("nvfp4_cold_page_encode", &nvfp4ColdPageEncode); - module.impl("nvfp4_cold_page_decode", &nvfp4ColdPageDecode); -} -#endif diff --git a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt b/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt deleted file mode 100644 index 2a95aba2807c..000000000000 --- a/cpp/tensorrt_llm/kv_cache_compression/CMakeLists.txt +++ /dev/null @@ -1,9 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. -# All rights reserved. SPDX-License-Identifier: Apache-2.0 - -add_library(kv_cache_compression_src OBJECT nativeColdPageCodec.cpp) -target_include_directories( - kv_cache_compression_src - PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) -set_property(TARGET kv_cache_compression_src PROPERTY POSITION_INDEPENDENT_CODE - ON) diff --git a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp index 8d4f7f42568c..352721a7d1ca 100644 --- a/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/kvCacheCompression/bindings.cpp @@ -17,7 +17,8 @@ */ #include "bindings.h" -#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/kernels/nvfp4ColdPageKernels.h" #include #include @@ -163,6 +164,43 @@ void initBindings(nb::module_& module) [](nb::handle provider, nb::handle codecState) -> std::unique_ptr { return std::make_unique(provider, codecState); }, nb::arg("provider"), nb::arg("codec_state")); + + // The Python provider traffics in raw KVCM addresses, so the launcher trampolines take scalar integers. + module.def("nvfp4_cold_page_encode", + [](std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, + std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, + std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) + { + tensorrt_llm::kernels::invokeNvfp4ColdPageEncode( + reinterpret_cast(static_cast(pageIndices)), + static_cast(numPages), + reinterpret_cast(static_cast(wide)), + reinterpret_cast(static_cast(integers)), + reinterpret_cast(static_cast(scales)), + static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), + static_cast(coldPageBytes), + static_cast(runtimeType), + reinterpret_cast(static_cast(coldBase)), + reinterpret_cast(static_cast(stream))); + }); + + module.def("nvfp4_cold_page_decode", + [](std::int64_t pageIndices, std::int64_t numPages, std::int64_t wide, std::int64_t integers, + std::int64_t scales, std::int64_t numBuffers, std::int64_t maxHalfGroupsPerTile, std::int64_t coldPageBytes, + std::int64_t runtimeType, std::int64_t coldBase, std::int64_t stream) + { + tensorrt_llm::kernels::invokeNvfp4ColdPageDecode( + reinterpret_cast(static_cast(pageIndices)), + static_cast(numPages), + reinterpret_cast(static_cast(wide)), + reinterpret_cast(static_cast(integers)), + reinterpret_cast(static_cast(scales)), + static_cast(numBuffers), static_cast(maxHalfGroupsPerTile), + static_cast(coldPageBytes), + static_cast(runtimeType), + reinterpret_cast(static_cast(coldBase)), + reinterpret_cast(static_cast(stream))); + }); } } // namespace tensorrt_llm::nanobind::kv_cache_compression diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index 13432ae7e6bc..cdf0a0228fac 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -146,12 +146,8 @@ add_library( trtllmGenQKVProcessOp.cpp inplaceSliceCopyOp.cpp mhcOp.cpp - compressorOp.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu) + compressorOp.cpp) set_property(TARGET th_common PROPERTY POSITION_INDEPENDENT_CODE ON) -set_source_files_properties( - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu - PROPERTIES COMPILE_DEFINITIONS TRTLLM_ENABLE_NVFP4_COLD_PAGE_TORCH_OP) target_link_libraries( th_common PRIVATE ${TORCH_LIBRARIES} th_utils ${Python3_LIBRARIES} ${SHARED_TARGET} pg_utils) diff --git a/cpp/tests/unit_tests/batch_manager/CMakeLists.txt b/cpp/tests/unit_tests/batch_manager/CMakeLists.txt index 468542c1dc03..0dc35bb3f753 100644 --- a/cpp/tests/unit_tests/batch_manager/CMakeLists.txt +++ b/cpp/tests/unit_tests/batch_manager/CMakeLists.txt @@ -38,6 +38,22 @@ add_gtest(kvCacheManagerV2ColdPageCopyTest kvCacheManagerV2ColdPageCopyTest.cu) target_include_directories( kvCacheManagerV2ColdPageCopyTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) +set(COLD_PAGE_CODEC_TEST_SRC + coldPageCodecTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) +add_gtest(coldPageCodecTest "${COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) +target_link_libraries(coldPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) +target_include_directories( + coldPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) add_gtest(kvCacheManagerV2TypedIndexTest kvCacheManagerV2TypedIndexTest.cpp) target_include_directories( kvCacheManagerV2TypedIndexTest diff --git a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp similarity index 88% rename from cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp rename to cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp index b853e7bc34be..ca0408559ec4 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp @@ -4,7 +4,9 @@ * SPDX-License-Identifier: Apache-2.0 */ -#include "tensorrt_llm/kv_cache_compression/nativeColdPageCodec.h" +#include "tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h" + +#include "kv_cache_manager_v2/utils/hostMem.h" #include @@ -189,6 +191,34 @@ TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{1}), kv::PageIndexLocation::kHost); } +TEST(NativeColdPageCodecTest, MirrorsLosslessFallbackHostMemRegistrationCapability) +{ + RecordingCodec codec{{0}}; + EXPECT_FALSE(codec.needsHostMemRegistration()); + + std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + auto const defaultCodec = kv::createDefaultKvCacheColdPageCodec(); + EXPECT_EQ(codec.needsHostMemRegistration(), defaultCodec->needsHostMemRegistration()); + EXPECT_EQ(kv::detail::needsHostMemRegistration(codec), codec.needsHostMemRegistration()); +} + +TEST(NativeColdPageCodecTest, ForwardsHostMemRegistrationToLosslessFallback) +{ + int deviceCount = 0; + if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) + { + GTEST_SKIP() << "HostMem registration requires a CUDA device"; + } + + RecordingCodec codec{{0}}; + std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + + kv::HostMem hostMem{4096}; + kv::detail::registerHostMem(codec, &hostMem); +} + TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) { { diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt index dbfe8cc5e023..00436d7058b0 100644 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt @@ -1,24 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. # All rights reserved. SPDX-License-Identifier: Apache-2.0 -set(COLD_PAGE_CODEC_TEST_SRC - coldPageCodecTest.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kv_cache_compression/nativeColdPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) - -add_gtest(coldPageCodecTest "${COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) -target_link_libraries(coldPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) -target_include_directories( - coldPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) - set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC ${PROJECT_SOURCE_DIR}/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu diff --git a/jenkins/L0_Test.groovy b/jenkins/L0_Test.groovy index 3eb41f075c15..2ecb733bdadd 100644 --- a/jenkins/L0_Test.groovy +++ b/jenkins/L0_Test.groovy @@ -5934,6 +5934,7 @@ def launchTestJobs(pipeline, testFilter, globalVars) "DGX_H100-4_GPUs-PyTorch-Others-2": ["auto:dgx-h100-x4", "l0_dgx_h100", 2, 2, 4], "DGX_H100-4_GPUs-PyTorch-Ray-1": ["auto:dgx-h100-x4", "l0_dgx_h100", 1, 1, 4], "DGX_H100-4_GPUs-PyTorch-Post-Merge-1": ["auto:dgx-h100-x4", "l0_dgx_h100", 1, 1, 4], + "DGX_B200-CPP-1": ["auto:dgx-b200-flex", "l0_b200", 1, 1, 1, 1, true], "DGX_B200-PyTorch-1": ["auto:dgx-b200-flex", "l0_b200", 1, 9, 1, 1, true], "DGX_B200-PyTorch-2": ["auto:dgx-b200-flex", "l0_b200", 2, 9, 1, 1, true], "DGX_B200-PyTorch-3": ["auto:dgx-b200-flex", "l0_b200", 3, 9, 1, 1, true], diff --git a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py index 44853ef736e0..429a904336b1 100644 --- a/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py +++ b/tensorrt_llm/_torch/kv_cache_compression/quantization_for_cold_page/nvfp4_quantization.py @@ -370,8 +370,10 @@ def encode_cold_pages( num_pages: int, stream: int, ) -> None: + from tensorrt_llm.bindings.internal import kv_cache_compression as native + metadata = codec_state.lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_encode( + native.nvfp4_cold_page_encode( page_indices, num_pages, metadata.wide.data_ptr(), @@ -394,8 +396,10 @@ def decode_cold_pages( num_pages: int, stream: int, ) -> None: + from tensorrt_llm.bindings.internal import kv_cache_compression as native + metadata = codec_state.lifecycle_metadata[lifecycle_index] - torch.ops.trtllm.nvfp4_cold_page_decode( + native.nvfp4_cold_page_decode( page_indices, num_pages, metadata.wide.data_ptr(), diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index eb98bc2a7b27..b3edb4073a14 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -1,5 +1,21 @@ version: 0.0.1 l0_b200: +- condition: + ranges: + system_gpu_count: + gte: 1 + lte: 1 + wildcards: + gpu: + - '*b100*' + - '*b200*' + linux_distribution_name: ubuntu* + terms: + stage: pre_merge + backend: cpp + tests: + # ------------- CPP tests --------------- + - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-100] - condition: ranges: system_gpu_count: diff --git a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py index 17374cd034f9..3d965a438241 100644 --- a/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py +++ b/tests/unittest/_torch/kv_cache_compression/test_quantization_for_cold_page.py @@ -78,6 +78,8 @@ def _native() -> tuple[SimpleNamespace, MagicMock]: ColdPageLifecycleProperties=lambda: SimpleNamespace(), ColdPageIndexLocation=SimpleNamespace(HOST="host"), create_python_cold_page_codec=MagicMock(return_value=codec), + nvfp4_cold_page_encode=MagicMock(), + nvfp4_cold_page_decode=MagicMock(), ) return module, codec @@ -328,14 +330,12 @@ def test_provider_creates_independent_state_per_kv_cache_manager() -> None: assert target_state.layer_ids == draft_state.layer_ids == (0,) -def test_provider_forwards_a_4096_page_batch_through_one_custom_op_call() -> None: +def test_provider_forwards_a_4096_page_batch_through_one_native_call() -> None: native, _ = _native() + encode = native.nvfp4_cold_page_encode + decode = native.nvfp4_cold_page_decode - with ( - patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native), - patch.object(torch.ops.trtllm, "nvfp4_cold_page_encode", create=True) as encode, - patch.object(torch.ops.trtllm, "nvfp4_cold_page_decode", create=True) as decode, - ): + with patch("tensorrt_llm.bindings.internal.kv_cache_compression", new=native): _manager().create_cold_page_codec( _cache_config((0, "attention")), runtime_dtype=DataType.BF16, From 58ec834c23d8e30aa68d5ff6a1984f50a027544d Mon Sep 17 00:00:00 2001 From: tianruih Date: Sat, 29 Aug 2026 08:19:33 -0700 Subject: [PATCH 25/29] [None][fix] Fix a sign comparison in the cold-page codec bridge The layer count from std::count_if is signed; comparing it against map::size() trips -Werror=sign-compare now that the file compiles inside the batch_manager static library. Signed-off-by: tianruih --- .../kv_cache_compression/nativeColdPageCodec.cpp | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp index e0dccd18f418..12065f0735cb 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp @@ -93,8 +93,8 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG for (auto const& variant : gpuDesc.slotDesc.variants) { auto resolved = resolveLifecycle(gpuDesc, variant); - auto const providerLayerCount = std::count_if(resolved.layers.begin(), resolved.layers.end(), - [this](auto const& layer) { return mLayerIds.count(layer.first) != 0U; }); + auto const providerLayerCount = static_cast(std::count_if(resolved.layers.begin(), + resolved.layers.end(), [this](auto const& layer) { return mLayerIds.count(layer.first) != 0U; })); LayerGroupState state; if (providerLayerCount == 0U) From e2c8284ee78d25a6faae71105acb3da776161f0e Mon Sep 17 00:00:00 2001 From: tianruih Date: Sat, 29 Aug 2026 14:43:26 -0700 Subject: [PATCH 26/29] [None][fix] Initialize a CUDA context in the HostMem forwarding test HostMem pins through the CUDA driver API, which requires a current context; cudaGetDeviceCount alone does not create one. Signed-off-by: tianruih --- cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp | 2 ++ 1 file changed, 2 insertions(+) diff --git a/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp index ca0408559ec4..cdd4edd60ab5 100644 --- a/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp @@ -210,6 +210,8 @@ TEST(NativeColdPageCodecTest, ForwardsHostMemRegistrationToLosslessFallback) { GTEST_SKIP() << "HostMem registration requires a CUDA device"; } + // HostMem pins through the CUDA driver API, which needs a current context. + ASSERT_EQ(cudaFree(nullptr), cudaSuccess); RecordingCodec codec{{0}}; std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; From a7156ed4527fc7d8627ce59756f4720e70c7a6b6 Mon Sep 17 00:00:00 2001 From: tianruih Date: Sun, 30 Aug 2026 22:11:38 -0700 Subject: [PATCH 27/29] [None][infra] Build only the standalone KV-cache compression gtests in CI The DGX_B200-CPP-1 stage timed out intermittently because the shared build_google_tests fixture rebuilds the full library and every gtest (~55 min on B200 runners) while the stage only runs two NO_TLLM_LINKAGE binaries (~40 s). Add a dedicated fixture that configures CMake and builds just coldPageCodecTest and nvfp4ColdPageKernelsTest, expose it as test_kv_cache_compression_unit_tests[80|100], and point the A30 and B200 list entries at it. The shared test_unit_tests parametrization is restored to its upstream form. Signed-off-by: tianruih --- tests/integration/defs/cpp/conftest.py | 42 +++++++++++++++++++ tests/integration/defs/cpp/test_unit_tests.py | 35 ++++++++++++++-- .../integration/test_lists/test-db/l0_a30.yml | 2 +- .../test_lists/test-db/l0_b200.yml | 2 +- 4 files changed, 76 insertions(+), 5 deletions(-) diff --git a/tests/integration/defs/cpp/conftest.py b/tests/integration/defs/cpp/conftest.py index 75d79fc463c5..c28875c18860 100644 --- a/tests/integration/defs/cpp/conftest.py +++ b/tests/integration/defs/cpp/conftest.py @@ -173,6 +173,48 @@ def build_google_tests(request, build_type): ) +@pytest.fixture(scope="session") +def build_kv_cache_compression_tests(request, build_type): + """Build only the standalone KV-cache compression gtests. + + Both binaries use NO_TLLM_LINKAGE, so this skips the full-library + build that build_google_tests pays for. + """ + cuda_arch = f"{request.param}-real" + + _logger.info(f"Using CUDA arch: {cuda_arch}") + + build_trt_llm( + build_type=build_type, + cuda_architectures=cuda_arch, + job_count=12, + use_ccache=True, + generator="Ninja", + nixl_root="/opt/nvidia/nvda_nixl", + skip_building_wheel=True, + configure_only=True, + ) + + build_dir = _cpp.find_build_dir(build_type) + _cpp.run_command( + [ + "cmake", + "--build", + str(build_dir), + "--config", + build_type, + "--parallel", + "12", + "--target", + "coldPageCodecTest", + "nvfp4ColdPageKernelsTest", + ], + cwd=build_dir, + env=_os.environ, + timeout=1800, + ) + + @pytest.fixture(scope="function", autouse=True) def keep_log_files(build_dir): """Backup previous cpp test results when run multiple ctest invocations.""" diff --git a/tests/integration/defs/cpp/test_unit_tests.py b/tests/integration/defs/cpp/test_unit_tests.py index 9989e70a355c..78f894c465a5 100644 --- a/tests/integration/defs/cpp/test_unit_tests.py +++ b/tests/integration/defs/cpp/test_unit_tests.py @@ -4,11 +4,11 @@ import pytest -@pytest.mark.parametrize("build_google_tests", ["80", "86", "89", "90", "100"], +@pytest.mark.parametrize("build_google_tests", ["80", "86", "89", "90"], indirect=True) @pytest.mark.parametrize("test_group", [ - "batch_manager", "common", "executor", "kernels", "kv_cache_compression", - "layers", "runtime", "thop" + "batch_manager", "common", "executor", "kernels", "layers", "runtime", + "thop" ]) def test_unit_tests(build_google_tests, test_group, build_dir, lora_setup): @@ -35,3 +35,32 @@ def test_unit_tests(build_google_tests, test_group, build_dir, lora_setup): env=cpp_env, timeout=2700, parallel=parallel) + + +@pytest.mark.parametrize("build_kv_cache_compression_tests", ["80", "100"], + indirect=True) +def test_kv_cache_compression_unit_tests(build_kv_cache_compression_tests, + build_dir): + + xml_name = "results-unit-tests-kv_cache_compression.xml" + + ctest_command = [ + "ctest", + "--output-on-failure", + "--test-dir", + f"{build_dir}/tests/unit_tests/kv_cache_compression", + "--output-junit", + f"{build_dir}/{xml_name}", + ] + + parallel = _cpp.default_test_parallel + if parallel_override := _os.environ.get("LLM_TEST_PARALLEL_OVERRIDE", None): + parallel = int(parallel_override) + + cpp_env = {**_os.environ} + + _cpp.parallel_run_ctest(ctest_command, + cwd=build_dir, + env=cpp_env, + timeout=2700, + parallel=parallel) diff --git a/tests/integration/test_lists/test-db/l0_a30.yml b/tests/integration/test_lists/test-db/l0_a30.yml index 7b124aa10ca3..3114a0a81a59 100644 --- a/tests/integration/test_lists/test-db/l0_a30.yml +++ b/tests/integration/test_lists/test-db/l0_a30.yml @@ -40,7 +40,7 @@ l0_a30: tests: # ------------- CPP tests --------------- - cpp/test_unit_tests.py::test_unit_tests[batch_manager-80] - - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-80] + - cpp/test_unit_tests.py::test_kv_cache_compression_unit_tests[80] - condition: ranges: system_gpu_count: diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index e9d8e5cb5b33..de19938875eb 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -15,7 +15,7 @@ l0_b200: backend: cpp tests: # ------------- CPP tests --------------- - - cpp/test_unit_tests.py::test_unit_tests[kv_cache_compression-100] + - cpp/test_unit_tests.py::test_kv_cache_compression_unit_tests[100] - condition: ranges: system_gpu_count: From 30e2a5b3d3dde0a9e9e91ddc0af92270e226d842 Mon Sep 17 00:00:00 2001 From: tianruih Date: Mon, 31 Aug 2026 10:49:33 -0700 Subject: [PATCH 28/29] [None][refactor] Address cold-page test layout and HostMem review feedback - Drop the dedicated kv_cache_compression unit-test directory; the NVFP4 kernel gtest registers under kernels/ and the CI test runs the binary directly with --gtest_output. - Simplify coldPageCodecTest to a plain add_gtest against the shared library now that nativeColdPageCodec.cpp compiles into batch_manager. - Revert the HostMem registration codec virtuals per KVCM owner review; NativeColdPageCodec instead fails closed at configure time when a lossless-fallback lifecycle meets chunked host registration (Linux 6.11-6.13), until KVCM moves those copies to kernels. Signed-off-by: tianruih --- .../nativeColdPageCodec.cpp | 34 +++++++++++------ .../nativeColdPageCodec.h | 6 --- .../kv_cache_manager_v2/coldPageCodec.cpp | 24 +++--------- .../kv_cache_manager_v2/coldPageCodec.h | 10 +---- cpp/tests/unit_tests/CMakeLists.txt | 1 - .../unit_tests/batch_manager/CMakeLists.txt | 15 +------- .../batch_manager/coldPageCodecTest.cpp | 38 ++++--------------- cpp/tests/unit_tests/kernels/CMakeLists.txt | 17 +++++++++ .../kv_cache_compression/CMakeLists.txt | 19 ---------- tests/integration/defs/cpp/conftest.py | 5 +-- tests/integration/defs/cpp/test_unit_tests.py | 31 ++++++--------- 11 files changed, 68 insertions(+), 132 deletions(-) delete mode 100644 cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp index 12065f0735cb..5beed9e998bb 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp @@ -18,7 +18,7 @@ #include "tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h" -#include "tensorrt_llm/common/assert.h" +#include "kv_cache_manager_v2/utils/hostMem.h" #include "tensorrt_llm/common/logger.h" #include @@ -131,6 +131,27 @@ bool NativeColdPageCodec::configure(kv::PoolGroupDesc const* gpuDescs, kv::PoolG throw std::invalid_argument("A provider layer is absent from all GPU descriptors"); } + // Fail closed until KVCM replaces the batched cuMemcpyBatchAsync copies with kernels: on host + // kernels that need chunked pinned-memory registration (Linux 6.11-6.13), the embedded lossless + // codec cannot split its copies at registration boundaries when wrapped by this codec. + bool hasFallbackLifecycle = false; + for (auto const& [lifeCycleId, state] : pendingGroups) + { + static_cast(lifeCycleId); + if (!state.lifecycleIndex) + { + hasFallbackLifecycle = true; + break; + } + } + if (hasFallbackLifecycle && kv::HostMem::shouldUseChunkedRegistration()) + { + throw std::invalid_argument( + "Cold-page compression is not supported for models with lossless-fallback lifecycles (SSM/GDN) on " + "this host kernel: chunked pinned-memory registration (Linux 6.11-6.13) breaks the fallback codec's " + "batched copies. Disable KV cache compression for this model or use a different host kernel."); + } + auto const properties = configureProvider(providerLifecycles); if (properties.size() != providerLifecycles.size()) { @@ -188,17 +209,6 @@ kv::PageIndexLocation NativeColdPageCodec::queryPageIndexLocation(kv::LayerGroup return state == nullptr ? kv::PageIndexLocation::kBadLocation : state->pageIndexLocation; } -bool NativeColdPageCodec::needsHostMemRegistration() const noexcept -{ - return mLosslessCodec != nullptr && mLosslessCodec->needsHostMemRegistration(); -} - -void NativeColdPageCodec::registerHostMem(kv::HostMem const* memory) -{ - TLLM_CHECK(mLosslessCodec != nullptr); - mLosslessCodec->registerHostMem(memory); -} - bool NativeColdPageCodec::encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept { diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h index 39d72950242e..34e6515cd295 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.h @@ -71,12 +71,6 @@ class NativeColdPageCodec : public kv::IKvCacheColdPageCodec [[nodiscard]] kv::PageIndexLocation queryPageIndexLocation(kv::LayerGroupId layerGroupId) const noexcept final; - //! Both forward to the embedded lossless codec so fallback lifecycles keep the batched-copy - //! registration-boundary workaround. - [[nodiscard]] bool needsHostMemRegistration() const noexcept final; - - void registerHostMem(kv::HostMem const* memory) final; - bool encode(kv::LayerGroupId layerGroupId, void* dstBasePtr, kv::PageIndexPair const* pageIndices, std::size_t numBasePages, cudaStream_t stream) noexcept final; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp index 0232009c4ea2..aa93eb6d1374 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp @@ -55,12 +55,7 @@ class ConcatKvCacheColdPageCodec final : public IKvCacheColdPageCodec TypedVec copyPlans; }; - [[nodiscard]] bool needsHostMemRegistration() const noexcept override - { - return HostMem::shouldUseChunkedRegistration(); - } - - void registerHostMem(HostMem const* memory) override + void registerHostMem(HostMem const* memory) { TLLM_CHECK( memory != nullptr && std::find(mHostMemories.begin(), mHostMemories.end(), memory) == mHostMemories.end()); @@ -340,12 +335,15 @@ namespace detail bool needsHostMemRegistration(IKvCacheColdPageCodec const& codec) noexcept { - return codec.needsHostMemRegistration(); + return HostMem::shouldUseChunkedRegistration() + && dynamic_cast(&codec) != nullptr; } void registerHostMem(IKvCacheColdPageCodec& codec, HostMem const* memory) { - codec.registerHostMem(memory); + auto* concatCodec = dynamic_cast(&codec); + TLLM_CHECK(concatCodec != nullptr); + concatCodec->registerHostMem(memory); } } // namespace detail @@ -358,16 +356,6 @@ LayerGroupId IKvCacheColdPageCodec::getBatchingLayerGroupId(LayerGroupId layerGr return layerGroupId; } -bool IKvCacheColdPageCodec::needsHostMemRegistration() const noexcept -{ - return false; -} - -void IKvCacheColdPageCodec::registerHostMem(HostMem const* /*memory*/) -{ - TLLM_CHECK_WITH_INFO(false, "This cold-page codec does not accept HostMem registration"); -} - std::unique_ptr createDefaultKvCacheColdPageCodec() { return std::make_unique(); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h index 84b2483e9a5e..0b3cb1c6a533 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.h @@ -106,14 +106,6 @@ class IKvCacheColdPageCodec //! PageIndexLocation::kBadLocation on failure or for an unknown layer group. [[nodiscard]] virtual PageIndexLocation queryPageIndexLocation(LayerGroupId layerGroupId) const noexcept = 0; - //! Returns whether this codec needs HostMem spans for the batched-copy registration-boundary workaround. - //! Defaults to false. - [[nodiscard]] virtual bool needsHostMemRegistration() const noexcept; - - //! Registers one KVCM-owned HostMem span (non-owning). Wrapper codecs must forward both methods to their - //! inner codec. - virtual void registerHostMem(HostMem const* memory); - //! Encodes hot pages into cold pages. //! //! The cold base pointer is GPU-accessible. The index-array location is selected by queryPageIndexLocation(). Host @@ -139,7 +131,7 @@ class IKvCacheColdPageCodec namespace detail { -//! Returns whether the codec needs HostMem spans for the batched-copy registration-boundary workaround. +//! Returns whether the default codec needs HostMem spans for the batched-copy registration-boundary workaround. [[nodiscard]] bool needsHostMemRegistration(IKvCacheColdPageCodec const& codec) noexcept; //! Registers KVCM-owned pinned memory after needsHostMemRegistration() returns true. diff --git a/cpp/tests/unit_tests/CMakeLists.txt b/cpp/tests/unit_tests/CMakeLists.txt index 44291864657c..554be48db1da 100644 --- a/cpp/tests/unit_tests/CMakeLists.txt +++ b/cpp/tests/unit_tests/CMakeLists.txt @@ -22,7 +22,6 @@ if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/executor) endif() add_subdirectory(common) -add_subdirectory(kv_cache_compression) add_subdirectory(kernels) add_subdirectory(multi_gpu) add_subdirectory(layers) diff --git a/cpp/tests/unit_tests/batch_manager/CMakeLists.txt b/cpp/tests/unit_tests/batch_manager/CMakeLists.txt index 0dc35bb3f753..e74699fbf418 100644 --- a/cpp/tests/unit_tests/batch_manager/CMakeLists.txt +++ b/cpp/tests/unit_tests/batch_manager/CMakeLists.txt @@ -38,20 +38,7 @@ add_gtest(kvCacheManagerV2ColdPageCopyTest kvCacheManagerV2ColdPageCopyTest.cu) target_include_directories( kvCacheManagerV2ColdPageCopyTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) -set(COLD_PAGE_CODEC_TEST_SRC - coldPageCodecTest.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCopy.cu - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/coldPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_compression/nativeColdPageCodec.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/cudaDriverWrapper.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) -add_gtest(coldPageCodecTest "${COLD_PAGE_CODEC_TEST_SRC}" NO_TLLM_LINKAGE) -target_link_libraries(coldPageCodecTest PRIVATE CUDA::cuda_driver CUDA::cudart) +add_gtest(coldPageCodecTest coldPageCodecTest.cpp) target_include_directories( coldPageCodecTest PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) add_gtest(kvCacheManagerV2TypedIndexTest kvCacheManagerV2TypedIndexTest.cpp) diff --git a/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp index cdd4edd60ab5..05ba3affa22a 100644 --- a/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/coldPageCodecTest.cpp @@ -184,6 +184,14 @@ TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) RecordingCodec codec{{0}}; std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; + if (kv::HostMem::shouldUseChunkedRegistration()) + { + // Fallback lifecycles fail closed on host kernels with chunked pinned-memory + // registration until KVCM replaces the batched copies with kernels. + EXPECT_FALSE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); + return; + } + ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); EXPECT_EQ(codec.resolved.size(), 1U); EXPECT_EQ(codec.queryColdPageBytes(kv::LayerGroupId{0}), 777U); @@ -191,36 +199,6 @@ TEST(NativeColdPageCodecTest, UnownedLifecycleUsesLosslessFallback) EXPECT_EQ(codec.queryPageIndexLocation(kv::LayerGroupId{1}), kv::PageIndexLocation::kHost); } -TEST(NativeColdPageCodecTest, MirrorsLosslessFallbackHostMemRegistrationCapability) -{ - RecordingCodec codec{{0}}; - EXPECT_FALSE(codec.needsHostMemRegistration()); - - std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; - ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - auto const defaultCodec = kv::createDefaultKvCacheColdPageCodec(); - EXPECT_EQ(codec.needsHostMemRegistration(), defaultCodec->needsHostMemRegistration()); - EXPECT_EQ(kv::detail::needsHostMemRegistration(codec), codec.needsHostMemRegistration()); -} - -TEST(NativeColdPageCodecTest, ForwardsHostMemRegistrationToLosslessFallback) -{ - int deviceCount = 0; - if (cudaGetDeviceCount(&deviceCount) != cudaSuccess || deviceCount == 0) - { - GTEST_SKIP() << "HostMem registration requires a CUDA device"; - } - // HostMem pins through the CUDA driver API, which needs a current context. - ASSERT_EQ(cudaFree(nullptr), cudaSuccess); - - RecordingCodec codec{{0}}; - std::array descs{makeAttentionDesc(), makeLosslessDesc(kv::PoolGroupIndex{1}, kv::LayerGroupId{1})}; - ASSERT_TRUE(codec.configure(descs.data(), kv::PoolGroupIndex{2})); - - kv::HostMem hostMem{4096}; - kv::detail::registerHostMem(codec, &hostMem); -} - TEST(NativeColdPageCodecTest, RejectsMixedMissingAndDuplicateLifecycleMappings) { { diff --git a/cpp/tests/unit_tests/kernels/CMakeLists.txt b/cpp/tests/unit_tests/kernels/CMakeLists.txt index 6b0e5a118211..9a4a4a471387 100644 --- a/cpp/tests/unit_tests/kernels/CMakeLists.txt +++ b/cpp/tests/unit_tests/kernels/CMakeLists.txt @@ -113,3 +113,20 @@ endif() add_gtest(eaglePackDataTest eaglePackDataTest.cpp) add_gtest(sparseKvCacheTest sparseKvCacheTest.cu) add_gtest(prepareCustomMaskTest prepareCustomMaskTest.cpp) + +set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC + nvfp4ColdPageKernelsTest.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu + ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp + ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) +add_gtest(nvfp4ColdPageKernelsTest "${NVFP4_COLD_PAGE_KERNEL_TEST_SRC}" + NO_TLLM_LINKAGE) +target_include_directories( + nvfp4ColdPageKernelsTest + PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) +target_link_libraries(nvfp4ColdPageKernelsTest PRIVATE CUDA::cudart + CUDA::cuda_driver) diff --git a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt b/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt deleted file mode 100644 index 00436d7058b0..000000000000 --- a/cpp/tests/unit_tests/kv_cache_compression/CMakeLists.txt +++ /dev/null @@ -1,19 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. -# All rights reserved. SPDX-License-Identifier: Apache-2.0 - -set(NVFP4_COLD_PAGE_KERNEL_TEST_SRC - ${PROJECT_SOURCE_DIR}/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/kernels/nvfp4ColdPageKernels.cu - ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/hostMem.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/assert.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/envUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/logger.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/stringUtils.cpp - ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/tllmException.cpp) -add_gtest(nvfp4ColdPageKernelsTest "${NVFP4_COLD_PAGE_KERNEL_TEST_SRC}" - NO_TLLM_LINKAGE) -target_include_directories( - nvfp4ColdPageKernelsTest - PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/batch_manager) -target_link_libraries(nvfp4ColdPageKernelsTest PRIVATE CUDA::cudart - CUDA::cuda_driver) diff --git a/tests/integration/defs/cpp/conftest.py b/tests/integration/defs/cpp/conftest.py index c28875c18860..a48b71e10a61 100644 --- a/tests/integration/defs/cpp/conftest.py +++ b/tests/integration/defs/cpp/conftest.py @@ -175,9 +175,9 @@ def build_google_tests(request, build_type): @pytest.fixture(scope="session") def build_kv_cache_compression_tests(request, build_type): - """Build only the standalone KV-cache compression gtests. + """Build only the standalone NVFP4 cold-page kernel gtest. - Both binaries use NO_TLLM_LINKAGE, so this skips the full-library + The binary uses NO_TLLM_LINKAGE, so this skips the full-library build that build_google_tests pays for. """ cuda_arch = f"{request.param}-real" @@ -206,7 +206,6 @@ def build_kv_cache_compression_tests(request, build_type): "--parallel", "12", "--target", - "coldPageCodecTest", "nvfp4ColdPageKernelsTest", ], cwd=build_dir, diff --git a/tests/integration/defs/cpp/test_unit_tests.py b/tests/integration/defs/cpp/test_unit_tests.py index 78f894c465a5..9b3f8ead0bf0 100644 --- a/tests/integration/defs/cpp/test_unit_tests.py +++ b/tests/integration/defs/cpp/test_unit_tests.py @@ -44,23 +44,14 @@ def test_kv_cache_compression_unit_tests(build_kv_cache_compression_tests, xml_name = "results-unit-tests-kv_cache_compression.xml" - ctest_command = [ - "ctest", - "--output-on-failure", - "--test-dir", - f"{build_dir}/tests/unit_tests/kv_cache_compression", - "--output-junit", - f"{build_dir}/{xml_name}", - ] - - parallel = _cpp.default_test_parallel - if parallel_override := _os.environ.get("LLM_TEST_PARALLEL_OVERRIDE", None): - parallel = int(parallel_override) - - cpp_env = {**_os.environ} - - _cpp.parallel_run_ctest(ctest_command, - cwd=build_dir, - env=cpp_env, - timeout=2700, - parallel=parallel) + # Run the binary directly: the lightweight fixture builds only this gtest, + # so a ctest directory scan would trip over unbuilt neighbors. + _cpp.run_command( + [ + f"{build_dir}/tests/unit_tests/kernels/nvfp4ColdPageKernelsTest", + f"--gtest_output=xml:{build_dir}/{xml_name}", + ], + cwd=build_dir, + env={**_os.environ}, + timeout=2700, + ) From ba7b21eba8373e60365c18d597d219a2e28c34d9 Mon Sep 17 00:00:00 2001 From: tianruih Date: Mon, 31 Aug 2026 10:49:34 -0700 Subject: [PATCH 29/29] DCO Remediation Commit for tianruih I, tianruih , hereby add my Signed-off-by to this commit: e5d6374ea3e8750fcf3336bf6e8f65975085a067 Signed-off-by: tianruih