diff --git a/cpp/tensorrt_llm/kernels/compressorKernels/compressorKernels.cu b/cpp/tensorrt_llm/kernels/compressorKernels/compressorKernels.cu index cf48dd04b720..24e88c604c17 100644 --- a/cpp/tensorrt_llm/kernels/compressorKernels/compressorKernels.cu +++ b/cpp/tensorrt_llm/kernels/compressorKernels/compressorKernels.cu @@ -64,6 +64,7 @@ #include "tensorrt_llm/kernels/compressorKernels/compressorKernels.h" #include "tensorrt_llm/common/assert.h" +#include #include #include #include @@ -610,49 +611,52 @@ __global__ void pagedKvCompressKernel(void const* __restrict__ kv_score_raw, flo } } -// Explicit instantiations for decode kernel (NUM_RED_WARPS defaults to 1) +// ============================================================================ +// Decode kernel configuration matrix — single source of truth. +// +// X-macro listing every supported (HD, KV_EB, STATE_EB, CR, NN, NRW) tuple. +// Both the explicit template instantiations below AND the runtime dispatcher +// in pagedKvCompressLaunch() walk this list, so adding/removing a config is +// a one-line edit. +// +// HD — HEAD_DIM in {128, 512} +// KV_EB — kv_score element bytes in {2 (bf16), 4 (fp32)} +// STATE_EB — paged state element bytes in {2 (bf16), 4 (fp32)} +// CR — COMPRESS_RATIO in {4, 128} +// NN — NEXT_N (new tokens / decode step) in {1..4} +// NRW — NUM_RED_WARPS — 4 only when CR=128 (multi-warp Phase 3 reduce +// hides DRAM latency for the heavier R=128 chunk); 1 otherwise. +// +// Multi-warp SMEM budget (per block): 3 * NRW * ELEM_PER_BLOCK * sizeof(float). +// HD=128: ELEM_PER_BLOCK=128 → 6 KB +// HD=512 bf16: ELEM_PER_BLOCK=256 → 12 KB +// HD=512 fp32: ELEM_PER_BLOCK=128 → 6 KB +// ============================================================================ + +// Per-axis fan-outs (used to keep the master list compact). +#define FOREACH_DECODE_NN(F, HD, KV, ST, CR, NRW) \ + F(HD, KV, ST, CR, 1, NRW) F(HD, KV, ST, CR, 2, NRW) F(HD, KV, ST, CR, 3, NRW) F(HD, KV, ST, CR, 4, NRW) +#define FOREACH_DECODE_DTYPE(F, HD, CR, NRW) \ + FOREACH_DECODE_NN(F, HD, 2, 2, CR, NRW) \ + FOREACH_DECODE_NN(F, HD, 2, 4, CR, NRW) \ + FOREACH_DECODE_NN(F, HD, 4, 2, CR, NRW) FOREACH_DECODE_NN(F, HD, 4, 4, CR, NRW) + +// Master list. Order does not matter; the dispatcher walks linearly. +// clang-format off +#define FOREACH_DECODE_CONFIG(F) \ + /* CR=4: single-warp only (small reduction; multi-warp would over-subscribe). */ \ + FOREACH_DECODE_DTYPE(F, 128, 4, 1) FOREACH_DECODE_DTYPE(F, 512, 4, 1) \ + /* CR=128: single-warp fallback (covers next_n>4 path which currently isn't reached). */ \ + FOREACH_DECODE_DTYPE(F, 128, 128, 1) FOREACH_DECODE_DTYPE(F, 512, 128, 1) \ + /* CR=128: multi-warp fast path. Used whenever next_n <= 4 (i.e. MTP-3 and below). */ \ + FOREACH_DECODE_DTYPE(F, 128, 128, 4) FOREACH_DECODE_DTYPE(F, 512, 128, 4) +// clang-format on + +// Generate explicit template instantiations. #define INST_DECODE(HD, KV_EB, STATE_EB, CR, NN, NRW) \ template __global__ void pagedKvCompressKernel(void const*, float const*, void*, \ void*, int32_t const*, int32_t const*, void*, int32_t const*, int32_t const*, int32_t const*, int, int, int); - -#define INST_DECODE_NN(HD, KV_EB, STATE_EB, CR) \ - INST_DECODE(HD, KV_EB, STATE_EB, CR, 1, 1) \ - INST_DECODE(HD, KV_EB, STATE_EB, CR, 2, 1) \ - INST_DECODE(HD, KV_EB, STATE_EB, CR, 3, 1) INST_DECODE(HD, KV_EB, STATE_EB, CR, 4, 1) - -#define INST_DECODE_DTYPES(HD, CR) \ - INST_DECODE_NN(HD, 2, 2, CR) \ - INST_DECODE_NN(HD, 2, 4, CR) INST_DECODE_NN(HD, 4, 2, CR) INST_DECODE_NN(HD, 4, 4, CR) - -INST_DECODE_DTYPES(128, 4) -INST_DECODE_DTYPES(128, 128) -INST_DECODE_DTYPES(512, 4) -INST_DECODE_DTYPES(512, 128) - -// 4-warp parallel reduction variants for large compress_ratio. -// HD=128: ELEM_PER_BLOCK=128, smem = 3*4*128*4 = 6 KB. -// HD=512 bf16: ELEM_PER_BLOCK=256, smem = 3*4*256*4 = 12 KB. -// HD=512 fp32: ELEM_PER_BLOCK=128, smem = 3*4*128*4 = 6 KB. -// NEXT_N=1 (single-token decode) and NEXT_N=2 (MTP speculative decode). -INST_DECODE(128, 2, 2, 128, 1, 4) -INST_DECODE(128, 2, 4, 128, 1, 4) -INST_DECODE(128, 4, 2, 128, 1, 4) -INST_DECODE(128, 4, 4, 128, 1, 4) -INST_DECODE(128, 2, 2, 128, 2, 4) -INST_DECODE(128, 2, 4, 128, 2, 4) -INST_DECODE(128, 4, 2, 128, 2, 4) -INST_DECODE(128, 4, 4, 128, 2, 4) -INST_DECODE(512, 2, 2, 128, 1, 4) -INST_DECODE(512, 2, 4, 128, 1, 4) -INST_DECODE(512, 4, 2, 128, 1, 4) -INST_DECODE(512, 4, 4, 128, 1, 4) -INST_DECODE(512, 2, 2, 128, 2, 4) -INST_DECODE(512, 2, 4, 128, 2, 4) -INST_DECODE(512, 4, 2, 128, 2, 4) -INST_DECODE(512, 4, 4, 128, 2, 4) - -#undef INST_DECODE_DTYPES -#undef INST_DECODE_NN +FOREACH_DECODE_CONFIG(INST_DECODE) #undef INST_DECODE // ============================================================================ @@ -692,7 +696,10 @@ void pagedKvCompressLaunch(void const* kv_score, float const* ape, void* paged_k // For large compress_ratio, use 4-warp parallel reduction to cut the serial // softmax loop from COMPRESS_RATIO iterations to COMPRESS_RATIO/4 per warp. - // Supported configs: CR=128, (HD=128 or HD=512), NEXT_N=1 or NEXT_N=2. + // Supported configs: CR=128, (HD=128 or HD=512), NEXT_N in 1..4. NEXT_N>2 + // is required for MTP-3 decode (each step accepts up to 4 tokens per request); + // without multi-warp the slow path is a single warp doing 128 serial paged + // loads, which is DRAM-latency-bound (no other warps to hide it). // // smem per block = 3 * MULTI_WARP * ELEM_PER_BLOCK * sizeof(float) // where ELEM_PER_BLOCK = nthreads_inner * vec = HEAD_DIM / HEAD_BLOCKS. @@ -700,7 +707,7 @@ void pagedKvCompressLaunch(void const* kv_score, float const* ape, void* paged_k // HD=512 with max elem size 2 (vec=8, HEAD_BLOCKS=2): ELEM_PER_BLOCK=256 → 12 KB. // HD=512 with max elem size 4 (vec=4, HEAD_BLOCKS=4): ELEM_PER_BLOCK=128 → 6 KB. constexpr int MULTI_WARP = 4; - bool const use_multi_warp = (compress_ratio == 128 && next_n <= 2); + bool const use_multi_warp = (compress_ratio == 128 && next_n <= 4); int const num_red_warps = use_multi_warp ? MULTI_WARP : 1; int const nthreads = nthreads_inner * num_red_warps; int const elem_per_block = nthreads_inner * vec; // = HEAD_DIM / HEAD_BLOCKS @@ -708,95 +715,33 @@ void pagedKvCompressLaunch(void const* kv_score, float const* ape, void* paged_k dim3 grid(batch_size, head_blocks); -#define LAUNCH_DECODE(HD, KV_EB, STATE_EB, CR, NN) \ - pagedKvCompressKernel<<>>(kv_score, ape, \ - paged_kv, paged_score, block_table_kv, block_table_score, output, kv_lens, cu_seq_lens, cu_kv_comp, page_size, \ - max_blocks, out_elem_bytes) + // Clamp the runtime next_n into the supported range; configs above 4 fall + // back to the NN=4 instantiation (matches the prior `default:` arm). + int const next_n_dispatch = std::min(next_n, 4); -#define LAUNCH_DECODE_MW(HD, KV_EB, STATE_EB, CR, NN) \ - pagedKvCompressKernel<<>>(kv_score, \ - ape, paged_kv, paged_score, block_table_kv, block_table_score, output, kv_lens, cu_seq_lens, cu_kv_comp, \ - page_size, max_blocks, out_elem_bytes) - -#define DISPATCH_NN_MW(HD, KV_EB, STATE_EB, CR) \ - switch (next_n) \ - { \ - case 2: LAUNCH_DECODE_MW(HD, KV_EB, STATE_EB, CR, 2); break; \ - default: LAUNCH_DECODE_MW(HD, KV_EB, STATE_EB, CR, 1); break; \ - } - -#define DISPATCH_NN(HD, KV_EB, STATE_EB, CR) \ - switch (next_n) \ + // Walk FOREACH_DECODE_CONFIG until we find a matching (HD, KV, ST, CR, NN, NRW) + // tuple, then launch that instantiation. Any unsupported tuple bails via TLLM_THROW. +#define TRY_LAUNCH(HD, KV_EB, STATE_EB, CR, NN, NRW) \ + if (head_dim == HD && kv_score_elem_bytes == KV_EB && state_elem_bytes == STATE_EB && compress_ratio == CR \ + && next_n_dispatch == NN && num_red_warps == NRW) \ { \ - case 1: LAUNCH_DECODE(HD, KV_EB, STATE_EB, CR, 1); break; \ - case 2: LAUNCH_DECODE(HD, KV_EB, STATE_EB, CR, 2); break; \ - case 3: LAUNCH_DECODE(HD, KV_EB, STATE_EB, CR, 3); break; \ - default: LAUNCH_DECODE(HD, KV_EB, STATE_EB, CR, 4); break; \ + pagedKvCompressKernel<<>>(kv_score, ape, \ + paged_kv, paged_score, block_table_kv, block_table_score, output, kv_lens, cu_seq_lens, cu_kv_comp, \ + page_size, max_blocks, out_elem_bytes); \ + return; \ } + FOREACH_DECODE_CONFIG(TRY_LAUNCH) +#undef TRY_LAUNCH -#define DISPATCH_DTYPE(HD, CR, DISPATCH_MACRO) \ - do \ - { \ - if (kv_score_elem_bytes == 4 && state_elem_bytes == 4) \ - { \ - DISPATCH_MACRO(HD, 4, 4, CR); \ - } \ - else if (kv_score_elem_bytes == 2 && state_elem_bytes == 4) \ - { \ - DISPATCH_MACRO(HD, 2, 4, CR); \ - } \ - else if (kv_score_elem_bytes == 4 && state_elem_bytes == 2) \ - { \ - DISPATCH_MACRO(HD, 4, 2, CR); \ - } \ - else \ - { \ - DISPATCH_MACRO(HD, 2, 2, CR); \ - } \ - } while (false) - - if (use_multi_warp) - { - // Multi-warp path: HD=128 or HD=512, NEXT_N=1 or NEXT_N=2, CR=128. - if (head_dim == 512) - { - DISPATCH_DTYPE(512, 128, DISPATCH_NN_MW); - } - else - { - DISPATCH_DTYPE(128, 128, DISPATCH_NN_MW); - } - } - else if (compress_ratio == 4) - { - if (head_dim == 512) - { - DISPATCH_DTYPE(512, 4, DISPATCH_NN); - } - else - { - DISPATCH_DTYPE(128, 4, DISPATCH_NN); - } - } - else - { - if (head_dim == 512) - { - DISPATCH_DTYPE(512, 128, DISPATCH_NN); - } - else - { - DISPATCH_DTYPE(128, 128, DISPATCH_NN); - } - } - -#undef DISPATCH_DTYPE -#undef DISPATCH_NN -#undef DISPATCH_NN_MW -#undef LAUNCH_DECODE_MW -#undef LAUNCH_DECODE + TLLM_THROW( + "pagedKvCompressLaunch: no matching instantiation for HD=%d, kv_eb=%d, state_eb=%d, CR=%d, NN=%d, NRW=%d", + head_dim, kv_score_elem_bytes, state_elem_bytes, compress_ratio, next_n_dispatch, num_red_warps); } +#undef FOREACH_DECODE_CONFIG +#undef FOREACH_DECODE_DTYPE +#undef FOREACH_DECODE_NN + // ============================================================================ // Prefill Kernel: prefillReductionKernel //