From 9093e6b81b73aab09b24a539976b6e9fee643619 Mon Sep 17 00:00:00 2001 From: Yihan Wang Date: Mon, 21 Sep 2026 21:29:00 -0700 Subject: [PATCH] [None][refactor] Split thop.attention into phased FMHA Replace the monolithic attention thop with context and generation methods on AttentionOp. Cache layer configuration in StaticAttentionConfig, derive runtime parameters from the attention layer and batch metadata, and generate native parameter declarations and bindings from the Python schema. Preserve legacy MLA KV totals, generation prompt-length slices, cross-attention beam counts, and workspace capacity for beam-expanded sequences. Keep native fallback phases eager and retain paged-context kernel checks after initialization. Gate speculative buffers by the active phase and clear inactive native inputs. Build position-offset views from the original buffer using the current query width, and update speculative workers and adapters to retain that width. Handle packed NVFP4 output dimensions and compare sparse MHA output against an unquantized FP32 reference to avoid amplifying FP8 rounding differences. Preserve grouped DSv4 fused-epilogue output buffers during workspace sizing and phase dispatch, with coverage for context and generation layouts. Update backend adapters, parameter builders, and regression tests for the phased parameter contract, including native AttentionOp workspace coverage. Route modeling_v2 catalog attention through FallbackFmha.attention and share zero-copy speculative offset views with phased FMHA. Forward KV normalization and skip-correction arguments, and update the existing catalog contracts for the native AttentionOp interface. Signed-off-by: Yihan Wang --- cpp/tensorrt_llm/CMakeLists.txt | 41 + cpp/tensorrt_llm/common/attentionOp.cpp | 3454 ------------ cpp/tensorrt_llm/common/attentionOp.h | 647 --- cpp/tensorrt_llm/common/attentionWorkspace.h | 37 +- cpp/tensorrt_llm/common/opUtils.h | 29 - cpp/tensorrt_llm/kernels/fmhaDispatcher.h | 6 +- cpp/tensorrt_llm/kernels/gptKernels.h | 7 +- .../kernels/unfusedAttentionKernels.cu | 270 - .../kernels/unfusedAttentionKernels.h | 18 - cpp/tensorrt_llm/kernels/xqaDispatcher.h | 7 +- cpp/tensorrt_llm/nanobind/CMakeLists.txt | 3 +- cpp/tensorrt_llm/nanobind/bindings.cpp | 60 + cpp/tensorrt_llm/nanobind/thop/bindings.cpp | 225 +- cpp/tensorrt_llm/thop/CMakeLists.txt | 2 + cpp/tensorrt_llm/thop/attentionOp.cpp | 4804 ++++++++++++----- cpp/tensorrt_llm/thop/attentionOp.h | 938 +++- cpp/tensorrt_llm/thop/dsv3RopeOp.cpp | 1 - .../thop/trtllmGenQKVProcessOp.cpp | 2 - .../common/attentionWorkspaceTest.cpp | 97 +- cpp/tests/unit_tests/thop/CMakeLists.txt | 10 +- .../thop/attentionOpWorkspaceTest.cpp | 162 + scripts/generate_fmha_params.py | 439 ++ .../catalog/attention/thop_attention.md | 24 +- .../catalog/attention/thop_attention.py | 23 +- .../modeling_v2/catalog/index.yaml | 4 +- .../r1_0528_nvfp4__sm_103__dep4/modeling.py | 4 +- .../gpt_oss_120b__sm_103__tp1/modeling.py | 8 +- .../attention/ATTENTION_DEVELOPER_GUIDE.md | 28 +- .../_torch/attention/backends/cpp_schema.py | 59 + .../_torch/attention/backends/flashinfer.py | 10 +- .../attention/backends/fmha/fallback.py | 691 ++- .../backends/fmha/flashinfer_trtllm_gen.py | 29 +- .../attention/backends/fmha/interface.py | 408 +- .../_torch/attention/backends/fmha/manager.py | 5 + .../_torch/attention/backends/fmha/phased.py | 111 +- .../_torch/attention/backends/fmha/utils.py | 46 + .../_torch/attention/backends/interface.py | 86 +- .../backends/sparse/deepseek_v4/backend.py | 3 +- .../_torch/attention/backends/sparse/hooks.py | 20 +- .../kernels/trtllm_gen_dense_decode.py | 6 +- .../attention/backends/sparse/params.py | 14 +- .../_torch/attention/backends/trtllm.py | 91 +- .../_torch/speculative/dflash_attention.py | 8 +- tensorrt_llm/_torch/speculative/eagle3.py | 14 +- .../_torch/speculative/eagle3_dynamic_tree.py | 6 +- .../_torch/speculative/mtp_dynamic_tree.py | 12 +- .../nvfp4_mla_kv_cache_gather.py | 8 +- .../_torch/attention/fmha_test_utils.py | 11 + .../sparse/dsa/test_req_idx_per_token.py | 1 + .../sparse/test_prims_ts_block_sparse.py | 3 + .../attention/sparse/test_sparse_attention.py | 42 +- .../attention/sparse/test_sparse_mha.py | 21 +- .../attention/sparse/test_sparse_mqa_gqa.py | 2 +- .../attention/test_attention_op_sync.py | 704 --- .../_torch/attention/test_combined_fmha.py | 35 +- .../test_context_fmha_kernel_presence.py | 29 +- .../test_flashinfer_trtllm_gen_fmha.py | 43 + .../_torch/attention/test_fmha_interface.py | 501 ++ .../_torch/attention/test_fmha_manager.py | 28 +- .../_torch/attention/test_fmha_page_index.py | 10 +- .../attention/test_fmha_params_codegen.py | 344 ++ .../_torch/attention/test_prims_ts_fmha.py | 45 +- .../test_modeling_v2_thop_attention.py | 19 +- .../test_modeling_v2_target_contract.py | 24 +- .../test_kimi_k3_dspark_semantics.py | 19 +- .../_torch/speculative/test_eagle3.py | 32 + tests/unittest/bindings/test_bindings_ut.py | 10 + 67 files changed, 7731 insertions(+), 7169 deletions(-) delete mode 100644 cpp/tensorrt_llm/common/attentionOp.cpp delete mode 100644 cpp/tensorrt_llm/common/attentionOp.h create mode 100644 cpp/tests/unit_tests/thop/attentionOpWorkspaceTest.cpp create mode 100644 scripts/generate_fmha_params.py create mode 100644 tensorrt_llm/_torch/attention/backends/cpp_schema.py delete mode 100644 tests/unittest/_torch/attention/test_attention_op_sync.py create mode 100644 tests/unittest/_torch/attention/test_fmha_params_codegen.py diff --git a/cpp/tensorrt_llm/CMakeLists.txt b/cpp/tensorrt_llm/CMakeLists.txt index e7e823200f62..eca304a8ee0c 100644 --- a/cpp/tensorrt_llm/CMakeLists.txt +++ b/cpp/tensorrt_llm/CMakeLists.txt @@ -22,6 +22,47 @@ set(API_INCLUDE_DIR ${PROJECT_SOURCE_DIR}/include) include_directories(${CMAKE_CURRENT_SOURCE_DIR}/cutlass_extensions/include ${API_INCLUDE_DIR}) +get_filename_component(TRTLLM_REPO_ROOT_DIR "${CMAKE_CURRENT_SOURCE_DIR}/../.." + ABSOLUTE) +set(TRTLLM_FMHA_PARAMS_GENERATOR + "${TRTLLM_REPO_ROOT_DIR}/scripts/generate_fmha_params.py") +# The schema classes live in the modules that own them; the generator parses all +# of them and resolves a nested member by the class its annotation names. +set(TRTLLM_FMHA_PARAMS_SCHEMAS + "${TRTLLM_REPO_ROOT_DIR}/tensorrt_llm/_torch/attention/backends/fmha/interface.py" + "${TRTLLM_REPO_ROOT_DIR}/tensorrt_llm/_torch/attention/backends/interface.py" + "${TRTLLM_REPO_ROOT_DIR}/tensorrt_llm/_torch/attention/backends/sparse/params.py" +) +set(TRTLLM_FMHA_PARAMS_SCHEMA_ARGS "") +foreach(_schema ${TRTLLM_FMHA_PARAMS_SCHEMAS}) + list(APPEND TRTLLM_FMHA_PARAMS_SCHEMA_ARGS --module-path "${_schema}") +endforeach() +set(TRTLLM_FMHA_PARAMS_GENERATED_INCLUDE_DIR + "${CMAKE_CURRENT_BINARY_DIR}/generated") +# Generate at configure time rather than as a build step: a configure-only tree, +# which is how the clangd compile database is produced, still needs these +# headers for attentionOp.h to parse. CONFIGURE_DEPENDS re-runs CMake when a +# schema or the generator changes, so editing a schema and building is enough to +# refresh them. +set_property( + DIRECTORY + APPEND + PROPERTY CMAKE_CONFIGURE_DEPENDS ${TRTLLM_FMHA_PARAMS_GENERATOR} + ${TRTLLM_FMHA_PARAMS_SCHEMAS}) +execute_process( + COMMAND + "${Python3_EXECUTABLE}" "${TRTLLM_FMHA_PARAMS_GENERATOR}" + ${TRTLLM_FMHA_PARAMS_SCHEMA_ARGS} --class-name FmhaParams --out-dir + "${TRTLLM_FMHA_PARAMS_GENERATED_INCLUDE_DIR}/tensorrt_llm/thop" + RESULT_VARIABLE TRTLLM_FMHA_PARAMS_RESULT + OUTPUT_VARIABLE TRTLLM_FMHA_PARAMS_OUTPUT + ERROR_VARIABLE TRTLLM_FMHA_PARAMS_OUTPUT) +if(NOT TRTLLM_FMHA_PARAMS_RESULT EQUAL 0) + message( + FATAL_ERROR + "Generating native FmhaParams failed:\n${TRTLLM_FMHA_PARAMS_OUTPUT}") +endif() + set(TARGET_ARCH "unknown") message(STATUS "CMAKE_SYSTEM_PROCESSOR: ${CMAKE_SYSTEM_PROCESSOR}") diff --git a/cpp/tensorrt_llm/common/attentionOp.cpp b/cpp/tensorrt_llm/common/attentionOp.cpp deleted file mode 100644 index 56f64b161c69..000000000000 --- a/cpp/tensorrt_llm/common/attentionOp.cpp +++ /dev/null @@ -1,3454 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & - * AFFILIATES. All rights reserved. SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ -#include "attentionOp.h" -#include "tensorrt_llm/common/assert.h" -#include "tensorrt_llm/common/attentionWorkspace.h" -#include "tensorrt_llm/common/config.h" -#include "tensorrt_llm/common/envUtils.h" -#include "tensorrt_llm/common/logger.h" -#include "tensorrt_llm/common/memoryUtils.h" -#include "tensorrt_llm/common/sageQuant.h" -#include "tensorrt_llm/common/tllmDataType.h" -#include "tensorrt_llm/kernels/decoderMaskedMultiheadAttention.h" -#include "tensorrt_llm/kernels/decoderMaskedMultiheadAttention/cascadeAttentionKernel.h" -#include "tensorrt_llm/kernels/flashMLA/flash_mla.h" -#include "tensorrt_llm/kernels/gptKernels.h" -#include "tensorrt_llm/kernels/kvCacheUtils.h" -#include "tensorrt_llm/kernels/multiHeadAttentionCommon.h" -#include "tensorrt_llm/kernels/sparseAttentionKernels.h" -#include "tensorrt_llm/kernels/unfusedAttentionKernels.h" -#include "tensorrt_llm/runtime/iBuffer.h" -#include "tensorrt_llm/runtime/utils/debugUtils.h" -#include "tensorrt_llm/runtime/utils/mpiUtils.h" -#include -#include -#include - -using namespace tensorrt_llm::kernels; -namespace tc = tensorrt_llm::common; -using tensorrt_llm::common::op::AttentionContextWorkspaceSizes; -using tensorrt_llm::common::op::AttentionFlashMlaWorkspaceSizes; -using tensorrt_llm::common::op::AttentionGenerationWorkspaceSizes; -using tensorrt_llm::common::op::AttentionOp; -using tensorrt_llm::common::op::AttentionWorkspaceManager; -using tensorrt_llm::common::op::AttentionXqaWorkspaceSizes; -using tensorrt_llm::common::op::KvCacheBuffers; - -template -struct SATypeConverter -{ - using Type = T; -}; - -template <> -struct SATypeConverter -{ - using Type = uint16_t; -}; - -template -struct FusedQKVMaskedAttentionDispatchParams -{ - T const* qkv_buf; - T const* qkv_bias; - T const* relative_attention_bias; - bool const* attention_mask; - float const* attention_sinks; - float const* logn_scaling_ptr; - int const* cache_indir; - void* context_buf; - bool const* finished; - int const* sequence_lengths; - int max_batch_size; - int inference_batch_size; - int beam_width; - int head_num; - int kv_head_num; - int size_per_head; - int rotary_embedding_dim; - float rotary_embedding_base; - RotaryScalingType rotary_embedding_scale_type; - float rotary_embedding_scale; - float const* rotary_embedding_inv_freq_cache; - float2 const* rotary_embedding_cos_sin_cache; - float rotary_embedding_short_m_scale; - float rotary_embedding_long_m_scale; - int rotary_embedding_max_positions; - int rotary_embedding_original_max_positions; - int rotary_cogvlm_vision_start; - int rotary_cogvlm_vision_length; - PositionEmbeddingType position_embedding_type; - bool position_shift_enabled; - - int chunked_attention_size; - int attention_mask_stride; - int max_attention_window_size; - int cyclic_attention_window_size; - int sink_token_length; - int const* input_lengths; - int timestep; - float q_scaling; - float attn_logit_softcapping_scale; - int relative_attention_bias_stride; - T const* linear_bias_slopes; - int const* ia3_tasks; - T const* ia3_key_weights; - T const* ia3_value_weights; - float const* qkv_scale_out; - bool fp8_context_fmha; - float const* attention_out_scale; - bool mUnfuseQkvGemm; - tc::QuantMode quant_option; - bool multi_block_mode; - int max_seq_len_tile; - int min_seq_len_tile; - T* partial_out; - float* partial_sum; - float* partial_max; - int* block_counter; - // Cascade attention prefix-side workspace (fp32). Sliced from the same - // generation workspace and forwarded into Multihead_attention_params so - // the cascade fast-path no longer needs its own cudaMalloc. - float* cascade_partial_out{}; - float* cascade_partial_max{}; - float* cascade_partial_sum{}; - float const* kv_scale_orig_quant; - float const* kv_scale_quant_orig; - tc::QuantMode kv_cache_quant_mode; - int multi_processor_count; - KVCacheBuffer kv_block_array; - KVLinearBuffer shift_k_cache_buffer; - bool cross_attention = false; - int const* memory_length_per_sample = nullptr; - int max_distance = 0; - bool block_sparse_attention = false; - BlockSparseParams block_sparse_params; - int32_t const* mrope_position_deltas; -}; - -template -struct ConvertMMHAToXQAParamsHelper -{ - static constexpr Data_type data_type = DATA_TYPE_FP16; - static constexpr bool supported = false; -}; - -template <> -struct ConvertMMHAToXQAParamsHelper<__half, KVLinearBuffer> -{ - static constexpr Data_type data_type = DATA_TYPE_FP16; - static constexpr bool supported = true; -}; - -template <> -struct ConvertMMHAToXQAParamsHelper<__half, KVBlockArray> -{ - static constexpr Data_type data_type = DATA_TYPE_FP16; - static constexpr bool supported = true; -}; - -#ifdef ENABLE_BF16 -template <> -struct ConvertMMHAToXQAParamsHelper<__nv_bfloat16, KVLinearBuffer> -{ - static constexpr Data_type data_type = DATA_TYPE_BF16; - static constexpr bool supported = true; -}; - -template <> -struct ConvertMMHAToXQAParamsHelper<__nv_bfloat16, KVBlockArray> -{ - static constexpr Data_type data_type = DATA_TYPE_BF16; - static constexpr bool supported = true; -}; -#endif - -template -bool AttentionOp::convertMMHAParamsToXQAParams(tensorrt_llm::kernels::XQAParams& xqaParams, - EnqueueGenerationParams const& generationsParams, bool forConfigurePlugin) -{ - bool retval = ConvertMMHAToXQAParamsHelper::supported; - if (!retval) - { - return false; - } - xqaParams = {}; - xqaParams.data_type = ConvertMMHAToXQAParamsHelper::data_type; - - xqaParams.num_q_heads = mNumAttnHeads; - xqaParams.num_kv_heads = mNumAttnKVHeads; - xqaParams.head_size = mHeadSize; - xqaParams.unidirectional = mUnidirectional; - xqaParams.q_scaling = mQScaling; - xqaParams.rotary_embedding_dim = mRotaryEmbeddingDim; - xqaParams.rotary_embedding_base = mRotaryEmbeddingBase; - xqaParams.rotary_embedding_scale_type = mRotaryEmbeddingScaleType; - xqaParams.rotary_embedding_scale = mRotaryEmbeddingScale; - xqaParams.rotary_embedding_max_positions = mRotaryEmbeddingMaxPositions; - xqaParams.rotary_vision_start = mVisionStart; - xqaParams.rotary_vision_length = mVisionLength; - xqaParams.rotary_cos_sin = generationsParams.rotary_cos_sin; - xqaParams.position_embedding_type = mPositionEmbeddingType; - xqaParams.position_shift_enabled = mPosShiftEnabled; - xqaParams.remove_padding = mRemovePadding; - xqaParams.mask_type = mMaskType; - xqaParams.paged_kv_cache = mPagedKVCache; - xqaParams.tokens_per_block = mTokensPerBlock; - xqaParams.kv_cache_quant_mode = mKVCacheQuantMode; - xqaParams.tp_size = mAttnTpSize; - xqaParams.tp_rank = mAttnTpRank; - xqaParams.qkv_bias_enabled = mQKVBiasEnabled; - xqaParams.cross_attention = mCrossAttention; - xqaParams.max_distance = mMaxDistance; - xqaParams.multi_block_mode = common::getEnvForceDeterministicAttention() ? false : mMultiBlockMode; - // Medusa mode will have multiple query tokens. - xqaParams.multi_query_tokens = mIsSpecDecodingEnabled && mUseSpecDecoding; - xqaParams.is_spec_dec_tree = mIsSpecDecTree; - xqaParams.force_prepare_spec_dec_tree_mask = mForcePrepareSpecDecTreeMask; - xqaParams.layer_idx = generationsParams.layer_idx; - - if (mKVCacheQuantMode.hasInt8KvCache()) - { - xqaParams.kv_cache_data_type = DATA_TYPE_INT8; - } - else if (mKVCacheQuantMode.hasFp8KvCache()) - { - // Inputs to MLA is FP8 instead of BF16/FP16 when using FP8 KV cache. - if (xqaParams.isMLA()) - { - xqaParams.data_type = DATA_TYPE_E4M3; - } - xqaParams.kv_cache_data_type = DATA_TYPE_E4M3; - } - else if (mKVCacheQuantMode.hasFp4KvCache()) - { - xqaParams.kv_cache_data_type = DATA_TYPE_E2M1; - } - else - { - xqaParams.kv_cache_data_type = xqaParams.data_type; - } - if (xqaParams.kv_cache_data_type == DATA_TYPE_INT8 - || (xqaParams.kv_cache_data_type == DATA_TYPE_E4M3 && (mSM < kSM_90 || mSM > kSM_120))) - { - xqaParams.multi_block_mode = false; - } - - xqaParams.output = generationsParams.context_buf; - xqaParams.qkv = generationsParams.attention_input; - xqaParams.cache_indir = generationsParams.cache_indir; - xqaParams.attention_sinks = generationsParams.attention_sinks; - xqaParams.kv_scale_orig_quant = generationsParams.kv_scale_orig_quant; - xqaParams.kv_scale_quant_orig = generationsParams.kv_scale_quant_orig; - xqaParams.host_past_key_value_lengths = generationsParams.host_past_key_value_lengths; - xqaParams.host_context_lengths = generationsParams.host_context_lengths; - xqaParams.semaphores = generationsParams.semaphores; - xqaParams.workspaces = generationsParams.workspace; - if (mCpSize > 1) - { - size_t const batch_beam = generationsParams.beam_width * generationsParams.num_requests; - size_t const cpMaxPaddedSequenceLength = (batch_beam + mCpSize - 1) / mCpSize * mCpSize; - size_t const cpWorkspaceSize - = 2 * sizeof(T) * cpMaxPaddedSequenceLength * (mNumHeads + 2 * mNumKVHeads) * mHeadSize; - xqaParams.workspaces - = reinterpret_cast(reinterpret_cast(xqaParams.workspaces) + cpWorkspaceSize); - } - xqaParams.batch_size = generationsParams.num_requests; - xqaParams.beam_width = generationsParams.beam_width; - // Speculative decoding mode has generation input_length > 1. - xqaParams.generation_input_length = generationsParams.input_seq_length; - xqaParams.chunked_attention_size - = mAttentionChunkSize && !tc::getEnvDisableChunkedAttentionInGenPhase() ? *mAttentionChunkSize : INT_MAX; - xqaParams.max_attention_window_size = generationsParams.max_attention_window_size; - xqaParams.cyclic_attention_window_size = generationsParams.cyclic_attention_window_size; - xqaParams.max_blocks_per_sequence = generationsParams.max_blocks_per_sequence; - xqaParams.sink_token_length = generationsParams.sink_token_length; - xqaParams.max_past_kv_length = generationsParams.max_past_kv_length; - xqaParams.qkv_bias = generationsParams.qkv_bias; - xqaParams.sequence_lengths = generationsParams.sequence_lengths; - xqaParams.context_lengths = generationsParams.context_lengths; - xqaParams.alibi_slopes = generationsParams.alibi_slopes; - // Pre-computed rotary inv freq when building the engines. - xqaParams.rotary_embedding_inv_freq_cache = generationsParams.rotary_inv_freq; - if (!forConfigurePlugin) - { - // Speculative decoding (need to take new generated ids into consideration). - TLLM_CHECK_WITH_INFO( - !(mIsSpecDecodingEnabled && mUseSpecDecoding) || generationsParams.spec_decoding_packed_mask != nullptr, - "Speculative decoding mode needs a valid packed_mask input tensor."); - } - xqaParams.spec_decoding_packed_mask = generationsParams.spec_decoding_packed_mask; - xqaParams.spec_decoding_position_offsets = generationsParams.spec_decoding_position_offsets; - xqaParams.spec_decoding_generation_lengths = generationsParams.spec_decoding_generation_lengths; - xqaParams.spec_decoding_is_generation_length_variable - = generationsParams.spec_decoding_is_generation_length_variable; - xqaParams.spec_decoding_max_generation_length = generationsParams.spec_decoding_max_generation_length; - xqaParams.spec_decoding_bl_tree_mask_offset = generationsParams.spec_decoding_bl_tree_mask_offset; - xqaParams.spec_decoding_bl_tree_mask = generationsParams.spec_decoding_bl_tree_mask; - xqaParams.spec_bl_tree_first_sparse_mask_offset_kv = generationsParams.spec_bl_tree_first_sparse_mask_offset_kv; - xqaParams.mrope_position_deltas = generationsParams.mrope_position_deltas; - xqaParams.helix_position_offsets = generationsParams.helix_position_offsets; - xqaParams.helix_is_inactive_rank = generationsParams.helix_is_inactive_rank; - xqaParams.softmax_stats = generationsParams.softmax_stats; - xqaParams.trtllm_gen_jit_warmup = generationsParams.trtllm_gen_jit_warmup; - xqaParams.trtllm_gen_jit_warmup_max_num_requests = mMaxNumRequests; - xqaParams.trtllm_gen_jit_warmup_max_seq_len_q = mMaxContextLength; - xqaParams.trtllm_gen_jit_warmup_max_seq_len_kv = mMaxSeqLen; - - xqaParams.logn_scaling_ptr = generationsParams.logn_scaling_ptr; - xqaParams.total_num_input_tokens = mCpSize > 1 ? generationsParams.num_requests : generationsParams.num_tokens; - xqaParams.is_fp8_output = mFP8AttenOutput; - xqaParams.fp8_out_scale = ((mFP8AttenOutput) ? generationsParams.attention_output_orig_quant : nullptr); - // Parameters required for FP4 output. - xqaParams.output_sf = generationsParams.context_buf_sf; - xqaParams.fp4_out_sf_scale = generationsParams.attention_output_sf_scale; - xqaParams.start_token_idx_sf = generationsParams.start_token_idx_sf; - // Parameters for sparse attention - xqaParams.sparse_params = mRuntimeSparseAttentionParams; - xqaParams.use_sparse_attention_gen_paged = useTllmGenSparseAttentionPaged(); - // Skip softmax threshold. - xqaParams.skip_softmax_threshold_scale_factor = mSkipSoftmaxThresholdScaleFactorDecode; -#ifdef SKIP_SOFTMAX_STAT - // Statistics of skip-softmax, pointers of device memory for output - xqaParams.skip_softmax_total_blocks = mSkipSoftmaxTotalBlocks; - xqaParams.skip_softmax_skipped_blocks = mSkipSoftmaxSkippedBlocks; -#endif - // Cross attention parameters. - xqaParams.encoder_input_lengths = generationsParams.encoder_input_lengths; - - return true; -} - -template -int AttentionOp::ulyssesContextPreprocess(T const* input, T* output, T* buffer, EnqueueContextParams const& params, - int const* cu_q_seqlens, int const* cu_cp_partial_seqlens, cudaStream_t stream) -{ - int32_t partialTokenNum = 0; - int32_t maxPartialLength = 0; - for (int32_t batchIdx = 0; batchIdx < params.batch_size; ++batchIdx) - { - int32_t partialLength = (params.host_context_lengths[batchIdx] + mCpSize - 1) / mCpSize; - maxPartialLength = std::max(maxPartialLength, partialLength); - partialTokenNum += partialLength; - } - auto const partialHeads = mNumAttnHeads + 2 * mNumAttnKVHeads; - - // full request: [bs, seqlen, head, headSize] - // - // input of cp: [bs, partialLength, head, headSize] - // view_1 as [bs, partialLength, cpSize_Head, partialHead, headSize] - // transpose_1 as [cpSize_Head, bs, partialLenth, partialHead, headSize] - // all-to-all to get [cpSize_Length, bs, partialLength, partialHead, headSize] - // transpose_2 to [bs, cpSize_Length, partialLength, partialHead, headSize] - // view_2 as [bs, totalLength, partialHead, headSize] - // and this is same to the input under TP. - // - // when we use remove_input_padding, bs and length are fused into numTokens. So, we need to - // insert the cpSize_Length dimension of transpose_2 into numTokens directly like - // input of cp: [partialNumTokens, head, headSize] - // view_1 as [partialNumTokens, cpSize_Head, partialHead, headSize] - // transpose_1 as [cpSize_Head, partialNumTokens, partialHead, headSize] - // all-to-all to get [cpSize_Length, partialNumTokens, partialHead, headSize] - // transpose_2 as [NumTokens, partialHead, headSize] - // and this is same to the input under TP. - - // view_1 + transpose_1 - invokeCpTranspose(output, buffer, input, partialTokenNum, mCpSize, mNumAttnHeads, mNumAttnKVHeads, - mUlyssesMQABroadcast, getHeadSize(), mCpRank, stream); - sync_check_cuda_error(stream); - - // Do all to all -#if ENABLE_MULTI_DEVICE - ncclGroupStart(); - for (int cpIdx = 0; cpIdx < mCpSize; cpIdx++) - { - if (cpIdx != mCpRank) - { - NCCLCHECK(ncclSend(output + cpIdx * (partialTokenNum * getHeadSize() * partialHeads), - (partialTokenNum * getHeadSize() * partialHeads), (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, - stream)); - NCCLCHECK(ncclRecv(buffer + cpIdx * (partialTokenNum * getHeadSize() * partialHeads), - (partialTokenNum * getHeadSize() * partialHeads), (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, - stream)); - } - } - ncclGroupEnd(); - sync_check_cuda_error(stream); -#endif // ENABLE_MULTI_DEVICE - - // transpose_2 + view_2 - invokeCpTranspose2(output, buffer, params.context_lengths, cu_q_seqlens, cu_cp_partial_seqlens, mCpSize, - maxPartialLength, params.batch_size, partialHeads, getHeadSize(), stream); - - return 0; -} - -template -int AttentionOp::ulyssesContextPostprocess(T* input, T* output, T* buffer, EnqueueContextParams const& params, - int const* cu_q_seqlens, int const* cu_cp_partial_seqlens, cudaStream_t stream) -{ - // After FMHA, we get result [numTokens(bs, cp, paritalLength), partialHead, headSize] - // transpose_2_reverse: [cpSize_Length, partialTokens(bs, partialLength), partialHead, headSize] - // all-to-all: [cpSize_Head, partialTokens, partialHead, headSize] - // transpose_1_reverse: [partialTokens, cpSize_Head, partialHead, headSize] - // view: [partialTokens, head, headSize] - - int32_t maxPartialLength = 0; - int32_t partialTokenNum = 0; - for (int32_t batchIdx = 0; batchIdx < params.batch_size; ++batchIdx) - { - int32_t partialLength = (params.host_context_lengths[batchIdx] + mCpSize - 1) / mCpSize; - maxPartialLength = std::max(maxPartialLength, partialLength); - partialTokenNum += partialLength; - } - - // transpose_2_reverse - if (mFP8AttenOutput) - { - invokeCpTransposeToSeqMajor2(reinterpret_cast<__nv_fp8_e4m3*>(buffer), - reinterpret_cast<__nv_fp8_e4m3 const*>(input), params.context_lengths, cu_q_seqlens, cu_cp_partial_seqlens, - mCpSize, maxPartialLength, params.batch_size, mNumAttnHeads, getHeadSize(), stream); - } - else - { - invokeCpTransposeToSeqMajor2(buffer, input, params.context_lengths, cu_q_seqlens, cu_cp_partial_seqlens, - mCpSize, maxPartialLength, params.batch_size, mNumAttnHeads, getHeadSize(), stream); - } - - // all-to-all -#if ENABLE_MULTI_DEVICE - size_t const elementNum = partialTokenNum * getHeadSize() * mNumAttnHeads; - ncclGroupStart(); - for (int cpIdx = 0; cpIdx < mCpSize; cpIdx++) - { - if (cpIdx != mCpRank) - { - if (mFP8AttenOutput) - { - NCCLCHECK(ncclSend(reinterpret_cast<__nv_fp8_e4m3*>(buffer) + cpIdx * elementNum, elementNum, ncclInt8, - cpIdx, *mCpNcclComm, stream)); - NCCLCHECK(ncclRecv(reinterpret_cast<__nv_fp8_e4m3*>(input) + cpIdx * elementNum, elementNum, ncclInt8, - cpIdx, *mCpNcclComm, stream)); - } - else - { - NCCLCHECK(ncclSend( - buffer + cpIdx * elementNum, elementNum, (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, stream)); - NCCLCHECK(ncclRecv( - input + cpIdx * elementNum, elementNum, (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, stream)); - } - } - } - ncclGroupEnd(); -#endif // ENABLE_MULTI_DEVICE - - // transpose_1_reverse + view - if (mFP8AttenOutput) - { - invokeCpTransposeToSeqMajor<__nv_fp8_e4m3>(reinterpret_cast<__nv_fp8_e4m3*>(output), - reinterpret_cast<__nv_fp8_e4m3 const*>(buffer), reinterpret_cast<__nv_fp8_e4m3 const*>(input), - partialTokenNum, mCpSize, mNumAttnHeads, getHeadSize(), mCpRank, stream); - } - else - { - invokeCpTransposeToSeqMajor( - (T*) output, buffer, input, partialTokenNum, mCpSize, mNumAttnHeads, getHeadSize(), mCpRank, stream); - } - return 0; -} - -template -int AttentionOp::ulyssesGenerationPreprocess( - T const* input, T* output, T* buffer, int32_t batch_beam, cudaStream_t stream) -{ - if (mCpSize <= 1) - return 0; - - auto const partialTokenNum = (batch_beam + mCpSize - 1) / mCpSize; - - // attention_input shape: [partialTokenNum, numHeads, headSize] - // view_1: [partialTokenNum, cpSize_Head, partialHeads, headSize] - // transpose_1: [cpSize_Head, partialTokenNum, partialHeads, headSize] - // all-to-all to get [cpSize_Length, partialTokenNum, partialHead, headSize] - // view_2 as [tokens, partialHead, headSize] - - // do transpose_1 - // [1, mNumHeads + 2*mNumKVHeads, headSize] - // -> (view) [1, cpSize * mNumAttnHeads + cpSize * mNumAttnKVHeads + cpSize * partilKVHeads, - // headSize] - // -> (transpose) [cpSize, 1, mNumAttnHeads + mNumAttnKVHeads + mNumAttnKVHeads, headSize] - invokeCpTranspose(buffer, output, input, partialTokenNum, mCpSize, mNumAttnHeads, mNumAttnKVHeads, - mUlyssesMQABroadcast, mHeadSize, mCpRank, stream); - sync_check_cuda_error(stream); - - // Do all to all -#if ENABLE_MULTI_DEVICE - auto const partialHeads = mNumAttnHeads + 2 * mNumAttnKVHeads; - - ncclGroupStart(); - for (int cpIdx = 0; cpIdx < mCpSize; cpIdx++) - { - if (cpIdx != mCpRank) - { - NCCLCHECK(ncclSend(buffer + cpIdx * (partialTokenNum * getHeadSize() * partialHeads), - (partialTokenNum * getHeadSize() * partialHeads), (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, - stream)); - NCCLCHECK(ncclRecv(output + cpIdx * (partialTokenNum * getHeadSize() * partialHeads), - (partialTokenNum * getHeadSize() * partialHeads), (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, - stream)); - } - } - ncclGroupEnd(); - sync_check_cuda_error(stream); -#endif // ENABLE_MULTI_DEVICE - return 0; -} - -template -int AttentionOp::ulyssesGenerationPostprocess(T* input, T* output, T* buffer, int32_t batch_beam, cudaStream_t stream) -{ - if (mCpSize <= 1) - return 0; - - // mmha output shape: [tokens, partialHead, headSize] - // view: [cpSize_Length, partialTokens, partialHead, headSize] - // all-to-all: [cpSize_Head, partialTokens, partialHead, headSize] - // transpose_1_reverse: [partialTokens, cpSize_Head, partialHead, headSize] - // view: [partialTokens, head, headSize] - - auto const partialTokenNum = (batch_beam + mCpSize - 1) / mCpSize; - - // do all-to-all -#if ENABLE_MULTI_DEVICE - size_t const elementNum = partialTokenNum * getHeadSize() * mNumAttnHeads; - ncclGroupStart(); - for (int cpIdx = 0; cpIdx < mCpSize; cpIdx++) - { - if (cpIdx != mCpRank) - { - if (mFP8AttenOutput) - { - NCCLCHECK(ncclSend(reinterpret_cast<__nv_fp8_e4m3*>(input) + cpIdx * elementNum, elementNum, ncclInt8, - cpIdx, *mCpNcclComm, stream)); - NCCLCHECK(ncclRecv(reinterpret_cast<__nv_fp8_e4m3*>(buffer) + cpIdx * elementNum, elementNum, ncclInt8, - cpIdx, *mCpNcclComm, stream)); - } - else - { - NCCLCHECK(ncclSend( - input + cpIdx * elementNum, elementNum, (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, stream)); - NCCLCHECK(ncclRecv( - buffer + cpIdx * elementNum, elementNum, (*getDtypeMap())[mType], cpIdx, *mCpNcclComm, stream)); - } - } - } - ncclGroupEnd(); -#endif // ENABLE_MULTI_DEVICE - - // do transpose_1_reverse - if (mFP8AttenOutput) - { - invokeCpTransposeToSeqMajor<__nv_fp8_e4m3>(reinterpret_cast<__nv_fp8_e4m3*>(output), - reinterpret_cast<__nv_fp8_e4m3 const*>(input), reinterpret_cast<__nv_fp8_e4m3 const*>(buffer), - partialTokenNum, mCpSize, mNumAttnHeads, getHeadSize(), mCpRank, stream); - } - else - { - invokeCpTransposeToSeqMajor( - (T*) output, input, buffer, partialTokenNum, mCpSize, mNumAttnHeads, getHeadSize(), mCpRank, stream); - } - return 0; -} - -template -void fusedQKV_masked_attention_dispatch(Multihead_attention_params& params, - FusedQKVMaskedAttentionDispatchParams const& input_params, cudaStream_t stream) -{ - using DataType = typename SATypeConverter::Type; - - // Prepare the parameters. - params = {}; - - int hidden_units = input_params.head_num * input_params.size_per_head; - int hidden_units_kv = input_params.kv_head_num * input_params.size_per_head; - if (input_params.qkv_bias != nullptr) - { - params.q_bias = reinterpret_cast(input_params.qkv_bias); - params.k_bias = reinterpret_cast(input_params.qkv_bias) + hidden_units; - params.v_bias = reinterpret_cast(input_params.qkv_bias) + hidden_units + hidden_units_kv; - } - else - { - params.q_bias = nullptr; - params.k_bias = nullptr; - params.v_bias = nullptr; - } - - // Set the output buffer. - params.out = input_params.context_buf; - - // Set the input buffers. - params.q = reinterpret_cast(input_params.qkv_buf); - params.k = reinterpret_cast(input_params.qkv_buf) + hidden_units; - params.v = reinterpret_cast(input_params.qkv_buf) + hidden_units + hidden_units_kv; - - params.int8_kv_cache = input_params.kv_cache_quant_mode.hasInt8KvCache(); - params.fp8_kv_cache = input_params.kv_cache_quant_mode.hasFp8KvCache(); - if (input_params.kv_cache_quant_mode.hasKvCacheQuant()) - { - params.kv_scale_orig_quant = input_params.kv_scale_orig_quant; - params.kv_scale_quant_orig = input_params.kv_scale_quant_orig; - } - - params.stride = hidden_units + 2 * hidden_units_kv; - params.finished = const_cast(input_params.finished); - - params.cache_indir = input_params.cache_indir; - params.batch_size = input_params.inference_batch_size; - params.beam_width = input_params.beam_width; - params.chunked_attention_size = input_params.chunked_attention_size; - if (input_params.chunked_attention_size != INT_MAX && !tc::getEnvDisableChunkedAttentionInGenPhase()) - { - TLLM_CHECK_WITH_INFO((input_params.chunked_attention_size & (input_params.chunked_attention_size - 1)) == 0, - "Attention chunk size should be a power of 2."); - params.chunked_attention_size_log2 = std::log2(input_params.chunked_attention_size); - } - else - { - params.chunked_attention_size_log2 = 0; - } - params.max_attention_window_size = input_params.max_attention_window_size; - params.cyclic_attention_window_size = input_params.cyclic_attention_window_size; - params.sink_token_length = input_params.sink_token_length; - params.length_per_sample = input_params.sequence_lengths; // max_input_length + current output length - // timestep for shared memory size calculation and rotary embedding computation - params.timestep = input_params.timestep; - params.num_heads = input_params.head_num; - params.num_kv_heads = input_params.kv_head_num; - params.hidden_size_per_head = input_params.size_per_head; - params.rotary_embedding_dim = input_params.rotary_embedding_dim; - params.rotary_embedding_base = input_params.rotary_embedding_base; - params.rotary_embedding_scale_type = input_params.rotary_embedding_scale_type; - params.rotary_embedding_scale = input_params.rotary_embedding_scale; - params.rotary_embedding_inv_freq_cache = input_params.rotary_embedding_inv_freq_cache; - params.rotary_embedding_cos_sin_cache = input_params.rotary_embedding_cos_sin_cache; - params.rotary_embedding_short_m_scale = input_params.rotary_embedding_short_m_scale; - params.rotary_embedding_long_m_scale = input_params.rotary_embedding_long_m_scale; - params.rotary_embedding_max_positions = input_params.rotary_embedding_max_positions; - params.rotary_embedding_original_max_positions = input_params.rotary_embedding_original_max_positions; - params.rotary_cogvlm_vision_start = input_params.rotary_cogvlm_vision_start; - params.rotary_cogvlm_vision_length = input_params.rotary_cogvlm_vision_length; - params.position_embedding_type = input_params.position_embedding_type; - params.position_shift_enabled = input_params.position_shift_enabled; - // Note: keep norm factor (sqrt(K_dim)) when adopting megatron T5 structure (may adjust) - params.inv_sqrt_dh = 1.F / (sqrtf((float) params.hidden_size_per_head) * input_params.q_scaling); - params.attn_logit_softcapping_scale = input_params.attn_logit_softcapping_scale; - params.attn_logit_softcapping_inverse_scale = 1.0f / input_params.attn_logit_softcapping_scale; - - params.logn_scaling_ptr = input_params.logn_scaling_ptr; - params.relative_attention_bias = reinterpret_cast(input_params.relative_attention_bias); - params.relative_attention_bias_stride = input_params.relative_attention_bias_stride; - params.max_distance = input_params.max_distance; - params.block_sparse_attention = input_params.block_sparse_attention; - params.block_sparse_params = input_params.block_sparse_params; - - // Attention mask input. - params.attention_mask = input_params.attention_mask; - params.attention_mask_stride = input_params.attention_mask_stride; - - // Attention sinks. - params.attention_sinks = input_params.attention_sinks; - - // The slope of linear position bias per head, e.g., ALiBi. - if (input_params.linear_bias_slopes != nullptr) - { - params.linear_bias_slopes = reinterpret_cast(input_params.linear_bias_slopes); - } - params.input_lengths = input_params.input_lengths; - - params.ia3_tasks = input_params.ia3_tasks; - params.ia3_key_weights = reinterpret_cast(input_params.ia3_key_weights); - params.ia3_value_weights = reinterpret_cast(input_params.ia3_value_weights); - - if (input_params.quant_option.hasStaticActivationScaling() || input_params.fp8_context_fmha) - { - // qkv_scale_out is nullptr currently (no scale). - params.qkv_scale_quant_orig = input_params.qkv_scale_out; - TLLM_CHECK_WITH_INFO(!input_params.fp8_context_fmha || input_params.attention_out_scale != nullptr, - "attention output scale should be provided."); - params.attention_out_scale_orig_quant = input_params.attention_out_scale; - } - - params.multi_block_mode = input_params.multi_block_mode; - // Cascade-attention partials must be wired regardless of multi_block_mode. - // Cascade decode runs with multi_block disabled (short-decode workloads have - // max_num_seq_len_tiles == 1, so enable_multi_block is structurally false). - // Gating these behind multi_block_mode leaves cascade_partial_* null and makes - // launch_cascade_attention fall back with "cascade workspace not provisioned". - params.cascade_partial_out = input_params.cascade_partial_out; - params.cascade_partial_max = input_params.cascade_partial_max; - params.cascade_partial_sum = input_params.cascade_partial_sum; - if (input_params.multi_block_mode) - { - params.min_seq_len_tile = input_params.min_seq_len_tile; - params.max_seq_len_tile = input_params.max_seq_len_tile; - - params.partial_out = reinterpret_cast(input_params.partial_out); - params.partial_sum = input_params.partial_sum; - params.partial_max = input_params.partial_max; - - params.block_counter = input_params.block_counter; - } - - params.multi_processor_count = input_params.multi_processor_count; - - // cross attn - params.memory_length_per_sample = input_params.memory_length_per_sample; - - params.mrope_position_deltas = input_params.mrope_position_deltas; - sync_check_cuda_error(stream); - - masked_multihead_attention(params, input_params.kv_block_array, input_params.shift_k_cache_buffer, stream); -} - -#define INSTANTIATE_MMHA_DISPATCH(T_MMHA, T) \ - template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ - FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ - template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ - FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ - template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ - FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ - template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ - FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); -INSTANTIATE_MMHA_DISPATCH(float, float) -INSTANTIATE_MMHA_DISPATCH(uint16_t, half) -#ifdef ENABLE_BF16 -INSTANTIATE_MMHA_DISPATCH(__nv_bfloat16, __nv_bfloat16) -#endif -#undef INSTANTIATE_MMHA_DISPATCH - -int AttentionOp::getHeadSize(bool checkInit) const -{ - if (checkInit) - { - TLLM_CHECK_WITH_INFO(mHeadSize > 0, "Trying to read mHeadSize before it's been initialized"); - } - return mHeadSize; -} - -size_t AttentionOp::getFmhaMultiCtasKvScratchSize() const noexcept -{ - static constexpr size_t kMultiCtasKvRowsPerCta = 256; - static constexpr size_t kMultiCtasKvStatsPerRow = 2; - static constexpr size_t kMultiCtasKvPartialOElementSize = 2; - - size_t const headDimV - = mIsMLAEnabled ? static_cast(mMLAParams.kv_lora_rank) : static_cast(getHeadSize()); - size_t const maxRows = kMultiCtasKvRowsPerCta * static_cast(mMultiProcessorCount); - size_t const partialStatsSize = sizeof(float) * kMultiCtasKvStatsPerRow * maxRows; - size_t const partialOSize = kMultiCtasKvPartialOElementSize * maxRows * headDimV; - - return partialStatsSize + partialOSize; -} - -size_t AttentionOp::contextMlaWorkspaceBytesPerToken(int32_t numAttnHeads, int32_t qkRopeHeadDim, int32_t qkNopeHeadDim, - int32_t vHeadDim, bool fp8ContextMla, bool separateQAndKvInput, bool sparseMla) noexcept -{ - // Only the fp8 context-MLA separate-Q/KV path stages total_kv_len-scaled K/V dequant buffers. - // Sparse MLA reads K/V directly from the paged KV cache (no staging), so its per-token cost is 0. - if (!fp8ContextMla || !separateQAndKvInput || sparseMla) - { - return 0; - } - // Mirror getWorkspaceSizeForContext's dim layout for the non-sparse fp8 branch: - // total_k_dim_all_heads = numAttnHeads * (qk_rope_head_dim + qk_nope_head_dim) - // total_v_dim_all_heads = numAttnHeads * v_head_dim - // The buffers are fp8 (1 byte/element), so bytes/token == element count. - int const dimKPerHead = qkRopeHeadDim + qkNopeHeadDim; - int const dimVPerHead = vHeadDim; - return static_cast(numAttnHeads) * static_cast(dimKPerHead + dimVPerHead); -} - -size_t AttentionOp::getWorkspaceSizeForContext(tensorrt_llm::DataType type, int32_t max_num_seq, - int32_t input_seq_length, int32_t cross_kv_length, int32_t max_num_tokens, int32_t total_kv_len) const noexcept -{ - if (max_num_tokens == 0) - { - return 0; - } - - int const local_hidden_units_qo = mNumAttnHeads * getHeadSize(); - int const local_hidden_units_kv = mNumAttnKVHeads * getHeadSize(); - - auto const size = tensorrt_llm::runtime::BufferDataType(type).getSize(); - - size_t context_workspace_size = 0; - - auto const batch_size = static_cast(max_num_seq); - auto const kv_seq_length = (isCrossAttention() ? cross_kv_length : input_seq_length); - // Unfused context attention operates on padded [batch, sequence] tensors, - // even when the input QKV is packed. Size those buffers from the padded - // token counts exactly as enqueueContext does; max_num_tokens remains the - // packed count used by the fused paths below. - size_t const padded_num_tokens = batch_size * static_cast(input_seq_length); - size_t const padded_kv_tokens = batch_size * static_cast(kv_seq_length); - size_t const attention_mask_size = mEnableContextFMHA ? 0 : size * padded_num_tokens * kv_seq_length; - size_t const cu_seqlens_size = sizeof(int) * (batch_size + 1); - size_t const rotary_inv_freq_size = sizeof(float) * batch_size * mRotaryEmbeddingDim / 2; - - size_t q_buf_2_size = 0; - if (!mEnableContextFMHA) - { - // Unfused mha - q_buf_2_size = size * batch_size * input_seq_length * local_hidden_units_qo; - } - else if (mFmhaDispatcher->isSeparateQAndKvInput()) - { - // Paged context fmha - q_buf_2_size = (mFP8ContextFMHA ? 1 : size) * max_num_tokens * local_hidden_units_qo; - } - - size_t const k_buf_2_size = mEnableContextFMHA ? 0 : size * batch_size * kv_seq_length * local_hidden_units_kv; - size_t const v_buf_2_size = mEnableContextFMHA ? 0 : size * batch_size * kv_seq_length * local_hidden_units_kv; - size_t const qk_buf_size - = mEnableContextFMHA ? 0 : size * batch_size * mNumHeads * input_seq_length * kv_seq_length; - size_t const qkv_buf_2_size = mEnableContextFMHA ? 0 : size * padded_num_tokens * local_hidden_units_qo; - size_t const qk_buf_float_size - = mEnableContextFMHA ? 0 : sizeof(float) * batch_size * mNumHeads * input_seq_length * kv_seq_length; - int dim_q_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); - int dim_k_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); - int dim_v_per_head = (mMLAParams.v_head_dim); - if (useSparseMLA()) - { - dim_q_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - dim_k_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - dim_v_per_head - = mMLAParams.rope_append ? mMLAParams.kv_lora_rank : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - } - - // Total dimension per token across all heads for Q, K, and V components respectively - int const total_q_dim_all_heads = mNumAttnHeads * dim_q_per_head; - int const total_k_dim_all_heads - = mNumAttnHeads * dim_k_per_head; // Assuming effective num_kv_heads = head_num for layout - int const total_v_dim_all_heads - = mNumAttnHeads * dim_v_per_head; // Assuming effective num_kv_heads = head_num for layout - bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mIsMLAEnabled && mFmhaDispatcher->isSeparateQAndKvInput() - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); - - // Packed fp8 qkv buffer size for normal fp8 context FMHA - size_t fp8_qkv_buffer_size = mFP8ContextFMHA && mEnableContextFMHA && !mFmhaDispatcher->isSeparateQAndKvInput() - ? max_num_tokens * size_t(local_hidden_units_qo + 2 * local_hidden_units_kv) - : 0; - // Separate fp8 q/k/v buffer size for fp8 context MLA - size_t fp8_q_buf_size = 0; - size_t fp8_k_buf_size = 0; - size_t fp8_v_buf_size = 0; - if (mEnableContextFMHA && mFP8ContextMLA && mFmhaDispatcher->isSeparateQAndKvInput()) - { - fp8_q_buf_size = max_num_tokens * static_cast(total_q_dim_all_heads); - - if (useSparseMLA()) - { - // Sparse MLA (absorption mode): K and V are stored directly in KV cache during MLA RoPE kernel. - // No separate FP8 buffers needed for K/V since they're read from paged KV cache (Q_PAGED_KV layout). - fp8_k_buf_size = 0; - fp8_v_buf_size = 0; - } - else - { - // Use total_kv_len when available (KV cache reuse causes total_kv_len >> max_num_tokens). - // enqueueContext sizes these buffers by total_kv_len, so workspace must match. - // NOTE: the per-token cost of these two buffers (total_k_dim_all_heads + total_v_dim_all_heads) is - // the single source of truth exposed via contextMlaWorkspaceBytesPerToken() for the KV-cache - // estimator's workspace reserve. Keep the two in sync if this dim layout changes. - size_t const kv_buf_tokens = std::max( - static_cast(total_kv_len), static_cast(mChunkPrefillBufferBatchSize) * max_num_tokens); - fp8_k_buf_size = kv_buf_tokens * static_cast(total_k_dim_all_heads); - fp8_v_buf_size = kv_buf_tokens * static_cast(total_v_dim_all_heads); - TLLM_CHECK(static_cast(total_k_dim_all_heads + total_v_dim_all_heads) - == contextMlaWorkspaceBytesPerToken(mNumAttnHeads, mMLAParams.qk_rope_head_dim, - mMLAParams.qk_nope_head_dim, mMLAParams.v_head_dim, mFP8ContextMLA, - /*separateQAndKvInput=*/true, useSparseMLA())); - } - } - else if (useSageAttnSeparateQkv) - { - fp8_q_buf_size = max_num_tokens * static_cast(local_hidden_units_qo); - fp8_k_buf_size = total_kv_len * static_cast(local_hidden_units_kv); - fp8_v_buf_size = total_kv_len * static_cast(local_hidden_units_kv); - } - - int32_t const q_max_n_blk - = mSageAttnNumEltsPerBlkQ > 0 ? tc::divUp(max_num_tokens, mSageAttnNumEltsPerBlkQ) + batch_size - 1 : 0; - int32_t const k_max_n_blk - = mSageAttnNumEltsPerBlkK > 0 ? tc::divUp(total_kv_len, mSageAttnNumEltsPerBlkK) + batch_size - 1 : 0; - size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); - size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); - size_t const sage_v_sfs_buffer_size = mSageAttnNumEltsPerBlkV > 0 - ? sizeof(float) * tc::divUp(local_hidden_units_kv, std::max(1, mSageAttnNumEltsPerBlkV)) - : 0; - - size_t const padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * padded_num_tokens; - size_t const encoder_padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * padded_kv_tokens; - // Each token holds (batch_idx, token_idx_in_seq) int2. - size_t const tokens_info_size = sizeof(int2) * max_num_tokens; - size_t const fmha_scheduler_counter = mEnableContextFMHA ? sizeof(uint32_t) : 0; - size_t const fmha_bmm1_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) * 2 : 0; - size_t const fmha_bmm2_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) : 0; - - // cp workspace size upper bound - size_t const cpMaxPaddedSequenceLength = max_num_tokens + batch_size * (mCpSize - 1); - size_t const cpWorkspaceSize = mCpSize == 1 - ? 0 - : (2 * size * cpMaxPaddedSequenceLength * getHeadSize() * (mNumHeads + 2 * mNumKVHeads) + cu_seqlens_size); - - size_t const fmha_multi_ctas_kv_scratch_size = useTllmGenSparseAttention() ? getFmhaMultiCtasKvScratchSize() : 0; - - AttentionContextWorkspaceSizes workspaceSizes{}; - workspaceSizes.attentionMask = attention_mask_size; - workspaceSizes.cuQSeqlens = cu_seqlens_size; - workspaceSizes.cuKvSeqlens = cu_seqlens_size; - workspaceSizes.cuMaskRows = cu_seqlens_size; - workspaceSizes.rotaryInvFreq = rotary_inv_freq_size; - workspaceSizes.qBuf = q_buf_2_size; - workspaceSizes.kBuf = k_buf_2_size; - workspaceSizes.vBuf = v_buf_2_size; - workspaceSizes.qkBuf = qk_buf_size; - workspaceSizes.qkvBuf = qkv_buf_2_size; - workspaceSizes.qkFloatBuf = qk_buf_float_size; - workspaceSizes.fp8QkvBuf = fp8_qkv_buffer_size; - workspaceSizes.fp8QBuf = fp8_q_buf_size; - workspaceSizes.fp8KBuf = fp8_k_buf_size; - workspaceSizes.fp8VBuf = fp8_v_buf_size; - workspaceSizes.paddingOffset = padding_offset_size; - workspaceSizes.encoderPaddingOffset = encoder_padding_offset_size; - workspaceSizes.tokensInfo = tokens_info_size; - workspaceSizes.fmhaTileCounter = fmha_scheduler_counter; - workspaceSizes.fmhaBmm1Scale = fmha_bmm1_scale_size; - workspaceSizes.fmhaBmm2Scale = fmha_bmm2_scale_size; - workspaceSizes.sageQScale = sage_q_sfs_buffer_size; - workspaceSizes.sageKScale = sage_k_sfs_buffer_size; - workspaceSizes.sageVScale = sage_v_sfs_buffer_size; - workspaceSizes.cpWorkspace = cpWorkspaceSize; - workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; - context_workspace_size = AttentionWorkspaceManager::buildContextLayout(workspaceSizes).totalSize; - - return context_workspace_size; -} - -size_t AttentionOp::getWorkspaceSizeForGeneration(tensorrt_llm::DataType type, int32_t max_num_seq, - int32_t max_attention_window_size, int32_t max_num_tokens, int32_t max_blocks_per_sequence) const noexcept -{ - if (max_num_tokens == 0) - { - return 0; - } - - auto const size = tensorrt_llm::runtime::BufferDataType(type).getSize(); - int const batch_beam = max_num_seq; - - // Compute the workspace size for MLA. - size_t fmha_v2_mla_workspace_size = 0; - if (mIsMLAEnabled) - { - size_t flash_mla_workspace_size = 0; - if (mUseGenFlashMLA) - { - static constexpr int TileSchedulerMetaDataSize = 8; - - int s_q = mMLAParams.predicted_tokens_per_seq; - - int num_q_heads = mNumHeads / mCpSize; - int num_kv_heads = mNumKVHeads; - int head_size_v = (mUseSparseAttention && !mMLAParams.rope_append) - ? mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim - : mMLAParams.kv_lora_rank; - - int num_sm_parts = getFlashMlaNumSmParts(s_q, num_q_heads, num_kv_heads, head_size_v); - - AttentionFlashMlaWorkspaceSizes flashMlaWorkspaceSizes{}; - flashMlaWorkspaceSizes.tileSchedulerMetadata = sizeof(int) * (num_sm_parts * TileSchedulerMetaDataSize); - flashMlaWorkspaceSizes.numSplits = sizeof(int) * (batch_beam + 1); - flashMlaWorkspaceSizes.softmaxLse = sizeof(float) * (batch_beam * s_q * num_q_heads); - flashMlaWorkspaceSizes.softmaxLseAccum = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q); - flashMlaWorkspaceSizes.outAccum - = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q * head_size_v); - flash_mla_workspace_size = AttentionWorkspaceManager::buildFlashMlaLayout(flashMlaWorkspaceSizes).totalSize; - } - - size_t const cu_seqlens_size = sizeof(int) * (max_num_seq + 1); - size_t const fmha_scheduler_counter = sizeof(uint32_t); - size_t const fmha_multi_ctas_kv_scratch_size = getFmhaMultiCtasKvScratchSize(); - - int const NUM_BUFFERS = 5; - size_t workspaces[NUM_BUFFERS]; - workspaces[0] = mIsGenerationMLA ? 0 : cu_seqlens_size; // cu_q_len - workspaces[1] = mIsGenerationMLA ? 0 : cu_seqlens_size; // cu_kv_len - workspaces[2] = mIsGenerationMLA ? 0 : fmha_scheduler_counter; - workspaces[3] = fmha_multi_ctas_kv_scratch_size; - workspaces[4] = flash_mla_workspace_size; - - fmha_v2_mla_workspace_size = tc::calculateTotalWorkspaceSize(workspaces, NUM_BUFFERS); - } - - size_t generation_workspace_size = 0; - // The minimum number of sequence length tiles (limited by the shared memory size). - int minSeqLenTile - = estimate_min_multi_block_count(max_attention_window_size, mMaxSharedMemoryPerBlockOptin - 2048, size); - int32_t const maxSeqLenTile - = std::max({minSeqLenTile, getMaxNumSeqLenTile(batch_beam), (int) tc::divUp(mMultiProcessorCount, mNumHeads)}); - - size_t const partial_out_size = size * batch_beam * mNumHeads * mHeadSize * maxSeqLenTile; - size_t const partial_sum_size = sizeof(float) * batch_beam * mNumHeads * maxSeqLenTile; - size_t const partial_max_size = sizeof(float) * batch_beam * mNumHeads * maxSeqLenTile; - size_t const shift_k_cache_size = (!mPosShiftEnabled || isCrossAttention()) - ? 0 - : size * batch_beam * mNumHeads * mHeadSize * max_attention_window_size; - size_t const cpMaxPaddedSequenceLength = (batch_beam + mCpSize - 1) / mCpSize * mCpSize; - size_t const cpWorkspaceSize - = mCpSize == 1 ? 0 : (2 * size * cpMaxPaddedSequenceLength * getHeadSize() * (mNumHeads + 2 * mNumKVHeads)); - - AttentionGenerationWorkspaceSizes generationWorkspaceSizes{}; - generationWorkspaceSizes.cpWorkspace = cpWorkspaceSize; - generationWorkspaceSizes.partialOut = partial_out_size; - generationWorkspaceSizes.partialSum = partial_sum_size; - generationWorkspaceSizes.partialMax = partial_max_size; - generationWorkspaceSizes.shiftKCache = shift_k_cache_size; - { - auto const cascadeSizes - = tensorrt_llm::kernels::mmha::cascade::getCascadeWorkspaceSizes(batch_beam, mNumHeads, mHeadSize); - generationWorkspaceSizes.cascadeOut = cascadeSizes.out; - generationWorkspaceSizes.cascadeMax = cascadeSizes.mMax; - generationWorkspaceSizes.cascadeSum = cascadeSizes.lSum; - } - generation_workspace_size = AttentionWorkspaceManager::buildGenerationLayout(generationWorkspaceSizes).totalSize; - - size_t xqa_workspace_size = 0; - if (mEnableXQA) - { - size_t const cu_seqlens_size = sizeof(int) * (batch_beam + 1); - size_t const cu_kv_seqlens_size = sizeof(int) * (batch_beam + 1); - size_t const rotary_inv_freq_size = sizeof(float) * batch_beam * mRotaryEmbeddingDim / 2; - // Two workspaces for sparse attention. One for the sequence lengths, and one for kv block offsets. - size_t const sparse_attn_cache_size = useTllmGenSparseAttentionPaged() - ? sizeof(int) * (batch_beam + batch_beam * 2 * max_blocks_per_sequence) * mNumKVHeads - : 0; - AttentionXqaWorkspaceSizes xqaWorkspaceSizes{}; - xqaWorkspaceSizes.cuSeqlens = cu_seqlens_size; - xqaWorkspaceSizes.cuKvSeqlens = cu_kv_seqlens_size; - xqaWorkspaceSizes.rotaryInvFreq = rotary_inv_freq_size; - xqaWorkspaceSizes.tokensInfo = max_num_tokens * sizeof(int2); - xqaWorkspaceSizes.bmm1Scale = sizeof(float) * 2; - xqaWorkspaceSizes.bmm2Scale = sizeof(float); - xqaWorkspaceSizes.sparseAttnCache = sparse_attn_cache_size; - xqaWorkspaceSizes.kernelWorkspace = mXqaDispatcher->getWorkspaceSize( - std::min(mSpecDecodingMaxGenerationLength * max_num_seq, max_num_tokens)); - xqa_workspace_size - = AttentionWorkspaceManager::buildXqaLayout(xqaWorkspaceSizes, mXqaDispatcher->getWorkspaceAlignment()) - .totalSize; - } - - return std::max(std::max(generation_workspace_size, xqa_workspace_size), fmha_v2_mla_workspace_size); -} - -int AttentionOp::getMaxNumSeqLenTile(int batch_beam_size) const -{ - if (mMultiBlockMode) - { - // And we allocate the buffer based on the maximum number of blocks per sequence (batch_beam_size = 1). - // Assume we can only have 1 block (large block size like 1024) in SM, and we only want one wave of blocks. - return tc::getEnvMmhaMultiblockDebug() ? std::max(kReservedMaxSeqLenTilePerSeq, getEnvMmhaBlocksPerSequence()) - : tc::divUp(mMultiProcessorCount, batch_beam_size * mNumHeads); - } - return 0; -} - -template -int AttentionOp::mlaGeneration( - MlaParams& params, EnqueueGenerationParams const& generation_params, cudaStream_t stream) -{ - TLLM_CHECK_WITH_INFO(params.seqQOffset != nullptr, "seqQOffset is nullptr."); - TLLM_CHECK_WITH_INFO(params.cache_seq_lens != nullptr, "cache_seq_lens is nullptr."); - TLLM_CHECK_WITH_INFO(params.fmha_tile_counter != nullptr, "fmha_tile_counter is nullptr."); - if (mFP8GenerationMLA) - { - TLLM_CHECK_WITH_INFO(params.quant_q_buf != nullptr, "quant_q_buf is nullptr."); - TLLM_CHECK_WITH_INFO(params.bmm1_scale != nullptr, "bmm1_scale is nullptr."); - TLLM_CHECK_WITH_INFO(params.bmm2_scale != nullptr, "bmm2_scale is nullptr."); - } - - int const num_kv_heads = 1; - int const head_size = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - int const head_size_v = (useSparseMLA() && !mMLAParams.rope_append) ? head_size : mMLAParams.kv_lora_rank; - int32_t const batch_beam = generation_params.beam_width * generation_params.num_requests; - - // The element size of the KV cache. - auto const elemSize = mFP8GenerationMLA ? sizeof(__nv_fp8_e4m3) : sizeof(T); - auto const sizePerToken = num_kv_heads * head_size * elemSize; - params.cache_type = (mFP8GenerationMLA ? KvCacheDataType::FP8 : KvCacheDataType::BASE); - - auto kv_cache_buffer = KVBlockArray(batch_beam, generation_params.max_blocks_per_sequence, mTokensPerBlock, - sizePerToken, generation_params.cyclic_attention_window_size, - generation_params.max_cyclic_attention_window_size, generation_params.sink_token_length, - generation_params.can_use_one_more_block, generation_params.host_primary_pool_pointer, - generation_params.host_secondary_pool_pointer, generation_params.block_offsets); - - // Static sparse NVFP4 MLA reads a separately dequantized FP8 scratch pool, - // so this paged-cache scale descriptor is not consumed by the attention kernel. - auto kv_scale_cache_buffer = KVBlockArray(); - - void* scratchPtr = params.workspace; - - params.quant_scale_o = generation_params.attention_output_orig_quant; - params.quant_scale_q = generation_params.kv_scale_orig_quant; - params.quant_scale_kv = generation_params.kv_scale_orig_quant; - params.dequant_scale_q = generation_params.kv_scale_quant_orig; - params.dequant_scale_kv = generation_params.kv_scale_quant_orig; - params.host_bmm1_scale - = 1 / (mQScaling * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); - - if (generation_params.runtime_perf_knobs) - { - int64_t multi_block_mode_val = generation_params.runtime_perf_knobs[0]; - mMultiBlockMode = multi_block_mode_val == 1; - int64_t enable_context_fmha_fp32_acc_val = generation_params.runtime_perf_knobs[1]; - mFMHAForceFP32Acc = mFMHAForceFP32Acc || enable_context_fmha_fp32_acc_val == 1; - } - - if (common::getEnvForceDeterministicAttention()) - { - mMultiBlockMode = false; - } - - if (mUseTllmGen) - { - TLLM_CHECK_WITH_INFO(mTllmGenFMHARunner.get(), "mTllmGenFMHARunner not initialized."); - TllmGenFmhaRunnerParams tllmRunnerParams{}; - - // Parameters to select kernels. - // MLA generation kernels use dense mask. For multi-token generation, TRTLLM-Gen applies causality by - // shrinking each token's effective KV length. - tllmRunnerParams.mMaskType = TrtllmGenAttentionMaskType::Dense; - tllmRunnerParams.mKernelType = FmhaKernelType::Generation; - tllmRunnerParams.mMultiCtasKvMode = mMultiBlockMode; - // Note that the tileScheduler and multiCtasKvMode will be automatically tuned when using multi_block mode. - // Otherwise, always enable the persistent scheduler for better performance. - tllmRunnerParams.mTileScheduler = mMultiBlockMode ? TileScheduler::Static : TileScheduler::Persistent; - - // Q buffer. - tllmRunnerParams.qPtr = mFP8GenerationMLA ? reinterpret_cast(params.quant_q_buf) - : reinterpret_cast(params.q_buf); - - // KV buffer - // Paged KV - tllmRunnerParams.mQkvLayout = QkvLayout::PagedKv; - tllmRunnerParams.kvPtr = kv_cache_buffer.mPrimaryPoolPtr; - tllmRunnerParams.kvPageIdxPtr = reinterpret_cast(kv_cache_buffer.data); - tllmRunnerParams.mMaxNumPagesPerSeqKv = kv_cache_buffer.mMaxBlocksPerSeq; - tllmRunnerParams.mNumTokensPerPage = kv_cache_buffer.mTokensPerBlock; - - // The partial buffers' pointers when the multiCtasKv mode is enabled. - tllmRunnerParams.multiCtasKvCounterPtr = generation_params.semaphores; - tllmRunnerParams.multiCtasKvScratchPtr = scratchPtr; - - // The sequence lengths for K/V. - tllmRunnerParams.seqLensKvPtr = params.cache_seq_lens; - - tllmRunnerParams.oPtr = reinterpret_cast(params.context_buf); - tllmRunnerParams.oSfPtr = generation_params.context_buf_sf; - if (params.dsv4_epilogue_fusion.enabled) - { - tllmRunnerParams.mDsv4EpilogueFusion.enabled = true; - tllmRunnerParams.mDsv4EpilogueFusion.cosSinCache = params.dsv4_epilogue_fusion.cos_sin_cache; - tllmRunnerParams.mDsv4EpilogueFusion.scaleBufM = params.dsv4_epilogue_fusion.scale_buf_m; - } - - // softmax stats if needed - tllmRunnerParams.softmaxStatsPtr = generation_params.softmax_stats; - - // Per-head attention sink added to the softmax denominator. - tllmRunnerParams.attentionSinksPtr = generation_params.attention_sinks; - - // MLA uses different head dimensions for Qk and V. - tllmRunnerParams.mHeadDimQk = head_size; - tllmRunnerParams.mHeadDimV = head_size_v; - - auto const num_q_heads = mNumAttnHeads; - tllmRunnerParams.mNumHeadsQ = num_q_heads; - tllmRunnerParams.mNumHeadsKv = num_kv_heads; - tllmRunnerParams.mNumHeadsQPerKv = num_q_heads / num_kv_heads; - - tllmRunnerParams.mBatchSize = batch_beam; - // It is used to construct contiguous kv cache TMA descriptors. - tllmRunnerParams.mMaxSeqLenCacheKv = generation_params.max_attention_window_size; - // This should be set to numDraftTokens + 1. - tllmRunnerParams.mMaxSeqLenQ = params.acc_q_len / batch_beam; - tllmRunnerParams.mMaxSeqLenKv = generation_params.max_past_kv_length; - tllmRunnerParams.mJITWarmup = generation_params.trtllm_gen_jit_warmup; - tllmRunnerParams.mJITWarmupMaxNumRequests = mMaxNumRequests; - tllmRunnerParams.mJITWarmupMaxSeqLenQ = mMaxContextLength; - tllmRunnerParams.mJITWarmupMaxSeqLenKv = mMaxSeqLen; - tllmRunnerParams.mSumOfSeqLensQ = int(batch_beam * tllmRunnerParams.mMaxSeqLenQ); - // Not used in the generation kernels as contiguous_kv or paged_kv layouts are used. - tllmRunnerParams.mSumOfSeqLensKv = int(batch_beam * tllmRunnerParams.mMaxSeqLenKv); - - // The attention window size. - tllmRunnerParams.mAttentionWindowSize = generation_params.cyclic_attention_window_size; - // The chunked attention size. - tllmRunnerParams.mChunkedAttentionSize = INT_MAX; - - // The scaleQ that will be applied to the BMM1 output. - tllmRunnerParams.mScaleQ = mQScaling * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) - / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); - - // Set it to INT_MAX as the kv cache pageOffsets will ensure that there is no out-of-bounds access. - tllmRunnerParams.mNumPagesInMemPool = INT_MAX; - tllmRunnerParams.mMultiProcessorCount = mMultiProcessorCount; - tllmRunnerParams.stream = stream; - tllmRunnerParams.mSfStartTokenIdx = generation_params.start_token_idx_sf; - tllmRunnerParams.mSkipCorrThreshold = mSkipCorrectionThreshold; - - // Scales for quantization - if (mFP8GenerationMLA) - { - static constexpr int bmm1_scale_offset = 1; - tllmRunnerParams.outputScalePtr = reinterpret_cast(params.bmm2_scale); - tllmRunnerParams.scaleSoftmaxLog2Ptr - = reinterpret_cast(params.bmm1_scale) + bmm1_scale_offset; - } - - // Set the following parameters if sparseAttention is used. - if (useSparseMLA()) - { - bool const useDynamicSparseMLA = mRuntimeSparseAttentionParams.sparse_attn_kv_lens != nullptr; - tllmRunnerParams.mSparseAttention - = useDynamicSparseMLA ? SparseType::DynamicTokenSparse : SparseType::StaticTokenSparse; - tllmRunnerParams.mSparseTopK = mRuntimeSparseAttentionParams.num_sparse_topk; - tllmRunnerParams.ptrSparseMlaTopKLens = mRuntimeSparseAttentionParams.sparse_attn_kv_lens; - tllmRunnerParams.kvPageIdxPtr = reinterpret_cast( - mRuntimeSparseAttentionParams.sparse_attn_indices); - if (useDynamicSparseMLA) - { - TLLM_CHECK_WITH_INFO(mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool != nullptr, - "SWA KV pool must be set for dynamic sparse MLA."); - // Dynamic sparse MLA always has an SWA pool. The compressed pool is optional; when it - // is absent (ratio == 1), use SWA as kvPtr only to keep TG's primary TMA descriptor valid. - tllmRunnerParams.kvPtr = mRuntimeSparseAttentionParams.sparse_kv_cache_pool != nullptr - ? mRuntimeSparseAttentionParams.sparse_kv_cache_pool - : mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool; - tllmRunnerParams.slidingWindowKvPoolBasePtr - = mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool; - } - else - { - tllmRunnerParams.kvPtr = mRuntimeSparseAttentionParams.sparse_kv_cache_pool; - - if (mUseNvfp4MlaKvCache) - { - // Static sparse MLA indexes a compact KV pool containing at most - // mSparseTopK rows per query. Do not let the original dense KV - // length drive kernel selection or launch geometry: for long - // sequences that can select a multi-CTA kernel which addresses - // beyond the compact page table. - TLLM_CHECK_WITH_INFO(tllmRunnerParams.mSparseTopK > 0, - "Static sparse MLA requires a positive TopK, got %d", tllmRunnerParams.mSparseTopK); - int32_t const originalMaxSeqLenKv = tllmRunnerParams.mMaxSeqLenKv; - int32_t const effectiveMaxSeqLenKv = std::min(originalMaxSeqLenKv, tllmRunnerParams.mSparseTopK); - tllmRunnerParams.mMaxSeqLenKv = effectiveMaxSeqLenKv; - tllmRunnerParams.mJITWarmupMaxSeqLenKv - = std::min(tllmRunnerParams.mJITWarmupMaxSeqLenKv, effectiveMaxSeqLenKv); - int64_t const sumOfSeqLensKv - = static_cast(tllmRunnerParams.mBatchSize) * effectiveMaxSeqLenKv; - TLLM_CHECK_WITH_INFO(sumOfSeqLensKv <= std::numeric_limits::max(), - "Static sparse MLA cumulative KV length exceeds int32 capacity: %ld", sumOfSeqLensKv); - tllmRunnerParams.mSumOfSeqLensKv = static_cast(sumOfSeqLensKv); - TLLM_LOG_DEBUG("Clamp static sparse MLA max KV length from %d to %d (TopK=%d)", originalMaxSeqLenKv, - effectiveMaxSeqLenKv, tllmRunnerParams.mSparseTopK); - } - } - } - - mTllmGenFMHARunner->run(tllmRunnerParams); - sync_check_cuda_error(stream); - } - else if (mUseGenFlashMLA) - { - static constexpr int TileSchedulerMetaDataSize = 8; - - int const num_q_heads = mNumHeads / mCpSize; - int const ngroups = num_q_heads / num_kv_heads; - - int const s_q = params.acc_q_len / batch_beam; - assert(s_q == mMLAParams.predicted_tokens_per_seq); - int const head_size_v = mMLAParams.kv_lora_rank; - int const num_sm_parts = getFlashMlaNumSmParts(s_q, num_q_heads, num_kv_heads, head_size_v); - - size_t const num_splits_size = sizeof(int) * (batch_beam + 1); - size_t const tile_scheduler_metadata_size = sizeof(int) * (num_sm_parts * TileSchedulerMetaDataSize); - size_t const softmax_lse_size = sizeof(float) * (batch_beam * s_q * num_q_heads * num_kv_heads); // softmax_lse - size_t const softmax_lse_accum_size = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q); - size_t const out_accum_size = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q * head_size_v); - - AttentionFlashMlaWorkspaceSizes flashMlaWorkspaceSizes{}; - flashMlaWorkspaceSizes.tileSchedulerMetadata = tile_scheduler_metadata_size; - flashMlaWorkspaceSizes.numSplits = num_splits_size; - flashMlaWorkspaceSizes.softmaxLse = softmax_lse_size; - flashMlaWorkspaceSizes.softmaxLseAccum = softmax_lse_accum_size; - flashMlaWorkspaceSizes.outAccum = out_accum_size; - auto const flashMlaWorkspaceLayout = AttentionWorkspaceManager::buildFlashMlaLayout(flashMlaWorkspaceSizes); - float* softmax_lse_ptr - = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.softmaxLse); - float* softmax_lse_accum_ptr - = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.softmaxLseAccum); - float* out_accum_ptr - = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.outAccum); - - // Metadata must always be pre-computed by Python (compute_flash_mla_metadata) and passed in. - TLLM_CHECK_WITH_INFO(params.flash_mla_tile_scheduler_metadata != nullptr, - "FlashMLA tile-scheduler metadata must be pre-computed by Python."); - TLLM_CHECK_WITH_INFO( - params.flash_mla_num_splits != nullptr, "FlashMLA num_splits must be pre-computed by Python."); - int* tile_scheduler_metadata_ptr = const_cast(params.flash_mla_tile_scheduler_metadata); - int* num_splits_ptr = const_cast(params.flash_mla_num_splits); - - Flash_fwd_mla_params flashMlaParams{}; - flashMlaParams.b = batch_beam; - flashMlaParams.seqlen_q = ngroups * s_q; - flashMlaParams.cu_seqlens_k = const_cast(params.cache_seq_lens); - flashMlaParams.h = 1; - flashMlaParams.h_h_k_ratio = 1; - - float softmax_scale - = 1.0f / (mQScaling * sqrtf((mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim) * 1.0f)); - - flashMlaParams.ngroups = ngroups; - flashMlaParams.is_causal = !(s_q == 1); - flashMlaParams.d = head_size; - flashMlaParams.d_v = head_size_v; - flashMlaParams.scale_softmax = softmax_scale; - flashMlaParams.scale_softmax_log2 = float(softmax_scale * M_LOG2E); - - flashMlaParams.q_ptr = mFP8GenerationMLA ? const_cast(reinterpret_cast(params.quant_q_buf)) - : const_cast(reinterpret_cast(params.q_buf)); - flashMlaParams.k_ptr = kv_cache_buffer.mPrimaryPoolPtr; - flashMlaParams.v_ptr = flashMlaParams.k_ptr; - flashMlaParams.o_ptr = reinterpret_cast(params.context_buf); - flashMlaParams.softmax_lse_ptr = softmax_lse_ptr; - - // since head_num_kv = 1 - flashMlaParams.q_batch_stride = head_size * params.head_num * s_q; - flashMlaParams.k_batch_stride = mTokensPerBlock * num_kv_heads * head_size * mMLAParams.num_layers; - flashMlaParams.o_batch_stride = s_q * num_q_heads * head_size_v; - flashMlaParams.q_row_stride = head_size; - flashMlaParams.k_row_stride = head_size; - flashMlaParams.o_row_stride = head_size_v; - flashMlaParams.q_head_stride = head_size; - flashMlaParams.k_head_stride = head_size; - flashMlaParams.o_head_stride = head_size_v; - - flashMlaParams.v_batch_stride = flashMlaParams.k_batch_stride; - flashMlaParams.v_row_stride = flashMlaParams.k_row_stride; - flashMlaParams.v_head_stride = flashMlaParams.k_head_stride; - - flashMlaParams.block_table = const_cast(params.block_ids_per_seq); - flashMlaParams.block_table_batch_stride = generation_params.max_blocks_per_sequence; - flashMlaParams.page_block_size = mTokensPerBlock; - - flashMlaParams.descale_q_ptr = const_cast(params.dequant_scale_q); - flashMlaParams.descale_k_ptr = const_cast(params.dequant_scale_kv); - - flashMlaParams.tile_scheduler_metadata_ptr = tile_scheduler_metadata_ptr; - flashMlaParams.num_sm_parts = num_sm_parts; - flashMlaParams.num_splits_ptr = num_splits_ptr; - - flashMlaParams.softmax_lseaccum_ptr = softmax_lse_accum_ptr; - flashMlaParams.oaccum_ptr = out_accum_ptr; - - if constexpr (std::is_same::value) - { - if (mFP8GenerationMLA) - { - TLLM_THROW("FP8 KV cache MLA is only supported for bf16 output"); - } - else - { - run_mha_fwd_splitkv_mla(flashMlaParams, stream); - } - } - else if constexpr (std::is_same::value) - { - if (mFP8GenerationMLA) - { - run_mha_fwd_splitkv_mla(flashMlaParams, stream); - } - else - { - run_mha_fwd_splitkv_mla(flashMlaParams, stream); - } - } - else - { - TLLM_THROW("Unsupported data type for FlashMLA"); - } - } - else - { - // Try XQA optimization first when CP is not used. - if (mCpSize == 1) - { - // NOTE: input_seq_length = num_medusa_tokens + 1 (new generated one from the original LM head) - // self attn - XQAParams xqaParams{}; - this->template convertMMHAParamsToXQAParams( - xqaParams, generation_params, /*forConfigurePlugin=*/false); - xqaParams.quant_q_buffer_ptr = params.quant_q_buf; - xqaParams.q_scaling - = 1 / (mQScaling * sqrtf((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); - if (mEnableXQA && mXqaDispatcher->shouldUse(xqaParams)) - { - TLLM_LOG_DEBUG("XQA kernels are selected in the generation phase."); - xqaParams.stream = stream; - mXqaDispatcher->run(xqaParams, kv_cache_buffer, kv_scale_cache_buffer); - return 0; - } - } - - // Use FMHA otherwise. - MHARunnerParams fmhaParams{}; - fmhaParams.b = batch_beam; - fmhaParams.numGroupedHeads = params.head_num; - fmhaParams.qSeqLen = params.head_num * (params.acc_q_len / batch_beam); - fmhaParams.kvSeqLen = generation_params.max_past_kv_length; - // Disable sliding window attention when it is not needed. - fmhaParams.slidingWindowSize = generation_params.cyclic_attention_window_size; - fmhaParams.totalQSeqLen = batch_beam * fmhaParams.qSeqLen; - // TODO: set it correctly for contiguous kv buffer (cross-attention). - // fmhaParams.totalKvSeqLen = params.num_tokens; - // Device buffer pointers. - // fmhaParams.qkvPtr = reinterpret_cast(params.attention_input); - fmhaParams.qPtr = mFP8GenerationMLA ? reinterpret_cast(params.quant_q_buf) - : reinterpret_cast(params.q_buf); - // TODO: add contiguous kv buffer (cross-attention). - fmhaParams.kvPtr = nullptr; - - fmhaParams.outputPtr = reinterpret_cast(params.context_buf); - - // fmhaParams.packedMaskPtr = params.fmha_custom_mask; - fmhaParams.pagedKvCache = kv_cache_buffer; - fmhaParams.cuQSeqLenPtr = params.seqQOffset; - fmhaParams.kvSeqLenPtr = params.cache_seq_lens; - fmhaParams.cuKvSeqLenPtr = params.cu_kv_seqlens; - fmhaParams.cuMaskRowsPtr = nullptr; // mla not support custorm mask right now - fmhaParams.tileCounterPtr = params.fmha_tile_counter; - fmhaParams.scaleBmm1Ptr = reinterpret_cast(params.bmm1_scale); - fmhaParams.scaleBmm2Ptr = reinterpret_cast(params.bmm2_scale); - fmhaParams.stream = stream; - fmhaParams.forceFp32Acc = mFMHAForceFP32Acc; - - // Sparse attention parameters - if (useSparseMLA()) - { - fmhaParams.sparse_params = mRuntimeSparseAttentionParams; - } - - // MLA does not support skip-softmax attention right now - - // Run the fmha kernel - mDecoderFMHARunner->run(fmhaParams); - } - - sync_check_cuda_error(stream); - return 0; -} - -#define MLA_FUNC_DEFINE(T) \ - template int AttentionOp::mlaGeneration( \ - MlaParams & params, EnqueueGenerationParams const& generation_params, cudaStream_t stream); - -MLA_FUNC_DEFINE(float) -MLA_FUNC_DEFINE(half) -#ifdef ENABLE_BF16 -MLA_FUNC_DEFINE(__nv_bfloat16) -#endif - -template -int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStream_t stream) -{ - int const headSize = getHeadSize(); - - int const local_hidden_units_qo = mNumHeads * headSize; - int const local_hidden_units_kv = mNumAttnKVHeads * headSize; - PositionEmbeddingType const position_embedding_type = mPositionEmbeddingType; - float const q_scaling = mQScaling; - - KVCacheBuffer kv_cache_buffer; - KVCacheBuffer kv_scale_cache_buffer; - - auto sizePerToken = mNumAttnKVHeads * headSize * getKvCacheElemSizeInBits() / 8 /*bits*/; - - if (useKVCache()) - { - auto buffers = buildKvCacheBuffers(params.batch_size, params.max_blocks_per_sequence, - mTokensPerBlock, sizePerToken, params.cyclic_attention_window_size, params.max_cyclic_attention_window_size, - params.sink_token_length, params.can_use_one_more_block, params.host_primary_pool_pointer, - params.host_secondary_pool_pointer, params.host_primary_block_scale_pool_pointer, - params.host_secondary_block_scale_pool_pointer, params.block_offsets, mKVCacheQuantMode.hasFp4KvCache(), - isCrossAttention() ? params.cross_kv_length : params.max_attention_window_size, params.key_value_cache); - kv_cache_buffer = buffers.kvCacheBuffer; - kv_scale_cache_buffer = buffers.kvScaleCacheBuffer; - } - - auto cublasHandle = mCublasWrapper->getCublasHandle(); - TLLM_CUDA_CHECK(cublasSetStream(cublasHandle, stream)); - mCublasWrapper->setStream(stream); - mCublasWrapper->setWorkspace(params.workspace); - if constexpr (std::is_same_v) - { - mCublasWrapper->setFP16GemmConfig(); - } - else if constexpr (std::is_same_v) - { - mCublasWrapper->setFP32GemmConfig(); - } -#ifdef ENABLE_BF16 - else if constexpr (std::is_same_v) - { - mCublasWrapper->setBF16GemmConfig(); - } -#endif - - size_t const kv_seq_length = (isCrossAttention() ? params.cross_kv_length : params.input_seq_length); - size_t const attention_mask_size - = mEnableContextFMHA ? 0 : sizeof(T) * params.batch_size * params.input_seq_length * kv_seq_length; - size_t const cu_seqlens_size = sizeof(int) * (params.batch_size + 1); - size_t const rotary_inv_freq_size = sizeof(float) * params.batch_size * mRotaryEmbeddingDim / 2; - size_t q_buf_2_size = 0; - if (!mEnableContextFMHA) - { - // Unfused mha - q_buf_2_size = sizeof(T) * params.batch_size * params.input_seq_length * local_hidden_units_qo; - } - else if (mFmhaDispatcher->isSeparateQAndKvInput()) - { - // Paged context fmha - q_buf_2_size = (mFP8ContextFMHA ? 1 : sizeof(T)) * params.num_tokens * local_hidden_units_qo; - } - - size_t const k_buf_2_size - = mEnableContextFMHA ? 0 : sizeof(T) * params.batch_size * kv_seq_length * local_hidden_units_kv; - size_t const v_buf_2_size - = mEnableContextFMHA ? 0 : sizeof(T) * params.batch_size * kv_seq_length * local_hidden_units_kv; - size_t const qk_buf_size - = mEnableContextFMHA ? 0 : sizeof(T) * params.batch_size * mNumHeads * params.input_seq_length * kv_seq_length; - size_t const qkv_buf_2_size - = mEnableContextFMHA ? 0 : sizeof(T) * params.batch_size * params.input_seq_length * local_hidden_units_qo; - size_t const qk_buf_float_size = mEnableContextFMHA - ? 0 - : sizeof(float) * params.batch_size * mNumHeads * params.input_seq_length * kv_seq_length; - int dim_q_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); - int dim_k_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); - int dim_v_per_head = (mMLAParams.v_head_dim); - if (useSparseMLA()) - { - dim_q_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - dim_k_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - dim_v_per_head - = mMLAParams.rope_append ? mMLAParams.kv_lora_rank : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - } - - // Total dimension per token across all heads for Q, K, and V components respectively - int const total_q_dim_all_heads = mNumAttnHeads * dim_q_per_head; - int const total_k_dim_all_heads - = mNumAttnHeads * dim_k_per_head; // Assuming effective num_kv_heads = head_num for layout - int const total_v_dim_all_heads - = mNumAttnHeads * dim_v_per_head; // Assuming effective num_kv_heads = head_num for layout - // Packed fp8 qkv buffer size for normal fp8 context FMHA - size_t fp8_qkv_buffer_size = mEnableContextFMHA && mFP8ContextFMHA && !mFmhaDispatcher->isSeparateQAndKvInput() - ? params.num_tokens * (local_hidden_units_qo + 2 * local_hidden_units_kv) - : 0; - // Separate fp8 q/k/v buffer size for fp8 context MLA - size_t fp8_q_buf_size = 0; - size_t fp8_k_buf_size = 0; - size_t fp8_v_buf_size = 0; - bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mIsMLAEnabled && mFmhaDispatcher->isSeparateQAndKvInput() - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); - if (mEnableContextFMHA && mFP8ContextMLA && mFmhaDispatcher->isSeparateQAndKvInput()) - { - fp8_q_buf_size = params.num_tokens * static_cast(total_q_dim_all_heads); - - if (useSparseMLA()) - { - // Sparse MLA (absorption mode): K and V are stored directly in KV cache during MLA RoPE kernel. - // No separate FP8 buffers needed for K/V since they're read from paged KV cache (Q_PAGED_KV layout). - fp8_k_buf_size = 0; - fp8_v_buf_size = 0; - } - else - { - fp8_k_buf_size = params.total_kv_len * static_cast(total_k_dim_all_heads); - fp8_v_buf_size = params.total_kv_len * static_cast(total_v_dim_all_heads); - } - } - else if (useSageAttnSeparateQkv) - { - fp8_q_buf_size = params.num_tokens * static_cast(local_hidden_units_qo); - fp8_k_buf_size = params.total_kv_len * static_cast(local_hidden_units_kv); - fp8_v_buf_size = params.total_kv_len * static_cast(local_hidden_units_kv); - } - - int32_t const q_max_n_blk = mSageAttnNumEltsPerBlkQ > 0 - ? tc::divUp(params.num_tokens, mSageAttnNumEltsPerBlkQ) + params.batch_size - 1 - : 0; - int32_t const k_max_n_blk = mSageAttnNumEltsPerBlkK > 0 - ? tc::divUp(params.total_kv_len, mSageAttnNumEltsPerBlkK) + params.batch_size - 1 - : 0; - int32_t const v_max_n_blk - = mSageAttnNumEltsPerBlkV > 0 ? tc::divUp(local_hidden_units_kv, mSageAttnNumEltsPerBlkV) : 0; - size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); - size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); - size_t const sage_v_sfs_buffer_size = sizeof(float) * v_max_n_blk; - - size_t const padding_offset_size - = mEnableContextFMHA ? 0 : sizeof(int) * params.batch_size * params.input_seq_length; - size_t const encoder_padding_offset_size - = mEnableContextFMHA ? 0 : sizeof(int) * params.batch_size * params.cross_kv_length; - // Each token holds (batch_idx, token_idx_in_seq) int2. - size_t const tokens_info_size = sizeof(int2) * params.num_tokens; - size_t const fmha_scheduler_counter = mEnableContextFMHA ? sizeof(uint32_t) : 0; - size_t const fmha_bmm1_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) * 2 : 0; - size_t const fmha_bmm2_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) : 0; - - // cp workspace size upper bound - size_t const cpMaxPadedSequenceLength = params.num_tokens + params.batch_size * (mCpSize - 1); - size_t const cpWorkspaceSize - = mCpSize == 1 ? 0 : 2 * sizeof(T) * cpMaxPadedSequenceLength * getHeadSize() * (mNumHeads + 2 * mNumKVHeads); - - size_t const fmha_multi_ctas_kv_scratch_size = useTllmGenSparseAttention() ? getFmhaMultiCtasKvScratchSize() : 0; - - bool const is_qk_buf_float_ = true; - - AttentionContextWorkspaceSizes workspaceSizes{}; - workspaceSizes.attentionMask = attention_mask_size; - workspaceSizes.cuQSeqlens = cu_seqlens_size; - workspaceSizes.cuKvSeqlens = cu_seqlens_size; - workspaceSizes.cuMaskRows = cu_seqlens_size; - workspaceSizes.rotaryInvFreq = rotary_inv_freq_size; - workspaceSizes.qBuf = q_buf_2_size; - workspaceSizes.kBuf = k_buf_2_size; - workspaceSizes.vBuf = v_buf_2_size; - workspaceSizes.qkBuf = qk_buf_size; - workspaceSizes.qkvBuf = qkv_buf_2_size; - workspaceSizes.qkFloatBuf = qk_buf_float_size; - workspaceSizes.fp8QkvBuf = fp8_qkv_buffer_size; - workspaceSizes.fp8QBuf = fp8_q_buf_size; - workspaceSizes.fp8KBuf = fp8_k_buf_size; - workspaceSizes.fp8VBuf = fp8_v_buf_size; - workspaceSizes.paddingOffset = padding_offset_size; - workspaceSizes.encoderPaddingOffset = encoder_padding_offset_size; - workspaceSizes.tokensInfo = tokens_info_size; - workspaceSizes.fmhaTileCounter = fmha_scheduler_counter; - workspaceSizes.fmhaBmm1Scale = fmha_bmm1_scale_size; - workspaceSizes.fmhaBmm2Scale = fmha_bmm2_scale_size; - workspaceSizes.sageQScale = sage_q_sfs_buffer_size; - workspaceSizes.sageKScale = sage_k_sfs_buffer_size; - workspaceSizes.sageVScale = sage_v_sfs_buffer_size; - workspaceSizes.cpWorkspace = cpWorkspaceSize; - workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; - auto const workspaceLayout = AttentionWorkspaceManager::buildContextLayout(workspaceSizes); - auto const workspaceViews = AttentionWorkspaceManager::materializeContext( - params.workspace, workspaceLayout, cpMaxPadedSequenceLength, getHeadSize(), mNumHeads, mNumKVHeads); - - auto* fp8QBuf = workspaceViews.fp8QBuf; - // Fused FP8-Q path: caller pre-fills the nope segment of `quant_q_buf`; - // route the context-MLA Q pointer to it so the fused RoPE kernel appends - // rope FP8 in place and the FMHA Q load reads the merged [nope|rope] buffer. - if (mIsMLAEnabled && params.mla_param != nullptr && params.mla_param->fuse_q_fp8_in_rope - && params.mla_param->quant_q_buf != nullptr) - { - fp8QBuf = reinterpret_cast<__nv_fp8_e4m3*>(params.mla_param->quant_q_buf); - } - - // build attention mask, cu_seqlens, and padding offset tensors - // Note: self attn and cross attn should use different params - // cross attn's seqlen info is from encoder input lengths, not decoder input lengths! - // moreover, attn mask for cross attn should be set separately (see below) - BuildDecoderInfoParams decoder_params{}; - int32_t const* precomputedCuQSeqlens = params.cu_q_seqlens; - int32_t const* precomputedCuKvSeqlens = params.cu_kv_seqlens != nullptr ? params.cu_kv_seqlens - : params.cu_q_seqlens != nullptr ? params.cu_q_seqlens - : nullptr; - decoder_params.seqQOffsets = workspaceViews.cuQSeqlens; - decoder_params.seqKVOffsets = workspaceViews.cuKvSeqlens; - decoder_params.precomputedSeqQOffsets = precomputedCuQSeqlens; - decoder_params.precomputedSeqKVOffsets = precomputedCuKvSeqlens; - decoder_params.seqCpPartialOffsets = workspaceViews.cuCpPartialSeqlens; - decoder_params.cpSize = mCpSize; - decoder_params.packedMaskRowOffsets = workspaceViews.cuMaskRows; - decoder_params.paddingOffsets = workspaceViews.paddingOffset; - decoder_params.tokensInfo = workspaceViews.tokensInfo; - // Cross attention takes offsets from encoder inputs. - decoder_params.encoderPaddingOffsets = isCrossAttention() ? workspaceViews.encoderPaddingOffset : nullptr; - // Manually set attention mask for unfused cross attention. - decoder_params.attentionMask = isCrossAttention() ? nullptr : workspaceViews.attentionMask; - // Fixed sequence length offset if not removing the padding (seqQOffsets[i] = i * seq_length). - decoder_params.seqQLengths = params.context_lengths; - decoder_params.seqKVLengths = isCrossAttention() ? params.encoder_input_lengths : params.sequence_lengths; - decoder_params.batchSize = params.batch_size; - decoder_params.maxQSeqLength = params.input_seq_length; - decoder_params.maxEncoderQSeqLength - = isCrossAttention() ? params.cross_kv_length : 0; // cross attention uses encoder seq length - decoder_params.attentionWindowSize = params.cyclic_attention_window_size; - decoder_params.sinkTokenLength = params.sink_token_length; - decoder_params.numTokens = params.num_tokens; - decoder_params.removePadding = mRemovePadding; - decoder_params.attentionMaskType = mMaskType; - decoder_params.blockSparseParams = mBlockSparseParams; - decoder_params.fmhaTileCounter = workspaceViews.fmhaTileCounter; - decoder_params.quantScaleO = params.attention_output_orig_quant; - decoder_params.dequantScaleQkv = params.kv_scale_quant_orig; - decoder_params.separateQkvScales = mKVCacheQuantMode.hasFp4KvCache(); - decoder_params.fmhaHostBmm1Scale = 1.0f / (sqrtf(getHeadSize() * 1.0f) * q_scaling); - decoder_params.fmhaBmm1Scale = workspaceViews.fmhaBmm1Scale; - decoder_params.fmhaBmm2Scale = workspaceViews.fmhaBmm2Scale; - // Rotary embedding inv_freq buffer. - decoder_params.rotaryEmbeddingScale = mRotaryEmbeddingScale; - decoder_params.rotaryEmbeddingBase = mRotaryEmbeddingBase; - decoder_params.rotaryEmbeddingDim = mRotaryEmbeddingDim; - decoder_params.rotaryScalingType = mRotaryEmbeddingScaleType; - // The inv freq might be updated during runtime with dynamic scaling type. - decoder_params.rotaryEmbeddingInvFreq = workspaceViews.rotaryInvFreq; - // This is pre-computed when building the engines. - decoder_params.rotaryEmbeddingInvFreqCache = params.rotary_inv_freq; - decoder_params.rotaryEmbeddingMaxPositions = mRotaryEmbeddingMaxPositions; - - invokeBuildDecoderInfo(decoder_params, stream); - sync_check_cuda_error(stream); - - int32_t const* contextCuQSeqlens - = precomputedCuQSeqlens != nullptr ? precomputedCuQSeqlens : workspaceViews.cuQSeqlens; - int32_t const* contextCuKvSeqlens - = precomputedCuKvSeqlens != nullptr ? precomputedCuKvSeqlens : workspaceViews.cuKvSeqlens; - - // In cross attention context phase, the attention mask should be a matrix of all ones. - // Override the attention mask produced by invokeBuildDecoderInfo(). - // also, invokeBuildDecoderInfo can only handle square mask, not cross B x q_len x kv_len mask - // TODO: put this logic in the kernel above. currently not much concern because q_len is mostly = 1 - if (isUnfusedCrossAttention()) - { - std::vector h_attention_mask(params.batch_size * params.input_seq_length * params.cross_kv_length, 1.); - std::vector h_encoder_input_lengths(params.batch_size); - tensorrt_llm::common::cudaMemcpyAsyncSanitized(h_encoder_input_lengths.data(), params.encoder_input_lengths, - sizeof(int32_t) * params.batch_size, cudaMemcpyDeviceToHost, stream); - sync_check_cuda_error(stream); - - for (int bi = 0; bi < params.batch_size; bi++) - { - int b_offset = bi * params.input_seq_length * params.cross_kv_length; - for (int qi = 0; qi < params.input_seq_length; qi++) - { - int q_offset = b_offset + qi * params.cross_kv_length; - if (h_encoder_input_lengths[bi] < params.cross_kv_length) - { - std::fill(h_attention_mask.begin() + q_offset + h_encoder_input_lengths[bi], - h_attention_mask.begin() + q_offset + params.cross_kv_length, 0.f); - } - } - } - cudaMemcpyAsync(workspaceViews.attentionMask, h_attention_mask.data(), - sizeof(T) * params.batch_size * params.cross_kv_length * params.input_seq_length, cudaMemcpyHostToDevice, - stream); - sync_check_cuda_error(stream); - } - - // FIXME: a temporary solution to make sure the padding part is 0. - if (!mRemovePadding) - { - cudaMemsetAsync(params.context_buf, 0, params.num_tokens * local_hidden_units_qo * sizeof(T), stream); - sync_check_cuda_error(stream); - } - - KvCacheDataType cache_type = cacheTypeFromQuantMode(mKVCacheQuantMode); - - cudaDataType_t const gemm_data_type = tc::CudaDataType::value; - int const attention_seq_len_1 = params.input_seq_length; // q length - int const attention_seq_len_2 = isCrossAttention() ? params.cross_kv_length : params.input_seq_length; // kv length - - // If the model has relative attentiona bias, q scaling should be applied in QK gemm stage and use 1 in - // softamax stage (because to get softmax[scale(Q*K) + rel pos bias] here, q_scaling can't be applied during - // softmax phase by qk_scale); otherwise, use 1 in gemm stage and apply scaling in softmax stage - float const qk_scale - = 1.0f / (sqrtf(getHeadSize() * 1.0f) * q_scaling); // q_scaling in denominator. by default q_scaling =1.0f - float const qk_scale_gemm = isRelativePosition() ? qk_scale : 1.0f; - T const qk_scale_softmax = static_cast(isRelativePosition() ? 1.0f : qk_scale); - - // in context phase, currently FMHA runner has two restrictions: - // 1. only apply to self attention. If want fused multi-head cross attention, FMHCA kernels and runner is needed - // 2. doesn't apply to MHA with relative attention bias, i.e. softmax(QK + bias) * V - // We update mEnableContextFMHA in constructor to check these conditions - if (mEnableContextFMHA) - { - // do all-to-all for params.attention_input, need to split on kv head - // [token_num // cp_size, kv_heads, head_size] -> [token_num, kv_heads // cp_size, head_size] - T* attention_input = const_cast(params.attention_input); - if (mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1) - { - this->template ulyssesContextPreprocess(attention_input, workspaceViews.gatherInBuffer, - workspaceViews.gatherOutBuffer, params, contextCuQSeqlens, workspaceViews.cuCpPartialSeqlens, stream); - attention_input = workspaceViews.gatherInBuffer; - sync_check_cuda_error(stream); - } - - bool const enablePagedKVContextFMHA = mPagedKVCache && mPagedContextFMHA; - TLLM_CHECK_WITH_INFO(!(mKVCacheQuantMode.hasInt8KvCache() && enablePagedKVContextFMHA), - "Paged Context FMHA doesn't work with int8 kv cache currently."); - TLLM_CHECK_WITH_INFO(!(params.sink_token_length > 0 && enablePagedKVContextFMHA), - "Cannot support StreamingLLM now when enabling paged KV context FMHA."); - - // The max_kv_seq_len comes from the encoder seqlen when cross attention is used. - int const max_kv_seq_len = isCrossAttention() ? params.cross_kv_length : params.max_past_kv_length; - - // Prepare QKV preprocessing parameters. - QKVPreprocessingParams preprocessingParams; - - // Buffers. - preprocessingParams.qkv_input = const_cast(attention_input); - preprocessingParams.cross_kv_input = const_cast(params.cross_kv); - preprocessingParams.quantized_qkv_output = workspaceViews.fp8QkvBuf; - preprocessingParams.q_output = workspaceViews.qBuf; - preprocessingParams.kv_cache_buffer = kv_cache_buffer; - preprocessingParams.kv_cache_block_scales_buffer = kv_scale_cache_buffer; - preprocessingParams.qkv_bias = params.qkv_bias; - preprocessingParams.tokens_info = decoder_params.tokensInfo; - preprocessingParams.seq_lens = params.context_lengths; - // For self-attention, cache_seq_lens indicates whether chunked context is used - // (i.e. cache_seq_len > seq_len). - // For cross-attention, callers do not consistently use sequence_lengths as decoder length; use decoder - // context lengths so the encoder KV-cache write gate opens. - preprocessingParams.cache_seq_lens = isCrossAttention() ? params.context_lengths : params.sequence_lengths; - - preprocessingParams.encoder_seq_lens = params.encoder_input_lengths; - preprocessingParams.cu_seq_lens = contextCuQSeqlens; - // Cross-attention only. - preprocessingParams.cu_kv_seq_lens = contextCuKvSeqlens; - preprocessingParams.rotary_embedding_inv_freq = workspaceViews.rotaryInvFreq; - preprocessingParams.rotary_coef_cache_buffer = params.rotary_cos_sin; - preprocessingParams.mrope_rotary_cos_sin = params.mrope_rotary_cos_sin; - preprocessingParams.qkv_scale_orig_quant = params.kv_scale_orig_quant; - preprocessingParams.spec_decoding_position_offsets = nullptr; - preprocessingParams.helix_position_offsets = params.helix_position_offsets; - preprocessingParams.helix_is_inactive_rank = params.helix_is_inactive_rank; - preprocessingParams.logn_scaling = params.logn_scaling_ptr; - - // Sparse KV write - preprocessingParams.sparse_kv_indices = mRuntimeSparseAttentionParams.sparse_kv_indices; - preprocessingParams.sparse_kv_offsets = mRuntimeSparseAttentionParams.sparse_kv_offsets; - - // Scalars - preprocessingParams.batch_size = params.batch_size; - preprocessingParams.max_input_seq_len = params.input_seq_length; - preprocessingParams.max_kv_seq_len = max_kv_seq_len; - preprocessingParams.cyclic_kv_cache_len - = isCrossAttention() ? params.cross_kv_length : params.cyclic_attention_window_size; - preprocessingParams.sink_token_len = params.sink_token_length; - preprocessingParams.token_num = params.num_tokens; - preprocessingParams.remove_padding = mRemovePadding; - preprocessingParams.cross_attention = isCrossAttention(); - preprocessingParams.head_num = mNumAttnHeads; - preprocessingParams.kv_head_num = mNumAttnKVHeads; - preprocessingParams.qheads_per_kv_head = mNumAttnHeads / mNumAttnKVHeads; - preprocessingParams.size_per_head = getHeadSize(); - preprocessingParams.rotary_embedding_dim = mRotaryEmbeddingDim; - preprocessingParams.rotary_embedding_base = mRotaryEmbeddingBase; - preprocessingParams.rotary_scale_type = mRotaryEmbeddingScaleType; - preprocessingParams.rotary_embedding_scale = mRotaryEmbeddingScale; - preprocessingParams.rotary_embedding_max_positions = mRotaryEmbeddingMaxPositions; - preprocessingParams.position_embedding_type = position_embedding_type; - preprocessingParams.position_shift_enabled = mPosShiftEnabled; - preprocessingParams.cache_type = cache_type; - preprocessingParams.separate_q_kv_output = enablePagedKVContextFMHA || isCrossAttention(); - preprocessingParams.quantized_fp8_output = mFP8ContextFMHA; - preprocessingParams.generation_phase = false; - preprocessingParams.multi_processor_count = mMultiProcessorCount; - - preprocessingParams.rotary_vision_start = mVisionStart; - preprocessingParams.rotary_vision_length = mVisionLength; - preprocessingParams.is_last_chunk - = !mAttentionChunkSize.has_value() || (params.input_seq_length == params.max_past_kv_length); - - if (!(mIsMLAEnabled && params.mla_param != nullptr && params.mla_param->q_rope_applied)) - { - std::string const beforeRopeStr = "ctx attention before RoPE at layer " + std::to_string(mLayerIdx); - TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(params.num_tokens, - (local_hidden_units_qo + 2 * local_hidden_units_kv), mType, - const_cast(attention_input), stream, beforeRopeStr) - == false, - "Found invalid number (NaN or Inf) in " + beforeRopeStr); - } - - if (mIsMLAEnabled) - { - TLLM_CHECK_WITH_INFO(params.mla_param != nullptr, "MLA param is nullptr"); - params.mla_param->cache_type = cache_type; - params.mla_param->cu_q_seqlens = const_cast(contextCuQSeqlens); - params.mla_param->cu_kv_seqlens = const_cast(contextCuKvSeqlens); - params.mla_param->quant_scale_kv = params.kv_scale_orig_quant; - // Set BMM scales for FP8 context computation - params.mla_param->bmm1_scale = workspaceViews.fmhaBmm1Scale; - params.mla_param->bmm2_scale = workspaceViews.fmhaBmm2Scale; - params.mla_param->quant_q_buf = mFP8ContextMLA ? fp8QBuf : nullptr; - params.mla_param->quant_k_buf = mFP8ContextMLA ? workspaceViews.fp8KBuf : nullptr; - params.mla_param->quant_v_buf = mFP8ContextMLA ? workspaceViews.fp8VBuf : nullptr; - // Set additional scales for context phase - params.mla_param->quant_scale_o = params.attention_output_orig_quant; - params.mla_param->quant_scale_q = params.kv_scale_orig_quant; - params.mla_param->quant_scale_kv = params.kv_scale_orig_quant; - params.mla_param->dequant_scale_q = params.kv_scale_quant_orig; - params.mla_param->dequant_scale_kv - = cache_type == KvCacheDataType::NVFP4 ? nullptr : params.kv_scale_quant_orig; - params.mla_param->host_bmm1_scale - = 1 / (mQScaling * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); - // The sparse MLA is in the absorption mode for the context phase. - params.mla_param->absorption_mode = useSparseMLA(); - // Fused FP8-Q-quant: RoPE kernel writes FP8 rope into `quant_q_buf`, - // so we skip the standalone invokeMLAContextFp8Quantize call below. - bool const useFusedQFp8 = params.mla_param->fuse_q_fp8_in_rope && mFP8ContextMLA - && params.mla_param->absorption_mode && cache_type == KvCacheDataType::FP8 - && params.mla_param->quant_q_buf != nullptr && params.mla_param->quant_scale_qkv != nullptr; - TLLM_CHECK_WITH_INFO(cache_type != KvCacheDataType::NVFP4 || params.mla_param->latent_cache == nullptr, - "NVFP4 sparse MLA context must append its latent cache before launching attention"); - if (params.mla_param->latent_cache != nullptr) - { - invokeMLARopeContext(*params.mla_param, kv_cache_buffer, stream); - } - if (mFP8ContextMLA && !useFusedQFp8) - { - invokeMLAContextFp8Quantize(*params.mla_param, params.total_kv_len, stream); - } - } - else if (useSageAttnSeparateQkv) - { - TLLM_CHECK_WITH_INFO(mFP8ContextFMHA, "SageAttention kernel runs under mFP8ContextFMHA option."); - TLLM_CHECK_WITH_INFO(mFmhaDispatcher->isSupported(), "SageAttention has no unfused fallback implemented."); - TLLM_CHECK_WITH_INFO(mMaskType == AttentionMaskType::PADDING, - "SageAttention only supports dense (padding) mask, got mask type %d.", static_cast(mMaskType)); - TLLM_CHECK_WITH_INFO( - mSageAttnNumEltsPerBlkQ > 0 && mSageAttnNumEltsPerBlkK > 0 && mSageAttnNumEltsPerBlkV == 1, - "SageQuant requires positive block sizes for Q and K while the block size for V must be 1."); - TLLM_CHECK_WITH_INFO(!params.kv_scale_quant_orig, - "SageAttention disregards the configured params.kv_scale_quant_orig, invalidating the result."); - check_cuda_error(cudaMemsetAsync(workspaceViews.sageVScale, 0, sage_v_sfs_buffer_size, stream)); - - // Common params for sageQuant - tc::SageQuantParams sageQuantParams{}; - sageQuantParams.headDim = getHeadSize(); - sageQuantParams.inputType = std::is_same_v ? DATA_TYPE_BF16 : DATA_TYPE_FP16; - sageQuantParams.quantType = mSageAttnQkInt8 ? DATA_TYPE_INT8 : DATA_TYPE_E4M3; - sageQuantParams.vStage = 0; - sageQuantParams.sumSeqLensV = params.total_kv_len; - sageQuantParams.numHeadsV = mNumAttnKVHeads; - sageQuantParams.ptrV = params.v_ptr; - sageQuantParams.ptrVQuant = workspaceViews.fp8VBuf; - sageQuantParams.ptrVScale = workspaceViews.sageVScale; - sageQuantParams.smCount = mMultiProcessorCount; - sageQuantParams.stream = stream; - - // Quantize into Fp8Q, SfsQ, SfsV - sageQuantParams.sumSeqLensQk = params.num_tokens; - sageQuantParams.batchSize = params.batch_size; - sageQuantParams.numHeads = mNumAttnHeads; - sageQuantParams.tokenBlockSize = mSageAttnNumEltsPerBlkQ; - sageQuantParams.ptrCuSeqLensQk = contextCuQSeqlens; - sageQuantParams.ptrQk = attention_input; - sageQuantParams.ptrQkQuant = workspaceViews.fp8QBuf; - sageQuantParams.ptrQkScale = workspaceViews.sageQScale; - sageQuantParams.vStage = 1; - tc::invokeSageQuant(sageQuantParams); - - // Quantize into Fp8K, SfsK, Fp8V - sageQuantParams.sumSeqLensQk = params.total_kv_len; - sageQuantParams.batchSize = params.batch_size; - sageQuantParams.numHeads = mNumAttnKVHeads; - sageQuantParams.tokenBlockSize = mSageAttnNumEltsPerBlkK; - sageQuantParams.ptrCuSeqLensQk = contextCuKvSeqlens; - sageQuantParams.ptrQk = params.k_ptr; - sageQuantParams.ptrQkQuant = workspaceViews.fp8KBuf; - sageQuantParams.ptrQkScale = workspaceViews.sageKScale; - sageQuantParams.vStage = 2; - tc::invokeSageQuant(sageQuantParams); - } - else - { - invokeQKVPreprocessing(preprocessingParams, stream); - } - sync_check_cuda_error(stream); - if (!(mIsMLAEnabled && params.mla_param != nullptr && params.mla_param->q_rope_applied)) - { - std::string const afterRopeStr = "ctx attention after RoPE at layer " + std::to_string(mLayerIdx); - TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(params.num_tokens, - (local_hidden_units_qo + 2 * local_hidden_units_kv), mType, - const_cast(attention_input), stream, afterRopeStr) - == false, - "Found invalid number (NaN or Inf) in " + afterRopeStr); - sync_check_cuda_error(stream); - } - - if (params.runtime_perf_knobs) - { - int64_t enable_context_fmha_fp32_acc_val = params.runtime_perf_knobs[1]; - mFMHAForceFP32Acc = mFMHAForceFP32Acc || enable_context_fmha_fp32_acc_val == 1; - } - - // Unified FMHA runner interface for both packed QKV FMHA, contiguous Q_KV, paged KV FMHA, and separate QKV - // FMHA. - // Page KV input layout: - // - q_ptr: [B, S, H, D], which supports variable sequence length - // - paged_kv_cache: paged kv buffer - // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. - // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. - // - // Contiguous KV input layout: - // - q_ptr: [B, S, H, D], which supports variable sequence length - // - kv_ptr: [B, S, 2, H, D], which supports variable sequence length - // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. - // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. - // - // Separate QKV input layout (only for context MLA now): - // - q_ptr: [B, S, H, D], which supports variable sequence length - // - k_ptr: [B, S, H_kv, D], which supports variable sequence length - // - v_ptr: [B, S, H_kv, D_v], which supports variable sequence length - // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. - // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. - // - total_kv_len: the total kv sequence length, needed for variable sequence length. - - // Construct the fmha params for running kernels. - MHARunnerParams fmhaParams{}; - fmhaParams.b = params.batch_size; - fmhaParams.qSeqLen = params.input_seq_length; - fmhaParams.kvSeqLen = max_kv_seq_len; - // Disable sliding window attention when it is not needed. - fmhaParams.slidingWindowSize - = (mDenseContextFMHA || isCrossAttention()) ? max_kv_seq_len : params.cyclic_attention_window_size; - fmhaParams.totalQSeqLen = params.num_tokens; - // TODO: set it correctly for contiguous kv buffer (cross-attention). - fmhaParams.totalKvSeqLen = isCrossAttention() ? params.num_encoder_tokens : params.total_kv_len; - // Device buffer pointers. - if (mIsMLAEnabled) - { - // separate QKV input for context MLA - if (mFP8ContextMLA) - { - TLLM_CHECK_WITH_INFO( - mFmhaDispatcher->isSeparateQAndKvInput(), "Separate QKV input is required for fp8 context MLA"); - TLLM_CHECK_WITH_INFO(fp8QBuf != nullptr, "FP8 q buffer is required for fp8 context MLA"); - // In sparse MLA (absorption mode), K and V are stored in KV cache, not as separate FP8 buffers - TLLM_CHECK_WITH_INFO(useSparseMLA() || workspaceViews.fp8KBuf != nullptr, - "FP8 k buffer is required for fp8 context MLA in non-sparse mode"); - TLLM_CHECK_WITH_INFO(useSparseMLA() || workspaceViews.fp8VBuf != nullptr, - "FP8 v buffer is required for fp8 context MLA in non-sparse mode"); - - fmhaParams.qPtr = reinterpret_cast(fp8QBuf); - fmhaParams.kPtr = useSparseMLA() ? nullptr : reinterpret_cast(workspaceViews.fp8KBuf); - fmhaParams.vPtr = useSparseMLA() ? nullptr : reinterpret_cast(workspaceViews.fp8VBuf); - } - else - { - fmhaParams.qPtr = attention_input; - fmhaParams.kPtr = params.k_ptr; - fmhaParams.vPtr = params.v_ptr; - } - } - else if (useSageAttnSeparateQkv) - { - // SageAttention: use quantized FP8/INT8 Q/K/V buffers as separate inputs. - TLLM_CHECK_WITH_INFO( - mFmhaDispatcher->isSeparateQAndKvInput(), "Separate QKV input is required for sage attention FMHA"); - fmhaParams.qkvPtr = nullptr; - fmhaParams.qPtr = reinterpret_cast(workspaceViews.fp8QBuf); - fmhaParams.kPtr = reinterpret_cast(workspaceViews.fp8KBuf); - fmhaParams.vPtr = reinterpret_cast(workspaceViews.fp8VBuf); - // Set sage attention scaling factor pointers. - fmhaParams.qScalePtr = workspaceViews.sageQScale; - fmhaParams.kScalePtr = workspaceViews.sageKScale; - fmhaParams.vScalePtr = workspaceViews.sageVScale; - } - else - { - fmhaParams.qkvPtr = mFP8ContextFMHA ? reinterpret_cast(workspaceViews.fp8QkvBuf) - : reinterpret_cast(attention_input); - fmhaParams.qPtr = reinterpret_cast(workspaceViews.qBuf); - } - // TODO: add contiguous kv buffer (cross-attention). - fmhaParams.kvPtr = nullptr; - if (isCrossAttention() && !useKVCache()) - { - fmhaParams.kvPtr = params.cross_kv; - } - // Only use [totalLength, h / cpSize, Dh]. - fmhaParams.outputPtr = mCpSize > 1 ? workspaceViews.gatherOutBuffer : params.context_buf; - fmhaParams.outputSfPtr = params.context_buf_sf; - if (params.mla_param != nullptr && params.mla_param->dsv4_epilogue_fusion.enabled) - { - fmhaParams.dsv4EpilogueFusion.enabled = true; - fmhaParams.dsv4EpilogueFusion.cosSinCache = params.mla_param->dsv4_epilogue_fusion.cos_sin_cache; - fmhaParams.dsv4EpilogueFusion.scaleBufM = params.mla_param->dsv4_epilogue_fusion.scale_buf_m; - } - fmhaParams.attentionSinksPtr = params.attention_sinks; - fmhaParams.packedMaskPtr = params.attention_packed_mask; - if constexpr (std::is_same_v) - { - fmhaParams.pagedKvCache = kv_cache_buffer; - fmhaParams.pagedKvSfCache = kv_scale_cache_buffer; - } - fmhaParams.cuQSeqLenPtr = contextCuQSeqlens; - fmhaParams.kvSeqLenPtr = decoder_params.seqKVLengths; - fmhaParams.cuKvSeqLenPtr = contextCuKvSeqlens; - fmhaParams.cuMaskRowsPtr = workspaceViews.cuMaskRows; - fmhaParams.tileCounterPtr = workspaceViews.fmhaTileCounter; - fmhaParams.scaleBmm1Ptr = workspaceViews.fmhaBmm1Scale; - fmhaParams.scaleBmm2Ptr = workspaceViews.fmhaBmm2Scale; - fmhaParams.oSfScalePtr = params.attention_output_sf_scale; - fmhaParams.stream = stream; - fmhaParams.forceFp32Acc = mFMHAForceFP32Acc; - fmhaParams.softmaxStatsPtr = params.softmax_stats; - fmhaParams.trtllmGenJITWarmup = params.trtllm_gen_jit_warmup; - fmhaParams.trtllmGenJITWarmupMaxNumRequests = mMaxNumRequests; - fmhaParams.trtllmGenJITWarmupMaxSeqLenQ = mMaxContextLength; - fmhaParams.trtllmGenJITWarmupMaxSeqLenKv = mMaxSeqLen; - - // Sparse attention parameters - if (useTllmGenSparseAttention()) - { - fmhaParams.sparse_params = mRuntimeSparseAttentionParams; - // Sparse context reuses generation-style trtllm-gen kernels; provide the scratch pool - // and per-CTA counter so the autotuner can select MultiCtasKv variants. - fmhaParams.multiCtasKvScratchPtr = workspaceViews.fmhaMultiCtasKvScratch; - fmhaParams.multiCtasKvCounterPtr = multiBlockSemaphores(); - } - - // Skip-softmax attention parameters - fmhaParams.skipSoftmaxThresholdScaleFactor = mSkipSoftmaxThresholdScaleFactorPrefill; - fmhaParams.skipCorrectionThreshold = mSkipCorrectionThreshold; -#ifdef SKIP_SOFTMAX_STAT - fmhaParams.skipSoftmaxTotalBlocks = mSkipSoftmaxTotalBlocks; - fmhaParams.skipSoftmaxSkippedBlocks = mSkipSoftmaxSkippedBlocks; -#else - if (tensorrt_llm::common::getEnvPrintSkipSoftmaxStat()) - { - TLLM_THROW("To print skip softmax stat, please run build_wheel.py with -DSKIP_SOFTMAX_STAT"); - } -#endif - - if (mAttentionChunkSize) - { - fmhaParams.chunkedAttentionSize = *mAttentionChunkSize; - } - - // Run the fmha kernel. - mFmhaDispatcher->run(fmhaParams); - sync_check_cuda_error(stream); - - if (mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1) - { - this->template ulyssesContextPostprocess(workspaceViews.gatherOutBuffer, - reinterpret_cast(params.context_buf), workspaceViews.gatherInBuffer, params, contextCuQSeqlens, - workspaceViews.cuCpPartialSeqlens, stream); - sync_check_cuda_error(stream); - } - - if (!mIsMLAEnabled) // Only for non-MLA attention - { - invokeKvCachePostprocessing(preprocessingParams, stream); - sync_check_cuda_error(stream); - } - } - else - { - TLLM_CHECK_DEBUG_WITH_INFO(params.logn_scaling_ptr == nullptr, "Unfused MHA does not support logn scaling"); - TLLM_CHECK_WITH_INFO(mAttentionChunkSize == std::nullopt, "Unfused MHA does not support chunked attention"); - // FIXME: a temporary solution to make sure the padding part of key/value buffer is 0 - // NOTE: pointer subtraction is used below since there could be some extra gap due to alignment. - // Otherwise, we could do cudaMemsetAsync(workspaceViews.kBuf, 0, k_buf_2_size + v_buf_2_size, stream). - // cudaMemsetAsync(workspaceViews.kBuf, 0, - // reinterpret_cast(workspaceViews.qkBuf) - reinterpret_cast(workspaceViews.kBuf), - // stream); - cudaMemsetAsync(workspaceViews.kBuf, 0, - reinterpret_cast(workspaceViews.vBuf) - reinterpret_cast(workspaceViews.kBuf) - + v_buf_2_size, - stream); - - if (!isCrossAttention()) - { - // self attention, write to from QKV to Q/K/V - invokeAddFusedQKVBiasTranspose(workspaceViews.qBuf, workspaceViews.kBuf, workspaceViews.vBuf, - const_cast(params.attention_input), const_cast(params.qkv_bias), params.context_lengths, - mRemovePadding ? workspaceViews.paddingOffset : nullptr, params.batch_size, params.input_seq_length, - params.num_tokens, mNumHeads, mNumKVHeads, getHeadSize(), mRotaryEmbeddingDim, mRotaryEmbeddingBase, - mRotaryEmbeddingScaleType, mRotaryEmbeddingScale, mRotaryEmbeddingMaxPositions, position_embedding_type, - (float*) nullptr, 0, stream); - sync_check_cuda_error(stream); - } - else - { - // cross attention, write from self QKV [*, head_num * head_size + 2 * kv_head_num * head_size]to Q, write - // from cross KV [*, 2 * kv_head_num * head_size] to K/V kernel modified accordingly to handle nullptr - // buffer - invokeAddFusedQKVBiasTranspose(workspaceViews.qBuf, (T*) nullptr, (T*) nullptr, - const_cast(params.attention_input), const_cast(params.qkv_bias), params.context_lengths, - mRemovePadding ? workspaceViews.paddingOffset : nullptr, params.batch_size, params.input_seq_length, - params.num_tokens, mNumHeads, mNumKVHeads, getHeadSize(), mRotaryEmbeddingDim, mRotaryEmbeddingBase, - mRotaryEmbeddingScaleType, mRotaryEmbeddingScale, mRotaryEmbeddingMaxPositions, position_embedding_type, - (float*) nullptr, 0, stream); - sync_check_cuda_error(stream); - - invokeAddFusedQKVBiasTranspose((T*) nullptr, workspaceViews.kBuf, workspaceViews.vBuf, - const_cast(params.cross_kv), const_cast(params.qkv_bias), params.encoder_input_lengths, - mRemovePadding ? workspaceViews.encoderPaddingOffset : nullptr, params.batch_size, - params.cross_kv_length, params.num_encoder_tokens, /*mNumHeads*/ 0, mNumKVHeads, getHeadSize(), - mRotaryEmbeddingDim, mRotaryEmbeddingBase, mRotaryEmbeddingScaleType, mRotaryEmbeddingScale, - mRotaryEmbeddingMaxPositions, position_embedding_type, (float*) nullptr, 0, stream); - sync_check_cuda_error(stream); - } - - // write KV to cache - if (useKVCache()) - { - invokeTranspose4dBatchMajor(workspaceViews.kBuf, workspaceViews.vBuf, kv_cache_buffer, params.batch_size, - isCrossAttention() ? params.cross_kv_length : params.input_seq_length, - isCrossAttention() ? params.cross_kv_length : params.cyclic_attention_window_size, getHeadSize(), - mNumKVHeads, cache_type, params.kv_scale_orig_quant, - isCrossAttention() ? params.encoder_input_lengths : params.context_lengths, stream); - } - sync_check_cuda_error(stream); - - T const* linear_bias_slopes = isALiBi() ? params.alibi_slopes : nullptr; - T const* relative_attention_bias = isRelativePosition() ? params.relative_attention_bias : nullptr; - int const relative_attention_bias_stride = isRelativePosition() ? params.relative_attention_bias_stride : 0; - int const max_distance = mMaxDistance; - cudaDataType_t gemm_out_data_type = is_qk_buf_float_ ? CUDA_R_32F : gemm_data_type; - void* gemm_out_buf_ = is_qk_buf_float_ ? static_cast(workspaceViews.qkFloatBuf) - : static_cast(workspaceViews.qkBuf); - if (mNumKVHeads == 1) // MQA - { - // Attn_weight[b, h*s_q, s_k] = Q[b, h*s_q, d] * K'[b, d, s_k] - // Attn_weight'[b, s_k, h*s_q] = K[b, s_k, d] * Q'[b, d, h*s_q] - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, - attention_seq_len_2, // n - attention_seq_len_1 * mNumHeads, // m - getHeadSize(), // k - qk_scale_gemm, workspaceViews.kBuf, gemm_data_type, - getHeadSize(), // k - attention_seq_len_2 * getHeadSize(), // n * k - workspaceViews.qBuf, gemm_data_type, - getHeadSize(), // k - attention_seq_len_1 * mNumHeads * getHeadSize(), // m * k - 0.0f, gemm_out_buf_, gemm_out_data_type, - attention_seq_len_2, // n - attention_seq_len_1 * mNumHeads * attention_seq_len_2, // m * n - params.batch_size, // global batch size - CUDA_R_32F); - } - else if (mNumKVHeads == mNumHeads) // MHA - { - // Attn_weight[b*h, s_q, s_k] = Q[b*h, s_q, d] * K'[b*h, d, s_k] - // Attn_weight'[b*h, s_k, s_q] = K[b*h, s_k, d] * Q'[b*h, d, s_q] - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, - attention_seq_len_2, // n - attention_seq_len_1, // m - getHeadSize(), // k - qk_scale_gemm, workspaceViews.kBuf, gemm_data_type, - getHeadSize(), // k - attention_seq_len_2 * getHeadSize(), // n * k - workspaceViews.qBuf, gemm_data_type, - getHeadSize(), // k - attention_seq_len_1 * getHeadSize(), // m * k - 0.0f, gemm_out_buf_, gemm_out_data_type, - attention_seq_len_2, // n - attention_seq_len_2 * attention_seq_len_1, - params.batch_size * mNumHeads, // global batch size - CUDA_R_32F); - } - else // GQA - { - // Some number of contiguous Q heads will share the same K/V head - // Since the KV stride is NOT fixed for all Q, we have 2 options: - // 1. Loop over stridedBatchedGemm for each KV head. (multiple API calls/cuda kernels) - // 2. Calculate the pointers and use batchedGemm() (extra device memory) ::TODO:: - int const num_qheads_per_kv_head = mNumHeads / mNumKVHeads; - for (int ki = 0; ki < mNumKVHeads; ++ki) - { - T* qptr = workspaceViews.qBuf + (ki * num_qheads_per_kv_head * attention_seq_len_1 * getHeadSize()); - T* kptr = workspaceViews.kBuf + (ki * attention_seq_len_2 * getHeadSize()); - int const qk_offset = ki * attention_seq_len_1 * num_qheads_per_kv_head * attention_seq_len_2; - void* qkptr = is_qk_buf_float_ ? static_cast(workspaceViews.qkFloatBuf + qk_offset) - : static_cast(workspaceViews.qkBuf + qk_offset); - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, - attention_seq_len_2, // n - attention_seq_len_1 * num_qheads_per_kv_head, // m - getHeadSize(), // k - qk_scale_gemm, kptr, gemm_data_type, - getHeadSize(), // k - mNumKVHeads * attention_seq_len_2 * getHeadSize(), // n * k - qptr, gemm_data_type, - getHeadSize(), // k - attention_seq_len_1 * mNumHeads * getHeadSize(), // m * k - 0.0f, qkptr, gemm_out_data_type, - attention_seq_len_2, // n - attention_seq_len_1 * mNumHeads * attention_seq_len_2, // m * n - params.batch_size, // global batch size - CUDA_R_32F); - } - } - - if (is_qk_buf_float_ == true) - { - // add relative position bias - if (isRelativePosition()) - { - // Add relative_attention_bias - // QK is (batch_size, local_head_num, q_length, k_length), relative_attention_bias is (1, - // local_head_num, max_output_len + 1, max_output_len + 1). broadcast along 1st dim. max_seq_len is - // already max_output_len + 1. In implicit mode, relative_attention_bias is relative_attention_table - // [num_heads, num_buckets], with necessary params (max_distance, num_buckets) passed at the end - invokeAddRelativeAttentionBiasUnaligned(workspaceViews.qkFloatBuf, relative_attention_bias, - params.batch_size, mNumHeads, attention_seq_len_1, - isCrossAttention() ? params.cross_kv_length : params.cyclic_attention_window_size, stream, - max_distance > 0, relative_attention_bias_stride, max_distance, false /* bidirectional */); - } - - MaskedSoftmaxParam param; - param.attention_score = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) - param.qk = workspaceViews.qkFloatBuf; // (batch_size, head_num, q_length, k_length) - param.attention_mask = workspaceViews.attentionMask; // (batch_size, q_length, k_length) - param.batch_size = params.batch_size; - param.q_length = attention_seq_len_1; - param.k_length = attention_seq_len_2; - param.num_heads = mNumHeads; - param.qk_scale = qk_scale_softmax; - param.attn_logit_softcapping_scale = mAttnLogitSoftcappingScale; - param.attn_logit_softcapping_inverse_scale = 1.0f / mAttnLogitSoftcappingScale; - param.linear_bias_slopes = const_cast(linear_bias_slopes); // (head_num,), optional - param.block_sparse_attn = mMaskType == AttentionMaskType::BLOCKSPARSE; - param.block_sparse_params = mBlockSparseParams; - param.q_seq_lengths = params.context_lengths; - invokeMaskedSoftmax(param, stream); - } - else - { - // add relative position bias - if (isRelativePosition()) - { - // Add relative_attention_bias - // QK is (batch_size, local_head_num, q_length, k_length), relative_attention_bias is (1, - // local_head_num, max_output_len + 1, max_output_len + 1). broadcast along 1st dim. max_seq_len is - // already max_output_len + 1. In implicit mode, relative_attention_bias is relative_attention_table - // [num_heads, num_buckets], with necessary params (max_distance, num_buckets) passed at the end - invokeAddRelativeAttentionBiasUnaligned(workspaceViews.qkBuf, relative_attention_bias, - params.batch_size, mNumHeads, attention_seq_len_1, - isCrossAttention() ? params.cross_kv_length : params.cyclic_attention_window_size, stream, - max_distance > 0, relative_attention_bias_stride, max_distance, false /* bidirectional */); - } - - MaskedSoftmaxParam param; - param.attention_score = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) - param.qk = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) - param.attention_mask = workspaceViews.attentionMask; // (batch_size, q_length, k_length) - param.batch_size = params.batch_size; - param.q_length = attention_seq_len_1; - param.k_length = attention_seq_len_2; - param.num_heads = mNumHeads; - param.qk_scale = qk_scale_softmax; - param.attn_logit_softcapping_scale = mAttnLogitSoftcappingScale; - param.attn_logit_softcapping_inverse_scale = 1.0f / mAttnLogitSoftcappingScale; - param.linear_bias_slopes = const_cast(linear_bias_slopes); // (head_num,), optional - param.block_sparse_attn = mMaskType == AttentionMaskType::BLOCKSPARSE; - param.block_sparse_params = mBlockSparseParams; - param.q_seq_lengths = params.context_lengths; - invokeMaskedSoftmax(param, stream); - } - - if (mNumKVHeads == 1) - { - // Attn_weight[b, h*s_q, s_k] - // O[b, h*s_q, d] = Attn_weight[b, h*s_q, s_k] * V[b, s_k, d] - // O'[b, d, h*s_q] = V'[b, d, s_k] * Attn_weight'[b, s_k, h*s_q] - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, - getHeadSize(), // n - mNumHeads * attention_seq_len_1, // m - attention_seq_len_2, // k - workspaceViews.vBuf, - getHeadSize(), // n - getHeadSize() * attention_seq_len_2, // n * k - workspaceViews.qkBuf, - attention_seq_len_2, // k - attention_seq_len_2 * mNumHeads * attention_seq_len_1, // m * k - workspaceViews.qkvBuf, - getHeadSize(), // n - getHeadSize() * mNumHeads * attention_seq_len_1, // n * m - params.batch_size // global batch size - ); - } - else if (mNumKVHeads == mNumHeads) // MHA - { - // O[b*h, s_q, d] = Attn_weight[b*h, s_q, s_k] * V[b*h, s_k, d] - // O'[b*h, d, s_q] = V'[b*h, d, s_k] * Attn_weight'[b*h, s_k, s_q] - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, getHeadSize(), attention_seq_len_1, - attention_seq_len_2, workspaceViews.vBuf, getHeadSize(), attention_seq_len_2 * getHeadSize(), - workspaceViews.qkBuf, attention_seq_len_2, attention_seq_len_1 * attention_seq_len_2, - workspaceViews.qkvBuf, getHeadSize(), attention_seq_len_1 * getHeadSize(), - params.batch_size * mNumHeads); - } - else // GQA - { - // Attn_weight[b, h*s_q, s_k] - // O[b, h*s_q, d] = Attn_weight[b, h*s_q, s_k] * V[b, s_k, d] - // O'[b, d, h*s_q] = V'[b, d, s_k] * Attn_weight'[b, s_k, h*s_q] - int const num_qheads_per_kv_head = mNumHeads / mNumKVHeads; - for (int ki = 0; ki < mNumKVHeads; ++ki) - { - T* qkptr - = workspaceViews.qkBuf + (ki * num_qheads_per_kv_head * attention_seq_len_1 * attention_seq_len_2); - T* vptr = workspaceViews.vBuf + (ki * attention_seq_len_2 * getHeadSize()); - T* qkvptr = workspaceViews.qkvBuf + (ki * attention_seq_len_1 * num_qheads_per_kv_head * getHeadSize()); - mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, - getHeadSize(), // n - num_qheads_per_kv_head * attention_seq_len_1, // m - attention_seq_len_2, // k - vptr, - getHeadSize(), // n - mNumKVHeads * getHeadSize() * attention_seq_len_2, // n * k - qkptr, - attention_seq_len_2, // k - attention_seq_len_2 * mNumHeads * attention_seq_len_1, // m * k - qkvptr, - getHeadSize(), // n - getHeadSize() * mNumHeads * attention_seq_len_1, // n * m - params.batch_size // global batch size - ); - } - } - - if (!mRemovePadding) - { - invokeTransposeQKV(static_cast(params.context_buf), workspaceViews.qkvBuf, params.batch_size, - attention_seq_len_1, mNumHeads, getHeadSize(), (float*) nullptr, 0, stream); - } - else - { - invokeTransposeAttentionOutRemovePadding(workspaceViews.qkvBuf, static_cast(params.context_buf), - params.num_tokens, params.batch_size, attention_seq_len_1, mNumHeads, getHeadSize(), - workspaceViews.paddingOffset, (float*) nullptr, 0, stream); - } - } - return 0; -} - -template int AttentionOp::enqueueContext( - EnqueueContextParams const& params, cudaStream_t stream); - -template int AttentionOp::enqueueContext( - EnqueueContextParams const& params, cudaStream_t stream); - -#ifdef ENABLE_BF16 -template int AttentionOp::enqueueContext<__nv_bfloat16, KVLinearBuffer>( - EnqueueContextParams<__nv_bfloat16> const& params, cudaStream_t stream); -#endif - -template int AttentionOp::enqueueContext( - EnqueueContextParams const& params, cudaStream_t stream); - -template int AttentionOp::enqueueContext( - EnqueueContextParams const& params, cudaStream_t stream); - -#ifdef ENABLE_BF16 -template int AttentionOp::enqueueContext<__nv_bfloat16, KVBlockArray>( - EnqueueContextParams<__nv_bfloat16> const& params, cudaStream_t stream); -#endif - -template -int AttentionOp::enqueueGeneration(EnqueueGenerationParams const& params, cudaStream_t stream) -{ - int const headSize = getHeadSize(); - float const q_scaling = mQScaling; - float const* logn_scaling_ptr = isLognScaling() ? params.logn_scaling_ptr : nullptr; - T const* relative_attention_bias = isRelativePosition() ? params.relative_attention_bias : nullptr; - int const relative_attention_bias_stride = isRelativePosition() ? params.relative_attention_bias_stride : 0; - int const max_distance = mMaxDistance; - bool const* finished = nullptr; - - auto const quant_option = tc::QuantMode{}; - float const* qkv_scale_out = nullptr; - - int const* ia3_tasks = nullptr; - T const* ia3_key_weights = nullptr; - T const* ia3_value_weights = nullptr; - - int32_t const batch_beam = params.beam_width * params.num_requests; - - KVCacheBuffer kv_cache_buffer; - KVCacheBuffer kv_scale_cache_buffer; - - auto const sizePerToken = mNumAttnKVHeads * headSize * getKvCacheElemSizeInBits() / 8 /*bits*/; - - if (useKVCache()) - { - auto buffers = buildKvCacheBuffers(batch_beam, params.max_blocks_per_sequence, mTokensPerBlock, - sizePerToken, params.cyclic_attention_window_size, params.max_cyclic_attention_window_size, - params.sink_token_length, params.can_use_one_more_block, params.host_primary_pool_pointer, - params.host_secondary_pool_pointer, params.host_primary_block_scale_pool_pointer, - params.host_secondary_block_scale_pool_pointer, params.block_offsets, mKVCacheQuantMode.hasFp4KvCache(), - params.max_attention_window_size, params.key_value_cache); - kv_cache_buffer = buffers.kvCacheBuffer; - kv_scale_cache_buffer = buffers.kvScaleCacheBuffer; - } - sync_check_cuda_error(stream); - -#ifndef NDEBUG - debugCheckSemaphores(stream); -#endif - - if (params.runtime_perf_knobs) - { - int64_t multi_block_mode_val = params.runtime_perf_knobs[0]; - mMultiBlockMode = multi_block_mode_val == 1; - if (common::getEnvForceDeterministicAttention()) - { - mMultiBlockMode = false; - } - } - - if (common::getEnvForceDeterministicAttention()) - { - mMultiBlockMode = false; - } - - // TODO only for debug usage - if (!mMultiBlockMode) - { - char* isForceMultiBlockModeChar = std::getenv("FORCE_MULTI_BLOCK_MODE"); - bool isForceMultiBlockMode - = (isForceMultiBlockModeChar != nullptr && std::string(isForceMultiBlockModeChar) == "ON"); - TLLM_CHECK_WITH_INFO(!(common::getEnvForceDeterministicAttention() && isForceMultiBlockMode), - "FORCE_MULTI_BLOCK_MODE and FORCE_DETERMINISTIC/FORCE_ATTENTION_KERNEL_DETERMINISTIC can not be set at " - "the same time."); - mMultiBlockMode = isForceMultiBlockMode; - } - - // Check that the chunked-attention and sliding-window-attention are not enabled at the same time. - TLLM_CHECK_WITH_INFO( - !mAttentionChunkSize.has_value() || params.cyclic_attention_window_size >= params.max_past_kv_length, - "Chunked-attention and sliding-window-attention should not be enabled at the same time."); - - size_t const cpMaxPaddedSequenceLength = (batch_beam + mCpSize - 1) / mCpSize * mCpSize; - size_t const cpWorkspaceSize - = mCpSize == 1 ? 0 : 2 * sizeof(T) * cpMaxPaddedSequenceLength * (mNumHeads + 2 * mNumKVHeads) * mHeadSize; - AttentionGenerationWorkspaceSizes cpWorkspaceSizes{}; - cpWorkspaceSizes.cpWorkspace = cpWorkspaceSize; - auto const cpWorkspaceLayout = AttentionWorkspaceManager::buildGenerationLayout(cpWorkspaceSizes); - auto const cpWorkspaceViews = AttentionWorkspaceManager::materializeGeneration( - params.workspace, cpWorkspaceLayout, cpMaxPaddedSequenceLength, mNumHeads, mNumKVHeads, mHeadSize); - - T* attention_input = const_cast(params.attention_input); - if (mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1) - { - this->template ulyssesGenerationPreprocess( - attention_input, cpWorkspaceViews.mhaInput, cpWorkspaceViews.mhaOutput, batch_beam, stream); - attention_input = cpWorkspaceViews.mhaInput; - sync_check_cuda_error(stream); - } - - // Try XQA optimization first. - { - // NOTE: input_seq_length = num_medusa_tokens + 1 (new generated one from the original LM head) - // self attn - XQAParams xqaParams{}; - this->template convertMMHAParamsToXQAParams(xqaParams, params, /*forConfigurePlugin=*/false); - - if (mEnableXQA && mXqaDispatcher->shouldUse(xqaParams)) - { - TLLM_LOG_DEBUG("XQA kernels are selected in the generation phase."); - xqaParams.stream = stream; - if (mCpSize > 1) - { - xqaParams.output = cpWorkspaceViews.mhaOutput; - xqaParams.qkv = attention_input; - } - { - mXqaDispatcher->run(xqaParams, kv_cache_buffer, kv_scale_cache_buffer); - } - if (mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1) - { - this->template ulyssesGenerationPostprocess(cpWorkspaceViews.mhaOutput, - reinterpret_cast(params.context_buf), cpWorkspaceViews.mhaInput, batch_beam, stream); - sync_check_cuda_error(stream); - } - return 0; - } - else if (mIsSpecDecodingEnabled && mUseSpecDecoding) - { - TLLM_CHECK_WITH_INFO(false, "No available XQA kernels are found for speculative decoding mode."); - } - else if (mFuseFp4Quant) - { - TLLM_CHECK_WITH_INFO(false, "No available kernels are found for FP4 output."); - } - else if (mKVCacheQuantMode.hasFp4KvCache()) - { - TLLM_CHECK_WITH_INFO(false, "No available kernels are found for FP4 KV cache."); - } - else - { - TLLM_LOG_DEBUG("XQA kernels are not selected in the generation phase."); - } - } - - // This is the number of kv tokens that q needs to visit, but excluding one as it will be processed before the kv - // loop. - int timestep = params.max_past_kv_length; - int const max_timesteps = std::min(timestep, params.cyclic_attention_window_size); - int estimated_min_multi_block_count - = estimate_min_multi_block_count(max_timesteps, mMaxSharedMemoryPerBlockOptin - 2048, sizeof(T)); - - if (!mMultiBlockMode && !mForceMultiBlockWarned && estimated_min_multi_block_count > 1) - { - mForceMultiBlockWarned = true; - TLLM_LOG_WARNING( - "Force using MultiBlockMode in MMHA as shared memory is not enough, " - "MultiBlockMode may have different accuracy compared to non-MultiBlockMode."); - } - - // estimate min block count to satisfy shared memory requirement to run kernel. - // Runtime check to see the actual number of blocks per sequence we need. - int32_t const max_num_seq_len_tiles = std::max(getMaxNumSeqLenTile(batch_beam), estimated_min_multi_block_count); - int32_t const min_num_seq_len_tiles = std::max(1, estimated_min_multi_block_count); - bool const enable_multi_block - = (mMultiBlockMode && max_num_seq_len_tiles > 1) || estimated_min_multi_block_count > 1; - size_t const partial_out_size - = enable_multi_block ? sizeof(T) * batch_beam * mNumHeads * mHeadSize * max_num_seq_len_tiles : 0; - size_t const partial_sum_size - = enable_multi_block ? sizeof(float) * batch_beam * mNumHeads * max_num_seq_len_tiles : 0; - size_t const partial_max_size - = enable_multi_block ? sizeof(float) * batch_beam * mNumHeads * max_num_seq_len_tiles : 0; - size_t const shift_k_cache_size = (!mPosShiftEnabled || isCrossAttention()) - ? 0 - : sizeof(T) * batch_beam * mNumHeads * mHeadSize * params.max_attention_window_size; - - AttentionGenerationWorkspaceSizes workspaceSizes{}; - workspaceSizes.cpWorkspace = cpWorkspaceSize; - workspaceSizes.partialOut = partial_out_size; - workspaceSizes.partialSum = partial_sum_size; - workspaceSizes.partialMax = partial_max_size; - workspaceSizes.shiftKCache = shift_k_cache_size; - { - auto const cascadeSizes - = tensorrt_llm::kernels::mmha::cascade::getCascadeWorkspaceSizes(batch_beam, mNumHeads, mHeadSize); - workspaceSizes.cascadeOut = cascadeSizes.out; - workspaceSizes.cascadeMax = cascadeSizes.mMax; - workspaceSizes.cascadeSum = cascadeSizes.lSum; - } - auto const workspaceLayout = AttentionWorkspaceManager::buildGenerationLayout(workspaceSizes); - auto const workspaceViews = AttentionWorkspaceManager::materializeGeneration( - params.workspace, workspaceLayout, cpMaxPaddedSequenceLength, mNumHeads, mNumKVHeads, mHeadSize); - - // Apply position embedding to the keys in the K cache - KVLinearBuffer shift_k_cache_buffer; - if (useKVCache() && mPosShiftEnabled && !isCrossAttention()) - { - shift_k_cache_buffer = KVLinearBuffer(batch_beam, params.max_attention_window_size, sizePerToken, - params.cyclic_attention_window_size, params.sink_token_length, true, - reinterpret_cast(workspaceViews.shiftKCache)); - sync_check_cuda_error(stream); - // KV cache type - KvCacheDataType const kv_cache_type = KvCacheDataType::BASE; - using DataType = typename SATypeConverter::Type; - invokeShiftKCache(kv_cache_buffer, shift_k_cache_buffer, kv_cache_type, getHeadSize(), - timestep, batch_beam, mNumKVHeads, params.beam_width, params.cyclic_attention_window_size, - params.sink_token_length, params.kv_scale_quant_orig, params.sequence_lengths, params.context_lengths, - mRotaryEmbeddingDim, mRotaryEmbeddingBase, mRotaryEmbeddingScaleType, mRotaryEmbeddingScale, - mRotaryEmbeddingMaxPositions, mPositionEmbeddingType, stream); - } - - FusedQKVMaskedAttentionDispatchParams dispatch_params{}; - dispatch_params.mUnfuseQkvGemm = mUnfuseQkvGemm; - dispatch_params.qkv_buf = attention_input; - dispatch_params.qkv_bias = params.qkv_bias; - dispatch_params.logn_scaling_ptr = logn_scaling_ptr; - dispatch_params.relative_attention_bias = relative_attention_bias; - dispatch_params.relative_attention_bias_stride = relative_attention_bias_stride; - dispatch_params.attention_mask = params.attention_mask; - dispatch_params.attention_mask_stride = params.attention_mask_stride; - dispatch_params.attention_sinks = params.attention_sinks; - dispatch_params.max_distance = max_distance; - dispatch_params.cache_indir = params.cache_indir; - dispatch_params.context_buf = mCpSize > 1 ? cpWorkspaceViews.mhaOutput : params.context_buf; // - dispatch_params.finished = finished; - dispatch_params.sequence_lengths - = params.sequence_lengths; // NOTE: current seq len including padding (fixed after meeting the finished id) - dispatch_params.max_batch_size = batch_beam; - dispatch_params.inference_batch_size = batch_beam; - dispatch_params.beam_width = params.beam_width; - dispatch_params.head_num = mNumAttnHeads; - dispatch_params.kv_head_num = mNumAttnKVHeads; - dispatch_params.size_per_head = getHeadSize(); - dispatch_params.rotary_embedding_dim = mRotaryEmbeddingDim; - dispatch_params.position_embedding_type = mPositionEmbeddingType; - dispatch_params.chunked_attention_size = mAttentionChunkSize ? *mAttentionChunkSize : INT_MAX; - dispatch_params.max_attention_window_size = params.max_attention_window_size; - dispatch_params.cyclic_attention_window_size = params.cyclic_attention_window_size; - dispatch_params.sink_token_length = isCrossAttention() ? 0 : params.sink_token_length; - dispatch_params.input_lengths = params.context_lengths; - dispatch_params.timestep = timestep; - dispatch_params.q_scaling = q_scaling; - dispatch_params.attn_logit_softcapping_scale = mAttnLogitSoftcappingScale; - dispatch_params.linear_bias_slopes = isALiBi() ? params.alibi_slopes : nullptr; - dispatch_params.ia3_tasks = ia3_tasks; - dispatch_params.ia3_key_weights = ia3_key_weights; - dispatch_params.ia3_value_weights = ia3_value_weights; - dispatch_params.qkv_scale_out = qkv_scale_out; - dispatch_params.fp8_context_fmha = mFP8ContextFMHA; - dispatch_params.attention_out_scale = params.attention_output_orig_quant; - dispatch_params.quant_option = quant_option; - dispatch_params.multi_block_mode = enable_multi_block; - dispatch_params.max_seq_len_tile = max_num_seq_len_tiles; - dispatch_params.min_seq_len_tile = min_num_seq_len_tiles; - dispatch_params.partial_out = workspaceViews.partialOut; - dispatch_params.partial_sum = workspaceViews.partialSum; - dispatch_params.partial_max = workspaceViews.partialMax; - dispatch_params.cascade_partial_out = workspaceViews.cascadeOut; - dispatch_params.cascade_partial_max = workspaceViews.cascadeMax; - dispatch_params.cascade_partial_sum = workspaceViews.cascadeSum; - dispatch_params.block_counter = mMultiBlockSemaphores.get(); - dispatch_params.kv_cache_quant_mode = mKVCacheQuantMode; - dispatch_params.kv_scale_orig_quant = params.kv_scale_orig_quant; - dispatch_params.kv_scale_quant_orig = params.kv_scale_quant_orig; - dispatch_params.kv_block_array = kv_cache_buffer; - dispatch_params.shift_k_cache_buffer = shift_k_cache_buffer; - dispatch_params.multi_processor_count = mMultiProcessorCount; - dispatch_params.rotary_embedding_base = mRotaryEmbeddingBase; - dispatch_params.rotary_embedding_scale_type = mRotaryEmbeddingScaleType; - dispatch_params.rotary_embedding_scale = mRotaryEmbeddingScale; - dispatch_params.rotary_embedding_inv_freq_cache = params.rotary_inv_freq; - dispatch_params.rotary_embedding_cos_sin_cache = params.rotary_cos_sin; - dispatch_params.rotary_embedding_short_m_scale = mRotaryEmbeddingShortMscale; - dispatch_params.rotary_embedding_long_m_scale = mRotaryEmbeddingLongMscale; - dispatch_params.rotary_embedding_max_positions = mRotaryEmbeddingMaxPositions; - dispatch_params.rotary_embedding_original_max_positions = mRotaryEmbeddingOriginalMaxPositions; - dispatch_params.position_shift_enabled = mPosShiftEnabled; - dispatch_params.rotary_cogvlm_vision_start = mVisionStart; - dispatch_params.rotary_cogvlm_vision_length = mVisionLength; - dispatch_params.cross_attention = isCrossAttention(); - dispatch_params.memory_length_per_sample = params.encoder_input_lengths; - dispatch_params.block_sparse_attention = mMaskType == AttentionMaskType::BLOCKSPARSE; - dispatch_params.block_sparse_params = mBlockSparseParams; - dispatch_params.mrope_position_deltas = params.mrope_position_deltas; - - using DataType = typename SATypeConverter::Type; - { - if (!isCrossAttention()) - { - // self attn - Masked_multihead_attention_params mmha_params; - fusedQKV_masked_attention_dispatch(mmha_params, dispatch_params, stream); - } - else - { - // cross attn - Cross_multihead_attention_params mmhca_params; - fusedQKV_masked_attention_dispatch(mmhca_params, dispatch_params, stream); - } - sync_check_cuda_error(stream); - } - - if (mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1) - { - this->template ulyssesGenerationPostprocess(cpWorkspaceViews.mhaOutput, - reinterpret_cast(params.context_buf), cpWorkspaceViews.mhaInput, batch_beam, stream); - sync_check_cuda_error(stream); - } - return 0; -} - -template int AttentionOp::enqueueGeneration( - EnqueueGenerationParams const& params, cudaStream_t stream); - -template int AttentionOp::enqueueGeneration( - EnqueueGenerationParams const& params, cudaStream_t stream); - -#ifdef ENABLE_BF16 -template int AttentionOp::enqueueGeneration<__nv_bfloat16, KVLinearBuffer>( - EnqueueGenerationParams<__nv_bfloat16> const& params, cudaStream_t stream); -#endif - -template int AttentionOp::enqueueGeneration( - EnqueueGenerationParams const& params, cudaStream_t stream); - -template int AttentionOp::enqueueGeneration( - EnqueueGenerationParams const& params, cudaStream_t stream); - -#ifdef ENABLE_BF16 -template int AttentionOp::enqueueGeneration<__nv_bfloat16, KVBlockArray>( - EnqueueGenerationParams<__nv_bfloat16> const& params, cudaStream_t stream); -#endif - -template -void AttentionOp::prepareEnqueueGeneration(EnqueueGenerationParams const& params) -{ - // self attn - if (mXqaDispatcher.get() != nullptr) - { - TLLM_LOG_TRACE("Preparing XQA kernels in prepareEnqueueGeneration."); - XQAParams xqaParams{}; - this->template convertMMHAParamsToXQAParams(xqaParams, params, /*forConfigurePlugin=*/true); - mXqaDispatcher->prepare(xqaParams); - } -} - -template void AttentionOp::prepareEnqueueGeneration(EnqueueGenerationParams const& params); - -template void AttentionOp::prepareEnqueueGeneration( - EnqueueGenerationParams const& params); - -#ifdef ENABLE_BF16 -template void AttentionOp::prepareEnqueueGeneration<__nv_bfloat16, KVLinearBuffer>( - EnqueueGenerationParams<__nv_bfloat16> const& params); -#endif - -template void AttentionOp::prepareEnqueueGeneration(EnqueueGenerationParams const& params); - -template void AttentionOp::prepareEnqueueGeneration(EnqueueGenerationParams const& params); - -#ifdef ENABLE_BF16 -template void AttentionOp::prepareEnqueueGeneration<__nv_bfloat16, KVBlockArray>( - EnqueueGenerationParams<__nv_bfloat16> const& params); -#endif - -template -KvCacheBuffers tensorrt_llm::common::op::buildKvCacheBuffers(int32_t batchSize, int32_t maxBlocksPerSeq, - int32_t tokensPerBlock, int32_t sizePerToken, int32_t cyclicAttentionWindowSize, - int32_t maxCyclicAttentionWindowSize, int32_t sinkTokenLen, bool canUseOneMoreBlock, void* primaryPoolPtr, - void* secondaryPoolPtr, void* primaryBlockScalePoolPtr, void* secondaryBlockScalePoolPtr, - KVBlockArray::DataType* blockOffsets, bool hasFp4KvCache, int32_t maxAttentionWindowSize, void* keyValueCache) -{ - KvCacheBuffers result; - if constexpr (std::is_same_v) - { - result.kvCacheBuffer = KVBlockArray(batchSize, maxBlocksPerSeq, tokensPerBlock, sizePerToken, - cyclicAttentionWindowSize, maxCyclicAttentionWindowSize, sinkTokenLen, canUseOneMoreBlock, primaryPoolPtr, - secondaryPoolPtr, blockOffsets); - if (hasFp4KvCache) - { - result.kvScaleCacheBuffer = KVBlockArray(batchSize, maxBlocksPerSeq, tokensPerBlock, sizePerToken / 8, - cyclicAttentionWindowSize, maxCyclicAttentionWindowSize, sinkTokenLen, canUseOneMoreBlock, - primaryBlockScalePoolPtr, secondaryBlockScalePoolPtr, blockOffsets); - } - } - else if constexpr (std::is_same_v) - { - TLLM_CHECK_WITH_INFO(!hasFp4KvCache, "FP4 KV cache only supports paged KV."); - TLLM_CHECK_WITH_INFO(keyValueCache != nullptr, "keyValueCache must not be null for linear KV cache."); - using BufferDataType = typename KVCacheBuffer::DataType; - result.kvCacheBuffer = KVLinearBuffer(batchSize, maxAttentionWindowSize, sizePerToken, - cyclicAttentionWindowSize, sinkTokenLen, false, reinterpret_cast(keyValueCache)); - } - return result; -} - -template KvCacheBuffers tensorrt_llm::common::op::buildKvCacheBuffers(int32_t, int32_t, - int32_t, int32_t, int32_t, int32_t, int32_t, bool, void*, void*, void*, void*, KVBlockArray::DataType*, bool, - int32_t, void*); - -template KvCacheBuffers tensorrt_llm::common::op::buildKvCacheBuffers(int32_t, int32_t, - int32_t, int32_t, int32_t, int32_t, int32_t, bool, void*, void*, void*, void*, KVBlockArray::DataType*, bool, - int32_t, void*); - -int AttentionOp::initialize() noexcept -{ - // use Ulysses for GPTAttentionPlugin - if (mAttnTpSize < 0 || mAttnCpSize < 0) - { - mAttnTpSize = mTpSize * mCpSize; - mAttnCpSize = 1; - } - mNumAttnHeads = mNumHeads * mTpSize / mAttnTpSize; - mNumAttnKVHeads = (mNumKVHeads * mTpSize + mAttnTpSize - 1) / mAttnTpSize; - - if (mCpSize != mAttnCpSize) - { - // mqa broadcast - mUlyssesMQABroadcast = (mAttnTpSize + mNumKVHeadsOrigin - 1) / mNumKVHeadsOrigin; - } - - // Pre-check whether FMHA is supported in order to save memory allocation. - if (mEnableContextFMHA) - { - mEnableContextFMHA = false; - if (!(mType == tensorrt_llm::DataType::kHALF || mType == tensorrt_llm::DataType::kBF16)) - { - TLLM_LOG_WARNING("Fall back to unfused MHA because of unsupported data type."); - } - else if (mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kRELATIVE) - { - TLLM_LOG_WARNING("Fall back to unfused MHA because of relative position embedding."); - } - else if (isCrossAttention() && useKVCache() && !mPagedKVCache) - { - // TODO: add the support for cross attention + contiguous kv cache. - TLLM_LOG_WARNING("Fall back to unfused MHA because of cross attention + contiguous kv cache."); - } - else - { - mEnableContextFMHA = true; - } - } - - // Pre-Check of FP8 Context FMHA. - if (mFP8ContextFMHA) - { - TLLM_CHECK_WITH_INFO(mEnableContextFMHA, "FP8 FMHA cannot be enabled because Context FMHA is not supported."); - TLLM_CHECK_WITH_INFO(mSM == 89 || mSM == 90 || tc::isSM100Family(mSM) || mSM == 120 || mSM == 121, - "FP8 FMHA can only be enabled on sm_89, sm_90, sm_100f, sm_120 or sm_121."); - } - - // Pre-Check of FP8 Generation MLA. - if (mFP8GenerationMLA) - { - TLLM_CHECK_WITH_INFO(mIsMLAEnabled, "FP8 Generation MLA cannot be enabled because MLA is not supported."); - TLLM_CHECK_WITH_INFO(mSM == 89 || mSM == 90 || tc::isSM100Family(mSM) || mSM == 120 || mSM == 121, - "FP8 Generation MLA is supported on Ada, Hopper or Blackwell architecture."); - } - - // Check requirements for FP4 output. - TLLM_CHECK_WITH_INFO(!mFuseFp4Quant || mEnableContextFMHA, "Context FMHA must enable if fuse_fp4_quant is enabled"); - TLLM_CHECK_WITH_INFO(!mFuseFp4Quant || tc::isSM100Family(mSM) || mSM == 120 || mSM == 121, - "fuse_fp4_quant only supports SM100f or SM120 or SM121 devices."); - - // Check requirements for FP4 KV cache. - TLLM_CHECK_WITH_INFO(!mKVCacheQuantMode.hasFp4KvCache() || mFP8ContextFMHA || mUseNvfp4MlaKvCache, - "FP4 KV cache requires FP8 context FMHA or static sparse MLA with an FP8 scratch pool"); - - TLLM_CHECK(isRoPE() == (mRotaryEmbeddingDim != 0)); - TLLM_CHECK_WITH_INFO((mSM >= 80) || (mType != tensorrt_llm::DataType::kBF16), - "Unsupported data type, pre SM 80 GPUs do not support bfloat16"); - - // Pre-check whether the head size is supported by MMHA. - // Support head size == 72 only for fmha kernels, so skip pre-check here. - if (getHeadSize() == 72) - { - ; - } - else if (!mmha_supported(getHeadSize()) && !mIsMLAEnabled) - { - TLLM_CHECK_WITH_INFO(false, "Head size %d is not supported by MMHA.", getHeadSize()); - } - - if (mIsMLAEnabled) - { - TLLM_CHECK_WITH_INFO(mEnableContextFMHA, "MLA(Deepseek v2) only support fmha"); - TLLM_CHECK_WITH_INFO(!mDenseContextFMHA, "MLA(Deepseek v2) currently not support dense fmha"); - TLLM_CHECK_WITH_INFO( - mPagedKVCache && mUseKVCache && mRemovePadding, "MLA(Deepseek v2) only support paged kv cache"); - TLLM_CHECK_WITH_INFO(!mCrossAttention, "MLA(Deepseek v2) do not support cross attention right now"); - TLLM_CHECK_WITH_INFO(mMaskType != tensorrt_llm::kernels::AttentionMaskType::CUSTOM_MASK, - "MLA(Deepseek v2) do not support custom mask right now"); - bool const mla_dims_supported = mMLAParams.qk_rope_head_dim == 64 - && ((mMLAParams.rope_append && mMLAParams.kv_lora_rank == 512) - || (!mMLAParams.rope_append && mMLAParams.kv_lora_rank == 448)); - TLLM_CHECK_WITH_INFO(mla_dims_supported, - "MLA(Deepseek v2) only supports qk_rope_head_dim=64 with kv_lora_rank=512 (rope_append=true) or " - "kv_lora_rank=448 (rope_append=false)."); - } - - mDriver = CUDADriverWrapper::getInstance(); - - auto cublasHandle = getCublasHandle(); - auto cublasLtHandle = getCublasLtHandle(); - - // Pre-warm getting environment variables - getEnvMmhaMultiblockDebug(); - getEnvMmhaBlocksPerSequence(); - - mCublasWrapper.reset(new tc::CublasMMWrapper(cublasHandle, cublasLtHandle, nullptr, nullptr)); - - if (mEnableContextFMHA) - { - // Construct the fmha runner. - MHARunnerFixedParams fmhaParams{}; - - bool const useSageAttn = mFP8ContextFMHA && !mIsMLAEnabled - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); - - // Pre-checked during constructing. - Data_type data_type, data_type_kv; - if (mType == tensorrt_llm::DataType::kHALF) - { - data_type = DATA_TYPE_FP16; - } - else if (mType == tensorrt_llm::DataType::kBF16) - { - data_type = DATA_TYPE_BF16; - } - else - { - TLLM_CHECK_WITH_INFO(false, "GPTAttentionPlugin received wrong data type."); - } - // The output dtype. - fmhaParams.dataTypeOut = mFP8AttenOutput ? DATA_TYPE_E4M3 : data_type; - data_type_kv = data_type; - - // FP8 FMHA should be used with fp8 workflow together. - if (mFP8ContextFMHA || mFP8ContextMLA) - { - if (mFP8ContextFMHA && useSageAttn && mSageAttnQkInt8) - { - data_type = DATA_TYPE_INT8; - data_type_kv = DATA_TYPE_KV_INT8_E4M3; - } - else - { - data_type = DATA_TYPE_E4M3; - data_type_kv = DATA_TYPE_E4M3; - } - } - - // The input dtype. - fmhaParams.dataType = data_type; - // The KV input data type. The default is same as dataType. - fmhaParams.dataTypeKv = data_type_kv; - // If the kernel must read from KV cache, set the dtype correctly. - if (mPagedKVCache && mPagedContextFMHA) - { - if (mKVCacheQuantMode.hasFp8KvCache()) - { - fmhaParams.dataTypeKv = DATA_TYPE_E4M3; - } - else if (mKVCacheQuantMode.hasFp4KvCache()) - { - fmhaParams.dataTypeKv = DATA_TYPE_E2M1; - } - } - if (mFuseFp4Quant) - { - // If FP4 quantization workflow is enabled, set output type to FP4. - fmhaParams.dataTypeOut = DATA_TYPE_E2M1; - } - if (mIsMLAEnabled) - { - // For FP8 MLA, currently context attention is performed in BF16. - fmhaParams.dataTypeOut = DATA_TYPE_BF16; - fmhaParams.dataTypeKv = DATA_TYPE_BF16; - } - if (mFP8ContextMLA) - { - fmhaParams.dataTypeKv = DATA_TYPE_E4M3; - fmhaParams.dataTypeOut = DATA_TYPE_BF16; - } - if (mFusesDsv4InvRopeFp8Quant) - { - fmhaParams.dataTypeOut = DATA_TYPE_E4M3; - } - // TODO: remove forceFp32Acc from MHARunnerFixedParams after adding host_runtime_perf_knobs to - // bertAttentionPlugin input tensors, so that we can change mLaunchParams.force_fp32_acc value in runtime. - fmhaParams.forceFp32Acc = false; - - // setting attention mask type based on the mask type - fmhaParams.setAttentionMaskType(static_cast(mMaskType)); - - if (isCrossAttention()) - { - // always use paged-kv-fmha if paged_kv cache is used. - fmhaParams.attentionInputLayout - = mPagedKVCache ? AttentionInputLayout::Q_PAGED_KV : AttentionInputLayout::Q_CONTIGUOUS_KV; - } - else if (!useKVCache()) - { - if (useSageAttn) - { - fmhaParams.attentionInputLayout = AttentionInputLayout::SEPARATE_Q_K_V; - } - else - { - fmhaParams.attentionInputLayout = AttentionInputLayout::PACKED_QKV; - } - } - else - { - fmhaParams.attentionInputLayout = (mPagedKVCache && mPagedContextFMHA) ? AttentionInputLayout::Q_PAGED_KV - : AttentionInputLayout::PACKED_QKV; - } - fmhaParams.isSPadded = !mRemovePadding; - fmhaParams.numQHeads = mNumAttnHeads; - fmhaParams.numKvHeads = mNumAttnKVHeads; - fmhaParams.numTokensPerBlock = mTokensPerBlock; - fmhaParams.headSize = mHeadSize; - fmhaParams.headSizeV = mHeadSize; - fmhaParams.qScaling = mQScaling; - - // mFmhaDispatcher is not used for generation MLA, but we still need to modify these values to avoid selecting - // the wrong kernel, no matter mIsGenerationMLA is true or false - if (mIsMLAEnabled) - { - if (useSparseMLA()) - { - fmhaParams.attentionInputLayout = AttentionInputLayout::Q_PAGED_KV; - fmhaParams.numKvHeads = 1; - fmhaParams.headSize = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - fmhaParams.headSizeV = mMLAParams.rope_append ? mMLAParams.kv_lora_rank - : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - fmhaParams.headSizeQkNope = mMLAParams.qk_nope_head_dim; - // Adjust the qScaling for the absorption mode. - fmhaParams.qScaling = mQScaling - * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) - / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); - } - else - { - // Context MLA always use separate_q_k_v layout - fmhaParams.attentionInputLayout = AttentionInputLayout::SEPARATE_Q_K_V; - // Context attention of MLA is different - fmhaParams.numKvHeads = mNumHeads; - fmhaParams.headSize = mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim; - // Ideally this should be mMLAParams.v_head_dim, but because we initialize both MLA - // context(v_head_dim=128) and gen(v_head_dim=512) runners in a single op, the headSizeV will be set to - // 512 when we create the gen attention op and that could fail to create the FmhaDispatcher for context - // phase. Luckily, for deepseek, qk_nope_head_dim is the same as v_head_dim in context phase. - fmhaParams.headSizeV = mMLAParams.qk_nope_head_dim; - fmhaParams.headSizeQkNope = mMLAParams.qk_nope_head_dim; - } - } - fmhaParams.attnLogitSoftcappingScale = mAttnLogitSoftcappingScale; - fmhaParams.hasAlibi = isALiBi(); - fmhaParams.scaleAlibi = isAliBiWithScale(); - fmhaParams.useSparseMLA = useSparseMLA(); - fmhaParams.useSpcompress = mUsesSpcompress; - fmhaParams.useTllmGenSparseAttention = useTllmGenSparseAttention(); - fmhaParams.fusesDsv4InvRopeFp8Quant = mFusesDsv4InvRopeFp8Quant; - - // SageAttention: set block sizes for sage quantization. - if (useSageAttn) - { - fmhaParams.sageBlockSizeQ = mSageAttnNumEltsPerBlkQ; - fmhaParams.sageBlockSizeK = mSageAttnNumEltsPerBlkK; - fmhaParams.sageBlockSizeV = mSageAttnNumEltsPerBlkV; - } - - // Load kernels from the pre-compiled cubins. - mFmhaDispatcher.reset(new FmhaDispatcher(fmhaParams)); - - // Deepseek-V2 Generation needs a differ fmha with different argumments - if (mIsMLAEnabled) - { - mEnableXQA = (mSM == kSM_120) && mIsGenerationMLA; - if (mUseTllmGen) - { - Data_type qDataType = DATA_TYPE_FP32; - Data_type kvDataType = DATA_TYPE_FP32; - Data_type outputDataType = DATA_TYPE_FP32; - - if (mType == tensorrt_llm::DataType::kHALF) - { - qDataType = DATA_TYPE_FP16; - kvDataType = DATA_TYPE_FP16; - outputDataType = DATA_TYPE_FP16; - } - else if (mType == tensorrt_llm::DataType::kBF16) - { - qDataType = DATA_TYPE_BF16; - kvDataType = DATA_TYPE_BF16; - outputDataType = DATA_TYPE_BF16; - } - else - { - TLLM_CHECK_WITH_INFO(false, "The data type is not supported."); - } - - if (mFP8GenerationMLA) - { - qDataType = DATA_TYPE_E4M3; - kvDataType = DATA_TYPE_E4M3; - } - if (mFusesDsv4InvRopeFp8Quant) - { - outputDataType = DATA_TYPE_E4M3; - } - - // Instantiate the mTllmGenFMHARunner used for MLA - mTllmGenFMHARunner.reset(new TllmGenFmhaRunner( - qDataType, kvDataType, kvDataType, outputDataType, 0, 0, 0, 0, mFusesDsv4InvRopeFp8Quant)); - } - else if (mIsGenerationMLA && !mUseGenFlashMLA) - { - // Construct the fmha runner for generation. - if (mFP8GenerationMLA) - { - data_type = DATA_TYPE_E4M3; - } - MHARunnerFixedParams fmhaParams{}; - fmhaParams.dataType = data_type; - fmhaParams.dataTypeKv = data_type; - fmhaParams.dataTypeOut = data_type; - // For FP8 MLA generation, the output type is BF16, and the quantization before o_proj is performed - // separately. - if (mFP8GenerationMLA) - { - fmhaParams.dataTypeOut = DATA_TYPE_BF16; - } - // TODO: remove forceFp32Acc from MHARunnerFixedParams after adding host_runtime_perf_knobs to - // bertAttentionPlugin input tensors, so that we can change mLaunchParams.force_fp32_acc value in - // runtime. - fmhaParams.forceFp32Acc = true; - fmhaParams.attentionMaskType - = useCustomMask() ? ContextAttentionMaskType::CUSTOM_MASK : ContextAttentionMaskType::PADDING; - // TODO: set it to Q_CONTIGUOUS_KV layout for cross-attention. - fmhaParams.attentionInputLayout = AttentionInputLayout::Q_PAGED_KV; - fmhaParams.isSPadded = !mRemovePadding; - fmhaParams.numQHeads = 1; - fmhaParams.numKvHeads = 1; - fmhaParams.headSize = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; - fmhaParams.headSizeV = mMLAParams.kv_lora_rank; - fmhaParams.qScaling = mQScaling - * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) - / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); - fmhaParams.attnLogitSoftcappingScale = mAttnLogitSoftcappingScale; - fmhaParams.hasAlibi = isALiBi(); - fmhaParams.scaleAlibi = isAliBiWithScale(); - fmhaParams.tpSize = mTpSize; - fmhaParams.tpRank = mTpRank; - mDecoderFMHARunner.reset(new FusedMHARunnerV2(fmhaParams)); - - // Only deepseek must using fmha in the generation phase when flash mla is not enabled. - if (!mUseGenFlashMLA) - { - TLLM_CHECK_WITH_INFO(mDecoderFMHARunner->isFmhaSupported(), - "Deepseek should be supported by fmha in generation part."); - } - } - if (!mIsGenerationMLA) - { - TLLM_CHECK_WITH_INFO( - mFmhaDispatcher->isSupported(), "Deepseek should be supported by fmha in context part."); - } - } - - // Fall back to unfused MHA kernels if not supported. - // Generation MLA reuses the context FMHA code path so set mEnableContextFMHA to true. - // However, do not check mFmhaDispatcher which is not used for generation MLA. - mEnableContextFMHA = mIsGenerationMLA || mFmhaDispatcher->isSupported(); - - // Only FMHA supports custom mask currently. - TLLM_CHECK_WITH_INFO( - !useCustomMask() || mEnableContextFMHA, "Only Context FMHA supports custom mask input currently."); - } - - mEnableXQA = (mEnableXQA || mIsSpecDecodingEnabled) - && (mType == tensorrt_llm::DataType::kHALF || mType == tensorrt_llm::DataType::kBF16) && mUseKVCache; - - if (mEnableXQA) - { - TLLM_LOG_DEBUG("Enabling XQA kernels for GPTAttention."); - - XqaFixedParams fixedParams{}; - fixedParams.isMLA = mIsGenerationMLA; - // TODO: support more combinations. - // Update Q and O dtype. - if (mType == tensorrt_llm::DataType::kHALF) - { - fixedParams.inputDataType = DATA_TYPE_FP16; - fixedParams.outputDataType = DATA_TYPE_FP16; - } - else if (mType == tensorrt_llm::DataType::kBF16) - { - fixedParams.inputDataType = DATA_TYPE_BF16; - fixedParams.outputDataType = DATA_TYPE_BF16; - } - // Update KV cache and math dtype. - if (mKVCacheQuantMode.hasInt8KvCache()) - { - fixedParams.kvDataType = DATA_TYPE_INT8; - fixedParams.mathDataType = fixedParams.inputDataType; - } - else if (mKVCacheQuantMode.hasFp8KvCache()) - { - fixedParams.kvDataType = DATA_TYPE_E4M3; - fixedParams.mathDataType = DATA_TYPE_E4M3; - } - else if (mKVCacheQuantMode.hasFp4KvCache()) - { - fixedParams.kvDataType = DATA_TYPE_E2M1; - fixedParams.mathDataType = DATA_TYPE_E4M3; - } - else - { - fixedParams.kvDataType = fixedParams.inputDataType; - fixedParams.mathDataType = fixedParams.inputDataType; - } - // If fuse_fp4_quant is enabled, set output data type to FP4. - if (mFuseFp4Quant) - { - fixedParams.outputDataType = DATA_TYPE_E2M1; - } - else if (mFP8AttenOutput) - { - fixedParams.outputDataType = DATA_TYPE_E4M3; - } - if (mIsSpecDecodingEnabled && !mUseTllmGen) - { - fixedParams.outputDataType = DATA_TYPE_E4M3; - TLLM_CHECK_WITH_INFO(mNumHeads % mNumKVHeads == 0, "mNumHeads should be multiples of mNumKVHeads."); - } - - fixedParams.numQHeads = mNumAttnHeads; - fixedParams.numKvHeads = mNumAttnKVHeads; - fixedParams.numTokensPerBlock = mTokensPerBlock; - fixedParams.headSize = mHeadSize; - fixedParams.qScaling = mQScaling; - fixedParams.multiBlockMode = mMultiBlockMode; - fixedParams.isPagedKv = mPagedKVCache; - fixedParams.isSpecDecoding = mIsSpecDecodingEnabled; - fixedParams.hasAlibi = isALiBi(); - fixedParams.useTllmGenSparseAttention = useTllmGenSparseAttention(); - fixedParams.specDecodingTargetMaxGenLen = mSpecDecodingTargetMaxGenLen; - - mXqaDispatcher.reset(new XqaDispatcher(fixedParams)); - - // Fall back to unfused MHA kernels if not supported. - mEnableXQA = mXqaDispatcher->isSupported(); - } - else if (mIsSpecDecodingEnabled) - { - TLLM_CHECK_WITH_INFO(false, "Speculative decoding mode doesn't support the data type or cross attention."); - } - - if (mNbMultiBlockSemaphores != 0) - { - reserveSemaphoreArray(mNbMultiBlockSemaphores); - } - -#if ENABLE_MULTI_DEVICE - if (mCpSize > 1 && COMM_SESSION.getSize() > 1) - { - TLLM_LOG_TRACE("%s start for rank %d", __PRETTY_FUNCTION__, COMM_SESSION.getRank()); - mCpNcclComm = getComm(mCpGroup); - TLLM_LOG_TRACE("%s stop for rank %d", __PRETTY_FUNCTION__, COMM_SESSION.getRank()); - } -#endif // ENABLE_MULTI_DEVICE - return 0; -} - -void AttentionOp::reserveSemaphoreArray(int32_t size) -{ - if (size == 0 || (size <= mNbMultiBlockSemaphores && mMultiBlockSemaphores != nullptr)) - { - return; - } - int32_t* ptr; - deviceMalloc(&ptr, size, false); - deviceMemSetZero(ptr, size); - mMultiBlockSemaphores.reset(ptr); - mNbMultiBlockSemaphores = size; -} - -void AttentionOp::debugCheckSemaphores(cudaStream_t stream) -{ -#ifdef NDEBUG - TLLM_CHECK_WITH_INFO(false, "debugCheckSemaphores should not be called in release build"); -#endif - - if (isCapturing(stream)) - { - // The sync for the d2h copy below won't work when we're capturing CUDA graphs. - return; - } - if (mNbMultiBlockSemaphores == 0) - { - return; - } - std::vector hostBuf(mNbMultiBlockSemaphores); - TLLM_CUDA_CHECK(tensorrt_llm::common::cudaMemcpyAsyncSanitized(hostBuf.data(), mMultiBlockSemaphores.get(), - sizeof(uint32_t) * mNbMultiBlockSemaphores, cudaMemcpyDeviceToHost, stream)); - TLLM_CUDA_CHECK(cudaStreamSynchronize(stream)); - TLLM_CHECK(std::count(hostBuf.begin(), hostBuf.end(), 0U) == mNbMultiBlockSemaphores); -} - -std::string AttentionOp::toString() const -{ - // member variables - std::stringstream ss; - ss << "gptAttentionCommon members ====================" << std::endl; - ss << "mNumHeads: " << mNumHeads << std::endl; - ss << "mNumKVHeads: " << mNumKVHeads << std::endl; - ss << "mNumKVHeadsOrigin: " << mNumKVHeadsOrigin << std::endl; - ss << "mHeadSize: " << mHeadSize << std::endl; - ss << "mUnidirectional: " << mUnidirectional << std::endl; - ss << "mQScaling: " << mQScaling << std::endl; - ss << "mRotaryEmbeddingDim: " << mRotaryEmbeddingDim << std::endl; - ss << "mRotaryEmbeddingBase: " << mRotaryEmbeddingBase << std::endl; - ss << "mRotaryEmbeddingScaleType: " << static_cast(mRotaryEmbeddingScaleType) << std::endl; - ss << "mRotaryEmbeddingScale: " << mRotaryEmbeddingScale << std::endl; - ss << "mRotaryEmbeddingMaxPositions: " << mRotaryEmbeddingMaxPositions << std::endl; - ss << "mPositionEmbeddingType: " << static_cast(mPositionEmbeddingType) << std::endl; - ss << "mUseLognScaling: " << std::boolalpha << mUseLognScaling << std::endl; - ss << "mRemovePadding: " << std::boolalpha << mRemovePadding << std::endl; - ss << "mMaskType: " << static_cast(mMaskType) << std::endl; - ss << "mPagedKVCache: " << std::boolalpha << mPagedKVCache << std::endl; - ss << "mTokensPerBlock: " << mTokensPerBlock << std::endl; - ss << "mKVCacheQuantMode: " << static_cast(mKVCacheQuantMode.value()) << std::endl; - ss << "mTpSize: " << mTpSize << std::endl; - ss << "mTpRank: " << mTpRank << std::endl; - ss << "mUnfuseQkvGemm: " << std::boolalpha << mUnfuseQkvGemm << std::endl; - ss << "mType: " << static_cast(mType) << std::endl; - ss << "mMaxContextLength: " << mMaxContextLength << std::endl; - ss << "mQKVBiasEnabled: " << std::boolalpha << mQKVBiasEnabled << std::endl; - ss << "mCrossAttention: " << std::boolalpha << mCrossAttention << std::endl; - ss << "mMaxDistance: " << mMaxDistance << std::endl; - ss << "mPosShiftEnabled: " << std::boolalpha << mPosShiftEnabled << std::endl; - ss << "mPagedContextFMHA: " << std::boolalpha << mPagedContextFMHA << std::endl; - ss << "mFP8ContextFMHA: " << std::boolalpha << mFP8ContextFMHA << std::endl; - ss << "mSageAttnNumEltsPerBlkQ: " << mSageAttnNumEltsPerBlkQ << std::endl; - ss << "mSageAttnNumEltsPerBlkK: " << mSageAttnNumEltsPerBlkK << std::endl; - ss << "mSageAttnNumEltsPerBlkV: " << mSageAttnNumEltsPerBlkV << std::endl; - ss << "mSageAttnQkInt8: " << std::boolalpha << mSageAttnQkInt8 << std::endl; - ss << "mFP8AttenOutput: " << std::boolalpha << mFP8AttenOutput << std::endl; - ss << "mFP8ContextMLA: " << std::boolalpha << mFP8ContextMLA << std::endl; - ss << "mDenseContextFMHA: " << std::boolalpha << mDenseContextFMHA << std::endl; - ss << "mEnableContextFMHA: " << std::boolalpha << mEnableContextFMHA << std::endl; - ss << "mFMHAForceFP32Acc: " << std::boolalpha << mFMHAForceFP32Acc << std::endl; - ss << "mSM: " << mSM << std::endl; - ss << "mUseTllmGen: " << mUseTllmGen << std::endl; - ss << "mIsGenerationMLA: " << std::boolalpha << mIsGenerationMLA << std::endl; - ss << "mUseGenFlashMLA: " << mUseGenFlashMLA << std::endl; - ss << "mMultiProcessorCount: " << mMultiProcessorCount << std::endl; - ss << "mMaxSharedMemoryPerBlockOptin: " << mMaxSharedMemoryPerBlockOptin << std::endl; - ss << "mMultiBlockMode: " << std::boolalpha << mMultiBlockMode << std::endl; - ss << "mEnableXQA: " << std::boolalpha << mEnableXQA << std::endl; - ss << "mUseKVCache: " << std::boolalpha << mUseKVCache << std::endl; - ss << "mForceMultiBlockWarned: " << mForceMultiBlockWarned << std::endl; - ss << "mSkipAttn: " << std::boolalpha << mSkipAttn << std::endl; - ss << "mFuseFp4Quant: " << std::boolalpha << mFuseFp4Quant << std::endl; - ss << "mCpSize: " << mCpSize << std::endl; - ss << "mCpRank: " << mCpRank << std::endl; - ss << "mCpGroup: ["; - for (auto it = mCpGroup.begin(); it != mCpGroup.end(); it++) - { - if (it != mCpGroup.begin()) - { - ss << ", "; - } - ss << *it; - } - ss << "]" << std::endl; - - return ss.str(); -} diff --git a/cpp/tensorrt_llm/common/attentionOp.h b/cpp/tensorrt_llm/common/attentionOp.h deleted file mode 100644 index bbc066b24c64..000000000000 --- a/cpp/tensorrt_llm/common/attentionOp.h +++ /dev/null @@ -1,647 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * SPDX-License-Identifier: Apache-2.0 - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ -#pragma once - -#include "tensorrt_llm/common/config.h" -#include "tensorrt_llm/common/cublasMMWrapper.h" -#include "tensorrt_llm/common/opUtils.h" -#include "tensorrt_llm/common/quantization.h" -#include "tensorrt_llm/common/tllmDataType.h" -#include "tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h" -#include "tensorrt_llm/kernels/cutlass_kernels/fp8_blockscale_gemm/fp8_blockscale_gemm.h" -#include "tensorrt_llm/kernels/decoderMaskedMultiheadAttention/decoderXQARunner.h" -#include "tensorrt_llm/kernels/fmhaDispatcher.h" -#include "tensorrt_llm/kernels/gptKernels.h" -#include "tensorrt_llm/kernels/kvCacheUtils.h" -#include "tensorrt_llm/kernels/mlaKernels.h" -#include "tensorrt_llm/kernels/sparseAttentionKernels.h" -#include "tensorrt_llm/kernels/xqaDispatcher.h" -#include -#include -#include -#include -#if ENABLE_MULTI_DEVICE -#include -#endif // ENABLE_MULTI_DEVICE - -TRTLLM_NAMESPACE_BEGIN - -namespace common::op -{ - -class AttentionOp -{ -public: - using RotaryScalingType = tensorrt_llm::kernels::RotaryScalingType; - using PositionEmbeddingType = tensorrt_llm::kernels::PositionEmbeddingType; - using AttentionMaskType = tensorrt_llm::kernels::AttentionMaskType; - - AttentionOp(){}; - ~AttentionOp() = default; - - int initialize() noexcept; - [[nodiscard]] size_t getFmhaMultiCtasKvScratchSize() const noexcept; - [[nodiscard]] int getHeadSize(bool checkInit = true) const; - [[nodiscard]] int getMaxNumSeqLenTile(int batch_beam_size = 1) const; - [[nodiscard]] size_t getWorkspaceSizeForContext(tensorrt_llm::DataType type, int32_t nbReq, - int32_t max_input_length, int32_t cross_kv_length = 0, int32_t max_num_tokens = 0, - int32_t total_kv_len = 0) const noexcept; - // Per-token byte cost of the context-MLA K/V dequant staging buffers, whose size scales with the summed - // attended KV length (`total_kv_len`). Only the fp8 context-MLA separate-Q/KV path stages these buffers; - // every other path (incl. sparse MLA, which reads K/V straight from the paged cache) returns 0. Single - // source of truth shared by getWorkspaceSizeForContext (runtime sizing) and the KV-cache estimator, so - // the two cannot drift. - [[nodiscard]] static size_t contextMlaWorkspaceBytesPerToken(int32_t numAttnHeads, int32_t qkRopeHeadDim, - int32_t qkNopeHeadDim, int32_t vHeadDim, bool fp8ContextMla, bool separateQAndKvInput, bool sparseMla) noexcept; - // total_num_seq is the sum of beam_width for multiple requests - [[nodiscard]] size_t getWorkspaceSizeForGeneration(tensorrt_llm::DataType type, int32_t total_num_seq, - int32_t max_attention_window_size, int32_t max_num_tokens, int32_t max_blocks_per_sequence) const noexcept; - - template - class EnqueueParams - { - public: - T const* attention_input = nullptr; - T const* qkv_bias = nullptr; - // Attention mask input, which has shape of [batch_size, attention_mask_stride]. - bool const* attention_mask = nullptr; - // Attention sinks with shape of [num_heads_q] float. - float const* attention_sinks = nullptr; - // Rotary inv_freq cache buffer to avoid re-computing. - float const* rotary_inv_freq = nullptr; - // Rotary cos sin cache buffer to avoid re-computing. - float2 const* rotary_cos_sin = nullptr; - // NOTE: input_seq_length might be larger than one in the medusa mode. - int32_t input_seq_length = 0; - int32_t max_past_kv_length = 0; - // By default, max_attention_window_size == cyclic_attention_window_size - // unless each layer has different cyclic kv cache length. - // Max cache capacity (used to allocate KV cache) - int32_t max_attention_window_size = 0; - // Cyclic kv cache capacity (used to get the cyclic kv cache position for new tokens) - int32_t cyclic_attention_window_size = 0; - int32_t max_cyclic_attention_window_size = 0; - bool can_use_one_more_block = false; - int32_t sink_token_length = 0; - float const* kv_scale_orig_quant = nullptr; - float const* kv_scale_quant_orig = nullptr; - float const* attention_output_orig_quant = nullptr; - float const* attention_output_sf_scale = nullptr; - T const* alibi_slopes = nullptr; - void* context_buf = nullptr; - void* context_buf_sf = nullptr; - void* key_value_cache = nullptr; - kernels::KVBlockArray::DataType* block_offsets = nullptr; - void* host_primary_pool_pointer = nullptr; - void* host_secondary_pool_pointer = nullptr; - void* host_primary_block_scale_pool_pointer = nullptr; - void* host_secondary_block_scale_pool_pointer = nullptr; - int32_t num_tokens = 0; - int32_t total_kv_len = 0; - int32_t max_blocks_per_sequence = 0; - int32_t const* sequence_lengths = nullptr; - int32_t const* context_lengths = nullptr; - int32_t const* host_context_lengths = nullptr; - void* workspace = nullptr; - // optional when logn scaling - float const* logn_scaling_ptr = nullptr; - // optional when relative position - T const* relative_attention_bias = nullptr; - int relative_attention_bias_stride = 0; - // optional when cross attention - int32_t const* encoder_input_lengths = nullptr; - int64_t const* runtime_perf_knobs = nullptr; - // optional when compute attention stats (MLA chunked prefill or Helix parallelism) - // this is a buffer of size [num_tokens, num_heads_q] with each element - // representing the max and LSE/denominator of the softmax values - float2* softmax_stats = nullptr; - // Optional SageAttention scaling factors. - float const* sage_attn_sfs_q = nullptr; - float const* sage_attn_sfs_k = nullptr; - float const* sage_attn_sfs_v = nullptr; - // Optional TRTLLM-Gen FMHA JIT warmup shape. - bool trtllm_gen_jit_warmup = false; - }; - - template - class EnqueueContextParams : public EnqueueParams - { - public: - // Attention packed mask input (used by context FMHA). - uint32_t const* attention_packed_mask = nullptr; - int32_t batch_size = 0; - float2 const* mrope_rotary_cos_sin = nullptr; - - // optional when cross attention - T const* cross_kv = nullptr; - int32_t cross_kv_length = 0; - int32_t num_encoder_tokens = 0; - kernels::MlaParams* mla_param = nullptr; - - // optional for separate QKV input, currently only used for context MLA - T const* k_ptr = nullptr; - T const* v_ptr = nullptr; - - // Helix parallelism params. - int32_t const* helix_position_offsets = nullptr; - bool const* helix_is_inactive_rank = nullptr; - - // Optional packed-varlen boundaries for context attention. When set, - // these describe attention sequences/segments and are used directly by - // FMHA instead of the workspace boundaries rebuilt from context_lengths. - int32_t const* cu_q_seqlens = nullptr; - int32_t const* cu_kv_seqlens = nullptr; - - std::string enqueueContextParamsToString() const - { - // variables from the params coming from the runtime - std::stringstream ss; - ss << "EnqueueContextParams ====================" << std::endl; - - ss << "attention_input: " << this->attention_input << std::endl; - ss << "qkv_bias: " << this->qkv_bias << std::endl; - ss << "attention_mask: " << this->attention_mask << std::endl; - ss << "attention_packed_mask: " << this->attention_packed_mask << std::endl; - ss << "rotary_inv_freq: " << this->rotary_inv_freq << std::endl; - ss << "rotary_cos_sin: " << this->rotary_cos_sin << std::endl; - ss << "input_seq_length: " << this->input_seq_length << std::endl; - ss << "max_past_kv_length: " << this->max_past_kv_length << std::endl; - ss << "max_attention_window_size: " << this->max_attention_window_size << std::endl; - ss << "cyclic_attention_window_size: " << this->cyclic_attention_window_size << std::endl; - ss << "max_cyclic_attention_window_size: " << this->max_cyclic_attention_window_size << std::endl; - ss << "can_use_one_more_block: " << (this->can_use_one_more_block ? "true" : "false") << std::endl; - ss << "sink_token_length: " << this->sink_token_length << std::endl; - if (this->context_lengths && batch_size > 0) - { - ss << "context_lengths: " - << *(runtime::ITensor::wrap((void*) this->context_lengths, tensorrt_llm::DataType::kINT32, - runtime::ITensor::makeShape({batch_size}))) - << std::endl; - } - if (this->sequence_lengths && batch_size > 0) - { - ss << "sequence_lengths: " - << *(runtime::ITensor::wrap((void*) this->sequence_lengths, tensorrt_llm::DataType::kINT32, - runtime::ITensor::makeShape({batch_size}))) - << std::endl; - } - ss << "kv_scale_orig_quant: " << this->kv_scale_orig_quant << std::endl; - ss << "kv_scale_quant_orig: " << this->kv_scale_quant_orig << std::endl; - ss << "attention_output_orig_quant: " << this->attention_output_orig_quant << std::endl; - ss << "alibi_slopes: " << this->alibi_slopes << std::endl; - ss << "context_buf: " << this->context_buf << std::endl; - ss << "context_buf_sf: " << this->context_buf_sf << std::endl; - ss << "key_value_cache: " << (half*) this->key_value_cache << std::endl; - ss << "block_offsets: " << this->block_offsets << std::endl; - ss << "host_primary_pool_pointer: " << this->host_primary_pool_pointer << std::endl; - ss << "host_secondary_pool_pointer: " << this->host_secondary_pool_pointer << std::endl; - ss << "batch_size: " << this->batch_size << std::endl; - ss << "num_tokens: " << this->num_tokens << std::endl; - ss << "total_kv_len: " << this->total_kv_len << std::endl; - ss << "max_blocks_per_sequence: " << this->max_blocks_per_sequence << std::endl; - ss << "workspace: " << this->workspace << std::endl; - ss << "logn_scaling_ptr: " << this->logn_scaling_ptr << std::endl; - ss << "relative_attention_bias: " << this->relative_attention_bias << std::endl; - ss << "relative_attention_bias_stride: " << this->relative_attention_bias_stride << std::endl; - ss << "cross_kv: " << this->cross_kv << std::endl; - ss << "cross_kv_length: " << this->cross_kv_length << std::endl; - ss << "encoder_input_lengths: " << this->encoder_input_lengths << std::endl; - ss << "num_encoder_tokens: " << this->num_encoder_tokens << std::endl; - ss << "softmaxStatsPtr: " << this->softmax_stats << std::endl; - ss << "k_ptr: " << this->k_ptr << std::endl; - ss << "v_ptr: " << this->v_ptr << std::endl; - ss << "cu_q_seqlens: " << this->cu_q_seqlens << std::endl; - ss << "cu_kv_seqlens: " << this->cu_kv_seqlens << std::endl; - return ss.str(); - } - }; - - template - int enqueueContext(EnqueueContextParams const& params, cudaStream_t stream); - - template - class EnqueueGenerationParams : public EnqueueParams - { - public: - int32_t beam_width = 1; - // Attention mask has shape of [batch_size, attention_mask_stride]. - int32_t attention_mask_stride = 0; - int32_t num_requests = 0; - int32_t const* cache_indir = nullptr; - int32_t* semaphores = nullptr; - int32_t const* host_past_key_value_lengths = nullptr; - int32_t const* mrope_position_deltas = nullptr; - - // optional when speculative decoding is used. - bool const* spec_decoding_mask = nullptr; - int32_t const* spec_decoding_packed_mask = nullptr; - int32_t const* spec_decoding_position_offsets = nullptr; - int32_t const* spec_decoding_generation_lengths = nullptr; - bool spec_decoding_is_generation_length_variable = false; - int32_t spec_decoding_max_generation_length = 1; - int64_t* spec_decoding_bl_tree_mask_offset = nullptr; - uint32_t* spec_decoding_bl_tree_mask = nullptr; - int32_t* spec_bl_tree_first_sparse_mask_offset_kv = nullptr; - // optional when fuse_fp4_quant is enabled - int32_t start_token_idx_sf = 0; - int32_t layer_idx = 0; - // Helix parallelism params. - int32_t const* helix_position_offsets = nullptr; - bool const* helix_is_inactive_rank = nullptr; - }; - - template - int enqueueGeneration(EnqueueGenerationParams const& params, cudaStream_t stream); - - template - int mlaGeneration( - kernels::MlaParams& params, EnqueueGenerationParams const& generation_params, cudaStream_t stream); - - int getFlashMlaNumSmParts(int s_q, int num_heads, int num_kv_heads, int head_size_v) const - { - static constexpr int block_size_m = 64; - int num_heads_per_head_k = s_q * num_heads / num_kv_heads; - int sm_cnt = mMultiProcessorCount; - int num_sm_parts = sm_cnt / num_kv_heads / cutlass::ceil_div(num_heads_per_head_k, block_size_m); - return num_sm_parts; - } - - static int getFlashMlaNumSmPartsStatic(int s_q, int num_heads, int num_kv_heads, int head_size_v) - { - static constexpr int block_size_m = 64; - int num_heads_per_head_k = s_q * num_heads / num_kv_heads; - int device; - cudaGetDevice(&device); - int sm_cnt; - cudaDeviceGetAttribute(&sm_cnt, cudaDevAttrMultiProcessorCount, device); - int num_sm_parts = sm_cnt / num_kv_heads / cutlass::ceil_div(num_heads_per_head_k, block_size_m); - return num_sm_parts; - } - - template - int getKvCacheElemSizeInBits() const - { - return getKvCacheElemSizeInBits(mKVCacheQuantMode, sizeof(T)); - } - - static int getKvCacheElemSizeInBits(tensorrt_llm::common::QuantMode quantMode, size_t dTypeSize) - { - if (quantMode.hasInt8KvCache() || quantMode.hasFp8KvCache()) - { - return 8; - } - else if (quantMode.hasFp4KvCache()) - { - return 4; - } - return dTypeSize * 8; - } - - // Called in configurePlugin(). - template - void prepareEnqueueGeneration(EnqueueGenerationParams const& params); - - template - bool convertMMHAParamsToXQAParams(tensorrt_llm::kernels::XQAParams& xqaParams, - EnqueueGenerationParams const& generationsParams, bool forConfigurePlugin); - - template - int ulyssesContextPreprocess(T const* input, T* output, T* buffer, EnqueueContextParams const& params, - int const* cu_q_seqlens, int const* cu_cp_partial_seqlens, cudaStream_t stream); - - template - int ulyssesContextPostprocess(T* input, T* output, T* buffer, EnqueueContextParams const& params, - int const* cu_q_seqlens, int const* cu_cp_partial_seqlens, cudaStream_t stream); - - template - int ulyssesGenerationPreprocess(T const* input, T* output, T* buffer, int32_t batch_beam, cudaStream_t stream); - - template - int ulyssesGenerationPostprocess(T* input, T* output, T* buffer, int32_t batch_beam, cudaStream_t stream); - - [[nodiscard]] bool isRelativePosition() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kRELATIVE; - } - - [[nodiscard]] bool isALiBi() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI - || mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI_WITH_SCALE; - } - - [[nodiscard]] bool isAliBiWithScale() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI_WITH_SCALE; - } - - [[nodiscard]] bool isRoPE() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_GPTJ - || mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_GPT_NEOX - || mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kLONG_ROPE - || mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kYARN - || mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_M; - } - - [[nodiscard]] bool isLongRoPE() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kLONG_ROPE; - } - - [[nodiscard]] bool isUnfusedCrossAttention() const - { - return !mEnableContextFMHA && mCrossAttention; - } - - [[nodiscard]] bool isMRoPE() const - { - return mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_M; - } - - [[nodiscard]] bool isLognScaling() const - { - return mUseLognScaling; - } - - [[nodiscard]] bool isCrossAttention() const - { - return mCrossAttention; - } - - [[nodiscard]] bool useKVCache() const - { - return mUseKVCache; - } - - [[nodiscard]] bool useCustomMask() const - { - return mMaskType == AttentionMaskType::CUSTOM_MASK; - } - - [[nodiscard]] bool useFullCustomMask() const - { - return useCustomMask() && mHasFullAttentionMask; - } - - [[nodiscard]] bool usePackedCustomMask() const - { - return useCustomMask() && mEnableContextFMHA; - } - - [[nodiscard]] bool isMLAEnabled() const - { - return mIsMLAEnabled; - } - - [[nodiscard]] bool useSparseAttention() const - { - return mUseSparseAttention && mPagedKVCache && mEnableXQA; - } - - [[nodiscard]] bool useTllmGenSparseAttentionPaged() const - { - return mUseTllmGenSparseAttentionPaged && useSparseAttention(); - } - - [[nodiscard]] bool useSparseMLA() const - { - return mUseSparseAttention && mUseTllmGen && mIsMLAEnabled; - } - - [[nodiscard]] bool useTllmGenSparseAttention() const - { - return useSparseMLA() || (mUseSparseAttention && mUseTllmGen && mUseTllmGenSparseAttention); - } - - [[nodiscard]] int smVersion() const - { - return mSM; - } - - [[nodiscard]] bool supportsNvFp4Output() const - { - bool needsUlyssesPostprocess = mCpSize > 1 && mAttnTpSize > 1 && mAttnCpSize == 1; - return mEnableContextFMHA && mEnableXQA && !needsUlyssesPostprocess; - } - - [[nodiscard]] int32_t* multiBlockSemaphores() const - { - return mMultiBlockSemaphores.get(); - } - - void reserveSemaphoreArray(int32_t size); - - void debugCheckSemaphores(cudaStream_t stream); - - [[nodiscard]] int getMultiProcessorCount() const - { - return mMultiProcessorCount; - } - - [[nodiscard]] std::string toString() const; - - int mLayerIdx = -1; - int mNumHeads = -1; - int mVisionStart = -1; - int mVisionLength = -1; - int mNumKVHeads = -1; - int mHeadSize = -1; - int mUnidirectional = 1; - float mQScaling = 1.0; - float mAttnLogitSoftcappingScale = 0.0; - int mRotaryEmbeddingDim = 0; - float mRotaryEmbeddingBase = 10000.0; - RotaryScalingType mRotaryEmbeddingScaleType = RotaryScalingType::kNONE; - float mRotaryEmbeddingScale = 1.0; - float mRotaryEmbeddingShortMscale = 1.0; - float mRotaryEmbeddingLongMscale = 1.0; - int mRotaryEmbeddingMaxPositions = 1024; - int mRotaryEmbeddingOriginalMaxPositions = 1024; - PositionEmbeddingType mPositionEmbeddingType = PositionEmbeddingType::kLEARNED_ABSOLUTE; - bool mUseLognScaling = false; - bool mRemovePadding = true; - AttentionMaskType mMaskType = AttentionMaskType::CAUSAL; - tensorrt_llm::kernels::BlockSparseParams mBlockSparseParams; - - // NOTE: default values for paged kv cache. - bool mPagedKVCache = true; - int mTokensPerBlock = 0; - tensorrt_llm::common::QuantMode mKVCacheQuantMode; - int mTpSize = 1; - int mTpRank = 0; - bool mUnfuseQkvGemm = false; - tensorrt_llm::DataType mType; - int32_t mMaxContextLength = 0; - int32_t mMaxSeqLen = 0; - int32_t mMaxNumRequests = 0; - bool mQKVBiasEnabled = false; - bool mCrossAttention = false; - int mMaxDistance = 0; - bool mPosShiftEnabled = false; - bool mPagedContextFMHA = false; - bool mFP8ContextFMHA = false; - bool mFP8AttenOutput = false; - bool mFP8ContextMLA = false; - bool mFP8GenerationMLA = false; - size_t mChunkPrefillBufferBatchSize = 1; - bool mDenseContextFMHA = false; - bool mHasFullAttentionMask = false; - bool mIsSpecDecodingEnabled = false; - bool mUseSpecDecoding = false; - bool mIsSpecDecTree = true; - bool mSpecDecodingIsGenerationLengthVariable = false; - int32_t mSpecDecodingMaxGenerationLength = 1; - // Static spec-dec tree length used by FMHA autotuning. - int32_t mSpecDecodingTargetMaxGenLen = 0; - bool mForcePrepareSpecDecTreeMask = false; - bool mIsMLAEnabled = false; - bool mIsGenerationMLA = false; - bool mUseGenFlashMLA = false; - bool mUseSparseAttention = false; - bool mUseTllmGenSparseAttentionPaged = false; - bool mUseTllmGenSparseAttention = false; - bool mUseNvfp4MlaKvCache = false; - tensorrt_llm::kernels::MlaMetaParams mMLAParams; - int mCpSize = 1; - int mCpRank = 0; - std::set mCpGroup = {}; - // These parameters are used to specifically configure the attention attributes when cp/tp_size are different - // between Attention and FFN(such as Ulysses) - int mNumAttnHeads = -1; - int mNumAttnKVHeads = -1; - int mNumKVHeadsOrigin = -1; - int mAttnTpSize = -1; - int mAttnTpRank = 0; - int mAttnCpSize = -1; - int mAttnCpRank = 0; - int mUlyssesMQABroadcast = 1; - - // fmha runner (enabled by default) - // flag: disabled = 0, enabled = 1, enabled with fp32 accumulation = 2 - bool mEnableContextFMHA = true; - bool mFMHAForceFP32Acc = false; - bool mMultiBlockMode = true; - bool mEnableXQA = true; - bool mUseKVCache = true; - bool mSkipAttn = false; - - // Whether to fuse FP4 quant into attention kernel. - bool mFuseFp4Quant = false; - // Whether to fuse DSv4 inverse-RoPE + FP8 output quant into trtllm-gen FMHA. - bool mFusesDsv4InvRopeFp8Quant = false; - - kernels::SparseAttentionParams mRuntimeSparseAttentionParams; - - // This is implementation details which we want to save when serializing, but not expose as - // a plugin field or a constructor parameter - int32_t mNbMultiBlockSemaphores = 0; - - // See [Chunked Attention] in _torch/attention/attention.py - std::optional mAttentionChunkSize = std::nullopt; - - // Skip softmax threshold scale factor. - float mSkipSoftmaxThresholdScaleFactorPrefill = 0; - float mSkipSoftmaxThresholdScaleFactorDecode = 0; - // Skip correction when the row-max increase is within this base-2 threshold. - float mSkipCorrectionThreshold = 0; - // Use spcompress (context phase, SM107 only). - bool mUsesSpcompress = false; - // Optional SageAttention block sizes. - // Currently, these are only consumed by the TllmGen backend path. - int mSageAttnNumEltsPerBlkQ = 0; - int mSageAttnNumEltsPerBlkK = 0; - int mSageAttnNumEltsPerBlkV = 0; - bool mSageAttnQkInt8 = false; -#ifdef SKIP_SOFTMAX_STAT - uint32_t* mSkipSoftmaxTotalBlocks; - uint32_t* mSkipSoftmaxSkippedBlocks; -#endif - - [[nodiscard]] auto data() const - { - return std::make_tuple(mLayerIdx, mNumHeads, mVisionStart, mVisionLength, mNumKVHeads, mHeadSize, - mUnidirectional, mQScaling, mAttnLogitSoftcappingScale, mRotaryEmbeddingDim, mRotaryEmbeddingBase, - (int8_t) mRotaryEmbeddingScaleType, mRotaryEmbeddingScale, mRotaryEmbeddingShortMscale, - mRotaryEmbeddingLongMscale, mRotaryEmbeddingMaxPositions, mRotaryEmbeddingOriginalMaxPositions, - (int8_t) mPositionEmbeddingType, mUseLognScaling, mRemovePadding, (int32_t) mMaskType, - mBlockSparseParams.data(), mPagedKVCache, mTokensPerBlock, mKVCacheQuantMode.value(), mTpSize, mTpRank, - mUnfuseQkvGemm, (int32_t) mType, mMaxContextLength, mMaxSeqLen, mMaxNumRequests, mQKVBiasEnabled, - mCrossAttention, mMaxDistance, mPosShiftEnabled, mPagedContextFMHA, mFP8ContextFMHA, mFP8AttenOutput, - mFP8ContextMLA, mFP8GenerationMLA, mChunkPrefillBufferBatchSize, mDenseContextFMHA, mHasFullAttentionMask, - mIsSpecDecodingEnabled, mUseSpecDecoding, mIsSpecDecTree, mSpecDecodingIsGenerationLengthVariable, - mSpecDecodingMaxGenerationLength, mSpecDecodingTargetMaxGenLen, mForcePrepareSpecDecTreeMask, mIsMLAEnabled, - mIsGenerationMLA, mUseGenFlashMLA, mUseSparseAttention, mUseTllmGenSparseAttentionPaged, - mUseTllmGenSparseAttention, mUseNvfp4MlaKvCache, mMLAParams.data(), mCpSize, mCpRank, mCpGroup, - mNumAttnHeads, mNumAttnKVHeads, mNumKVHeadsOrigin, mAttnTpSize, mAttnTpRank, mAttnCpSize, mAttnCpRank, - mUlyssesMQABroadcast, mEnableContextFMHA, mFMHAForceFP32Acc, mMultiBlockMode, mEnableXQA, mUseKVCache, - mSkipAttn, mFuseFp4Quant, mFusesDsv4InvRopeFp8Quant, mNbMultiBlockSemaphores, - mAttentionChunkSize.value_or(-1), mSkipSoftmaxThresholdScaleFactorPrefill, - mSkipSoftmaxThresholdScaleFactorDecode, mSkipCorrectionThreshold, mUsesSpcompress, mSageAttnNumEltsPerBlkQ, - mSageAttnNumEltsPerBlkK, mSageAttnNumEltsPerBlkV, mSageAttnQkInt8); - }; - -private: - static constexpr int kReservedMaxSeqLenTilePerSeq = 64; - - int mSM = tensorrt_llm::common::getSMVersion(); - bool mUseTllmGen = (mSM >= 100) && (mSM != 120) && (mSM != 121); - bool mForceMultiBlockWarned = false; - int mMultiProcessorCount = tensorrt_llm::common::getMultiProcessorCount(); - int mMaxSharedMemoryPerBlockOptin = tensorrt_llm::common::getMaxSharedMemoryPerBlockOptin(); - // The default copy constructor will leave it as nullptr. clone() shall initialize it. - std::shared_ptr mDriver; - UniqPtrWNullCopy mDecoderFMHARunner; - UniqPtrWNullCopy mFmhaDispatcher; - UniqPtrWNullCopy mXqaDispatcher; - UniqPtrWNullCopy mTllmGenFMHARunner; - - // The default copy constructor will leave it as nullptr. clone() shall initialize it. - UniqPtrWNullCopy mCublasWrapper; - -#if ENABLE_MULTI_DEVICE - std::shared_ptr mCpNcclComm; -#endif // ENABLE_MULTI_DEVICE - - struct Deleter - { - void operator()(void* ptr) - { - cudaFree(ptr); - } - }; - - UniqPtrWNullCopy mMultiBlockSemaphores = {}; -}; - -template -struct KvCacheBuffers -{ - KVCacheBuffer kvCacheBuffer; - KVCacheBuffer kvScaleCacheBuffer; -}; - -template -KvCacheBuffers buildKvCacheBuffers(int32_t batchSize, int32_t maxBlocksPerSeq, int32_t tokensPerBlock, - int32_t sizePerToken, int32_t cyclicAttentionWindowSize, int32_t maxCyclicAttentionWindowSize, int32_t sinkTokenLen, - bool canUseOneMoreBlock, void* primaryPoolPtr, void* secondaryPoolPtr, void* primaryBlockScalePoolPtr, - void* secondaryBlockScalePoolPtr, kernels::KVBlockArray::DataType* blockOffsets, bool hasFp4KvCache, - int32_t maxAttentionWindowSize = 0, void* keyValueCache = nullptr); - -} // namespace common::op - -TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/common/attentionWorkspace.h b/cpp/tensorrt_llm/common/attentionWorkspace.h index 3ba53da3ff34..18186b79b622 100644 --- a/cpp/tensorrt_llm/common/attentionWorkspace.h +++ b/cpp/tensorrt_llm/common/attentionWorkspace.h @@ -65,7 +65,6 @@ struct AttentionContextWorkspaceSizes size_t sageQScale{}; size_t sageKScale{}; size_t sageVScale{}; - size_t cpWorkspace{}; size_t fmhaMultiCtasKvScratch{}; }; @@ -96,7 +95,6 @@ struct AttentionContextWorkspaceLayout WorkspaceSlice sageQScale{}; WorkspaceSlice sageKScale{}; WorkspaceSlice sageVScale{}; - WorkspaceSlice cpWorkspace{}; WorkspaceSlice fmhaMultiCtasKvScratch{}; size_t totalSize{}; }; @@ -129,15 +127,11 @@ struct AttentionContextWorkspaceViews float* sageQScale{}; float* sageKScale{}; float* sageVScale{}; - T* gatherInBuffer{}; - T* gatherOutBuffer{}; - int* cuCpPartialSeqlens{}; void* fmhaMultiCtasKvScratch{}; }; struct AttentionGenerationWorkspaceSizes { - size_t cpWorkspace{}; size_t partialOut{}; size_t partialSum{}; size_t partialMax{}; @@ -152,7 +146,6 @@ struct AttentionGenerationWorkspaceSizes struct AttentionGenerationWorkspaceLayout { - WorkspaceSlice cpWorkspace{}; WorkspaceSlice partialOut{}; WorkspaceSlice partialSum{}; WorkspaceSlice partialMax{}; @@ -166,8 +159,6 @@ struct AttentionGenerationWorkspaceLayout template struct AttentionGenerationWorkspaceViews { - T* mhaOutput{}; - T* mhaInput{}; T* partialOut{}; float* partialSum{}; float* partialMax{}; @@ -254,16 +245,14 @@ class AttentionWorkspaceManager layout.sageQScale = nextSlice(offset, sizes.sageQScale, alignment); layout.sageKScale = nextSlice(offset, sizes.sageKScale, alignment); layout.sageVScale = nextSlice(offset, sizes.sageVScale, alignment); - layout.cpWorkspace = nextSlice(offset, sizes.cpWorkspace, alignment); layout.fmhaMultiCtasKvScratch = nextSlice(offset, sizes.fmhaMultiCtasKvScratch, alignment); layout.totalSize = offset; return layout; } template - static AttentionContextWorkspaceViews materializeContext(void* workspace, - AttentionContextWorkspaceLayout const& layout, size_t cpMaxPaddedSequenceLength, int headSize, int numHeads, - int numKvHeads) + static AttentionContextWorkspaceViews materializeContext( + void* workspace, AttentionContextWorkspaceLayout const& layout) { AttentionContextWorkspaceViews views{}; views.cublasWorkspace = ptr(workspace, layout.cublasWorkspace); @@ -291,15 +280,6 @@ class AttentionWorkspaceManager views.sageQScale = ptr(workspace, layout.sageQScale); views.sageKScale = ptr(workspace, layout.sageKScale); views.sageVScale = ptr(workspace, layout.sageVScale); - - views.gatherInBuffer = ptr(workspace, layout.cpWorkspace); - if (views.gatherInBuffer != nullptr) - { - auto const cpBufferElements - = cpMaxPaddedSequenceLength * static_cast(headSize) * (numHeads + 2 * numKvHeads); - views.gatherOutBuffer = views.gatherInBuffer + cpBufferElements; - views.cuCpPartialSeqlens = reinterpret_cast(views.gatherOutBuffer + cpBufferElements); - } views.fmhaMultiCtasKvScratch = ptr(workspace, layout.fmhaMultiCtasKvScratch); return views; } @@ -309,7 +289,6 @@ class AttentionWorkspaceManager { AttentionGenerationWorkspaceLayout layout{}; size_t offset = 0; - layout.cpWorkspace = nextSlice(offset, sizes.cpWorkspace, alignment); layout.partialOut = nextSlice(offset, sizes.partialOut, alignment); layout.partialSum = nextSlice(offset, sizes.partialSum, alignment); layout.partialMax = nextSlice(offset, sizes.partialMax, alignment); @@ -322,18 +301,10 @@ class AttentionWorkspaceManager } template - static AttentionGenerationWorkspaceViews materializeGeneration(void* workspace, - AttentionGenerationWorkspaceLayout const& layout, size_t cpMaxPaddedSequenceLength, int numHeads, - int numKvHeads, int headSize) + static AttentionGenerationWorkspaceViews materializeGeneration( + void* workspace, AttentionGenerationWorkspaceLayout const& layout) { AttentionGenerationWorkspaceViews views{}; - views.mhaOutput = ptr(workspace, layout.cpWorkspace); - if (views.mhaOutput != nullptr) - { - auto const cpBufferElements - = cpMaxPaddedSequenceLength * (numHeads + 2 * numKvHeads) * static_cast(headSize); - views.mhaInput = views.mhaOutput + cpBufferElements; - } views.partialOut = ptr(workspace, layout.partialOut); views.partialSum = ptr(workspace, layout.partialSum); views.partialMax = ptr(workspace, layout.partialMax); diff --git a/cpp/tensorrt_llm/common/opUtils.h b/cpp/tensorrt_llm/common/opUtils.h index fb0f2aca3d62..9e1995e26e7d 100644 --- a/cpp/tensorrt_llm/common/opUtils.h +++ b/cpp/tensorrt_llm/common/opUtils.h @@ -75,35 +75,6 @@ inline cudaDataType_t trtToCublasDtype(tensorrt_llm::DataType type) } } -// Like std::unique_ptr, but does not prevent generation of default copy constructor when used as class members. -// The copy constructor produces nullptr. So the plugin default copy constructor will not really copy this, and -// your clone() implementation is responsible for initializing such data members. -// With this we can simplify clone() implementation when there are many data members including at least one unique_ptr. -template > -class UniqPtrWNullCopy : public std::unique_ptr -{ -public: - using std::unique_ptr::unique_ptr; - - // for compatibility with std::make_unique - explicit UniqPtrWNullCopy(std::unique_ptr&& src) - : std::unique_ptr::unique_ptr{std::move(src)} - { - } - - // copy constructor produces nullptr - UniqPtrWNullCopy(UniqPtrWNullCopy const&) - : std::unique_ptr::unique_ptr{} - { - } - - // copy assignment copies nothing - UniqPtrWNullCopy& operator=(UniqPtrWNullCopy const&) - { - return *this; - } -}; - namespace { diff --git a/cpp/tensorrt_llm/kernels/fmhaDispatcher.h b/cpp/tensorrt_llm/kernels/fmhaDispatcher.h index f1291d6bc1cb..f178452da4ce 100644 --- a/cpp/tensorrt_llm/kernels/fmhaDispatcher.h +++ b/cpp/tensorrt_llm/kernels/fmhaDispatcher.h @@ -22,7 +22,7 @@ #include "tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h" #include "tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h" -using tensorrt_llm::common::op::UniqPtrWNullCopy; +#include TRTLLM_NAMESPACE_BEGIN @@ -58,9 +58,9 @@ class FmhaDispatcher // Whether to enable trtllm-gen kernels. bool mUseTllmGen; // Runner for fmha v2 kernels (for SM <= 90) - UniqPtrWNullCopy mFMHARunner; + std::unique_ptr mFMHARunner; // Runner for trtllm-gen fmha kernels (for SM == 100) - UniqPtrWNullCopy mTllmGenFMHARunner; + std::unique_ptr mTllmGenFMHARunner; // Cached SM count to avoid repeated cudaDeviceGetAttribute calls in the per-iter // FMHA dispatch hot path (isSupported / run). int mMultiProcessorCount{0}; diff --git a/cpp/tensorrt_llm/kernels/gptKernels.h b/cpp/tensorrt_llm/kernels/gptKernels.h index d855aade79c2..d9b79e1c25a8 100644 --- a/cpp/tensorrt_llm/kernels/gptKernels.h +++ b/cpp/tensorrt_llm/kernels/gptKernels.h @@ -1,5 +1,5 @@ /* - * Copyright (c) 2022-2024, NVIDIA CORPORATION. All rights reserved. + * Copyright (c) 2022-2026, NVIDIA CORPORATION. All rights reserved. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -62,6 +62,7 @@ enum class PositionEmbeddingType : int8_t kCHATGLM = 7, kYARN = 8, kROPE_M = 9, + kDEFERRED = 10, }; enum class RotaryScalingType : int8_t @@ -70,7 +71,9 @@ enum class RotaryScalingType : int8_t kLINEAR = 1, kDYNAMIC = 2, kLONG = 3, - kLLAMA3 = 4 + kLLAMA3 = 4, + kYARN = 5, + kMROPE = 6 }; struct BlockSparseParams diff --git a/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.cu b/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.cu index 1db236fc47f2..2dd6518001b7 100644 --- a/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.cu +++ b/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.cu @@ -2138,194 +2138,6 @@ __global__ void convertData(Dst* dst, Src const* src, int64_t size, float const* } } -template -__global__ void runCpTranspose(T* dst, T* dst2, T const* src, int64_t partialTokenNum, int64_t cpSize, - int64_t partialQHeads, int64_t partialKVHeads, int64_t mqaBroadcast, int64_t headSize, int64_t rank) -{ - // Do transpose from - // [partialTokenNum, mNumHeads + 2*mNumKVHeads, headSize] - // -> (view) [partialTokenNum, cpSize * partialQHeads + cpSize * partialKVHeads + cpSize * partilKVHeads, headSize] - // -> (transpose) [cpSize, partialTokenNum, partialQHeads + partialKvHeads + partialKVHeads, headSize] - using VecType = int4; - static constexpr int32_t kStep = static_cast(sizeof(VecType) / sizeof(T)); - int64_t hiddenSize = static_cast(headSize / kStep); - int64_t hiddenRestSize = static_cast(headSize % kStep); - - if (threadIdx.x >= hiddenSize + hiddenRestSize) - return; - - int64_t seqIdx = blockIdx.x; - int64_t cpIdx = blockIdx.y; - int64_t headIdx = blockIdx.z; - - auto srcHeadIdx = 0; - if (headIdx < partialQHeads) - { - srcHeadIdx = cpIdx * partialQHeads + headIdx; - } - else if (headIdx < partialQHeads + partialKVHeads) - { - srcHeadIdx = cpSize * partialQHeads + cpIdx / mqaBroadcast * partialKVHeads + (headIdx - partialQHeads); - } - else - { - srcHeadIdx = cpSize * partialQHeads + cpSize / mqaBroadcast * partialKVHeads - + cpIdx / mqaBroadcast * partialKVHeads + (headIdx - partialQHeads - partialKVHeads); - } - - if (cpIdx == rank) - { - dst = dst2; - } - VecType* out = reinterpret_cast(dst - + (cpIdx * partialTokenNum * (partialQHeads + 2 * partialKVHeads) - + seqIdx * (partialQHeads + 2 * partialKVHeads) + headIdx) - * headSize); - VecType const* in = reinterpret_cast( - src + (seqIdx * (partialQHeads * cpSize + 2 * partialKVHeads * cpSize / mqaBroadcast) + srcHeadIdx) * headSize); - - for (int hiddenIdx = threadIdx.x; hiddenIdx < hiddenSize + hiddenRestSize; hiddenIdx += blockDim.x) - { - if (hiddenIdx < hiddenSize) - out[hiddenIdx] = in[hiddenIdx]; - else - reinterpret_cast(out + hiddenSize)[hiddenIdx - hiddenSize] - = reinterpret_cast(in + hiddenSize)[hiddenIdx - hiddenSize]; - } -} - -template -__global__ void runCpTransposeToSeqMajor(T* dst, T const* srcMyRank, T const* srcOtherRank, int64_t partialLength, - int64_t cpSize, int64_t newPartialHeads, int64_t headSize, int64_t rank) -{ - // Do transpose from - // [totalLength, mNumHeads / cp, Dh] - // -> (view) [cp, partialLength, mNumHeads / cp, Dh] - // -> (transpose) [partialLength, mNumHeads, Dh] - using VecType = int4; - static constexpr int32_t kStep = static_cast(sizeof(VecType) / sizeof(T)); - int64_t hiddenSize = static_cast(headSize * newPartialHeads / kStep); - int64_t hiddenRestSize = static_cast(headSize * newPartialHeads % kStep); - - if (threadIdx.x >= hiddenSize + hiddenRestSize) - return; - - int64_t cpIdx = blockIdx.x; - int64_t seqIdx = blockIdx.y; - T const* src; - if (cpIdx == rank) - { - src = srcMyRank; - } - else - { - src = srcOtherRank; - } - VecType const* in - = reinterpret_cast(src + (cpIdx * partialLength + seqIdx) * headSize * newPartialHeads); - VecType* out = reinterpret_cast(dst + (seqIdx * cpSize + cpIdx) * headSize * newPartialHeads); - for (int hiddenIdx = threadIdx.x; hiddenIdx < hiddenSize + hiddenRestSize; hiddenIdx += blockDim.x) - { - if (hiddenIdx < hiddenSize) - out[hiddenIdx] = in[hiddenIdx]; - else - reinterpret_cast(out + hiddenSize)[hiddenIdx - hiddenSize] - = reinterpret_cast(in + hiddenSize)[hiddenIdx - hiddenSize]; - } -} - -template -__global__ void runCpTranspose2(T* dst, T const* src, int32_t const* q_seq_lengths, int32_t const* cu_q_seqlens, - int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, int64_t batchSize, - int64_t partialHeads, int64_t headSize) -{ - // Do transpose from - // [cpSize_Length, bs, partialLength, partialHead, headSize] - // -> (transpose) [tokens(bs, cpSize_Length, partialLength), partialHead, headSize] - // paddings of partial length are removed here - using VecType = int4; - static constexpr int32_t kStep = static_cast(sizeof(VecType) / sizeof(T)); - int64_t hiddenSize = static_cast(headSize * partialHeads / kStep); - int64_t hiddenRestSize = static_cast(headSize * partialHeads % kStep); - - if (threadIdx.x >= hiddenSize + hiddenRestSize) - return; - - int64_t cpIdx = blockIdx.x; - int64_t tokenIdx = blockIdx.y; - int64_t seqIdx = blockIdx.z; - int64_t length = q_seq_lengths[seqIdx]; - int64_t partialLength = (length + cpSize - 1) / cpSize; - int64_t partialLengthOutIdx = cu_q_seqlens[seqIdx] + partialLength * cpIdx; // cpMajor - int64_t partialLengthInIdx = cu_cp_partial_seqlens[batchSize] * cpIdx + cu_cp_partial_seqlens[seqIdx]; // bsMajor - if (cpIdx + 1 == cpSize) - { - partialLength = length - partialLength * (cpSize - 1); - } - // for (int partialTokenIdx = blockIdx.y % maxPartalLength; partialTokenIdx < partialLength; - for (int partialTokenIdx = tokenIdx; partialTokenIdx < partialLength; partialTokenIdx += maxPartalLength) - { - - VecType* out - = reinterpret_cast(dst + (partialLengthOutIdx + partialTokenIdx) * partialHeads * headSize); - VecType const* in - = reinterpret_cast(src + (partialLengthInIdx + partialTokenIdx) * partialHeads * headSize); - for (int hiddenIdx = threadIdx.x; hiddenIdx < hiddenSize + hiddenRestSize; hiddenIdx += blockDim.x) - { - if (hiddenIdx < hiddenSize) - out[hiddenIdx] = in[hiddenIdx]; - else - reinterpret_cast(out + hiddenSize)[hiddenIdx - hiddenSize] - = reinterpret_cast(in + hiddenSize)[hiddenIdx - hiddenSize]; - } - } -} - -template -__global__ void runCpTransposeToSeqMajor2(T* dst, T const* src, int32_t const* q_seq_lengths, - int32_t const* cu_q_seqlens, int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, - int64_t batchSize, int64_t partialHeads, int64_t headSize) -{ - // Do transpose from - // [tokens(bs, cp, paritalLength), partialHeads, headSize] - // -> (transpose) [cp, partialTokens(bs, partialLength), partialHeads, headSize] - // paddings of partial length are added here - using VecType = int4; - static constexpr int32_t kStep = static_cast(sizeof(VecType) / sizeof(T)); - int64_t hiddenSize = static_cast(headSize * partialHeads / kStep); - int64_t hiddenRestSize = static_cast(headSize * partialHeads % kStep); - - if (threadIdx.x >= hiddenSize + hiddenRestSize) - return; - - int64_t cpIdx = blockIdx.x; - int64_t tokenIdx = blockIdx.y; - int64_t seqIdx = blockIdx.z; - int64_t length = q_seq_lengths[seqIdx]; - int64_t partialLength = (length + cpSize - 1) / cpSize; - int64_t partialLengthOutIdx = cu_q_seqlens[seqIdx] + partialLength * cpIdx; // cpMajor - int64_t partialLengthInIdx = cu_cp_partial_seqlens[batchSize] * cpIdx + cu_cp_partial_seqlens[seqIdx]; // bsMajor - if (cpIdx + 1 == cpSize) - { - partialLength = length - partialLength * (cpSize - 1); - } - for (int partialTokenIdx = tokenIdx; partialTokenIdx < partialLength; partialTokenIdx += maxPartalLength) - { - VecType* out - = reinterpret_cast(dst + (partialLengthInIdx + partialTokenIdx) * partialHeads * headSize); - VecType const* in - = reinterpret_cast(src + (partialLengthOutIdx + partialTokenIdx) * partialHeads * headSize); - for (int hiddenIdx = threadIdx.x; hiddenIdx < hiddenSize + hiddenRestSize; hiddenIdx += blockDim.x) - { - if (hiddenIdx < hiddenSize) - out[hiddenIdx] = in[hiddenIdx]; - else - reinterpret_cast(out + hiddenSize)[hiddenIdx - hiddenSize] - = reinterpret_cast(in + hiddenSize)[hiddenIdx - hiddenSize]; - } - } -} - } // unnamed namespace template @@ -2346,88 +2158,6 @@ INSTANTIATE_invokeConversion(__nv_fp8_e4m3, half); INSTANTIATE_invokeConversion(__nv_fp8_e4m3, __nv_bfloat16); #undef INSTANTIATE_invokeConversion -template -void invokeCpTranspose(T* dst, T* dst2, T const* src, int64_t partialTokenNum, int64_t cpSize, int64_t partialQHeads, - int64_t partialKVHeads, int64_t mqaBroadcast, int64_t headSize, int64_t rank, cudaStream_t stream) -{ - dim3 grid(partialTokenNum, cpSize, partialQHeads + 2 * partialKVHeads); - dim3 block(128); - runCpTranspose<<>>( - dst, dst2, src, partialTokenNum, cpSize, partialQHeads, partialKVHeads, mqaBroadcast, headSize, rank); -} - -#define INSTANTIATE_invokeCpTranspose(T) \ - template void invokeCpTranspose(T * dst, T * dst2, T const* src, int64_t partialLength, int64_t cpSize, \ - int64_t partialQHeads, int64_t partialKVHeads, int64_t mqaBroadcast, int64_t headSize, int64_t rank, \ - cudaStream_t stream) -INSTANTIATE_invokeCpTranspose(float); -INSTANTIATE_invokeCpTranspose(half); -INSTANTIATE_invokeCpTranspose(__nv_bfloat16); -#undef INSTANTIATE_invokeCpTranspose - -template -void invokeCpTransposeToSeqMajor(T* dst, T const* srcMyRank, T const* srcOtherRank, int64_t partialLength, - int64_t cpSize, int64_t newPartialHeads, int64_t headSize, int64_t rank, cudaStream_t stream) -{ - dim3 grid(cpSize, partialLength); - dim3 block(128); - runCpTransposeToSeqMajor<<>>( - dst, srcMyRank, srcOtherRank, partialLength, cpSize, newPartialHeads, headSize, rank); -} - -#define INSTANTIATE_invokeCpTransposeToSeqMajor(T) \ - template void invokeCpTransposeToSeqMajor(T * dst, T const* srcMyRank, T const* srcOtherRank, \ - int64_t partialLength, int64_t cpSize, int64_t newPartialHeads, int64_t headSize, int64_t rank, \ - cudaStream_t stream) -INSTANTIATE_invokeCpTransposeToSeqMajor(float); -INSTANTIATE_invokeCpTransposeToSeqMajor(half); -INSTANTIATE_invokeCpTransposeToSeqMajor(__nv_bfloat16); -INSTANTIATE_invokeCpTransposeToSeqMajor(__nv_fp8_e4m3); -#undef INSTANTIATE_invokeCpTransposeToSeqMajor - -template -void invokeCpTranspose2(T* dst, T const* src, int32_t const* q_seq_lengths, int32_t const* cu_q_seqlens, - int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, int64_t batchSize, - int64_t partialHeads, int64_t headSize, cudaStream_t stream) -{ - int64_t clipedMaxPartialLength = min(static_cast(maxPartalLength), 512); - dim3 grid(cpSize, clipedMaxPartialLength, batchSize); - dim3 block(128); - runCpTranspose2<<>>(dst, src, q_seq_lengths, cu_q_seqlens, cu_cp_partial_seqlens, cpSize, - clipedMaxPartialLength, batchSize, partialHeads, headSize); -} - -#define INSTANTIATE_invokeCpTranspose2(T) \ - template void invokeCpTranspose2(T * dst, T const* src, int32_t const* q_seq_lengths, \ - int32_t const* cu_q_seqlens, int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, \ - int64_t batchSize, int64_t partialHeads, int64_t headSize, cudaStream_t stream) -INSTANTIATE_invokeCpTranspose2(float); -INSTANTIATE_invokeCpTranspose2(half); -INSTANTIATE_invokeCpTranspose2(__nv_bfloat16); -#undef INSTANTIATE_invokeCpTranspose2 - -template -void invokeCpTransposeToSeqMajor2(T* dst, T const* src, int32_t const* q_seq_lengths, int32_t const* cu_q_seqlens, - int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, int64_t batchSize, - int64_t partialHeads, int64_t headSize, cudaStream_t stream) -{ - int64_t clipedMaxPartialLength = min(static_cast(maxPartalLength), 512); - dim3 grid(cpSize, clipedMaxPartialLength, batchSize); - dim3 block(128); - runCpTransposeToSeqMajor2<<>>(dst, src, q_seq_lengths, cu_q_seqlens, - cu_cp_partial_seqlens, cpSize, clipedMaxPartialLength, batchSize, partialHeads, headSize); -} - -#define INSTANTIATE_invokeCpTransposeToSeqMajor2(T) \ - template void invokeCpTransposeToSeqMajor2(T * dst, T const* src, int32_t const* q_seq_lengths, \ - int32_t const* cu_q_seqlens, int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, \ - int64_t batchSize, int64_t partialHeads, int64_t headSize, cudaStream_t stream) -INSTANTIATE_invokeCpTransposeToSeqMajor2(float); -INSTANTIATE_invokeCpTransposeToSeqMajor2(half); -INSTANTIATE_invokeCpTransposeToSeqMajor2(__nv_bfloat16); -INSTANTIATE_invokeCpTransposeToSeqMajor2(__nv_fp8_e4m3); -#undef INSTANTIATE_invokeCpTransposeToSeqMajor2 - } // namespace kernels TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.h b/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.h index 81fbe2d7b3ae..4b98c28bf7d0 100644 --- a/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.h +++ b/cpp/tensorrt_llm/kernels/unfusedAttentionKernels.h @@ -473,24 +473,6 @@ void invokeShiftKCache(KVCacheBuffer const& kvCacheBuffer, KVLinearBuffer const& template void invokeConversion(Dst* dst, Src const* src, int64_t size, float const* __restrict__ scale, cudaStream_t stream); -template -void invokeCpTranspose(T* dst, T* dst2, T const* src, int64_t partialLength, int64_t cpSize, int64_t partialQHeads, - int64_t partialKVHeads, int64_t mqaBroadcast, int64_t headSize, int64_t rank, cudaStream_t stream); - -template -void invokeCpTransposeToSeqMajor(T* dst, T const* srcMyRank, T const* srcOtherRank, int64_t partialLength, - int64_t cpSize, int64_t newPartialHeads, int64_t headSize, int64_t rank, cudaStream_t stream); - -template -void invokeCpTranspose2(T* dst, T const* src, int32_t const* q_seq_lengths, int32_t const* cu_q_seqlens, - int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, int64_t batchSize, - int64_t partialHeads, int64_t headSize, cudaStream_t stream); - -template -void invokeCpTransposeToSeqMajor2(T* dst, T const* src, int32_t const* q_seq_lengths, int32_t const* cu_q_seqlens, - int32_t const* cu_cp_partial_seqlens, int64_t cpSize, int64_t maxPartalLength, int64_t batchSize, - int64_t partialHeads, int64_t headSize, cudaStream_t stream); - } // namespace kernels TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kernels/xqaDispatcher.h b/cpp/tensorrt_llm/kernels/xqaDispatcher.h index 83eb3febe46b..a581ce6a7155 100644 --- a/cpp/tensorrt_llm/kernels/xqaDispatcher.h +++ b/cpp/tensorrt_llm/kernels/xqaDispatcher.h @@ -23,8 +23,9 @@ #include "tensorrt_llm/kernels/multiHeadAttentionCommon.h" #include "tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunner.h" +#include + using namespace tensorrt_llm::common; -using tensorrt_llm::common::op::UniqPtrWNullCopy; TRTLLM_NAMESPACE_BEGIN @@ -107,9 +108,9 @@ class XqaDispatcher // The multi-processor count. int mMultiProcessorCount; // Runner for decoder XQA kernels (for SM <= 90) - UniqPtrWNullCopy mDecoderXqaRunner; + std::unique_ptr mDecoderXqaRunner; // Runner for trtllm-gen XQA kernels (for SM == 100) - UniqPtrWNullCopy mTllmGenFMHARunner; + std::unique_ptr mTllmGenFMHARunner; protected: template diff --git a/cpp/tensorrt_llm/nanobind/CMakeLists.txt b/cpp/tensorrt_llm/nanobind/CMakeLists.txt index b192ecc53188..55b84e3df514 100755 --- a/cpp/tensorrt_llm/nanobind/CMakeLists.txt +++ b/cpp/tensorrt_llm/nanobind/CMakeLists.txt @@ -51,7 +51,8 @@ endif() # compilation of kvCacheManagerV2.cpp. target_include_directories( ${TRTLLM_NB_MODULE} PRIVATE ${PROJECT_SOURCE_DIR}/tensorrt_llm/common/sha256) - +target_include_directories(${TRTLLM_NB_MODULE} + PRIVATE ${TRTLLM_FMHA_PARAMS_GENERATED_INCLUDE_DIR}) set_property(TARGET ${TRTLLM_NB_MODULE} PROPERTY POSITION_INDEPENDENT_CODE ON) target_link_directories(${TRTLLM_NB_MODULE} PUBLIC diff --git a/cpp/tensorrt_llm/nanobind/bindings.cpp b/cpp/tensorrt_llm/nanobind/bindings.cpp index b48915d5f83d..8e733e6b0f8c 100644 --- a/cpp/tensorrt_llm/nanobind/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/bindings.cpp @@ -32,6 +32,8 @@ #include "tensorrt_llm/batch_manager/peftCacheManagerConfig.h" #include "tensorrt_llm/common/quantization.h" #include "tensorrt_llm/common/tllmDataType.h" +#include "tensorrt_llm/kernels/gptKernels.h" +#include "tensorrt_llm/kernels/mlaKernels.h" #include "tensorrt_llm/nanobind/batch_manager/algorithms.h" #include "tensorrt_llm/nanobind/batch_manager/bindings.h" #include "tensorrt_llm/nanobind/batch_manager/cacheTransceiver.h" @@ -60,6 +62,7 @@ namespace nb = nanobind; namespace tb = tensorrt_llm::batch_manager; namespace tpb = tensorrt_llm::nanobind::batch_manager; namespace tc = tensorrt_llm::common; +namespace tk = tensorrt_llm::kernels; namespace tr = tensorrt_llm::runtime; namespace tle = tensorrt_llm::executor; using SizeType32 = tr::SizeType32; @@ -185,6 +188,59 @@ NB_MODULE(TRTLLM_NB_MODULE, m) .value("NVFP4", tensorrt_llm::DataType::kFP4) .export_values(); + // Attention kernel enums and parameter structs. Bound here rather than in the thop + // submodule so every binding in the module presents them the same way; nanobind's + // enum caster still accepts a plain int for a valid enumerator, so Python callers may + // pass either these or the matching tensorrt_llm.functional IntEnum. + nb::enum_(m, "AttentionMaskType", nb::is_arithmetic()) + .value("PADDING", tk::AttentionMaskType::PADDING) + .value("CAUSAL", tk::AttentionMaskType::CAUSAL) + .value("SLIDING_WINDOW_CAUSAL", tk::AttentionMaskType::SLIDING_WINDOW_CAUSAL) + .value("BIDIRECTIONAL", tk::AttentionMaskType::BIDIRECTIONAL) + .value("BIDIRECTIONALGLM", tk::AttentionMaskType::BIDIRECTIONALGLM) + .value("BLOCKSPARSE", tk::AttentionMaskType::BLOCKSPARSE) + .value("CUSTOM_MASK", tk::AttentionMaskType::CUSTOM_MASK); + + nb::enum_(m, "PositionEmbeddingType", nb::is_arithmetic()) + .value("LEARNED_ABSOLUTE", tk::PositionEmbeddingType::kLEARNED_ABSOLUTE) + .value("ROPE_GPTJ", tk::PositionEmbeddingType::kROPE_GPTJ) + .value("ROPE_GPT_NEOX", tk::PositionEmbeddingType::kROPE_GPT_NEOX) + .value("LONG_ROPE", tk::PositionEmbeddingType::kLONG_ROPE) + .value("ALIBI", tk::PositionEmbeddingType::kALIBI) + .value("ALIBI_WITH_SCALE", tk::PositionEmbeddingType::kALIBI_WITH_SCALE) + .value("RELATIVE", tk::PositionEmbeddingType::kRELATIVE) + .value("CHATGLM", tk::PositionEmbeddingType::kCHATGLM) + .value("YARN", tk::PositionEmbeddingType::kYARN) + .value("ROPE_M", tk::PositionEmbeddingType::kROPE_M) + .value("DEFERRED", tk::PositionEmbeddingType::kDEFERRED); + + nb::enum_(m, "RotaryScalingType", nb::is_arithmetic()) + .value("NONE", tk::RotaryScalingType::kNONE) + .value("LINEAR", tk::RotaryScalingType::kLINEAR) + .value("DYNAMIC", tk::RotaryScalingType::kDYNAMIC) + .value("LONG", tk::RotaryScalingType::kLONG) + .value("LLAMA3", tk::RotaryScalingType::kLLAMA3) + .value("YARN", tk::RotaryScalingType::kYARN) + .value("MROPE", tk::RotaryScalingType::kMROPE); + + nb::class_(m, "BlockSparseParams") + .def(nb::init<>()) + .def_rw("block_size", &tk::BlockSparseParams::block_size) + .def_rw("homo_head_pattern", &tk::BlockSparseParams::homo_head_pattern) + .def_rw("num_local_blocks", &tk::BlockSparseParams::num_local_blocks) + .def_rw("vertical_stride", &tk::BlockSparseParams::vertical_stride); + + nb::class_(m, "MlaMetaParams") + .def(nb::init<>()) + .def_rw("q_lora_rank", &tk::MlaMetaParams::q_lora_rank) + .def_rw("kv_lora_rank", &tk::MlaMetaParams::kv_lora_rank) + .def_rw("qk_nope_head_dim", &tk::MlaMetaParams::qk_nope_head_dim) + .def_rw("qk_rope_head_dim", &tk::MlaMetaParams::qk_rope_head_dim) + .def_rw("v_head_dim", &tk::MlaMetaParams::v_head_dim) + .def_rw("predicted_tokens_per_seq", &tk::MlaMetaParams::predicted_tokens_per_seq) + .def_rw("num_layers", &tk::MlaMetaParams::num_layers) + .def_rw("rope_append", &tk::MlaMetaParams::rope_append); + nb::enum_(m, "GptModelVariant") .value("GPT", tr::ModelConfig::ModelVariant::kGpt) .value("GLM", tr::ModelConfig::ModelVariant::kGlm) @@ -249,6 +305,7 @@ NB_MODULE(TRTLLM_NB_MODULE, m) nb::arg("mamba_in_proj_size") = 0, nb::arg("mamba_inner_size") = 0, nb::arg("moe_latent_size") = 0); nb::class_(m, "QuantMode") + .def(nb::init(), nb::arg("value")) .def_static("none", &tc::QuantMode::none) .def_static("int4_weights", &tc::QuantMode::int4Weights) .def_static("int8_weights", &tc::QuantMode::int8Weights) @@ -298,6 +355,9 @@ NB_MODULE(TRTLLM_NB_MODULE, m) .def(nb::self == nb::self) .def(nb::self != nb::self); + // Python's QuantMode is an IntFlag, so let its value cross as one. + nb::implicitly_convertible(); + nb::class_(m, "ModelConfig") .def(nb::init(), nb::arg("vocab_size"), nb::arg("num_layers"), nb::arg("num_attention_layers"), nb::arg("num_rnn_layers"), diff --git a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp index 518c80671132..03975b410982 100644 --- a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp @@ -17,9 +17,9 @@ #include "bindings.h" #include #include +#include #include #include -#include #include #include #include @@ -29,6 +29,8 @@ #include #include #include +#include +#include namespace nb = nanobind; @@ -38,89 +40,6 @@ namespace tensorrt_llm::nanobind::thop namespace { -template -nb::object optionalToObject(std::optional const& value) -{ - if (value.has_value()) - { - return nb::cast(*value); - } - return nb::none(); -} - -nb::tuple trtllmGenContextPreprocessBinding(torch::Tensor qkv_input, torch::Tensor workspace, - torch::Tensor sequence_lengths, torch::Tensor context_lengths, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, - std::optional kv_scale_orig_quant, std::optional kv_scale_quant_orig, - std::optional attention_output_orig_quant, std::optional rotary_inv_freq, - std::optional rotary_cos_sin, std::optional mrope_rotary_cos_sin, int64_t layer_idx, - int64_t num_heads, int64_t num_kv_heads, int64_t head_size, int64_t tokens_per_block, int64_t mask_type, - int64_t kv_cache_quant_mode, int64_t max_attention_window_size, int64_t cyclic_attention_window_size, - int64_t num_tokens, int64_t batch_size, int64_t input_seq_length, int64_t max_past_kv_length, - int64_t rotary_embedding_dim, double rotary_embedding_base, int64_t rotary_embedding_scale_type, - double rotary_embedding_scale, int64_t rotary_embedding_max_positions, int64_t position_embedding_type, - double bmm1_scale, double bmm2_scale, int64_t attention_chunk_size, bool fp8_context_fmha, bool paged_context_fmha, - bool is_mla_enable, int64_t multi_processor_count, int64_t total_num_blocks, int64_t kv_factor, - bool need_build_kv_cache_metadata, std::optional cross_kv, bool cross_attention, - bool skip_fmha_workspace) -{ - auto result = [&]() - { - nb::gil_scoped_release release; - return torch_ext::trtllmGenContextPreprocess(qkv_input, workspace, sequence_lengths, context_lengths, - kv_cache_block_offsets, host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, kv_scale_orig_quant, - kv_scale_quant_orig, attention_output_orig_quant, rotary_inv_freq, rotary_cos_sin, mrope_rotary_cos_sin, - layer_idx, num_heads, num_kv_heads, head_size, tokens_per_block, mask_type, kv_cache_quant_mode, - max_attention_window_size, cyclic_attention_window_size, num_tokens, batch_size, input_seq_length, - max_past_kv_length, rotary_embedding_dim, rotary_embedding_base, rotary_embedding_scale_type, - rotary_embedding_scale, rotary_embedding_max_positions, position_embedding_type, bmm1_scale, bmm2_scale, - attention_chunk_size, fp8_context_fmha, paged_context_fmha, is_mla_enable, multi_processor_count, - total_num_blocks, kv_factor, need_build_kv_cache_metadata, cross_kv, cross_attention, skip_fmha_workspace); - }(); - - return nb::make_tuple(std::get<0>(result), optionalToObject(std::get<1>(result)), - optionalToObject(std::get<2>(result)), optionalToObject(std::get<3>(result)), - optionalToObject(std::get<4>(result)), optionalToObject(std::get<5>(result)), std::get<6>(result), - std::get<7>(result), std::get<8>(result), std::get<9>(result), std::get<10>(result), std::get<11>(result)); -} - -nb::tuple trtllmGenGenerationPreprocessBinding(torch::Tensor qkv_input, torch::Tensor workspace, - torch::Tensor sequence_lengths, std::optional spec_decoding_generation_lengths, - std::optional spec_decoding_position_offsets, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, - std::optional kv_scale_orig_quant, std::optional kv_scale_quant_orig, - std::optional attention_output_orig_quant, std::optional rotary_inv_freq, - std::optional rotary_cos_sin, std::optional mrope_position_deltas, int64_t layer_idx, - int64_t seq_offset, int64_t num_heads, int64_t num_kv_heads, int64_t head_size, int64_t tokens_per_block, - int64_t kv_cache_quant_mode, int64_t max_attention_window_size, int64_t cyclic_attention_window_size, - int64_t num_tokens, int64_t batch_beam, int64_t input_seq_length, int64_t max_past_kv_length, - int64_t rotary_embedding_dim, double rotary_embedding_base, int64_t rotary_embedding_scale_type, - double rotary_embedding_scale, int64_t rotary_embedding_max_positions, int64_t position_embedding_type, - double bmm1_scale, double bmm2_scale, bool fp8_context_fmha, int64_t predicted_tokens_per_seq, - int64_t attention_chunk_size, int64_t multi_processor_count, int64_t total_num_blocks, int64_t kv_factor, - bool need_build_kv_cache_metadata, bool cross_attention, bool skip_fmha_workspace) -{ - auto result = [&]() - { - nb::gil_scoped_release release; - return torch_ext::trtllmGenGenerationPreprocess(qkv_input, workspace, sequence_lengths, - spec_decoding_generation_lengths, spec_decoding_position_offsets, kv_cache_block_offsets, - host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, kv_scale_orig_quant, kv_scale_quant_orig, - attention_output_orig_quant, rotary_inv_freq, rotary_cos_sin, mrope_position_deltas, layer_idx, seq_offset, - num_heads, num_kv_heads, head_size, tokens_per_block, kv_cache_quant_mode, max_attention_window_size, - cyclic_attention_window_size, num_tokens, batch_beam, input_seq_length, max_past_kv_length, - rotary_embedding_dim, rotary_embedding_base, rotary_embedding_scale_type, rotary_embedding_scale, - rotary_embedding_max_positions, position_embedding_type, bmm1_scale, bmm2_scale, fp8_context_fmha, - predicted_tokens_per_seq, attention_chunk_size, multi_processor_count, total_num_blocks, kv_factor, - need_build_kv_cache_metadata, cross_attention, skip_fmha_workspace); - }(); - - return nb::make_tuple(std::get<0>(result), optionalToObject(std::get<1>(result)), - optionalToObject(std::get<2>(result)), optionalToObject(std::get<3>(result)), std::get<4>(result), - std::get<5>(result), std::get<6>(result), optionalToObject(std::get<7>(result)), std::get<8>(result), - std::get<9>(result), std::get<10>(result), std::get<11>(result)); -} - bool fusedContextFmhaKernelExists( int headSize, tensorrt_llm::DataType kvCacheDtype, int tokensPerBlock, tensorrt_llm::DataType outputDtype) { @@ -214,57 +133,58 @@ void initBindings(nb::module_& m) m.attr(kv.first) = kv.second; } - m.def("attention", &torch_ext::attention, - // Parameters with default values: using std::nullopt for trailing optional arguments (omittable) - // and .none() for optional arguments followed by parameters without default values (not omittable). - nb::arg("q"), nb::arg("k").none(), nb::arg("v").none(), nb::arg("output"), nb::arg("output_sf").none(), - nb::arg("workspace_").none(), nb::arg("sequence_length"), nb::arg("host_past_key_value_lengths"), - nb::arg("host_total_kv_lens"), nb::arg("context_lengths"), nb::arg("host_context_lengths"), - nb::arg("host_request_types"), nb::arg("max_context_q_len_override").none(), - nb::arg("kv_cache_block_offsets").none(), nb::arg("host_kv_cache_pool_pointers").none(), - nb::arg("host_kv_cache_pool_mapping").none(), nb::arg("cache_indirection").none(), - nb::arg("kv_scale_orig_quant").none(), nb::arg("kv_scale_quant_orig").none(), nb::arg("out_scale").none(), - nb::arg("rotary_inv_freq").none(), nb::arg("rotary_cos_sin").none(), nb::arg("latent_cache").none(), - nb::arg("q_pe").none(), nb::arg("block_ids_per_seq").none(), nb::arg("attention_sinks").none(), - nb::arg("is_fused_qkv"), nb::arg("update_kv_cache"), nb::arg("predicted_tokens_per_seq"), - nb::arg("local_layer_idx"), nb::arg("num_heads"), nb::arg("num_kv_heads"), nb::arg("head_size"), - nb::arg("tokens_per_block").none(), nb::arg("max_num_requests"), nb::arg("max_context_length"), - nb::arg("max_seq_len"), nb::arg("attention_window_size"), nb::arg("beam_width"), nb::arg("mask_type"), - nb::arg("quant_mode"), nb::arg("q_scaling"), nb::arg("position_embedding_type"), nb::arg("rope_dim"), - nb::arg("rope_base"), nb::arg("rope_scale_type"), nb::arg("rope_scale"), nb::arg("rope_short_m_scale"), - nb::arg("rope_long_m_scale"), nb::arg("rope_max_positions"), nb::arg("rope_original_max_positions"), - nb::arg("use_paged_context_fmha"), nb::arg("attention_input_type").none(), nb::arg("is_mla_enable"), - nb::arg("chunked_prefill_buffer_batch_size").none(), nb::arg("q_lora_rank").none(), - nb::arg("kv_lora_rank").none(), nb::arg("qk_nope_head_dim").none(), nb::arg("qk_rope_head_dim").none(), - nb::arg("v_head_dim").none(), nb::arg("rope_append").none(), nb::arg("mrope_rotary_cos_sin").none(), - nb::arg("mrope_position_deltas").none(), nb::arg("helix_position_offsets").none(), - nb::arg("helix_is_inactive_rank").none(), nb::arg("attention_chunk_size").none(), - nb::arg("softmax_stats_tensor").none(), nb::arg("is_spec_decoding_enabled"), nb::arg("use_spec_decoding"), - nb::arg("is_spec_dec_tree"), nb::arg("spec_decoding_generation_lengths").none(), - nb::arg("spec_decoding_position_offsets_for_cpp").none(), nb::arg("spec_decoding_packed_mask").none(), - nb::arg("spec_decoding_bl_tree_mask_offset").none(), nb::arg("spec_decoding_bl_tree_mask").none(), - nb::arg("spec_bl_tree_first_sparse_mask_offset_kv").none(), nb::arg("sparse_kv_indices").none(), - nb::arg("sparse_kv_offsets").none(), nb::arg("sparse_attn_indices").none(), - nb::arg("sparse_attn_offsets").none(), nb::arg("sparse_attn_indices_block_size"), - nb::arg("num_sparse_topk") = std::nullopt, nb::arg("sparse_attn_kv_lens") = std::nullopt, - nb::arg("skip_softmax_threshold_scale_factor_prefill") = std::nullopt, - nb::arg("skip_softmax_threshold_scale_factor_decode") = std::nullopt, - nb::arg("skip_softmax_stat") = std::nullopt, nb::arg("cu_q_seqlens") = std::nullopt, - nb::arg("cu_kv_seqlens") = std::nullopt, nb::arg("fmha_scheduler_counter") = std::nullopt, - nb::arg("mla_bmm1_scale") = std::nullopt, nb::arg("mla_bmm2_scale") = std::nullopt, - nb::arg("quant_q_buffer") = std::nullopt, nb::arg("flash_mla_tile_scheduler_metadata") = std::nullopt, - nb::arg("flash_mla_num_splits") = std::nullopt, nb::arg("sage_attn_num_elts_per_blk_q") = 0, - nb::arg("sage_attn_num_elts_per_blk_k") = 0, nb::arg("sage_attn_num_elts_per_blk_v") = 0, - nb::arg("sage_attn_qk_int8") = false, nb::arg("num_contexts") = 0, nb::arg("num_ctx_tokens") = 0, - nb::arg("trtllm_gen_jit_warmup") = false, nb::arg("aux_kv_cache_pool_ptr") = std::nullopt, - nb::arg("is_cross") = false, nb::arg("cross_kv") = std::nullopt, - nb::arg("relative_attention_bias") = std::nullopt, nb::arg("relative_attention_max_distance") = 0, - nb::arg("spec_decoding_target_max_draft_tokens") = std::nullopt, nb::arg("quant_scale_qkv") = std::nullopt, - nb::arg("dsv4_inv_rope_cos_sin_cache") = std::nullopt, nb::arg("enable_dsv4_epilogue_fusion") = false, - nb::arg("force_prepare_spec_dec_tree_mask") = false, nb::arg("max_num_sequences") = std::nullopt, - nb::arg("kv_norm_weight") = std::nullopt, nb::arg("kv_norm_eps") = 1e-6, - nb::arg("skip_correction_threshold") = 0.0, nb::arg("uses_spcompress") = std::nullopt, - "Multi-head attention operation", nb::call_guard()); + // ---- Phased attention API ---- + // Bind the nested structs before FmhaParams so its `fwd` member has a registered + // type. def_rw hands back a reference into the parent, which is what lets Python + // fill a nested struct in place. + auto sparseBackendForwardArgs + = nb::class_(m, "SparseBackendForwardArgs").def(nb::init<>()); +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) \ + sparseBackendForwardArgs.def_rw(#name, &torch_ext::SparseBackendForwardArgs::name); +#include "tensorrt_llm/thop/sparse_backend_forward_args_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + auto sparseRuntimeParams = nb::class_(m, "SparseRuntimeParams").def(nb::init<>()); +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) \ + sparseRuntimeParams.def_rw(#name, &torch_ext::SparseRuntimeParams::name); +#include "tensorrt_llm/thop/sparse_runtime_params_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + auto attentionForwardArgs + = nb::class_(m, "AttentionForwardArgs").def(nb::init<>()); +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) \ + attentionForwardArgs.def_rw(#name, &torch_ext::AttentionForwardArgs::name); +#include "tensorrt_llm/thop/attention_forward_args_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + auto staticAttentionConfig + = nb::class_(m, "StaticAttentionConfig").def(nb::init<>()); +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) \ + staticAttentionConfig.def_rw(#name, &torch_ext::StaticAttentionConfig::name); +#include "tensorrt_llm/thop/static_attention_config_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + auto fmhaParams = nb::class_(m, "FmhaParams").def(nb::init<>()); +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) fmhaParams.def_rw(#name, &torch_ext::FmhaParams::name); +#include "tensorrt_llm/thop/fmha_params_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + nb::class_(m, "AttentionOp", + "One attention layer's native op. Reuse the same instance for every call from that layer so its " + "kernel runners are built once; drop it when the layer's quantization config changes and on " + "teardown, since it owns a cuBLAS handle, the runners, and the context-parallel communicator.") + .def(nb::init(), nb::arg("config")) + .def("run_context", &torch_ext::AttentionOp::runContext, nb::arg("params"), "Phased attention context pass.", + nb::call_guard()) + .def("run_generation", &torch_ext::AttentionOp::runGeneration, nb::arg("params"), + "Phased attention generation pass.", nb::call_guard()) + .def("run_mla_generation", &torch_ext::AttentionOp::runMlaGeneration, nb::arg("params"), + "Phased attention MLA generation pass.", nb::call_guard()) + .def("get_attention_workspace_size", &torch_ext::AttentionOp::getAttentionWorkspaceSize, nb::arg("params"), + nb::arg("num_tokens"), nb::arg("max_attention_window_size"), nb::arg("num_gen_tokens"), + nb::arg("max_blocks_per_sequence"), + "Max of the context/generation workspace byte requirements for sizing FmhaParams.workspace.", + nb::call_guard()); m.def( "get_helix_workspace_size_per_rank", @@ -272,7 +192,7 @@ void initBindings(nb::module_& m) nb::arg("cp_size"), "Get helix all-to-all workspace size per rank in bytes"); m.def("get_context_mla_workspace_bytes_per_token", - &tensorrt_llm::common::op::AttentionOp::contextMlaWorkspaceBytesPerToken, nb::arg("num_attn_heads"), + &tensorrt_llm::torch_ext::AttentionOp::contextMlaWorkspaceBytesPerToken, nb::arg("num_attn_heads"), nb::arg("qk_rope_head_dim"), nb::arg("qk_nope_head_dim"), nb::arg("v_head_dim"), nb::arg("fp8_context_mla"), nb::arg("separate_q_and_kv_input"), nb::arg("sparse_mla"), "Per-token byte cost of the context-MLA K/V dequant staging buffers (scales with summed attended KV " @@ -365,7 +285,7 @@ void initBindings(nb::module_& m) nb::arg("use_sparse_attention") = false, nb::arg("skip_fmha_workspace") = false, "Return the C++ trtllm-gen generation workspace layout."); - m.def("trtllm_gen_context_preprocess", &trtllmGenContextPreprocessBinding, nb::arg("qkv_input"), + m.def("trtllm_gen_context_preprocess", &torch_ext::trtllmGenContextPreprocess, nb::arg("qkv_input"), nb::arg("workspace"), nb::arg("sequence_lengths"), nb::arg("context_lengths"), nb::arg("kv_cache_block_offsets").none(), nb::arg("host_kv_cache_pool_pointers").none(), nb::arg("host_kv_cache_pool_mapping").none(), nb::arg("kv_scale_orig_quant").none(), @@ -382,7 +302,7 @@ void initBindings(nb::module_& m) nb::arg("is_mla_enable"), nb::arg("multi_processor_count"), nb::arg("total_num_blocks"), nb::arg("kv_factor"), nb::arg("need_build_kv_cache_metadata") = true, nb::arg("cross_kv").none() = nb::none(), nb::arg("cross_attention") = false, nb::arg("skip_fmha_workspace") = false, - "Fused nanobind context preprocess for trtllm-gen attention."); + "Fused nanobind context preprocess for trtllm-gen attention.", nb::call_guard()); m.def("trtllm_gen_context_postprocess", &torch_ext::trtllmGenContextPostprocess, nb::arg("qkv_input"), nb::arg("workspace"), nb::arg("sequence_lengths"), nb::arg("context_lengths"), @@ -405,28 +325,24 @@ void initBindings(nb::module_& m) [](torch::Tensor host_kv_cache_pool_pointers, torch::Tensor host_kv_cache_pool_mapping, torch::Tensor kv_cache_block_offsets, int64_t layer_idx, int64_t num_kv_heads, int64_t tokens_per_block, int64_t head_dim, int64_t kv_factor, int64_t total_num_blocks, int64_t kv_cache_quant_mode, - int64_t batch_start, int64_t batch_size, at::ScalarType dtype) -> nb::tuple + int64_t batch_start, int64_t batch_size, + at::ScalarType dtype) -> std::tuple> { - at::Tensor kvPool; - std::optional kvScalePool; - at::Tensor blockTables; - { - nb::gil_scoped_release release; - std::tie(kvPool, kvScalePool) = torch_ext::buildFlashinferTrtllmGenPagedKvCacheBuffers( - host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, layer_idx, num_kv_heads, tokens_per_block, - head_dim, kv_factor, total_num_blocks, kv_cache_quant_mode, dtype); - auto const mapping = torch_ext::readKvCachePoolMapping(host_kv_cache_pool_mapping, layer_idx); - blockTables = kv_cache_block_offsets.select(0, mapping.poolIndex).narrow(0, batch_start, batch_size); - } - return nb::make_tuple(nb::cast(kvPool), nb::cast(blockTables), optionalToObject(kvScalePool)); + auto [kvPool, kvScalePool] = torch_ext::buildFlashinferTrtllmGenPagedKvCacheBuffers( + host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, layer_idx, num_kv_heads, tokens_per_block, + head_dim, kv_factor, total_num_blocks, kv_cache_quant_mode, dtype); + auto const mapping = torch_ext::readKvCachePoolMapping(host_kv_cache_pool_mapping, layer_idx); + auto blockTables = kv_cache_block_offsets.select(0, mapping.poolIndex).narrow(0, batch_start, batch_size); + return {std::move(kvPool), std::move(blockTables), std::move(kvScalePool)}; }, nb::arg("host_kv_cache_pool_pointers"), nb::arg("host_kv_cache_pool_mapping"), nb::arg("kv_cache_block_offsets"), nb::arg("layer_idx"), nb::arg("num_kv_heads"), nb::arg("tokens_per_block"), nb::arg("head_dim"), nb::arg("kv_factor"), nb::arg("total_num_blocks"), nb::arg("kv_cache_quant_mode"), nb::arg("batch_start"), nb::arg("batch_size"), nb::arg("dtype"), - "Build flashinfer-style KV cache pool view and slice block tables for a given layer."); + "Build flashinfer-style KV cache pool view and slice block tables for a given layer.", + nb::call_guard()); - m.def("trtllm_gen_generation_preprocess", &trtllmGenGenerationPreprocessBinding, nb::arg("qkv_input"), + m.def("trtllm_gen_generation_preprocess", &torch_ext::trtllmGenGenerationPreprocess, nb::arg("qkv_input"), nb::arg("workspace"), nb::arg("sequence_lengths"), nb::arg("spec_decoding_generation_lengths").none(), nb::arg("spec_decoding_position_offsets").none(), nb::arg("kv_cache_block_offsets").none(), nb::arg("host_kv_cache_pool_pointers").none(), nb::arg("host_kv_cache_pool_mapping").none(), @@ -442,6 +358,7 @@ void initBindings(nb::module_& m) nb::arg("bmm2_scale"), nb::arg("fp8_context_fmha"), nb::arg("predicted_tokens_per_seq"), nb::arg("attention_chunk_size"), nb::arg("multi_processor_count"), nb::arg("total_num_blocks"), nb::arg("kv_factor"), nb::arg("need_build_kv_cache_metadata") = true, nb::arg("cross_attention") = false, - nb::arg("skip_fmha_workspace") = false, "Fused nanobind generation preprocess for trtllm-gen attention."); + nb::arg("skip_fmha_workspace") = false, "Fused nanobind generation preprocess for trtllm-gen attention.", + nb::call_guard()); } } // namespace tensorrt_llm::nanobind::thop diff --git a/cpp/tensorrt_llm/thop/CMakeLists.txt b/cpp/tensorrt_llm/thop/CMakeLists.txt index 04fe40d2baf1..602890d28c4b 100644 --- a/cpp/tensorrt_llm/thop/CMakeLists.txt +++ b/cpp/tensorrt_llm/thop/CMakeLists.txt @@ -136,6 +136,8 @@ add_library( # th_common with target_sources() so the list stays next to the files. add_subdirectory(moe) set_property(TARGET th_common PROPERTY POSITION_INDEPENDENT_CODE ON) +target_include_directories(th_common + PRIVATE ${TRTLLM_FMHA_PARAMS_GENERATED_INCLUDE_DIR}) target_link_libraries( th_common PRIVATE ${TORCH_LIBRARIES} th_utils ${Python3_LIBRARIES} ${SHARED_TARGET} pg_utils) diff --git a/cpp/tensorrt_llm/thop/attentionOp.cpp b/cpp/tensorrt_llm/thop/attentionOp.cpp index ffd09ef927a8..f2693282c14a 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.cpp +++ b/cpp/tensorrt_llm/thop/attentionOp.cpp @@ -15,34 +15,45 @@ * limitations under the License. */ -#include "tensorrt_llm/common/attentionOp.h" +#include "tensorrt_llm/thop/attentionOp.h" +#include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/attentionWorkspace.h" -#include "tensorrt_llm/common/dataType.h" +#include "tensorrt_llm/common/envUtils.h" +#include "tensorrt_llm/common/logger.h" +#include "tensorrt_llm/common/memoryUtils.h" +#include "tensorrt_llm/common/sageQuant.h" #include "tensorrt_llm/common/tllmDataType.h" +#include "tensorrt_llm/kernels/decoderMaskedMultiheadAttention.h" +#include "tensorrt_llm/kernels/decoderMaskedMultiheadAttention/cascadeAttentionKernel.h" #include "tensorrt_llm/kernels/flashMLA/flash_mla.h" #include "tensorrt_llm/kernels/gptKernels.h" +#include "tensorrt_llm/kernels/kvCacheUtils.h" #include "tensorrt_llm/kernels/mlaKernels.h" +#include "tensorrt_llm/kernels/multiHeadAttentionCommon.h" #include "tensorrt_llm/kernels/sparseAttentionKernels.h" +#include "tensorrt_llm/kernels/unfusedAttentionKernels.h" +#include "tensorrt_llm/runtime/iBuffer.h" #include "tensorrt_llm/runtime/torchUtils.h" #include "tensorrt_llm/runtime/utils/debugUtils.h" -#include "tensorrt_llm/thop/attentionOp.h" -#include "tensorrt_llm/thop/thUtils.h" +#include "tensorrt_llm/runtime/utils/mpiUtils.h" +#include #include -#include -#include +#include #include -#include #include -#include + +using namespace tensorrt_llm::kernels; +namespace tc = tensorrt_llm::common; TRTLLM_NAMESPACE_BEGIN namespace torch_ext { -using tensorrt_llm::common::op::AttentionOp; +using tensorrt_llm::common::op::AttentionContextWorkspaceSizes; +using tensorrt_llm::common::op::AttentionFlashMlaWorkspaceSizes; +using tensorrt_llm::common::op::AttentionGenerationWorkspaceSizes; using tensorrt_llm::common::op::AttentionWorkspaceManager; -using tensorrt_llm::common::op::OpCustomHash; -using tensorrt_llm::runtime::RequestType; +using tensorrt_llm::common::op::AttentionXqaWorkspaceSizes; namespace { @@ -322,1259 +333,3633 @@ TrtllmGenGenerationWorkspaceViews TrtllmAttentionWorkspaceManager::materializeGe return materializeGenerationWorkspace(workspace, layout); } -namespace trtllm::attention +template +struct SATypeConverter { -using tensorrt_llm::kernels::KVBlockArray; -using tensorrt_llm::kernels::MlaParams; -using tensorrt_llm::kernels::SparseAttentionParams; -using tensorrt_llm::torch_ext::KvCachePoolPointers; -using tensorrt_llm::torch_ext::buildKvCachePoolPointers; + using Type = T; +}; -enum class AttentionInputType : int8_t +template <> +struct SATypeConverter { - Mixed, - ContextOnly, - GenerationOnly, + using Type = uint16_t; }; -class RunnerBase +template +struct FusedQKVMaskedAttentionDispatchParams { -public: - int32_t beam_width; - int32_t max_num_requests; - int32_t max_num_sequences; - int32_t attention_window_size; - - auto data() const - { - return std::make_tuple(beam_width, max_num_requests, max_num_sequences, attention_window_size); - }; + T const* qkv_buf; + T const* qkv_bias; + T const* relative_attention_bias; + bool const* attention_mask; + float const* attention_sinks; + float const* logn_scaling_ptr; + int const* cache_indir; + void* context_buf; + bool const* finished; + int const* sequence_lengths; + int max_batch_size; + int inference_batch_size; + int beam_width; + int head_num; + int kv_head_num; + int size_per_head; + int rotary_embedding_dim; + float rotary_embedding_base; + RotaryScalingType rotary_embedding_scale_type; + float rotary_embedding_scale; + float const* rotary_embedding_inv_freq_cache; + float2 const* rotary_embedding_cos_sin_cache; + float rotary_embedding_short_m_scale; + float rotary_embedding_long_m_scale; + int rotary_embedding_max_positions; + int rotary_embedding_original_max_positions; + int rotary_cogvlm_vision_start; + int rotary_cogvlm_vision_length; + PositionEmbeddingType position_embedding_type; + bool position_shift_enabled; + + int chunked_attention_size; + int attention_mask_stride; + int max_attention_window_size; + int cyclic_attention_window_size; + int sink_token_length; + int const* input_lengths; + int timestep; + float q_scaling; + float attn_logit_softcapping_scale; + int relative_attention_bias_stride; + T const* linear_bias_slopes; + int const* ia3_tasks; + T const* ia3_key_weights; + T const* ia3_value_weights; + float const* qkv_scale_out; + bool fp8_context_fmha; + float const* attention_out_scale; + bool mUnfuseQkvGemm; + tc::QuantMode quant_option; + bool multi_block_mode; + int max_seq_len_tile; + int min_seq_len_tile; + T* partial_out; + float* partial_sum; + float* partial_max; + int* block_counter; + // Cascade attention prefix-side workspace (fp32). Sliced from the same + // generation workspace and forwarded into Multihead_attention_params, so the + // cascade fast-path allocates nothing of its own. + float* cascade_partial_out{}; + float* cascade_partial_max{}; + float* cascade_partial_sum{}; + float const* kv_scale_orig_quant; + float const* kv_scale_quant_orig; + tc::QuantMode kv_cache_quant_mode; + int multi_processor_count; + KVCacheBuffer kv_block_array; + KVLinearBuffer shift_k_cache_buffer; + bool cross_attention = false; + int const* memory_length_per_sample = nullptr; + int max_distance = 0; + bool block_sparse_attention = false; + BlockSparseParams block_sparse_params; + int32_t const* mrope_position_deltas; +}; - virtual ~RunnerBase() = default; - virtual void prepare(AttentionOp& op) const = 0; - virtual int64_t getWorkspaceSize(AttentionOp const& op, int const num_tokens, int const max_attention_window_size, - int const num_gen_tokens, int const max_blocks_per_sequence, int const ctx_total_kv_len = 0, - int const maxCrossKvLength = 0) const - = 0; - // typically, we use single qkv input, but for context MLA, we use separate qkv inputs - virtual void run(AttentionOp& op, bool const is_context, int32_t const seq_offset, int32_t const num_seqs, - int32_t const token_offset, int32_t const num_tokens, int32_t const predicted_tokens_per_seq, - torch::Tensor workspace, torch::Tensor output, torch::optional output_sf, torch::Tensor qkv_or_q, - torch::optional k, torch::optional v, torch::Tensor sequence_length, - torch::Tensor host_past_key_value_lengths, int32_t const total_kv_len, torch::Tensor context_lengths, - torch::Tensor host_context_lengths, std::optional max_context_q_len_override, - torch::optional kv_cache_block_offsets, - torch::optional host_kv_cache_pool_pointers, - torch::optional host_kv_cache_pool_mapping, torch::optional cache_indirection, - torch::optional kv_scale_orig_quant, torch::optional kv_scale_quant_orig, - torch::optional out_scale, torch::optional rotary_inv_freq, - torch::optional rotary_cos_sin, torch::optional latent_cache, - torch::optional q_pe, torch::optional block_ids_per_seq, - torch::optional mrope_rotary_cos_sin, torch::optional mrope_position_deltas, - std::optional helix_position_offsets, std::optional helix_is_inactive_rank, - torch::optional softmax_stats_tensor, - std::optional spec_decoding_generation_lengths, - std::optional spec_decoding_position_offsets_for_cpp, - std::optional spec_decoding_packed_mask, - std::optional spec_decoding_bl_tree_mask_offset, - std::optional spec_decoding_bl_tree_mask, - std::optional spec_bl_tree_first_sparse_mask_offset_kv, - torch::optional attention_sinks, torch::optional sparse_kv_indices, - torch::optional sparse_kv_offsets, torch::optional sparse_attn_indices, - torch::optional sparse_attn_offsets, int64_t const sparse_attn_indices_block_size, - int32_t const num_sparse_topk, std::optional sparse_attn_kv_lens, - std::optional cu_q_seqlens, std::optional cu_kv_seqlens, - std::optional fmha_scheduler_counter, std::optional mla_bmm1_scale, - std::optional mla_bmm2_scale, std::optional quant_q_buffer, - std::optional flash_mla_tile_scheduler_metadata, - std::optional flash_mla_num_splits, bool trtllm_gen_jit_warmup, - std::optional aux_kv_cache_pool_ptr, bool const is_cross, std::optional cross_kv, - std::optional relative_attention_bias, - std::optional quant_scale_qkv = std::nullopt, - std::optional dsv4_inv_rope_cos_sin_cache = std::nullopt, - bool enable_dsv4_epilogue_fusion = false, std::optional kv_norm_weight = std::nullopt, - double kv_norm_eps = 1e-6) const - = 0; +template +struct ConvertMMHAToXQAParamsHelper +{ + static constexpr Data_type data_type = DATA_TYPE_FP16; + static constexpr bool supported = false; }; -template -class Runner : public RunnerBase +template <> +struct ConvertMMHAToXQAParamsHelper<__half, KVLinearBuffer> { -public: - void prepare(AttentionOp& op) const override - { - AttentionOp::EnqueueGenerationParams enqueueParams; - enqueueParams.max_attention_window_size = attention_window_size; - enqueueParams.cyclic_attention_window_size = attention_window_size; - enqueueParams.max_cyclic_attention_window_size = attention_window_size; - enqueueParams.beam_width = beam_width; - enqueueParams.num_requests = max_num_requests; - - op.prepareEnqueueGeneration(enqueueParams); - - // Always reserve SemaphoreArray (for multi-block mode) as MMHA may enable multi-block mode when shared memory - // is not enough. - // The attention kernel might split the heads into multiple blocks, so we might need to reserve more semaphores. - // Use mMultiProcessorCount as the lower-bound to make sure we reserve enough semaphores. - op.reserveSemaphoreArray(std::max(op.mNumHeads * max_num_sequences, op.getMultiProcessorCount())); - } - - int64_t getWorkspaceSize(AttentionOp const& op, int const num_tokens, int const max_attention_window_size, - int const num_gen_tokens, int const max_blocks_per_sequence, int const ctx_total_kv_len = 0, - int const maxCrossKvLength = 0) const override - { - size_t const context_workspace_size = op.getWorkspaceSizeForContext( - op.mType, max_num_requests, op.mMaxContextLength, maxCrossKvLength, num_tokens, ctx_total_kv_len); - size_t const generation_workspace_size = op.getWorkspaceSizeForGeneration( - op.mType, max_num_sequences, max_attention_window_size, num_gen_tokens, max_blocks_per_sequence); - - return std::max(context_workspace_size, generation_workspace_size); - } - - void run(AttentionOp& op, bool const is_context, int32_t const seq_offset, int32_t const num_seqs, - int32_t const token_offset, int32_t const num_tokens, int32_t const predicted_tokens_per_seq, - torch::Tensor workspace, torch::Tensor output, torch::optional output_sf, torch::Tensor qkv_or_q, - torch::optional k, torch::optional v, torch::Tensor sequence_length, - torch::Tensor host_past_key_value_lengths, int32_t const total_kv_len, torch::Tensor context_lengths, - torch::Tensor host_context_lengths, std::optional max_context_q_len_override, - torch::optional kv_cache_block_offsets, - torch::optional host_kv_cache_pool_pointers, - torch::optional host_kv_cache_pool_mapping, torch::optional cache_indirection, - torch::optional kv_scale_orig_quant, torch::optional kv_scale_quant_orig, - torch::optional out_scale, torch::optional rotary_inv_freq, - torch::optional rotary_cos_sin, torch::optional latent_cache, - torch::optional q_pe, torch::optional block_ids_per_seq, - torch::optional mrope_rotary_cos_sin, torch::optional mrope_position_deltas, - std::optional helix_position_offsets, std::optional helix_is_inactive_rank, - torch::optional softmax_stats_tensor, - std::optional spec_decoding_generation_lengths, - std::optional spec_decoding_position_offsets_for_cpp, - std::optional spec_decoding_packed_mask, - std::optional spec_decoding_bl_tree_mask_offset, - std::optional spec_decoding_bl_tree_mask, - std::optional spec_bl_tree_first_sparse_mask_offset_kv, - torch::optional attention_sinks, torch::optional sparse_kv_indices, - torch::optional sparse_kv_offsets, torch::optional sparse_attn_indices, - torch::optional sparse_attn_offsets, int64_t const sparse_attn_indices_block_size, - int32_t const num_sparse_topk, std::optional sparse_attn_kv_lens, - std::optional cu_q_seqlens, std::optional cu_kv_seqlens, - std::optional fmha_scheduler_counter, std::optional mla_bmm1_scale, - std::optional mla_bmm2_scale, std::optional quant_q_buffer, - std::optional flash_mla_tile_scheduler_metadata, - std::optional flash_mla_num_splits, bool trtllm_gen_jit_warmup, - std::optional aux_kv_cache_pool_ptr, bool const is_cross, std::optional cross_kv, - std::optional relative_attention_bias, std::optional quant_scale_qkv, - std::optional dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion, - std::optional kv_norm_weight, double kv_norm_eps) const override - { - auto stream = at::cuda::getCurrentCUDAStream(qkv_or_q.get_device()); - T* attention_input = static_cast(qkv_or_q.slice(0, token_offset).data_ptr()); - T* k_ptr = nullptr; - T* v_ptr = nullptr; - AttentionOutT* context_buf = static_cast(output.slice(0, token_offset).data_ptr()); - TORCH_CHECK(!op.mFuseFp4Quant || output_sf.has_value()); - TORCH_CHECK(!enable_dsv4_epilogue_fusion || output_sf.has_value()); - void* context_buf_sf = (op.mFuseFp4Quant || enable_dsv4_epilogue_fusion) ? output_sf->data_ptr() : nullptr; - - // Rotary inv_freq, cos_sin cache to avoid re-computing. - float const* rotary_inv_freq_ptr = nullptr; - float2 const* rotary_cos_sin_ptr = nullptr; - - if (op.isRoPE()) - { - if (rotary_inv_freq.has_value()) - { - rotary_inv_freq_ptr = rotary_inv_freq.value().data_ptr(); - } - if (rotary_cos_sin.has_value()) - { - rotary_cos_sin_ptr = static_cast(rotary_cos_sin.value().data_ptr()); - } - } + static constexpr Data_type data_type = DATA_TYPE_FP16; + static constexpr bool supported = true; +}; - void* workspace_ptr = workspace.data_ptr(); - [[maybe_unused]] MlaParams mla_params; - if (op.isMLAEnabled()) - { - if (is_context && op.mUseSparseAttention) - { - if (latent_cache.has_value()) - { - mla_params.latent_cache = static_cast(latent_cache->data_ptr()); - } - else - { - // kv cache reuse / chunked context cases, latent_cache is not used - mla_params.latent_cache = nullptr; - } - TORCH_CHECK(q_pe.has_value()); - TORCH_CHECK(q_pe->dim() == 3); - TORCH_CHECK(q_pe->strides()[2] == 1); - - mla_params.q_pe = static_cast(q_pe->data_ptr()); - mla_params.q_pe_ld = q_pe->strides()[1]; - mla_params.q_pe_stride = q_pe->strides()[0]; - - // Fused FP8-Q path: forward caller's quant_q_buffer / scale so - // applyMLARopeAndAssignQKVKernelOptContext - // appends rope FP8 in place and the standalone quantize is - // skipped. Without this wiring the sparse-MLA context branch - // runs the legacy quantize over the bf16 placeholder q. - mla_params.bmm1_scale = mla_bmm1_scale.has_value() - ? reinterpret_cast(mla_bmm1_scale.value().data_ptr()) - : nullptr; - mla_params.bmm2_scale = mla_bmm2_scale.has_value() - ? reinterpret_cast(mla_bmm2_scale.value().data_ptr()) - : nullptr; - mla_params.quant_q_buf - = quant_q_buffer.has_value() ? reinterpret_cast(quant_q_buffer.value().data_ptr()) : nullptr; - mla_params.quant_scale_qkv = quant_scale_qkv.has_value() - ? reinterpret_cast(quant_scale_qkv.value().data_ptr()) - : nullptr; - mla_params.fuse_q_fp8_in_rope = (quant_q_buffer.has_value() && quant_scale_qkv.has_value()); - - // Fused kv_a_layernorm: the norm weight implies `latent_cache` is the - // RAW kv_a_proj output, with the caller's RMSNorm and concat dropped. - if (kv_norm_weight.has_value()) - { - TORCH_CHECK(kv_norm_weight->is_cuda(), "kv_norm_weight must be a CUDA tensor"); - TORCH_CHECK(kv_norm_weight->is_contiguous(), "kv_norm_weight must be contiguous"); - TORCH_CHECK(kv_norm_weight->scalar_type() == qkv_or_q.scalar_type(), - "kv_norm_weight dtype must match the activation dtype"); - TORCH_CHECK(latent_cache.has_value(), - "fused kv-norm needs latent_cache (the raw kv_a_proj output) to be provided"); - // The kernel norms the whole latent row, so a narrower weight would - // read out of bounds. dsv3RopeOp.cpp checks the same on the - // generation side. - // `mla_params.meta` is not assigned until further below; read the op's copy. - auto const kv_norm_width = op.mMLAParams.kv_lora_rank + op.mMLAParams.qk_rope_head_dim; - TORCH_CHECK(kv_norm_weight->numel() == kv_norm_width, - "kv_norm_weight must span kv_lora_rank + qk_rope_head_dim (", kv_norm_width, "), got ", - kv_norm_weight->numel()); - // A last-dim slice, so rows are wider than the row itself. Forward - // the real stride; only the innermost dim must be unit-stride. - TORCH_CHECK(latent_cache->dim() == 2, "latent_cache must be 2D for fused kv-norm, got ", - latent_cache->dim(), "D"); - TORCH_CHECK(latent_cache->stride(1) == 1, "latent_cache must be unit-stride in its last dim"); - // The kernel walks rows with 16-byte vector loads, so a row start - // that is not 16B-aligned faults with a bare misaligned-address - // error far from here. - auto const kEltsPer16B = 16 / latent_cache->element_size(); - TORCH_CHECK(latent_cache->stride(0) % kEltsPer16B == 0, "latent_cache row stride (", - latent_cache->stride(0), ") must be a multiple of ", kEltsPer16B, - " for the fused kv-norm 16B vector loads"); - TORCH_CHECK(reinterpret_cast(latent_cache->data_ptr()) % 16 == 0, - "latent_cache must be 16B-aligned for the fused kv-norm vector loads"); - mla_params.latent_row_stride = static_cast(latent_cache->stride(0)); - mla_params.kv_norm_weight = static_cast(kv_norm_weight->data_ptr()); - mla_params.kv_norm_eps = static_cast(kv_norm_eps); - mla_params.fuse_kv_norm_in_rope = true; - } - } - else if (is_context) - { - if (latent_cache.has_value()) - { - mla_params.latent_cache = static_cast(latent_cache->data_ptr()); - } - else - { - // kv cache reuse / chunked context cases, latent_cache is not used - mla_params.latent_cache = nullptr; - } - TORCH_CHECK(k.has_value()); - TORCH_CHECK(v.has_value()); - TORCH_CHECK(k->dim() == 2); - TORCH_CHECK(v->dim() == 2); - TORCH_CHECK(k->strides()[1] == 1); - TORCH_CHECK(v->strides()[1] == 1); - - k_ptr = static_cast(k->slice(0, token_offset).data_ptr()); - v_ptr = static_cast(v->slice(0, token_offset).data_ptr()); - mla_params.k_buf = k_ptr; - mla_params.v_buf = v_ptr; - - // For generation, helix position is in ropeOp - if (helix_position_offsets.has_value()) - { - mla_params.helix_position_offsets = helix_position_offsets->data_ptr(); - } - if (helix_is_inactive_rank.has_value()) - { - mla_params.helix_is_inactive_rank = helix_is_inactive_rank->data_ptr(); - } - } - else - { - TORCH_CHECK(latent_cache.has_value()); - mla_params.latent_cache = static_cast(latent_cache->data_ptr()); - TORCH_CHECK(q_pe.has_value()); - TORCH_CHECK(q_pe->dim() == 3); - TORCH_CHECK(q_pe->strides()[2] == 1); - - mla_params.q_pe = static_cast(q_pe->data_ptr()); - mla_params.q_pe_ld = q_pe->strides()[1]; - mla_params.q_pe_stride = q_pe->strides()[0]; - - mla_params.seqQOffset - = cu_q_seqlens.has_value() ? reinterpret_cast(cu_q_seqlens.value().data_ptr()) : nullptr; - mla_params.cu_kv_seqlens - = cu_kv_seqlens.has_value() ? reinterpret_cast(cu_kv_seqlens.value().data_ptr()) : nullptr; - mla_params.fmha_tile_counter = fmha_scheduler_counter.has_value() - ? reinterpret_cast(fmha_scheduler_counter.value().data_ptr()) - : nullptr; - mla_params.bmm1_scale = mla_bmm1_scale.has_value() - ? reinterpret_cast(mla_bmm1_scale.value().data_ptr()) - : nullptr; - mla_params.bmm2_scale = mla_bmm2_scale.has_value() - ? reinterpret_cast(mla_bmm2_scale.value().data_ptr()) - : nullptr; - mla_params.quant_q_buf - = quant_q_buffer.has_value() ? reinterpret_cast(quant_q_buffer.value().data_ptr()) : nullptr; - mla_params.quant_scale_qkv = quant_scale_qkv.has_value() - ? reinterpret_cast(quant_scale_qkv.value().data_ptr()) - : nullptr; - // Request the fused FP8-Q path; common/attentionOp.cpp gates the - // actual skip on FP8 KV cache + absorption mode. - mla_params.fuse_q_fp8_in_rope = (quant_q_buffer.has_value() && quant_scale_qkv.has_value()); - } - mla_params.q_buf = attention_input; - mla_params.context_buf = reinterpret_cast(context_buf); +template <> +struct ConvertMMHAToXQAParamsHelper<__half, KVBlockArray> +{ + static constexpr Data_type data_type = DATA_TYPE_FP16; + static constexpr bool supported = true; +}; - mla_params.cos_sin_cache = rotary_cos_sin_ptr; - if (enable_dsv4_epilogue_fusion) - { - TORCH_CHECK(dsv4_inv_rope_cos_sin_cache.has_value(), - "DSv4 fused epilogue requires inverse-RoPE cos/sin cache."); - auto const& cos_sin_cache = dsv4_inv_rope_cos_sin_cache.value(); - auto const& output_sf_tensor = output_sf.value(); - TORCH_CHECK(cos_sin_cache.scalar_type() == torch::kFloat32, - "DSv4 fused epilogue cos/sin cache must be float32."); - TORCH_CHECK( - output.scalar_type() == torch::kFloat8_e4m3fn, "DSv4 fused epilogue output must be float8_e4m3fn."); - TORCH_CHECK(output.dim() == 3 && output.is_contiguous(), - "DSv4 fused epilogue output must be contiguous [groups, tokens, K]."); - TORCH_CHECK(output_sf_tensor.scalar_type() == torch::kFloat32, - "DSv4 fused epilogue output_sf must be float32."); - TORCH_CHECK(output_sf_tensor.dim() == 3 && output_sf_tensor.is_contiguous(), - "DSv4 fused epilogue output_sf must be contiguous [groups, K/128, padded_tokens]."); - TORCH_CHECK(output.size(1) >= num_tokens, "DSv4 fused epilogue output token dimension is too small."); - TORCH_CHECK(op.mMLAParams.v_head_dim > 0 && op.mMLAParams.v_head_dim % 128 == 0, - "DSv4 fused epilogue requires v_head_dim to be a positive multiple of 128."); - TORCH_CHECK(output_sf_tensor.size(2) >= num_tokens, - "DSv4 fused epilogue output_sf token dimension is too small."); - - mla_params.dsv4_epilogue_fusion.enabled = true; - mla_params.dsv4_epilogue_fusion.cos_sin_cache = static_cast(cos_sin_cache.data_ptr()); - mla_params.dsv4_epilogue_fusion.scale_buf_m = static_cast(output_sf_tensor.size(2)); - } - mla_params.batch_size = num_seqs; - mla_params.acc_q_len = num_tokens; - mla_params.head_num = op.mNumHeads; - mla_params.meta = op.mMLAParams; - - mla_params.workspace = workspace_ptr; - } - // Extract K/V pointers for sage attention (separate Q/K/V inputs). - else if (is_context - && (op.mSageAttnNumEltsPerBlkQ > 0 || op.mSageAttnNumEltsPerBlkK > 0 || op.mSageAttnNumEltsPerBlkV > 0)) - { - TORCH_CHECK(k.has_value() && v.has_value(), "SageAttention demands separate K and V buffers"); - k_ptr = static_cast(k->slice(0, token_offset).data_ptr()); - v_ptr = static_cast(v->slice(0, token_offset).data_ptr()); - } - - int const* context_lengths_ptr = context_lengths.slice(0, seq_offset).data_ptr(); - int const* sequence_lengths_ptr = sequence_length.slice(0, seq_offset).data_ptr(); - // Note we still need context length during generation for MMHA optimization. - // For encoder CUDA graphs compatibility, allow the caller to override the - // max context Q length so FMHA kernel launch params (mMaxSeqLenQ-driven grid - // and cluster dims) are stable across graph replays even when actual per-batch - // sequence lengths vary. - int32_t const max_context_q_len_computed - = host_context_lengths.slice(0, seq_offset, seq_offset + num_seqs).max().item(); - int32_t const max_past_kv_length_computed - = host_past_key_value_lengths.slice(0, seq_offset, seq_offset + num_seqs).max().item(); - - if (max_context_q_len_override.has_value()) - { - int32_t const override_value = static_cast(max_context_q_len_override.value()); - TORCH_CHECK(override_value >= max_context_q_len_computed, - "max_context_q_len_override (%d) must be >= computed max context q length (%d).", override_value, - max_context_q_len_computed); - TORCH_CHECK(override_value >= max_past_kv_length_computed, - "max_context_q_len_override (%d) must be >= computed max past kv length (%d).", override_value, - max_past_kv_length_computed); - } - - int32_t const max_context_q_len = max_context_q_len_override.has_value() - ? static_cast(max_context_q_len_override.value()) - : max_context_q_len_computed; - // Override the max_past_kv_length as well for encoder CUDA graph compatibility - int32_t const max_past_kv_length = max_context_q_len_override.has_value() - ? static_cast(max_context_q_len_override.value()) - : max_past_kv_length_computed; - - // Commonly, cyclic_attention_window_size, and max_attention_window_size will be the same - // unless each layer has different attention window sizes. - int const max_attention_window_size = beam_width == 1 ? attention_window_size - : cache_indirection.has_value() ? cache_indirection.value().size(2) - : attention_window_size; - // The cyclic_attention_window_size will determine the cyclic kv cache position of new tokens. - // Note that this cyclic_attention_window_size might be smaller than the actual kv cache capactity. - int const cyclic_attention_window_size = attention_window_size; - bool const can_use_one_more_block = beam_width > 1; - - int max_blocks_per_sequence = 0; - int32_t pool_index = 0; - int32_t layer_idx_in_cache_pool = 0; - KVBlockArray::DataType* block_offsets = nullptr; - bool use_kv_cache = false; - KvCachePoolPointers pool_pointers; - max_blocks_per_sequence - = op.useKVCache() && kv_cache_block_offsets.has_value() ? kv_cache_block_offsets.value().size(-1) : 0; - pool_index = op.useKVCache() && host_kv_cache_pool_mapping.has_value() - ? host_kv_cache_pool_mapping.value().index({op.mLayerIdx, 0}).item() - : 0; - layer_idx_in_cache_pool = op.useKVCache() && host_kv_cache_pool_mapping.has_value() - ? host_kv_cache_pool_mapping.value().index({op.mLayerIdx, 1}).item() - : 0; - block_offsets = static_cast(op.useKVCache() && kv_cache_block_offsets.has_value() - ? kv_cache_block_offsets.value().index({pool_index, seq_offset}).data_ptr() - : nullptr); +#ifdef ENABLE_BF16 +template <> +struct ConvertMMHAToXQAParamsHelper<__nv_bfloat16, KVLinearBuffer> +{ + static constexpr Data_type data_type = DATA_TYPE_BF16; + static constexpr bool supported = true; +}; - // The cache element size in bits. - int cache_elem_bits = op.getKvCacheElemSizeInBits(); - auto const block_size = op.mTokensPerBlock * op.mNumKVHeads * op.mHeadSize; - auto const bytes_per_block = block_size * cache_elem_bits / 8 /*bits*/; - int32_t const kv_factor = op.isMLAEnabled() ? 1 : 2; - auto const intra_pool_offset = layer_idx_in_cache_pool * kv_factor * bytes_per_block; +template <> +struct ConvertMMHAToXQAParamsHelper<__nv_bfloat16, KVBlockArray> +{ + static constexpr Data_type data_type = DATA_TYPE_BF16; + static constexpr bool supported = true; +}; +#endif - // Build KV cache pool pointers from the host tensor. - use_kv_cache = op.useKVCache() && host_kv_cache_pool_pointers.has_value(); - if (use_kv_cache) +template +bool AttentionOp::convertMMHAParamsToXQAParams( + tensorrt_llm::kernels::XQAParams& xqaParams, FmhaParams const& p, bool forConfigurePlugin) +{ + bool retval = ConvertMMHAToXQAParamsHelper::supported; + if (!retval) + { + return false; + } + xqaParams = {}; + xqaParams.data_type = ConvertMMHAToXQAParamsHelper::data_type; + + xqaParams.num_q_heads = mNumAttnHeads; + xqaParams.num_kv_heads = mNumAttnKVHeads; + xqaParams.head_size = mHeadSize; + xqaParams.unidirectional = p.unidirectional; + xqaParams.q_scaling = mCfg.q_scaling; + xqaParams.rotary_embedding_dim = mCfg.rotary_embedding_dim; + xqaParams.rotary_embedding_base = p.rotary_embedding_base; + xqaParams.rotary_embedding_scale_type = p.rotary_embedding_scale_type; + xqaParams.rotary_embedding_scale = p.rotary_embedding_scale; + xqaParams.rotary_embedding_max_positions = p.rotary_embedding_max_positions; + xqaParams.rotary_vision_start = p.vision_start; + xqaParams.rotary_vision_length = p.vision_length; + xqaParams.rotary_cos_sin = p.getRotaryCosSin(); + xqaParams.position_embedding_type = mCfg.position_embedding_type; + xqaParams.position_shift_enabled = p.pos_shift_enabled; + xqaParams.remove_padding = mCfg.remove_padding; + xqaParams.mask_type = mCfg.mask_type; + xqaParams.paged_kv_cache = mPagedKVCache; + xqaParams.tokens_per_block = mCfg.tokens_per_block; + xqaParams.kv_cache_quant_mode = mCfg.quant_mode; + xqaParams.tp_size = 1; + xqaParams.tp_rank = 0; + xqaParams.qkv_bias_enabled = p.qkv_bias_enabled; + xqaParams.cross_attention = p.is_cross; + xqaParams.max_distance = static_cast(p.fwd.relative_attention_max_distance); + xqaParams.multi_block_mode = common::getEnvForceDeterministicAttention() ? false : mMultiBlockMode; + // Medusa mode will have multiple query tokens. + xqaParams.multi_query_tokens = mCfg.is_spec_decoding_enabled && p.use_spec_decoding; + xqaParams.is_spec_dec_tree = p.is_spec_dec_tree; + xqaParams.force_prepare_spec_dec_tree_mask = p.force_prepare_spec_dec_tree_mask; + xqaParams.layer_idx = static_cast(p.local_layer_idx); + + if (mCfg.quant_mode.hasInt8KvCache()) + { + xqaParams.kv_cache_data_type = DATA_TYPE_INT8; + } + else if (mCfg.quant_mode.hasFp8KvCache()) + { + // Inputs to MLA is FP8 instead of BF16/FP16 when using FP8 KV cache. + if (xqaParams.isMLA()) { - pool_pointers = buildKvCachePoolPointers(host_kv_cache_pool_pointers.value(), pool_index, intra_pool_offset, - block_size, layer_idx_in_cache_pool, kv_factor, op.mKVCacheQuantMode.hasFp4KvCache()); + xqaParams.data_type = DATA_TYPE_E4M3; } + xqaParams.kv_cache_data_type = DATA_TYPE_E4M3; + } + else if (mCfg.quant_mode.hasFp4KvCache()) + { + xqaParams.kv_cache_data_type = DATA_TYPE_E2M1; + } + else + { + xqaParams.kv_cache_data_type = xqaParams.data_type; + } + if (xqaParams.kv_cache_data_type == DATA_TYPE_INT8 + || (xqaParams.kv_cache_data_type == DATA_TYPE_E4M3 && (mSM < kSM_90 || mSM > kSM_120))) + { + xqaParams.multi_block_mode = false; + } - float const* kv_scale_orig_quant_ptr = nullptr; - float const* kv_scale_quant_orig_ptr = nullptr; - if (op.mKVCacheQuantMode.hasKvCacheQuant() && kv_scale_orig_quant.has_value() - && kv_scale_quant_orig.has_value()) - { - if (op.mKVCacheQuantMode.hasFp4KvCache()) - { - if (op.isMLAEnabled()) - { - auto const& origQuantScale = kv_scale_orig_quant.value(); - auto const& quantOrigScale = kv_scale_quant_orig.value(); - TORCH_CHECK(origQuantScale.scalar_type() == torch::kFloat32, - "kv_scale_orig_quant must have float32 dtype for MLA with FP4 KV cache"); - TORCH_CHECK(quantOrigScale.scalar_type() == torch::kFloat32, - "kv_scale_quant_orig must have float32 dtype for MLA with FP4 KV cache"); - TORCH_CHECK(origQuantScale.is_contiguous(), - "kv_scale_orig_quant must be contiguous for MLA with FP4 KV cache"); - TORCH_CHECK(quantOrigScale.is_contiguous(), - "kv_scale_quant_orig must be contiguous for MLA with FP4 KV cache"); - TORCH_CHECK(origQuantScale.dim() == 1 && origQuantScale.size(0) == 1, - "kv_scale_orig_quant must have shape [1] for MLA with FP4 KV cache"); - TORCH_CHECK(quantOrigScale.dim() == 1 && quantOrigScale.size(0) == 1, - "kv_scale_quant_orig must have shape [1] for MLA with FP4 KV cache"); - } - else - { - TORCH_CHECK(kv_scale_orig_quant.value().size(0) == 3); - TORCH_CHECK(kv_scale_quant_orig.value().size(0) == 3); - } - } - kv_scale_orig_quant_ptr = kv_scale_orig_quant.value().data_ptr(); - kv_scale_quant_orig_ptr = kv_scale_quant_orig.value().data_ptr(); - } - // For FP8 output, out_scale represents the output scale. - float const* out_scale_ptr = (op.mFP8ContextFMHA && !op.mFuseFp4Quant && out_scale.has_value()) - ? out_scale.value().data_ptr() - : nullptr; - // For NVFP4 output, out_scale holds the global scale for scaling factors. - float const* out_sf_scale_ptr - = op.mFuseFp4Quant && out_scale.has_value() ? out_scale.value().data_ptr() : nullptr; + xqaParams.output = p.getOutput(); + xqaParams.qkv = p.getQkvOrQ(); + xqaParams.cache_indir = p.getCacheIndirection(); + xqaParams.attention_sinks = p.getAttentionSinks(); + xqaParams.kv_scale_orig_quant = p.getKvScaleOrigQuant(); + xqaParams.kv_scale_quant_orig = p.getKvScaleQuantOrig(); + xqaParams.host_past_key_value_lengths = p.getHostPastKeyValueLengths(); + xqaParams.host_context_lengths = p.getHostContextLengths(); + xqaParams.semaphores = static_cast(p.getMultiCtasKvCounter()); + xqaParams.workspaces = p.getWorkspace(); + xqaParams.batch_size = p.num_requests; + xqaParams.beam_width = p.beam_width; + // Speculative decoding mode has generation input_length > 1. + xqaParams.generation_input_length = p.input_seq_length; + xqaParams.chunked_attention_size + = p.attention_chunk_size && !tc::getEnvDisableChunkedAttentionInGenPhase() ? *p.attention_chunk_size : INT_MAX; + xqaParams.max_attention_window_size = p.max_attention_window_size; + xqaParams.cyclic_attention_window_size = p.cyclic_attention_window_size; + xqaParams.max_blocks_per_sequence = p.max_blocks_per_sequence; + xqaParams.sink_token_length = p.sink_token_length; + xqaParams.max_past_kv_length = p.max_past_kv_length; + xqaParams.qkv_bias = p.getQkvBias(); + xqaParams.sequence_lengths = p.getSequenceLength(); + xqaParams.context_lengths = p.getContextLengths(); + xqaParams.alibi_slopes = p.getAlibiSlopes(); + // Pre-computed rotary inv freq when building the engines. + xqaParams.rotary_embedding_inv_freq_cache = p.getRotaryInvFreq(); + if (!forConfigurePlugin) + { + // Speculative decoding (need to take new generated ids into consideration). + TLLM_CHECK_WITH_INFO( + !(mCfg.is_spec_decoding_enabled && p.use_spec_decoding) || p.getSpecDecodingPackedMask() != nullptr, + "Speculative decoding mode needs a valid packed_mask input tensor."); + } + xqaParams.spec_decoding_packed_mask = p.getSpecDecodingPackedMask(); + xqaParams.spec_decoding_position_offsets = p.getSpecDecodingPositionOffsets(); + xqaParams.spec_decoding_generation_lengths = p.getSpecDecodingGenerationLengths(); + xqaParams.spec_decoding_is_generation_length_variable = p.spec_decoding_is_generation_length_variable; + xqaParams.spec_decoding_max_generation_length = p.spec_decoding_max_generation_length; + xqaParams.spec_decoding_bl_tree_mask_offset = p.getSpecDecodingBlTreeMaskOffset(); + xqaParams.spec_decoding_bl_tree_mask = p.getSpecDecodingBlTreeMask(); + xqaParams.spec_bl_tree_first_sparse_mask_offset_kv = p.getSpecBlTreeFirstSparseMaskOffsetKv(); + xqaParams.mrope_position_deltas = p.getMropePositionDeltas(); + xqaParams.helix_position_offsets = p.getHelixPositionOffsets(); + xqaParams.helix_is_inactive_rank = p.getHelixIsInactiveRank(); + xqaParams.softmax_stats = p.getSoftmaxStatsTensor(); + xqaParams.trtllm_gen_jit_warmup = p.trtllm_gen_jit_warmup; + xqaParams.trtllm_gen_jit_warmup_max_num_requests = p.max_num_requests; + xqaParams.trtllm_gen_jit_warmup_max_seq_len_q = p.max_context_length; + xqaParams.trtllm_gen_jit_warmup_max_seq_len_kv = p.max_seq_len; + + xqaParams.logn_scaling_ptr = p.getLognScalingPtr(); + xqaParams.total_num_input_tokens = p.num_tokens; + xqaParams.is_fp8_output = mFP8AttenOutput; + xqaParams.fp8_out_scale = ((mFP8AttenOutput) ? p.getOutScale() : nullptr); + // Parameters required for FP4 output. + xqaParams.output_sf = p.getOutputSf(); + xqaParams.fp4_out_sf_scale = p.getOutSfScale(); + xqaParams.start_token_idx_sf = p.token_offset; + // Parameters for sparse attention + xqaParams.sparse_params = p.sparse_params; + xqaParams.use_sparse_attention_gen_paged = useTllmGenSparseAttentionPaged(p); + // Skip softmax threshold. + xqaParams.skip_softmax_threshold_scale_factor + = static_cast(p.fwd.sparse_runtime_params.threshold_scale_factor_decode); +#ifdef SKIP_SOFTMAX_STAT + // Statistics of skip-softmax, pointers of device memory for output + xqaParams.skip_softmax_total_blocks = mSkipSoftmaxTotalBlocks; + xqaParams.skip_softmax_skipped_blocks = mSkipSoftmaxSkippedBlocks; +#endif + // Cross attention parameters. + xqaParams.encoder_input_lengths = p.getEncoderInputLengths(); - // The attention_sinks is a float tensor with shape [num_heads_q] - float const* attention_sinks_ptr = nullptr; - if (attention_sinks.has_value()) - { - TORCH_CHECK( - attention_sinks.value().dtype() == torch::kFloat32, "Expected attention_sinks to have float dtype"); - attention_sinks_ptr = attention_sinks.value().data_ptr(); - } - T const* relative_attention_bias_ptr = nullptr; - int relative_attention_bias_stride = 0; - if (relative_attention_bias.has_value()) - { - auto const& relative_attention_bias_tensor = relative_attention_bias.value(); - TORCH_CHECK(relative_attention_bias_tensor.dim() == 2 || relative_attention_bias_tensor.dim() == 3, - "relative_attention_bias must be [num_heads, num_buckets] for implicit mode or " - "[num_heads, max_seq_len, max_seq_len] for explicit mode"); - TORCH_CHECK(relative_attention_bias_tensor.is_contiguous(), "relative_attention_bias must be contiguous"); - TORCH_CHECK(relative_attention_bias_tensor.scalar_type() == qkv_or_q.scalar_type(), - "relative_attention_bias dtype must match attention input dtype"); - relative_attention_bias_ptr = static_cast(relative_attention_bias_tensor.data_ptr()); - relative_attention_bias_stride = static_cast(relative_attention_bias_tensor.size(1)); - } - - // Prepare sparse attention parameters - op.mRuntimeSparseAttentionParams.sparse_kv_indices - = sparse_kv_indices.has_value() ? sparse_kv_indices.value().data_ptr() : nullptr; - op.mRuntimeSparseAttentionParams.sparse_kv_offsets - = sparse_kv_offsets.has_value() ? sparse_kv_offsets.value().data_ptr() : nullptr; - op.mRuntimeSparseAttentionParams.sparse_attn_indices - = sparse_attn_indices.has_value() ? sparse_attn_indices.value().data_ptr() : nullptr; - op.mRuntimeSparseAttentionParams.sparse_attn_offsets - = sparse_attn_offsets.has_value() ? sparse_attn_offsets.value().data_ptr() : nullptr; - op.mRuntimeSparseAttentionParams.sparse_attn_indices_block_size = sparse_attn_indices_block_size; - op.mRuntimeSparseAttentionParams.sparse_attn_indices_stride - = sparse_attn_indices.has_value() ? sparse_attn_indices.value().size(-1) : 0; - op.mRuntimeSparseAttentionParams.num_sparse_topk = num_sparse_topk; - op.mRuntimeSparseAttentionParams.sparse_attn_kv_lens - = sparse_attn_kv_lens.has_value() ? sparse_attn_kv_lens.value().data_ptr() : nullptr; - op.mRuntimeSparseAttentionParams.sparse_kv_cache_pool = nullptr; - op.mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool = nullptr; - if (op.mUseSparseAttention && use_kv_cache) - { - if (host_kv_cache_pool_pointers.has_value()) - { - auto* kvCachePool = reinterpret_cast(host_kv_cache_pool_pointers.value().dim() == 3 - ? host_kv_cache_pool_pointers.value().index({pool_index, 0, 0}).item() - : host_kv_cache_pool_pointers.value().index({pool_index, 0}).item()); - if (sparse_attn_kv_lens.has_value()) - { - // Deepseek V4 dynamic sparse MLA always uses the SWA pool for now. - op.mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool = kvCachePool; - if (aux_kv_cache_pool_ptr.has_value()) - { - op.mRuntimeSparseAttentionParams.sparse_kv_cache_pool - = reinterpret_cast(aux_kv_cache_pool_ptr.value()); - } - } - else - { - op.mRuntimeSparseAttentionParams.sparse_kv_cache_pool = aux_kv_cache_pool_ptr.has_value() - ? reinterpret_cast(aux_kv_cache_pool_ptr.value()) - : kvCachePool; - } - } - } + return true; +} - AttentionOp::EnqueueParams common_enqueue_params; - common_enqueue_params.attention_input = attention_input; - common_enqueue_params.attention_sinks = attention_sinks_ptr; - common_enqueue_params.rotary_inv_freq = rotary_inv_freq_ptr; - common_enqueue_params.rotary_cos_sin = rotary_cos_sin_ptr; - common_enqueue_params.relative_attention_bias = relative_attention_bias_ptr; - common_enqueue_params.relative_attention_bias_stride = relative_attention_bias_stride; - common_enqueue_params.max_past_kv_length = max_past_kv_length; - common_enqueue_params.max_attention_window_size = max_attention_window_size; - common_enqueue_params.cyclic_attention_window_size = cyclic_attention_window_size; - common_enqueue_params.max_cyclic_attention_window_size = cyclic_attention_window_size; - common_enqueue_params.can_use_one_more_block = can_use_one_more_block; - common_enqueue_params.kv_scale_orig_quant = kv_scale_orig_quant_ptr; - common_enqueue_params.kv_scale_quant_orig = kv_scale_quant_orig_ptr; - common_enqueue_params.attention_output_orig_quant = out_scale_ptr; - common_enqueue_params.attention_output_sf_scale = out_sf_scale_ptr; - common_enqueue_params.context_buf = context_buf; - common_enqueue_params.context_buf_sf = context_buf_sf; - common_enqueue_params.block_offsets = block_offsets; - common_enqueue_params.host_primary_pool_pointer = pool_pointers.primaryPoolPtr; - common_enqueue_params.host_secondary_pool_pointer = pool_pointers.secondaryPoolPtr; - common_enqueue_params.host_primary_block_scale_pool_pointer = pool_pointers.primaryBlockScalePoolPtr; - common_enqueue_params.host_secondary_block_scale_pool_pointer = pool_pointers.secondaryBlockScalePoolPtr; - common_enqueue_params.num_tokens = num_tokens; - common_enqueue_params.total_kv_len = total_kv_len; - common_enqueue_params.max_blocks_per_sequence = max_blocks_per_sequence; - common_enqueue_params.sequence_lengths = sequence_lengths_ptr; - common_enqueue_params.context_lengths = context_lengths_ptr; - common_enqueue_params.host_context_lengths = host_context_lengths.data_ptr(); - common_enqueue_params.workspace = workspace_ptr; - common_enqueue_params.trtllm_gen_jit_warmup = trtllm_gen_jit_warmup; - if (is_cross) - { - // For cross attention, the KV (encoder) sequence lengths are passed in via - // `sequence_length` (already sliced into `sequence_lengths_ptr`), so reuse - // it directly instead of a redundant `encoder_input_lengths` tensor. - common_enqueue_params.encoder_input_lengths = sequence_lengths_ptr; - } - if (softmax_stats_tensor.has_value()) - { - TLLM_CHECK_WITH_INFO(softmax_stats_tensor.value().scalar_type() == at::ScalarType::Float, - "softmax_stats_tensor must have float type"); - TLLM_CHECK_WITH_INFO(softmax_stats_tensor.value().size(0) >= num_tokens, - "softmax_stats_tensor must have first dimension >= num_tokens"); - TLLM_CHECK_WITH_INFO(softmax_stats_tensor.value().size(1) >= op.mNumHeads, - "softmax_stats_tensor must have second dimension >= num_heads"); - TLLM_CHECK_WITH_INFO( - softmax_stats_tensor.value().size(2) == 2, "softmax_stats_tensor must have third dimension == 2"); - common_enqueue_params.softmax_stats = static_cast(softmax_stats_tensor.value().data_ptr()); - } +template +void fusedQKV_masked_attention_dispatch(Multihead_attention_params& params, + FusedQKVMaskedAttentionDispatchParams const& input_params, cudaStream_t stream) +{ + using DataType = typename SATypeConverter::Type; - // Shared helper to wire helix params into the enqueue params. - // Works for both EnqueueContextParams and EnqueueGenerationParams since both have - // helix_position_offsets and helix_is_inactive_rank fields. - auto const extractHelixParams = [&helix_position_offsets, &helix_is_inactive_rank](auto& params) - { - if (helix_position_offsets.has_value()) - { - params.helix_position_offsets = helix_position_offsets->data_ptr(); - } - if (helix_is_inactive_rank.has_value()) - { - params.helix_is_inactive_rank = helix_is_inactive_rank->data_ptr(); - } - }; + // Prepare the parameters. + params = {}; - if (is_context) // context stage - { - common_enqueue_params.input_seq_length = max_context_q_len; - AttentionOp::EnqueueContextParams enqueue_params{common_enqueue_params}; - enqueue_params.batch_size = num_seqs; - enqueue_params.k_ptr = k_ptr; - enqueue_params.v_ptr = v_ptr; - if (cu_q_seqlens.has_value()) - { - TORCH_CHECK(cu_q_seqlens->dim() == 1, "cu_q_seqlens must be a 1-D tensor."); - TORCH_CHECK(cu_q_seqlens->is_cuda(), "cu_q_seqlens must be a CUDA tensor."); - TORCH_CHECK(cu_q_seqlens->scalar_type() == at::ScalarType::Int, "cu_q_seqlens must be int32."); - TORCH_CHECK( - cu_q_seqlens->size(0) >= num_seqs + 1, "cu_q_seqlens must have at least num_seqs + 1 elements."); - enqueue_params.cu_q_seqlens = cu_q_seqlens->data_ptr(); - } - if (cu_kv_seqlens.has_value()) - { - TORCH_CHECK(cu_kv_seqlens->dim() == 1, "cu_kv_seqlens must be a 1-D tensor."); - TORCH_CHECK(cu_kv_seqlens->is_cuda(), "cu_kv_seqlens must be a CUDA tensor."); - TORCH_CHECK(cu_kv_seqlens->scalar_type() == at::ScalarType::Int, "cu_kv_seqlens must be int32."); - TORCH_CHECK( - cu_kv_seqlens->size(0) >= num_seqs + 1, "cu_kv_seqlens must have at least num_seqs + 1 elements."); - enqueue_params.cu_kv_seqlens = cu_kv_seqlens->data_ptr(); - } - if (is_cross && cross_kv.has_value()) - { - auto const& cross_kv_tensor = cross_kv.value(); - enqueue_params.cross_kv = static_cast(cross_kv_tensor.data_ptr()); - enqueue_params.num_encoder_tokens = static_cast(cross_kv_tensor.size(0)); - // Kept in step with maxCrossKvLength in attention(), which sizes the workspace carved here. - enqueue_params.cross_kv_length - = host_past_key_value_lengths.slice(0, seq_offset, seq_offset + num_seqs).max().item(); - } - else if (is_cross) - { - // Later chunks of a chunked decoder prefill carry no encoder K/V and must read the cross KV cache. - // Only the fused kernel can do that; the unfused cross path builds K and V from cross_kv and would - // read a null pointer here. get_attention_op exempts cross attention from its paged-context check - // for exactly this reason, so this is the one place the requirement is enforced. - TLLM_CHECK_WITH_INFO(!op.isUnfusedCrossAttention(), - "Cross attention without encoder K/V input requires a fused context FMHA kernel to read the " - "cached cross KV, and this build has none for this configuration. Disable chunked prefill for " - "this model, or use a build whose --cuda_architectures includes this device's SM."); - } + int hidden_units = input_params.head_num * input_params.size_per_head; + int hidden_units_kv = input_params.kv_head_num * input_params.size_per_head; + if (input_params.qkv_bias != nullptr) + { + params.q_bias = reinterpret_cast(input_params.qkv_bias); + params.k_bias = reinterpret_cast(input_params.qkv_bias) + hidden_units; + params.v_bias = reinterpret_cast(input_params.qkv_bias) + hidden_units + hidden_units_kv; + } + else + { + params.q_bias = nullptr; + params.k_bias = nullptr; + params.v_bias = nullptr; + } - if (op.isMLAEnabled()) - { - mla_params.cache_seq_lens = sequence_lengths_ptr; - mla_params.max_input_seq_len = max_context_q_len; - enqueue_params.mla_param = &mla_params; - } - if (op.isMRoPE() && mrope_rotary_cos_sin.has_value()) - { - enqueue_params.mrope_rotary_cos_sin - = static_cast(mrope_rotary_cos_sin.value().data_ptr()); - } - extractHelixParams(enqueue_params); - op.enqueueContext(enqueue_params, stream); - } - else // generation stage - { - int32_t const batch_beam = num_seqs; - TLLM_CHECK(batch_beam % beam_width == 0); - int32_t const num_requests = batch_beam / beam_width; - - TLLM_CHECK_WITH_INFO(num_tokens % num_seqs == 0, - "seq_len should be same for all generation requests, num_tokens=%d, num_seqs=%d", num_tokens, num_seqs); - int32_t const input_seq_length = num_tokens / num_seqs; - - common_enqueue_params.input_seq_length = input_seq_length; - AttentionOp::EnqueueGenerationParams enqueue_params{common_enqueue_params}; - enqueue_params.layer_idx = op.mLayerIdx; - enqueue_params.beam_width = beam_width; - enqueue_params.num_requests = num_requests; - enqueue_params.cache_indir = beam_width == 1 - ? nullptr - : (cache_indirection.has_value() ? cache_indirection.value().data_ptr() : nullptr); - enqueue_params.semaphores = op.multiBlockSemaphores(); - enqueue_params.host_past_key_value_lengths = host_past_key_value_lengths.data_ptr(); - enqueue_params.start_token_idx_sf = token_offset; - - if (op.isMRoPE() && mrope_position_deltas.has_value()) - { - enqueue_params.mrope_position_deltas = mrope_position_deltas.value().data_ptr(); - } - if (op.mIsSpecDecodingEnabled && op.mUseSpecDecoding) - { - bool useTllmGen = tensorrt_llm::common::isSM100Family(); - TORCH_CHECK(spec_decoding_generation_lengths.has_value(), - "Expecting spec_decoding_generation_lengths in spec-dec mode."); - TORCH_CHECK(spec_decoding_position_offsets_for_cpp.has_value(), - "Expecting spec_decoding_position_offsets_for_cpp in spec-dec mode."); - TORCH_CHECK( - spec_decoding_packed_mask.has_value(), "Expecting spec_decoding_packed_mask in spec-dec mode."); - if (useTllmGen) - { - TORCH_CHECK(spec_decoding_bl_tree_mask_offset.has_value(), - "Expecting spec_decoding_bl_tree_mask_offset in trtllm-gen spec-dec mode."); - TORCH_CHECK(spec_decoding_bl_tree_mask.has_value(), - "Expecting spec_decoding_bl_tree_mask in trtllm-gen spec-dec mode."); - TORCH_CHECK(spec_bl_tree_first_sparse_mask_offset_kv.has_value(), - "Expecting spec_bl_tree_first_sparse_mask_offset_kv in trtllm-gen spec-dec mode."); - enqueue_params.spec_decoding_bl_tree_mask_offset - = spec_decoding_bl_tree_mask_offset->data_ptr(); - enqueue_params.spec_decoding_bl_tree_mask = spec_decoding_bl_tree_mask->data_ptr(); - enqueue_params.spec_bl_tree_first_sparse_mask_offset_kv - = spec_bl_tree_first_sparse_mask_offset_kv->data_ptr(); - } - enqueue_params.spec_decoding_generation_lengths = spec_decoding_generation_lengths->data_ptr(); - enqueue_params.spec_decoding_position_offsets - = spec_decoding_position_offsets_for_cpp->data_ptr(); - enqueue_params.spec_decoding_packed_mask = spec_decoding_packed_mask->data_ptr(); - enqueue_params.spec_decoding_is_generation_length_variable = true; - TLLM_CHECK(spec_decoding_position_offsets_for_cpp->dim() == 2); // [batch_size, max_draft_len + 1] - if (useTllmGen) - { - // Blackwell uses the padded packed-mask row dim as the mask stride. - TLLM_CHECK(spec_decoding_packed_mask->dim() == 3); - enqueue_params.spec_decoding_max_generation_length = spec_decoding_packed_mask->sizes()[1]; - } - else - { - enqueue_params.spec_decoding_max_generation_length - = spec_decoding_position_offsets_for_cpp->sizes()[1]; - } - } + // Set the output buffer. + params.out = input_params.context_buf; - // Current mlaGeneration will using fmha to do attention, so we don't go into enqueueGeneration - if (op.isMLAEnabled()) - { - if (op.mUseGenFlashMLA == true) - { - TORCH_CHECK(block_ids_per_seq.has_value()); - int const* block_ids_per_seq_ptr = static_cast(block_ids_per_seq->data_ptr()); - mla_params.block_ids_per_seq = block_ids_per_seq_ptr; - // Use pre-computed metadata if provided. - if (flash_mla_tile_scheduler_metadata.has_value()) - { - TORCH_CHECK(flash_mla_num_splits.has_value(), - "flash_mla_num_splits must be provided when flash_mla_tile_scheduler_metadata is set."); - mla_params.flash_mla_tile_scheduler_metadata - = flash_mla_tile_scheduler_metadata->data_ptr(); - mla_params.flash_mla_num_splits = flash_mla_num_splits->data_ptr(); - } - } - mla_params.cache_seq_lens = sequence_lengths_ptr; - { - op.mlaGeneration(mla_params, enqueue_params, stream); - } - } - else - { - extractHelixParams(enqueue_params); - { - op.enqueueGeneration(enqueue_params, stream); - } - } + // Set the input buffers. + params.q = reinterpret_cast(input_params.qkv_buf); + params.k = reinterpret_cast(input_params.qkv_buf) + hidden_units; + params.v = reinterpret_cast(input_params.qkv_buf) + hidden_units + hidden_units_kv; - { - std::string const afterGenStr = "gen attention at layer " + std::to_string(op.mLayerIdx); - { - TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(num_tokens, - output.size(1), op.mType, context_buf, stream, afterGenStr) - == false, - "Found invalid number (NaN or Inf) in " + afterGenStr); - } - } - } - sync_check_cuda_error(stream); + params.int8_kv_cache = input_params.kv_cache_quant_mode.hasInt8KvCache(); + params.fp8_kv_cache = input_params.kv_cache_quant_mode.hasFp8KvCache(); + if (input_params.kv_cache_quant_mode.hasKvCacheQuant()) + { + params.kv_scale_orig_quant = input_params.kv_scale_orig_quant; + params.kv_scale_quant_orig = input_params.kv_scale_quant_orig; } -}; -template class Runner; -template class Runner; -template class Runner; -#ifdef ENABLE_BF16 -template class Runner<__nv_bfloat16>; -template class Runner<__nv_bfloat16, __nv_fp8_e4m3>; -#endif + params.stride = hidden_units + 2 * hidden_units_kv; + params.finished = const_cast(input_params.finished); -} // namespace trtllm::attention + params.cache_indir = input_params.cache_indir; + params.batch_size = input_params.inference_batch_size; + params.beam_width = input_params.beam_width; + params.chunked_attention_size = input_params.chunked_attention_size; + if (input_params.chunked_attention_size != INT_MAX && !tc::getEnvDisableChunkedAttentionInGenPhase()) + { + TLLM_CHECK_WITH_INFO((input_params.chunked_attention_size & (input_params.chunked_attention_size - 1)) == 0, + "Attention chunk size should be a power of 2."); + params.chunked_attention_size_log2 = std::log2(input_params.chunked_attention_size); + } + else + { + params.chunked_attention_size_log2 = 0; + } + params.max_attention_window_size = input_params.max_attention_window_size; + params.cyclic_attention_window_size = input_params.cyclic_attention_window_size; + params.sink_token_length = input_params.sink_token_length; + params.length_per_sample = input_params.sequence_lengths; // max_input_length + current output length + // timestep for shared memory size calculation and rotary embedding computation + params.timestep = input_params.timestep; + params.num_heads = input_params.head_num; + params.num_kv_heads = input_params.kv_head_num; + params.hidden_size_per_head = input_params.size_per_head; + params.rotary_embedding_dim = input_params.rotary_embedding_dim; + params.rotary_embedding_base = input_params.rotary_embedding_base; + params.rotary_embedding_scale_type = input_params.rotary_embedding_scale_type; + params.rotary_embedding_scale = input_params.rotary_embedding_scale; + params.rotary_embedding_inv_freq_cache = input_params.rotary_embedding_inv_freq_cache; + params.rotary_embedding_cos_sin_cache = input_params.rotary_embedding_cos_sin_cache; + params.rotary_embedding_short_m_scale = input_params.rotary_embedding_short_m_scale; + params.rotary_embedding_long_m_scale = input_params.rotary_embedding_long_m_scale; + params.rotary_embedding_max_positions = input_params.rotary_embedding_max_positions; + params.rotary_embedding_original_max_positions = input_params.rotary_embedding_original_max_positions; + params.rotary_cogvlm_vision_start = input_params.rotary_cogvlm_vision_start; + params.rotary_cogvlm_vision_length = input_params.rotary_cogvlm_vision_length; + params.position_embedding_type = input_params.position_embedding_type; + params.position_shift_enabled = input_params.position_shift_enabled; + // Note: keep norm factor (sqrt(K_dim)) when adopting megatron T5 structure (may adjust) + params.inv_sqrt_dh = 1.F / (sqrtf((float) params.hidden_size_per_head) * input_params.q_scaling); + params.attn_logit_softcapping_scale = input_params.attn_logit_softcapping_scale; + params.attn_logit_softcapping_inverse_scale = 1.0f / input_params.attn_logit_softcapping_scale; + + params.logn_scaling_ptr = input_params.logn_scaling_ptr; + params.relative_attention_bias = reinterpret_cast(input_params.relative_attention_bias); + params.relative_attention_bias_stride = input_params.relative_attention_bias_stride; + params.max_distance = input_params.max_distance; + params.block_sparse_attention = input_params.block_sparse_attention; + params.block_sparse_params = input_params.block_sparse_params; + + // Attention mask input. + params.attention_mask = input_params.attention_mask; + params.attention_mask_stride = input_params.attention_mask_stride; + + // Attention sinks. + params.attention_sinks = input_params.attention_sinks; + + // The slope of linear position bias per head, e.g., ALiBi. + if (input_params.linear_bias_slopes != nullptr) + { + params.linear_bias_slopes = reinterpret_cast(input_params.linear_bias_slopes); + } + params.input_lengths = input_params.input_lengths; -using RunnerPtr = std::shared_ptr; -using torch_ext::trtllm::attention::Runner; -using torch_ext::trtllm::attention::AttentionInputType; + params.ia3_tasks = input_params.ia3_tasks; + params.ia3_key_weights = reinterpret_cast(input_params.ia3_key_weights); + params.ia3_value_weights = reinterpret_cast(input_params.ia3_value_weights); -static std::shared_ptr get_attention_op( - RunnerPtr const& runner, std::shared_ptr& op, int64_t local_layer_idx) -{ - auto cache_key = std::make_tuple(op->data(), runner->data()); - using CacheKey = decltype(cache_key); - static std::unordered_map, OpCustomHash> op_cache; - static std::shared_mutex op_cache_mutex; - - std::shared_lock read_lock{op_cache_mutex}; - - if (auto iter = op_cache.find(cache_key); iter != op_cache.end()) - { - TLLM_LOG_TRACE("Attention op for layer %lld is cached", local_layer_idx); - return iter->second; - } - - read_lock.unlock(); - TLLM_LOG_TRACE( - "Attention op for layer %lld is not cached, cache key: %s", local_layer_idx, to_string(cache_key).c_str()); - std::unique_lock lock{op_cache_mutex}; - op->initialize(); - // initialize() ends with mEnableContextFMHA = mIsGenerationMLA || mFmhaDispatcher->isSupported(), so the flag - // reflects the exact Q/KV/output precision, mask type and page size this op will run with. Checking here rather - // than inside initialize() keeps the throw out of that noexcept function. - // - // Paged-context attention exists to attend to KV already in the cache. The unfused self-attention fallback - // builds K and V from the current chunk alone, so it drops the cached prefix and then overwrites it: without - // these checks a missing kernel produces a plausible wrong answer instead of an error. - // - // Cross attention is exempt from both checks. Its unfused path builds K and V from the encoder output - // (params.cross_kv) rather than from a cached prefix, so it is correct whenever cross_kv is supplied. The one - // cross case that must read the cache (later chunks of a chunked decoder prefill, which arrive without - // cross_kv) is checked per call in the context-stage enqueue (the is_cross branch without cross_kv). - bool const needs_fused_paged_context - = op->mPagedContextFMHA && op->mPagedKVCache && !op->mIsMLAEnabled && !op->mCrossAttention; - // Relative position embedding (T5) has no fused context FMHA implementation at all: initialize() clears - // mEnableContextFMHA for it before the kernel table is consulted ("Fall back to unfused MHA because of relative - // position embedding"). Report that as an unsupported feature combination, not as a missing kernel; no - // --cuda_architectures list can supply one. This check runs first so its message wins. - TLLM_CHECK_WITH_INFO(!needs_fused_paged_context || !op->isRelativePosition(), - "Paged-context attention (chunked prefill, KV cache reuse or speculative draft tokens) is not supported with " - "relative position embedding: that attention always runs unfused, and the unfused path cannot attend to " - "cached KV. Disable chunked prefill, KV cache reuse and speculative decoding for this model. " - "Attention configuration: %s", - to_string(cache_key).c_str()); - TLLM_CHECK_WITH_INFO(!needs_fused_paged_context || op->mEnableContextFMHA, - "Paged-context attention requires a fused context FMHA kernel, and this build has none for this " - "configuration. The unfused fallback cannot attend to cached KV. If the device's SM is not named in the " - "build's --cuda_architectures then the build carries no kernels for it at all; check that first. Otherwise " - "use another attention backend, or disable chunked prefill, KV cache reuse and speculative decoding. " - "Attention configuration: %s", - to_string(cache_key).c_str()); - runner->prepare(*op); - auto [iter, _] = op_cache.try_emplace(cache_key, op); - return iter->second; + if (input_params.quant_option.hasStaticActivationScaling() || input_params.fp8_context_fmha) + { + // qkv_scale_out is nullptr currently (no scale). + params.qkv_scale_quant_orig = input_params.qkv_scale_out; + TLLM_CHECK_WITH_INFO(!input_params.fp8_context_fmha || input_params.attention_out_scale != nullptr, + "attention output scale should be provided."); + params.attention_out_scale_orig_quant = input_params.attention_out_scale; + } + + params.multi_block_mode = input_params.multi_block_mode; + // Cascade-attention partials must be wired regardless of multi_block_mode. + // Cascade decode runs with multi_block disabled (short-decode workloads have + // max_num_seq_len_tiles == 1, so enable_multi_block is structurally false). + // Gating these behind multi_block_mode leaves cascade_partial_* null and makes + // launch_cascade_attention fall back with "cascade workspace not provisioned". + params.cascade_partial_out = input_params.cascade_partial_out; + params.cascade_partial_max = input_params.cascade_partial_max; + params.cascade_partial_sum = input_params.cascade_partial_sum; + if (input_params.multi_block_mode) + { + params.min_seq_len_tile = input_params.min_seq_len_tile; + params.max_seq_len_tile = input_params.max_seq_len_tile; + + params.partial_out = reinterpret_cast(input_params.partial_out); + params.partial_sum = input_params.partial_sum; + params.partial_max = input_params.partial_max; + + params.block_counter = input_params.block_counter; + } + + params.multi_processor_count = input_params.multi_processor_count; + + // cross attn + params.memory_length_per_sample = input_params.memory_length_per_sample; + + params.mrope_position_deltas = input_params.mrope_position_deltas; + sync_check_cuda_error(stream); + + masked_multihead_attention(params, input_params.kv_block_array, input_params.shift_k_cache_buffer, stream); } -void attention(torch::Tensor q, std::optional k, std::optional v, torch::Tensor& output, - std::optional output_sf, std::optional workspace_, torch::Tensor sequence_length, - torch::Tensor host_past_key_value_lengths, torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, - torch::Tensor host_context_lengths, torch::Tensor host_request_types, - std::optional max_context_q_len_override, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, - std::optional cache_indirection, std::optional kv_scale_orig_quant, - std::optional kv_scale_quant_orig, std::optional out_scale, - std::optional rotary_inv_freq, std::optional rotary_cos_sin, - std::optional latent_cache, std::optional q_pe, - std::optional block_ids_per_seq, std::optional attention_sinks, - bool const is_fused_qkv, bool const update_kv_cache, int64_t const predicted_tokens_per_seq, - int64_t const local_layer_idx, int64_t const num_heads, int64_t const num_kv_heads, int64_t const head_size, - std::optional const tokens_per_block, int64_t const max_num_requests, int64_t const max_context_length, - int64_t const max_seq_len, int64_t const attention_window_size, int64_t const beam_width, int64_t const mask_type, - int64_t const quant_mode, double const q_scaling, int64_t const position_embedding_type, int64_t const rope_dim, - double const rope_base, int64_t const rope_scale_type, double const rope_scale, double const rope_short_m_scale, - double const rope_long_m_scale, int64_t const rope_max_positions, int64_t const rope_original_max_positions, - bool const use_paged_context_fmha, std::optional attention_input_type, bool is_mla_enable, - std::optional chunked_prefill_buffer_batch_size, std::optional q_lora_rank, - std::optional kv_lora_rank, std::optional qk_nope_head_dim, - std::optional qk_rope_head_dim, std::optional v_head_dim, std::optional rope_append, - std::optional mrope_rotary_cos_sin, std::optional mrope_position_deltas, - std::optional helix_position_offsets, std::optional helix_is_inactive_rank, - std::optional attention_chunk_size, std::optional softmax_stats_tensor, - bool const is_spec_decoding_enabled, bool const use_spec_decoding, bool const is_spec_dec_tree, - std::optional spec_decoding_generation_lengths, - std::optional spec_decoding_position_offsets_for_cpp, - std::optional spec_decoding_packed_mask, - std::optional spec_decoding_bl_tree_mask_offset, - std::optional spec_decoding_bl_tree_mask, - std::optional spec_bl_tree_first_sparse_mask_offset_kv, - std::optional sparse_kv_indices, std::optional sparse_kv_offsets, - std::optional sparse_attn_indices, std::optional sparse_attn_offsets, - int64_t const sparse_attn_indices_block_size, std::optional num_sparse_topk, - std::optional sparse_attn_kv_lens, std::optional skip_softmax_threshold_scale_factor_prefill, - std::optional skip_softmax_threshold_scale_factor_decode, std::optional skip_softmax_stat, - std::optional cu_q_seqlens, std::optional cu_kv_seqlens, - std::optional fmha_scheduler_counter, std::optional mla_bmm1_scale, - std::optional mla_bmm2_scale, std::optional quant_q_buffer, - std::optional flash_mla_tile_scheduler_metadata, std::optional flash_mla_num_splits, - int64_t sage_attn_num_elts_per_blk_q, int64_t sage_attn_num_elts_per_blk_k, int64_t sage_attn_num_elts_per_blk_v, - bool sage_attn_qk_int8, int64_t num_contexts, int64_t num_ctx_tokens, bool trtllm_gen_jit_warmup, - std::optional aux_kv_cache_pool_ptr, bool const is_cross, std::optional cross_kv, - std::optional relative_attention_bias, int64_t relative_attention_max_distance, - std::optional spec_decoding_target_max_draft_tokens, std::optional quant_scale_qkv, - std::optional dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion, - bool const force_prepare_spec_dec_tree_mask, std::optional const max_num_sequences, - std::optional kv_norm_weight, double kv_norm_eps, double skip_correction_threshold, - std::optional uses_spcompress) -{ - TLLM_LOG_TRACE("Attention op starts at layer %d", local_layer_idx); - // Use these tensors to infer if the attention is using KV cache - bool const use_kv_cache = kv_cache_block_offsets.has_value() && host_kv_cache_pool_pointers.has_value() - && host_kv_cache_pool_mapping.has_value(); +#define INSTANTIATE_MMHA_DISPATCH(T_MMHA, T) \ + template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ + FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ + template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ + FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ + template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ + FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); \ + template void fusedQKV_masked_attention_dispatch(Multihead_attention_params&, \ + FusedQKVMaskedAttentionDispatchParams const&, cudaStream_t stream); +INSTANTIATE_MMHA_DISPATCH(float, float) +INSTANTIATE_MMHA_DISPATCH(uint16_t, half) +#ifdef ENABLE_BF16 +INSTANTIATE_MMHA_DISPATCH(__nv_bfloat16, __nv_bfloat16) +#endif +#undef INSTANTIATE_MMHA_DISPATCH - bool const use_sage_attn - = sage_attn_num_elts_per_blk_q > 0 || sage_attn_num_elts_per_blk_k > 0 || sage_attn_num_elts_per_blk_v > 0; - TLLM_CHECK_WITH_INFO(is_mla_enable || is_fused_qkv || use_sage_attn || is_cross, - "For non-MLA, non-cross, non-SageAttention attention, only fused QKV is supported now."); - TLLM_CHECK_WITH_INFO( - update_kv_cache || is_cross, "KV cache update cannot be disabled now (except for cross attention)."); - auto qkv_or_q = q; - // MLA validates its separate Q/K/V or latent-cache inputs in Runner::run. - if (!is_mla_enable && is_fused_qkv) +int AttentionOp::getHeadSize(bool checkInit) const +{ + if (checkInit) { - TLLM_CHECK_WITH_INFO(!k.has_value(), "The k tensor should be null if using fused QKV"); - TLLM_CHECK_WITH_INFO(!v.has_value(), "The v tensor should be null if using fused QKV"); + TLLM_CHECK_WITH_INFO(mHeadSize > 0, "Trying to read mHeadSize before it's been initialized"); } - if (!is_mla_enable && !is_fused_qkv && update_kv_cache && !is_cross) + return mHeadSize; +} + +size_t AttentionOp::getFmhaMultiCtasKvScratchSize(FmhaParams const& p) const noexcept +{ + static constexpr size_t kMultiCtasKvRowsPerCta = 256; + static constexpr size_t kMultiCtasKvStatsPerRow = 2; + static constexpr size_t kMultiCtasKvPartialOElementSize = 2; + + size_t const headDimV + = mCfg.is_mla_enable ? static_cast(mMLAParams.kv_lora_rank) : static_cast(getHeadSize()); + size_t const maxRows = kMultiCtasKvRowsPerCta * static_cast(mMultiProcessorCount); + size_t const partialStatsSize = sizeof(float) * kMultiCtasKvStatsPerRow * maxRows; + size_t const partialOSize = kMultiCtasKvPartialOElementSize * maxRows * headDimV; + + return partialStatsSize + partialOSize; +} + +size_t AttentionOp::contextMlaWorkspaceBytesPerToken(int32_t numAttnHeads, int32_t qkRopeHeadDim, int32_t qkNopeHeadDim, + int32_t vHeadDim, bool fp8ContextMla, bool separateQAndKvInput, bool sparseMla) noexcept +{ + // Only the fp8 context-MLA separate-Q/KV path stages total_kv_len-scaled K/V dequant buffers. + // Sparse MLA reads K/V directly from the paged KV cache (no staging), so its per-token cost is 0. + if (!fp8ContextMla || !separateQAndKvInput || sparseMla) { - TLLM_CHECK_WITH_INFO(k.has_value(), "The k tensor should be provided if updating KV cache with unfused K/V"); - TLLM_CHECK_WITH_INFO(v.has_value(), "The v tensor should be provided if updating KV cache with unfused K/V"); + return 0; } - if (use_sage_attn) + // Mirror getWorkspaceSizeForContext's dim layout for the non-sparse fp8 branch: + // total_k_dim_all_heads = numAttnHeads * (qk_rope_head_dim + qk_nope_head_dim) + // total_v_dim_all_heads = numAttnHeads * v_head_dim + // The buffers are fp8 (1 byte/element), so bytes/token == element count. + int const dimKPerHead = qkRopeHeadDim + qkNopeHeadDim; + int const dimVPerHead = vHeadDim; + return static_cast(numAttnHeads) * static_cast(dimKPerHead + dimVPerHead); +} + +size_t AttentionOp::getWorkspaceSizeForContext(FmhaParams const& p, int32_t max_num_seq, int32_t input_seq_length, + int32_t cross_kv_length, int32_t max_num_tokens, int32_t total_kv_len) const noexcept +{ + if (max_num_tokens == 0) { - TLLM_CHECK_WITH_INFO( - !is_fused_qkv, "SageAttention requires separate q/k/v tensors (is_fused_qkv must be false)."); - TLLM_CHECK_WITH_INFO(k.has_value(), "SageAttention requires k tensor to be provided."); - TLLM_CHECK_WITH_INFO(v.has_value(), "SageAttention requires v tensor to be provided."); + return 0; } - auto const dtype = tensorrt_llm::runtime::TorchUtils::dataType(qkv_or_q.scalar_type()); - auto const out_dtype = output.scalar_type(); - bool const is_fp8_out = out_dtype == torch::kFloat8_e4m3fn; - // Torch does not support native nvfp4 type. - bool const is_fp4_out = out_dtype == torch::kUInt8; + int const local_hidden_units_qo = mNumAttnHeads * getHeadSize(); + int const local_hidden_units_kv = mNumAttnKVHeads * getHeadSize(); + + auto const size = tensorrt_llm::runtime::BufferDataType(p.getType()).getSize(); + + size_t context_workspace_size = 0; + + auto const batch_size = static_cast(max_num_seq); + auto const kv_seq_length = (isCrossAttention(p) ? cross_kv_length : input_seq_length); + // The unfused-MHA buffers below must upper-bound the enqueueContext carve, which sizes them by + // batch_size * input_seq_length (not num_tokens): with padding removal the actual token count can be + // smaller than batch_size * max(context q length), so sizing by max_num_tokens underestimates. + size_t const attention_mask_size = mEnableContextFMHA ? 0 : size * batch_size * input_seq_length * kv_seq_length; + size_t const cu_seqlens_size = sizeof(int) * (batch_size + 1); + size_t const rotary_inv_freq_size = sizeof(float) * batch_size * mCfg.rotary_embedding_dim / 2; - RunnerPtr runner; - if (dtype == tensorrt_llm::DataType::kHALF) + size_t q_buf_2_size = 0; + if (!mEnableContextFMHA) { - if (is_fp8_out) - { - runner = std::make_shared>(); - } - else if (is_fp4_out) - { - runner = std::make_shared>(); - } - else - { - TLLM_CHECK(out_dtype == torch::kFloat16); - runner = std::make_shared>(); - } + // Unfused mha + q_buf_2_size = size * batch_size * input_seq_length * local_hidden_units_qo; } - else if (dtype == tensorrt_llm::DataType::kFLOAT) + else if (mFmhaDispatcher->isSeparateQAndKvInput()) { - TLLM_CHECK(out_dtype == torch::kFloat32); - runner = std::make_shared>(); + // Paged context fmha + q_buf_2_size = (mFP8ContextFMHA ? 1 : size) * max_num_tokens * local_hidden_units_qo; } -#ifdef ENABLE_BF16 - else if (dtype == tensorrt_llm::DataType::kBF16) + + size_t const k_buf_2_size = mEnableContextFMHA ? 0 : size * batch_size * kv_seq_length * local_hidden_units_kv; + size_t const v_buf_2_size = mEnableContextFMHA ? 0 : size * batch_size * kv_seq_length * local_hidden_units_kv; + size_t const qk_buf_size + = mEnableContextFMHA ? 0 : size * batch_size * mCfg.num_heads * input_seq_length * kv_seq_length; + size_t const qkv_buf_2_size = mEnableContextFMHA ? 0 : size * batch_size * input_seq_length * local_hidden_units_qo; + size_t const qk_buf_float_size + = mEnableContextFMHA ? 0 : sizeof(float) * batch_size * mCfg.num_heads * input_seq_length * kv_seq_length; + int dim_q_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); + int dim_k_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); + int dim_v_per_head = (mMLAParams.v_head_dim); + if (useSparseMLA(p)) { - if (is_fp8_out) - { - runner = std::make_shared>(); - } - else if (is_fp4_out) + dim_q_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + dim_k_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + dim_v_per_head + = mMLAParams.rope_append ? mMLAParams.kv_lora_rank : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + } + + // Total dimension per token across all heads for Q, K, and V components respectively + int const total_q_dim_all_heads = mNumAttnHeads * dim_q_per_head; + int const total_k_dim_all_heads + = mNumAttnHeads * dim_k_per_head; // Assuming effective num_kv_heads = head_num for layout + int const total_v_dim_all_heads + = mNumAttnHeads * dim_v_per_head; // Assuming effective num_kv_heads = head_num for layout + bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mCfg.is_mla_enable + && mFmhaDispatcher->isSeparateQAndKvInput() + && (p.fwd.sage_attn_num_elts_per_blk_q > 0 || p.fwd.sage_attn_num_elts_per_blk_k > 0 + || p.fwd.sage_attn_num_elts_per_blk_v > 0); + + // Packed fp8 qkv buffer size for normal fp8 context FMHA + size_t fp8_qkv_buffer_size = mFP8ContextFMHA && mEnableContextFMHA && !mFmhaDispatcher->isSeparateQAndKvInput() + ? max_num_tokens * (local_hidden_units_qo + (2ULL * local_hidden_units_kv)) + : 0; + // Separate fp8 q/k/v buffer size for fp8 context MLA + size_t fp8_q_buf_size = 0; + size_t fp8_k_buf_size = 0; + size_t fp8_v_buf_size = 0; + if (mEnableContextFMHA && mFP8ContextMLA && mFmhaDispatcher->isSeparateQAndKvInput()) + { + fp8_q_buf_size = max_num_tokens * static_cast(total_q_dim_all_heads); + + if (useSparseMLA(p)) { - runner = std::make_shared>(); + // Sparse MLA (absorption mode): K and V are stored directly in KV cache during MLA RoPE kernel. + // No separate FP8 buffers needed for K/V since they're read from paged KV cache (Q_PAGED_KV layout). + fp8_k_buf_size = 0; + fp8_v_buf_size = 0; } else { - TLLM_CHECK(out_dtype == torch::kBFloat16); - runner = std::make_shared>(); + // Use total_kv_len when available (KV cache reuse causes total_kv_len >> max_num_tokens). + // enqueueContext sizes these buffers by total_kv_len, so workspace must match. + // NOTE: the per-token cost of these two buffers (total_k_dim_all_heads + total_v_dim_all_heads) is + // the single source of truth exposed via contextMlaWorkspaceBytesPerToken() for the KV-cache + // estimator's workspace reserve. Keep the two in sync if this dim layout changes. + size_t const kv_buf_tokens = std::max(static_cast(total_kv_len), + static_cast(p.fwd.chunked_prefill_buffer_batch_size) * max_num_tokens); + fp8_k_buf_size = kv_buf_tokens * static_cast(total_k_dim_all_heads); + fp8_v_buf_size = kv_buf_tokens * static_cast(total_v_dim_all_heads); + TLLM_CHECK(static_cast(total_k_dim_all_heads + total_v_dim_all_heads) + == contextMlaWorkspaceBytesPerToken(mNumAttnHeads, mMLAParams.qk_rope_head_dim, + mMLAParams.qk_nope_head_dim, mMLAParams.v_head_dim, mFP8ContextMLA, + /*separateQAndKvInput=*/true, useSparseMLA(p))); } } -#endif - runner->beam_width = beam_width; - runner->max_num_requests = max_num_requests; - runner->max_num_sequences = max_num_sequences.value_or(max_num_requests); - runner->attention_window_size = attention_window_size; - - auto op = std::make_shared(); - op->mType = dtype; - op->mFMHAForceFP32Acc = dtype == tensorrt_llm::DataType::kBF16; - op->mLayerIdx = local_layer_idx; - op->mNumHeads = num_heads; - op->mNumKVHeads = num_kv_heads; - op->mHeadSize = head_size; - op->mMaskType = static_cast(int32_t(mask_type)); - op->mKVCacheQuantMode = tensorrt_llm::common::QuantMode(uint32_t(quant_mode)); - op->mUseKVCache = use_kv_cache; - op->mPagedKVCache = op->mPagedKVCache && use_kv_cache; // update mPagedKVCache based on use_kv_cache - op->mTokensPerBlock = tokens_per_block.value_or(0); - op->mFP8GenerationMLA = false; - op->mFuseFp4Quant = is_fp4_out; - op->mFusesDsv4InvRopeFp8Quant = enable_dsv4_epilogue_fusion; - op->mMaxContextLength = max_context_length; - op->mMaxSeqLen = max_seq_len; - op->mMaxNumRequests = max_num_requests; - op->mQScaling = q_scaling; - op->mPositionEmbeddingType - = static_cast(int8_t(position_embedding_type)); - if (relative_attention_bias.has_value()) - { - auto const relative_attention_bias_dim = relative_attention_bias.value().dim(); - TORCH_CHECK(relative_attention_bias_dim == 2 || relative_attention_bias_dim == 3, - "relative_attention_bias must be [num_heads, num_buckets] for implicit mode or " - "[num_heads, max_seq_len, max_seq_len] for explicit mode"); - TORCH_CHECK(relative_attention_bias_dim != 2 || relative_attention_max_distance > 0, - "relative_attention_max_distance must be positive when relative_attention_bias is a bucket table"); - TORCH_CHECK(relative_attention_bias_dim != 3 || relative_attention_max_distance == 0, - "relative_attention_max_distance must be 0 when relative_attention_bias is precomputed"); - TLLM_CHECK_WITH_INFO(op->mPositionEmbeddingType == tensorrt_llm::kernels::PositionEmbeddingType::kRELATIVE, - "relative_attention_bias requires position_embedding_type to be relative."); - op->mMaxDistance = static_cast(relative_attention_max_distance); - } - op->mRotaryEmbeddingDim = rope_dim; - op->mRotaryEmbeddingBase = rope_base; - op->mRotaryEmbeddingScaleType = static_cast(int8_t(rope_scale_type)); - op->mRotaryEmbeddingScale = rope_scale; - op->mRotaryEmbeddingShortMscale = rope_short_m_scale; - op->mRotaryEmbeddingLongMscale = rope_long_m_scale; - op->mRotaryEmbeddingMaxPositions = rope_max_positions; - op->mRotaryEmbeddingOriginalMaxPositions = rope_original_max_positions; - op->mFP8ContextFMHA = is_fp8_out || is_fp4_out || (op->mKVCacheQuantMode.hasFp8KvCache() && use_paged_context_fmha) - || use_sage_attn; - // SageAttention block sizes and quantization mode. - op->mSageAttnNumEltsPerBlkQ = static_cast(sage_attn_num_elts_per_blk_q); - op->mSageAttnNumEltsPerBlkK = static_cast(sage_attn_num_elts_per_blk_k); - op->mSageAttnNumEltsPerBlkV = static_cast(sage_attn_num_elts_per_blk_v); - op->mSageAttnQkInt8 = sage_attn_qk_int8; - op->mFP8AttenOutput = is_fp8_out; - op->mPagedContextFMHA = use_paged_context_fmha; - op->mCrossAttention = is_cross; - - op->mAttentionChunkSize = attention_chunk_size; - op->mSkipSoftmaxThresholdScaleFactorPrefill - = static_cast(skip_softmax_threshold_scale_factor_prefill.value_or(0)); - op->mSkipSoftmaxThresholdScaleFactorDecode - = static_cast(skip_softmax_threshold_scale_factor_decode.value_or(0)); - auto constexpr maxSkipCorrectionThreshold = 32.0; - TORCH_CHECK(skip_correction_threshold >= 0.0 && skip_correction_threshold <= maxSkipCorrectionThreshold, - "skip_correction_threshold must be in the range (0, 32] when enabled, or 0 when disabled."); - int const smVersion = op->smVersion(); - bool const applySkipCorrection = is_mla_enable && (smVersion == 100 || smVersion == 103); - op->mSkipCorrectionThreshold = applySkipCorrection ? static_cast(skip_correction_threshold) : 0.0F; -#ifdef SKIP_SOFTMAX_STAT - op->mSkipSoftmaxTotalBlocks = reinterpret_cast(skip_softmax_stat.value().data_ptr()); - op->mSkipSoftmaxSkippedBlocks = op->mSkipSoftmaxTotalBlocks + 1; -#endif - op->mIsSpecDecodingEnabled = is_spec_decoding_enabled; - op->mUseSpecDecoding = use_spec_decoding; - op->mIsSpecDecTree = is_spec_dec_tree; - // Include the tree length in the AttentionOp cache key. - if (spec_decoding_target_max_draft_tokens.has_value() && op->mSpecDecodingTargetMaxGenLen == 0) + else if (useSageAttnSeparateQkv) { - op->mSpecDecodingTargetMaxGenLen = static_cast(spec_decoding_target_max_draft_tokens.value()) + 1; + fp8_q_buf_size = max_num_tokens * static_cast(local_hidden_units_qo); + fp8_k_buf_size = total_kv_len * static_cast(local_hidden_units_kv); + fp8_v_buf_size = total_kv_len * static_cast(local_hidden_units_kv); } - op->mForcePrepareSpecDecTreeMask = force_prepare_spec_dec_tree_mask; - op->mUseSparseAttention = false; - op->mUseTllmGenSparseAttentionPaged = false; - op->mUseTllmGenSparseAttention = false; - if ((sparse_kv_indices.has_value() && sparse_kv_indices.value().numel() > 0) - || (sparse_attn_indices.has_value() && sparse_attn_indices.value().numel() > 0)) - { - op->mUseSparseAttention = true; - if (sparse_attn_indices.has_value() && sparse_attn_indices.value().numel() > 0) - { - // Dispatch based on sparse_attn_offsets presence: - // - sparse_attn_offsets provided → generation paged sparse attention - // - sparse_attn_offsets absent → context sparse attention - if (sparse_attn_offsets.has_value() && sparse_attn_offsets.value().numel() > 0) - { - op->mUseTllmGenSparseAttentionPaged = true; - } - else - { - op->mUseTllmGenSparseAttention = true; - } - } - } - int32_t const num_sparse_topk_value = num_sparse_topk.has_value() ? num_sparse_topk.value() : 0; + int32_t const q_max_n_blk = p.fwd.sage_attn_num_elts_per_blk_q > 0 + ? tc::divUp(max_num_tokens, p.fwd.sage_attn_num_elts_per_blk_q) + batch_size - 1 + : 0; + int32_t const k_max_n_blk = p.fwd.sage_attn_num_elts_per_blk_k > 0 + ? tc::divUp(total_kv_len, p.fwd.sage_attn_num_elts_per_blk_k) + batch_size - 1 + : 0; + size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); + size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); + size_t const sage_v_sfs_buffer_size = p.fwd.sage_attn_num_elts_per_blk_v > 0 + ? sizeof(float) * tc::divUp(local_hidden_units_kv, std::max(1, p.fwd.sage_attn_num_elts_per_blk_v)) + : 0; - if (is_mla_enable) + size_t const padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * batch_size * input_seq_length; + size_t const encoder_padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * batch_size * cross_kv_length; + // Each token holds (batch_idx, token_idx_in_seq) int2. + size_t const tokens_info_size = sizeof(int2) * max_num_tokens; + size_t const fmha_scheduler_counter = mEnableContextFMHA ? sizeof(uint32_t) : 0; + size_t const fmha_bmm1_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) * 2 : 0; + size_t const fmha_bmm2_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) : 0; + + size_t const fmha_multi_ctas_kv_scratch_size = useTllmGenSparseAttention(p) ? getFmhaMultiCtasKvScratchSize(p) : 0; + + AttentionContextWorkspaceSizes workspaceSizes{}; + workspaceSizes.attentionMask = attention_mask_size; + workspaceSizes.cuQSeqlens = cu_seqlens_size; + workspaceSizes.cuKvSeqlens = cu_seqlens_size; + workspaceSizes.cuMaskRows = cu_seqlens_size; + workspaceSizes.rotaryInvFreq = rotary_inv_freq_size; + workspaceSizes.qBuf = q_buf_2_size; + workspaceSizes.kBuf = k_buf_2_size; + workspaceSizes.vBuf = v_buf_2_size; + workspaceSizes.qkBuf = qk_buf_size; + workspaceSizes.qkvBuf = qkv_buf_2_size; + workspaceSizes.qkFloatBuf = qk_buf_float_size; + workspaceSizes.fp8QkvBuf = fp8_qkv_buffer_size; + workspaceSizes.fp8QBuf = fp8_q_buf_size; + workspaceSizes.fp8KBuf = fp8_k_buf_size; + workspaceSizes.fp8VBuf = fp8_v_buf_size; + workspaceSizes.paddingOffset = padding_offset_size; + workspaceSizes.encoderPaddingOffset = encoder_padding_offset_size; + workspaceSizes.tokensInfo = tokens_info_size; + workspaceSizes.fmhaTileCounter = fmha_scheduler_counter; + workspaceSizes.fmhaBmm1Scale = fmha_bmm1_scale_size; + workspaceSizes.fmhaBmm2Scale = fmha_bmm2_scale_size; + workspaceSizes.sageQScale = sage_q_sfs_buffer_size; + workspaceSizes.sageKScale = sage_k_sfs_buffer_size; + workspaceSizes.sageVScale = sage_v_sfs_buffer_size; + workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; + context_workspace_size = AttentionWorkspaceManager::buildContextLayout(workspaceSizes).totalSize; + + return context_workspace_size; +} + +size_t AttentionOp::getWorkspaceSizeForGeneration(FmhaParams const& p, int32_t max_num_seq, + int32_t max_attention_window_size, int32_t max_num_tokens, int32_t max_blocks_per_sequence) const noexcept +{ + if (max_num_tokens == 0) { - // MLA does not support NVFP4 output yet. - TLLM_CHECK(!is_fp4_out); + return 0; + } - TLLM_CHECK(host_kv_cache_pool_mapping.has_value()); - int32_t const layer_num = host_kv_cache_pool_mapping.value().size(0); - bool const rope_append_value = rope_append.value_or(true); + auto const size = tensorrt_llm::runtime::BufferDataType(p.getType()).getSize(); + int const batch_beam = max_num_seq; - if (num_sparse_topk_value > 0 && sparse_attn_indices.has_value() && sparse_attn_indices.value().numel() > 0) + // Compute the workspace size for MLA. + size_t fmha_v2_mla_workspace_size = 0; + if (mCfg.is_mla_enable) + { + size_t flash_mla_workspace_size = 0; + if (mUseGenFlashMLA) { - op->mUseSparseAttention = true; + static constexpr int TileSchedulerMetaDataSize = 8; + + int s_q = mMLAParams.predicted_tokens_per_seq; + + int num_q_heads = mCfg.num_heads; + int num_kv_heads = mNumKVHeads; + int head_size_v = (p.use_sparse_attention && !mMLAParams.rope_append) + ? mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim + : mMLAParams.kv_lora_rank; + + int num_sm_parts = getFlashMlaNumSmParts(s_q, num_q_heads, num_kv_heads, head_size_v); + + AttentionFlashMlaWorkspaceSizes flashMlaWorkspaceSizes{}; + flashMlaWorkspaceSizes.tileSchedulerMetadata = sizeof(int) * (num_sm_parts * TileSchedulerMetaDataSize); + flashMlaWorkspaceSizes.numSplits = sizeof(int) * (batch_beam + 1); + flashMlaWorkspaceSizes.softmaxLse = sizeof(float) * (batch_beam * s_q * num_q_heads); + flashMlaWorkspaceSizes.softmaxLseAccum = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q); + flashMlaWorkspaceSizes.outAccum + = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q * head_size_v); + flash_mla_workspace_size = AttentionWorkspaceManager::buildFlashMlaLayout(flashMlaWorkspaceSizes).totalSize; } - op->mIsMLAEnabled = true; - op->mMLAParams = {static_cast(q_lora_rank.value()), static_cast(kv_lora_rank.value()), - static_cast(qk_nope_head_dim.value()), static_cast(qk_rope_head_dim.value()), - static_cast(v_head_dim.value()), static_cast(predicted_tokens_per_seq), - static_cast(layer_num), static_cast(rope_append_value)}; - - op->mUseNvfp4MlaKvCache = op->mKVCacheQuantMode.hasFp4KvCache() && op->mUseTllmGenSparseAttention - && !sparse_attn_kv_lens.has_value() && aux_kv_cache_pool_ptr.has_value(); - op->mFP8ContextMLA - = (tensorrt_llm::common::getSMVersion() == 90 || tensorrt_llm::common::getSMVersion() == 100 - || tensorrt_llm::common::getSMVersion() == 103 || tensorrt_llm::common::getSMVersion() == 107 - || tensorrt_llm::common::getSMVersion() == 120) - && (op->mKVCacheQuantMode.hasFp8KvCache() || op->mUseNvfp4MlaKvCache); - op->mIsGenerationMLA = head_size == op->mMLAParams.kv_lora_rank + op->mMLAParams.qk_rope_head_dim; - op->mFP8GenerationMLA = op->mKVCacheQuantMode.hasFp8KvCache() || op->mUseNvfp4MlaKvCache; - // only enable flash mla on sm90 and head_size == 576 and tokens_per_block == 64 - op->mUseGenFlashMLA = tensorrt_llm::common::getSMVersion() == 90 && tokens_per_block == 64 && head_size == 576; + size_t const cu_seqlens_size = sizeof(int) * (max_num_seq + 1); + size_t const fmha_scheduler_counter = sizeof(uint32_t); + size_t const fmha_multi_ctas_kv_scratch_size = getFmhaMultiCtasKvScratchSize(p); - // The following two parameters are used to compute kvcache related parameters such as kvcache block_size. So - // they need to be set to 1 and 512 + 64 for both context and generation. For MLA attention kernel configs, - // mNumKVHeads/mHeadSize are overwritten in common/attentionOp.cpp. - op->mNumKVHeads = 1; - op->mHeadSize = op->mMLAParams.kv_lora_rank + op->mMLAParams.qk_rope_head_dim; + int const NUM_BUFFERS = 5; + size_t workspaces[NUM_BUFFERS]; + workspaces[0] = mIsGenerationMLA ? 0 : cu_seqlens_size; // cu_q_len + workspaces[1] = mIsGenerationMLA ? 0 : cu_seqlens_size; // cu_kv_len + workspaces[2] = mIsGenerationMLA ? 0 : fmha_scheduler_counter; + workspaces[3] = fmha_multi_ctas_kv_scratch_size; + workspaces[4] = flash_mla_workspace_size; - // For chunked prefill MLA, we need larger buffer size for k and v - op->mChunkPrefillBufferBatchSize - = chunked_prefill_buffer_batch_size.has_value() ? chunked_prefill_buffer_batch_size.value() : 1; + fmha_v2_mla_workspace_size = tc::calculateTotalWorkspaceSize(workspaces, NUM_BUFFERS); } - op->mUsesSpcompress = uses_spcompress.value_or(false); - if (op->mUsesSpcompress) + size_t generation_workspace_size = 0; + // The minimum number of sequence length tiles (limited by the shared memory size). + int minSeqLenTile + = estimate_min_multi_block_count(max_attention_window_size, mMaxSharedMemoryPerBlockOptin - 2048, size); + int32_t const maxSeqLenTile = std::max( + {minSeqLenTile, getMaxNumSeqLenTile(p, batch_beam), (int) tc::divUp(mMultiProcessorCount, mCfg.num_heads)}); + + size_t const partial_out_size = size * batch_beam * mCfg.num_heads * mHeadSize * maxSeqLenTile; + size_t const partial_sum_size = sizeof(float) * batch_beam * mCfg.num_heads * maxSeqLenTile; + size_t const partial_max_size = sizeof(float) * batch_beam * mCfg.num_heads * maxSeqLenTile; + size_t const shift_k_cache_size = (!p.pos_shift_enabled || isCrossAttention(p)) + ? 0 + : size * batch_beam * mCfg.num_heads * mHeadSize * max_attention_window_size; + AttentionGenerationWorkspaceSizes generationWorkspaceSizes{}; + generationWorkspaceSizes.partialOut = partial_out_size; + generationWorkspaceSizes.partialSum = partial_sum_size; + generationWorkspaceSizes.partialMax = partial_max_size; + generationWorkspaceSizes.shiftKCache = shift_k_cache_size; { - int const smVersionSpcompress = op->smVersion(); - TORCH_CHECK(smVersionSpcompress == 107, - "uses_spcompress is only supported on SM107. Got SM version: ", smVersionSpcompress); - TORCH_CHECK( - op->mFP8ContextFMHA || op->mFP8ContextMLA, "uses_spcompress requires FP8 context FMHA or FP8 context MLA."); + auto const cascadeSizes + = tensorrt_llm::kernels::mmha::cascade::getCascadeWorkspaceSizes(batch_beam, mCfg.num_heads, mHeadSize); + generationWorkspaceSizes.cascadeOut = cascadeSizes.out; + generationWorkspaceSizes.cascadeMax = cascadeSizes.mMax; + generationWorkspaceSizes.cascadeSum = cascadeSizes.lSum; } + generation_workspace_size = AttentionWorkspaceManager::buildGenerationLayout(generationWorkspaceSizes).totalSize; - op = get_attention_op(runner, op, local_layer_idx); + size_t xqa_workspace_size = 0; + if (mEnableXQA) + { + size_t const cu_seqlens_size = sizeof(int) * (batch_beam + 1); + size_t const cu_kv_seqlens_size = sizeof(int) * (batch_beam + 1); + size_t const rotary_inv_freq_size = sizeof(float) * batch_beam * mCfg.rotary_embedding_dim / 2; + // Two workspaces for sparse attention. One for the sequence lengths, and one for kv block offsets. + size_t const sparse_attn_cache_size = useTllmGenSparseAttentionPaged(p) + ? sizeof(int) * (batch_beam + batch_beam * 2 * max_blocks_per_sequence) * mNumKVHeads + : 0; + AttentionXqaWorkspaceSizes xqaWorkspaceSizes{}; + xqaWorkspaceSizes.cuSeqlens = cu_seqlens_size; + xqaWorkspaceSizes.cuKvSeqlens = cu_kv_seqlens_size; + xqaWorkspaceSizes.rotaryInvFreq = rotary_inv_freq_size; + xqaWorkspaceSizes.tokensInfo = max_num_tokens * sizeof(int2); + xqaWorkspaceSizes.bmm1Scale = sizeof(float) * 2; + xqaWorkspaceSizes.bmm2Scale = sizeof(float); + xqaWorkspaceSizes.sparseAttnCache = sparse_attn_cache_size; + xqaWorkspaceSizes.kernelWorkspace = mXqaDispatcher->getWorkspaceSize( + std::min(p.spec_decoding_max_generation_length * max_num_seq, max_num_tokens)); + xqa_workspace_size + = AttentionWorkspaceManager::buildXqaLayout(xqaWorkspaceSizes, mXqaDispatcher->getWorkspaceAlignment()) + .totalSize; + } - int32_t const num_seqs = host_context_lengths.size(0); - RequestType const* request_types = static_cast(host_request_types.data_ptr()); + return std::max(std::max(generation_workspace_size, xqa_workspace_size), fmha_v2_mla_workspace_size); +} - AttentionInputType attn_input_type = AttentionInputType::Mixed; - if (attention_input_type.has_value()) +int AttentionOp::getMaxNumSeqLenTile(FmhaParams const& p, int batch_beam_size) const +{ + if (mMultiBlockMode) { - attn_input_type = static_cast(attention_input_type.value()); + // And we allocate the buffer based on the maximum number of blocks per sequence (batch_beam_size = 1). + // Assume we can only have 1 block (large block size like 1024) in SM, and we only want one wave of blocks. + return tc::getEnvMmhaMultiblockDebug() ? std::max(kReservedMaxSeqLenTilePerSeq, getEnvMmhaBlocksPerSequence()) + : tc::divUp(mMultiProcessorCount, batch_beam_size * mCfg.num_heads); } - bool const is_gen_only = attn_input_type == AttentionInputType::GenerationOnly; - - int32_t const num_generations = num_seqs - static_cast(num_contexts); - int32_t const num_tokens = qkv_or_q.size(0); - int32_t const num_gen_tokens = is_gen_only ? num_tokens : num_tokens - static_cast(num_ctx_tokens); - auto const ctx_total_kv_len = host_total_kv_lens.index({0}).item(); - auto const gen_total_kv_len = host_total_kv_lens.index({1}).item(); + return 0; +} - for (int32_t idx = num_contexts; idx < num_seqs; idx++) +template +int AttentionOp::mlaGeneration(MlaParams& params, FmhaParams const& p, cudaStream_t stream) +{ + TLLM_CHECK_WITH_INFO(params.seqQOffset != nullptr, "seqQOffset is nullptr."); + TLLM_CHECK_WITH_INFO(params.cache_seq_lens != nullptr, "cache_seq_lens is nullptr."); + TLLM_CHECK_WITH_INFO(params.fmha_tile_counter != nullptr, "fmha_tile_counter is nullptr."); + if (mFP8GenerationMLA) { - TLLM_CHECK(request_types[idx] == RequestType::kGENERATION); + TLLM_CHECK_WITH_INFO(params.quant_q_buf != nullptr, "quant_q_buf is nullptr."); + TLLM_CHECK_WITH_INFO(params.bmm1_scale != nullptr, "bmm1_scale is nullptr."); + TLLM_CHECK_WITH_INFO(params.bmm2_scale != nullptr, "bmm2_scale is nullptr."); } - int32_t const max_attention_window_size - = beam_width == 1 ? attention_window_size : cache_indirection.value().size(2); - int32_t const max_blocks_per_sequence - = use_kv_cache && kv_cache_block_offsets.has_value() ? kv_cache_block_offsets.value().size(-1) : 0; - // For cross-attention, several unfused-path context buffers scale with the encoder KV length. - // Mirror the context-stage enqueue, which uses the max past-KV length over the context sequences - // as cross_kv_length; sizing with 0 here under-allocates the workspace and the carved views in - // enqueueContext land past the end of the allocation. The enqueue also gates on cross_kv.has_value(), - // so this can over-allocate relative to the carve; that is safe. - int32_t maxCrossKvLength = 0; - if (op->isCrossAttention() && num_contexts > 0) + int const num_kv_heads = 1; + int const head_size = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + int const head_size_v = (useSparseMLA(p) && !mMLAParams.rope_append) ? head_size : mMLAParams.kv_lora_rank; + int32_t const batch_beam = p.beam_width * p.num_requests; + + // The element size of the KV cache. + auto const elemSize = mFP8GenerationMLA ? sizeof(__nv_fp8_e4m3) : sizeof(T); + auto const sizePerToken = num_kv_heads * head_size * elemSize; + params.cache_type = (mFP8GenerationMLA ? KvCacheDataType::FP8 : KvCacheDataType::BASE); + + int32_t const kvCachePoolIndex = p.getKvCachePoolIndex(p.local_layer_idx); + auto kv_cache_buffer = KVBlockArray(batch_beam, p.getMaxBlocksPerSequence(), mCfg.tokens_per_block, sizePerToken, + p.cyclic_attention_window_size, p.max_cyclic_attention_window_size, p.sink_token_length, + p.can_use_one_more_block, p.getHostPrimaryPoolPtr(), p.getHostSecondaryPoolPtr(), + p.getKvCacheBlockOffsets(kvCachePoolIndex)); + + // Static sparse NVFP4 MLA reads a separately dequantized FP8 scratch pool, + // so this paged-cache scale descriptor is not consumed by the attention kernel. + auto kv_scale_cache_buffer = KVBlockArray(); + + void* scratchPtr = params.workspace; + + params.quant_scale_o = p.getOutScale(); + params.quant_scale_q = p.getKvScaleOrigQuant(); + params.quant_scale_kv = p.getKvScaleOrigQuant(); + params.dequant_scale_q = p.getKvScaleQuantOrig(); + params.dequant_scale_kv = p.getKvScaleQuantOrig(); + params.host_bmm1_scale + = 1 / (mCfg.q_scaling * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); + + if (p.runtime_perf_knobs.has_value()) { - maxCrossKvLength = host_past_key_value_lengths.slice(0, 0, num_contexts).max().item(); + int64_t const* const runtimePerfKnobs = p.getRuntimePerfKnobs(); + int64_t multi_block_mode_val = runtimePerfKnobs[0]; + mMultiBlockMode = multi_block_mode_val == 1; + int64_t enable_context_fmha_fp32_acc_val = runtimePerfKnobs[1]; + mFMHAForceFP32Acc = mFMHAForceFP32Acc || enable_context_fmha_fp32_acc_val == 1; } - int64_t const workspace_size = runner->getWorkspaceSize(*op, num_tokens, max_attention_window_size, num_gen_tokens, - max_blocks_per_sequence, ctx_total_kv_len, maxCrossKvLength); - TLLM_LOG_TRACE("Expected workspace size is %ld bytes", workspace_size); - torch::Tensor workspace; - if (workspace_.has_value()) + if (common::getEnvForceDeterministicAttention()) { - if (workspace_.value().numel() < workspace_size) + mMultiBlockMode = false; + } + + if (mUseTllmGen) + { + TLLM_CHECK_WITH_INFO(mTllmGenFMHARunner.get(), "mTllmGenFMHARunner not initialized."); + TllmGenFmhaRunnerParams tllmRunnerParams{}; + + // Parameters to select kernels. + // MLA generation kernels use dense mask. For multi-token generation, TRTLLM-Gen applies causality by + // shrinking each token's effective KV length. + tllmRunnerParams.mMaskType = TrtllmGenAttentionMaskType::Dense; + tllmRunnerParams.mKernelType = FmhaKernelType::Generation; + tllmRunnerParams.mMultiCtasKvMode = mMultiBlockMode; + // Note that the tileScheduler and multiCtasKvMode will be automatically tuned when using multi_block mode. + // Otherwise, always enable the persistent scheduler for better performance. + tllmRunnerParams.mTileScheduler = mMultiBlockMode ? TileScheduler::Static : TileScheduler::Persistent; + + // Q buffer. + tllmRunnerParams.qPtr = mFP8GenerationMLA ? reinterpret_cast(params.quant_q_buf) + : reinterpret_cast(params.q_buf); + + // KV buffer + // Paged KV + tllmRunnerParams.mQkvLayout = QkvLayout::PagedKv; + tllmRunnerParams.kvPtr = kv_cache_buffer.mPrimaryPoolPtr; + tllmRunnerParams.kvPageIdxPtr = reinterpret_cast(kv_cache_buffer.data); + tllmRunnerParams.mMaxNumPagesPerSeqKv = kv_cache_buffer.mMaxBlocksPerSeq; + tllmRunnerParams.mNumTokensPerPage = kv_cache_buffer.mTokensPerBlock; + + // The partial buffers' pointers when the multiCtasKv mode is enabled. + tllmRunnerParams.multiCtasKvCounterPtr = static_cast(p.getMultiCtasKvCounter()); + tllmRunnerParams.multiCtasKvScratchPtr = scratchPtr; + + // The sequence lengths for K/V. + tllmRunnerParams.seqLensKvPtr = params.cache_seq_lens; + + tllmRunnerParams.oPtr = reinterpret_cast(params.context_buf); + tllmRunnerParams.oSfPtr = p.getOutputSf(); + if (params.dsv4_epilogue_fusion.enabled) + { + tllmRunnerParams.mDsv4EpilogueFusion.enabled = true; + tllmRunnerParams.mDsv4EpilogueFusion.cosSinCache = params.dsv4_epilogue_fusion.cos_sin_cache; + tllmRunnerParams.mDsv4EpilogueFusion.scaleBufM = params.dsv4_epilogue_fusion.scale_buf_m; + } + + // softmax stats if needed + tllmRunnerParams.softmaxStatsPtr = p.getSoftmaxStatsTensor(); + + // Per-head attention sink added to the softmax denominator. + tllmRunnerParams.attentionSinksPtr = p.getAttentionSinks(); + + // MLA uses different head dimensions for Qk and V. + tllmRunnerParams.mHeadDimQk = head_size; + tllmRunnerParams.mHeadDimV = head_size_v; + + auto const num_q_heads = mNumAttnHeads; + tllmRunnerParams.mNumHeadsQ = num_q_heads; + tllmRunnerParams.mNumHeadsKv = num_kv_heads; + tllmRunnerParams.mNumHeadsQPerKv = num_q_heads / num_kv_heads; + + tllmRunnerParams.mBatchSize = batch_beam; + // It is used to construct contiguous kv cache TMA descriptors. + tllmRunnerParams.mMaxSeqLenCacheKv = p.max_attention_window_size; + // This should be set to numDraftTokens + 1. + tllmRunnerParams.mMaxSeqLenQ = params.acc_q_len / batch_beam; + tllmRunnerParams.mMaxSeqLenKv = p.max_past_kv_length; + tllmRunnerParams.mJITWarmup = p.trtllm_gen_jit_warmup; + tllmRunnerParams.mJITWarmupMaxNumRequests = p.max_num_requests; + tllmRunnerParams.mJITWarmupMaxSeqLenQ = p.max_context_length; + tllmRunnerParams.mJITWarmupMaxSeqLenKv = p.max_seq_len; + tllmRunnerParams.mSumOfSeqLensQ = int(batch_beam * tllmRunnerParams.mMaxSeqLenQ); + // Not used in the generation kernels as contiguous_kv or paged_kv layouts are used. + tllmRunnerParams.mSumOfSeqLensKv = int(batch_beam * tllmRunnerParams.mMaxSeqLenKv); + + // The attention window size. + tllmRunnerParams.mAttentionWindowSize = p.cyclic_attention_window_size; + // The chunked attention size. + tllmRunnerParams.mChunkedAttentionSize = INT_MAX; + + // The scaleQ that will be applied to the BMM1 output. + tllmRunnerParams.mScaleQ = mCfg.q_scaling + * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) + / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); + + // Set it to INT_MAX as the kv cache pageOffsets will ensure that there is no out-of-bounds access. + tllmRunnerParams.mNumPagesInMemPool = INT_MAX; + tllmRunnerParams.mMultiProcessorCount = mMultiProcessorCount; + tllmRunnerParams.stream = stream; + tllmRunnerParams.mSfStartTokenIdx = p.token_offset; + + // Scales for quantization + if (mFP8GenerationMLA) + { + static constexpr int bmm1_scale_offset = 1; + tllmRunnerParams.outputScalePtr = reinterpret_cast(params.bmm2_scale); + tllmRunnerParams.scaleSoftmaxLog2Ptr + = reinterpret_cast(params.bmm1_scale) + bmm1_scale_offset; + } + + // Set the following parameters if sparseAttention is used. + if (useSparseMLA(p)) + { + bool const useDynamicSparseMLA = p.sparse_params.sparse_attn_kv_lens != nullptr; + tllmRunnerParams.mSparseAttention + = useDynamicSparseMLA ? SparseType::DynamicTokenSparse : SparseType::StaticTokenSparse; + tllmRunnerParams.mSkipCorrThreshold = mSkipCorrectionThreshold; + tllmRunnerParams.mSparseTopK = p.sparse_params.num_sparse_topk; + tllmRunnerParams.ptrSparseMlaTopKLens = p.sparse_params.sparse_attn_kv_lens; + tllmRunnerParams.kvPageIdxPtr + = reinterpret_cast(p.sparse_params.sparse_attn_indices); + if (useDynamicSparseMLA) + { + TLLM_CHECK_WITH_INFO(p.sparse_params.sliding_window_kv_cache_pool != nullptr, + "SWA KV pool must be set for dynamic sparse MLA."); + // Dynamic sparse MLA always has an SWA pool. The compressed pool is optional; when it + // is absent (ratio == 1), use SWA as kvPtr only to keep TG's primary TMA descriptor valid. + tllmRunnerParams.kvPtr = p.sparse_params.sparse_kv_cache_pool != nullptr + ? p.sparse_params.sparse_kv_cache_pool + : p.sparse_params.sliding_window_kv_cache_pool; + tllmRunnerParams.slidingWindowKvPoolBasePtr = p.sparse_params.sliding_window_kv_cache_pool; + } + else + { + tllmRunnerParams.kvPtr = p.sparse_params.sparse_kv_cache_pool; + } + + if (!useDynamicSparseMLA && mUseNvfp4MlaKvCache) + { + // Static sparse MLA indexes a compact KV pool containing at most + // mSparseTopK rows per query. Do not let the original dense KV + // length drive kernel selection or launch geometry: for long + // sequences that can select a multi-CTA kernel which addresses + // beyond the compact page table. + TLLM_CHECK_WITH_INFO(tllmRunnerParams.mSparseTopK > 0, + "Static sparse MLA requires a positive TopK, got %d", tllmRunnerParams.mSparseTopK); + int32_t const originalMaxSeqLenKv = tllmRunnerParams.mMaxSeqLenKv; + int32_t const effectiveMaxSeqLenKv = std::min(originalMaxSeqLenKv, tllmRunnerParams.mSparseTopK); + tllmRunnerParams.mMaxSeqLenKv = effectiveMaxSeqLenKv; + tllmRunnerParams.mJITWarmupMaxSeqLenKv + = std::min(tllmRunnerParams.mJITWarmupMaxSeqLenKv, effectiveMaxSeqLenKv); + int64_t const sumOfSeqLensKv = static_cast(tllmRunnerParams.mBatchSize) * effectiveMaxSeqLenKv; + TLLM_CHECK_WITH_INFO(sumOfSeqLensKv <= std::numeric_limits::max(), + "Static sparse MLA cumulative KV length exceeds int32 capacity: %ld", sumOfSeqLensKv); + tllmRunnerParams.mSumOfSeqLensKv = static_cast(sumOfSeqLensKv); + TLLM_LOG_DEBUG("Clamp static sparse MLA max KV length from %d to %d (TopK=%d)", originalMaxSeqLenKv, + effectiveMaxSeqLenKv, tllmRunnerParams.mSparseTopK); + } + } + + mTllmGenFMHARunner->run(tllmRunnerParams); + sync_check_cuda_error(stream); + } + else if (mUseGenFlashMLA) + { + static constexpr int TileSchedulerMetaDataSize = 8; + + int const num_q_heads = mCfg.num_heads; + int const ngroups = num_q_heads / num_kv_heads; + + int const s_q = params.acc_q_len / batch_beam; + assert(s_q == mMLAParams.predicted_tokens_per_seq); + int const head_size_v = mMLAParams.kv_lora_rank; + int const num_sm_parts = getFlashMlaNumSmParts(s_q, num_q_heads, num_kv_heads, head_size_v); + + size_t const num_splits_size = sizeof(int) * (batch_beam + 1); + size_t const tile_scheduler_metadata_size = sizeof(int) * (num_sm_parts * TileSchedulerMetaDataSize); + size_t const softmax_lse_size = sizeof(float) * (batch_beam * s_q * num_q_heads * num_kv_heads); // softmax_lse + size_t const softmax_lse_accum_size = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q); + size_t const out_accum_size = sizeof(float) * ((batch_beam + num_sm_parts) * num_q_heads * s_q * head_size_v); + + AttentionFlashMlaWorkspaceSizes flashMlaWorkspaceSizes{}; + flashMlaWorkspaceSizes.tileSchedulerMetadata = tile_scheduler_metadata_size; + flashMlaWorkspaceSizes.numSplits = num_splits_size; + flashMlaWorkspaceSizes.softmaxLse = softmax_lse_size; + flashMlaWorkspaceSizes.softmaxLseAccum = softmax_lse_accum_size; + flashMlaWorkspaceSizes.outAccum = out_accum_size; + auto const flashMlaWorkspaceLayout = AttentionWorkspaceManager::buildFlashMlaLayout(flashMlaWorkspaceSizes); + float* softmax_lse_ptr + = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.softmaxLse); + float* softmax_lse_accum_ptr + = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.softmaxLseAccum); + float* out_accum_ptr + = AttentionWorkspaceManager::ptr(params.workspace, flashMlaWorkspaceLayout.outAccum); + + // Metadata must always be pre-computed by Python (compute_flash_mla_metadata) and passed in. + TLLM_CHECK_WITH_INFO(params.flash_mla_tile_scheduler_metadata != nullptr, + "FlashMLA tile-scheduler metadata must be pre-computed by Python."); + TLLM_CHECK_WITH_INFO( + params.flash_mla_num_splits != nullptr, "FlashMLA num_splits must be pre-computed by Python."); + int* tile_scheduler_metadata_ptr = const_cast(params.flash_mla_tile_scheduler_metadata); + int* num_splits_ptr = const_cast(params.flash_mla_num_splits); + + Flash_fwd_mla_params flashMlaParams{}; + flashMlaParams.b = batch_beam; + flashMlaParams.seqlen_q = ngroups * s_q; + flashMlaParams.cu_seqlens_k = const_cast(params.cache_seq_lens); + flashMlaParams.h = 1; + flashMlaParams.h_h_k_ratio = 1; + + float softmax_scale + = 1.0f / (mCfg.q_scaling * sqrtf((mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim) * 1.0f)); + + flashMlaParams.ngroups = ngroups; + flashMlaParams.is_causal = !(s_q == 1); + flashMlaParams.d = head_size; + flashMlaParams.d_v = head_size_v; + flashMlaParams.scale_softmax = softmax_scale; + flashMlaParams.scale_softmax_log2 = float(softmax_scale * M_LOG2E); + + flashMlaParams.q_ptr = mFP8GenerationMLA ? const_cast(reinterpret_cast(params.quant_q_buf)) + : const_cast(reinterpret_cast(params.q_buf)); + flashMlaParams.k_ptr = kv_cache_buffer.mPrimaryPoolPtr; + flashMlaParams.v_ptr = flashMlaParams.k_ptr; + flashMlaParams.o_ptr = reinterpret_cast(params.context_buf); + flashMlaParams.softmax_lse_ptr = softmax_lse_ptr; + + // since head_num_kv = 1 + flashMlaParams.q_batch_stride = head_size * params.head_num * s_q; + flashMlaParams.k_batch_stride = mCfg.tokens_per_block * num_kv_heads * head_size * mMLAParams.num_layers; + flashMlaParams.o_batch_stride = s_q * num_q_heads * head_size_v; + flashMlaParams.q_row_stride = head_size; + flashMlaParams.k_row_stride = head_size; + flashMlaParams.o_row_stride = head_size_v; + flashMlaParams.q_head_stride = head_size; + flashMlaParams.k_head_stride = head_size; + flashMlaParams.o_head_stride = head_size_v; + + flashMlaParams.v_batch_stride = flashMlaParams.k_batch_stride; + flashMlaParams.v_row_stride = flashMlaParams.k_row_stride; + flashMlaParams.v_head_stride = flashMlaParams.k_head_stride; + + flashMlaParams.block_table = const_cast(params.block_ids_per_seq); + flashMlaParams.block_table_batch_stride = p.max_blocks_per_sequence; + flashMlaParams.page_block_size = mCfg.tokens_per_block; + + flashMlaParams.descale_q_ptr = const_cast(params.dequant_scale_q); + flashMlaParams.descale_k_ptr = const_cast(params.dequant_scale_kv); + + flashMlaParams.tile_scheduler_metadata_ptr = tile_scheduler_metadata_ptr; + flashMlaParams.num_sm_parts = num_sm_parts; + flashMlaParams.num_splits_ptr = num_splits_ptr; + + flashMlaParams.softmax_lseaccum_ptr = softmax_lse_accum_ptr; + flashMlaParams.oaccum_ptr = out_accum_ptr; + + if constexpr (std::is_same::value) + { + if (mFP8GenerationMLA) + { + TLLM_THROW("FP8 KV cache MLA is only supported for bf16 output"); + } + else + { + run_mha_fwd_splitkv_mla(flashMlaParams, stream); + } + } + else if constexpr (std::is_same::value) + { + if (mFP8GenerationMLA) + { + run_mha_fwd_splitkv_mla(flashMlaParams, stream); + } + else + { + run_mha_fwd_splitkv_mla(flashMlaParams, stream); + } + } + else + { + TLLM_THROW("Unsupported data type for FlashMLA"); + } + } + else + { + // Try XQA optimization first. + // NOTE: input_seq_length = num_medusa_tokens + 1 (new generated one from the original LM head) + // self attn + XQAParams xqaParams{}; + this->template convertMMHAParamsToXQAParams( + xqaParams, p, /*forConfigurePlugin=*/false); + xqaParams.quant_q_buffer_ptr = params.quant_q_buf; + xqaParams.q_scaling + = 1 / (mCfg.q_scaling * sqrtf((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); + if (mEnableXQA && mXqaDispatcher->shouldUse(xqaParams)) + { + TLLM_LOG_DEBUG("XQA kernels are selected in the generation phase."); + xqaParams.stream = stream; + mXqaDispatcher->run(xqaParams, kv_cache_buffer, kv_scale_cache_buffer); + return 0; + } + + // Use FMHA otherwise. + MHARunnerParams fmhaParams{}; + fmhaParams.b = batch_beam; + fmhaParams.numGroupedHeads = params.head_num; + fmhaParams.qSeqLen = params.head_num * (params.acc_q_len / batch_beam); + fmhaParams.kvSeqLen = p.max_past_kv_length; + // Disable sliding window attention when it is not needed. + fmhaParams.slidingWindowSize = p.cyclic_attention_window_size; + fmhaParams.totalQSeqLen = batch_beam * fmhaParams.qSeqLen; + // TODO: set it correctly for contiguous kv buffer (cross-attention). + // fmhaParams.totalKvSeqLen = params.num_tokens; + // Device buffer pointers. + // fmhaParams.qkvPtr = reinterpret_cast(params.attention_input); + fmhaParams.qPtr = mFP8GenerationMLA ? reinterpret_cast(params.quant_q_buf) + : reinterpret_cast(params.q_buf); + // TODO: add contiguous kv buffer (cross-attention). + fmhaParams.kvPtr = nullptr; + + fmhaParams.outputPtr = reinterpret_cast(params.context_buf); + + // fmhaParams.packedMaskPtr = params.fmha_custom_mask; + fmhaParams.pagedKvCache = kv_cache_buffer; + fmhaParams.cuQSeqLenPtr = params.seqQOffset; + fmhaParams.kvSeqLenPtr = params.cache_seq_lens; + fmhaParams.cuKvSeqLenPtr = params.cu_kv_seqlens; + fmhaParams.cuMaskRowsPtr = nullptr; // mla not support custorm mask right now + fmhaParams.tileCounterPtr = params.fmha_tile_counter; + fmhaParams.scaleBmm1Ptr = reinterpret_cast(params.bmm1_scale); + fmhaParams.scaleBmm2Ptr = reinterpret_cast(params.bmm2_scale); + fmhaParams.stream = stream; + fmhaParams.forceFp32Acc = mFMHAForceFP32Acc; + + // Sparse attention parameters + if (useSparseMLA(p)) + { + fmhaParams.sparse_params = p.sparse_params; + } + + // MLA does not support skip-softmax attention right now + + // Run the fmha kernel + mDecoderFMHARunner->run(fmhaParams); + } + + sync_check_cuda_error(stream); + return 0; +} + +#define MLA_FUNC_DEFINE(T) \ + template int AttentionOp::mlaGeneration(MlaParams & params, FmhaParams const& p, cudaStream_t stream); + +MLA_FUNC_DEFINE(float) +MLA_FUNC_DEFINE(half) +#ifdef ENABLE_BF16 +MLA_FUNC_DEFINE(__nv_bfloat16) +#endif + +template +int AttentionOp::enqueueContext(FmhaParams const& p, MlaParams* mlaParam, cudaStream_t stream) +{ + int const headSize = getHeadSize(); + + int const local_hidden_units_qo = mCfg.num_heads * headSize; + int const local_hidden_units_kv = mNumAttnKVHeads * headSize; + PositionEmbeddingType const position_embedding_type = mCfg.position_embedding_type; + float const q_scaling = mCfg.q_scaling; + + KVCacheBuffer kv_cache_buffer; + KVCacheBuffer kv_scale_cache_buffer; + + auto sizePerToken = mNumAttnKVHeads * headSize * getKvCacheElemSizeInBits(p) / 8 /*bits*/; + + if (useKVCache(p)) + { + auto buffers = buildKvCacheBuffers(p.num_seqs, p.getMaxBlocksPerSequence(), + mCfg.tokens_per_block, sizePerToken, p.cyclic_attention_window_size, p.max_cyclic_attention_window_size, + p.sink_token_length, p.can_use_one_more_block, p.getHostPrimaryPoolPtr(), p.getHostSecondaryPoolPtr(), + p.getHostPrimaryBlockScalePoolPtr(), p.getHostSecondaryBlockScalePoolPtr(), + p.getKvCacheBlockOffsets(p.getKvCachePoolIndex(p.local_layer_idx)), mCfg.quant_mode.hasFp4KvCache(), + isCrossAttention(p) ? p.cross_kv_length : p.max_attention_window_size, p.getKeyValueCache()); + kv_cache_buffer = buffers.kvCacheBuffer; + kv_scale_cache_buffer = buffers.kvScaleCacheBuffer; + } + + auto cublasHandle = mCublasWrapper->getCublasHandle(); + TLLM_CUDA_CHECK(cublasSetStream(cublasHandle, stream)); + mCublasWrapper->setStream(stream); + mCublasWrapper->setWorkspace(p.getWorkspace()); + if constexpr (std::is_same_v) + { + mCublasWrapper->setFP16GemmConfig(); + } + else if constexpr (std::is_same_v) + { + mCublasWrapper->setFP32GemmConfig(); + } +#ifdef ENABLE_BF16 + else if constexpr (std::is_same_v) + { + mCublasWrapper->setBF16GemmConfig(); + } +#endif + + size_t const kv_seq_length = (isCrossAttention(p) ? p.cross_kv_length : p.input_seq_length); + size_t const attention_mask_size + = mEnableContextFMHA ? 0 : sizeof(T) * p.num_seqs * p.input_seq_length * kv_seq_length; + size_t const cu_seqlens_size = sizeof(int) * (p.num_seqs + 1); + size_t const rotary_inv_freq_size = sizeof(float) * p.num_seqs * mCfg.rotary_embedding_dim / 2; + size_t q_buf_2_size = 0; + if (!mEnableContextFMHA) + { + // Unfused mha + q_buf_2_size = sizeof(T) * p.num_seqs * p.input_seq_length * local_hidden_units_qo; + } + else if (mFmhaDispatcher->isSeparateQAndKvInput()) + { + // Paged context fmha + q_buf_2_size = (mFP8ContextFMHA ? 1 : sizeof(T)) * p.num_tokens * local_hidden_units_qo; + } + + size_t const k_buf_2_size = mEnableContextFMHA ? 0 : sizeof(T) * p.num_seqs * kv_seq_length * local_hidden_units_kv; + size_t const v_buf_2_size = mEnableContextFMHA ? 0 : sizeof(T) * p.num_seqs * kv_seq_length * local_hidden_units_kv; + size_t const qk_buf_size + = mEnableContextFMHA ? 0 : sizeof(T) * p.num_seqs * mCfg.num_heads * p.input_seq_length * kv_seq_length; + size_t const qkv_buf_2_size + = mEnableContextFMHA ? 0 : sizeof(T) * p.num_seqs * p.input_seq_length * local_hidden_units_qo; + size_t const qk_buf_float_size + = mEnableContextFMHA ? 0 : sizeof(float) * p.num_seqs * mCfg.num_heads * p.input_seq_length * kv_seq_length; + int dim_q_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); + int dim_k_per_head = (mMLAParams.qk_rope_head_dim + mMLAParams.qk_nope_head_dim); + int dim_v_per_head = (mMLAParams.v_head_dim); + if (useSparseMLA(p)) + { + dim_q_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + dim_k_per_head = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + dim_v_per_head + = mMLAParams.rope_append ? mMLAParams.kv_lora_rank : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + } + + // Total dimension per token across all heads for Q, K, and V components respectively + int const total_q_dim_all_heads = mNumAttnHeads * dim_q_per_head; + int const total_k_dim_all_heads + = mNumAttnHeads * dim_k_per_head; // Assuming effective num_kv_heads = head_num for layout + int const total_v_dim_all_heads + = mNumAttnHeads * dim_v_per_head; // Assuming effective num_kv_heads = head_num for layout + // Packed fp8 qkv buffer size for normal fp8 context FMHA + size_t fp8_qkv_buffer_size = mEnableContextFMHA && mFP8ContextFMHA && !mFmhaDispatcher->isSeparateQAndKvInput() + ? p.num_tokens * (local_hidden_units_qo + 2 * local_hidden_units_kv) + : 0; + // Separate fp8 q/k/v buffer size for fp8 context MLA + size_t fp8_q_buf_size = 0; + size_t fp8_k_buf_size = 0; + size_t fp8_v_buf_size = 0; + bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mCfg.is_mla_enable + && mFmhaDispatcher->isSeparateQAndKvInput() + && (p.fwd.sage_attn_num_elts_per_blk_q > 0 || p.fwd.sage_attn_num_elts_per_blk_k > 0 + || p.fwd.sage_attn_num_elts_per_blk_v > 0); + if (mEnableContextFMHA && mFP8ContextMLA && mFmhaDispatcher->isSeparateQAndKvInput()) + { + fp8_q_buf_size = p.num_tokens * static_cast(total_q_dim_all_heads); + + if (useSparseMLA(p)) + { + // Sparse MLA (absorption mode): K and V are stored directly in KV cache during MLA RoPE kernel. + // No separate FP8 buffers needed for K/V since they're read from paged KV cache (Q_PAGED_KV layout). + fp8_k_buf_size = 0; + fp8_v_buf_size = 0; + } + else + { + fp8_k_buf_size = p.total_kv_len * static_cast(total_k_dim_all_heads); + fp8_v_buf_size = p.total_kv_len * static_cast(total_v_dim_all_heads); + } + } + else if (useSageAttnSeparateQkv) + { + fp8_q_buf_size = p.num_tokens * static_cast(local_hidden_units_qo); + fp8_k_buf_size = p.total_kv_len * static_cast(local_hidden_units_kv); + fp8_v_buf_size = p.total_kv_len * static_cast(local_hidden_units_kv); + } + + int const q_max_n_blk = p.fwd.sage_attn_num_elts_per_blk_q > 0 + ? static_cast(tc::divUp(p.num_tokens, p.fwd.sage_attn_num_elts_per_blk_q) + p.num_seqs - 1) + : 0; + int const k_max_n_blk = p.fwd.sage_attn_num_elts_per_blk_k > 0 + ? static_cast(tc::divUp(p.total_kv_len, p.fwd.sage_attn_num_elts_per_blk_k) + p.num_seqs - 1) + : 0; + int const v_max_n_blk = p.fwd.sage_attn_num_elts_per_blk_v > 0 + ? static_cast(tc::divUp(local_hidden_units_kv, p.fwd.sage_attn_num_elts_per_blk_v)) + : 0; + size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); + size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); + size_t const sage_v_sfs_buffer_size = sizeof(float) * v_max_n_blk; + + size_t const padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * p.num_seqs * p.input_seq_length; + size_t const encoder_padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * p.num_seqs * p.cross_kv_length; + // Each token holds (batch_idx, token_idx_in_seq) int2. + size_t const tokens_info_size = sizeof(int2) * p.num_tokens; + size_t const fmha_scheduler_counter = mEnableContextFMHA ? sizeof(uint32_t) : 0; + size_t const fmha_bmm1_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) * 2 : 0; + size_t const fmha_bmm2_scale_size = (mFP8ContextFMHA || mFP8ContextMLA) ? sizeof(float) : 0; + + size_t const fmha_multi_ctas_kv_scratch_size = useTllmGenSparseAttention(p) ? getFmhaMultiCtasKvScratchSize(p) : 0; + + bool const is_qk_buf_float_ = true; + + AttentionContextWorkspaceSizes workspaceSizes{}; + workspaceSizes.attentionMask = attention_mask_size; + workspaceSizes.cuQSeqlens = cu_seqlens_size; + workspaceSizes.cuKvSeqlens = cu_seqlens_size; + workspaceSizes.cuMaskRows = cu_seqlens_size; + workspaceSizes.rotaryInvFreq = rotary_inv_freq_size; + workspaceSizes.qBuf = q_buf_2_size; + workspaceSizes.kBuf = k_buf_2_size; + workspaceSizes.vBuf = v_buf_2_size; + workspaceSizes.qkBuf = qk_buf_size; + workspaceSizes.qkvBuf = qkv_buf_2_size; + workspaceSizes.qkFloatBuf = qk_buf_float_size; + workspaceSizes.fp8QkvBuf = fp8_qkv_buffer_size; + workspaceSizes.fp8QBuf = fp8_q_buf_size; + workspaceSizes.fp8KBuf = fp8_k_buf_size; + workspaceSizes.fp8VBuf = fp8_v_buf_size; + workspaceSizes.paddingOffset = padding_offset_size; + workspaceSizes.encoderPaddingOffset = encoder_padding_offset_size; + workspaceSizes.tokensInfo = tokens_info_size; + workspaceSizes.fmhaTileCounter = fmha_scheduler_counter; + workspaceSizes.fmhaBmm1Scale = fmha_bmm1_scale_size; + workspaceSizes.fmhaBmm2Scale = fmha_bmm2_scale_size; + workspaceSizes.sageQScale = sage_q_sfs_buffer_size; + workspaceSizes.sageKScale = sage_k_sfs_buffer_size; + workspaceSizes.sageVScale = sage_v_sfs_buffer_size; + workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; + auto const workspaceLayout = AttentionWorkspaceManager::buildContextLayout(workspaceSizes); + auto const workspaceViews = AttentionWorkspaceManager::materializeContext(p.getWorkspace(), workspaceLayout); + + auto* fp8QBuf = workspaceViews.fp8QBuf; + // Fused FP8-Q path: caller pre-fills the nope segment of `quant_q_buf`; + // route the context-MLA Q pointer to it so the fused RoPE kernel appends + // rope FP8 in place and the FMHA Q load reads the merged [nope|rope] buffer. + if (mCfg.is_mla_enable && p.fwd.quant_q_buffer.has_value() && p.fwd.quant_scale_qkv.has_value() + && p.getQuantQBuffer() != nullptr) + { + fp8QBuf = reinterpret_cast<__nv_fp8_e4m3*>(p.getQuantQBuffer()); + } + + // build attention mask, cu_seqlens, and padding offset tensors + // Note: self attn and cross attn should use different p + // cross attn's seqlen info is from encoder input lengths, not decoder input lengths! + // moreover, attn mask for cross attn should be set separately (see below) + BuildDecoderInfoParams decoder_params{}; + int32_t const* precomputedCuQSeqlens = p.getCuQSeqlens(); + int32_t const* precomputedCuKvSeqlens = p.getCuKvSeqlens() != nullptr ? p.getCuKvSeqlens() + : p.getCuQSeqlens() != nullptr ? p.getCuQSeqlens() + : nullptr; + decoder_params.seqQOffsets = workspaceViews.cuQSeqlens; + decoder_params.seqKVOffsets = workspaceViews.cuKvSeqlens; + decoder_params.precomputedSeqQOffsets = precomputedCuQSeqlens; + decoder_params.precomputedSeqKVOffsets = precomputedCuKvSeqlens; + decoder_params.seqCpPartialOffsets = nullptr; + decoder_params.cpSize = 1; + decoder_params.packedMaskRowOffsets = workspaceViews.cuMaskRows; + decoder_params.paddingOffsets = workspaceViews.paddingOffset; + decoder_params.tokensInfo = workspaceViews.tokensInfo; + // Cross attention takes offsets from encoder inputs. + decoder_params.encoderPaddingOffsets = isCrossAttention(p) ? workspaceViews.encoderPaddingOffset : nullptr; + // Manually set attention mask for unfused cross attention. + decoder_params.attentionMask = isCrossAttention(p) ? nullptr : workspaceViews.attentionMask; + // Fixed sequence length offset if not removing the padding (seqQOffsets[i] = i * seq_length). + decoder_params.seqQLengths = p.getContextLengths(); + decoder_params.seqKVLengths = isCrossAttention(p) ? p.getEncoderInputLengths() : p.getSequenceLength(); + decoder_params.batchSize = p.num_seqs; + decoder_params.maxQSeqLength = p.input_seq_length; + decoder_params.maxEncoderQSeqLength + = isCrossAttention(p) ? p.cross_kv_length : 0; // cross attention uses encoder seq length + decoder_params.attentionWindowSize = p.cyclic_attention_window_size; + decoder_params.sinkTokenLength = p.sink_token_length; + decoder_params.numTokens = p.num_tokens; + decoder_params.removePadding = mCfg.remove_padding; + decoder_params.attentionMaskType = mCfg.mask_type; + decoder_params.blockSparseParams = p.block_sparse_params; + decoder_params.fmhaTileCounter = workspaceViews.fmhaTileCounter; + decoder_params.quantScaleO = p.getOutScale(); + decoder_params.dequantScaleQkv = p.getKvScaleQuantOrig(); + decoder_params.separateQkvScales = mCfg.quant_mode.hasFp4KvCache(); + decoder_params.fmhaHostBmm1Scale = 1.0f / (sqrtf(getHeadSize() * 1.0f) * q_scaling); + decoder_params.fmhaBmm1Scale = workspaceViews.fmhaBmm1Scale; + decoder_params.fmhaBmm2Scale = workspaceViews.fmhaBmm2Scale; + // Rotary embedding inv_freq buffer. + decoder_params.rotaryEmbeddingScale = p.rotary_embedding_scale; + decoder_params.rotaryEmbeddingBase = p.rotary_embedding_base; + decoder_params.rotaryEmbeddingDim = mCfg.rotary_embedding_dim; + decoder_params.rotaryScalingType = p.rotary_embedding_scale_type; + // The inv freq might be updated during runtime with dynamic scaling type. + decoder_params.rotaryEmbeddingInvFreq = workspaceViews.rotaryInvFreq; + // This is pre-computed when building the engines. + decoder_params.rotaryEmbeddingInvFreqCache = p.getRotaryInvFreq(); + decoder_params.rotaryEmbeddingMaxPositions = p.rotary_embedding_max_positions; + + invokeBuildDecoderInfo(decoder_params, stream); + sync_check_cuda_error(stream); + + int32_t const* contextCuQSeqlens + = precomputedCuQSeqlens != nullptr ? precomputedCuQSeqlens : workspaceViews.cuQSeqlens; + int32_t const* contextCuKvSeqlens + = precomputedCuKvSeqlens != nullptr ? precomputedCuKvSeqlens : workspaceViews.cuKvSeqlens; + + // In cross attention context phase, the attention mask should be a matrix of all ones. + // Override the attention mask produced by invokeBuildDecoderInfo(). + // also, invokeBuildDecoderInfo can only handle square mask, not cross B x q_len x kv_len mask + // TODO: put this logic in the kernel above. currently not much concern because q_len is mostly = 1 + if (isUnfusedCrossAttention(p)) + { + std::vector h_attention_mask(p.num_seqs * p.input_seq_length * p.cross_kv_length, 1.); + std::vector h_encoder_input_lengths(p.num_seqs); + tensorrt_llm::common::cudaMemcpyAsyncSanitized(h_encoder_input_lengths.data(), p.getEncoderInputLengths(), + sizeof(int32_t) * p.num_seqs, cudaMemcpyDeviceToHost, stream); + sync_check_cuda_error(stream); + + for (int bi = 0; bi < p.num_seqs; bi++) + { + int b_offset = bi * p.input_seq_length * p.cross_kv_length; + for (int qi = 0; qi < p.input_seq_length; qi++) + { + int q_offset = b_offset + qi * p.cross_kv_length; + if (h_encoder_input_lengths[bi] < p.cross_kv_length) + { + std::fill(h_attention_mask.begin() + q_offset + h_encoder_input_lengths[bi], + h_attention_mask.begin() + q_offset + p.cross_kv_length, 0.f); + } + } + } + cudaMemcpyAsync(workspaceViews.attentionMask, h_attention_mask.data(), + sizeof(T) * p.num_seqs * p.cross_kv_length * p.input_seq_length, cudaMemcpyHostToDevice, stream); + sync_check_cuda_error(stream); + } + + // FIXME: a temporary solution to make sure the padding part is 0. + if (!mCfg.remove_padding) + { + cudaMemsetAsync(p.getOutput(), 0, p.num_tokens * local_hidden_units_qo * sizeof(T), stream); + sync_check_cuda_error(stream); + } + + KvCacheDataType cache_type = cacheTypeFromQuantMode(mCfg.quant_mode); + + cudaDataType_t const gemm_data_type = tc::CudaDataType::value; + int const attention_seq_len_1 = p.input_seq_length; // q length + int const attention_seq_len_2 = isCrossAttention(p) ? p.cross_kv_length : p.input_seq_length; // kv length + + // If the model has relative attentiona bias, q scaling should be applied in QK gemm stage and use 1 in + // softamax stage (because to get softmax[scale(Q*K) + rel pos bias] here, q_scaling can't be applied during + // softmax phase by qk_scale); otherwise, use 1 in gemm stage and apply scaling in softmax stage + float const qk_scale + = 1.0f / (sqrtf(getHeadSize() * 1.0f) * q_scaling); // q_scaling in denominator. by default q_scaling =1.0f + float const qk_scale_gemm = isRelativePosition() ? qk_scale : 1.0f; + T const qk_scale_softmax = static_cast(isRelativePosition() ? 1.0f : qk_scale); + + // in context phase, currently FMHA runner has two restrictions: + // 1. only apply to self attention. If want fused multi-head cross attention, FMHCA kernels and runner is needed + // 2. doesn't apply to MHA with relative attention bias, i.e. softmax(QK + bias) * V + // We update mEnableContextFMHA in constructor to check these conditions + if (mEnableContextFMHA) + { + T* attention_input = p.getQkvOrQ(); + bool const enablePagedKVContextFMHA = mPagedKVCache && mCfg.paged_context_fmha; + TLLM_CHECK_WITH_INFO(!(mCfg.quant_mode.hasInt8KvCache() && enablePagedKVContextFMHA), + "Paged Context FMHA doesn't work with int8 kv cache currently."); + TLLM_CHECK_WITH_INFO(!(p.sink_token_length > 0 && enablePagedKVContextFMHA), + "Cannot support StreamingLLM now when enabling paged KV context FMHA."); + + // The max_kv_seq_len comes from the encoder seqlen when cross attention is used. + int const max_kv_seq_len = isCrossAttention(p) ? p.cross_kv_length : p.max_past_kv_length; + + // Prepare QKV preprocessing parameters. + QKVPreprocessingParams preprocessingParams; + + // Buffers. + preprocessingParams.qkv_input = const_cast(attention_input); + preprocessingParams.cross_kv_input = p.getCrossKv(); + preprocessingParams.quantized_qkv_output = workspaceViews.fp8QkvBuf; + preprocessingParams.q_output = workspaceViews.qBuf; + preprocessingParams.kv_cache_buffer = kv_cache_buffer; + preprocessingParams.kv_cache_block_scales_buffer = kv_scale_cache_buffer; + preprocessingParams.qkv_bias = p.getQkvBias(); + preprocessingParams.tokens_info = decoder_params.tokensInfo; + preprocessingParams.seq_lens = p.getContextLengths(); + // For self-attention, cache_seq_lens indicates whether chunked context is used + // (i.e. cache_seq_len > seq_len). + // For cross-attention, callers do not consistently use sequence_lengths as decoder length; use decoder + // context lengths so the encoder KV-cache write gate opens. + preprocessingParams.cache_seq_lens = isCrossAttention(p) ? p.getContextLengths() : p.getSequenceLength(); + + preprocessingParams.encoder_seq_lens = p.getEncoderInputLengths(); + preprocessingParams.cu_seq_lens = contextCuQSeqlens; + // Cross-attention only. + preprocessingParams.cu_kv_seq_lens = contextCuKvSeqlens; + preprocessingParams.rotary_embedding_inv_freq = workspaceViews.rotaryInvFreq; + preprocessingParams.rotary_coef_cache_buffer = p.getRotaryCosSin(); + preprocessingParams.mrope_rotary_cos_sin = p.getMropeRotaryCosSin(); + preprocessingParams.qkv_scale_orig_quant = p.getKvScaleOrigQuant(); + preprocessingParams.spec_decoding_position_offsets = nullptr; + preprocessingParams.helix_position_offsets = p.getHelixPositionOffsets(); + preprocessingParams.helix_is_inactive_rank = p.getHelixIsInactiveRank(); + preprocessingParams.logn_scaling = p.getLognScalingPtr(); + + // Sparse KV write + preprocessingParams.sparse_kv_indices = p.sparse_params.sparse_kv_indices; + preprocessingParams.sparse_kv_offsets = p.sparse_params.sparse_kv_offsets; + + // Scalars + preprocessingParams.batch_size = p.num_seqs; + preprocessingParams.max_input_seq_len = p.input_seq_length; + preprocessingParams.max_kv_seq_len = max_kv_seq_len; + preprocessingParams.cyclic_kv_cache_len + = isCrossAttention(p) ? p.cross_kv_length : p.cyclic_attention_window_size; + preprocessingParams.sink_token_len = p.sink_token_length; + preprocessingParams.token_num = p.num_tokens; + preprocessingParams.remove_padding = mCfg.remove_padding; + preprocessingParams.cross_attention = isCrossAttention(p); + preprocessingParams.head_num = mNumAttnHeads; + preprocessingParams.kv_head_num = mNumAttnKVHeads; + preprocessingParams.qheads_per_kv_head = mNumAttnHeads / mNumAttnKVHeads; + preprocessingParams.size_per_head = getHeadSize(); + preprocessingParams.rotary_embedding_dim = mCfg.rotary_embedding_dim; + preprocessingParams.rotary_embedding_base = p.rotary_embedding_base; + preprocessingParams.rotary_scale_type = p.rotary_embedding_scale_type; + preprocessingParams.rotary_embedding_scale = p.rotary_embedding_scale; + preprocessingParams.rotary_embedding_max_positions = p.rotary_embedding_max_positions; + preprocessingParams.position_embedding_type = position_embedding_type; + preprocessingParams.position_shift_enabled = p.pos_shift_enabled; + preprocessingParams.cache_type = cache_type; + preprocessingParams.separate_q_kv_output = enablePagedKVContextFMHA || isCrossAttention(p); + preprocessingParams.quantized_fp8_output = mFP8ContextFMHA; + preprocessingParams.generation_phase = false; + preprocessingParams.multi_processor_count = mMultiProcessorCount; + + preprocessingParams.rotary_vision_start = p.vision_start; + preprocessingParams.rotary_vision_length = p.vision_length; + preprocessingParams.is_last_chunk + = !p.attention_chunk_size.has_value() || (p.input_seq_length == p.max_past_kv_length); + + { + std::string const beforeRopeStr = "ctx attention before RoPE at layer " + std::to_string(p.layer_idx); + TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(p.num_tokens, + (local_hidden_units_qo + 2 * local_hidden_units_kv), p.getType(), + const_cast(attention_input), stream, beforeRopeStr) + == false, + "Found invalid number (NaN or Inf) in " + beforeRopeStr); + } + + if (mCfg.is_mla_enable) + { + TLLM_CHECK_WITH_INFO(mlaParam != nullptr, "MLA param is nullptr"); + mlaParam->cache_type = cache_type; + mlaParam->cu_q_seqlens = const_cast(contextCuQSeqlens); + mlaParam->cu_kv_seqlens = const_cast(contextCuKvSeqlens); + mlaParam->quant_scale_kv = p.getKvScaleOrigQuant(); + // Set BMM scales for FP8 context computation + mlaParam->bmm1_scale = workspaceViews.fmhaBmm1Scale; + mlaParam->bmm2_scale = workspaceViews.fmhaBmm2Scale; + mlaParam->quant_q_buf = mFP8ContextMLA ? fp8QBuf : nullptr; + mlaParam->quant_k_buf = mFP8ContextMLA ? workspaceViews.fp8KBuf : nullptr; + mlaParam->quant_v_buf = mFP8ContextMLA ? workspaceViews.fp8VBuf : nullptr; + // Set additional scales for context phase + mlaParam->quant_scale_o = p.getOutScale(); + mlaParam->quant_scale_q = p.getKvScaleOrigQuant(); + mlaParam->quant_scale_kv = p.getKvScaleOrigQuant(); + mlaParam->dequant_scale_q = p.getKvScaleQuantOrig(); + mlaParam->dequant_scale_kv = cache_type == KvCacheDataType::NVFP4 ? nullptr : p.getKvScaleQuantOrig(); + mlaParam->host_bmm1_scale + = 1 / (mCfg.q_scaling * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim))); + // The sparse MLA is in the absorption mode for the context phase. + mlaParam->absorption_mode = useSparseMLA(p); + // Fused FP8-Q-quant: RoPE kernel writes FP8 rope into `quant_q_buf`, + // so we skip the standalone invokeMLAContextFp8Quantize call below. + bool const useFusedQFp8 = mlaParam->fuse_q_fp8_in_rope && mFP8ContextMLA && mlaParam->absorption_mode + && cache_type == KvCacheDataType::FP8 && mlaParam->quant_q_buf != nullptr + && mlaParam->quant_scale_qkv != nullptr; + TLLM_CHECK_WITH_INFO(cache_type != KvCacheDataType::NVFP4 || mlaParam->latent_cache == nullptr, + "NVFP4 sparse MLA context must append its latent cache before launching attention"); + if (mlaParam->latent_cache != nullptr) + { + invokeMLARopeContext(*mlaParam, kv_cache_buffer, stream); + } + if (mFP8ContextMLA && !useFusedQFp8) + { + invokeMLAContextFp8Quantize(*mlaParam, p.total_kv_len, stream); + } + } + else if (useSageAttnSeparateQkv) + { + TLLM_CHECK_WITH_INFO(mFP8ContextFMHA, "SageAttention kernel runs under mFP8ContextFMHA option."); + TLLM_CHECK_WITH_INFO(mFmhaDispatcher->isSupported(), "SageAttention has no unfused fallback implemented."); + TLLM_CHECK_WITH_INFO(mCfg.mask_type == AttentionMaskType::PADDING, + "SageAttention only supports dense (padding) mask, got mask type %d.", + static_cast(mCfg.mask_type)); + TLLM_CHECK_WITH_INFO(p.fwd.sage_attn_num_elts_per_blk_q > 0 && p.fwd.sage_attn_num_elts_per_blk_k > 0 + && p.fwd.sage_attn_num_elts_per_blk_v == 1, + "SageQuant requires positive block sizes for Q and K while the block size for V must be 1."); + TLLM_CHECK_WITH_INFO(!p.fwd.kv_scale_quant_orig, + "SageAttention disregards the configured p.fwd.kv_scale_quant_orig, invalidating the result."); + check_cuda_error(cudaMemsetAsync(workspaceViews.sageVScale, 0, sage_v_sfs_buffer_size, stream)); + + // Common p for sageQuant + tc::SageQuantParams sageQuantParams{}; + sageQuantParams.headDim = getHeadSize(); + sageQuantParams.inputType = std::is_same_v ? DATA_TYPE_BF16 : DATA_TYPE_FP16; + sageQuantParams.quantType = p.fwd.sage_attn_qk_int8 ? DATA_TYPE_INT8 : DATA_TYPE_E4M3; + sageQuantParams.vStage = 0; + sageQuantParams.sumSeqLensV = p.total_kv_len; + sageQuantParams.numHeadsV = mNumAttnKVHeads; + sageQuantParams.ptrV = p.getV(); + sageQuantParams.ptrVQuant = workspaceViews.fp8VBuf; + sageQuantParams.ptrVScale = workspaceViews.sageVScale; + sageQuantParams.smCount = mMultiProcessorCount; + sageQuantParams.stream = stream; + + // Quantize into Fp8Q, SfsQ, SfsV + sageQuantParams.sumSeqLensQk = p.num_tokens; + sageQuantParams.batchSize = p.num_seqs; + sageQuantParams.numHeads = mNumAttnHeads; + sageQuantParams.tokenBlockSize = p.fwd.sage_attn_num_elts_per_blk_q; + sageQuantParams.ptrCuSeqLensQk = contextCuQSeqlens; + sageQuantParams.ptrQk = attention_input; + sageQuantParams.ptrQkQuant = workspaceViews.fp8QBuf; + sageQuantParams.ptrQkScale = workspaceViews.sageQScale; + sageQuantParams.vStage = 1; + tc::invokeSageQuant(sageQuantParams); + + // Quantize into Fp8K, SfsK, Fp8V + sageQuantParams.sumSeqLensQk = p.total_kv_len; + sageQuantParams.batchSize = p.num_seqs; + sageQuantParams.numHeads = mNumAttnKVHeads; + sageQuantParams.tokenBlockSize = p.fwd.sage_attn_num_elts_per_blk_k; + sageQuantParams.ptrCuSeqLensQk = contextCuKvSeqlens; + sageQuantParams.ptrQk = p.getK(); + sageQuantParams.ptrQkQuant = workspaceViews.fp8KBuf; + sageQuantParams.ptrQkScale = workspaceViews.sageKScale; + sageQuantParams.vStage = 2; + tc::invokeSageQuant(sageQuantParams); + } + else + { + invokeQKVPreprocessing(preprocessingParams, stream); + } + sync_check_cuda_error(stream); + { + std::string const afterRopeStr = "ctx attention after RoPE at layer " + std::to_string(p.layer_idx); + TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(p.num_tokens, + (local_hidden_units_qo + 2 * local_hidden_units_kv), p.getType(), + const_cast(attention_input), stream, afterRopeStr) + == false, + "Found invalid number (NaN or Inf) in " + afterRopeStr); + sync_check_cuda_error(stream); + } + + if (p.runtime_perf_knobs.has_value()) + { + int64_t const enable_context_fmha_fp32_acc_val = p.getRuntimePerfKnobs()[1]; + mFMHAForceFP32Acc = mFMHAForceFP32Acc || enable_context_fmha_fp32_acc_val == 1; + } + + // Unified FMHA runner interface for both packed QKV FMHA, contiguous Q_KV, paged KV FMHA, and separate QKV + // FMHA. + // Page KV input layout: + // - q_ptr: [B, S, H, D], which supports variable sequence length + // - paged_kv_cache: paged kv buffer + // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. + // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. + // + // Contiguous KV input layout: + // - q_ptr: [B, S, H, D], which supports variable sequence length + // - kv_ptr: [B, S, 2, H, D], which supports variable sequence length + // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. + // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. + // + // Separate QKV input layout (only for context MLA now): + // - q_ptr: [B, S, H, D], which supports variable sequence length + // - k_ptr: [B, S, H_kv, D], which supports variable sequence length + // - v_ptr: [B, S, H_kv, D_v], which supports variable sequence length + // - cu_q_seqlens: the cumulative query sequence lengths, needed for variable sequence length. + // - cu_kv_seqlens: the cumulative kv sequence lengths, needed for variable sequence length. + // - total_kv_len: the total kv sequence length, needed for variable sequence length. + + // Construct the fmha p for running kernels. + MHARunnerParams fmhaParams{}; + fmhaParams.b = p.num_seqs; + fmhaParams.qSeqLen = p.input_seq_length; + fmhaParams.kvSeqLen = max_kv_seq_len; + // Disable sliding window attention when it is not needed. + fmhaParams.slidingWindowSize + = (mCfg.dense_context_fmha || isCrossAttention(p)) ? max_kv_seq_len : p.cyclic_attention_window_size; + fmhaParams.totalQSeqLen = p.num_tokens; + // TODO: set it correctly for contiguous kv buffer (cross-attention). + fmhaParams.totalKvSeqLen = isCrossAttention(p) ? p.num_encoder_tokens : p.total_kv_len; + // Device buffer pointers. + if (mCfg.is_mla_enable) + { + // separate QKV input for context MLA + if (mFP8ContextMLA) + { + TLLM_CHECK_WITH_INFO( + mFmhaDispatcher->isSeparateQAndKvInput(), "Separate QKV input is required for fp8 context MLA"); + TLLM_CHECK_WITH_INFO(fp8QBuf != nullptr, "FP8 q buffer is required for fp8 context MLA"); + // In sparse MLA (absorption mode), K and V are stored in KV cache, not as separate FP8 buffers + TLLM_CHECK_WITH_INFO(useSparseMLA(p) || workspaceViews.fp8KBuf != nullptr, + "FP8 k buffer is required for fp8 context MLA in non-sparse mode"); + TLLM_CHECK_WITH_INFO(useSparseMLA(p) || workspaceViews.fp8VBuf != nullptr, + "FP8 v buffer is required for fp8 context MLA in non-sparse mode"); + + fmhaParams.qPtr = reinterpret_cast(fp8QBuf); + fmhaParams.kPtr = useSparseMLA(p) ? nullptr : reinterpret_cast(workspaceViews.fp8KBuf); + fmhaParams.vPtr = useSparseMLA(p) ? nullptr : reinterpret_cast(workspaceViews.fp8VBuf); + } + else + { + fmhaParams.qPtr = attention_input; + fmhaParams.kPtr = p.getK(); + fmhaParams.vPtr = p.getV(); + } + } + else if (useSageAttnSeparateQkv) + { + // SageAttention: use quantized FP8/INT8 Q/K/V buffers as separate inputs. + TLLM_CHECK_WITH_INFO( + mFmhaDispatcher->isSeparateQAndKvInput(), "Separate QKV input is required for sage attention FMHA"); + fmhaParams.qkvPtr = nullptr; + fmhaParams.qPtr = reinterpret_cast(workspaceViews.fp8QBuf); + fmhaParams.kPtr = reinterpret_cast(workspaceViews.fp8KBuf); + fmhaParams.vPtr = reinterpret_cast(workspaceViews.fp8VBuf); + // Set sage attention scaling factor pointers. + fmhaParams.qScalePtr = workspaceViews.sageQScale; + fmhaParams.kScalePtr = workspaceViews.sageKScale; + fmhaParams.vScalePtr = workspaceViews.sageVScale; + } + else + { + fmhaParams.qkvPtr = mFP8ContextFMHA ? reinterpret_cast(workspaceViews.fp8QkvBuf) + : reinterpret_cast(attention_input); + fmhaParams.qPtr = reinterpret_cast(workspaceViews.qBuf); + } + // TODO: add contiguous kv buffer (cross-attention). + fmhaParams.kvPtr = nullptr; + if (isCrossAttention(p) && !useKVCache(p)) + { + fmhaParams.kvPtr = p.getCrossKv(); + } + // Only use [totalLength, h / cpSize, Dh]. + fmhaParams.outputPtr = p.getOutput(); + fmhaParams.outputSfPtr = p.getOutputSf(); + if (mlaParam != nullptr && mlaParam->dsv4_epilogue_fusion.enabled) + { + fmhaParams.dsv4EpilogueFusion.enabled = true; + fmhaParams.dsv4EpilogueFusion.cosSinCache = mlaParam->dsv4_epilogue_fusion.cos_sin_cache; + fmhaParams.dsv4EpilogueFusion.scaleBufM = mlaParam->dsv4_epilogue_fusion.scale_buf_m; + } + fmhaParams.attentionSinksPtr = p.getAttentionSinks(); + fmhaParams.packedMaskPtr = p.getAttentionPackedMask(); + if constexpr (std::is_same_v) + { + fmhaParams.pagedKvCache = kv_cache_buffer; + fmhaParams.pagedKvSfCache = kv_scale_cache_buffer; + } + fmhaParams.cuQSeqLenPtr = contextCuQSeqlens; + fmhaParams.kvSeqLenPtr = decoder_params.seqKVLengths; + fmhaParams.cuKvSeqLenPtr = contextCuKvSeqlens; + fmhaParams.cuMaskRowsPtr = workspaceViews.cuMaskRows; + fmhaParams.tileCounterPtr = workspaceViews.fmhaTileCounter; + fmhaParams.scaleBmm1Ptr = workspaceViews.fmhaBmm1Scale; + fmhaParams.scaleBmm2Ptr = workspaceViews.fmhaBmm2Scale; + fmhaParams.oSfScalePtr = p.getOutSfScale(); + fmhaParams.stream = stream; + fmhaParams.forceFp32Acc = mFMHAForceFP32Acc; + fmhaParams.skipCorrectionThreshold = mSkipCorrectionThreshold; + fmhaParams.softmaxStatsPtr = p.getSoftmaxStatsTensor(); + fmhaParams.trtllmGenJITWarmup = p.trtllm_gen_jit_warmup; + fmhaParams.trtllmGenJITWarmupMaxNumRequests = p.max_num_requests; + fmhaParams.trtllmGenJITWarmupMaxSeqLenQ = p.max_context_length; + fmhaParams.trtllmGenJITWarmupMaxSeqLenKv = p.max_seq_len; + + // Sparse attention parameters + if (useTllmGenSparseAttention(p)) + { + fmhaParams.sparse_params = p.sparse_params; + // Sparse context reuses generation-style trtllm-gen kernels; provide the scratch pool + // and per-CTA counter so the autotuner can select MultiCtasKv variants. + fmhaParams.multiCtasKvScratchPtr = workspaceViews.fmhaMultiCtasKvScratch; + fmhaParams.multiCtasKvCounterPtr = static_cast(p.getMultiCtasKvCounter()); + } + + // Skip-softmax attention parameters + fmhaParams.skipSoftmaxThresholdScaleFactor = p.fwd.sparse_runtime_params.threshold_scale_factor_prefill; +#ifdef SKIP_SOFTMAX_STAT + fmhaParams.skipSoftmaxTotalBlocks = mSkipSoftmaxTotalBlocks; + fmhaParams.skipSoftmaxSkippedBlocks = mSkipSoftmaxSkippedBlocks; +#else + if (tensorrt_llm::common::getEnvPrintSkipSoftmaxStat()) + { + TLLM_THROW("To print skip softmax stat, please run build_wheel.py with -DSKIP_SOFTMAX_STAT"); + } +#endif + + if (p.attention_chunk_size) + { + fmhaParams.chunkedAttentionSize = *p.attention_chunk_size; + } + + // Run the fmha kernel. + mFmhaDispatcher->run(fmhaParams); + sync_check_cuda_error(stream); + + if (!mCfg.is_mla_enable) // Only for non-MLA attention + { + invokeKvCachePostprocessing(preprocessingParams, stream); + sync_check_cuda_error(stream); + } + } + else + { + TLLM_CHECK_DEBUG_WITH_INFO(p.getLognScalingPtr() == nullptr, "Unfused MHA does not support logn scaling"); + TLLM_CHECK_WITH_INFO(p.attention_chunk_size == std::nullopt, "Unfused MHA does not support chunked attention"); + // FIXME: a temporary solution to make sure the padding part of key/value buffer is 0 + // NOTE: pointer subtraction is used below since there could be some extra gap due to alignment. + // Otherwise, we could do cudaMemsetAsync(workspaceViews.kBuf, 0, k_buf_2_size + v_buf_2_size, stream). + // cudaMemsetAsync(workspaceViews.kBuf, 0, + // reinterpret_cast(workspaceViews.qkBuf) - reinterpret_cast(workspaceViews.kBuf), + // stream); + cudaMemsetAsync(workspaceViews.kBuf, 0, + reinterpret_cast(workspaceViews.vBuf) - reinterpret_cast(workspaceViews.kBuf) + + v_buf_2_size, + stream); + + if (!isCrossAttention(p)) + { + // self attention, write to from QKV to Q/K/V + invokeAddFusedQKVBiasTranspose(workspaceViews.qBuf, workspaceViews.kBuf, workspaceViews.vBuf, + p.getQkvOrQ(), p.getQkvBias(), p.getContextLengths(), + mCfg.remove_padding ? workspaceViews.paddingOffset : nullptr, p.num_seqs, p.input_seq_length, + p.num_tokens, mCfg.num_heads, mNumKVHeads, getHeadSize(), mCfg.rotary_embedding_dim, + p.rotary_embedding_base, p.rotary_embedding_scale_type, p.rotary_embedding_scale, + p.rotary_embedding_max_positions, position_embedding_type, (float*) nullptr, 0, stream); + sync_check_cuda_error(stream); + } + else + { + // cross attention, write from self QKV [*, head_num * head_size + 2 * kv_head_num * head_size]to Q, write + // from cross KV [*, 2 * kv_head_num * head_size] to K/V kernel modified accordingly to handle nullptr + // buffer + invokeAddFusedQKVBiasTranspose(workspaceViews.qBuf, (T*) nullptr, (T*) nullptr, p.getQkvOrQ(), + p.getQkvBias(), p.getContextLengths(), mCfg.remove_padding ? workspaceViews.paddingOffset : nullptr, + p.num_seqs, p.input_seq_length, p.num_tokens, mCfg.num_heads, mNumKVHeads, getHeadSize(), + mCfg.rotary_embedding_dim, p.rotary_embedding_base, p.rotary_embedding_scale_type, + p.rotary_embedding_scale, p.rotary_embedding_max_positions, position_embedding_type, (float*) nullptr, + 0, stream); + sync_check_cuda_error(stream); + + invokeAddFusedQKVBiasTranspose((T*) nullptr, workspaceViews.kBuf, workspaceViews.vBuf, p.getCrossKv(), + p.getQkvBias(), p.getEncoderInputLengths(), + mCfg.remove_padding ? workspaceViews.encoderPaddingOffset : nullptr, p.num_seqs, p.cross_kv_length, + p.num_encoder_tokens, /*mCfg.num_heads*/ 0, mNumKVHeads, getHeadSize(), mCfg.rotary_embedding_dim, + p.rotary_embedding_base, p.rotary_embedding_scale_type, p.rotary_embedding_scale, + p.rotary_embedding_max_positions, position_embedding_type, (float*) nullptr, 0, stream); + sync_check_cuda_error(stream); + } + + // write KV to cache + if (useKVCache(p)) + { + invokeTranspose4dBatchMajor(workspaceViews.kBuf, workspaceViews.vBuf, kv_cache_buffer, p.num_seqs, + isCrossAttention(p) ? p.cross_kv_length : p.input_seq_length, + isCrossAttention(p) ? p.cross_kv_length : p.cyclic_attention_window_size, getHeadSize(), mNumKVHeads, + cache_type, p.getKvScaleOrigQuant(), + isCrossAttention(p) ? p.getEncoderInputLengths() : p.getContextLengths(), stream); + } + sync_check_cuda_error(stream); + + T const* linear_bias_slopes = isALiBi() ? p.getAlibiSlopes() : nullptr; + T const* relative_attention_bias = isRelativePosition() ? p.getRelativeAttentionBias() : nullptr; + int const relative_attention_bias_stride = isRelativePosition() ? p.relative_attention_bias_stride : 0; + int const max_distance = p.fwd.relative_attention_max_distance; + cudaDataType_t gemm_out_data_type = is_qk_buf_float_ ? CUDA_R_32F : gemm_data_type; + void* gemm_out_buf_ = is_qk_buf_float_ ? static_cast(workspaceViews.qkFloatBuf) + : static_cast(workspaceViews.qkBuf); + if (mNumKVHeads == 1) // MQA + { + // Attn_weight[b, h*s_q, s_k] = Q[b, h*s_q, d] * K'[b, d, s_k] + // Attn_weight'[b, s_k, h*s_q] = K[b, s_k, d] * Q'[b, d, h*s_q] + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, + attention_seq_len_2, // n + attention_seq_len_1 * mCfg.num_heads, // m + getHeadSize(), // k + qk_scale_gemm, workspaceViews.kBuf, gemm_data_type, + getHeadSize(), // k + static_cast(attention_seq_len_2) * getHeadSize(), // n * k + workspaceViews.qBuf, gemm_data_type, + getHeadSize(), // k + static_cast(attention_seq_len_1) * mCfg.num_heads * getHeadSize(), // m * k + 0.0f, gemm_out_buf_, gemm_out_data_type, + attention_seq_len_2, // n + static_cast(attention_seq_len_1) * mCfg.num_heads * attention_seq_len_2, // m * n + p.num_seqs, // global batch size + CUDA_R_32F); + } + else if (mNumKVHeads == mCfg.num_heads) // MHA + { + // Attn_weight[b*h, s_q, s_k] = Q[b*h, s_q, d] * K'[b*h, d, s_k] + // Attn_weight'[b*h, s_k, s_q] = K[b*h, s_k, d] * Q'[b*h, d, s_q] + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, + attention_seq_len_2, // n + attention_seq_len_1, // m + getHeadSize(), // k + qk_scale_gemm, workspaceViews.kBuf, gemm_data_type, + getHeadSize(), // k + static_cast(attention_seq_len_2) * getHeadSize(), // n * k + workspaceViews.qBuf, gemm_data_type, + getHeadSize(), // k + static_cast(attention_seq_len_1) * getHeadSize(), // m * k + 0.0f, gemm_out_buf_, gemm_out_data_type, + attention_seq_len_2, // n + static_cast(attention_seq_len_2) * attention_seq_len_1, + p.num_seqs * mCfg.num_heads, // global batch size + CUDA_R_32F); + } + else // GQA + { + // Some number of contiguous Q heads will share the same K/V head + // Since the KV stride is NOT fixed for all Q, we have 2 options: + // 1. Loop over stridedBatchedGemm for each KV head. (multiple API calls/cuda kernels) + // 2. Calculate the pointers and use batchedGemm() (extra device memory) ::TODO:: + int const num_qheads_per_kv_head = mCfg.num_heads / mNumKVHeads; + for (int ki = 0; ki < mNumKVHeads; ++ki) + { + int64_t const q_offset + = static_cast(ki) * num_qheads_per_kv_head * attention_seq_len_1 * getHeadSize(); + int64_t const k_offset = static_cast(ki) * attention_seq_len_2 * getHeadSize(); + int64_t const qk_offset + = static_cast(ki) * attention_seq_len_1 * num_qheads_per_kv_head * attention_seq_len_2; + T* qptr = workspaceViews.qBuf + q_offset; + T* kptr = workspaceViews.kBuf + k_offset; + void* qkptr = is_qk_buf_float_ ? static_cast(workspaceViews.qkFloatBuf + qk_offset) + : static_cast(workspaceViews.qkBuf + qk_offset); + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_T, CUBLAS_OP_N, + attention_seq_len_2, // n + attention_seq_len_1 * num_qheads_per_kv_head, // m + getHeadSize(), // k + qk_scale_gemm, kptr, gemm_data_type, + getHeadSize(), // k + static_cast(mNumKVHeads) * attention_seq_len_2 * getHeadSize(), // n * k + qptr, gemm_data_type, + getHeadSize(), // k + static_cast(attention_seq_len_1) * mCfg.num_heads * getHeadSize(), // m * k + 0.0f, qkptr, gemm_out_data_type, + attention_seq_len_2, // n + static_cast(attention_seq_len_1) * mCfg.num_heads * attention_seq_len_2, // m * n + p.num_seqs, // global batch size + CUDA_R_32F); + } + } + + if (is_qk_buf_float_ == true) + { + // add relative position bias + if (isRelativePosition()) + { + // Add relative_attention_bias + // QK is (batch_size, local_head_num, q_length, k_length), relative_attention_bias is (1, + // local_head_num, max_output_len + 1, max_output_len + 1). broadcast along 1st dim. max_seq_len is + // already max_output_len + 1. In implicit mode, relative_attention_bias is relative_attention_table + // [num_heads, num_buckets], with necessary p (max_distance, num_buckets) passed at the end + invokeAddRelativeAttentionBiasUnaligned(workspaceViews.qkFloatBuf, relative_attention_bias, p.num_seqs, + mCfg.num_heads, attention_seq_len_1, + isCrossAttention(p) ? p.cross_kv_length : p.cyclic_attention_window_size, stream, max_distance > 0, + relative_attention_bias_stride, max_distance, false /* bidirectional */); + } + + MaskedSoftmaxParam param; + param.attention_score = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) + param.qk = workspaceViews.qkFloatBuf; // (batch_size, head_num, q_length, k_length) + param.attention_mask = workspaceViews.attentionMask; // (batch_size, q_length, k_length) + param.batch_size = p.num_seqs; + param.q_length = attention_seq_len_1; + param.k_length = attention_seq_len_2; + param.num_heads = mCfg.num_heads; + param.qk_scale = qk_scale_softmax; + param.attn_logit_softcapping_scale = mCfg.attn_logit_softcapping_scale; + param.attn_logit_softcapping_inverse_scale = 1.0f / mCfg.attn_logit_softcapping_scale; + param.linear_bias_slopes = const_cast(linear_bias_slopes); // (head_num,), optional + param.block_sparse_attn = mCfg.mask_type == AttentionMaskType::BLOCKSPARSE; + param.block_sparse_params = p.block_sparse_params; + param.q_seq_lengths = p.getContextLengths(); + invokeMaskedSoftmax(param, stream); + } + else + { + // add relative position bias + if (isRelativePosition()) + { + // Add relative_attention_bias + // QK is (batch_size, local_head_num, q_length, k_length), relative_attention_bias is (1, + // local_head_num, max_output_len + 1, max_output_len + 1). broadcast along 1st dim. max_seq_len is + // already max_output_len + 1. In implicit mode, relative_attention_bias is relative_attention_table + // [num_heads, num_buckets], with necessary p (max_distance, num_buckets) passed at the end + invokeAddRelativeAttentionBiasUnaligned(workspaceViews.qkBuf, relative_attention_bias, p.num_seqs, + mCfg.num_heads, attention_seq_len_1, + isCrossAttention(p) ? p.cross_kv_length : p.cyclic_attention_window_size, stream, max_distance > 0, + relative_attention_bias_stride, max_distance, false /* bidirectional */); + } + + MaskedSoftmaxParam param; + param.attention_score = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) + param.qk = workspaceViews.qkBuf; // (batch_size, head_num, q_length, k_length) + param.attention_mask = workspaceViews.attentionMask; // (batch_size, q_length, k_length) + param.batch_size = p.num_seqs; + param.q_length = attention_seq_len_1; + param.k_length = attention_seq_len_2; + param.num_heads = mCfg.num_heads; + param.qk_scale = qk_scale_softmax; + param.attn_logit_softcapping_scale = mCfg.attn_logit_softcapping_scale; + param.attn_logit_softcapping_inverse_scale = 1.0f / mCfg.attn_logit_softcapping_scale; + param.linear_bias_slopes = const_cast(linear_bias_slopes); // (head_num,), optional + param.block_sparse_attn = mCfg.mask_type == AttentionMaskType::BLOCKSPARSE; + param.block_sparse_params = p.block_sparse_params; + param.q_seq_lengths = p.getContextLengths(); + invokeMaskedSoftmax(param, stream); + } + + if (mNumKVHeads == 1) + { + // Attn_weight[b, h*s_q, s_k] + // O[b, h*s_q, d] = Attn_weight[b, h*s_q, s_k] * V[b, s_k, d] + // O'[b, d, h*s_q] = V'[b, d, s_k] * Attn_weight'[b, s_k, h*s_q] + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, + getHeadSize(), // n + mCfg.num_heads * attention_seq_len_1, // m + attention_seq_len_2, // k + workspaceViews.vBuf, + getHeadSize(), // n + static_cast(getHeadSize()) * attention_seq_len_2, // n * k + workspaceViews.qkBuf, + attention_seq_len_2, // k + static_cast(attention_seq_len_2) * mCfg.num_heads * attention_seq_len_1, // m * k + workspaceViews.qkvBuf, + getHeadSize(), // n + static_cast(getHeadSize()) * mCfg.num_heads * attention_seq_len_1, // n * m + p.num_seqs // global batch size + ); + } + else if (mNumKVHeads == mCfg.num_heads) // MHA + { + // O[b*h, s_q, d] = Attn_weight[b*h, s_q, s_k] * V[b*h, s_k, d] + // O'[b*h, d, s_q] = V'[b*h, d, s_k] * Attn_weight'[b*h, s_k, s_q] + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, getHeadSize(), attention_seq_len_1, + attention_seq_len_2, workspaceViews.vBuf, getHeadSize(), + static_cast(attention_seq_len_2) * getHeadSize(), workspaceViews.qkBuf, attention_seq_len_2, + static_cast(attention_seq_len_1) * attention_seq_len_2, workspaceViews.qkvBuf, getHeadSize(), + static_cast(attention_seq_len_1) * getHeadSize(), p.num_seqs * mCfg.num_heads); + } + else // GQA + { + // Attn_weight[b, h*s_q, s_k] + // O[b, h*s_q, d] = Attn_weight[b, h*s_q, s_k] * V[b, s_k, d] + // O'[b, d, h*s_q] = V'[b, d, s_k] * Attn_weight'[b, s_k, h*s_q] + int const num_qheads_per_kv_head = mCfg.num_heads / mNumKVHeads; + for (int ki = 0; ki < mNumKVHeads; ++ki) + { + int64_t const qk_offset + = static_cast(ki) * num_qheads_per_kv_head * attention_seq_len_1 * attention_seq_len_2; + int64_t const v_offset = static_cast(ki) * attention_seq_len_2 * getHeadSize(); + int64_t const qkv_offset + = static_cast(ki) * attention_seq_len_1 * num_qheads_per_kv_head * getHeadSize(); + T* qkptr = workspaceViews.qkBuf + qk_offset; + T* vptr = workspaceViews.vBuf + v_offset; + T* qkvptr = workspaceViews.qkvBuf + qkv_offset; + mCublasWrapper->stridedBatchedGemm(CUBLAS_OP_N, CUBLAS_OP_N, + getHeadSize(), // n + num_qheads_per_kv_head * attention_seq_len_1, // m + attention_seq_len_2, // k + vptr, + getHeadSize(), // n + static_cast(mNumKVHeads) * getHeadSize() * attention_seq_len_2, // n * k + qkptr, + attention_seq_len_2, // k + static_cast(attention_seq_len_2) * mCfg.num_heads * attention_seq_len_1, // m * k + qkvptr, + getHeadSize(), // n + static_cast(getHeadSize()) * mCfg.num_heads * attention_seq_len_1, // n * m + p.num_seqs // global batch size + ); + } + } + + if (!mCfg.remove_padding) + { + invokeTransposeQKV(static_cast(p.getOutput()), workspaceViews.qkvBuf, p.num_seqs, attention_seq_len_1, + mCfg.num_heads, getHeadSize(), (float*) nullptr, 0, stream); + } + else + { + invokeTransposeAttentionOutRemovePadding(workspaceViews.qkvBuf, static_cast(p.getOutput()), + p.num_tokens, p.num_seqs, attention_seq_len_1, mCfg.num_heads, getHeadSize(), + workspaceViews.paddingOffset, (float*) nullptr, 0, stream); + } + } + return 0; +} + +template int AttentionOp::enqueueContext( + FmhaParams const& p, MlaParams* mlaParam, cudaStream_t stream); + +template int AttentionOp::enqueueContext( + FmhaParams const& p, MlaParams* mlaParam, cudaStream_t stream); + +#ifdef ENABLE_BF16 +template int AttentionOp::enqueueContext<__nv_bfloat16, KVLinearBuffer>( + FmhaParams const& p, MlaParams<__nv_bfloat16>* mlaParam, cudaStream_t stream); +#endif + +template int AttentionOp::enqueueContext( + FmhaParams const& p, MlaParams* mlaParam, cudaStream_t stream); + +template int AttentionOp::enqueueContext( + FmhaParams const& p, MlaParams* mlaParam, cudaStream_t stream); + +#ifdef ENABLE_BF16 +template int AttentionOp::enqueueContext<__nv_bfloat16, KVBlockArray>( + FmhaParams const& p, MlaParams<__nv_bfloat16>* mlaParam, cudaStream_t stream); +#endif + +template +int AttentionOp::enqueueGeneration(FmhaParams const& p, cudaStream_t stream) +{ + int const headSize = getHeadSize(); + float const q_scaling = mCfg.q_scaling; + float const* logn_scaling_ptr = isLognScaling(p) ? p.getLognScalingPtr() : nullptr; + T const* relative_attention_bias = isRelativePosition() ? p.getRelativeAttentionBias() : nullptr; + int const relative_attention_bias_stride = isRelativePosition() ? p.relative_attention_bias_stride : 0; + int const max_distance = p.fwd.relative_attention_max_distance; + bool const* finished = nullptr; + + auto const quant_option = tc::QuantMode{}; + float const* qkv_scale_out = nullptr; + + int const* ia3_tasks = nullptr; + T const* ia3_key_weights = nullptr; + T const* ia3_value_weights = nullptr; + + int32_t const batch_beam = p.beam_width * p.num_requests; + + KVCacheBuffer kv_cache_buffer; + KVCacheBuffer kv_scale_cache_buffer; + + auto const sizePerToken = mNumAttnKVHeads * headSize * getKvCacheElemSizeInBits(p) / 8 /*bits*/; + + if (useKVCache(p)) + { + auto buffers = buildKvCacheBuffers(batch_beam, p.getMaxBlocksPerSequence(), + mCfg.tokens_per_block, sizePerToken, p.cyclic_attention_window_size, p.max_cyclic_attention_window_size, + p.sink_token_length, p.can_use_one_more_block, p.getHostPrimaryPoolPtr(), p.getHostSecondaryPoolPtr(), + p.getHostPrimaryBlockScalePoolPtr(), p.getHostSecondaryBlockScalePoolPtr(), + p.getKvCacheBlockOffsets(p.getKvCachePoolIndex(p.local_layer_idx)), mCfg.quant_mode.hasFp4KvCache(), + p.max_attention_window_size, p.getKeyValueCache()); + kv_cache_buffer = buffers.kvCacheBuffer; + kv_scale_cache_buffer = buffers.kvScaleCacheBuffer; + } + sync_check_cuda_error(stream); + + if (p.runtime_perf_knobs.has_value()) + { + int64_t const multi_block_mode_val = p.getRuntimePerfKnobs()[0]; + mMultiBlockMode = multi_block_mode_val == 1; + if (common::getEnvForceDeterministicAttention()) + { + mMultiBlockMode = false; + } + } + + if (common::getEnvForceDeterministicAttention()) + { + mMultiBlockMode = false; + } + + // TODO only for debug usage + if (!mMultiBlockMode) + { + char* isForceMultiBlockModeChar = std::getenv("FORCE_MULTI_BLOCK_MODE"); + bool isForceMultiBlockMode + = (isForceMultiBlockModeChar != nullptr && std::string(isForceMultiBlockModeChar) == "ON"); + TLLM_CHECK_WITH_INFO(!(common::getEnvForceDeterministicAttention() && isForceMultiBlockMode), + "FORCE_MULTI_BLOCK_MODE and FORCE_DETERMINISTIC/FORCE_ATTENTION_KERNEL_DETERMINISTIC can not be set at " + "the same time."); + mMultiBlockMode = isForceMultiBlockMode; + } + + // Check that the chunked-attention and sliding-window-attention are not enabled at the same time. + TLLM_CHECK_WITH_INFO(!p.attention_chunk_size.has_value() || p.cyclic_attention_window_size >= p.max_past_kv_length, + "Chunked-attention and sliding-window-attention should not be enabled at the same time."); + + T* attention_input = p.getQkvOrQ(); + // Try XQA optimization first. + { + // NOTE: input_seq_length = num_medusa_tokens + 1 (new generated one from the original LM head) + // self attn + XQAParams xqaParams{}; + this->template convertMMHAParamsToXQAParams(xqaParams, p, /*forConfigurePlugin=*/false); + + if (mEnableXQA && mXqaDispatcher->shouldUse(xqaParams)) + { + TLLM_LOG_DEBUG("XQA kernels are selected in the generation phase."); + xqaParams.stream = stream; + { + mXqaDispatcher->run(xqaParams, kv_cache_buffer, kv_scale_cache_buffer); + } + return 0; + } + else if (mCfg.is_spec_decoding_enabled && p.use_spec_decoding) + { + TLLM_CHECK_WITH_INFO(false, "No available XQA kernels are found for speculative decoding mode."); + } + else if (mFuseFp4Quant) + { + TLLM_CHECK_WITH_INFO(false, "No available kernels are found for FP4 output."); + } + else if (mCfg.quant_mode.hasFp4KvCache()) + { + TLLM_CHECK_WITH_INFO(false, "No available kernels are found for FP4 KV cache."); + } + else + { + TLLM_LOG_DEBUG("XQA kernels are not selected in the generation phase."); + } + } + + // This is the number of kv tokens that q needs to visit, but excluding one as it will be processed before the kv + // loop. + int timestep = p.max_past_kv_length; + int const max_timesteps = std::min(timestep, static_cast(p.cyclic_attention_window_size)); + int estimated_min_multi_block_count + = estimate_min_multi_block_count(max_timesteps, mMaxSharedMemoryPerBlockOptin - 2048, sizeof(T)); + + if (!mMultiBlockMode && !mForceMultiBlockWarned && estimated_min_multi_block_count > 1) + { + mForceMultiBlockWarned = true; + TLLM_LOG_WARNING( + "Force using MultiBlockMode in MMHA as shared memory is not enough, " + "MultiBlockMode may have different accuracy compared to non-MultiBlockMode."); + } + + // estimate min block count to satisfy shared memory requirement to run kernel. + // Runtime check to see the actual number of blocks per sequence we need. + int32_t const max_num_seq_len_tiles = std::max(getMaxNumSeqLenTile(p, batch_beam), estimated_min_multi_block_count); + int32_t const min_num_seq_len_tiles = std::max(1, estimated_min_multi_block_count); + bool const enable_multi_block + = (mMultiBlockMode && max_num_seq_len_tiles > 1) || estimated_min_multi_block_count > 1; + size_t const partial_out_size + = enable_multi_block ? sizeof(T) * batch_beam * mCfg.num_heads * mHeadSize * max_num_seq_len_tiles : 0; + size_t const partial_sum_size + = enable_multi_block ? sizeof(float) * batch_beam * mCfg.num_heads * max_num_seq_len_tiles : 0; + size_t const partial_max_size + = enable_multi_block ? sizeof(float) * batch_beam * mCfg.num_heads * max_num_seq_len_tiles : 0; + size_t const shift_k_cache_size = (!p.pos_shift_enabled || isCrossAttention(p)) + ? 0 + : sizeof(T) * batch_beam * mCfg.num_heads * mHeadSize * p.max_attention_window_size; + + AttentionGenerationWorkspaceSizes workspaceSizes{}; + workspaceSizes.partialOut = partial_out_size; + workspaceSizes.partialSum = partial_sum_size; + workspaceSizes.partialMax = partial_max_size; + workspaceSizes.shiftKCache = shift_k_cache_size; + { + auto const cascadeSizes + = tensorrt_llm::kernels::mmha::cascade::getCascadeWorkspaceSizes(batch_beam, mCfg.num_heads, mHeadSize); + workspaceSizes.cascadeOut = cascadeSizes.out; + workspaceSizes.cascadeMax = cascadeSizes.mMax; + workspaceSizes.cascadeSum = cascadeSizes.lSum; + } + auto const workspaceLayout = AttentionWorkspaceManager::buildGenerationLayout(workspaceSizes); + auto const workspaceViews = AttentionWorkspaceManager::materializeGeneration(p.getWorkspace(), workspaceLayout); + + // Apply position embedding to the keys in the K cache + KVLinearBuffer shift_k_cache_buffer; + if (useKVCache(p) && p.pos_shift_enabled && !isCrossAttention(p)) + { + shift_k_cache_buffer + = KVLinearBuffer(batch_beam, p.max_attention_window_size, sizePerToken, p.cyclic_attention_window_size, + p.sink_token_length, true, reinterpret_cast(workspaceViews.shiftKCache)); + sync_check_cuda_error(stream); + // KV cache type + KvCacheDataType const kv_cache_type = KvCacheDataType::BASE; + using DataType = typename SATypeConverter::Type; + invokeShiftKCache(kv_cache_buffer, shift_k_cache_buffer, kv_cache_type, getHeadSize(), + timestep, batch_beam, mNumKVHeads, p.beam_width, p.cyclic_attention_window_size, p.sink_token_length, + p.getKvScaleQuantOrig(), p.getSequenceLength(), p.getContextLengths(), mCfg.rotary_embedding_dim, + p.rotary_embedding_base, p.rotary_embedding_scale_type, p.rotary_embedding_scale, + p.rotary_embedding_max_positions, mCfg.position_embedding_type, stream); + } + + FusedQKVMaskedAttentionDispatchParams dispatch_params{}; + dispatch_params.mUnfuseQkvGemm = p.unfuse_qkv_gemm; + dispatch_params.qkv_buf = attention_input; + dispatch_params.qkv_bias = p.getQkvBias(); + dispatch_params.logn_scaling_ptr = logn_scaling_ptr; + dispatch_params.relative_attention_bias = relative_attention_bias; + dispatch_params.relative_attention_bias_stride = relative_attention_bias_stride; + dispatch_params.attention_mask = p.getAttentionMask(); + dispatch_params.attention_mask_stride = p.attention_mask_stride; + dispatch_params.attention_sinks = p.getAttentionSinks(); + dispatch_params.max_distance = max_distance; + dispatch_params.cache_indir = p.getCacheIndirection(); + dispatch_params.context_buf = p.getOutput(); // + dispatch_params.finished = finished; + dispatch_params.sequence_lengths + = p.getSequenceLength(); // NOTE: current seq len including padding (fixed after meeting the finished id) + dispatch_params.max_batch_size = batch_beam; + dispatch_params.inference_batch_size = batch_beam; + dispatch_params.beam_width = p.beam_width; + dispatch_params.head_num = mNumAttnHeads; + dispatch_params.kv_head_num = mNumAttnKVHeads; + dispatch_params.size_per_head = getHeadSize(); + dispatch_params.rotary_embedding_dim = mCfg.rotary_embedding_dim; + dispatch_params.position_embedding_type = mCfg.position_embedding_type; + dispatch_params.chunked_attention_size = p.attention_chunk_size ? *p.attention_chunk_size : INT_MAX; + dispatch_params.max_attention_window_size = p.max_attention_window_size; + dispatch_params.cyclic_attention_window_size = p.cyclic_attention_window_size; + dispatch_params.sink_token_length = isCrossAttention(p) ? 0 : p.sink_token_length; + dispatch_params.input_lengths = p.getContextLengths(); + dispatch_params.timestep = timestep; + dispatch_params.q_scaling = q_scaling; + dispatch_params.attn_logit_softcapping_scale = mCfg.attn_logit_softcapping_scale; + dispatch_params.linear_bias_slopes = isALiBi() ? p.getAlibiSlopes() : nullptr; + dispatch_params.ia3_tasks = ia3_tasks; + dispatch_params.ia3_key_weights = ia3_key_weights; + dispatch_params.ia3_value_weights = ia3_value_weights; + dispatch_params.qkv_scale_out = qkv_scale_out; + dispatch_params.fp8_context_fmha = mFP8ContextFMHA; + dispatch_params.attention_out_scale = p.getOutScale(); + dispatch_params.quant_option = quant_option; + dispatch_params.multi_block_mode = enable_multi_block; + dispatch_params.max_seq_len_tile = max_num_seq_len_tiles; + dispatch_params.min_seq_len_tile = min_num_seq_len_tiles; + dispatch_params.partial_out = workspaceViews.partialOut; + dispatch_params.partial_sum = workspaceViews.partialSum; + dispatch_params.partial_max = workspaceViews.partialMax; + dispatch_params.cascade_partial_out = workspaceViews.cascadeOut; + dispatch_params.cascade_partial_max = workspaceViews.cascadeMax; + dispatch_params.cascade_partial_sum = workspaceViews.cascadeSum; + dispatch_params.block_counter = static_cast(p.getMultiCtasKvCounter()); + dispatch_params.kv_cache_quant_mode = mCfg.quant_mode; + dispatch_params.kv_scale_orig_quant = p.getKvScaleOrigQuant(); + dispatch_params.kv_scale_quant_orig = p.getKvScaleQuantOrig(); + dispatch_params.kv_block_array = kv_cache_buffer; + dispatch_params.shift_k_cache_buffer = shift_k_cache_buffer; + dispatch_params.multi_processor_count = mMultiProcessorCount; + dispatch_params.rotary_embedding_base = p.rotary_embedding_base; + dispatch_params.rotary_embedding_scale_type = p.rotary_embedding_scale_type; + dispatch_params.rotary_embedding_scale = p.rotary_embedding_scale; + dispatch_params.rotary_embedding_inv_freq_cache = p.getRotaryInvFreq(); + dispatch_params.rotary_embedding_cos_sin_cache = p.getRotaryCosSin(); + dispatch_params.rotary_embedding_short_m_scale = p.rotary_embedding_short_mscale; + dispatch_params.rotary_embedding_long_m_scale = p.rotary_embedding_long_mscale; + dispatch_params.rotary_embedding_max_positions = p.rotary_embedding_max_positions; + dispatch_params.rotary_embedding_original_max_positions = p.rotary_embedding_original_max_positions; + dispatch_params.position_shift_enabled = p.pos_shift_enabled; + dispatch_params.rotary_cogvlm_vision_start = p.vision_start; + dispatch_params.rotary_cogvlm_vision_length = p.vision_length; + dispatch_params.cross_attention = isCrossAttention(p); + dispatch_params.memory_length_per_sample = p.getEncoderInputLengths(); + dispatch_params.block_sparse_attention = mCfg.mask_type == AttentionMaskType::BLOCKSPARSE; + dispatch_params.block_sparse_params = p.block_sparse_params; + dispatch_params.mrope_position_deltas = p.getMropePositionDeltas(); + + using DataType = typename SATypeConverter::Type; + { + if (!isCrossAttention(p)) + { + // self attn + Masked_multihead_attention_params mmha_params; + fusedQKV_masked_attention_dispatch(mmha_params, dispatch_params, stream); + } + else + { + // cross attn + Cross_multihead_attention_params mmhca_params; + fusedQKV_masked_attention_dispatch(mmhca_params, dispatch_params, stream); + } + sync_check_cuda_error(stream); + } + + return 0; +} + +template int AttentionOp::enqueueGeneration(FmhaParams const& p, cudaStream_t stream); + +template int AttentionOp::enqueueGeneration(FmhaParams const& p, cudaStream_t stream); + +#ifdef ENABLE_BF16 +template int AttentionOp::enqueueGeneration<__nv_bfloat16, KVLinearBuffer>(FmhaParams const& p, cudaStream_t stream); +#endif + +template int AttentionOp::enqueueGeneration(FmhaParams const& p, cudaStream_t stream); + +template int AttentionOp::enqueueGeneration(FmhaParams const& p, cudaStream_t stream); + +#ifdef ENABLE_BF16 +template int AttentionOp::enqueueGeneration<__nv_bfloat16, KVBlockArray>(FmhaParams const& p, cudaStream_t stream); +#endif + +template +void AttentionOp::prepareEnqueueGeneration(FmhaParams const& p) +{ + // self attn + if (mXqaDispatcher.get() != nullptr) + { + TLLM_LOG_TRACE("Preparing XQA kernels in prepareEnqueueGeneration."); + XQAParams xqaParams{}; + this->template convertMMHAParamsToXQAParams(xqaParams, p, /*forConfigurePlugin=*/true); + mXqaDispatcher->prepare(xqaParams); + } +} + +template void AttentionOp::prepareEnqueueGeneration(FmhaParams const& p); + +template void AttentionOp::prepareEnqueueGeneration(FmhaParams const& p); + +#ifdef ENABLE_BF16 +template void AttentionOp::prepareEnqueueGeneration<__nv_bfloat16, KVLinearBuffer>(FmhaParams const& p); +#endif + +template void AttentionOp::prepareEnqueueGeneration(FmhaParams const& p); + +template void AttentionOp::prepareEnqueueGeneration(FmhaParams const& p); + +#ifdef ENABLE_BF16 +template void AttentionOp::prepareEnqueueGeneration<__nv_bfloat16, KVBlockArray>(FmhaParams const& p); +#endif + +template +KvCacheBuffers buildKvCacheBuffers(int32_t batchSize, int32_t maxBlocksPerSeq, int32_t tokensPerBlock, + int32_t sizePerToken, int32_t cyclicAttentionWindowSize, int32_t maxCyclicAttentionWindowSize, int32_t sinkTokenLen, + bool canUseOneMoreBlock, void* primaryPoolPtr, void* secondaryPoolPtr, void* primaryBlockScalePoolPtr, + void* secondaryBlockScalePoolPtr, KVBlockArray::DataType* blockOffsets, bool hasFp4KvCache, + int32_t maxAttentionWindowSize, void* keyValueCache) +{ + KvCacheBuffers result; + if constexpr (std::is_same_v) + { + result.kvCacheBuffer = KVBlockArray(batchSize, maxBlocksPerSeq, tokensPerBlock, sizePerToken, + cyclicAttentionWindowSize, maxCyclicAttentionWindowSize, sinkTokenLen, canUseOneMoreBlock, primaryPoolPtr, + secondaryPoolPtr, blockOffsets); + if (hasFp4KvCache) + { + result.kvScaleCacheBuffer = KVBlockArray(batchSize, maxBlocksPerSeq, tokensPerBlock, sizePerToken / 8, + cyclicAttentionWindowSize, maxCyclicAttentionWindowSize, sinkTokenLen, canUseOneMoreBlock, + primaryBlockScalePoolPtr, secondaryBlockScalePoolPtr, blockOffsets); + } + } + else if constexpr (std::is_same_v) + { + TLLM_CHECK_WITH_INFO(!hasFp4KvCache, "FP4 KV cache only supports paged KV."); + TLLM_CHECK_WITH_INFO(keyValueCache != nullptr, "keyValueCache must not be null for linear KV cache."); + using BufferDataType = typename KVCacheBuffer::DataType; + result.kvCacheBuffer = KVLinearBuffer(batchSize, maxAttentionWindowSize, sizePerToken, + cyclicAttentionWindowSize, sinkTokenLen, false, reinterpret_cast(keyValueCache)); + } + return result; +} + +template KvCacheBuffers buildKvCacheBuffers(int32_t, int32_t, int32_t, int32_t, int32_t, + int32_t, int32_t, bool, void*, void*, void*, void*, KVBlockArray::DataType*, bool, int32_t, void*); + +template KvCacheBuffers buildKvCacheBuffers(int32_t, int32_t, int32_t, int32_t, int32_t, + int32_t, int32_t, bool, void*, void*, void*, void*, KVBlockArray::DataType*, bool, int32_t, void*); + +std::string AttentionOp::toString() const +{ + // Only the op's own state. The per-call parameters are the caller's to log; they + // change every call, whereas these are what the op derived once and dispatches on. + std::stringstream ss; + ss << std::boolalpha; +#define TRTLLM_DUMP_MEMBER(M) ss << #M ": " << M << "\n"; + TRTLLM_DUMP_MEMBER(mNumKVHeads) + TRTLLM_DUMP_MEMBER(mNumAttnHeads) + TRTLLM_DUMP_MEMBER(mNumAttnKVHeads) + TRTLLM_DUMP_MEMBER(mHeadSize) + TRTLLM_DUMP_MEMBER(mPagedKVCache) + TRTLLM_DUMP_MEMBER(mEnableContextFMHA) + TRTLLM_DUMP_MEMBER(mFMHAForceFP32Acc) + TRTLLM_DUMP_MEMBER(mMultiBlockMode) + TRTLLM_DUMP_MEMBER(mEnableXQA) + TRTLLM_DUMP_MEMBER(mFP8ContextFMHA) + TRTLLM_DUMP_MEMBER(mFP8AttenOutput) + TRTLLM_DUMP_MEMBER(mFP8ContextMLA) + TRTLLM_DUMP_MEMBER(mFP8GenerationMLA) + TRTLLM_DUMP_MEMBER(mFuseFp4Quant) + TRTLLM_DUMP_MEMBER(mIsGenerationMLA) + TRTLLM_DUMP_MEMBER(mUseGenFlashMLA) + TRTLLM_DUMP_MEMBER(mSM) + TRTLLM_DUMP_MEMBER(mUseTllmGen) + TRTLLM_DUMP_MEMBER(mMultiProcessorCount) + TRTLLM_DUMP_MEMBER(mMaxSharedMemoryPerBlockOptin) + TRTLLM_DUMP_MEMBER(mForceMultiBlockWarned) +#undef TRTLLM_DUMP_MEMBER + return ss.str(); +} + +namespace trtllm::attention +{ +using tensorrt_llm::kernels::KVBlockArray; +using tensorrt_llm::kernels::MlaParams; +using tensorrt_llm::torch_ext::AttentionOp; +using tensorrt_llm::torch_ext::KvCachePoolPointers; + +#ifdef ENABLE_BF16 +#define _DISPATCH_ON_DTYPE_BF16(FN, ...) \ + case tensorrt_llm::DataType::kBF16: FN<__nv_bfloat16>(__VA_ARGS__); break; +#else +#define _DISPATCH_ON_DTYPE_BF16(FN, ...) +#endif +#ifdef ENABLE_BF16 +#define _DISPATCH_ON_TORCH_DTYPE_BF16(FN, ...) \ + case torch::kBFloat16: FN<__nv_bfloat16>(__VA_ARGS__); break; +#else +#define _DISPATCH_ON_TORCH_DTYPE_BF16(FN, ...) +#endif + +// Dispatches straight off the tensor because the entry points run before prepare() +// derives FmhaParams::type. +#define DISPATCH_ON_TORCH_DTYPE(SCALAR_TYPE, FN, ...) \ + do \ + { \ + switch (SCALAR_TYPE) \ + { \ + case torch::kFloat32: FN(__VA_ARGS__); break; \ + case torch::kFloat16: \ + FN(__VA_ARGS__); \ + break; \ + _DISPATCH_ON_TORCH_DTYPE_BF16(FN, __VA_ARGS__) \ + default: TLLM_CHECK_WITH_INFO(false, "Unsupported attention dtype"); break; \ + } \ + } while (0) + +#define DISPATCH_ON_DTYPE(DTYPE, FN, ...) \ + do \ + { \ + switch (DTYPE) \ + { \ + case tensorrt_llm::DataType::kFLOAT: FN(__VA_ARGS__); break; \ + case tensorrt_llm::DataType::kHALF: \ + FN(__VA_ARGS__); \ + break; \ + _DISPATCH_ON_DTYPE_BF16(FN, __VA_ARGS__) \ + default: TLLM_CHECK_WITH_INFO(false, "Unsupported attention dtype"); break; \ + } \ + } while (0) + +} // namespace trtllm::attention + +template +void FmhaParams::finalizeMlaParams(MlaParams& mla, AttentionOp const& op) const +{ + mla.q_buf = getQkvOrQ(); + mla.context_buf = static_cast(getOutput()); + // MlaParams has no default member initializers for these two, and the MLA RoPE + // kernels use max_input_seq_len as a grid dimension and dereference cache_seq_lens. + // Both phases must set them. + mla.cache_seq_lens = getSequenceLength(); + mla.max_input_seq_len = input_seq_length; + + mla.cos_sin_cache = op.isRoPE() ? getRotaryCosSin() : nullptr; + if (fwd.enable_dsv4_epilogue_fusion) + { + TORCH_CHECK( + fwd.dsv4_inv_rope_cos_sin_cache.has_value(), "DSv4 fused epilogue requires inverse-RoPE cos/sin cache."); + auto const& cosSinCache = fwd.dsv4_inv_rope_cos_sin_cache.value(); + auto const& outputSfTensor = fwd.output_sf.value(); + TORCH_CHECK(cosSinCache.scalar_type() == torch::kFloat32, "DSv4 fused epilogue cos/sin cache must be float32."); + TORCH_CHECK(output.scalar_type() == torch::kFloat8_e4m3fn, "DSv4 fused epilogue output must be float8_e4m3fn."); + TORCH_CHECK(output.dim() == 3 && output.is_contiguous(), + "DSv4 fused epilogue output must be contiguous [groups, tokens, K]."); + TORCH_CHECK( + outputSfTensor.scalar_type() == torch::kFloat32, "DSv4 fused epilogue fwd.output_sf must be float32."); + TORCH_CHECK(outputSfTensor.dim() == 3 && outputSfTensor.is_contiguous(), + "DSv4 fused epilogue fwd.output_sf must be contiguous [groups, K/128, padded_tokens]."); + TORCH_CHECK(output.size(1) >= num_tokens, "DSv4 fused epilogue output token dimension is too small."); + TORCH_CHECK(op.mlaMeta().v_head_dim > 0 && op.mlaMeta().v_head_dim % 128 == 0, + "DSv4 fused epilogue requires v_head_dim to be a positive multiple of 128."); + TORCH_CHECK( + outputSfTensor.size(2) >= num_tokens, "DSv4 fused epilogue fwd.output_sf token dimension is too small."); + + mla.dsv4_epilogue_fusion.enabled = true; + mla.dsv4_epilogue_fusion.cos_sin_cache = getDsv4InvRopeCosSinCache(); + mla.dsv4_epilogue_fusion.scale_buf_m = static_cast(outputSfTensor.size(2)); + } + mla.batch_size = num_seqs; + mla.acc_q_len = num_tokens; + mla.head_num = op.config().num_heads; + mla.meta = op.mlaMeta(); + + mla.workspace = getWorkspace(); +} + +template +MlaParams FmhaParams::buildContextMlaParams(AttentionOp const& op) const +{ + MlaParams mla{}; + if (use_sparse_attention) + { + mla.latent_cache = getLatentCache(); + TORCH_CHECK(fwd.q_pe.has_value()); + TORCH_CHECK(fwd.q_pe->dim() == 3); + TORCH_CHECK(fwd.q_pe->strides()[2] == 1); + + mla.q_pe = getQPe(); + mla.q_pe_ld = fwd.q_pe->strides()[1]; + mla.q_pe_stride = fwd.q_pe->strides()[0]; + + // Fused FP8-Q path: forward caller's fwd.quant_q_buffer / scale so + // applyMLARopeAndAssignQKVKernelOptContext + // appends rope FP8 in place and the standalone quantize is + // skipped. Without this wiring the sparse-MLA context branch + // runs the legacy quantize over the bf16 placeholder q. + mla.bmm1_scale = getMlaBmm1Scale(); + mla.bmm2_scale = getMlaBmm2Scale(); + mla.quant_q_buf = getQuantQBuffer(); + mla.quant_scale_qkv = getQuantScaleQkv(); + mla.fuse_q_fp8_in_rope = (fwd.quant_q_buffer.has_value() && fwd.quant_scale_qkv.has_value()); + + // Fused kv_a_layernorm: the norm weight implies `latent_cache` is the + // RAW kv_a_proj output, with the caller's RMSNorm and concat dropped. + if (fwd.kv_norm_weight) + { + auto const& kvNormWeight = fwd.kv_norm_weight.value(); + TORCH_CHECK(kvNormWeight.is_cuda(), "kv_norm_weight must be a CUDA tensor"); + TORCH_CHECK(kvNormWeight.is_contiguous(), "kv_norm_weight must be contiguous"); + TORCH_CHECK(kvNormWeight.scalar_type() == qkv_or_q.scalar_type(), + "kv_norm_weight dtype must match the activation dtype"); + TORCH_CHECK(fwd.latent_cache, "fused kv-norm needs latent_cache (the raw kv_a_proj output) to be provided"); + // The kernel norms the whole latent row, so a narrower weight would + // read out of bounds. dsv3RopeOp.cpp checks the same on the + // generation side. + auto const kvNormWidth = op.mlaMeta().kv_lora_rank + op.mlaMeta().qk_rope_head_dim; + TORCH_CHECK(kvNormWeight.numel() == kvNormWidth, + "kv_norm_weight must span kv_lora_rank + qk_rope_head_dim (", kvNormWidth, "), got ", + kvNormWeight.numel()); + // A last-dim slice, so rows are wider than the row itself. Forward + // the real stride; only the innermost dim must be unit-stride. + auto const& latentCache = fwd.latent_cache.value(); + TORCH_CHECK( + latentCache.dim() == 2, "latent_cache must be 2D for fused kv-norm, got ", latentCache.dim(), "D"); + TORCH_CHECK(latentCache.stride(1) == 1, "latent_cache must be unit-stride in its last dim"); + // The kernel walks rows with 16-byte vector loads, so a row start + // that is not 16B-aligned faults with a bare misaligned-address + // error far from here. + auto const kEltsPer16B = 16 / latentCache.element_size(); + TORCH_CHECK(latentCache.stride(0) % kEltsPer16B == 0, "latent_cache row stride (", latentCache.stride(0), + ") must be a multiple of ", kEltsPer16B, " for the fused kv-norm 16B vector loads"); + TORCH_CHECK(reinterpret_cast(latentCache.data_ptr()) % 16 == 0, + "latent_cache must be 16B-aligned for the fused kv-norm vector loads"); + mla.latent_row_stride = static_cast(latentCache.stride(0)); + mla.kv_norm_weight = static_cast(kvNormWeight.data_ptr()); + mla.kv_norm_eps = static_cast(fwd.kv_norm_eps); + mla.fuse_kv_norm_in_rope = true; + } + } + else + { + mla.latent_cache = getLatentCache(); + TORCH_CHECK(k.has_value()); + TORCH_CHECK(v.has_value()); + TORCH_CHECK(k->dim() == 2); + TORCH_CHECK(v->dim() == 2); + TORCH_CHECK(k->strides()[1] == 1); + TORCH_CHECK(v->strides()[1] == 1); + + mla.k_buf = getK(); + mla.v_buf = getV(); + + mla.helix_position_offsets = getHelixPositionOffsets(); + mla.helix_is_inactive_rank = getHelixIsInactiveRank(); + } + finalizeMlaParams(mla, op); + return mla; +} + +template +MlaParams FmhaParams::buildGenerationMlaParams(AttentionOp const& op) const +{ + MlaParams mla{}; + TORCH_CHECK(fwd.latent_cache.has_value()); + mla.latent_cache = getLatentCache(); + TORCH_CHECK(fwd.q_pe.has_value()); + TORCH_CHECK(fwd.q_pe->dim() == 3); + TORCH_CHECK(fwd.q_pe->strides()[2] == 1); + + mla.q_pe = getQPe(); + mla.q_pe_ld = fwd.q_pe->strides()[1]; + mla.q_pe_stride = fwd.q_pe->strides()[0]; + + mla.seqQOffset = const_cast(getCuQSeqlens()); + mla.cu_kv_seqlens = const_cast(getCuKvSeqlens()); + mla.fmha_tile_counter = getFmhaSchedulerCounter(); + mla.bmm1_scale = getMlaBmm1Scale(); + mla.bmm2_scale = getMlaBmm2Scale(); + mla.quant_q_buf = getQuantQBuffer(); + mla.quant_scale_qkv = getQuantScaleQkv(); + mla.fuse_q_fp8_in_rope = (fwd.quant_q_buffer.has_value() && fwd.quant_scale_qkv.has_value()); + finalizeMlaParams(mla, op); + return mla; +} + +template +void FmhaParams::addFlashMlaGenerationParams(MlaParams& mla) const +{ + TORCH_CHECK(block_ids_per_seq.has_value()); + mla.block_ids_per_seq = getBlockIdsPerSeq(); + if (flash_mla_tile_scheduler_metadata.has_value()) + { + TORCH_CHECK(flash_mla_num_splits.has_value(), + "flash_mla_num_splits must be provided when flash_mla_tile_scheduler_metadata is set."); + mla.flash_mla_tile_scheduler_metadata = getFlashMlaTileSchedulerMetadata(); + mla.flash_mla_num_splits = getFlashMlaNumSplits(); + } +} + +template +void AttentionOp::runContextImpl(FmhaParams& p) +{ + prepare(p, /*isGen=*/false); + // Without encoder K/V input, cross attention must read the cached encoder KV. + TLLM_CHECK_WITH_INFO(p.fwd.cross_kv.has_value() || !isUnfusedCrossAttention(p), + "Cross attention without encoder K/V input requires a fused context FMHA kernel to read the " + "cached cross KV, and this build has none for this configuration. Disable chunked prefill for " + "this model, or use a build whose --cuda_architectures includes this device's SM."); + auto const stream = at::cuda::getCurrentCUDAStream(p.qkv_or_q.get_device()); + MlaParams mla{}; + MlaParams* mlaParam = nullptr; + if (isMLAEnabled()) + { + mla = p.buildContextMlaParams(*this); + mlaParam = &mla; + } + enqueueContext(p, mlaParam, stream); + sync_check_cuda_error(stream); +} + +template +void AttentionOp::runGenerationImpl(FmhaParams& p) +{ + prepare(p, /*isGen=*/true); + auto const stream = at::cuda::getCurrentCUDAStream(p.qkv_or_q.get_device()); + enqueueGeneration(p, stream); + { + std::string const afterGenStr = "gen attention at layer " + std::to_string(p.layer_idx); + TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(p.num_tokens, + mCfg.num_heads * mCfg.head_size, p.getType(), p.getOutput(), stream, afterGenStr) + == false, + "Found invalid number (NaN or Inf) in " + afterGenStr); + } + sync_check_cuda_error(stream); +} + +template +void AttentionOp::runMlaGenerationImpl(FmhaParams& p) +{ + prepare(p, /*isGen=*/true); + auto const stream = at::cuda::getCurrentCUDAStream(p.qkv_or_q.get_device()); + auto mla = p.buildGenerationMlaParams(*this); + if (mUseGenFlashMLA) + { + p.addFlashMlaGenerationParams(mla); + } + mlaGeneration(mla, p, stream); + { + std::string const afterGenStr = "mla gen attention at layer " + std::to_string(p.layer_idx); + TLLM_CHECK_DEBUG_WITH_INFO(tensorrt_llm::runtime::utils::tensorHasInvalid(p.num_tokens, + mCfg.num_heads * mCfg.head_size, p.getType(), p.getOutput(), stream, afterGenStr) + == false, + "Found invalid number (NaN or Inf) in " + afterGenStr); + } + sync_check_cuda_error(stream); +} + +AttentionOp::AttentionOp(StaticAttentionConfig const& cfg) + : mDriver(CUDADriverWrapper::getInstance()) + , mCublasWrapper(new tc::CublasMMWrapper(getCublasHandle(), getCublasLtHandle(), nullptr, nullptr)) +{ + // Reading these caches the parsed values, keeping getenv() off the per-call path. + getEnvMmhaMultiblockDebug(); + getEnvMmhaBlocksPerSequence(); + + mNumKVHeads = static_cast(cfg.num_kv_heads); + mHeadSize = static_cast(cfg.head_size); + if (cfg.is_mla_enable) + { + mIsGenerationMLA = cfg.head_size == cfg.kv_lora_rank + cfg.qk_rope_head_dim; + mUseGenFlashMLA = mSM == 90 && cfg.tokens_per_block == 64 && cfg.head_size == 576; + mNumKVHeads = 1; + mHeadSize = cfg.kv_lora_rank + cfg.qk_rope_head_dim; + } + + auto constexpr kMaxSkipCorrectionThreshold = 32.0; + TLLM_CHECK_WITH_INFO( + cfg.skip_correction_threshold >= 0.0 && cfg.skip_correction_threshold <= kMaxSkipCorrectionThreshold, + "skip_correction_threshold must be in the range (0, 32] when enabled, or 0 when disabled."); + bool const applySkipCorrection = cfg.is_mla_enable && (mSM == 100 || mSM == 103); + mSkipCorrectionThreshold = applySkipCorrection ? static_cast(cfg.skip_correction_threshold) : 0.0F; + + // One rank per attention op: neither tensor nor context parallelism applies here. + mNumAttnHeads = static_cast(cfg.num_heads); + mNumAttnKVHeads = mNumKVHeads; + + mCfg = cfg; + mMLAParams = {static_cast(cfg.q_lora_rank), static_cast(cfg.kv_lora_rank), + static_cast(cfg.qk_nope_head_dim), static_cast(cfg.qk_rope_head_dim), + static_cast(cfg.v_head_dim), static_cast(cfg.predicted_tokens_per_seq), + static_cast(cfg.mla_layer_num), static_cast(cfg.rope_append)}; + mUseNvfp4MlaKvCache = cfg.use_nvfp4_mla_kv_cache; + initialize(); + + // Unfused self-attention cannot read a cached prefix. Cross attention can run + // unfused when encoder K/V is supplied; runContextImpl checks that per call. + bool const needsFusedPagedContext + = cfg.paged_context_fmha && cfg.use_kv_cache && !cfg.is_mla_enable && !cfg.cross_attention; + TLLM_CHECK_WITH_INFO(!needsFusedPagedContext || !isRelativePosition(), + "Paged-context attention (chunked prefill, KV cache reuse or speculative draft tokens) is not supported with " + "relative position embedding: that attention always runs unfused, and the unfused path cannot attend to " + "cached KV. Disable chunked prefill, KV cache reuse and speculative decoding for this model."); + TLLM_CHECK_WITH_INFO(!needsFusedPagedContext || mEnableContextFMHA, + "Paged-context attention requires a fused context FMHA kernel, and this build has none for this " + "configuration. The unfused fallback cannot attend to cached KV. If the device's SM is not named in the " + "build's --cuda_architectures then the build carries no kernels for it at all; check that first. Otherwise " + "use another attention backend, or disable chunked prefill, KV cache reuse and speculative decoding."); +} + +int AttentionOp::prepare(FmhaParams& p, bool isGen) +{ + + // Phase extents. Python sends the phase's token and row counts, which only it can + // know before the phase split; the rest is a reduction over the host length arrays, + // which the op can do itself. + p.max_past_kv_length = p.getMaxHostPastKeyValueLength(p.seq_offset, p.num_seqs); + p.total_kv_len = p.getTotalKvLen(p.seq_offset, p.num_seqs, isGen); + p.input_seq_length = isGen ? p.num_tokens / std::max(p.num_seqs, 1) + : p.getMaxHostContextLength(p.seq_offset, p.num_seqs); + // Encoder CUDA graphs capture a padded context extent; the override widens both. + if (!isGen && p.max_context_q_len_override.has_value()) + { + p.input_seq_length = p.max_context_q_len_override.value(); + p.max_past_kv_length = p.max_context_q_len_override.value(); + } + + if (p.spec_decoding_target_max_draft_tokens.has_value() && p.spec_decoding_target_max_gen_len == 0) + { + p.spec_decoding_target_max_gen_len = static_cast(p.spec_decoding_target_max_draft_tokens.value()) + 1; + } + + bool const hasSparseAttnIndices = p.fwd.sparse_runtime_params.sparse_attn_indices.has_value() + && p.fwd.sparse_runtime_params.sparse_attn_indices.value().numel() > 0; + p.use_sparse_attention = (p.fwd.sparse_runtime_params.sparse_kv_indices.has_value() + && p.fwd.sparse_runtime_params.sparse_kv_indices.value().numel() > 0) + || hasSparseAttnIndices; + + if (mCfg.is_mla_enable) + { + if (p.num_sparse_topk > 0 && hasSparseAttnIndices) + { + p.use_sparse_attention = true; + } + TLLM_CHECK(!p.isFp4Out()); + mMLAParams = {static_cast(mCfg.q_lora_rank), static_cast(mCfg.kv_lora_rank), + static_cast(mCfg.qk_nope_head_dim), static_cast(mCfg.qk_rope_head_dim), + static_cast(mCfg.v_head_dim), static_cast(mCfg.predicted_tokens_per_seq), + static_cast(p.getMlaLayerNum()), static_cast(mCfg.rope_append)}; + } + + // Commonly identical to cyclic_attention_window_size unless layers use different + // attention window sizes; beam search may consume one extra block. + p.max_cyclic_attention_window_size = p.cyclic_attention_window_size; + p.can_use_one_more_block = p.beam_width > 1; + + // Block-table stride, i.e. the trailing dimension of kv_cache_block_offsets. + p.max_blocks_per_sequence = p.getMaxBlocksPerSequence(); + + // Generation consumers dereference the counter without a null check -- MMHA at + // `params.block_counter[bhi]`, XQA at `semaphores[idxSeq]` -- so a caller that + // forgets it produces an illegal memory access with no host-side stack. The phased + // run_* ops are a new caller-facing boundary, so name the missing buffer here. + TLLM_CHECK_WITH_INFO(!isGen || p.multi_ctas_kv_counter.has_value(), + "Generation attention requires semaphores; the FMHA library owns this buffer."); + + // `out_scale` is the output scale for FP8 output, but the global scale for the + // scaling factors when the NVFP4 quant epilogue is fused; the two feed different + // kernel arguments and must not be mixed up. + if (p.isFp4Out()) + { + p.out_sf_scale = p.fwd.out_scale; + p.fwd.out_scale = std::nullopt; + } + + if (p.fwd.attention_sinks.has_value()) + { + TORCH_CHECK(p.fwd.attention_sinks.value().scalar_type() == torch::kFloat32, + "Expected attention_sinks to have float dtype"); + } + if (mCfg.quant_mode.hasFp4KvCache()) + { + TORCH_CHECK(!p.fwd.kv_scale_orig_quant.has_value() || p.fwd.kv_scale_orig_quant.value().size(0) == 3, + "FP4 KV cache expects kv_scale_orig_quant to have 3 elements."); + TORCH_CHECK(!p.fwd.kv_scale_quant_orig.has_value() || p.fwd.kv_scale_quant_orig.value().size(0) == 3, + "FP4 KV cache expects kv_scale_quant_orig to have 3 elements."); + } + + // The KV-cache scales are meaningful only when the cache is actually quantized. + // `TrtllmAttention` hands them over unconditionally so the XQA kernels always see a + // valid pointer when `isKVCacheQuantized` is true; consumers in turn read a null + // pointer as "no KV-cache dequant", and some branch on exactly that. Drop the pair + // here when the quant mode says otherwise, so the whole op -- including the + // `!kv_scale_quant_orig` precondition on the SageAttention path -- keeps seeing a + // consistent pair. They travel together: a lone scale describes only half of the + // conversion and is never usable. + if (!mCfg.quant_mode.hasKvCacheQuant() || !p.fwd.kv_scale_orig_quant.has_value() + || !p.fwd.kv_scale_quant_orig.has_value()) + { + p.fwd.kv_scale_orig_quant.reset(); + p.fwd.kv_scale_quant_orig.reset(); + } + else if (mCfg.quant_mode.hasFp4KvCache()) + { + // An FP4 cache reads the scales as raw float pointers, so the shape and layout + // are part of the ABI rather than something the kernels can check. + auto const& origQuantScale = p.fwd.kv_scale_orig_quant.value(); + auto const& quantOrigScale = p.fwd.kv_scale_quant_orig.value(); + if (mCfg.is_mla_enable) + { + TORCH_CHECK(origQuantScale.scalar_type() == torch::kFloat32, + "kv_scale_orig_quant must have float32 dtype for MLA with FP4 KV cache"); + TORCH_CHECK(quantOrigScale.scalar_type() == torch::kFloat32, + "kv_scale_quant_orig must have float32 dtype for MLA with FP4 KV cache"); + TORCH_CHECK( + origQuantScale.is_contiguous(), "kv_scale_orig_quant must be contiguous for MLA with FP4 KV cache"); + TORCH_CHECK( + quantOrigScale.is_contiguous(), "kv_scale_quant_orig must be contiguous for MLA with FP4 KV cache"); + TORCH_CHECK(origQuantScale.dim() == 1 && origQuantScale.size(0) == 1, + "kv_scale_orig_quant must have shape [1] for MLA with FP4 KV cache"); + TORCH_CHECK(quantOrigScale.dim() == 1 && quantOrigScale.size(0) == 1, + "kv_scale_quant_orig must have shape [1] for MLA with FP4 KV cache"); + } + else + { + TORCH_CHECK(origQuantScale.size(0) == 3, "kv_scale_orig_quant must have 3 entries for FP4 KV cache"); + TORCH_CHECK(quantOrigScale.size(0) == 3, "kv_scale_quant_orig must have 3 entries for FP4 KV cache"); + } + } + auto checkCuSeqlens = [&p](std::optional const& cuSeqlens, char const* name) + { + if (!cuSeqlens.has_value()) + { + return; + } + auto const& tensor = cuSeqlens.value(); + TORCH_CHECK(tensor.dim() == 1, name, " must be a 1-D tensor."); + TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor."); + TORCH_CHECK(tensor.scalar_type() == at::ScalarType::Int, name, " must be int32."); + TORCH_CHECK(tensor.size(0) >= p.num_seqs + 1, name, " must have at least num_seqs + 1 elements."); + }; + if (!isGen) + { + checkCuSeqlens(p.fwd.cu_q_seqlens, "cu_q_seqlens"); + checkCuSeqlens(p.fwd.cu_kv_seqlens, "cu_kv_seqlens"); + } + + p.relative_attention_bias_stride = 0; + if (p.fwd.relative_attention_bias.has_value()) + { + auto const& bias = p.fwd.relative_attention_bias.value(); + TORCH_CHECK(bias.dim() == 2 || bias.dim() == 3, + "relative_attention_bias must be [num_heads, num_buckets] for implicit mode or " + "[num_heads, max_seq_len, max_seq_len] for explicit mode"); + TORCH_CHECK(bias.is_contiguous(), "relative_attention_bias must be contiguous"); + TORCH_CHECK(bias.scalar_type() == p.qkv_or_q.scalar_type(), + "relative_attention_bias dtype must match attention input dtype"); + p.relative_attention_bias_stride = static_cast(bias.size(1)); + } + + // Cross attention addresses the encoder KV by its own extents. The encoder lengths are + // the phase-local sequence lengths, so no separate tensor is needed. + p.num_encoder_tokens = 0; + p.cross_kv_length = 0; + if (p.is_cross) + { + p.encoder_input_lengths = p.sequence_length; + if (p.fwd.cross_kv.has_value() && p.num_seqs > 0) + { + p.num_encoder_tokens = p.getCrossKvNumTokens(); + p.cross_kv_length = p.getMaxHostPastKeyValueLength(p.seq_offset, p.num_seqs); + } + } + + // Speculative decoding extents come from the mask / position-offset tensor shapes and + // only apply to the generation phase. + if (isGen && mCfg.is_spec_decoding_enabled && p.use_spec_decoding) + { + TORCH_CHECK(p.spec_decoding_generation_lengths.has_value(), + "Expecting spec_decoding_generation_lengths in spec-dec mode."); + TORCH_CHECK( + p.spec_decoding_position_offsets.has_value(), "Expecting spec_decoding_position_offsets in spec-dec mode."); + TORCH_CHECK(p.spec_decoding_packed_mask.has_value(), "Expecting spec_decoding_packed_mask in spec-dec mode."); + + auto const& positionOffsets = p.spec_decoding_position_offsets.value(); + // [batch_size, max_draft_len + 1] + TORCH_CHECK(positionOffsets.dim() == 2, "spec_decoding_position_offsets must be 2-D."); + p.spec_decoding_is_generation_length_variable = true; + + if (tensorrt_llm::common::isSM100Family()) + { + TORCH_CHECK(p.spec_decoding_bl_tree_mask_offset.has_value(), + "Expecting spec_decoding_bl_tree_mask_offset in trtllm-gen spec-dec mode."); + TORCH_CHECK(p.spec_decoding_bl_tree_mask.has_value(), + "Expecting spec_decoding_bl_tree_mask in trtllm-gen spec-dec mode."); + TORCH_CHECK(p.spec_bl_tree_first_sparse_mask_offset_kv.has_value(), + "Expecting spec_bl_tree_first_sparse_mask_offset_kv in trtllm-gen spec-dec mode."); + // Blackwell uses the padded packed-mask row dim as the mask stride. + auto const& packedMask = p.spec_decoding_packed_mask.value(); + TORCH_CHECK(packedMask.dim() == 3, "spec_decoding_packed_mask must be 3-D in trtllm-gen spec-dec mode."); + p.spec_decoding_max_generation_length = static_cast(packedMask.size(1)); + } + else + { + p.spec_decoding_max_generation_length = static_cast(positionOffsets.size(1)); + } + } + else + { + // Callers can retain speculative buffers between steps. Inactive calls must not + // expose them as kernel inputs: QKV preprocessing checks the length and position + // pointers independently of the speculative-decoding flags. + p.spec_decoding_generation_lengths.reset(); + p.spec_decoding_position_offsets.reset(); + p.spec_decoding_packed_mask.reset(); + p.spec_decoding_bl_tree_mask_offset.reset(); + p.spec_decoding_bl_tree_mask.reset(); + p.spec_bl_tree_first_sparse_mask_offset_kv.reset(); + // Restore derived extents when preparing a reused native carrier. + p.spec_decoding_is_generation_length_variable = false; + p.spec_decoding_max_generation_length = 1; + } + + return finishPrepare(p, isGen); +} + +void AttentionOp::initialize() +{ + // Derived op state. + mFMHAForceFP32Acc = mCfg.type == tensorrt_llm::DataType::kBF16; + if (mCfg.is_mla_enable) + { + mFP8ContextMLA = (mSM == 90 || tensorrt_llm::common::isSM100Family(mSM) || mSM == 107 || mSM == 120) + && (mCfg.quant_mode.hasFp8KvCache() || mUseNvfp4MlaKvCache); + mFP8GenerationMLA = mCfg.quant_mode.hasFp8KvCache() || mUseNvfp4MlaKvCache; + } + mPagedKVCache = mPagedKVCache && mCfg.use_kv_cache; + bool const use_sage_attn = mCfg.sage_attn_num_elts_per_blk_q > 0 || mCfg.sage_attn_num_elts_per_blk_k > 0 + || mCfg.sage_attn_num_elts_per_blk_v > 0; + mFP8ContextFMHA = mCfg.is_fp8_out || mCfg.is_fp4_out || (mCfg.quant_mode.hasFp8KvCache() && mCfg.paged_context_fmha) + || use_sage_attn; + mFP8AttenOutput = mCfg.is_fp8_out; + mFuseFp4Quant = mCfg.is_fp4_out; + // Pre-check whether FMHA is supported in order to save memory allocation. + if (mEnableContextFMHA) + { + mEnableContextFMHA = false; + if (!(mCfg.type == tensorrt_llm::DataType::kHALF || mCfg.type == tensorrt_llm::DataType::kBF16)) + { + TLLM_LOG_WARNING("Fall back to unfused MHA because of unsupported data type."); + } + else if (mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kRELATIVE) + { + TLLM_LOG_WARNING("Fall back to unfused MHA because of relative position embedding."); + } + else if (mCfg.cross_attention && mCfg.use_kv_cache && !mPagedKVCache) + { + // TODO: add the support for cross attention + contiguous kv cache. + TLLM_LOG_WARNING("Fall back to unfused MHA because of cross attention + contiguous kv cache."); + } + else + { + mEnableContextFMHA = true; + } + } + + // Pre-Check of FP8 Context FMHA. + if (mFP8ContextFMHA) + { + TLLM_CHECK_WITH_INFO(mEnableContextFMHA, "FP8 FMHA cannot be enabled because Context FMHA is not supported."); + TLLM_CHECK_WITH_INFO( + mSM == 89 || mSM == 90 || tensorrt_llm::common::isSM100Family(mSM) || mSM == 120 || mSM == 121, + "FP8 FMHA can only be enabled on sm_89, sm_90, sm_100f, sm_120 or sm_121."); + } + + // Pre-Check of FP8 Generation MLA. + if (mFP8GenerationMLA) + { + TLLM_CHECK_WITH_INFO(mCfg.is_mla_enable, "FP8 Generation MLA cannot be enabled because MLA is not supported."); + TLLM_CHECK_WITH_INFO( + mSM == 89 || mSM == 90 || tensorrt_llm::common::isSM100Family(mSM) || mSM == 120 || mSM == 121, + "FP8 Generation MLA is supported on Ada, Hopper or Blackwell architecture."); + } + + // Check requirements for FP4 output. + TLLM_CHECK_WITH_INFO(!mFuseFp4Quant || mEnableContextFMHA, "Context FMHA must enable if fuse_fp4_quant is enabled"); + TLLM_CHECK_WITH_INFO(!mFuseFp4Quant || tensorrt_llm::common::isSM100Family(mSM) || mSM == 120 || mSM == 121, + "fuse_fp4_quant only supports SM100f or SM120 or SM121 devices."); + + // Check requirements for FP4 KV cache. + TLLM_CHECK_WITH_INFO(!mCfg.quant_mode.hasFp4KvCache() || mFP8ContextFMHA || mUseNvfp4MlaKvCache, + "FP4 KV cache requires FP8 context FMHA or static sparse MLA with an FP8 scratch pool"); + + TLLM_CHECK(isRoPE() == (mCfg.rotary_embedding_dim != 0)); + TLLM_CHECK_WITH_INFO((mSM >= 80) || (mCfg.type != tensorrt_llm::DataType::kBF16), + "Unsupported data type, pre SM 80 GPUs do not support bfloat16"); + + // Pre-check whether the head size is supported by MMHA. + // Support head size == 72 only for fmha kernels, so skip pre-check here. + if (getHeadSize() == 72) + { + ; + } + else if (!mmha_supported(getHeadSize()) && !mCfg.is_mla_enable) + { + TLLM_CHECK_WITH_INFO(false, "Head size %d is not supported by MMHA.", getHeadSize()); + } + + if (mCfg.is_mla_enable) + { + TLLM_CHECK_WITH_INFO(mEnableContextFMHA, "MLA(Deepseek v2) only support fmha"); + TLLM_CHECK_WITH_INFO(!mCfg.dense_context_fmha, "MLA(Deepseek v2) currently not support dense fmha"); + TLLM_CHECK_WITH_INFO( + mPagedKVCache && mCfg.use_kv_cache && mCfg.remove_padding, "MLA(Deepseek v2) only support paged kv cache"); + TLLM_CHECK_WITH_INFO(!mCfg.cross_attention, "MLA(Deepseek v2) do not support cross attention right now"); + TLLM_CHECK_WITH_INFO(mCfg.mask_type != tensorrt_llm::kernels::AttentionMaskType::CUSTOM_MASK, + "MLA(Deepseek v2) do not support custom mask right now"); + bool const mla_dims_supported = mMLAParams.qk_rope_head_dim == 64 + && ((mMLAParams.rope_append && mMLAParams.kv_lora_rank == 512) + || (!mMLAParams.rope_append && mMLAParams.kv_lora_rank == 448)); + TLLM_CHECK_WITH_INFO(mla_dims_supported, + "MLA(Deepseek v2) only supports qk_rope_head_dim=64 with kv_lora_rank=512 " + "(rope_append=true) or " + "kv_lora_rank=448 (rope_append=false)."); + } + if (mEnableContextFMHA) + { + // Construct the fmha runner. + MHARunnerFixedParams fmhaParams{}; + + bool const useSageAttn = mFP8ContextFMHA && !mCfg.is_mla_enable + && (mCfg.sage_attn_num_elts_per_blk_q > 0 || mCfg.sage_attn_num_elts_per_blk_k > 0 + || mCfg.sage_attn_num_elts_per_blk_v > 0); + + // Pre-checked during constructing. + Data_type data_type, data_type_kv; + if (mCfg.type == tensorrt_llm::DataType::kHALF) + { + data_type = DATA_TYPE_FP16; + } + else if (mCfg.type == tensorrt_llm::DataType::kBF16) + { + data_type = DATA_TYPE_BF16; + } + else + { + TLLM_CHECK_WITH_INFO(false, "GPTAttentionPlugin received wrong data type."); + } + // The output dtype. + fmhaParams.dataTypeOut = mFP8AttenOutput ? DATA_TYPE_E4M3 : data_type; + data_type_kv = data_type; + + // FP8 FMHA should be used with fp8 workflow together. + if (mFP8ContextFMHA || mFP8ContextMLA) + { + if (mFP8ContextFMHA && useSageAttn && mCfg.sage_attn_qk_int8) + { + data_type = DATA_TYPE_INT8; + data_type_kv = DATA_TYPE_KV_INT8_E4M3; + } + else + { + data_type = DATA_TYPE_E4M3; + data_type_kv = DATA_TYPE_E4M3; + } + } + + // The input dtype. + fmhaParams.dataType = data_type; + // The KV input data type. The default is same as dataType. + fmhaParams.dataTypeKv = data_type_kv; + // If the kernel must read from KV cache, set the dtype correctly. + if (mPagedKVCache && mCfg.paged_context_fmha) + { + if (mCfg.quant_mode.hasFp8KvCache()) + { + fmhaParams.dataTypeKv = DATA_TYPE_E4M3; + } + else if (mCfg.quant_mode.hasFp4KvCache()) + { + fmhaParams.dataTypeKv = DATA_TYPE_E2M1; + } + } + if (mFuseFp4Quant) + { + // If FP4 quantization workflow is enabled, set output type to FP4. + fmhaParams.dataTypeOut = DATA_TYPE_E2M1; + } + if (mCfg.is_mla_enable) + { + // For FP8 MLA, currently context attention is performed in BF16. + fmhaParams.dataTypeOut = DATA_TYPE_BF16; + fmhaParams.dataTypeKv = DATA_TYPE_BF16; + } + if (mFP8ContextMLA) + { + fmhaParams.dataTypeKv = DATA_TYPE_E4M3; + fmhaParams.dataTypeOut = DATA_TYPE_BF16; + } + if (mCfg.fuses_dsv4_inv_rope_fp8_quant) + { + fmhaParams.dataTypeOut = DATA_TYPE_E4M3; + } + // TODO: remove forceFp32Acc from MHARunnerFixedParams after adding host_runtime_perf_knobs to + // bertAttentionPlugin input tensors, so that we can change mLaunchParams.force_fp32_acc value + // in runtime. + fmhaParams.forceFp32Acc = false; + + // setting attention mask type based on the mask type + fmhaParams.setAttentionMaskType(static_cast(mCfg.mask_type)); + + if (mCfg.cross_attention) + { + // always use paged-kv-fmha if paged_kv cache is used. + fmhaParams.attentionInputLayout + = mPagedKVCache ? AttentionInputLayout::Q_PAGED_KV : AttentionInputLayout::Q_CONTIGUOUS_KV; + } + else if (!mCfg.use_kv_cache) + { + if (useSageAttn) + { + fmhaParams.attentionInputLayout = AttentionInputLayout::SEPARATE_Q_K_V; + } + else + { + fmhaParams.attentionInputLayout = AttentionInputLayout::PACKED_QKV; + } + } + else + { + fmhaParams.attentionInputLayout = (mPagedKVCache && mCfg.paged_context_fmha) + ? AttentionInputLayout::Q_PAGED_KV + : AttentionInputLayout::PACKED_QKV; + } + fmhaParams.isSPadded = !mCfg.remove_padding; + fmhaParams.numQHeads = mNumAttnHeads; + fmhaParams.numKvHeads = mNumAttnKVHeads; + fmhaParams.numTokensPerBlock = mCfg.tokens_per_block; + fmhaParams.headSize = mHeadSize; + fmhaParams.headSizeV = mHeadSize; + fmhaParams.qScaling = mCfg.q_scaling; + + // mFmhaDispatcher is not used for generation MLA, but we still need to modify these values to + // avoid selecting the wrong kernel, no matter mIsGenerationMLA is true or false + if (mCfg.is_mla_enable) + { + if (cfgUseSparseMLA()) + { + fmhaParams.attentionInputLayout = AttentionInputLayout::Q_PAGED_KV; + fmhaParams.numKvHeads = 1; + fmhaParams.headSize = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + fmhaParams.headSizeV = mMLAParams.rope_append ? mMLAParams.kv_lora_rank + : mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + fmhaParams.headSizeQkNope = mMLAParams.qk_nope_head_dim; + // Adjust the qScaling for the absorption mode. + fmhaParams.qScaling = mCfg.q_scaling + * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) + / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); + } + else + { + // Context MLA always use separate_q_k_v layout + fmhaParams.attentionInputLayout = AttentionInputLayout::SEPARATE_Q_K_V; + // Context attention of MLA is different + fmhaParams.numKvHeads = mCfg.num_heads; + fmhaParams.headSize = mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim; + // Ideally this should be mMLAParams.v_head_dim, but because we initialize both MLA + // context(v_head_dim=128) and gen(v_head_dim=512) runners in a single op, the headSizeV + // will be set to 512 when we create the gen attention op and that could fail to create + // the FmhaDispatcher for context phase. Luckily, for deepseek, qk_nope_head_dim is the + // same as v_head_dim in context phase. + fmhaParams.headSizeV = mMLAParams.qk_nope_head_dim; + fmhaParams.headSizeQkNope = mMLAParams.qk_nope_head_dim; + } + } + fmhaParams.attnLogitSoftcappingScale = mCfg.attn_logit_softcapping_scale; + fmhaParams.hasAlibi = isALiBi(); + fmhaParams.scaleAlibi = isAliBiWithScale(); + fmhaParams.useSparseMLA = cfgUseSparseMLA(); + fmhaParams.useTllmGenSparseAttention = cfgUseTllmGenSparseAttention(); + fmhaParams.fusesDsv4InvRopeFp8Quant = mCfg.fuses_dsv4_inv_rope_fp8_quant; + + // SageAttention: set block sizes for sage quantization. + if (useSageAttn) + { + fmhaParams.sageBlockSizeQ = mCfg.sage_attn_num_elts_per_blk_q; + fmhaParams.sageBlockSizeK = mCfg.sage_attn_num_elts_per_blk_k; + fmhaParams.sageBlockSizeV = mCfg.sage_attn_num_elts_per_blk_v; + } + + // Load kernels from the pre-compiled cubins. + mFmhaDispatcher.reset(new FmhaDispatcher(fmhaParams)); + + // Deepseek-V2 Generation needs a differ fmha with different argumments + if (mCfg.is_mla_enable) { - auto const capacity = workspace_.value().storage().nbytes(); - if (capacity < static_cast(workspace_size)) + mEnableXQA = (mSM == kSM_120) && mIsGenerationMLA; + if (mUseTllmGen) + { + Data_type qDataType = DATA_TYPE_FP32; + Data_type kvDataType = DATA_TYPE_FP32; + Data_type outputDataType = DATA_TYPE_FP32; + + if (mCfg.type == tensorrt_llm::DataType::kHALF) + { + qDataType = DATA_TYPE_FP16; + kvDataType = DATA_TYPE_FP16; + outputDataType = DATA_TYPE_FP16; + } + else if (mCfg.type == tensorrt_llm::DataType::kBF16) + { + qDataType = DATA_TYPE_BF16; + kvDataType = DATA_TYPE_BF16; + outputDataType = DATA_TYPE_BF16; + } + else + { + TLLM_CHECK_WITH_INFO(false, "The data type is not supported."); + } + + if (mFP8GenerationMLA) + { + qDataType = DATA_TYPE_E4M3; + kvDataType = DATA_TYPE_E4M3; + } + if (mCfg.fuses_dsv4_inv_rope_fp8_quant) + { + outputDataType = DATA_TYPE_E4M3; + } + + // Instantiate the mTllmGenFMHARunner used for MLA + mTllmGenFMHARunner.reset(new TllmGenFmhaRunner( + qDataType, kvDataType, kvDataType, outputDataType, 0, 0, 0, 0, mCfg.fuses_dsv4_inv_rope_fp8_quant)); + } + else if (mIsGenerationMLA && !mUseGenFlashMLA) + { + // Construct the fmha runner for generation. + if (mFP8GenerationMLA) + { + data_type = DATA_TYPE_E4M3; + } + MHARunnerFixedParams fmhaParams{}; + fmhaParams.dataType = data_type; + fmhaParams.dataTypeKv = data_type; + fmhaParams.dataTypeOut = data_type; + // For FP8 MLA generation, the output type is BF16, and the quantization before o_proj + // is performed separately. + if (mFP8GenerationMLA) + { + fmhaParams.dataTypeOut = DATA_TYPE_BF16; + } + // TODO: remove forceFp32Acc from MHARunnerFixedParams after adding + // host_runtime_perf_knobs to bertAttentionPlugin input tensors, so that we can change + // mLaunchParams.force_fp32_acc value in runtime. + fmhaParams.forceFp32Acc = true; + fmhaParams.attentionMaskType + = useCustomMask() ? ContextAttentionMaskType::CUSTOM_MASK : ContextAttentionMaskType::PADDING; + // TODO: set it to Q_CONTIGUOUS_KV layout for cross-attention. + fmhaParams.attentionInputLayout = AttentionInputLayout::Q_PAGED_KV; + fmhaParams.isSPadded = !mCfg.remove_padding; + fmhaParams.numQHeads = 1; + fmhaParams.numKvHeads = 1; + fmhaParams.headSize = mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim; + fmhaParams.headSizeV = mMLAParams.kv_lora_rank; + fmhaParams.qScaling = mCfg.q_scaling + * sqrt((float) (mMLAParams.qk_nope_head_dim + mMLAParams.qk_rope_head_dim)) + / sqrtf((float) (mMLAParams.kv_lora_rank + mMLAParams.qk_rope_head_dim)); + fmhaParams.attnLogitSoftcappingScale = mCfg.attn_logit_softcapping_scale; + fmhaParams.hasAlibi = isALiBi(); + fmhaParams.scaleAlibi = isAliBiWithScale(); + fmhaParams.tpSize = 1; + fmhaParams.tpRank = 0; + mDecoderFMHARunner.reset(new FusedMHARunnerV2(fmhaParams)); + + // Only deepseek must using fmha in the generation phase when flash mla is not enabled. + if (!mUseGenFlashMLA) + { + TLLM_CHECK_WITH_INFO(mDecoderFMHARunner->isFmhaSupported(), + "Deepseek should be supported by fmha in generation part."); + } + } + if (!mIsGenerationMLA) { - TLLM_LOG_WARNING( - "Attention workspace size is not enough, increase the size from %zu bytes to %ld bytes", capacity, - workspace_size); + TLLM_CHECK_WITH_INFO( + mFmhaDispatcher->isSupported(), "Deepseek should be supported by fmha in context part."); } - workspace_.value().resize_({workspace_size}); } - workspace = workspace_.value(); + + // Fall back to unfused MHA kernels if not supported. + // Generation MLA reuses the context FMHA code path so set mEnableContextFMHA to true. + // However, do not check mFmhaDispatcher which is not used for generation MLA. + mEnableContextFMHA = mIsGenerationMLA || mFmhaDispatcher->isSupported(); + + // Only FMHA supports custom mask currently. + TLLM_CHECK_WITH_INFO( + !useCustomMask() || mEnableContextFMHA, "Only Context FMHA supports custom mask input currently."); } - else + + mEnableXQA = (mEnableXQA || mCfg.is_spec_decoding_enabled) + && (mCfg.type == tensorrt_llm::DataType::kHALF || mCfg.type == tensorrt_llm::DataType::kBF16) + && mCfg.use_kv_cache; + + if (mEnableXQA) + { + TLLM_LOG_DEBUG("Enabling XQA kernels for GPTAttention."); + + XqaFixedParams fixedParams{}; + fixedParams.isMLA = mIsGenerationMLA; + // TODO: support more combinations. + // Update Q and O dtype. + if (mCfg.type == tensorrt_llm::DataType::kHALF) + { + fixedParams.inputDataType = DATA_TYPE_FP16; + fixedParams.outputDataType = DATA_TYPE_FP16; + } + else if (mCfg.type == tensorrt_llm::DataType::kBF16) + { + fixedParams.inputDataType = DATA_TYPE_BF16; + fixedParams.outputDataType = DATA_TYPE_BF16; + } + // Update KV cache and math dtype. + if (mCfg.quant_mode.hasInt8KvCache()) + { + fixedParams.kvDataType = DATA_TYPE_INT8; + fixedParams.mathDataType = fixedParams.inputDataType; + } + else if (mCfg.quant_mode.hasFp8KvCache()) + { + fixedParams.kvDataType = DATA_TYPE_E4M3; + fixedParams.mathDataType = DATA_TYPE_E4M3; + } + else if (mCfg.quant_mode.hasFp4KvCache()) + { + fixedParams.kvDataType = DATA_TYPE_E2M1; + fixedParams.mathDataType = DATA_TYPE_E4M3; + } + else + { + fixedParams.kvDataType = fixedParams.inputDataType; + fixedParams.mathDataType = fixedParams.inputDataType; + } + // If fuse_fp4_quant is enabled, set output data type to FP4. + if (mFuseFp4Quant) + { + fixedParams.outputDataType = DATA_TYPE_E2M1; + } + else if (mFP8AttenOutput) + { + fixedParams.outputDataType = DATA_TYPE_E4M3; + } + if (mCfg.is_spec_decoding_enabled && !mUseTllmGen) + { + fixedParams.outputDataType = DATA_TYPE_E4M3; + TLLM_CHECK_WITH_INFO( + mCfg.num_heads % mNumKVHeads == 0, "mCfg.num_heads should be multiples of mNumKVHeads."); + } + + fixedParams.numQHeads = mNumAttnHeads; + fixedParams.numKvHeads = mNumAttnKVHeads; + fixedParams.numTokensPerBlock = mCfg.tokens_per_block; + fixedParams.headSize = mHeadSize; + fixedParams.qScaling = mCfg.q_scaling; + fixedParams.multiBlockMode = mMultiBlockMode; + fixedParams.isPagedKv = mPagedKVCache; + fixedParams.isSpecDecoding = mCfg.is_spec_decoding_enabled; + fixedParams.hasAlibi = isALiBi(); + fixedParams.useTllmGenSparseAttention = cfgUseTllmGenSparseAttention(); + fixedParams.specDecodingTargetMaxGenLen = mCfg.spec_decoding_target_max_gen_len; + + mXqaDispatcher.reset(new XqaDispatcher(fixedParams)); + + // Fall back to unfused MHA kernels if not supported. + mEnableXQA = mXqaDispatcher->isSupported(); + } + else if (mCfg.is_spec_decoding_enabled) + { + TLLM_CHECK_WITH_INFO(false, "Speculative decoding mode doesn't support the data type or cross attention."); + } +} + +int AttentionOp::finishPrepare(FmhaParams& p, bool isGen) +{ + +#if ENABLE_MULTI_DEVICE +#endif // ENABLE_MULTI_DEVICE + // JIT the generation XQA kernel only for a generation call. Its cubin key carries + // `input_seq_length`, `num_requests` and the window extents, which only the + // generation phase fills in; the workspace-sizing pass runs `prepare` before the + // phase split, and compiling from those zeros fails inside NVRTC on SM90. + if (isGen) + { + DISPATCH_ON_DTYPE(p.getType(), prepareEnqueueGeneration, p); + } + + p.sparse_params = {}; + p.sparse_params.sparse_kv_indices = p.getSparseKvIndices(); + p.sparse_params.sparse_kv_offsets = p.getSparseKvOffsets(); + p.sparse_params.sparse_attn_indices = p.getSparseAttnIndices(); + p.sparse_params.sparse_attn_offsets = p.getSparseAttnOffsets(); + p.sparse_params.sparse_attn_indices_block_size = p.fwd.sparse_runtime_params.sparse_attn_indices_block_size; + p.sparse_params.sparse_attn_indices_stride = p.getSparseAttnIndicesStride(); + p.sparse_params.num_sparse_topk = p.num_sparse_topk; + p.sparse_params.sparse_attn_kv_lens = p.getSparseAttnKvLens(); + + p.kv_cache_pool_pointers = {}; + if (!useKVCache(p) || !p.hasKvCache()) + { + return 0; + } + + int32_t const poolIndex = p.getKvCachePoolIndex(p.local_layer_idx); + int32_t const layerIdxInCachePool = p.getLayerIdxInCachePool(p.local_layer_idx); + size_t const dTypeSize = p.getType() == tensorrt_llm::DataType::kFLOAT ? sizeof(float) : sizeof(half); + int const cacheElemBits = AttentionOp::getKvCacheElemSizeInBits(mCfg.quant_mode, dTypeSize); + auto const blockSize = static_cast(mCfg.tokens_per_block) * mNumKVHeads * mHeadSize; + auto const bytesPerBlock = blockSize * cacheElemBits / CHAR_BIT; + int32_t const kvFactor = isMLAEnabled() ? 1 : 2; + auto const intraPoolOffset = layerIdxInCachePool * kvFactor * bytesPerBlock; + + p.kv_cache_pool_pointers = buildKvCachePoolPointers(p.getHostKvCachePoolPointers(), poolIndex, intraPoolOffset, + blockSize, layerIdxInCachePool, kvFactor, mCfg.quant_mode.hasFp4KvCache()); + + if (p.use_sparse_attention) + { + auto* kvCachePool = p.getSparseKvCachePool(poolIndex); + if (kvCachePool != nullptr) + { + if (p.fwd.sparse_runtime_params.sparse_attn_kv_lens.has_value()) + { + // Deepseek V4 dynamic sparse MLA always uses the SWA pool for now. + p.sparse_params.sliding_window_kv_cache_pool = kvCachePool; + if (p.fwd.sparse_runtime_params.aux_kv_cache_pool_ptr.has_value()) + { + p.sparse_params.sparse_kv_cache_pool + = (void*) (intptr_t) p.fwd.sparse_runtime_params.aux_kv_cache_pool_ptr.value(); + } + } + else + { + p.sparse_params.sparse_kv_cache_pool = kvCachePool; + } + } + } + return 0; +} + +int64_t AttentionOp::getAttentionWorkspaceSize(FmhaParams const& p, int64_t num_tokens, + int64_t max_attention_window_size, int64_t num_gen_tokens, int64_t max_blocks_per_sequence) +{ + TORCH_CHECK(num_gen_tokens >= 0 && num_gen_tokens <= num_tokens, "num_gen_tokens must be in [0, num_tokens]"); + auto params = p; + AttentionOp& op = *this; + op.prepare(params, /*isGen=*/false); + int64_t const numContextTokens = num_tokens - num_gen_tokens; + int32_t const numContexts = numContextTokens > 0 ? static_cast(params.num_seqs) : 0; + int32_t const maxContextLength = numContexts > 0 ? static_cast(params.input_seq_length) : 0; + // For cross-attention, several unfused-path context buffers scale with the encoder KV length. + // Mirror the context-stage enqueue, which uses the max past-KV length over the context sequences + // as cross_kv_length; sizing with 0 here under-allocates the workspace and the carved views in + // enqueueContext land past the end of the allocation. The enqueue also gates on + // cross_kv.has_value(), so this can over-allocate relative to the carve; that is safe. + int32_t maxCrossKvLength = 0; + if (isCrossAttention(params) && numContexts > 0) { - TLLM_LOG_TRACE("Allocate new attention workspace with size %ld bytes", workspace_size); - workspace = torch::empty({workspace_size}, torch::dtype(torch::kByte).device(qkv_or_q.device())); - } - - if ((num_contexts > 0) && (attn_input_type != AttentionInputType::GenerationOnly)) - { - auto seq_offset = 0; - auto token_offset = 0; - runner->run(*op, - /*is_context=*/true, seq_offset, - /*num_seqs=*/num_contexts, token_offset, - /*num_tokens=*/num_ctx_tokens, predicted_tokens_per_seq, workspace, output, output_sf, qkv_or_q, k, v, - sequence_length, host_past_key_value_lengths, ctx_total_kv_len, context_lengths, host_context_lengths, - max_context_q_len_override, kv_cache_block_offsets, host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, - cache_indirection, kv_scale_orig_quant, kv_scale_quant_orig, out_scale, rotary_inv_freq, rotary_cos_sin, - latent_cache, q_pe, block_ids_per_seq, mrope_rotary_cos_sin, mrope_position_deltas, helix_position_offsets, - helix_is_inactive_rank, softmax_stats_tensor, spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset, - spec_decoding_bl_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv, attention_sinks, sparse_kv_indices, - sparse_kv_offsets, sparse_attn_indices, sparse_attn_offsets, sparse_attn_indices_block_size, - num_sparse_topk_value, sparse_attn_kv_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter, - mla_bmm1_scale, mla_bmm2_scale, quant_q_buffer, flash_mla_tile_scheduler_metadata, flash_mla_num_splits, - trtllm_gen_jit_warmup, aux_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias, quant_scale_qkv, - dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion, kv_norm_weight, kv_norm_eps); - } - - if ((num_generations > 0) && (attn_input_type != AttentionInputType::ContextOnly)) - { - - auto seq_offset = num_contexts; - auto token_offset = is_gen_only ? 0 : num_ctx_tokens; - runner->run(*op, - /*is_context=*/false, seq_offset, - /*num_seqs=*/num_generations, token_offset, - /*num_tokens=*/num_gen_tokens, predicted_tokens_per_seq, workspace, output, output_sf, qkv_or_q, k, v, - sequence_length, host_past_key_value_lengths, gen_total_kv_len, context_lengths, host_context_lengths, - max_context_q_len_override, kv_cache_block_offsets, host_kv_cache_pool_pointers, host_kv_cache_pool_mapping, - cache_indirection, kv_scale_orig_quant, kv_scale_quant_orig, out_scale, rotary_inv_freq, rotary_cos_sin, - latent_cache, q_pe, block_ids_per_seq, mrope_rotary_cos_sin, mrope_position_deltas, helix_position_offsets, - helix_is_inactive_rank, softmax_stats_tensor, spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset, - spec_decoding_bl_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv, attention_sinks, sparse_kv_indices, - sparse_kv_offsets, sparse_attn_indices, sparse_attn_offsets, sparse_attn_indices_block_size, - num_sparse_topk_value, sparse_attn_kv_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter, - mla_bmm1_scale, mla_bmm2_scale, quant_q_buffer, flash_mla_tile_scheduler_metadata, flash_mla_num_splits, - trtllm_gen_jit_warmup, aux_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias, quant_scale_qkv, - dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion, - // Context-only here; generation gets the fusion from mla_rope_generation. - /*kv_norm_weight=*/std::nullopt, kv_norm_eps); - } - - TLLM_LOG_TRACE("Attention op stops at layer %d", local_layer_idx); + maxCrossKvLength = params.getMaxHostPastKeyValueLength(params.seq_offset, numContexts); + } + size_t const contextWorkspaceSize = op.getWorkspaceSizeForContext(params, numContexts, maxContextLength, + maxCrossKvLength, static_cast(numContextTokens), static_cast(params.total_kv_len)); + // The generation workspace is sized per sequence (max_num_requests * beam_width), not + // per request; they only coincide when beam_width == 1. + int64_t const maxNumSequences = params.max_num_sequences > 0 ? params.max_num_sequences : params.max_num_requests; + size_t const generationWorkspaceSize = op.getWorkspaceSizeForGeneration(params, static_cast(maxNumSequences), + static_cast(max_attention_window_size), static_cast(num_gen_tokens), + static_cast(max_blocks_per_sequence)); + return static_cast(std::max(contextWorkspaceSize, generationWorkspaceSize)); +} + +void AttentionOp::runContext(FmhaParams& p) +{ + DISPATCH_ON_TORCH_DTYPE(p.qkv_or_q.scalar_type(), runContextImpl, p); +} + +void AttentionOp::runGeneration(FmhaParams& p) +{ + DISPATCH_ON_TORCH_DTYPE(p.qkv_or_q.scalar_type(), runGenerationImpl, p); +} + +void AttentionOp::runMlaGeneration(FmhaParams& p) +{ + DISPATCH_ON_TORCH_DTYPE(p.qkv_or_q.scalar_type(), runMlaGenerationImpl, p); } bool attention_supports_nvfp4_output(int64_t const num_heads, int64_t const num_kv_heads, int64_t const head_size, @@ -1594,36 +3979,23 @@ bool attention_supports_nvfp4_output(int64_t const num_heads, int64_t const num_ return false; } - auto op = std::make_shared(); - op->mType = tensorrt_llm::DataType::kHALF; - op->mNumHeads = num_heads; - op->mNumKVHeads = num_kv_heads; - op->mHeadSize = head_size; - op->mMaskType = static_cast(int32_t(mask_type)); - op->mKVCacheQuantMode = tensorrt_llm::common::QuantMode(uint32_t(quant_mode)); - op->mFP8ContextFMHA = op->mKVCacheQuantMode.hasFp8KvCache() || op->mKVCacheQuantMode.hasFp4KvCache(); - op->mUseKVCache = true; - op->mPagedKVCache = true; - op->mTokensPerBlock = tokens_per_block.value_or(0); - op->mFuseFp4Quant = true; - op->mPagedContextFMHA = use_paged_context_fmha; - - auto cache_key = op->data(); - using CacheKey = decltype(cache_key); - static std::unordered_map> op_cache; - if (auto it = op_cache.find(cache_key); it != op_cache.end()) - { - TLLM_LOG_TRACE("Attention op runtime check is cached"); - return it->second; - } - else - { - TLLM_LOG_TRACE("Caching attention op runtime check with cache key: %s", to_string(cache_key).c_str()); - op->initialize(); - op_cache[cache_key] = op->supportsNvFp4Output(); - } - - return op->supportsNvFp4Output(); + StaticAttentionConfig cfg{}; + cfg.num_heads = num_heads; + cfg.num_kv_heads = num_kv_heads; + cfg.head_size = head_size; + cfg.tokens_per_block = tokens_per_block.value_or(0); + cfg.type = tensorrt_llm::DataType::kHALF; + cfg.is_fp4_out = true; + cfg.use_kv_cache = true; + cfg.paged_context_fmha = use_paged_context_fmha; + cfg.mask_type = static_cast(int32_t(mask_type)); + cfg.q_scaling = 1.0; + cfg.remove_padding = true; + cfg.is_mla_enable = is_mla_enable; + cfg.quant_mode = tensorrt_llm::common::QuantMode(uint32_t(quant_mode)); + + AttentionOp op{cfg}; + return op.supportsNvFp4Output(); } KvCachePoolPointers buildKvCachePoolPointers(at::Tensor const& hostKvCachePoolPointers, int32_t poolIndex, @@ -1633,7 +4005,8 @@ KvCachePoolPointers buildKvCachePoolPointers(at::Tensor const& hostKvCachePoolPo if (isFp4KvCache) { // For NVFP4 KV cache, extra block scales are stored in separate pools. - // The layout of host_kv_cache_pool_pointers is [num_pools, 2 (primary and secondary), 2 (data and scale)]. + // The layout of host_kv_cache_pool_pointers is [num_pools, 2 (primary and secondary), 2 (data + // and scale)]. TORCH_CHECK(hostKvCachePoolPointers.dim() == 3); pointers.primaryPoolPtr = reinterpret_cast( reinterpret_cast( @@ -1671,7 +4044,7 @@ KvCachePoolPointers buildKvCachePoolPointers(at::Tensor const& hostKvCachePoolPo return pointers; } -common::op::KvCacheBuffers buildPagedKvCacheBuffers( +KvCacheBuffers buildPagedKvCacheBuffers( std::optional const& kv_cache_block_offsets, std::optional const& host_kv_cache_pool_pointers, std::optional const& host_kv_cache_pool_mapping, common::QuantMode quantMode, int64_t layer_idx, @@ -1694,7 +4067,7 @@ common::op::KvCacheBuffers buildPagedKvCacheBuffers( auto* blockOffsets = tensorPtr2D( kv_cache_block_offsets.value(), poolIndex, static_cast(seq_offset), "kv_cache_block_offsets"); - int cacheElemBits = common::op::AttentionOp::getKvCacheElemSizeInBits(quantMode, elem_size); + int cacheElemBits = AttentionOp::getKvCacheElemSizeInBits(quantMode, elem_size); auto const blockSize = tokens_per_block * kv_head_num * size_per_head; auto const bytesPerBlock = blockSize * cacheElemBits / CHAR_BIT; @@ -1706,9 +4079,8 @@ common::op::KvCacheBuffers buildPagedKvCacheBuffers( blockSize, layerIdxInCachePool, kvFactor, quantMode.hasFp4KvCache()); int32_t const maxBlocksPerSequence = static_cast(kv_cache_block_offsets->size(-1)); - return common::op::buildKvCacheBuffers(static_cast(batch_size), - maxBlocksPerSequence, static_cast(tokens_per_block), sizePerToken, - static_cast(cyclic_attention_window_size), + return buildKvCacheBuffers(static_cast(batch_size), maxBlocksPerSequence, + static_cast(tokens_per_block), sizePerToken, static_cast(cyclic_attention_window_size), static_cast(std::max(cyclic_attention_window_size, max_attention_window_size)), /*sink_token_length=*/0, beam_width > 1, poolPointers.primaryPoolPtr, poolPointers.secondaryPoolPtr, poolPointers.primaryBlockScalePoolPtr, poolPointers.secondaryBlockScalePoolPtr, blockOffsets, @@ -1728,7 +4100,7 @@ std::tuple> buildFlashinferTrtllmGenPagedK bool const isFp4 = quantMode.hasFp4KvCache(); size_t const inputElemSize = isFp4 ? 1 : (quantMode.hasFp8KvCache() || quantMode.hasInt8KvCache() ? 1 : 2); - int const cacheElemBits = common::op::AttentionOp::getKvCacheElemSizeInBits(quantMode, inputElemSize); + int const cacheElemBits = AttentionOp::getKvCacheElemSizeInBits(quantMode, inputElemSize); auto const blockSize = tokens_per_block * num_kv_heads * head_dim; auto const bytesPerBlock = blockSize * cacheElemBits / CHAR_BIT; @@ -1777,7 +4149,7 @@ void computeFlashMlaMetadata(torch::Tensor seqlens_k, torch::Tensor tile_schedul cudaStream_t stream = at::cuda::getCurrentCUDAStream(seqlens_k.get_device()); static constexpr int block_size_n = 64; static constexpr int fixed_overhead_num_blocks = 5; - int const num_sm_parts = tensorrt_llm::common::op::AttentionOp::getFlashMlaNumSmPartsStatic(static_cast(s_q), + int const num_sm_parts = tensorrt_llm::torch_ext::AttentionOp::getFlashMlaNumSmPartsStatic(static_cast(s_q), static_cast(num_q_heads), static_cast(num_kv_heads), static_cast(head_size_v)); Mla_metadata_params params = {}; params.seqlens_k_ptr = seqlens_k.data_ptr(); diff --git a/cpp/tensorrt_llm/thop/attentionOp.h b/cpp/tensorrt_llm/thop/attentionOp.h index 0897b0eff367..bd2bf7ee3836 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.h +++ b/cpp/tensorrt_llm/thop/attentionOp.h @@ -16,89 +16,38 @@ #pragma once -#include -#include -#include -#include - -#include "tensorrt_llm/common/attentionOp.h" #include "tensorrt_llm/common/config.h" +#include "tensorrt_llm/common/cublasMMWrapper.h" #include "tensorrt_llm/common/cudaUtils.h" #include "tensorrt_llm/common/quantization.h" +#include "tensorrt_llm/kernels/fmhaDispatcher.h" +#include "tensorrt_llm/kernels/gptKernels.h" #include "tensorrt_llm/kernels/kvCacheUtils.h" -#include "tensorrt_llm/kernels/unfusedAttentionKernels.h" +#include "tensorrt_llm/kernels/mlaKernels.h" +#include "tensorrt_llm/kernels/sparseAttentionKernels.h" +#include "tensorrt_llm/kernels/xqaDispatcher.h" +#include "tensorrt_llm/runtime/torchUtils.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +#if ENABLE_MULTI_DEVICE +#include +#endif // ENABLE_MULTI_DEVICE TRTLLM_NAMESPACE_BEGIN namespace torch_ext { -/** - * @brief Attention operation for TensorRT-LLM - * - * This function performs multi-head attention computation in-place, supporting both - * context and generation phases with various optimization features including: - * - Fused QKV processing - * - KV cache management - * - Multiple position embedding types (RoPE, ALiBi, etc.) - * - Quantization support (FP8, FP4, etc.) - * - Multi-layer attention (MLA) - * - Speculative decoding - */ -void attention(torch::Tensor q, std::optional k, std::optional v, torch::Tensor& output, - std::optional output_sf, std::optional workspace_, torch::Tensor sequence_length, - torch::Tensor host_past_key_value_lengths, torch::Tensor host_total_kv_lens, torch::Tensor context_lengths, - torch::Tensor host_context_lengths, torch::Tensor host_request_types, - std::optional max_context_q_len_override, std::optional kv_cache_block_offsets, - std::optional host_kv_cache_pool_pointers, std::optional host_kv_cache_pool_mapping, - std::optional cache_indirection, std::optional kv_scale_orig_quant, - std::optional kv_scale_quant_orig, std::optional out_scale, - std::optional rotary_inv_freq, std::optional rotary_cos_sin, - std::optional latent_cache, std::optional q_pe, - std::optional block_ids_per_seq, std::optional attention_sinks, - bool const is_fused_qkv, bool const update_kv_cache, int64_t const predicted_tokens_per_seq, - int64_t const local_layer_idx, int64_t const num_heads, int64_t const num_kv_heads, int64_t const head_size, - std::optional const tokens_per_block, int64_t const max_num_requests, int64_t const max_context_length, - int64_t const max_seq_len, int64_t const attention_window_size, int64_t const beam_width, int64_t const mask_type, - int64_t const quant_mode, double const q_scaling, int64_t const position_embedding_type, int64_t const rope_dim, - double const rope_base, int64_t const rope_scale_type, double const rope_scale, double const rope_short_m_scale, - double const rope_long_m_scale, int64_t const rope_max_positions, int64_t const rope_original_max_positions, - bool const use_paged_context_fmha, std::optional attention_input_type, bool is_mla_enable, - std::optional chunked_prefill_buffer_batch_size, std::optional q_lora_rank, - std::optional kv_lora_rank, std::optional qk_nope_head_dim, - std::optional qk_rope_head_dim, std::optional v_head_dim, std::optional rope_append, - std::optional mrope_rotary_cos_sin, std::optional mrope_position_deltas, - std::optional helix_position_offsets, std::optional helix_is_inactive_rank, - std::optional attention_chunk_size, std::optional softmax_stats_tensor, - bool const is_spec_decoding_enabled, bool const use_spec_decoding, bool const is_spec_dec_tree, - std::optional spec_decoding_generation_lengths, - std::optional spec_decoding_position_offsets_for_cpp, - std::optional spec_decoding_packed_mask, - std::optional spec_decoding_bl_tree_mask_offset, - std::optional spec_decoding_bl_tree_mask, - std::optional spec_bl_tree_first_sparse_mask_offset_kv, - std::optional sparse_kv_indices, std::optional sparse_kv_offsets, - std::optional sparse_attn_indices, std::optional sparse_attn_offsets, - int64_t const sparse_attn_indices_block_size, std::optional num_sparse_topk, - std::optional sparse_attn_kv_lens, std::optional skip_softmax_threshold_scale_factor_prefill, - std::optional skip_softmax_threshold_scale_factor_decode, std::optional skip_softmax_stat, - std::optional cu_q_seqlens, std::optional cu_kv_seqlens, - std::optional fmha_scheduler_counter, std::optional mla_bmm1_scale, - std::optional mla_bmm2_scale, std::optional quant_q_buffer, - std::optional flash_mla_tile_scheduler_metadata = std::nullopt, - std::optional flash_mla_num_splits = std::nullopt, int64_t sage_attn_num_elts_per_blk_q = 0, - int64_t sage_attn_num_elts_per_blk_k = 0, int64_t sage_attn_num_elts_per_blk_v = 0, bool sage_attn_qk_int8 = false, - int64_t num_contexts = 0, int64_t num_ctx_tokens = 0, bool trtllm_gen_jit_warmup = false, - std::optional aux_kv_cache_pool_ptr = std::nullopt, bool const is_cross = false, - std::optional cross_kv = std::nullopt, - std::optional relative_attention_bias = std::nullopt, int64_t relative_attention_max_distance = 0, - std::optional spec_decoding_target_max_draft_tokens = std::nullopt, - std::optional quant_scale_qkv = std::nullopt, - std::optional dsv4_inv_rope_cos_sin_cache = std::nullopt, bool enable_dsv4_epilogue_fusion = false, - bool const force_prepare_spec_dec_tree_mask = false, std::optional const max_num_sequences = std::nullopt, - std::optional kv_norm_weight = std::nullopt, double kv_norm_eps = 1e-6, - double skip_correction_threshold = 0.0, std::optional uses_spcompress = std::nullopt); - struct KvCachePoolPointers { void* primaryPoolPtr{nullptr}; @@ -107,6 +56,833 @@ struct KvCachePoolPointers void* secondaryBlockScalePoolPtr{nullptr}; }; +#define TRTLLM_FMHA_PARAM_FIELD(name, cpp_type) cpp_type name{}; + +/// An attention layer's fixed shape, handed to the op once at construction. +/// Generated from StaticAttentionConfig. +struct StaticAttentionConfig +{ +#include "tensorrt_llm/thop/static_attention_config_fields.inc" +}; + +/// Sparse inputs an attention module hands to its backend. Generated from +/// SparseBackendForwardArgs; see FmhaParams for how a schema class is written. +struct SparseBackendForwardArgs +{ +#include "tensorrt_llm/thop/sparse_backend_forward_args_accessors.inc" +#include "tensorrt_llm/thop/sparse_backend_forward_args_fields.inc" +}; + +/// Sparse inputs a backend hands to the attention op. Generated from +/// SparseRuntimeParams. +struct SparseRuntimeParams +{ +#include "tensorrt_llm/thop/sparse_runtime_params_accessors.inc" +#include "tensorrt_llm/thop/sparse_runtime_params_fields.inc" +}; + +/// The arguments that vary per forward pass. Generated from AttentionForwardArgs and +/// held by value in FmhaParams, so the native layout mirrors the Python one. +struct AttentionForwardArgs +{ +#include "tensorrt_llm/thop/attention_forward_args_accessors.inc" +#include "tensorrt_llm/thop/attention_forward_args_fields.inc" + + // The attention entry point's dtype dispatch supplies T. + template + T* getLatentCache() const + { + return latent_cache.has_value() ? static_cast(latent_cache.value().data_ptr()) : nullptr; + } + + template + T* getQPe() const + { + return q_pe.has_value() ? static_cast(q_pe.value().data_ptr()) : nullptr; + } + + template + T* getCrossKv() const + { + return cross_kv.has_value() ? static_cast(cross_kv.value().data_ptr()) : nullptr; + } + + template + T* getRelativeAttentionBias() const + { + return relative_attention_bias.has_value() ? static_cast(relative_attention_bias.value().data_ptr()) + : nullptr; + } + + // Handwritten for the same reason as in FmhaParams: the native view of these + // buffers is not their dtype. + void* getOutputSf() const + { + return output_sf.has_value() ? output_sf.value().data_ptr() : nullptr; + } + + void* getQuantQBuffer() const + { + return quant_q_buffer.has_value() ? quant_q_buffer.value().data_ptr() : nullptr; + } + + float2* getMropeRotaryCosSin() const + { + return mrope_rotary_cos_sin.has_value() ? static_cast(mrope_rotary_cos_sin.value().data_ptr()) + : nullptr; + } + + float2* getSoftmaxStatsTensor() const + { + return softmax_stats_tensor.has_value() ? static_cast(softmax_stats_tensor.value().data_ptr()) + : nullptr; + } +}; + +/// The unify attention parameter struct: every phased entry point and every enqueue path +/// consumes it directly. Data members are generated from the Python schema +/// (fmha/interface.py, via scripts/generate_fmha_params.py), which also states the offset +/// contract in full. In short: device tensors arrive pre-sliced for the phase, while host +/// tensors, KV-cache block offsets and FP4 scaling factors arrive whole-batch: a pointer +/// accessor applies `seq_offset` / `token_offset` itself, so call sites never pass one. +/// Fixed-dtype pointer accessors are generated. Runtime-dispatched types and semantic +/// views stay handwritten so dtype dispatch, validation and offsets remain explicit. +/// +/// A C++-only field, for state the op derives rather than receives, is declared below the +/// generated block and filled during initialization; it stays invisible to Python. A field +/// visible to both is declared in the Python schema instead, and the build regenerates the +/// member and its binding. +class AttentionOp; + +struct FmhaParams +{ +#include "tensorrt_llm/thop/fmha_params_fields.inc" +#undef TRTLLM_FMHA_PARAM_FIELD + + // ---- handwritten derived state (deliberately outside the generated schema) ---- + // Both are filled by AttentionOp::prepare(): they are projections of the fields + // above plus the cache-layout arithmetic the op alone can do. + kernels::SparseAttentionParams sparse_params{}; + KvCachePoolPointers kv_cache_pool_pointers{}; + + /// Read straight off the inputs rather than mirrored into a field: a mirror can + /// disagree with the tensor it describes, and these are all cheap scalar queries. + tensorrt_llm::DataType getType() const + { + return tensorrt_llm::runtime::TorchUtils::dataType(qkv_or_q.scalar_type()); + } + + bool isFp8Out() const + { + return output.scalar_type() == torch::kFloat8_e4m3fn; + } + + bool isFp4Out() const + { + return output.scalar_type() == torch::kUInt8; + } + + bool hasSparseAttnIndices() const + { + return fwd.sparse_runtime_params.sparse_attn_indices.has_value() + && fwd.sparse_runtime_params.sparse_attn_indices.value().numel() > 0; + } + + bool hasPagedSparseAttnIndices() const + { + return hasSparseAttnIndices() && fwd.sparse_runtime_params.sparse_attn_offsets.has_value() + && fwd.sparse_runtime_params.sparse_attn_offsets.value().numel() > 0; + } + + // Accessors for the members above. These mirror what the generator emitted while the + // fields were part of the schema; nothing assigns them, so they resolve to nullptr. + template + T* getQkvBias() const + { + return qkv_bias.has_value() ? static_cast(qkv_bias.value().data_ptr()) : nullptr; + } + + template + T* getAlibiSlopes() const + { + return alibi_slopes.has_value() ? static_cast(alibi_slopes.value().data_ptr()) : nullptr; + } + + bool* getAttentionMask() const + { + return attention_mask.has_value() ? attention_mask.value().data_ptr() : nullptr; + } + + std::uint32_t* getAttentionPackedMask() const + { + return attention_packed_mask.has_value() ? attention_packed_mask.value().data_ptr() : nullptr; + } + + void* getKeyValueCache() const + { + return key_value_cache.has_value() ? key_value_cache.value().data_ptr() : nullptr; + } + + float* getLognScalingPtr() const + { + return logn_scaling_ptr.has_value() ? logn_scaling_ptr.value().data_ptr() : nullptr; + } + + float* getOutSfScale() const + { + return out_sf_scale.has_value() ? out_sf_scale.value().data_ptr() : nullptr; + } + + std::int32_t* getEncoderInputLengths() const + { + return encoder_input_lengths.has_value() ? encoder_input_lengths.value().data_ptr() : nullptr; + } + + std::int64_t* getRuntimePerfKnobs() const + { + return runtime_perf_knobs.has_value() ? runtime_perf_knobs.value().data_ptr() : nullptr; + } + + bool hasUnpagedSparseAttnIndices() const + { + return hasSparseAttnIndices() && !hasPagedSparseAttnIndices(); + } + + std::int64_t attention_mask_stride = 0; + tensorrt_llm::kernels::BlockSparseParams block_sparse_params{}; + bool can_use_one_more_block = false; + std::int64_t cross_kv_length = 0; + std::optional encoder_input_lengths{}; + bool has_full_attention_mask = false; + std::int64_t input_seq_length = 0; + std::int64_t max_blocks_per_sequence = 0; + std::int64_t max_cyclic_attention_window_size = 0; + std::int64_t max_past_kv_length = 0; + std::int64_t num_encoder_tokens = 0; + std::optional out_sf_scale{}; + std::optional qkv_bias{}; + std::optional alibi_slopes{}; + std::optional key_value_cache{}; + std::optional logn_scaling_ptr{}; + std::optional attention_mask{}; + std::optional attention_packed_mask{}; + bool pos_shift_enabled = false; + bool qkv_bias_enabled = false; + std::int64_t relative_attention_bias_stride = 0; + std::optional runtime_perf_knobs{}; + std::int64_t sink_token_length = 0; + bool spec_decoding_is_generation_length_variable = false; + std::int64_t spec_decoding_max_generation_length = 1; + std::int64_t spec_decoding_target_max_gen_len = 0; + std::int64_t total_kv_len = 0; + bool unfuse_qkv_gemm = false; + std::int64_t unidirectional = 1; + bool use_logn_scaling = false; + bool use_sparse_attention = false; + std::int64_t vision_length = -1; + std::int64_t vision_start = -1; + + /// Build the MlaParams the MLA kernels take. Returned by value: the kernels hold a + /// pointer to it for the duration of the launch, so the caller owns the storage. + /// Defined out of line because they reach for AttentionOp, which is declared below. + template + kernels::MlaParams buildContextMlaParams(AttentionOp const& op) const; + + template + kernels::MlaParams buildGenerationMlaParams(AttentionOp const& op) const; + + /// Fills the fields flash-MLA generation adds on top of buildGenerationMlaParams(). + template + void addFlashMlaGenerationParams(kernels::MlaParams& mla) const; + + /// The tail both builders share. + template + /// Takes the op because the MLA dimensions, the head count and the position-embedding + /// type are layer configuration, which FmhaParams deliberately does not carry. + void finalizeMlaParams(kernels::MlaParams& mla, AttentionOp const& op) const; + + // ---- generated accessors: fixed-dtype tensor views ---- +#include "tensorrt_llm/thop/fmha_params_accessors.inc" + + // ---- generated forwarding: reaches through `fwd` so call sites stay flat ---- +#include "tensorrt_llm/thop/fmha_params_forwarding.inc" + + // The attention entry point's dtype dispatch supplies T. + template + T* getQkvOrQ() const + { + return static_cast(qkv_or_q.data_ptr()); + } + + template + T* getK() const + { + return k.has_value() ? static_cast(k.value().data_ptr()) : nullptr; + } + + template + T* getV() const + { + return v.has_value() ? static_cast(v.value().data_ptr()) : nullptr; + } + + template + T* getLatentCache() const + { + return fwd.getLatentCache(); + } + + template + T* getQPe() const + { + return fwd.getQPe(); + } + + template + T* getCrossKv() const + { + return fwd.getCrossKv(); + } + + template + T* getRelativeAttentionBias() const + { + return fwd.getRelativeAttentionBias(); + } + + // ---- hand-written accessors: the native view differs from the tensor's dtype ---- + // These read the buffer as something the schema cannot name: an opaque pointer, or + // a float32 buffer consumed as pairs. + void* getWorkspace() const + { + return workspace.data_ptr(); + } + + /// The multi-CTA-KV counter semaphores are an opaque byte-sized scratch buffer. + void* getMultiCtasKvCounter() const + { + return multi_ctas_kv_counter.has_value() ? multi_ctas_kv_counter.value().data_ptr() : nullptr; + } + + void* getOutput() const + { + return output.data_ptr(); + } + + void* getOutputSf() const + { + return fwd.getOutputSf(); + } + + void* getQuantQBuffer() const + { + return fwd.getQuantQBuffer(); + } + + float2* getRotaryCosSin() const + { + return rotary_cos_sin.has_value() ? static_cast(rotary_cos_sin.value().data_ptr()) : nullptr; + } + + float2* getMropeRotaryCosSin() const + { + return fwd.getMropeRotaryCosSin(); + } + + float2* getSoftmaxStatsTensor() const + { + return fwd.getSoftmaxStatsTensor(); + } + + // ---- hand-written accessors: sizes, derived pointers, and anything taking an index ---- + // Sequence lengths, windows, and paged-KV metadata. + /// Block offsets for this phase's sequences; `seq_offset` is applied here. + kernels::KVBlockArray::DataType* getKvCacheBlockOffsets(int32_t poolIndex) const + { + return kv_cache_block_offsets.has_value() ? static_cast( + kv_cache_block_offsets.value().index({poolIndex, seq_offset}).data_ptr()) + : nullptr; + } + + // The KV-cache pool base pointers are derived state, not inputs: they are the + // int64 values stored in `host_kv_cache_pool_pointers` shifted by this layer's + // intra-pool offset, resolved once the op knows the cache element size. Never + // Python-facing. + void* getHostPrimaryPoolPtr() const + { + return kv_cache_pool_pointers.primaryPoolPtr; + } + + void* getHostSecondaryPoolPtr() const + { + return kv_cache_pool_pointers.secondaryPoolPtr; + } + + void* getHostPrimaryBlockScalePoolPtr() const + { + return kv_cache_pool_pointers.primaryBlockScalePoolPtr; + } + + void* getHostSecondaryBlockScalePoolPtr() const + { + return kv_cache_pool_pointers.secondaryBlockScalePoolPtr; + } + + // Quantization scales and output quantization. + // RoPE, ALiBi, and logn data. + // MLA input/cache data. + // MRoPE and Helix position data. + // Context chunking, Helix reduction, and softmax statistics. + // Speculative decoding masks and offsets. + // Attention sinks. + // Sparse attention, sparse MLA, and SageAttention runtime data. + // Packed-varlen context boundaries and scheduler state. + // MLA scales, quantized-Q buffers, and FlashMLA metadata. + // Cross attention. + // DeepSeek-V4 FP8-Q/epilogue fusion. + /// Host past-KV lengths for this phase's sequences. + /// + /// `seq_offset` is applied here, so the returned array is phase-local and lines up + /// with the device tensors, which the caller slices before handing them over. + int32_t* getHostPastKeyValueLengths() const + { + // A None source lowers to an undefined tensor rather than an error, so name the + // field here instead of failing inside data_ptr with only a dtype to go on. + TORCH_CHECK(host_past_key_value_lengths.defined(), "FmhaParams.host_past_key_value_lengths is unset."); + return static_cast(host_past_key_value_lengths.data_ptr()) + seq_offset; + } + + /// Host context lengths for this phase's sequences, offset as above. + int32_t* getHostContextLengths() const + { + TORCH_CHECK(host_context_lengths.defined(), "FmhaParams.host_context_lengths is unset."); + return static_cast(host_context_lengths.data_ptr()) + seq_offset; + } + + /// A phase with no sequences has no extent; at::max() rejects an empty input, so the + /// empty range yields 0, matching the sum-based reduction below. + int32_t getMaxHostPastKeyValueLength(int64_t seqOffset, int64_t numSeqs) const + { + if (numSeqs <= 0) + { + return 0; + } + return host_past_key_value_lengths.slice(0, seqOffset, seqOffset + numSeqs).max().item(); + } + + int32_t getMaxHostContextLength(int64_t seqOffset, int64_t numSeqs) const + { + if (numSeqs <= 0) + { + return 0; + } + return host_context_lengths.slice(0, seqOffset, seqOffset + numSeqs).max().item(); + } + + /// Total attended-KV length for this phase. Prefer explicit phase totals when + /// supplied; otherwise sum the host KV lengths over the phase's sequences. + int32_t getTotalKvLen(int64_t seqOffset, int64_t numSeqs, bool isGen) const + { + if (numSeqs <= 0) + { + return 0; + } + if (host_total_kv_lens.has_value()) + { + return host_total_kv_lens.value().select(0, isGen ? 1 : 0).item(); + } + return host_past_key_value_lengths.slice(0, seqOffset, seqOffset + numSeqs).sum().item(); + } + + int getCacheIndirectionWindowSize(int defaultValue) const + { + return cache_indirection.has_value() ? static_cast(cache_indirection.value().size(2)) : defaultValue; + } + + bool hasKvCache() const + { + return kv_cache_block_offsets.has_value() && host_kv_cache_pool_pointers.has_value() + && host_kv_cache_pool_mapping.has_value(); + } + + torch::Tensor const& getHostKvCachePoolPointers() const + { + return host_kv_cache_pool_pointers.value(); + } + + int getMaxBlocksPerSequence() const + { + return kv_cache_block_offsets.has_value() ? static_cast(kv_cache_block_offsets.value().size(-1)) : 0; + } + + // NOTE: `host_kv_cache_pool_mapping` is indexed by the layer's index *within its own + // KV-cache manager* (`local_layer_idx`), not by the model-global `layer_idx`. The two + // differ for draft / MTP models, whose mapping holds only their own layers. + int32_t getKvCachePoolIndex(int64_t localLayerIdx) const + { + return host_kv_cache_pool_mapping.has_value() + ? checkedPoolMapping(localLayerIdx).index({localLayerIdx, 0}).item() + : 0; + } + + int32_t getLayerIdxInCachePool(int64_t localLayerIdx) const + { + return host_kv_cache_pool_mapping.has_value() + ? checkedPoolMapping(localLayerIdx).index({localLayerIdx, 1}).item() + : 0; + } + + torch::Tensor const& checkedPoolMapping(int64_t localLayerIdx) const + { + auto const& mapping = host_kv_cache_pool_mapping.value(); + TORCH_CHECK(localLayerIdx >= 0 && localLayerIdx < mapping.size(0), "local_layer_idx ", localLayerIdx, + " is out of range for host_kv_cache_pool_mapping with ", mapping.size(0), + " layers. This index must be the layer's position within its own KV-cache manager, " + "not the model-global layer index."); + return mapping; + } + + int64_t getMlaLayerNum() const + { + return host_kv_cache_pool_mapping.has_value() ? host_kv_cache_pool_mapping.value().size(0) : 0; + } + + int64_t getSparseAttnIndicesStride() const + { + return fwd.sparse_runtime_params.sparse_attn_indices.has_value() + ? fwd.sparse_runtime_params.sparse_attn_indices.value().size(-1) + : 0; + } + + int32_t getCrossKvNumTokens() const + { + return fwd.cross_kv.has_value() ? static_cast(fwd.cross_kv.value().size(0)) : 0; + } + + char* getSparseKvCachePool(int32_t poolIndex) const + { + return host_kv_cache_pool_pointers.has_value() + ? reinterpret_cast(host_kv_cache_pool_pointers.value().index({poolIndex, 0}).item()) + : nullptr; + } +}; + +class AttentionOp +{ +public: + /// Acquires the handles the op keeps for its lifetime and derives everything that + /// follows from the layer's fixed shape. Deliberately not per call: cublasCreate() + /// is illegal while a stream is capturing, so the op has to be built outside + /// capture and reused for every call from its layer. + explicit AttentionOp(StaticAttentionConfig const& cfg); + + using RotaryScalingType = tensorrt_llm::kernels::RotaryScalingType; + using PositionEmbeddingType = tensorrt_llm::kernels::PositionEmbeddingType; + using AttentionMaskType = tensorrt_llm::kernels::AttentionMaskType; + + /// Derives the per-call state into \p p. Kernel runners are fixed by the constructor; + /// `mMultiBlockMode` remains a per-call runtime knob read from + /// `p.runtime_perf_knobs` at generation time. + int prepare(FmhaParams& p, bool isGen); + + /// One per entry point, instantiated per activation dtype. prepare() is called from + /// here so the dtype-independent and dtype-dependent halves of the setup sit + /// together, and so the MlaParams this scope owns outlives the launch. + template + void runContextImpl(FmhaParams& p); + template + void runGenerationImpl(FmhaParams& p); + template + void runMlaGenerationImpl(FmhaParams& p); + + /// Phased attention entry points. Reuse one op per layer: it owns the cuBLAS handle, + /// which is expensive to create and cannot be created while a stream is capturing. + /// The caller's params are filled in place: Python builds a fresh holder for every + /// call, so there is nothing to preserve and nothing to copy. + void runContext(FmhaParams& params); + void runGeneration(FmhaParams& params); + void runMlaGeneration(FmhaParams& params); + + /// Max of the context and generation workspace byte requirements. + int64_t getAttentionWorkspaceSize(FmhaParams const& params, int64_t numTokens, int64_t maxAttentionWindowSize, + int64_t numGenTokens, int64_t maxBlocksPerSequence); + + [[nodiscard]] size_t getFmhaMultiCtasKvScratchSize(FmhaParams const& p) const noexcept; + [[nodiscard]] int getHeadSize(bool checkInit = true) const; + [[nodiscard]] int getMaxNumSeqLenTile(FmhaParams const& p, int batch_beam_size = 1) const; + [[nodiscard]] size_t getWorkspaceSizeForContext(FmhaParams const& p, int32_t nbReq, int32_t max_input_length, + int32_t cross_kv_length = 0, int32_t max_num_tokens = 0, int32_t total_kv_len = 0) const noexcept; + // Per-token byte cost of the context-MLA K/V dequant staging buffers, whose size scales with the summed + // attended KV length (`total_kv_len`). Only the fp8 context-MLA separate-Q/KV path stages these buffers; + // every other path (incl. sparse MLA, which reads K/V straight from the paged cache) returns 0. Single + // source of truth shared by getWorkspaceSizeForContext (runtime sizing) and the KV-cache estimator, so + // the two cannot drift. + [[nodiscard]] static size_t contextMlaWorkspaceBytesPerToken(int32_t numAttnHeads, int32_t qkRopeHeadDim, + int32_t qkNopeHeadDim, int32_t vHeadDim, bool fp8ContextMla, bool separateQAndKvInput, bool sparseMla) noexcept; + // total_num_seq is the sum of beam_width for multiple requests + [[nodiscard]] size_t getWorkspaceSizeForGeneration(FmhaParams const& p, int32_t total_num_seq, + int32_t max_attention_window_size, int32_t max_num_tokens, int32_t max_blocks_per_sequence) const noexcept; + + template + int enqueueContext(FmhaParams const& p, kernels::MlaParams* mlaParam, cudaStream_t stream); + + template + int enqueueGeneration(FmhaParams const& p, cudaStream_t stream); + + template + int mlaGeneration(kernels::MlaParams& params, FmhaParams const& p, cudaStream_t stream); + + int getFlashMlaNumSmParts(int s_q, int num_heads, int num_kv_heads, int head_size_v) const + { + static constexpr int block_size_m = 64; + int num_heads_per_head_k = s_q * num_heads / num_kv_heads; + int sm_cnt = mMultiProcessorCount; + int num_sm_parts = sm_cnt / num_kv_heads / cutlass::ceil_div(num_heads_per_head_k, block_size_m); + return num_sm_parts; + } + + static int getFlashMlaNumSmPartsStatic(int s_q, int num_heads, int num_kv_heads, int head_size_v) + { + static constexpr int block_size_m = 64; + int num_heads_per_head_k = s_q * num_heads / num_kv_heads; + int device; + cudaGetDevice(&device); + int sm_cnt; + cudaDeviceGetAttribute(&sm_cnt, cudaDevAttrMultiProcessorCount, device); + int num_sm_parts = sm_cnt / num_kv_heads / cutlass::ceil_div(num_heads_per_head_k, block_size_m); + return num_sm_parts; + } + + template + int getKvCacheElemSizeInBits(FmhaParams const& p) const + { + return getKvCacheElemSizeInBits(mCfg.quant_mode, sizeof(T)); + } + + static int getKvCacheElemSizeInBits(tensorrt_llm::common::QuantMode quantMode, size_t dTypeSize) + { + if (quantMode.hasInt8KvCache() || quantMode.hasFp8KvCache()) + { + return 8; + } + else if (quantMode.hasFp4KvCache()) + { + return 4; + } + return dTypeSize * 8; + } + + /// Defaulted so the dtype dispatch can name this directly; paged KV is the only + /// layout the generation path prepares for. + template + void prepareEnqueueGeneration(FmhaParams const& p); + + template + bool convertMMHAParamsToXQAParams( + tensorrt_llm::kernels::XQAParams& xqaParams, FmhaParams const& p, bool forConfigurePlugin); + + /// The op's own derived state, for debugging a dispatch decision. + [[nodiscard]] std::string toString() const; + + // --------------------------------------------------------------------------- + // Predicates over the layer's configuration, read from mCfg, and over the per-call + // state, which is passed explicitly. + // --------------------------------------------------------------------------- + [[nodiscard]] bool isRelativePosition() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kRELATIVE; + } + + [[nodiscard]] bool isALiBi() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI + || mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI_WITH_SCALE; + } + + [[nodiscard]] bool isAliBiWithScale() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kALIBI_WITH_SCALE; + } + + [[nodiscard]] bool isRoPE() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_GPTJ + || mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_GPT_NEOX + || mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kLONG_ROPE + || mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kYARN + || mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_M; + } + + [[nodiscard]] bool isLongRoPE() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kLONG_ROPE; + } + + [[nodiscard]] bool isUnfusedCrossAttention(FmhaParams const& p) const + { + return !mEnableContextFMHA && p.is_cross; + } + + [[nodiscard]] bool isMRoPE() const + { + return mCfg.position_embedding_type == tensorrt_llm::kernels::PositionEmbeddingType::kROPE_M; + } + + [[nodiscard]] bool isLognScaling(FmhaParams const& p) const + { + return p.use_logn_scaling; + } + + [[nodiscard]] bool isCrossAttention(FmhaParams const& p) const + { + return p.is_cross; + } + + [[nodiscard]] bool useKVCache(FmhaParams const& p) const + { + return p.hasKvCache(); + } + + [[nodiscard]] bool useCustomMask() const + { + return mCfg.mask_type == AttentionMaskType::CUSTOM_MASK; + } + + [[nodiscard]] bool useFullCustomMask(FmhaParams const& p) const + { + return useCustomMask() && p.has_full_attention_mask; + } + + [[nodiscard]] bool usePackedCustomMask(FmhaParams const& p) const + { + return useCustomMask() && mEnableContextFMHA; + } + + [[nodiscard]] bool isMLAEnabled() const + { + return mCfg.is_mla_enable; + } + + [[nodiscard]] bool useSparseAttention(FmhaParams const& p) const + { + return p.use_sparse_attention && mPagedKVCache && mEnableXQA; + } + + [[nodiscard]] bool useTllmGenSparseAttentionPaged(FmhaParams const& p) const + { + return p.hasPagedSparseAttnIndices() && useSparseAttention(p); + } + + // Config-only variants, for initialize(): it runs at construction and has no call to + // read, so the same questions are answered from mCfg. + [[nodiscard]] bool cfgUseSparseMLA() const + { + return mCfg.use_sparse_attention && mUseTllmGen && mCfg.is_mla_enable; + } + + [[nodiscard]] bool cfgUseTllmGenSparseAttention() const + { + return cfgUseSparseMLA() || (mCfg.use_sparse_attention && mUseTllmGen && mCfg.use_tllm_gen_sparse_attention); + } + + [[nodiscard]] bool useSparseMLA(FmhaParams const& p) const + { + return p.use_sparse_attention && mUseTllmGen && mCfg.is_mla_enable; + } + + [[nodiscard]] bool useTllmGenSparseAttention(FmhaParams const& p) const + { + return useSparseMLA(p) || (p.use_sparse_attention && mUseTllmGen && p.hasUnpagedSparseAttnIndices()); + } + + [[nodiscard]] StaticAttentionConfig const& config() const + { + return mCfg; + } + + [[nodiscard]] tensorrt_llm::kernels::MlaMetaParams const& mlaMeta() const + { + return mMLAParams; + } + + [[nodiscard]] int smVersion() const + { + return mSM; + } + + [[nodiscard]] bool supportsNvFp4Output() const + { + return mEnableContextFMHA && mEnableXQA; + } + + [[nodiscard]] int getMultiProcessorCount() const + { + return mMultiProcessorCount; + } + + // --------------------------------------------------------------------------- + // Op state: the derived, hardware and dispatcher members. Per-call configuration + // scalars are read off `FmhaParams` directly and are not mirrored here. + // --------------------------------------------------------------------------- + + int mNumKVHeads = -1; + int mHeadSize = -1; + + bool mPagedKVCache = true; + bool mFP8ContextFMHA = false; + bool mFP8AttenOutput = false; + bool mFP8ContextMLA = false; + bool mFP8GenerationMLA = false; + bool mIsGenerationMLA = false; + bool mUseGenFlashMLA = false; + // Static sparse MLA reads a separately dequantized FP8 scratch pool, so an NVFP4 + // paged cache still runs the FP8 kernels. + bool mUseNvfp4MlaKvCache = false; + // Skip correction when the row-max increase is within this base-2 threshold. + float mSkipCorrectionThreshold = 0.0F; + + // Equal to the full head counts: this op always runs on a single rank. + int mNumAttnHeads = -1; + // The layer's fixed configuration, captured at construction. Every value the + // kernel-runner selection and the per-call paths read that does not vary between + // calls lives here rather than being resent on each FmhaParams. + StaticAttentionConfig mCfg{}; + tensorrt_llm::kernels::MlaMetaParams mMLAParams{}; + int mNumAttnKVHeads = -1; + + // fmha runner (enabled by default) + // flag: disabled = 0, enabled = 1, enabled with fp32 accumulation = 2 + bool mEnableContextFMHA = true; + bool mFMHAForceFP32Acc = false; + bool mMultiBlockMode = true; + bool mEnableXQA = true; + + bool mFuseFp4Quant = false; + +#ifdef SKIP_SOFTMAX_STAT + uint32_t* mSkipSoftmaxTotalBlocks; + uint32_t* mSkipSoftmaxSkippedBlocks; +#endif + +private: + void initialize(); + int finishPrepare(FmhaParams& p, bool isGen); + + static constexpr int kReservedMaxSeqLenTilePerSeq = 64; + + int mSM = tensorrt_llm::common::getSMVersion(); + bool mUseTllmGen = (mSM >= 100) && (mSM != 120) && (mSM != 121); + bool mForceMultiBlockWarned = false; + int mMultiProcessorCount = tensorrt_llm::common::getMultiProcessorCount(); + int mMaxSharedMemoryPerBlockOptin = tensorrt_llm::common::getMaxSharedMemoryPerBlockOptin(); + std::shared_ptr mDriver; + std::unique_ptr mDecoderFMHARunner; + std::unique_ptr mFmhaDispatcher; + std::unique_ptr mXqaDispatcher; + std::unique_ptr mTllmGenFMHARunner; + std::unique_ptr mCublasWrapper; +}; + struct KvCachePoolMapping { int32_t poolIndex{0}; @@ -118,7 +894,21 @@ KvCachePoolMapping readKvCachePoolMapping(at::Tensor const& hostKvCachePoolMappi KvCachePoolPointers buildKvCachePoolPointers(at::Tensor const& hostKvCachePoolPointers, int32_t poolIndex, int64_t intraPoolOffset, int64_t blockSize, int32_t layerIdxInCachePool, int32_t kvFactor, bool isFp4KvCache); -common::op::KvCacheBuffers buildPagedKvCacheBuffers( +template +struct KvCacheBuffers +{ + KVCacheBuffer kvCacheBuffer; + KVCacheBuffer kvScaleCacheBuffer; +}; + +template +KvCacheBuffers buildKvCacheBuffers(int32_t batchSize, int32_t maxBlocksPerSeq, int32_t tokensPerBlock, + int32_t sizePerToken, int32_t cyclicAttentionWindowSize, int32_t maxCyclicAttentionWindowSize, int32_t sinkTokenLen, + bool canUseOneMoreBlock, void* primaryPoolPtr, void* secondaryPoolPtr, void* primaryBlockScalePoolPtr, + void* secondaryBlockScalePoolPtr, kernels::KVBlockArray::DataType* blockOffsets, bool hasFp4KvCache, + int32_t maxAttentionWindowSize = 0, void* keyValueCache = nullptr); + +KvCacheBuffers buildPagedKvCacheBuffers( std::optional const& kv_cache_block_offsets, std::optional const& host_kv_cache_pool_pointers, std::optional const& host_kv_cache_pool_mapping, common::QuantMode quantMode, int64_t layer_idx, diff --git a/cpp/tensorrt_llm/thop/dsv3RopeOp.cpp b/cpp/tensorrt_llm/thop/dsv3RopeOp.cpp index e3a2eb2ef200..a11f0a8956ca 100644 --- a/cpp/tensorrt_llm/thop/dsv3RopeOp.cpp +++ b/cpp/tensorrt_llm/thop/dsv3RopeOp.cpp @@ -15,7 +15,6 @@ * limitations under the License. */ -#include "tensorrt_llm/common/attentionOp.h" #include "tensorrt_llm/common/dataType.h" #include "tensorrt_llm/common/memoryUtils.h" #include "tensorrt_llm/kernels/gptKernels.h" diff --git a/cpp/tensorrt_llm/thop/trtllmGenQKVProcessOp.cpp b/cpp/tensorrt_llm/thop/trtllmGenQKVProcessOp.cpp index 5e61bc0c28a5..f1cf2d2b9e60 100644 --- a/cpp/tensorrt_llm/thop/trtllmGenQKVProcessOp.cpp +++ b/cpp/tensorrt_llm/thop/trtllmGenQKVProcessOp.cpp @@ -15,7 +15,6 @@ * limitations under the License. */ -#include "tensorrt_llm/common/attentionOp.h" #include "tensorrt_llm/common/cudaUtils.h" #include "tensorrt_llm/common/quantization.h" #include "tensorrt_llm/common/tllmDataType.h" @@ -36,7 +35,6 @@ TRTLLM_NAMESPACE_BEGIN namespace torch_ext { -using tensorrt_llm::common::op::AttentionOp; using tensorrt_llm::kernels::AttentionMaskType; using tensorrt_llm::kernels::BlockSparseParams; using tensorrt_llm::kernels::BuildDecoderInfoParams; diff --git a/cpp/tests/unit_tests/common/attentionWorkspaceTest.cpp b/cpp/tests/unit_tests/common/attentionWorkspaceTest.cpp index b0f2dc0dcd1a..3883e1c5b5c7 100644 --- a/cpp/tests/unit_tests/common/attentionWorkspaceTest.cpp +++ b/cpp/tests/unit_tests/common/attentionWorkspaceTest.cpp @@ -15,7 +15,6 @@ */ #include "tensorrt_llm/common/attentionWorkspace.h" -#include "tensorrt_llm/common/attentionOp.h" #include "tensorrt_llm/common/workspace.h" @@ -57,73 +56,8 @@ void expectNextSlice(char const* name, Slice const& slice, size_t size, size_t& expectedOffset += tc::alignSize(size, kAlignment); } -constexpr int32_t kBatchSize = 2; -constexpr int32_t kInputSequenceLength = 11; -constexpr int32_t kCrossKvLength = 7; -constexpr int32_t kPackedTokenCount = 14; -constexpr int32_t kHeadSize = 8; - -void configureUnfusedAttention(tcop::AttentionOp& op, bool crossAttention) -{ - op.mNumHeads = 1; - op.mNumKVHeads = 1; - op.mHeadSize = kHeadSize; - op.mNumAttnHeads = 1; - op.mNumAttnKVHeads = 1; - op.mEnableContextFMHA = false; - op.mCrossAttention = crossAttention; -} - -size_t expectedUnfusedContextWorkspace(bool crossAttention) -{ - constexpr size_t kElementSize = sizeof(half); - size_t const batchSize = kBatchSize; - size_t const inputSequenceLength = kInputSequenceLength; - size_t const kvSequenceLength = crossAttention ? kCrossKvLength : kInputSequenceLength; - size_t const paddedTokenCount = batchSize * inputSequenceLength; - size_t const paddedKvTokenCount = batchSize * kvSequenceLength; - - tcop::AttentionContextWorkspaceSizes sizes{}; - sizes.attentionMask = kElementSize * paddedTokenCount * kvSequenceLength; - sizes.cuQSeqlens = sizeof(int) * (batchSize + 1); - sizes.cuKvSeqlens = sizes.cuQSeqlens; - sizes.cuMaskRows = sizes.cuQSeqlens; - sizes.qBuf = kElementSize * paddedTokenCount * kHeadSize; - sizes.kBuf = kElementSize * paddedKvTokenCount * kHeadSize; - sizes.vBuf = sizes.kBuf; - sizes.qkBuf = kElementSize * batchSize * inputSequenceLength * kvSequenceLength; - sizes.qkvBuf = kElementSize * paddedTokenCount * kHeadSize; - sizes.qkFloatBuf = sizeof(float) * batchSize * inputSequenceLength * kvSequenceLength; - sizes.paddingOffset = sizeof(int) * paddedTokenCount; - sizes.encoderPaddingOffset = sizeof(int) * paddedKvTokenCount; - sizes.tokensInfo = sizeof(int2) * kPackedTokenCount; - return tcop::AttentionWorkspaceManager::buildContextLayout(sizes).totalSize; -} - -size_t getUnfusedContextWorkspace(tcop::AttentionOp const& op) -{ - return op.getWorkspaceSizeForContext( - tensorrt_llm::DataType::kHALF, kBatchSize, kInputSequenceLength, kCrossKvLength, kPackedTokenCount); -} - } // namespace -TEST(AttentionWorkspaceManagerTest, RaggedUnfusedSelfAttentionUsesPaddedTokenCounts) -{ - tcop::AttentionOp op; - configureUnfusedAttention(op, false); - - EXPECT_EQ(getUnfusedContextWorkspace(op), expectedUnfusedContextWorkspace(false)); -} - -TEST(AttentionWorkspaceManagerTest, RaggedUnfusedCrossAttentionUsesPaddedTokenCounts) -{ - tcop::AttentionOp op; - configureUnfusedAttention(op, true); - - EXPECT_EQ(getUnfusedContextWorkspace(op), expectedUnfusedContextWorkspace(true)); -} - TEST(AttentionWorkspaceManagerTest, ContextLayoutMatchesAttentionOpOrdering) { tcop::AttentionContextWorkspaceSizes sizes{}; @@ -152,7 +86,6 @@ TEST(AttentionWorkspaceManagerTest, ContextLayoutMatchesAttentionOpOrdering) sizes.sageQScale = 101; sizes.sageKScale = 103; sizes.sageVScale = 107; - sizes.cpWorkspace = 109; auto const layout = tcop::AttentionWorkspaceManager::buildContextLayout(sizes); @@ -182,17 +115,11 @@ TEST(AttentionWorkspaceManagerTest, ContextLayoutMatchesAttentionOpOrdering) expectNextSlice("sageQScale", layout.sageQScale, sizes.sageQScale, expectedOffset); expectNextSlice("sageKScale", layout.sageKScale, sizes.sageKScale, expectedOffset); expectNextSlice("sageVScale", layout.sageVScale, sizes.sageVScale, expectedOffset); - expectNextSlice("cpWorkspace", layout.cpWorkspace, sizes.cpWorkspace, expectedOffset); EXPECT_EQ(layout.totalSize, expectedOffset); } TEST(AttentionWorkspaceManagerTest, MaterializeContextReturnsTypedViewsAndNullZeroSlices) { - constexpr size_t kCpMaxPaddedSequenceLength = 2; - constexpr int kHeadSize = 4; - constexpr int kNumHeads = 2; - constexpr int kNumKvHeads = 1; - constexpr size_t kCpBufferElements = kCpMaxPaddedSequenceLength * kHeadSize * (kNumHeads + 2 * kNumKvHeads); tcop::AttentionContextWorkspaceSizes sizes{}; sizes.cublasWorkspace = 0; @@ -205,7 +132,6 @@ TEST(AttentionWorkspaceManagerTest, MaterializeContextReturnsTypedViewsAndNullZe sizes.tokensInfo = sizeof(int2) * 2; sizes.fmhaTileCounter = sizeof(uint32_t); sizes.fmhaBmm1Scale = sizeof(float) * 2; - sizes.cpWorkspace = 2 * kCpBufferElements * sizeof(float) + sizeof(int) * 3; auto const layout = tcop::AttentionWorkspaceManager::buildContextLayout(sizes); @@ -213,8 +139,7 @@ TEST(AttentionWorkspaceManagerTest, MaterializeContextReturnsTypedViewsAndNullZe alignas(kAlignment) std::array workspace{}; auto* base = workspace.data(); - auto const views = tcop::AttentionWorkspaceManager::materializeContext( - workspace.data(), layout, kCpMaxPaddedSequenceLength, kHeadSize, kNumHeads, kNumKvHeads); + auto const views = tcop::AttentionWorkspaceManager::materializeContext(workspace.data(), layout); EXPECT_EQ(views.cublasWorkspace, nullptr); expectPtrAt(views.attentionMask, base, layout.attentionMask); @@ -226,25 +151,12 @@ TEST(AttentionWorkspaceManagerTest, MaterializeContextReturnsTypedViewsAndNullZe expectPtrAt(views.tokensInfo, base, layout.tokensInfo); expectPtrAt(views.fmhaTileCounter, base, layout.fmhaTileCounter); expectPtrAt(views.fmhaBmm1Scale, base, layout.fmhaBmm1Scale); - expectPtrAt(views.gatherInBuffer, base, layout.cpWorkspace); - - auto* const expectedGatherOutBuffer - = reinterpret_cast(base + layout.cpWorkspace.offset) + kCpBufferElements; - auto* const expectedCuCpPartialSeqlens = reinterpret_cast(expectedGatherOutBuffer + kCpBufferElements); - EXPECT_EQ(static_cast(views.gatherOutBuffer), static_cast(expectedGatherOutBuffer)); - EXPECT_EQ(static_cast(views.cuCpPartialSeqlens), static_cast(expectedCuCpPartialSeqlens)); } TEST(AttentionWorkspaceManagerTest, GenerationLayoutPlacesCpWorkspaceBeforePartialBuffers) { - constexpr size_t kCpMaxPaddedSequenceLength = 3; - constexpr int kNumHeads = 2; - constexpr int kNumKvHeads = 1; - constexpr int kHeadSize = 4; - constexpr size_t kCpBufferElements = kCpMaxPaddedSequenceLength * (kNumHeads + 2 * kNumKvHeads) * kHeadSize; tcop::AttentionGenerationWorkspaceSizes sizes{}; - sizes.cpWorkspace = 2 * kCpBufferElements * sizeof(float); sizes.partialOut = 33; sizes.partialSum = 9; sizes.partialMax = 0; @@ -253,7 +165,6 @@ TEST(AttentionWorkspaceManagerTest, GenerationLayoutPlacesCpWorkspaceBeforeParti auto const layout = tcop::AttentionWorkspaceManager::buildGenerationLayout(sizes); size_t expectedOffset = 0; - expectNextSlice("cpWorkspace", layout.cpWorkspace, sizes.cpWorkspace, expectedOffset); expectNextSlice("partialOut", layout.partialOut, sizes.partialOut, expectedOffset); expectNextSlice("partialSum", layout.partialSum, sizes.partialSum, expectedOffset); expectNextSlice("partialMax", layout.partialMax, sizes.partialMax, expectedOffset); @@ -264,12 +175,8 @@ TEST(AttentionWorkspaceManagerTest, GenerationLayoutPlacesCpWorkspaceBeforeParti alignas(kAlignment) std::array workspace{}; auto* base = workspace.data(); - auto const views = tcop::AttentionWorkspaceManager::materializeGeneration( - workspace.data(), layout, kCpMaxPaddedSequenceLength, kNumHeads, kNumKvHeads, kHeadSize); + auto const views = tcop::AttentionWorkspaceManager::materializeGeneration(workspace.data(), layout); - expectPtrAt(views.mhaOutput, base, layout.cpWorkspace); - auto* const expectedMhaInput = reinterpret_cast(base + layout.cpWorkspace.offset) + kCpBufferElements; - EXPECT_EQ(static_cast(views.mhaInput), static_cast(expectedMhaInput)); expectPtrAt(views.partialOut, base, layout.partialOut); expectPtrAt(views.partialSum, base, layout.partialSum); EXPECT_EQ(views.partialMax, nullptr); diff --git a/cpp/tests/unit_tests/thop/CMakeLists.txt b/cpp/tests/unit_tests/thop/CMakeLists.txt index 67b02d325be8..d83207049587 100644 --- a/cpp/tests/unit_tests/thop/CMakeLists.txt +++ b/cpp/tests/unit_tests/thop/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 @@ -14,6 +14,14 @@ # the License. if(${BUILD_PYT}) + add_gtest(attentionOpWorkspaceTest attentionOpWorkspaceTest.cpp) + target_include_directories( + attentionOpWorkspaceTest + PRIVATE ${PROJECT_BINARY_DIR}/tensorrt_llm/generated) + target_link_libraries( + attentionOpWorkspaceTest PRIVATE th_common ${Python3_LIBRARIES} + ${TORCH_LIBRARIES}) + add_gtest(thUtilsTest thUtilsTest.cpp) target_link_libraries(thUtilsTest PUBLIC th_utils ${Python3_LIBRARIES} ${TORCH_LIBRARIES}) diff --git a/cpp/tests/unit_tests/thop/attentionOpWorkspaceTest.cpp b/cpp/tests/unit_tests/thop/attentionOpWorkspaceTest.cpp new file mode 100644 index 000000000000..560599613113 --- /dev/null +++ b/cpp/tests/unit_tests/thop/attentionOpWorkspaceTest.cpp @@ -0,0 +1,162 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * 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/common/attentionWorkspace.h" +#include "tensorrt_llm/thop/attentionOp.h" + +#include + +namespace tcop = tensorrt_llm::common::op; +namespace tt = tensorrt_llm::torch_ext; + +namespace +{ + +constexpr int32_t kBatchSize = 2; +constexpr int32_t kInputSequenceLength = 11; +constexpr int32_t kCrossKvLength = 7; +constexpr int32_t kPackedTokenCount = 14; +constexpr int32_t kHeadSize = 64; + +tt::StaticAttentionConfig unfusedAttentionConfig(bool crossAttention) +{ + tt::StaticAttentionConfig cfg{}; + cfg.num_heads = 1; + cfg.num_kv_heads = 1; + cfg.head_size = kHeadSize; + cfg.type = tensorrt_llm::DataType::kHALF; + cfg.tokens_per_block = 64; + cfg.q_scaling = 1.0; + cfg.remove_padding = true; + cfg.cross_attention = crossAttention; + cfg.position_embedding_type = tt::AttentionOp::PositionEmbeddingType::kRELATIVE; + return cfg; +} + +size_t expectedUnfusedContextWorkspace(bool crossAttention) +{ + constexpr size_t kElementSize = sizeof(half); + size_t const batchSize = kBatchSize; + size_t const inputSequenceLength = kInputSequenceLength; + size_t const kvSequenceLength = crossAttention ? kCrossKvLength : kInputSequenceLength; + size_t const paddedTokenCount = batchSize * inputSequenceLength; + size_t const paddedKvTokenCount = batchSize * kvSequenceLength; + + tcop::AttentionContextWorkspaceSizes sizes{}; + sizes.attentionMask = kElementSize * paddedTokenCount * kvSequenceLength; + sizes.cuQSeqlens = sizeof(int) * (batchSize + 1); + sizes.cuKvSeqlens = sizes.cuQSeqlens; + sizes.cuMaskRows = sizes.cuQSeqlens; + sizes.qBuf = kElementSize * paddedTokenCount * kHeadSize; + sizes.kBuf = kElementSize * paddedKvTokenCount * kHeadSize; + sizes.vBuf = sizes.kBuf; + sizes.qkBuf = kElementSize * batchSize * inputSequenceLength * kvSequenceLength; + sizes.qkvBuf = kElementSize * paddedTokenCount * kHeadSize; + sizes.qkFloatBuf = sizeof(float) * batchSize * inputSequenceLength * kvSequenceLength; + sizes.paddingOffset = sizeof(int) * paddedTokenCount; + sizes.encoderPaddingOffset = sizeof(int) * paddedKvTokenCount; + sizes.tokensInfo = sizeof(int2) * kPackedTokenCount; + return tcop::AttentionWorkspaceManager::buildContextLayout(sizes).totalSize; +} + +size_t getUnfusedContextWorkspace(tt::AttentionOp const& op, bool crossAttention) +{ + tt::FmhaParams params{}; + params.qkv_or_q = at::empty({0}, at::kHalf); + params.is_cross = crossAttention; + return op.getWorkspaceSizeForContext(params, kBatchSize, kInputSequenceLength, kCrossKvLength, kPackedTokenCount); +} + +} // namespace + +TEST(AttentionOpWorkspaceTest, RaggedUnfusedSelfAttentionUsesPaddedTokenCounts) +{ + tt::AttentionOp op(unfusedAttentionConfig(false)); + + EXPECT_EQ(getUnfusedContextWorkspace(op, false), expectedUnfusedContextWorkspace(false)); +} + +TEST(AttentionOpWorkspaceTest, RaggedUnfusedCrossAttentionUsesPaddedTokenCounts) +{ + tt::AttentionOp op(unfusedAttentionConfig(true)); + + EXPECT_EQ(getUnfusedContextWorkspace(op, true), expectedUnfusedContextWorkspace(true)); +} + +class AttentionOpSpecDecodingTest : public ::testing::TestWithParam> +{ +}; + +TEST_P(AttentionOpSpecDecodingTest, OnlyActiveGenerationExposesSpeculativeInputs) +{ + auto const [isGen, enabled, active] = GetParam(); + auto cfg = unfusedAttentionConfig(false); + cfg.num_heads = 8; + cfg.head_size = 128; + cfg.use_kv_cache = true; + cfg.position_embedding_type = tt::AttentionOp::PositionEmbeddingType::kLEARNED_ABSOLUTE; + cfg.mask_type = tt::AttentionOp::AttentionMaskType::CAUSAL; + cfg.is_spec_decoding_enabled = enabled; + cfg.spec_decoding_target_max_gen_len = 4; + tt::AttentionOp op(cfg); + + bool const useSpecDecoding = isGen && enabled && active; + tt::FmhaParams params{}; + params.num_seqs = params.num_requests = params.max_num_requests = params.max_num_sequences = 1; + params.num_tokens = useSpecDecoding ? 4 : 1; + params.beam_width = 1; + params.max_context_length = params.max_seq_len = 64; + params.max_attention_window_size = params.cyclic_attention_window_size = 64; + params.qkv_or_q = at::empty({params.num_tokens, 1280}, at::kHalf); + params.output = at::empty({params.num_tokens, 1024}, at::kHalf); + params.workspace = at::empty({0}, at::kByte); + params.multi_ctas_kv_counter = at::zeros({64}, at::kByte); + params.host_past_key_value_lengths = at::full({1}, 4, at::kInt); + params.host_context_lengths = at::full({1}, 1, at::kInt); + params.sequence_length = params.host_past_key_value_lengths; + params.context_lengths = params.host_context_lengths; + params.fwd.is_fused_qkv = true; + params.use_spec_decoding = active; + params.spec_decoding_target_max_draft_tokens = 3; + params.spec_decoding_generation_lengths = at::full({1}, 4, at::kInt); + params.spec_decoding_position_offsets = at::zeros({1, 4}, at::kInt); + params.spec_decoding_packed_mask = at::zeros({1, 4, 1}, at::kInt); + params.spec_decoding_bl_tree_mask_offset = at::zeros({1}, at::kLong); + params.spec_decoding_bl_tree_mask = at::zeros({1}, at::kUInt32); + params.spec_bl_tree_first_sparse_mask_offset_kv = at::zeros({1}, at::kInt); + params.spec_decoding_is_generation_length_variable = true; + params.spec_decoding_max_generation_length = 4; + + auto const source = params; + ASSERT_EQ(op.prepare(params, isGen), 0); + EXPECT_EQ(params.spec_decoding_is_generation_length_variable, useSpecDecoding); + EXPECT_EQ(params.spec_decoding_max_generation_length, useSpecDecoding ? 4 : 1); + EXPECT_EQ(params.spec_decoding_target_max_gen_len, 4); + EXPECT_EQ(params.getSpecDecodingGenerationLengths(), + useSpecDecoding ? source.getSpecDecodingGenerationLengths() : nullptr); + EXPECT_EQ( + params.getSpecDecodingPositionOffsets(), useSpecDecoding ? source.getSpecDecodingPositionOffsets() : nullptr); + EXPECT_EQ(params.getSpecDecodingPackedMask(), useSpecDecoding ? source.getSpecDecodingPackedMask() : nullptr); + EXPECT_EQ( + params.getSpecDecodingBlTreeMaskOffset(), useSpecDecoding ? source.getSpecDecodingBlTreeMaskOffset() : nullptr); + EXPECT_EQ(params.getSpecDecodingBlTreeMask(), useSpecDecoding ? source.getSpecDecodingBlTreeMask() : nullptr); + EXPECT_EQ(params.getSpecBlTreeFirstSparseMaskOffsetKv(), + useSpecDecoding ? source.getSpecBlTreeFirstSparseMaskOffsetKv() : nullptr); + EXPECT_TRUE(source.spec_decoding_generation_lengths.has_value()); +} + +INSTANTIATE_TEST_SUITE_P(PhaseAndRuntimeFlags, AttentionOpSpecDecodingTest, + ::testing::Combine(::testing::Bool(), ::testing::Bool(), ::testing::Bool())); diff --git a/scripts/generate_fmha_params.py b/scripts/generate_fmha_params.py new file mode 100644 index 000000000000..a00ea935ace7 --- /dev/null +++ b/scripts/generate_fmha_params.py @@ -0,0 +1,439 @@ +#!/usr/bin/env python3 +# 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. +"""Generate native FmhaParams declarations from the Python schema. + +The schema is parsed, never imported: it lives inside the tensorrt_llm package, +whose import pulls in the very bindings this generator runs before. Parsing also +means a field's annotation is read exactly as written, which is where optionality +and the non-dtype types come from. + +Each field's C++ type comes from its annotation, and ``dtype`` supplies the +element type the annotation cannot carry -- a tensor's dtype, or a set's. +Supported named types, scalars, and tensors, including Optional forms, need no +cpp_metadata: ordinary defaults and default_factory are independent of codegen. Unmarked +tensors have no generated getter. Native schema classes are registered explicitly; +their inclusion does not depend on any field carrying cpp_metadata. +""" + +from __future__ import annotations + +import argparse +import ast +import difflib +import re +import sys +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parent.parent +ROOT_CLASS_NAME = "FmhaParams" +METADATA_FACTORY = "cpp_metadata" +# Structs declared in attentionOp.h, independent of their fields' dtype metadata. +NATIVE_STRUCT_NAMES = ( + "StaticAttentionConfig", + "AttentionForwardArgs", + "SparseBackendForwardArgs", + "SparseRuntimeParams", +) + +# The schema classes are spread over the modules that own them, so every source +# is parsed and a nested member is resolved by the class its annotation names. +DEFAULT_MODULE_PATHS = ( + REPO_ROOT / "tensorrt_llm" / "_torch" / "attention" / "backends" / "fmha" / "interface.py", + REPO_ROOT / "tensorrt_llm" / "_torch" / "attention" / "backends" / "interface.py", + REPO_ROOT / "tensorrt_llm" / "_torch" / "attention" / "backends" / "sparse" / "params.py", +) + +_TORCH_DTYPE_CPP = { + "torch.bool": "bool", + "torch.uint8": "std::uint8_t", + "torch.uint32": "std::uint32_t", + "torch.int32": "std::int32_t", + "torch.int64": "std::int64_t", + "torch.float32": "float", + "torch.float64": "double", +} + +# Python spells one integer and one floating type, so every scalar widens to the +# larger native type rather than carrying a redundant dtype. +_ANNOTATION_CPP = { + "bool": "bool", + "int": "std::int64_t", + "float": "double", + "AttentionMaskType": "tensorrt_llm::kernels::AttentionMaskType", + "PositionEmbeddingType": "tensorrt_llm::kernels::PositionEmbeddingType", + "RotaryScalingType": "tensorrt_llm::kernels::RotaryScalingType", + "QuantMode": "tensorrt_llm::common::QuantMode", + "DataType": "tensorrt_llm::DataType", + "BlockSparseParams": "tensorrt_llm::kernels::BlockSparseParams", + "MlaMetaParams": "tensorrt_llm::kernels::MlaMetaParams", + # A Python IntEnum with no native counterpart crosses as its integer value. + "AttentionInputType": "std::int64_t", +} + + +@dataclass(frozen=True) +class Field: + name: str + annotation: str + dtype: str | None # a torch dtype, or None (no generated accessor) + + +@dataclass(frozen=True) +class Struct: + name: str + fields: tuple[Field, ...] + + +@dataclass(frozen=True) +class Schema: + source: str + structs: dict[str, Struct] + + def nested(self, field: Field) -> Struct | None: + """Return the struct a field's annotation names, if it is one.""" + return self.structs.get(_optional_inner(field.annotation) or field.annotation) + + +def _optional_inner(annotation: str) -> str | None: + """Return X for Optional[X], else None.""" + if annotation.startswith("Optional[") and annotation.endswith("]"): + return annotation[len("Optional[") : -1] + return None + + +def _set_element(annotation: str) -> str | None: + """Return X for Set[X], else None.""" + if annotation.startswith("Set[") and annotation.endswith("]"): + return annotation[len("Set[") : -1] + return None + + +def _has_cpp_metadata(statement: ast.stmt) -> bool: + return ( + isinstance(statement, ast.AnnAssign) + and isinstance(statement.value, ast.Call) + and getattr(statement.value.func, "id", None) == METADATA_FACTORY + ) + + +def _read_struct(class_def: ast.ClassDef, schema_names: set[str]) -> Struct | None: + """Collect native fields, including references to known nested schemas.""" + fields = [] + implicit_types = schema_names | _ANNOTATION_CPP.keys() | {"torch.Tensor"} + for statement in class_def.body: + if not isinstance(statement, ast.AnnAssign) or not isinstance(statement.target, ast.Name): + continue + annotation = ast.unparse(statement.annotation) + explicit = _has_cpp_metadata(statement) + if not explicit and (_optional_inner(annotation) or annotation) not in implicit_types: + continue + dtype = None + if explicit: + for keyword in statement.value.keywords: + if keyword.arg not in {"dtype", "default"}: + raise ValueError( + f"{statement.target.id}: unsupported {METADATA_FACTORY} argument {keyword.arg}" + ) + if keyword.arg == "dtype": + dtype = ( + keyword.value.value + if isinstance(keyword.value, ast.Constant) + else ast.unparse(keyword.value) + ) + if dtype is None: + raise ValueError( + f"{statement.target.id}: {METADATA_FACTORY} needs a dtype; " + "use a plain default when no generated getter is needed" + ) + fields.append( + Field( + name=statement.target.id, + annotation=annotation, + dtype=dtype, + ) + ) + return Struct(class_def.name, tuple(fields)) if fields else None + + +def load_schemas( + source_paths: Sequence[Path], + root_class: str, + struct_names: Sequence[str] = NATIVE_STRUCT_NAMES, +) -> Schema: + """Read every schema class out of the given sources without importing them.""" + structs: dict[str, Struct] = {} + classes: dict[str, ast.ClassDef] = {} + selected = {root_class, *struct_names} + for path in source_paths: + for node in ast.parse(path.read_text()).body: + if not isinstance(node, ast.ClassDef) or node.name not in selected: + continue + classes[node.name] = node + # Discover schema names first so plain nested fields work across modules, + # regardless of which module or class is declared first. + for node in classes.values(): + struct = _read_struct(node, selected) + if struct is not None: + structs[struct.name] = struct + if root_class not in structs: + raise ValueError(f"no class named {root_class} declares native fields") + return Schema(source=", ".join(str(p) for p in source_paths), structs=structs) + + +def _element_cpp(field: Field) -> str: + if field.dtype is None: + raise ValueError(f"{field.name}: this type needs an explicit dtype") + try: + return _TORCH_DTYPE_CPP[field.dtype] + except KeyError: + raise ValueError(f"{field.name}: unsupported dtype {field.dtype}") from None + + +def render_field_type(schema: Schema, field: Field, annotation: str | None = None) -> str: + """Render the C++ type of one field from its annotation.""" + annotation = field.annotation if annotation is None else annotation + inner = _optional_inner(annotation) + if inner is not None: + # A nested struct is held by value; C++ mirrors the Python layout, and an + # absent one is the value-initialized struct rather than an empty optional. + if annotation in schema.structs or inner in schema.structs: + return render_field_type(schema, field, inner) + return f"std::optional<{render_field_type(schema, field, inner)}>" + if annotation in schema.structs: + return annotation + if annotation == "torch.Tensor": + return "torch::Tensor" + if _set_element(annotation) is not None: + return f"std::set<{_element_cpp(field)}>" + try: + return _ANNOTATION_CPP[annotation] + except KeyError: + raise ValueError(f"{field.name}: unsupported annotation {annotation}") from None + + +def _header_lines(schema: Schema, struct: Struct) -> list[str]: + return [ + "// Generated file; do not edit.", + "// Generated by scripts/generate_fmha_params.py.", + f"// Source: {struct.name} in {schema.source}.", + "", + ] + + +def render_fields(schema: Schema, struct: Struct) -> str: + """Render one struct's member declarations. + + Python owns semantic initialization; the members are only value-initialized + here so none of them start out indeterminate. + """ + lines = _header_lines(schema, struct) + [ + "#if !defined(TRTLLM_FMHA_PARAM_FIELD)", + '# error "Define TRTLLM_FMHA_PARAM_FIELD before including this file"', + "#endif", + "", + ] + for field in struct.fields: + lines.append(f"TRTLLM_FMHA_PARAM_FIELD({field.name}, {render_field_type(schema, field)})") + return "\n".join(lines) + "\n" + + +def _accessor_name(field_name: str) -> str: + return "get" + "".join(part.capitalize() for part in field_name.split("_")) + + +def _is_tensor(field: Field) -> bool: + return "torch.Tensor" in (field.annotation, _optional_inner(field.annotation)) + + +def _render_accessor(field: Field) -> list[str]: + """Render one tensor getter: the pointer the kernels expect, or nullptr if absent.""" + pointer = _element_cpp(field) + read = "static_cast<" + pointer + "*>({0}.data_ptr())" + + name = field.name + if _optional_inner(field.annotation) is not None: + body = f"return {name}.has_value() ? {read.format(name + '.value()')} : nullptr;" + else: + body = f"return {read.format(name)};" + + return [f"{pointer}* {_accessor_name(name)}() const", "{", f" {body}", "}"] + + +def render_accessors(schema: Schema, struct: Struct) -> str: + """Render the getters that are a plain typed view of a tensor field. + + Unmarked tensors use handwritten getters for runtime-dispatched pointer + types and semantic views that need offsets or a different buffer layout. + """ + lines = _header_lines(schema, struct) + for field in struct.fields: + if not _is_tensor(field) or field.dtype is None: + continue + lines.extend(_render_accessor(field)) + lines.append("") + return "\n".join(lines).rstrip("\n") + "\n" + + +def _accessor_names(schema: Schema, struct: Struct) -> set[str]: + """Names a struct answers itself, generated or forwarded.""" + names = {_accessor_name(f.name) for f in struct.fields if _is_tensor(f) and f.dtype is not None} + for field in struct.fields: + nested = schema.nested(field) + if nested is not None: + names |= _accessor_names(schema, nested) + return names + + +def render_forwarding(schema: Schema, struct: Struct) -> str: + """Render getters that reach through a nested member. + + Keeps call sites reading `p.getX()` even after X moved into a nested struct. A + name the outer struct defines itself is left alone: `output` and + `attention_mask` mean different things at the two levels. + """ + lines = _header_lines(schema, struct) + own = {_accessor_name(f.name) for f in struct.fields if _is_tensor(f) and f.dtype is not None} + seen = set(own) + for field in struct.fields: + nested = schema.nested(field) + if nested is None: + continue + for inner, prefix in _forwardable(schema, nested, field.name): + name = _accessor_name(inner.name) + if name in seen: + continue + seen.add(name) + pointer = _element_cpp(inner) + lines += [ + f"{pointer}* {name}() const", + "{", + f" return {prefix}.{name}();", + "}", + "", + ] + return "\n".join(lines).rstrip("\n") + "\n" + + +def _forwardable(schema: Schema, struct: Struct, prefix: str): + """Yield (field, access path) for every getter reachable through a nested member.""" + for field in struct.fields: + nested = schema.nested(field) + if nested is not None: + yield from _forwardable(schema, nested, f"{prefix}.{field.name}") + elif _is_tensor(field) and field.dtype is not None: + yield field, prefix + + +def _filename(struct_name: str, kind: str) -> str: + snake = re.sub(r"(? bool: + old_content = path.read_text() if path.exists() else None + if old_content == content: + return False + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(content) + return True + + +def _check_file(path: Path, content: str) -> bool: + try: + old_content = path.read_text() + except FileNotFoundError: + print(f"{path}: missing generated file", file=sys.stderr) + return False + if old_content == content: + return True + sys.stderr.writelines( + difflib.unified_diff( + old_content.splitlines(keepends=True), + content.splitlines(keepends=True), + fromfile=str(path), + tofile=f"{path} (expected)", + ) + ) + return False + + +def generate(schema: Schema, out_dir: Path, check: bool, root_class: str) -> int: + status = 0 + outputs = [ + (_filename(struct.name, kind), render(schema, struct)) + for struct in schema.structs.values() + for kind, render in BACKENDS + ] + outputs.append( + (_filename(root_class, "forwarding"), render_forwarding(schema, schema.structs[root_class])) + ) + for filename, content in outputs: + path = out_dir / filename + if check: + status |= 0 if _check_file(path, content) else 1 + continue + action = "updated" if _write_if_changed(path, content) else "unchanged" + print(f"{action}: {path}") + return status + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--module-path", + type=Path, + action="append", + dest="module_paths", + help="Python schema source to parse; repeatable (default: the schema modules).", + ) + parser.add_argument( + "--class-name", + default=ROOT_CLASS_NAME, + help="Outermost schema class (default: %(default)s).", + ) + parser.add_argument( + "--out-dir", + type=Path, + required=True, + help="Directory that receives the generated includes.", + ) + parser.add_argument( + "--check", + action="store_true", + help="Compare generated output against --out-dir without writing.", + ) + return parser.parse_args() + + +def main() -> int: + args = parse_args() + try: + paths = args.module_paths or list(DEFAULT_MODULE_PATHS) + schema = load_schemas(paths, args.class_name) + return generate(schema, args.out_dir, args.check, args.class_name) + except (OSError, SyntaxError, ValueError) as error: + print(f"FmhaParams generation failed: {error}", file=sys.stderr) + return 2 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md index 94f95df61644..146db2c96345 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md @@ -5,10 +5,13 @@ receipts: # thop_attention -**Wraps** `tensorrt_llm.bindings.internal.thop.attention` (one call). +**Wraps** `FallbackFmha.attention`, which builds native parameters and calls +`AttentionOp.run_context`, `run_generation`, or `run_mla_generation` for +the active phases. Native runners are cached by static configuration, +device, and thread; batch tensors and workspace remain caller-owned. -This is a pybind binding, not a `torch.ops` op; its inclusion in the catalog -is an approved policy exception to the `torch.ops.trtllm.*` entry shape. +This entry uses native bindings through the Python FMHA adapter; its inclusion +in the catalog is an approved policy exception to the `torch.ops.trtllm.*` entry shape. ## Semantics @@ -496,7 +499,7 @@ page-32), and the no-append flavor's chunked partial-pass pattern with `softmax_stats_tensor` — the sweep takes its one-shot cached-KV pattern. Every observable came back **bitwise identical** at both head counts: the `output` rows, the whole paged latent pool, the in-place-roped `q`/`k` of -the fresh-prefill flavor, and the byte count the op resized `workspace_` +the fresh-prefill flavor, and the byte count the op resized `workspace` to. The repeated `1536` is the run-to-run determinism control that makes those comparisons mean something. Kernel selection does not see it either: a process running only the two sweeps compiled exactly one decode kernel @@ -505,8 +508,7 @@ values and one `...VarSeqQ8...` serving the four `H = 8` ones — where a real selection axis (the head count) compiles once per value; see *Notes*. `None` is **not** an accepted value on the MLA path even though the -signature admits it: the C++ unwraps the optional unconditionally and the -call raises `RuntimeError: bad optional access`. Pass the checkpoint's rank, +signature admits it: the catalog wrapper raises `ValueError`. Pass the checkpoint's rank, or `0` when it has none. **Context phase — fresh prefill** (`attention_input_type=1`, separate @@ -828,7 +830,7 @@ flavor (measured-wrong, see above, rather than untested). ```python def thop_attention( - q, k, v, output, output_sf, workspace_, + q, k, v, output, output_sf, workspace, sequence_length, host_past_key_value_lengths, host_total_kv_lens, context_lengths, host_context_lengths, host_request_types, max_context_q_len_override, @@ -850,7 +852,7 @@ def thop_attention( helix_position_offsets, helix_is_inactive_rank, attention_chunk_size, softmax_stats_tensor, is_spec_decoding_enabled, use_spec_decoding, is_spec_dec_tree, spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, + spec_decoding_position_offsets, spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset, spec_decoding_bl_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv, sparse_kv_indices, sparse_kv_offsets, sparse_attn_indices, @@ -987,7 +989,7 @@ here promises the next version reads them the same way. | Argument | Shape / value | Dtype | Device | |---|---|---|---| | `output_sf` | `None` (NVFP4-output path not certified) | — | — | -| `workspace_` | persistent scratch tensor, any length (0 ok); the op grows it **in place** via `resize_()` when too small — 33 MB to 109 MB across the certified shapes. The MLA context requirement scales with tokens x heads and still sets the peak: 113 924 096 B for an `H = 128` context call over 545 tokens, against 52 579 072 B for the same call over 129 tokens and 100 024 320 B for the standard-configuration head-geometry sweep. The MLA **generation** call is well below that and, measured, does not grow with `predicted_tokens_per_seq`: 39 100 416 B at `H = 128`, page 32 for every `(G, P, L)` tried — `G` 2 and 4, `P` 1 and 4, `L` 50 and 500 — so a `P`-times-taller query block moves nothing here. Pass a plain resizable tensor, not a view; reuse it across calls to avoid re-allocation | int8 | CUDA | +| `workspace` | persistent scratch tensor, any length (0 ok); the op grows it **in place** via `resize_()` when too small — 33 MB to 109 MB across the certified shapes. The MLA context requirement scales with tokens x heads and still sets the peak: 113 924 096 B for an `H = 128` context call over 545 tokens, against 52 579 072 B for the same call over 129 tokens and 100 024 320 B for the standard-configuration head-geometry sweep. The MLA **generation** call is well below that and, measured, does not grow with `predicted_tokens_per_seq`: 39 100 416 B at `H = 128`, page 32 for every `(G, P, L)` tried — `G` 2 and 4, `P` 1 and 4, `L` 50 and 500 — so a `P`-times-taller query block moves nothing here. Pass a plain resizable tensor, not a view; reuse it across calls to avoid re-allocation | int8 | CUDA | ### Batch state (the "prepared metadata" of this op — all int32) @@ -1399,7 +1401,7 @@ plays for the registered-layer entry points.) unsatisfiable on a mixed one — see *Semantics*. - Every MLA call needs an int `q_lora_rank` — `0` for a checkpoint without a q-LoRA — even though the value is never read: `None` raises - `RuntimeError: bad optional access` (see *MLA q-LoRA rank*). + `ValueError` (see *MLA q-LoRA rank*). - MLA generation additionally requires, before the call: every generation sequence's full latent history `[0, L_g)` resident in the pool — the context call's append covers the prefill rows; this step's `P` rows at @@ -1535,7 +1537,7 @@ plays for the registered-layer entry points.) `HVPerCta` segment at all. So a target sweeping `max_draft_len` should expect one decode JIT compile per `P` it runs, not one for the model. - First call logs "Attention workspace size is not enough" and resizes - `workspace_` in place — expected when starting from an empty tensor. The + `workspace` in place — expected when starting from an empty tensor. The size tracks the call's tokens x heads: ~33 MB standard, ~39 MB MLA and ~35 MB no-append MLA context at the 8/16/32-head shapes, but ~50 MB for an `H = 128` MLA context call over 129 tokens and ~109 MB over 545, which is diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py index 17815fdf2c96..3874a9b309b8 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py @@ -2,18 +2,17 @@ # SPDX-License-Identifier: Apache-2.0 """Fused attention core with fully explicit state: paged KV-cache append + masked FMHA. -Wraps the pybind binding ``tensorrt_llm.bindings.internal.thop.attention`` -(approved policy exception to the ``torch.ops.trtllm.*`` entry shape): the -same C++ attention op behind the TRTLLM backend, but every piece of batch -state and layer config arrives as an explicit argument — no registered -layers, no thread-local metadata. +Wraps ``FallbackFmha.attention``, which dispatches to the native +``AttentionOp`` context and generation methods. Every piece of batch state +and layer config arrives as an explicit argument; native runners are cached +by static configuration, device, and thread. """ from typing import Optional import torch -from tensorrt_llm.bindings.internal import thop +from tensorrt_llm._torch.attention.backends.fmha.fallback import FallbackFmha def thop_attention( @@ -22,7 +21,7 @@ def thop_attention( v: Optional[torch.Tensor], output: torch.Tensor, output_sf: Optional[torch.Tensor], - workspace_: Optional[torch.Tensor], + workspace: Optional[torch.Tensor], sequence_length: torch.Tensor, host_past_key_value_lengths: torch.Tensor, host_total_kv_lens: torch.Tensor, @@ -88,7 +87,7 @@ def thop_attention( use_spec_decoding: bool, is_spec_dec_tree: bool, spec_decoding_generation_lengths: Optional[torch.Tensor], - spec_decoding_position_offsets_for_cpp: Optional[torch.Tensor], + spec_decoding_position_offsets: Optional[torch.Tensor], spec_decoding_packed_mask: Optional[torch.Tensor], spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor], spec_decoding_bl_tree_mask: Optional[torch.Tensor], @@ -140,6 +139,8 @@ def thop_attention( append + masked FMHA (standard), or the MLA context/generation phase selected by attention_input_type; writes rows [:num_tokens] of output. Returns None.""" + if is_mla_enable and q_lora_rank is None: + raise ValueError("q_lora_rank must be an int for MLA; use 0 when there is no q-LoRA.") # num_heads must be an integer multiple of num_kv_heads. A context-only # call with a non-multiple returns without raising on sm_100 (observed at # 6q/4kv d128: only the first (num_heads // num_kv_heads) * num_kv_heads @@ -161,13 +162,13 @@ def thop_attention( f"{tuple(attention_sinks.shape)} contiguous=" f"{attention_sinks.is_contiguous()}" ) - thop.attention( + FallbackFmha.attention( q=q, k=k, v=v, output=output, output_sf=output_sf, - workspace_=workspace_, + workspace=workspace, sequence_length=sequence_length, host_past_key_value_lengths=host_past_key_value_lengths, host_total_kv_lens=host_total_kv_lens, @@ -233,7 +234,7 @@ def thop_attention( use_spec_decoding=use_spec_decoding, is_spec_dec_tree=is_spec_dec_tree, spec_decoding_generation_lengths=spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp=spec_decoding_position_offsets_for_cpp, + spec_decoding_position_offsets=spec_decoding_position_offsets, spec_decoding_packed_mask=spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset=spec_decoding_bl_tree_mask_offset, spec_decoding_bl_tree_mask=spec_decoding_bl_tree_mask, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 7382f31d99a1..bafe7d653286 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -169,8 +169,8 @@ entries: summary: "MLA context-prefill preprocessing: in-place GPT-J RoPE of q_pe (in q) and k_pe (in latent_cache) at each new token's absolute position + [compressed_kv | rope(k_pe)] paged-latent-cache append; an fp8-e4m3 pool quantizes only the appended row, by kv_scale_orig_quant" - path: attention/thop_attention.py - impl: tensorrt_llm.bindings.internal.thop.attention - summary: "Full attention core with fully explicit state (pybind binding, approved exception): paged KV-cache append + causal/padding-masked GQA FMHA over a caller-owned pool (bf16, or fp8-e4m3 with per-tensor kv scales) and explicit length/offset tensors, written into a caller buffer, on either context execution path (use_paged_context_fmha selects packed-QKV context FMHA, or paged-KV context FMHA so a context call may run over a cached prefix — KV-cache reuse and chunked prefill), with optional per-query-head attention sinks (one extra softmax-denominator logit, dropped from the output) and an optional per-call sliding window (attention_window_size keys ending at the query's absolute position; a pure mask — the append stays at absolute positions, so cyclic pool reuse is the caller's page mapping); MLA mode runs context prefill (in-kernel RoPE + latent append), no-append context over explicit K/V (latent_cache=None: cached-KV prefixes and chunked partial passes with softmax-stats output), and generation latent-MQA decode as separate calls over a paged latent pool (bf16 or fp8-e4m3)" + impl: tensorrt_llm._torch.attention.backends.fmha.fallback.FallbackFmha.attention + summary: "Full attention core with fully explicit state (phased native FMHA adapter, approved exception): paged KV-cache append + causal/padding-masked GQA FMHA over a caller-owned pool (bf16, or fp8-e4m3 with per-tensor kv scales) and explicit length/offset tensors, written into a caller buffer, on either context execution path (use_paged_context_fmha selects packed-QKV context FMHA, or paged-KV context FMHA so a context call may run over a cached prefix — KV-cache reuse and chunked prefill), with optional per-query-head attention sinks (one extra softmax-denominator logit, dropped from the output) and an optional per-call sliding window (attention_window_size keys ending at the query's absolute position; a pure mask — the append stays at absolute positions, so cyclic pool reuse is the caller's page mapping); MLA mode runs context prefill (in-kernel RoPE + latent append), no-append context over explicit K/V (latent_cache=None: cached-KV prefixes and chunked partial passes with softmax-stats output), and generation latent-MQA decode as separate calls over a paged latent pool (bf16 or fp8-e4m3)" # ─── moe ─────────────────────────────────────────────────────── - path: moe/noaux_tc_op.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py index 65e903151a7a..9becdd1747bb 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py @@ -219,7 +219,7 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: kv_cache_block_offsets=md.kv_cache_block_offsets, host_kv_cache_pool_pointers=md.host_kv_cache_pool_pointers, host_kv_cache_pool_mapping=md.host_kv_cache_pool_mapping, - workspace_=md.effective_workspace, + workspace=md.effective_workspace, tokens_per_block=md.tokens_per_block, max_num_requests=md.max_num_requests, max_context_length=md.max_context_length, @@ -297,7 +297,7 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: use_spec_decoding=False, is_spec_dec_tree=False, spec_decoding_generation_lengths=None, - spec_decoding_position_offsets_for_cpp=None, + spec_decoding_position_offsets=None, spec_decoding_packed_mask=None, spec_decoding_bl_tree_mask_offset=None, spec_decoding_bl_tree_mask=None, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py index 6adc9e2b0798..4c315a489b60 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py @@ -64,6 +64,7 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.embedding import embedding from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.empty import empty from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.reshape import reshape +from tensorrt_llm._torch.attention.backends.fmha.phased import get_spec_decoding_position_offsets from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.model_config import ModelConfig @@ -126,7 +127,8 @@ "use_spec_decoding", "is_spec_dec_tree", "spec_decoding_generation_lengths", - "spec_decoding_position_offsets_for_cpp", + "spec_decoding_position_offsets", + "spec_decoding_query_len", "spec_decoding_packed_mask", "spec_decoding_bl_tree_mask_offset", "spec_decoding_bl_tree_mask", @@ -161,7 +163,7 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: kv_cache_block_offsets=md.kv_cache_block_offsets, host_kv_cache_pool_pointers=md.host_kv_cache_pool_pointers, host_kv_cache_pool_mapping=md.host_kv_cache_pool_mapping, - workspace_=md.effective_workspace, + workspace=md.effective_workspace, tokens_per_block=md.tokens_per_block, max_num_requests=md.max_num_requests, max_context_length=md.max_context_length, @@ -179,7 +181,7 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: use_spec_decoding=md.use_spec_decoding, is_spec_dec_tree=md.is_spec_dec_tree, spec_decoding_generation_lengths=md.spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp=md.spec_decoding_position_offsets_for_cpp, + spec_decoding_position_offsets=get_spec_decoding_position_offsets(md), spec_decoding_packed_mask=md.spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset=md.spec_decoding_bl_tree_mask_offset, spec_decoding_bl_tree_mask=md.spec_decoding_bl_tree_mask, diff --git a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md index a44632837d7c..08f5eb173b59 100644 --- a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md +++ b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md @@ -378,6 +378,19 @@ starting with an empty selection cache. `TrtllmAttention` prepares the complete per-forward state, passes itself to the manager for selection, and then executes the selected library. +`FallbackFmha` owns its lazy native `AttentionOp` cache. The stateless +`attention()` compatibility adapter supports direct MHA and MLA calls and owns +a separate Python cache. Both caches key by `StaticAttentionConfig`, CUDA device, and host +thread, never by per-call tensors or sequence lengths. Streams are not part of +the key so CUDA graph capture can reuse runners initialized during warmup. +Quantization updates release the old manager's implementation-owned resources. + +`prepare_workspace(q, k, v, metadata, forward_args, workspace)` runs on the whole +batch before phase execution. `CombinedFmha` delegates preparation to both +implementations before either runs, so the shared workspace covers both phases. +`FallbackFmha` builds its sizing parameters internally; other libraries do not +need to accept `FmhaParams` for workspace preparation. + `TLLM_FMHA_LIBS` controls the ordered selection. Dense PrimTS is opt-in because it may add host overhead; use `TLLM_FMHA_LIBS=+prims_ts` to add it to the defaults or `TLLM_FMHA_LIBS=fallback` to force the fallback path. Generic @@ -387,10 +400,12 @@ default membership and follow canonical registry order, while an exact list preserves the user-specified order. Each FMHA library exposes `is_available()` for module/static environment checks and `is_supported()` for per-forward request checks. `AttentionForwardArgs.sparse_runtime_params` is the sole -per-call lowered sparse runtime carrier and defaults to an empty -`SparseRuntimeParams()`. The core forward overwrites that field with the carrier -that `prepare_sparse_runtime_params` builds from the caller's carrier plus the -hook results. The carrier holds both flat `AttentionOp` parameters and optional +per-call lowered sparse runtime carrier and defaults to `None` at the backend +boundary. Before selecting or calling FMHA, the core forward populates that +field through `prepare_sparse_runtime_params`, which creates a per-call carrier +and preserves caller-provided fields plus the hook results. Direct callers of +FMHA must supply a prepared carrier, including an empty `SparseRuntimeParams()` +for dense attention. The carrier holds both flat `AttentionOp` parameters and optional `BlockSparseForwardInputs` in its nested `block_sparse_inputs` field. `Fmha.is_supported()` rejects a request that carries routes for every library that does not declare `supports_block_sparse_inputs`, so no dense kernel can @@ -500,11 +515,10 @@ chunk alone: the cached prefix is dropped from attention and then overwritten by the chunk's write-back, which turns a missing kernel into a plausible wrong answer rather than an error. -`get_attention_op` in `thop/attentionOp.cpp` therefore refuses a non-MLA, +The `AttentionOp` constructor in `thop/attentionOp.cpp` therefore refuses a non-MLA, non-cross paged-context configuration whose initialization produced no context FMHA kernel. The check runs after `initialize()`, because only the initialized -op reflects the exact Q/KV/output precision, mask type and page size, and -outside `initialize()` itself, which is `noexcept`. +op reflects the exact Q/KV/output precision, mask type and page size. The refusal has three distinct causes, each with its own message, and the distinction matters when triaging: diff --git a/tensorrt_llm/_torch/attention/backends/cpp_schema.py b/tensorrt_llm/_torch/attention/backends/cpp_schema.py new file mode 100644 index 000000000000..0a07589a6bc4 --- /dev/null +++ b/tensorrt_llm/_torch/attention/backends/cpp_schema.py @@ -0,0 +1,59 @@ +# 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. +"""Field metadata shared by the Python classes that generate native structs. + +This lives on its own because every schema class needs it, and those classes sit +in modules that already import one another. +""" + +from __future__ import annotations + +import dataclasses +from dataclasses import dataclass + +# Namespaced so the native schema metadata cannot collide with other users of +# dataclasses.field(metadata=...). +CPP_METADATA_KEY = "fmha.cpp" + + +@dataclass(frozen=True) +class CppMetadata: + """Controls a field's representation in the native FmhaParams struct. + + The generator reads the schema's source, not this object, so the value here is + kept only for introspection. See cpp_metadata() for what ``dtype`` means. + """ + + dtype: object + + +def cpp_metadata(*, dtype: object, default: object = None) -> dataclasses.Field[object]: + """Customize a dataclass field's representation in the native struct. + + The annotation gives the shape -- tensor, optional, set, scalar, or a named + type -- and ``dtype`` gives its fixed element type as a ``torch`` dtype. + + The default is ``None`` unless explicitly supplied. Scalars annotated as + bool, int, float, or their Optional forms do not need this helper: their + native types are bool, std::int64_t, and double, respectively. + + A tensor without this helper is a native field with no generated getter, + for example when its pointer type is supplied by C++ dtype dispatch or its + view needs offsets. These getters are handwritten in attentionOp.h. + Registered native structs and supported named types are recognized by their + annotations and need no marker either. Defaults and default factories remain + Python-owned. + """ + return dataclasses.field(default=default, metadata={CPP_METADATA_KEY: CppMetadata(dtype)}) diff --git a/tensorrt_llm/_torch/attention/backends/flashinfer.py b/tensorrt_llm/_torch/attention/backends/flashinfer.py index b800b0cd453f..0e42ea7a59c5 100644 --- a/tensorrt_llm/_torch/attention/backends/flashinfer.py +++ b/tensorrt_llm/_torch/attention/backends/flashinfer.py @@ -2808,7 +2808,7 @@ def decode_forward(plan_params: PlanParams, out: torch.Tensor): kv_dtype=kv_cache.dtype, q_scaling=self.q_scaling, attention_window_size=None, - attention_mask_type=int(AttentionMaskType.causal), + attention_mask_type=AttentionMaskType.causal, attention_mask_data=None, flashinfer_backend=self.flashinfer_backend) decode_forward(decode_plan_params, output[num_ctx_tokens:, :]) @@ -2823,7 +2823,7 @@ def decode_forward(plan_params: PlanParams, out: torch.Tensor): and attention_mask_data is not None): logger.warning_once("Falling back to causal attention", key="trtllm_gen_unsupported_custom_mask") - effective_mask_type = int(AttentionMaskType.causal) + effective_mask_type = AttentionMaskType.causal effective_mask_data = None plan_params = metadata.plan( @@ -2859,14 +2859,14 @@ def forward(self, latent_cache = forward_args.latent_cache if forward_args.attention_mask == CustomAttentionMask.CUSTOM: assert attention_mask_data is not None, "attention_mask_data is required for custom attention mask." - attention_mask_type = int(AttentionMaskType.custom_mask) + attention_mask_type = AttentionMaskType.custom_mask attention_mask_data = attention_mask_data if attention_mask_data.ndim == 1 else attention_mask_data.flatten( ) elif forward_args.attention_mask == PredefinedAttentionMask.CAUSAL: - attention_mask_type = int(AttentionMaskType.causal) + attention_mask_type = AttentionMaskType.causal attention_mask_data = None elif forward_args.attention_mask == PredefinedAttentionMask.FULL: - attention_mask_type = int(AttentionMaskType.padding) + attention_mask_type = AttentionMaskType.padding attention_mask_data = None else: raise ValueError("Unexpected attention mask type") diff --git a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py index 5631950e9e06..f34ac88c817e 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py @@ -13,18 +13,27 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import TYPE_CHECKING, Optional +from dataclasses import fields +from threading import Thread, current_thread +from typing import TYPE_CHECKING, ClassVar, Mapping, Optional import torch from tensorrt_llm._torch.attention.backends.interface import ( AttentionForwardArgs, + AttentionInputType, CustomAttentionMask, + PredefinedAttentionMask, ) from tensorrt_llm._utils import get_sm_version from tensorrt_llm.bindings.internal import thop +from tensorrt_llm.functional import AttentionMaskType -from .interface import Fmha, FmhaPhase +from ..sparse.params import SparseRuntimeParams +from . import interface +from .interface import FmhaPhase, StaticAttentionConfig, build_op_params +from .phased import FmhaParams, PhasedFmha +from .utils import get_multi_ctas_kv_counter if TYPE_CHECKING: from tensorrt_llm._torch.attention.backends.trtllm import ( @@ -33,33 +42,112 @@ ) -# ``AttentionForwardArgs`` fields that this backend does not consume. -# Sync test (test_attention_op_sync.py) requires every other field to map to a -# kwarg name, a @property on the dataclass, or a field that some @property -# transitively reads; entries here are exempt. -_THOP_EXCLUDED_FIELDS: frozenset = frozenset( - { - "sparse_backend_args", # consumed by sparse prediction before the attention op - "block_sparse_inputs", # consumed by the selected block-sparse FMHA - "attention_mask_data", # custom-mask code path - "out_scale_sf", # promoted into ``out_scale`` in ``TrtllmAttention.forward`` for NVFP4 path - "skip_mla_rope_generation", # handled in ``TrtllmAttention.forward`` for the test-only MLA path - "timestep", # consumed by sparse prediction before FMHA dispatch +class _AttentionOpCache: + """Reuse initialized runners, without retaining any per-call tensors.""" + + def __init__(self) -> None: + self._ops: dict[tuple[Thread, torch.device, StaticAttentionConfig], thop.AttentionOp] = {} + + def get(self, config: StaticAttentionConfig, device: torch.device) -> "thop.AttentionOp": + # Native runners and cuBLAS state must not be shared by concurrent host threads. + # Do not key by stream: CUDA graph capture must reuse the op initialized during + # warmup on another stream. Execution passes the current stream and scratch anew. + key = (current_thread(), device, config) + op = self._ops.get(key) + if op is None: + with torch.cuda.device(device): + op = thop.AttentionOp(config.to_thop_config()) + self._ops[key] = op + return op + + def clear(self) -> None: + self._ops.clear() + + +_FORWARD_ARG_NAMES = frozenset(field.name for field in fields(AttentionForwardArgs)) +_SPARSE_ARG_NAMES = frozenset(field.name for field in fields(SparseRuntimeParams)) + + +def _legacy_forward_args( + arguments: Mapping[str, object], **overrides: object +) -> AttentionForwardArgs: + """Translate flat compatibility arguments into the shared nested carriers.""" + sparse = { + name: arguments[name] for name in _SPARSE_ARG_NAMES if arguments.get(name) is not None } -) + sparse.setdefault("sparse_attn_indices_block_size", 1) + sparse["threshold_scale_factor_prefill"] = ( + arguments.get("skip_softmax_threshold_scale_factor_prefill") or 0.0 + ) + sparse["threshold_scale_factor_decode"] = ( + arguments.get("skip_softmax_threshold_scale_factor_decode") or 0.0 + ) + values = { + name: arguments[name] for name in _FORWARD_ARG_NAMES if arguments.get(name) is not None + } + mask_type = arguments["mask_type"] + if mask_type == AttentionMaskType.causal: + values["attention_mask"] = PredefinedAttentionMask.CAUSAL + elif mask_type == AttentionMaskType.padding: + values["attention_mask"] = PredefinedAttentionMask.FULL + else: + raise ValueError(f"Unsupported legacy attention mask type: {mask_type}") + values["sparse_runtime_params"] = SparseRuntimeParams(**sparse) + values.update(overrides) + return AttentionForwardArgs(**values) + -# ``thop.attention`` kwargs hard-wired to a literal at the call site (no -# rich object owns them). Sync test enforces both the kwarg name and the -# literal value. -_THOP_LITERALS: dict = {} +def _set_context_workspace_shape( + params: FmhaParams, *, num_contexts: int, num_ctx_tokens: int +) -> None: + """Point the phase counters at the context extent for native workspace sizing. + Only the counts: the op reduces the host length arrays over `seq_offset` and + `batch_size` itself, so the extents no longer cross the boundary. + """ + active = num_contexts > 0 and num_ctx_tokens > 0 + params.seq_offset = 0 + params.batch_size = num_contexts if active else 0 + params.num_requests = num_contexts if active else 0 + params.num_tokens = num_ctx_tokens if active else 0 -class FallbackFmha(Fmha): - """Fallback FMHA implementation using the fused TRT-LLM thop attention op.""" +def _phase_query(params: FmhaParams) -> torch.Tensor: + """The phase's query tensor, whichever of the two mutually exclusive fields holds it.""" + query = params.qkv_input if params.qkv_input is not None else params.query_input + if query is None: + raise RuntimeError("FallbackFmha requires qkv_input or query_input.") + return query + + +class FallbackFmha(PhasedFmha): + """Fallback FMHA implementation over the phased TRT-LLM thop ops.""" + + REQUIRES_PAGED_KV = False + NEEDS_BLOCK_EXTENT = False + # The flat compatibility entry point is a classmethod with no layer instance, so its + # runners hang off the class instead. Kept across calls and captures like the layer's + # own cache, and cleared before CUDA teardown by the atexit hook below. + _compat_attention_ops: ClassVar[_AttentionOpCache] = _AttentionOpCache() supports_skip_correction = True supports_workspace_reclamation = True + def __init__(self, attn: "TrtllmAttention"): + super().__init__(attn) + self._multi_ctas_kv_counter: Optional[torch.Tensor] = None + # Construct lazily: the layer can be initialized on meta before CUDA is ready. + self._attention_ops = _AttentionOpCache() + + def attention_op(self, params: FmhaParams) -> "thop.AttentionOp": + config = StaticAttentionConfig.from_params( + params, skip_correction_threshold=self.attn.skip_correction_threshold + ) + return self._attention_ops.get(config, _phase_query(params).device) + + def release(self) -> None: + self._attention_ops.clear() + self._multi_ctas_kv_counter = None + @classmethod def _is_available(cls, attn: "TrtllmAttention") -> bool: sparse_algorithm = getattr(attn.sparse_params, "algorithm", None) @@ -105,145 +193,444 @@ def _is_supported( return False return True - def forward( + @classmethod + def attention( + cls, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + output: torch.Tensor, + output_sf: Optional[torch.Tensor], + workspace: Optional[torch.Tensor], + sequence_length: torch.Tensor, + host_past_key_value_lengths: torch.Tensor, + host_total_kv_lens: torch.Tensor, + context_lengths: torch.Tensor, + host_context_lengths: torch.Tensor, + host_request_types: Optional[torch.Tensor], + max_context_q_len_override: Optional[int], + kv_cache_block_offsets: Optional[torch.Tensor], + host_kv_cache_pool_pointers: Optional[torch.Tensor], + host_kv_cache_pool_mapping: Optional[torch.Tensor], + cache_indirection: Optional[torch.Tensor], + kv_scale_orig_quant: Optional[torch.Tensor], + kv_scale_quant_orig: Optional[torch.Tensor], + out_scale: Optional[torch.Tensor], + rotary_inv_freq: Optional[torch.Tensor], + rotary_cos_sin: Optional[torch.Tensor], + latent_cache: Optional[torch.Tensor], + q_pe: Optional[torch.Tensor], + block_ids_per_seq: Optional[torch.Tensor], + attention_sinks: Optional[torch.Tensor], + is_fused_qkv: bool, + update_kv_cache: bool, + predicted_tokens_per_seq: int, + local_layer_idx: int, + num_heads: int, + num_kv_heads: int, + head_size: int, + tokens_per_block: Optional[int], + max_num_requests: int, + max_context_length: int, + max_seq_len: int, + attention_window_size: int, + beam_width: int, + mask_type: int, + quant_mode: int, + q_scaling: float, + position_embedding_type: int, + rope_dim: int, + rope_base: float, + rope_scale_type: int, + rope_scale: float, + rope_short_m_scale: float, + rope_long_m_scale: float, + rope_max_positions: int, + rope_original_max_positions: int, + use_paged_context_fmha: bool, + attention_input_type: Optional[int], + is_mla_enable: bool, + chunked_prefill_buffer_batch_size: Optional[int], + q_lora_rank: Optional[int], + kv_lora_rank: Optional[int], + qk_nope_head_dim: Optional[int], + qk_rope_head_dim: Optional[int], + v_head_dim: Optional[int], + rope_append: Optional[bool], + mrope_rotary_cos_sin: Optional[torch.Tensor], + mrope_position_deltas: Optional[torch.Tensor], + helix_position_offsets: Optional[torch.Tensor], + helix_is_inactive_rank: Optional[torch.Tensor], + attention_chunk_size: Optional[int], + softmax_stats_tensor: Optional[torch.Tensor], + is_spec_decoding_enabled: bool, + use_spec_decoding: bool, + is_spec_dec_tree: bool, + spec_decoding_generation_lengths: Optional[torch.Tensor], + spec_decoding_position_offsets: Optional[torch.Tensor], + spec_decoding_packed_mask: Optional[torch.Tensor], + spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor], + spec_decoding_bl_tree_mask: Optional[torch.Tensor], + spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor], + sparse_kv_indices: Optional[torch.Tensor], + sparse_kv_offsets: Optional[torch.Tensor], + sparse_attn_indices: Optional[torch.Tensor], + sparse_attn_offsets: Optional[torch.Tensor], + sparse_attn_indices_block_size: int, + num_sparse_topk: Optional[int] = None, + sparse_attn_kv_lens: Optional[torch.Tensor] = None, + skip_softmax_threshold_scale_factor_prefill: Optional[float] = None, + skip_softmax_threshold_scale_factor_decode: Optional[float] = None, + skip_softmax_stat: Optional[torch.Tensor] = None, + cu_q_seqlens: Optional[torch.Tensor] = None, + cu_kv_seqlens: Optional[torch.Tensor] = None, + fmha_scheduler_counter: Optional[torch.Tensor] = None, + mla_bmm1_scale: Optional[torch.Tensor] = None, + mla_bmm2_scale: Optional[torch.Tensor] = None, + quant_q_buffer: Optional[torch.Tensor] = None, + flash_mla_tile_scheduler_metadata: Optional[torch.Tensor] = None, + flash_mla_num_splits: Optional[torch.Tensor] = None, + sage_attn_num_elts_per_blk_q: int = 0, + sage_attn_num_elts_per_blk_k: int = 0, + sage_attn_num_elts_per_blk_v: int = 0, + sage_attn_qk_int8: bool = False, + num_contexts: int = 0, + num_ctx_tokens: int = 0, + trtllm_gen_jit_warmup: bool = False, + aux_kv_cache_pool_ptr: Optional[int] = None, + is_cross: bool = False, + cross_kv: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, + spec_decoding_target_max_draft_tokens: Optional[int] = None, + quant_scale_qkv: Optional[torch.Tensor] = None, + dsv4_inv_rope_cos_sin_cache: Optional[torch.Tensor] = None, + enable_dsv4_epilogue_fusion: bool = False, + force_prepare_spec_dec_tree_mask: bool = False, + max_num_sequences: Optional[int] = None, + kv_norm_weight: Optional[torch.Tensor] = None, + kv_norm_eps: float = 1e-6, + skip_correction_threshold: float = 0.0, + ) -> None: + """Shared MHA/MLA replacement for the removed monolithic ``thop.attention``. + + Builds native FMHA parameters and dispatches to the phased + ``AttentionOp.run_context`` / ``run_generation`` / ``run_mla_generation`` + ops. Direct callers supply the flat attention arguments and workspace. + """ + arguments = locals().copy() + del host_request_types, update_kv_cache, skip_softmax_stat + + if workspace is None: + raise RuntimeError("FallbackFmha.attention requires workspace.") + if output is None: + raise RuntimeError("FallbackFmha.attention requires output.") + + num_tokens = q.size(0) + if attention_input_type is None: + attention_input_type = AttentionInputType.mixed + is_gen_only = attention_input_type == AttentionInputType.generation_only + is_ctx_only = attention_input_type == AttentionInputType.context_only + if is_gen_only: + # Q is decode-only, but host metadata / KV blocks can still include + # a prefill prefix. Preserve num_contexts as the sequence offset. + num_ctx_tokens = 0 + total_seqs = host_context_lengths.size(0) + num_generations = 0 if is_ctx_only else total_seqs - num_contexts + if ( + enable_dsv4_epilogue_fusion + and num_contexts > 0 + and num_generations > 0 + and not is_gen_only + ): + raise ValueError("DSv4 fused epilogue requires separate context and generation calls.") + num_gen_tokens = num_tokens - num_ctx_tokens + if num_gen_tokens < 0: + raise RuntimeError( + f"Invalid FMHA token counts: num_tokens={num_tokens}, " + f"num_ctx_tokens={num_ctx_tokens}." + ) + + # Generation requires an FMHA-owned multi-CTA scratch buffer in addition to + # MLA's caller-owned scheduler counter. + multi_ctas_kv_counter = None + if num_generations > 0: + multi_ctas_kv_counter = get_multi_ctas_kv_counter( + None, q.device, num_heads, max_num_sequences or max_num_requests + ) + + effective_beam_width = 1 if is_cross else beam_width + sizing_num_contexts = num_contexts if num_ctx_tokens > 0 and not is_gen_only else 0 + params = interface.FmhaParams._from_arguments( + arguments, + qkv_or_q=q, + layer_idx=local_layer_idx, + num_seqs=sizing_num_contexts, + num_requests=sizing_num_contexts, + num_tokens=num_ctx_tokens if sizing_num_contexts > 0 else 0, + max_num_sequences=max_num_sequences or max_num_requests, + beam_width=effective_beam_width, + multi_ctas_kv_counter=multi_ctas_kv_counter, + fwd=_legacy_forward_args( + arguments, + chunked_prefill_buffer_batch_size=chunked_prefill_buffer_batch_size or 1, + ), + rotary_embedding_base=rope_base, + rotary_embedding_scale_type=rope_scale_type, + rotary_embedding_scale=rope_scale, + rotary_embedding_short_mscale=rope_short_m_scale, + rotary_embedding_long_mscale=rope_long_m_scale, + rotary_embedding_max_positions=rope_max_positions, + rotary_embedding_original_max_positions=rope_original_max_positions, + num_sparse_topk=num_sparse_topk or 0, + cyclic_attention_window_size=attention_window_size, + max_attention_window_size=( + attention_window_size + if effective_beam_width == 1 or cache_indirection is None + else cache_indirection.size(2) + ), + ) + tp = params.to_op_params() + op = cls._compat_attention_ops.get( + StaticAttentionConfig.from_legacy_arguments(arguments), arguments["q"].device + ) + + max_blocks_per_sequence = ( + kv_cache_block_offsets.size(-1) if kv_cache_block_offsets is not None else 0 + ) + workspace_size = op.get_attention_workspace_size( + tp, + num_tokens, + attention_window_size, + num_gen_tokens, + max_blocks_per_sequence, + ) + if workspace.numel() < workspace_size: + workspace.resize_(workspace_size) + + if num_contexts > 0 and not is_gen_only: + # Context phase. The context-MLA path is handled inside run_context, so + # both MLA and non-MLA go through run_context. + if max_context_q_len_override is not None: + max_context_q_len = int(host_context_lengths[:num_contexts].max()) + max_past_kv_len = int(host_past_key_value_lengths[:num_contexts].max()) + override = int(max_context_q_len_override) + if override < max_context_q_len or override < max_past_kv_len: + raise ValueError( + f"max_context_q_len_override ({override}) must be >= the computed max " + f"context q length ({max_context_q_len}) and max past kv length " + f"({max_past_kv_len})." + ) + tp.qkv_or_q = q[:num_ctx_tokens] + if k is not None: + tp.k = k[:num_ctx_tokens] + if v is not None: + tp.v = v[:num_ctx_tokens] + tp.output = output if enable_dsv4_epilogue_fusion else output[:num_ctx_tokens] + tp.sequence_length = sequence_length[:num_contexts] + tp.context_lengths = context_lengths[:num_contexts] + tp.seq_offset = 0 + tp.num_seqs = num_contexts + tp.num_requests = num_contexts + tp.token_offset = 0 + tp.num_tokens = num_ctx_tokens + op.run_context(tp) + + if num_generations > 0 and not is_ctx_only: + # Native preparation mutates its carrier; each phase needs fresh parameters. + if num_contexts > 0 and not is_gen_only: + tp = params.to_op_params() + seq_offset = num_contexts + tp.qkv_or_q = q[num_ctx_tokens:] + if k is not None: + tp.k = k[num_ctx_tokens:] + if v is not None: + tp.v = v[num_ctx_tokens:] + tp.output = output if enable_dsv4_epilogue_fusion else output[num_ctx_tokens:] + tp.sequence_length = sequence_length[seq_offset:] + tp.context_lengths = context_lengths[seq_offset:] + tp.seq_offset = seq_offset + tp.num_seqs = num_generations + tp.num_requests = num_generations // effective_beam_width + # The tensors above are phase-local; token_offset only indexes the whole-batch + # FP4 scaling-factor output. + tp.token_offset = num_ctx_tokens + tp.num_tokens = num_gen_tokens + if is_mla_enable: + op.run_mla_generation(tp) + else: + op.run_generation(tp) + + def _to_op_params(self, params: FmhaParams) -> "thop.FmhaParams": + """Validate and lower one phase's Python parameters. + + The native struct is the only place these values are gathered: FmhaParams keeps + the phase slice, and the layer and batch state stay behind `attn` and `meta` + rather than being mirrored into a second carrier. + """ + fwd, attn, meta = params.fwd, params.attn, params.meta + if fwd is None: + raise RuntimeError("FallbackFmha requires forward args.") + if params.output is None: + raise RuntimeError("FallbackFmha requires output.") + if params.workspace is None: + raise RuntimeError("FallbackFmha requires workspace.") + query = _phase_query(params) + + tp = thop.FmhaParams() + # Apply phase-local tensors and extents after the shared batch state. + build_op_params(tp, meta, params) + + tp.qkv_or_q = query + tp.sequence_length = params.sequence_lengths + # Direct phase callers may omit the prompt-length view. + tp.context_lengths = ( + params.context_lengths + if params.context_lengths is not None + else meta.prompt_lens_cuda_runtime[params.seq_offset :] + ) + tp.num_seqs = params.batch_size + # Encoder KV is request-scoped; each decoder beam reads it independently. + tp.beam_width = meta.effective_beam_width + tp.num_requests = params.batch_size if meta.is_cross else params.num_requests + tp.host_past_key_value_lengths = meta.kv_lens_runtime + tp.host_context_lengths = meta.prompt_lens_cpu_runtime + tp.max_context_length = meta.max_context_length + tp.max_num_requests = meta.max_num_requests + tp.max_num_sequences = meta.max_num_sequences or meta.max_num_requests + tp.layer_idx = attn.layer_idx + tp.local_layer_idx = attn.get_local_layer_idx(meta) + + # Leave a native field alone when the source is None: it already holds the empty + # value, and its setter takes the field's own type, as build_op_params assumes. + for name, value in ( + ("k", params.key_input), + ("v", params.value_input), + ("host_kv_cache_pool_pointers", meta.host_kv_cache_pool_pointers), + ("host_kv_cache_pool_mapping", meta.host_kv_cache_pool_mapping), + ("cache_indirection", meta.cache_indirection), + ("spec_decoding_target_max_draft_tokens", meta.max_total_draft_tokens), + ("attention_chunk_size", attn.attention_chunk_size), + ("rotary_inv_freq", attn.rotary_inv_freq), + ("rotary_cos_sin", attn.rotary_cos_sin), + ("multi_ctas_kv_counter", self._multi_ctas_kv_counter), + ): + if value is not None: + setattr(tp, name, value) + + tp.has_fp8_kv_cache = bool(getattr(attn, "has_fp8_kv_cache", False)) + + rope_params = attn.rope_params + if rope_params is not None: + # `rotary_embedding_dim` is layer configuration and reaches the op through + # StaticAttentionConfig, so it is not resent here. + tp.rotary_embedding_base = rope_params.theta + tp.rotary_embedding_scale_type = rope_params.scale_type + tp.rotary_embedding_scale = rope_params.scale + tp.rotary_embedding_short_mscale = rope_params.short_m_scale + tp.rotary_embedding_long_mscale = rope_params.long_m_scale + tp.rotary_embedding_max_positions = rope_params.max_positions + tp.rotary_embedding_original_max_positions = rope_params.original_max_positions + return tp + + # Keep nanobind calls eager, including when CombinedFmha delegates individual phases. + @torch.compiler.disable + def prepare_workspace( self, q: torch.Tensor, k: Optional[torch.Tensor], v: Optional[torch.Tensor], metadata: "TrtllmAttentionMetadata", forward_args: AttentionForwardArgs, + workspace: torch.Tensor, ) -> None: + if metadata.num_generations > 0: + self._multi_ctas_kv_counter = get_multi_ctas_kv_counter( + self._multi_ctas_kv_counter, + q.device, + self.attn.num_heads, + metadata.max_num_sequences or metadata.max_num_requests, + ) + # The native sizing call wants a parameter carrier, but the phase split has not + # run yet: build one covering the whole batch and point the phase counters at the + # context extent, which is what bounds the context workspace. attn = self.attn - - # Every kwarg sources from ``attn`` / ``metadata`` / ``forward_args`` - # (with ``forward_args.sparse_runtime_params`` for sparse inputs), - # or a literal allowlisted in ``_THOP_LITERALS``. - # ``test_attention_op_sync.py`` enforces this statically. - thop.attention( - q=q, - k=k, - v=v, - output=forward_args.output, - output_sf=forward_args.output_sf, - workspace_=metadata.effective_workspace, - # --- Per-step batch state (TrtllmAttentionMetadata) --- - sequence_length=metadata.kv_lens_cuda_runtime, - host_past_key_value_lengths=metadata.kv_lens_runtime, - host_total_kv_lens=metadata.host_total_kv_lens, + output = forward_args.output + if output is None: + raise RuntimeError("FallbackFmha requires output.") + num_tokens = q.size(0) + is_gen_only = forward_args.attention_input_type == AttentionInputType.generation_only + num_gen_tokens = num_tokens if is_gen_only else num_tokens - metadata.num_ctx_tokens + attention_window_size = forward_args.attention_window_size + cache_indirection = metadata.cache_indirection + max_attention_window_size = ( + attention_window_size + if metadata.effective_beam_width == 1 + else ( + cache_indirection.size(2) + if cache_indirection is not None + else attention_window_size + ) + ) + is_fused_qkv = forward_args.is_fused_qkv + params = FmhaParams( + attn=attn, + meta=metadata, + fwd=forward_args, + workspace=workspace, + qkv_input=q if is_fused_qkv else None, + query_input=None if is_fused_qkv else q, + key_input=k, + value_input=v, + # Workspace sizing only needs the output dtype; preserve the caller's layout. + output=output, + sequence_lengths=metadata.kv_lens_cuda_runtime, context_lengths=metadata.prompt_lens_cuda_runtime, - host_context_lengths=metadata.prompt_lens_cpu_runtime, - host_request_types=metadata.host_request_types_runtime, - max_context_q_len_override=metadata.max_context_q_len_override, - kv_cache_block_offsets=metadata.kv_cache_block_offsets, - host_kv_cache_pool_pointers=metadata.host_kv_cache_pool_pointers, - host_kv_cache_pool_mapping=metadata.host_kv_cache_pool_mapping, - cache_indirection=metadata.cache_indirection, - block_ids_per_seq=metadata.block_ids_per_seq, - tokens_per_block=metadata.tokens_per_block, - max_num_requests=metadata.max_num_requests, - max_num_sequences=metadata.max_num_sequences, - beam_width=metadata.effective_beam_width, - use_paged_context_fmha=metadata.use_paged_context_fmha, - helix_position_offsets=metadata.helix_position_offsets, - helix_is_inactive_rank=metadata.helix_is_inactive_rank, - is_spec_decoding_enabled=metadata.is_spec_decoding_enabled, - use_spec_decoding=metadata.use_spec_decoding, - is_spec_dec_tree=metadata.is_spec_dec_tree, - spec_decoding_generation_lengths=metadata.spec_decoding_generation_lengths, - spec_decoding_position_offsets_for_cpp=metadata.spec_decoding_position_offsets_for_cpp, - spec_decoding_packed_mask=metadata.spec_decoding_packed_mask, - spec_decoding_bl_tree_mask_offset=metadata.spec_decoding_bl_tree_mask_offset, - spec_decoding_bl_tree_mask=metadata.spec_decoding_bl_tree_mask, - spec_decoding_target_max_draft_tokens=metadata.max_total_draft_tokens, - force_prepare_spec_dec_tree_mask=metadata.force_prepare_spec_dec_tree_mask, - spec_bl_tree_first_sparse_mask_offset_kv=metadata.spec_bl_tree_first_sparse_mask_offset_kv, - num_sparse_topk=metadata.num_sparse_topk, - flash_mla_tile_scheduler_metadata=metadata.flash_mla_tile_scheduler_metadata, - flash_mla_num_splits=metadata.flash_mla_num_splits, - num_contexts=metadata.num_contexts, - num_ctx_tokens=metadata.num_ctx_tokens, - max_context_length=metadata.max_context_length, - max_seq_len=metadata.max_seq_len, - trtllm_gen_jit_warmup=metadata.trtllm_gen_jit_warmup, - is_cross=metadata.is_cross, - # --- Per-call (AttentionForwardArgs) --- - out_scale=forward_args.out_scale, - kv_scale_orig_quant=forward_args.kv_scale_orig_quant, - kv_scale_quant_orig=forward_args.kv_scale_quant_orig, - latent_cache=forward_args.latent_cache, - q_pe=forward_args.q_pe, - attention_sinks=forward_args.attention_sinks, - mask_type=forward_args.mask_type, - attention_input_type=int(forward_args.attention_input_type), - attention_window_size=forward_args.attention_window_size, - chunked_prefill_buffer_batch_size=forward_args.chunked_prefill_buffer_batch_size, - mrope_rotary_cos_sin=forward_args.mrope_rotary_cos_sin, - mrope_position_deltas=forward_args.mrope_position_deltas, - softmax_stats_tensor=forward_args.softmax_stats_tensor, - cu_q_seqlens=forward_args.cu_q_seqlens, - cu_kv_seqlens=forward_args.cu_kv_seqlens, - fmha_scheduler_counter=forward_args.fmha_scheduler_counter, - mla_bmm1_scale=forward_args.mla_bmm1_scale, - mla_bmm2_scale=forward_args.mla_bmm2_scale, - quant_q_buffer=forward_args.quant_q_buffer, - quant_scale_qkv=forward_args.quant_scale_qkv, - dsv4_inv_rope_cos_sin_cache=forward_args.dsv4_inv_rope_cos_sin_cache, - enable_dsv4_epilogue_fusion=forward_args.enable_dsv4_epilogue_fusion, - kv_norm_weight=forward_args.kv_norm_weight, - kv_norm_eps=forward_args.kv_norm_eps, - sage_attn_num_elts_per_blk_q=forward_args.sage_attn_num_elts_per_blk_q, - sage_attn_num_elts_per_blk_k=forward_args.sage_attn_num_elts_per_blk_k, - sage_attn_num_elts_per_blk_v=forward_args.sage_attn_num_elts_per_blk_v, - sage_attn_qk_int8=forward_args.sage_attn_qk_int8, - is_fused_qkv=forward_args.is_fused_qkv, - update_kv_cache=forward_args.update_kv_cache, - cross_kv=forward_args.cross_kv, - relative_attention_bias=forward_args.relative_attention_bias, - relative_attention_max_distance=forward_args.relative_attention_max_distance, - # --- Module config (TrtllmAttention) --- - rotary_inv_freq=attn.rotary_inv_freq, - rotary_cos_sin=attn.rotary_cos_sin, - predicted_tokens_per_seq=attn.predicted_tokens_per_seq, - local_layer_idx=attn.local_layer_idx, - num_heads=attn.num_heads, - num_kv_heads=attn.num_kv_heads, - head_size=attn.head_dim, - quant_mode=attn.quant_mode, - q_scaling=attn.q_scaling, - position_embedding_type=attn.position_embedding_type, - rope_dim=attn.rope_dim, - rope_base=attn.rope_base, - rope_scale_type=attn.rope_scale_type, - rope_scale=attn.rope_scale, - rope_short_m_scale=attn.rope_short_m_scale, - rope_long_m_scale=attn.rope_long_m_scale, - rope_max_positions=attn.rope_max_positions, - rope_original_max_positions=attn.rope_original_max_positions, - is_mla_enable=attn.is_mla_enable, - q_lora_rank=attn.q_lora_rank, - kv_lora_rank=attn.kv_lora_rank, - qk_nope_head_dim=attn.qk_nope_head_dim, - qk_rope_head_dim=attn.qk_rope_head_dim, - v_head_dim=attn.v_head_dim, - rope_append=attn.rope_append, - attention_chunk_size=attn.attention_chunk_size, - skip_softmax_stat=attn.skip_softmax_stat, - skip_correction_threshold=attn.skip_correction_threshold, - uses_spcompress=attn.uses_spcompress, - # --- Sparse runtime parameters --- - sparse_kv_indices=forward_args.sparse_runtime_params.sparse_kv_indices, - sparse_kv_offsets=forward_args.sparse_runtime_params.sparse_kv_offsets, - sparse_attn_indices=forward_args.sparse_runtime_params.sparse_attn_indices, - sparse_attn_offsets=forward_args.sparse_runtime_params.sparse_attn_offsets, - sparse_attn_indices_block_size=( - forward_args.sparse_runtime_params.sparse_attn_indices_block_size - ), - sparse_attn_kv_lens=forward_args.sparse_runtime_params.sparse_attn_kv_lens, - aux_kv_cache_pool_ptr=forward_args.sparse_runtime_params.aux_kv_cache_pool_ptr, - skip_softmax_threshold_scale_factor_prefill=( - forward_args.sparse_runtime_params.threshold_scale_factor_prefill - ), - skip_softmax_threshold_scale_factor_decode=( - forward_args.sparse_runtime_params.threshold_scale_factor_decode + max_attention_window_size=max_attention_window_size, + cyclic_attention_window_size=attention_window_size, + tokens_per_block=( + metadata.tokens_per_block if metadata.tokens_per_block is not None else 64 ), + kv_factor=self.kv_factor, + is_cross=metadata.is_cross, ) + _set_context_workspace_shape( + params, + num_contexts=0 if is_gen_only else metadata.num_contexts, + num_ctx_tokens=0 if is_gen_only else metadata.num_ctx_tokens, + ) + tp = self._to_op_params(params) + + kv_cache_block_offsets = metadata.kv_cache_block_offsets + use_kv_cache = kv_cache_block_offsets is not None + max_blocks_per_sequence = kv_cache_block_offsets.size(-1) if use_kv_cache else 0 + max_attention_window_size = ( + params.cyclic_attention_window_size + if metadata.effective_beam_width == 1 + else params.max_attention_window_size + ) + workspace_size = self.attention_op(params).get_attention_workspace_size( + tp, + num_tokens, + max_attention_window_size, + num_gen_tokens, + max_blocks_per_sequence, + ) + if workspace.numel() < workspace_size: + workspace.resize_(workspace_size) + + @torch.compiler.disable + def run_context(self, params: FmhaParams) -> None: + self.attention_op(params).run_context(self._to_op_params(params)) + + @torch.compiler.disable + def run_mla_context(self, params: FmhaParams) -> None: + self.attention_op(params).run_context(self._to_op_params(params)) + + @torch.compiler.disable + def run_generation(self, params: FmhaParams) -> None: + self.attention_op(params).run_generation(self._to_op_params(params)) + + @torch.compiler.disable + def run_mla_generation(self, params: FmhaParams) -> None: + self.attention_op(params).run_mla_generation(self._to_op_params(params)) diff --git a/tensorrt_llm/_torch/attention/backends/fmha/flashinfer_trtllm_gen.py b/tensorrt_llm/_torch/attention/backends/fmha/flashinfer_trtllm_gen.py index 1f43ae3bb385..020e3a8a6027 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/flashinfer_trtllm_gen.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/flashinfer_trtllm_gen.py @@ -66,6 +66,7 @@ from .utils import ( get_attention_chunk_size, get_bmm1_scale, + get_multi_ctas_kv_counter_size, get_multi_processor_count_for_device, get_trtllm_gen_context_workspace_size, ) @@ -217,23 +218,6 @@ def _cached_build( _install_flashinfer_mla_decode_tuning_config_cache() -_MULTI_CTAS_KV_COUNTER_ALIGNMENT = 8 - - -def _get_multi_ctas_kv_counter_size( - num_heads: int, - max_num_sequences: int, - multi_processor_count: int, -) -> int: - num_counters = max(num_heads * max_num_sequences, multi_processor_count) - aligned_num_counters = ( - (num_counters + _MULTI_CTAS_KV_COUNTER_ALIGNMENT - 1) - // _MULTI_CTAS_KV_COUNTER_ALIGNMENT - * _MULTI_CTAS_KV_COUNTER_ALIGNMENT - ) - return aligned_num_counters * torch.int32.itemsize - - def _get_bmm1_scale_log2(bmm1_scale: torch.Tensor) -> torch.Tensor: if bmm1_scale.numel() < 2: raise RuntimeError("trtllm-gen bmm1_scale workspace must contain raw and log2 scales.") @@ -815,9 +799,10 @@ def prepare_workspace( # One counter per head per decoder sequence; beam search expands each # request into ``beam_width`` sequences. - required_counter_size = _get_multi_ctas_kv_counter_size( + max_num_sequences = metadata.max_num_sequences or metadata.max_num_requests + required_counter_size = get_multi_ctas_kv_counter_size( attn.num_heads, - metadata.max_num_sequences or metadata.max_num_requests, + max_num_sequences, self._multi_processor_count, ) counter_buffer = self._multi_ctas_kv_counter_buffer @@ -846,8 +831,8 @@ def prepare_workspace( raise RuntimeError(f"{type(self).__name__} requires output.") fp8_context_fmha = self._use_fp8_context_fmha(output, attention_input_type) - workspace_max_tokens = max(num_tokens, metadata.max_context_length) - workspace_max_gen_tokens = max(num_gen_tokens, metadata.max_num_requests) + workspace_max_tokens = max(num_tokens, metadata.max_context_length, max_num_sequences) + workspace_max_gen_tokens = max(num_gen_tokens, max_num_sequences) required_workspace_size = _get_workspace_size( dtype=q.dtype, num_tokens=workspace_max_tokens, @@ -855,7 +840,7 @@ def prepare_workspace( num_heads=attn.num_heads, num_kv_heads=attn.num_kv_heads, head_size=attn.head_dim, - max_num_requests=metadata.max_num_requests, + max_num_requests=max_num_sequences, rotary_embedding_dim=attn.rope_dim, fp8_context_fmha=fp8_context_fmha, ) diff --git a/tensorrt_llm/_torch/attention/backends/fmha/interface.py b/tensorrt_llm/_torch/attention/backends/fmha/interface.py index 2d1bd5c86742..54cef16de5bc 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/interface.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/interface.py @@ -13,13 +13,19 @@ # See the License for the specific language governing permissions and # limitations under the License. +from __future__ import annotations + +import dataclasses import weakref from abc import ABC, abstractmethod +from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, ClassVar, Optional, final +from functools import cache +from typing import TYPE_CHECKING, Any, ClassVar, Mapping, Optional, final import torch +from tensorrt_llm._torch.attention.backends.cpp_schema import cpp_metadata from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs from tensorrt_llm.logger import logger @@ -28,6 +34,403 @@ TrtllmAttention, TrtllmAttentionMetadata, ) + from tensorrt_llm.bindings import DataType + from tensorrt_llm.functional import AttentionMaskType, PositionEmbeddingType, RotaryScalingType + from tensorrt_llm.quantization.mode import QuantMode + + +_STATIC_CONFIG_DIRECT_ARGS = ( + "num_heads", + "num_kv_heads", + "head_size", + "tokens_per_block", + "position_embedding_type", + "mask_type", + "q_scaling", + "is_spec_decoding_enabled", + "is_mla_enable", + "q_lora_rank", + "kv_lora_rank", + "qk_nope_head_dim", + "qk_rope_head_dim", + "v_head_dim", + "predicted_tokens_per_seq", + "rope_append", + "sage_attn_num_elts_per_blk_q", + "sage_attn_num_elts_per_blk_k", + "sage_attn_num_elts_per_blk_v", + "sage_attn_qk_int8", + "skip_correction_threshold", +) + + +@dataclass(kw_only=True, slots=True, frozen=True) +class StaticAttentionConfig: + """The attention-layer configuration used to build native kernel runners.""" + + num_heads: int = 0 + num_kv_heads: int = 0 + head_size: int = 0 + tokens_per_block: int = 0 + type: DataType = None + is_fp8_out: bool = False + is_fp4_out: bool = False + use_kv_cache: bool = False + paged_context_fmha: bool = False + position_embedding_type: PositionEmbeddingType = 0 + mask_type: AttentionMaskType = 1 + q_scaling: float = 1.0 + rotary_embedding_dim: int = 0 + attn_logit_softcapping_scale: float = 0.0 + remove_padding: bool = True + cross_attention: bool = False + dense_context_fmha: bool = False + fuses_dsv4_inv_rope_fp8_quant: bool = False + use_sparse_attention: bool = False + use_tllm_gen_sparse_attention: bool = False + use_nvfp4_mla_kv_cache: bool = False + is_spec_decoding_enabled: bool = False + spec_decoding_target_max_gen_len: int = 0 + is_mla_enable: bool = False + q_lora_rank: int = 0 + kv_lora_rank: int = 0 + qk_nope_head_dim: int = 0 + qk_rope_head_dim: int = 0 + v_head_dim: int = 0 + predicted_tokens_per_seq: int = 1 + mla_layer_num: int = 0 + rope_append: bool = True + sage_attn_num_elts_per_blk_q: int = 0 + sage_attn_num_elts_per_blk_k: int = 0 + sage_attn_num_elts_per_blk_v: int = 0 + sage_attn_qk_int8: bool = False + quant_mode: QuantMode = 0 + skip_correction_threshold: float = 0.0 + + @classmethod + def from_params( + cls, + params: FmhaParams, + *, + skip_correction_threshold: float = 0.0, + ) -> StaticAttentionConfig: + """Capture the fixed runner-selection inputs from one attention call. + + Sourced from the layer and the batch metadata rather than from the phase + carrier: everything here is fixed for the layer, so reading it off `attn` and + `meta` keeps the phase carrier free of a second copy. + """ + from tensorrt_llm._utils import torch_dtype_to_binding + from tensorrt_llm.quantization.mode import QuantMode + + attn, meta, fwd = params.attn, params.meta, params.fwd + if fwd is None: + raise RuntimeError("StaticAttentionConfig requires forward args.") + if params.output is None: + raise RuntimeError("StaticAttentionConfig requires output.") + query = params.qkv_input if params.qkv_input is not None else params.query_input + if query is None: + raise RuntimeError("StaticAttentionConfig requires qkv_input or query_input.") + + quant_mode = QuantMode(attn.quant_mode) + sparse = fwd.sparse_runtime_params + has_sparse_attn_indices = ( + sparse.sparse_attn_indices is not None and sparse.sparse_attn_indices.numel() > 0 + ) + has_sparse_attention = ( + sparse.sparse_kv_indices is not None and sparse.sparse_kv_indices.numel() > 0 + ) or has_sparse_attn_indices + has_paged_sparse_attention = ( + has_sparse_attn_indices + and sparse.sparse_attn_offsets is not None + and sparse.sparse_attn_offsets.numel() > 0 + ) + use_tllm_gen_sparse_attention = has_sparse_attn_indices and not has_paged_sparse_attention + pool_mapping = meta.host_kv_cache_pool_mapping + use_kv_cache = ( + meta.kv_cache_block_offsets is not None + and meta.host_kv_cache_pool_pointers is not None + and pool_mapping is not None + ) + target_max_gen_len = 0 + if meta.max_total_draft_tokens is not None: + target_max_gen_len = meta.max_total_draft_tokens + 1 + + return cls( + num_heads=attn.num_heads, + num_kv_heads=attn.num_kv_heads, + head_size=attn.head_dim, + tokens_per_block=params.tokens_per_block, + type=torch_dtype_to_binding(query.dtype), + is_fp8_out=params.output.dtype == torch.float8_e4m3fn, + is_fp4_out=params.output.dtype == torch.uint8, + use_kv_cache=use_kv_cache, + paged_context_fmha=meta.use_paged_context_fmha, + position_embedding_type=attn.position_embedding_type, + mask_type=fwd.mask_type, + q_scaling=attn.q_scaling, + rotary_embedding_dim=attn.rope_params.dim if attn.rope_params is not None else 0, + cross_attention=params.is_cross, + fuses_dsv4_inv_rope_fp8_quant=fwd.enable_dsv4_epilogue_fusion, + use_sparse_attention=has_sparse_attention, + use_tllm_gen_sparse_attention=use_tllm_gen_sparse_attention, + use_nvfp4_mla_kv_cache=( + quant_mode.has_fp4_kv_cache() + and use_tllm_gen_sparse_attention + and sparse.sparse_attn_kv_lens is None + and sparse.aux_kv_cache_pool_ptr is not None + ), + is_spec_decoding_enabled=meta.is_spec_decoding_enabled, + spec_decoding_target_max_gen_len=target_max_gen_len, + is_mla_enable=attn.is_mla_enable, + q_lora_rank=attn.q_lora_rank or 0, + kv_lora_rank=attn.kv_lora_rank or 0, + qk_nope_head_dim=attn.qk_nope_head_dim or 0, + qk_rope_head_dim=attn.qk_rope_head_dim or 0, + v_head_dim=attn.v_head_dim or 0, + predicted_tokens_per_seq=attn.predicted_tokens_per_seq, + mla_layer_num=pool_mapping.size(0) if pool_mapping is not None else 0, + rope_append=attn.rope_append is not False, + sage_attn_num_elts_per_blk_q=fwd.sage_attn_num_elts_per_blk_q, + sage_attn_num_elts_per_blk_k=fwd.sage_attn_num_elts_per_blk_k, + sage_attn_num_elts_per_blk_v=fwd.sage_attn_num_elts_per_blk_v, + sage_attn_qk_int8=fwd.sage_attn_qk_int8, + quant_mode=quant_mode, + skip_correction_threshold=skip_correction_threshold, + ) + + @classmethod + def from_legacy_arguments(cls, arguments: Mapping[str, Any]) -> StaticAttentionConfig: + """Build the config for the flat compatibility entry point. + + That entry point has no layer object to read from, so the values come from its + own arguments; `from_params` covers every other caller. + """ + from tensorrt_llm._utils import torch_dtype_to_binding + from tensorrt_llm.quantization.mode import QuantMode + + def arg(name: str, default: Any = None) -> Any: + value = arguments.get(name, default) + return default if value is None else value + + quant_mode = QuantMode(arg("quant_mode", 0)) + attn_indices = arg("sparse_attn_indices") + has_attn_indices = attn_indices is not None and attn_indices.numel() > 0 + kv_indices = arg("sparse_kv_indices") + attn_offsets = arg("sparse_attn_offsets") + use_tllm_gen_sparse = has_attn_indices and not ( + attn_offsets is not None and attn_offsets.numel() > 0 + ) + pool_mapping = arg("host_kv_cache_pool_mapping") + output = arguments["output"] + direct = { + name: arguments[name] + for name in _STATIC_CONFIG_DIRECT_ARGS + if arguments.get(name) is not None + } + return cls( + **direct, + type=torch_dtype_to_binding(arguments["q"].dtype), + is_fp8_out=output.dtype == torch.float8_e4m3fn, + is_fp4_out=output.dtype == torch.uint8, + use_kv_cache=( + arg("kv_cache_block_offsets") is not None + and arg("host_kv_cache_pool_pointers") is not None + and pool_mapping is not None + ), + paged_context_fmha=bool(arg("use_paged_context_fmha", False)), + rotary_embedding_dim=int(arg("rope_dim", 0)), + cross_attention=bool(arg("is_cross", False)), + fuses_dsv4_inv_rope_fp8_quant=bool(arg("enable_dsv4_epilogue_fusion", False)), + use_sparse_attention=(kv_indices is not None and kv_indices.numel() > 0) + or has_attn_indices, + use_tllm_gen_sparse_attention=use_tllm_gen_sparse, + use_nvfp4_mla_kv_cache=( + quant_mode.has_fp4_kv_cache() + and use_tllm_gen_sparse + and arg("sparse_attn_kv_lens") is None + and arg("aux_kv_cache_pool_ptr") is not None + ), + mla_layer_num=pool_mapping.size(0) if pool_mapping is not None else 0, + quant_mode=quant_mode, + ) + + def to_thop_config(self) -> Any: + """Build the native constructor configuration.""" + from tensorrt_llm.bindings.internal import thop + + target = thop.StaticAttentionConfig() + build_op_params(target, self) + return target + + +@dataclass(slots=True) +class FmhaParams: + """Attention parameters shared by Python, DSL, Triton, and native FMHA paths. + + Offset contract, relied on by the native side. It splits by memory space: + + * **Device** per-token and per-sequence tensors are phase-local views. Slice them + here for the context or generation phase; C++ never re-slices them. + * **Host** tensors, the KV-cache block offsets and the FP4 scaling factors stay + whole-batch. C++ indexes them with ``seq_offset`` / ``token_offset``: pointer + accessors apply the offset themselves, so call sites never pass it; only the + explicit max-over-range queries take a range. + + Applying an offset to the first group double-counts it; omitting it for the second + shifts every sequence by the number of context requests. + """ + + fwd: AttentionForwardArgs = None + # Objects whose types have no native schema are Python-only. + # FMHA backends that need layer/metadata state the flat schema does not carry + # (e.g. the Triton custom-mask backend) read them from here. + local_layer_idx: int = -1 + has_fp8_kv_cache: bool = False + kv_pool: Optional[torch.Tensor] = None + use_paged_context_fmha: bool = False + kv_factor: int = 1 + seq_offset: int = 0 + num_seqs: int = 0 + num_tokens: int = 0 + # First query token of this phase on the axis of the q handed to the + # library. Phase tensors are already sliced; separate per-token inputs + # such as sparse block tables use this offset to select the same phase. + token_offset: int = 0 + num_requests: int = 0 + + layer_idx: int = -1 + rotary_embedding_base: float = 10000.0 + rotary_embedding_scale_type: RotaryScalingType = 0 + rotary_embedding_scale: float = 1.0 + rotary_embedding_short_mscale: float = 1.0 + rotary_embedding_long_mscale: float = 1.0 + rotary_embedding_max_positions: int = 1024 + rotary_embedding_original_max_positions: int = 1024 + max_context_length: int = 0 + max_seq_len: int = 0 + max_num_requests: int = 0 + # Total sequence rows (max_num_requests * beam_width). The generation workspace and + # the multi-block counter are sized per sequence, not per request. + max_num_sequences: int = 0 + # Total number of sequences, i.e. max_num_requests * beam_width. The generation + # workspace and the multi-block counter are sized per sequence, not per request. + beam_width: int = 1 + use_spec_decoding: bool = False + is_spec_dec_tree: bool = True + force_prepare_spec_dec_tree_mask: bool = False + attention_chunk_size: Optional[int] = None + + spec_decoding_target_max_draft_tokens: Optional[int] = None + + workspace: torch.Tensor = None + output: torch.Tensor = None + qkv_or_q: torch.Tensor = None + k: Optional[torch.Tensor] = None + v: Optional[torch.Tensor] = None + + sequence_length: torch.Tensor = cpp_metadata(dtype=torch.int32) + host_past_key_value_lengths: torch.Tensor = None + # CPU totals [context, generation] for callers whose past lengths exclude new tokens. + host_total_kv_lens: Optional[torch.Tensor] = None + context_lengths: torch.Tensor = cpp_metadata(dtype=torch.int32) + host_context_lengths: torch.Tensor = None + max_context_q_len_override: Optional[int] = None + kv_cache_block_offsets: Optional[torch.Tensor] = None + host_kv_cache_pool_pointers: Optional[torch.Tensor] = None + host_kv_cache_pool_mapping: Optional[torch.Tensor] = None + cache_indirection: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + max_attention_window_size: int = 0 + cyclic_attention_window_size: int = 0 + + rotary_inv_freq: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + rotary_cos_sin: Optional[torch.Tensor] = None + + block_ids_per_seq: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + + helix_position_offsets: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + helix_is_inactive_rank: Optional[torch.Tensor] = cpp_metadata(dtype=torch.bool) + + spec_decoding_generation_lengths: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + spec_decoding_position_offsets: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + spec_decoding_packed_mask: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int64) + spec_decoding_bl_tree_mask: Optional[torch.Tensor] = cpp_metadata(dtype=torch.uint32) + spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.int32 + ) + + num_sparse_topk: int = 0 + + flash_mla_tile_scheduler_metadata: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + flash_mla_num_splits: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + + trtllm_gen_jit_warmup: bool = False + + is_cross: bool = False + + # Fused kv_a_layernorm for the DSv4 sparse context path: when set, `latent_cache` + # is the raw kv_a_proj output and the context RoPE kernel norms it in place. + + # Mechanical state consumed by handwritten C++ lowering hooks. Defaults + # remain Python-owned; the generated C++ holder is only value-initialized. + # NOTE: the KV-cache pool base pointers are deliberately absent. They are derived from + # host_kv_cache_pool_pointers plus a per-layer intra-pool byte offset, which depends on + # the resolved KV-cache element size, so they live in handwritten C++ lowering + # (FmhaParams::kv_cache_pool_pointers) rather than in this schema. + multi_ctas_kv_counter: Optional[torch.Tensor] = None + + @classmethod + def _from_arguments(cls, arguments: Mapping[str, object], /, **overrides: object) -> FmhaParams: + """Build compatibility parameters from same-named legacy arguments.""" + params_fields = dataclasses.fields(cls) + fields_by_name = {field.name: field for field in params_fields} + unknown = overrides.keys() - fields_by_name.keys() + if unknown: + names = ", ".join(sorted(unknown)) + raise ValueError(f"FmhaParams has no field(s): {names}") + values = { + field.name: arguments[field.name] for field in params_fields if field.name in arguments + } + values.update(overrides) + return cls(**values) + + def to_op_params(self, context: object = None) -> Any: + """Build native parameters from this Python interface.""" + from tensorrt_llm.bindings.internal import thop + + target = thop.FmhaParams() + build_op_params(target, self) + return target + + +@cache +def _native_field_names(source_type: type, native_type: type) -> tuple[str, ...]: + """Resolve the schema boundary once, independently of per-call field values.""" + return tuple( + field.name for field in dataclasses.fields(source_type) if hasattr(native_type, field.name) + ) + + +def build_op_params(target: Any, *sources: object) -> None: + """Build native parameters from schema dataclasses, with later sources taking priority. + + The generated native bindings define which fields cross the boundary; + Python-only fields are left alone. A later None suppresses an earlier value + and leaves the native default intact, including for value-initialized nested structs. + """ + values = { + name: getattr(source, name) + for source in sources + for name in _native_field_names(type(source), type(target)) + } + for name, value in values.items(): + if value is None: + continue + if dataclasses.is_dataclass(value): + build_op_params(getattr(target, name), value) + else: + setattr(target, name, value) class FmhaPhase(str, Enum): @@ -48,6 +451,9 @@ class Fmha(ABC): def __init__(self, attn: "TrtllmAttention"): self._attn_ref: weakref.ReferenceType["TrtllmAttention"] = weakref.ref(attn) + def release(self) -> None: + """Release implementation-owned resources before CUDA teardown.""" + @property def attn(self) -> "TrtllmAttention": attn = self._attn_ref() diff --git a/tensorrt_llm/_torch/attention/backends/fmha/manager.py b/tensorrt_llm/_torch/attention/backends/fmha/manager.py index c9afff66e1d8..95d5531f675f 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/manager.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/manager.py @@ -312,6 +312,11 @@ def __init__(self, attn: TrtllmAttention) -> None: if fmha_cls.is_available(attn): self.fmha_libs.append(fmha_cls(attn)) + def release(self) -> None: + """Release each owned implementation, including those used by CombinedFmha.""" + for fmha in self.fmha_libs: + fmha.release() + def _make_cache_key( self, q: torch.Tensor, diff --git a/tensorrt_llm/_torch/attention/backends/fmha/phased.py b/tensorrt_llm/_torch/attention/backends/fmha/phased.py index f53aac9569a7..4d44f191abe9 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/phased.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/phased.py @@ -26,6 +26,7 @@ ) from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2, Role from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._utils import get_sm_version from .interface import Fmha @@ -73,26 +74,52 @@ class FmhaParams: batch_size: int = 0 # Number of logical requests in the active phase. num_requests: int = 0 + use_spec_decoding: bool = False spec_decoding_generation_lengths: Optional[torch.Tensor] = None spec_decoding_position_offsets: Optional[torch.Tensor] = None + spec_decoding_packed_mask: Optional[torch.Tensor] = None + spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor] = None + spec_decoding_bl_tree_mask: Optional[torch.Tensor] = None + spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor] = None is_cross: bool = False +def get_spec_decoding_position_offsets( + metadata: "TrtllmAttentionMetadata", +) -> Optional[torch.Tensor]: + """Return the active int32 offsets as a zero-copy [requests, query width] view.""" + if not (metadata.is_spec_decoding_enabled and metadata.use_spec_decoding): + return None + offsets = metadata.spec_decoding_position_offsets + if offsets is not None and offsets.dim() == 1: + if not metadata.is_sm_version_trtllm_gen_kernel(sm=get_sm_version()): + # Hopper offsets use the current compact query width, which can be + # smaller than the persistent buffer's capacity. + query_len = metadata.spec_decoding_query_len + if query_len <= 0: + raise ValueError( + "1-D speculative position offsets require a positive query length." + ) + offsets = offsets[: metadata.max_num_requests * query_len].view( + metadata.max_num_requests, query_len + ) + else: + offsets = offsets.view(metadata.max_num_requests, -1) + return offsets + + class PhasedFmha(Fmha): """FMHA helper for paged-KV libraries that split work by request phase.""" REQUIRES_PAGED_KV = True + # Required by backends that construct tensor views over the KV pool. + NEEDS_BLOCK_EXTENT = True def __init__(self, attn: "TrtllmAttention"): super().__init__(attn) self.kv_factor = 1 if attn.is_mla_enable else 2 - kv_lora_rank = attn.kv_lora_rank or 0 - self.generation_out_head_size = ( - kv_lora_rank if attn.is_mla_enable and kv_lora_rank else attn.head_dim - ) - self.context_out_head_size = ( - attn.v_head_dim if attn.is_mla_enable and attn.v_head_dim else attn.head_dim - ) + self.generation_out_head_size = attn.out_head_size(is_gen_only=True) + self.context_out_head_size = attn.out_head_size(is_gen_only=False) self._v1_total_num_blocks_cache: Optional[tuple[object, int, int]] = None def _get_total_num_blocks( @@ -188,6 +215,13 @@ def forward( num_contexts = metadata.num_contexts num_ctx_tokens = metadata.num_ctx_tokens num_generations = metadata.num_generations + has_context = num_contexts > 0 and not is_gen_only + has_generation = ( + num_generations > 0 and attention_input_type != AttentionInputType.context_only + ) + fused_epilogue = forward_args.enable_dsv4_epilogue_fusion + if fused_epilogue and has_context and has_generation: + raise ValueError("DSv4 fused epilogue requires separate context and generation calls.") num_gen_tokens = num_tokens if is_gen_only else num_tokens - num_ctx_tokens if num_gen_tokens < 0: raise RuntimeError( @@ -204,8 +238,17 @@ def forward( workspace, ) - out_head_size = self.generation_out_head_size if is_gen_only else self.context_out_head_size - out_tensor = output.view(num_tokens, attn.num_heads, out_head_size) + if fused_epilogue: + # DSv4 owns a contiguous [groups, tokens, K] output buffer for each phase. + out_tensor = output + else: + out_head_size = ( + self.generation_out_head_size if is_gen_only else self.context_out_head_size + ) + if output.dtype == torch.uint8: + # NVFP4 stores two output values in each byte. + out_head_size //= 2 + out_tensor = output.view(num_tokens, attn.num_heads, out_head_size) attention_window_size = forward_args.attention_window_size cache_indirection = metadata.cache_indirection @@ -232,14 +275,16 @@ def forward( cyclic_attention_window_size=attention_window_size, tokens_per_block=tokens_per_block, kv_factor=self.kv_factor, - total_num_blocks=self._get_total_num_blocks(metadata), + total_num_blocks=( + self._get_total_num_blocks(metadata) if self.NEEDS_BLOCK_EXTENT else 0 + ), is_cross=metadata.is_cross, ) sequence_length = metadata.kv_lens_cuda_runtime host_past_key_value_lengths = metadata.kv_lens_runtime - if num_contexts > 0 and attention_input_type != AttentionInputType.generation_only: + if has_context: seq_offset = 0 token_offset = 0 num_seqs = num_contexts @@ -260,7 +305,11 @@ def forward( params.value_input = ( v[token_offset : token_offset + num_ctx_tokens] if v is not None else None ) - params.output = out_tensor[token_offset : token_offset + num_ctx_tokens] + params.output = ( + out_tensor + if fused_epilogue + else out_tensor[token_offset : token_offset + num_ctx_tokens] + ) params.sequence_lengths = sequence_length[seq_offset:] params.context_lengths = context_lengths[seq_offset:] params.max_past_kv_length = max_past_kv_len @@ -275,7 +324,7 @@ def forward( else: self.run_context(params) - if num_generations > 0 and attention_input_type != AttentionInputType.context_only: + if has_generation: seq_offset = num_contexts token_offset = 0 if is_gen_only else num_ctx_tokens num_seqs = num_generations @@ -285,17 +334,20 @@ def forward( ) input_seq_length = num_gen_tokens // num_seqs if num_seqs > 0 else 1 - predicted_tokens_per_seq = attn.predicted_tokens_per_seq - spec_gen_lengths = None - spec_pos_offsets = None - if metadata.is_spec_decoding_enabled and predicted_tokens_per_seq > 1: - spec_gen_lengths = metadata.spec_decoding_generation_lengths - position_offsets_for_cpp = metadata.spec_decoding_position_offsets_for_cpp - if position_offsets_for_cpp is not None and position_offsets_for_cpp.dim() == 1: - position_offsets_for_cpp = position_offsets_for_cpp.view( - metadata.max_num_requests, -1 - ) - spec_pos_offsets = position_offsets_for_cpp + params.use_spec_decoding = ( + metadata.is_spec_decoding_enabled and metadata.use_spec_decoding + ) + if params.use_spec_decoding: + params.spec_decoding_generation_lengths = metadata.spec_decoding_generation_lengths + params.spec_decoding_position_offsets = get_spec_decoding_position_offsets(metadata) + params.spec_decoding_packed_mask = metadata.spec_decoding_packed_mask + params.spec_decoding_bl_tree_mask_offset = ( + metadata.spec_decoding_bl_tree_mask_offset + ) + params.spec_decoding_bl_tree_mask = metadata.spec_decoding_bl_tree_mask + params.spec_bl_tree_first_sparse_mask_offset_kv = ( + metadata.spec_bl_tree_first_sparse_mask_offset_kv + ) phase_input = q[token_offset : token_offset + num_gen_tokens] params.qkv_input = phase_input if is_fused_qkv else None @@ -306,8 +358,15 @@ def forward( params.value_input = ( v[token_offset : token_offset + num_gen_tokens] if v is not None else None ) - params.output = out_tensor[token_offset : token_offset + num_gen_tokens] + params.output = ( + out_tensor + if fused_epilogue + else out_tensor[token_offset : token_offset + num_gen_tokens] + ) params.sequence_lengths = sequence_length[seq_offset:] + params.context_lengths = metadata.prompt_lens_cuda_runtime[ + seq_offset : seq_offset + num_seqs + ] params.max_past_kv_length = max_past_kv_len params.num_tokens = num_gen_tokens params.seq_offset = seq_offset @@ -315,8 +374,6 @@ def forward( params.input_seq_length = input_seq_length params.batch_size = num_seqs params.num_requests = num_seqs // metadata.beam_width - params.spec_decoding_generation_lengths = spec_gen_lengths - params.spec_decoding_position_offsets = spec_pos_offsets if attn.is_mla_enable: self.run_mla_generation(params) else: diff --git a/tensorrt_llm/_torch/attention/backends/fmha/utils.py b/tensorrt_llm/_torch/attention/backends/fmha/utils.py index b296f10dae3e..2b1e6f3c0fad 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/utils.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/utils.py @@ -122,6 +122,52 @@ def get_trtllm_gen_context_workspace_size( return int(layout["total_size"]) +# The trtllm-gen kernels index the counters in groups of eight. +_MULTI_CTAS_KV_COUNTER_ALIGNMENT = 8 + + +def get_multi_ctas_kv_counter_size( + num_heads: int, + max_num_sequences: int, + multi_processor_count: int, +) -> int: + """Bytes the trtllm-gen multi-CTA KV counter needs. + + One counter per head per decoder sequence, floored at the SM count for the kernels + that split a single head across blocks, then rounded up to the group size. + """ + num_counters = max(num_heads * max_num_sequences, multi_processor_count) + aligned_num_counters = ( + (num_counters + _MULTI_CTAS_KV_COUNTER_ALIGNMENT - 1) + // _MULTI_CTAS_KV_COUNTER_ALIGNMENT + * _MULTI_CTAS_KV_COUNTER_ALIGNMENT + ) + return aligned_num_counters * torch.int32.itemsize + + +def get_multi_ctas_kv_counter( + counter: Optional[torch.Tensor], + device: torch.device, + num_heads: int, + max_num_sequences: int, +) -> torch.Tensor: + """Return a zeroed counter buffer, reusing `counter` when it already fits. + + The buffer is per FMHA library, not shared: two libraries running different phases + of one batch each need their own, and the kernels reset it at the end of a launch. + Generation consumers dereference it without a null check, so a library dispatching a + generation phase must supply one. + """ + index = device.index if device.index is not None else torch.cuda.current_device() + size = get_multi_ctas_kv_counter_size( + num_heads, max_num_sequences, get_multi_processor_count_for_device(index) + ) + if counter is None or counter.device != device or counter.numel() < size: + counter = torch.empty(size, dtype=torch.uint8, device=device) + counter.zero_() + return counter + + @lru_cache(maxsize=None) def get_multi_processor_count_for_device(device_index: int) -> int: return torch.cuda.get_device_properties(device_index).multi_processor_count diff --git a/tensorrt_llm/_torch/attention/backends/interface.py b/tensorrt_llm/_torch/attention/backends/interface.py index 55b88face9d4..8f30aa809de7 100644 --- a/tensorrt_llm/_torch/attention/backends/interface.py +++ b/tensorrt_llm/_torch/attention/backends/interface.py @@ -32,6 +32,7 @@ from ...pyexecutor.resource_manager import KVCacheManager from ...pyexecutor.trace_log_utils import log_tensor_size from ...utils import get_model_extra_attrs +from .cpp_schema import cpp_metadata from .sparse.params import (SparseBackendForwardArgs, SparseMetadataParams, SparseRuntimeParams) @@ -904,48 +905,35 @@ class CustomAttentionMask(str, Enum): class AttentionForwardArgs: """Per-forward optional arguments for attention backends.""" + # Caller-facing output buffer, allocated here when absent. Native kernels + # consume the per-phase slice in ``FmhaParams.output``, not this full buffer. output: Optional[torch.Tensor] = None output_sf: Optional[torch.Tensor] = None - - out_scale: Optional[torch.Tensor] = None - out_scale_sf: Optional[torch.Tensor] = None - kv_scale_orig_quant: Optional[torch.Tensor] = None - kv_scale_quant_orig: Optional[torch.Tensor] = None - - attention_mask: AttentionMask = PredefinedAttentionMask.CAUSAL - attention_input_type: AttentionInputType = AttentionInputType.mixed - attention_window_size: Optional[int] = None - attention_mask_data: Optional[torch.Tensor] = None - attention_sinks: Optional[torch.Tensor] = None - relative_attention_bias: Optional[torch.Tensor] = None - relative_attention_max_distance: int = 0 - cross_kv: Optional[torch.Tensor] = None - + kv_scale_orig_quant: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.float32) + kv_scale_quant_orig: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.float32) + out_scale: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + out_scale_sf: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) latent_cache: Optional[torch.Tensor] = None q_pe: Optional[torch.Tensor] = None mrope_rotary_cos_sin: Optional[torch.Tensor] = None - mrope_position_deltas: Optional[torch.Tensor] = None - + mrope_position_deltas: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.int32) softmax_stats_tensor: Optional[torch.Tensor] = None - chunked_prefill_buffer_batch_size: int = 1 - - cu_q_seqlens: Optional[torch.Tensor] = None - cu_kv_seqlens: Optional[torch.Tensor] = None - fmha_scheduler_counter: Optional[torch.Tensor] = None - # Testing only: skip the RoPE step of MLA generation (the standalone harness - # feeds a pre-RoPE'd fused_q). The TRTLLM backend then appends the new latent - # and inits the trtllm-gen scheduler buffers itself. - skip_mla_rope_generation: bool = False - - mla_bmm1_scale: Optional[torch.Tensor] = None - mla_bmm2_scale: Optional[torch.Tensor] = None + attention_sinks: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + cu_q_seqlens: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + cu_kv_seqlens: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + fmha_scheduler_counter: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.uint32) + mla_bmm1_scale: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + mla_bmm2_scale: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) quant_q_buffer: Optional[torch.Tensor] = None - # Per-tensor FP8 scale (fp32 [1]) for the fused DSv4 FP8-Q-quant path. - # When non-None alongside `quant_q_buffer`, the C++ op skips - # `quantizeCopyInputToFp8Kernel`. - quant_scale_qkv: Optional[torch.Tensor] = None - - dsv4_inv_rope_cos_sin_cache: Optional[torch.Tensor] = None + cross_kv: Optional[torch.Tensor] = None + relative_attention_bias: Optional[torch.Tensor] = None + quant_scale_qkv: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + dsv4_inv_rope_cos_sin_cache: Optional[torch.Tensor] = cpp_metadata( + dtype=torch.float32) enable_dsv4_epilogue_fusion: bool = False # Fused kv_a_layernorm, DSv4 sparse context path. When set, `latent_cache` is the @@ -954,6 +942,13 @@ class AttentionForwardArgs: kv_norm_weight: Optional[torch.Tensor] = None kv_norm_eps: float = 1e-6 + attention_mask: AttentionMask = PredefinedAttentionMask.CAUSAL + attention_input_type: AttentionInputType = AttentionInputType.mixed + attention_window_size: Optional[int] = None + attention_mask_data: Optional[torch.Tensor] = cpp_metadata(dtype=torch.bool) + relative_attention_max_distance: int = 0 + chunked_prefill_buffer_batch_size: int = 1 + skip_mla_rope_generation: bool = False sage_attn_num_elts_per_blk_q: int = 0 sage_attn_num_elts_per_blk_k: int = 0 sage_attn_num_elts_per_blk_v: int = 0 @@ -962,21 +957,18 @@ class AttentionForwardArgs: # Packed QKV for non-MLA attention. MLA always passes a separate query. is_fused_qkv: bool = False update_kv_cache: bool = True - # Optional normalized diffusion timestep for timestep-varying sparse attention. - timestep: Optional[torch.Tensor] = None - + timestep: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) sparse_backend_args: Optional[SparseBackendForwardArgs] = None sparse_runtime_params: SparseRuntimeParams = field( default_factory=SparseRuntimeParams) @property - def mask_type(self) -> int: - """Integer mask type accepted by the C++ attention op - (``causal`` or ``padding``).""" + def mask_type(self) -> AttentionMaskType: + """Return the mask type this forward pass asks the native attention op for.""" if self.attention_mask == PredefinedAttentionMask.CAUSAL: - return int(AttentionMaskType.causal) + return AttentionMaskType.causal if self.attention_mask == PredefinedAttentionMask.FULL: - return int(AttentionMaskType.padding) + return AttentionMaskType.padding raise ValueError( f"Unexpected attention mask type: {self.attention_mask!r}") @@ -986,23 +978,19 @@ def mask_type(self) -> int: def merge_attention_forward_args( - forward_args: Optional[AttentionForwardArgs], - kwargs: Dict[str, Any], -) -> AttentionForwardArgs: + forward_args: Optional[AttentionForwardArgs], + kwargs: Dict[str, Any]) -> AttentionForwardArgs: """Merge legacy attention kwargs into explicit forward arguments.""" - unknown_kwargs = sorted(set(kwargs) - _ATTENTION_FORWARD_ARGS_FIELDS) if unknown_kwargs: raise ValueError( f"Unknown attention forward arguments: {unknown_kwargs}") - if forward_args is not None: if kwargs: raise ValueError( "Pass attention forward options either through forward_args " f"or as legacy kwargs, not both: {sorted(kwargs)}") return forward_args - return AttentionForwardArgs(**kwargs) diff --git a/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/backend.py b/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/backend.py index e42de4ea4ccc..b8ca3589797d 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/backend.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/backend.py @@ -192,7 +192,8 @@ def _prepare_sparse_forward_args( start_idx = 0 end_idx = metadata.num_tokens - sparse_args = forward_args.sparse_runtime_params + sparse_args = replace(forward_args.sparse_runtime_params) + forward_args.sparse_runtime_params = sparse_args sparse_args.sparse_attn_kv_lens = metadata.sparse_mla_topk_lens[self.compress_ratio][ start_idx:end_idx ] diff --git a/tensorrt_llm/_torch/attention/backends/sparse/hooks.py b/tensorrt_llm/_torch/attention/backends/sparse/hooks.py index abddc8ad0b2d..7193bddd3c25 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/hooks.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/hooks.py @@ -237,22 +237,24 @@ def prepare_sparse_runtime_params( ``SparseRuntimeParams`` built from ``forward_args.sparse_runtime_params`` plus the hook results. Fields a backend writes into that carrier outside the hooks, such as an auxiliary pool pointer, are carried over. SkipSoftmax - backends receive their threshold schedule last. + backends receive their threshold schedule last. The per-call carrier is + available before the hooks run so they can also populate auxiliary fields. """ + runtime_params = replace(forward_args.sparse_runtime_params) + forward_args.sparse_runtime_params = runtime_params kv_indices, kv_offsets = backend.sparse_kv_predict(q, k, metadata, forward_args) attn_indices, attn_offsets = backend.sparse_attn_predict(q, k, metadata, forward_args) block_sparse_inputs = backend.block_sparse_attn_predict(q, k, v, metadata, forward_args) has_attn_indices = attn_indices is not None or attn_offsets is not None sparse_params = backend.sparse_params - runtime_params = replace( - forward_args.sparse_runtime_params, - sparse_kv_indices=kv_indices, - sparse_kv_offsets=kv_offsets, - sparse_attn_indices=attn_indices, - sparse_attn_offsets=attn_offsets, - sparse_attn_indices_block_size=sparse_params.indices_block_size if has_attn_indices else 0, - block_sparse_inputs=block_sparse_inputs, + runtime_params.sparse_kv_indices = kv_indices + runtime_params.sparse_kv_offsets = kv_offsets + runtime_params.sparse_attn_indices = attn_indices + runtime_params.sparse_attn_offsets = attn_offsets + runtime_params.sparse_attn_indices_block_size = ( + sparse_params.indices_block_size if has_attn_indices else 0 ) + runtime_params.block_sparse_inputs = block_sparse_inputs if isinstance(sparse_params, SkipSoftmaxParams): runtime_params = sparse_params.scheduler.get_runtime_params( runtime_params=runtime_params, diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py index afaad2ba8430..31111854a4a2 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py @@ -37,12 +37,10 @@ def _counter_size(num_heads: int, max_num_requests: int, device_index: int) -> i run, while computing it reaches C++ through a deferred import and a device-properties query on every dense layer of every step. """ - from tensorrt_llm._torch.attention.backends.fmha.flashinfer_trtllm_gen import ( - _get_multi_ctas_kv_counter_size, - ) + from tensorrt_llm._torch.attention.backends.fmha.utils import get_multi_ctas_kv_counter_size multi_processor_count = torch.cuda.get_device_properties(device_index).multi_processor_count - return int(_get_multi_ctas_kv_counter_size(num_heads, max_num_requests, multi_processor_count)) + return int(get_multi_ctas_kv_counter_size(num_heads, max_num_requests, multi_processor_count)) def _device_index(device: torch.device) -> int: diff --git a/tensorrt_llm/_torch/attention/backends/sparse/params.py b/tensorrt_llm/_torch/attention/backends/sparse/params.py index 2c87709bdc5a..128f3ac8f392 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/params.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/params.py @@ -20,6 +20,8 @@ import torch +from ..cpp_schema import cpp_metadata + _INDEXER_MQA_LOGITS_DEFAULT_ELEM_BUDGET = 1 << 31 _INDEXER_MQA_LOGITS_BYTES_PER_ELEMENT = 4 @@ -101,7 +103,7 @@ class SparseBackendForwardArgs: """Sparse inputs passed from an attention module to its backend.""" # Shared by algorithms that accept precomputed top-k indices. - topk_indices: Optional[torch.Tensor] = None + topk_indices: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) # Complete block-sparse routing payload predicted by the module before the # core forward; the default backend hook hands it through unchanged. block_sparse_inputs: Optional["BlockSparseForwardInputs"] = None @@ -154,14 +156,14 @@ class SparseRuntimeParams: """Complete per-attention sparse runtime state consumed by FMHA/``AttentionOp``.""" # Sparse index inputs shared by multiple algorithms. - sparse_kv_indices: Optional[torch.Tensor] = None - sparse_kv_offsets: Optional[torch.Tensor] = None - sparse_attn_indices: Optional[torch.Tensor] = None + sparse_kv_indices: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + sparse_kv_offsets: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + sparse_attn_indices: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) # Per-query offsets, or backend-specific secondary sparse indices # (DeepSeek-V4 fp8_ds_mla compressed-pool indices). - sparse_attn_offsets: Optional[torch.Tensor] = None + sparse_attn_offsets: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) sparse_attn_indices_block_size: int = 0 - sparse_attn_kv_lens: Optional[torch.Tensor] = None + sparse_attn_kv_lens: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) aux_kv_cache_pool_ptr: Optional[int] = None # SkipSoftmax prefill threshold; kernels divide it by context length. diff --git a/tensorrt_llm/_torch/attention/backends/trtllm.py b/tensorrt_llm/_torch/attention/backends/trtllm.py index 44f4ef79e8a0..dcf96794e27c 100644 --- a/tensorrt_llm/_torch/attention/backends/trtllm.py +++ b/tensorrt_llm/_torch/attention/backends/trtllm.py @@ -31,7 +31,8 @@ from tensorrt_llm._utils import get_sm_version, maybe_pin_memory, prefer_pinned from tensorrt_llm.bindings.internal import thop -from tensorrt_llm.functional import AttentionMaskType +from tensorrt_llm.functional import (AttentionMaskType, PositionEmbeddingType, + RotaryScalingType) from tensorrt_llm.logger import logger from tensorrt_llm.math_utils import ceil_div from tensorrt_llm.models.modeling_utils import QuantConfig @@ -172,11 +173,8 @@ def effective_beam_width(self) -> int: # parameters required for spec-dec mode max_total_draft_tokens: Optional[int] = None spec_decoding_position_offsets: Optional[torch.Tensor] = None - # C++ attention op requires a 2-D position_offsets tensor and reads - # sizes()[1] as the generation length / packed-mask row stride. - spec_decoding_position_offsets_cpp: Optional[torch.Tensor] = None - # Compact Hopper C++ row stride for 1D dynamic-tree offsets. - position_offsets_stride: int = 0 + # Current query width of the compact prefix in a 1-D dynamic-tree offsets buffer. + spec_decoding_query_len: int = 0 spec_decoding_packed_mask: Optional[torch.Tensor] = None spec_decoding_generation_lengths: Optional[torch.Tensor] = None spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor] = None @@ -260,19 +258,6 @@ def effective_workspace(self) -> Optional[torch.Tensor]: """Attention-kernel workspace, switching to the CUDA-graph copy under capture.""" return self.cuda_graph_workspace if self.is_cuda_graph else self.workspace - @property - def spec_decoding_position_offsets_for_cpp(self) -> Optional[torch.Tensor]: - """``spec_decoding_position_offsets`` reshaped to the 2D layout the C++ - kernel expects.""" - offsets = self.spec_decoding_position_offsets - if offsets is not None and offsets.dim() == 1: - if (self.spec_decoding_position_offsets_cpp is not None - and not self.is_sm_version_trtllm_gen_kernel( - sm=get_sm_version())): - return self.spec_decoding_position_offsets_cpp - return offsets.view(self.max_num_requests, -1) - return offsets - @property def max_context_length(self) -> int: """ @@ -348,23 +333,6 @@ def __post_init__(self) -> None: ) if self.runtime_features is not None else False self._post_init_with_buffers(self.cuda_graph_buffers) - def update_position_offsets_for_cpp(self, query_len: int) -> None: - """Refresh the C++ view of spec-dec position offsets.""" - offsets = self.spec_decoding_position_offsets - if offsets is None or offsets.dim() != 1: - self.spec_decoding_position_offsets_cpp = offsets - self.position_offsets_stride = 0 - return - - if self.max_num_requests > 0 and query_len > 0: - self.position_offsets_stride = query_len - total = self.max_num_requests * query_len - self.spec_decoding_position_offsets_cpp = offsets[:total].view( - self.max_num_requests, query_len) - else: - self.spec_decoding_position_offsets_cpp = offsets - self.position_offsets_stride = 0 - def _post_init_with_buffers(self, buffers) -> None: # Set a default value, as max_num_sequences is not always set. @@ -1496,8 +1464,6 @@ def update_spec_dec_param( self.spec_decoding_bl_tree_mask = None self.spec_bl_tree_first_sparse_mask_offset_kv = None - cpp_query_len = 0 - # Case 1: dynamic tree — copy per-request params from spec_tree_manager. if self.is_spec_dec_dynamic_tree: assert spec_tree_manager is not None, "spec_tree_manager is required for dynamic tree" @@ -1540,7 +1506,7 @@ def update_spec_dec_param( mask_src.reshape(-1), non_blocking=True) self.spec_decoding_generation_lengths[:batch_size].fill_(n_dt) - cpp_query_len = n_dt + self.spec_decoding_query_len = n_dt # Case 2: linear tree else: @@ -1558,8 +1524,7 @@ def update_spec_dec_param( self.max_num_requests, runtime_draft_token_buffer_width) self.spec_decoding_packed_mask = generate_spec_decoding_packed_mask( self.max_num_requests, runtime_draft_token_buffer_width) - - self.update_position_offsets_for_cpp(cpp_query_len) + self.spec_decoding_query_len = runtime_draft_token_buffer_width + 1 def generate_spec_decoding_generation_length(self, runtime_draft_len): self.spec_decoding_generation_lengths[:self.max_num_requests].fill_( @@ -1663,8 +1628,9 @@ def __init__( self.rotary_inv_freq, self.rotary_cos_sin = self.rope_params.create_rope_const_params( ) - self.position_embedding_type = int( - pos_embd_params.type) if pos_embd_params is not None else 0 + self.position_embedding_type = (pos_embd_params.type + if pos_embd_params is not None else + PositionEmbeddingType.learned_absolute) self.skip_softmax_stat = torch.zeros(2, dtype=torch.uint32, device='cuda') @@ -1687,7 +1653,7 @@ def __init__( self.kv_scale_orig_quant = 1.0 / self.kv_cache_scaling_factor self.local_layer_idx: Optional[int] = None - self._fmha_manager: FmhaManager + self._fmha_manager: Optional[FmhaManager] = None if not skip_create_weights_in_init: self.update_quant_config(self.quant_config) @@ -1697,9 +1663,16 @@ def uses_fp4_mla_attention(self) -> bool: return (self.is_mla_enable and self.has_fp4_kv_cache and self.sparse_params is None) + def release(self) -> None: + """Release implementation-owned resources before CUDA teardown.""" + if self._fmha_manager is not None: + self._fmha_manager.release() + def update_quant_config(self, new_quant_config: Optional[QuantConfig]): self.quant_config = new_quant_config or QuantConfig() - self.quant_mode = int(self.quant_config.layer_quant_mode) + self.quant_mode = self.quant_config.layer_quant_mode + # The op's kernel selection is derived from the quantization mode. + self.release() self.has_fp8_qdq = self.has_fp8_kv_cache = self.has_nvfp4 = False if self.quant_config is not None: @@ -1947,6 +1920,21 @@ def _is_nvfp4_output_kernel_available( is_mla_enable, ) + def out_head_size(self, is_gen_only: bool) -> int: + """Per-head width of the attention output for this phase. + + Single source of truth for the output buffer `create_output` allocates and + the view a phased FMHA library takes over it; the two must not drift. MLA + generation writes the latent (plus the rope part when it is not appended in + place), while MLA context writes the already-projected V head. + """ + if not self.is_mla_enable: + return self.head_dim + if not is_gen_only: + return self.v_head_dim + return (self.kv_lora_rank if self.rope_append else self.kv_lora_rank + + self.qk_rope_head_dim) + def create_output(self, q, *, is_quantize_output: bool, metadata: TrtllmAttentionMetadata, attention_mask: AttentionMask, is_gen_only: bool, @@ -1960,13 +1948,7 @@ def create_output(self, q, *, is_quantize_output: bool, num_tokens = q.size(0) if out_dtype is None: out_dtype = q.dtype - v_head_size = self.head_dim - if self.is_mla_enable: - if is_gen_only: - v_head_size = self.kv_lora_rank if self.rope_append else ( - self.kv_lora_rank + self.qk_rope_head_dim) - else: - v_head_size = self.v_head_dim + v_head_size = self.out_head_size(is_gen_only) if use_nvfp4_output: num_nvfp4_elements_per_container = 2 scaling_vector_size = 16 @@ -1993,8 +1975,8 @@ def rope_base(self) -> float: return self.rope_params.theta @property - def rope_scale_type(self) -> int: - return int(self.rope_params.scale_type) + def rope_scale_type(self) -> RotaryScalingType: + return self.rope_params.scale_type @property def rope_scale(self) -> float: @@ -2317,6 +2299,7 @@ def forward( assert metadata.kv_cache_manager is None assert metadata.num_contexts == metadata.num_seqs + assert self._fmha_manager is not None fmha = self._fmha_manager.select(self, q, k, v, metadata, forward_args) if fmha is None: diff --git a/tensorrt_llm/_torch/speculative/dflash_attention.py b/tensorrt_llm/_torch/speculative/dflash_attention.py index 5759211742e1..3b59c84879d7 100644 --- a/tensorrt_llm/_torch/speculative/dflash_attention.py +++ b/tensorrt_llm/_torch/speculative/dflash_attention.py @@ -135,16 +135,14 @@ def get_dflash_trtllm_gen_ops() -> DFlashTrtllmGenOps: import flashinfer - from ..attention.backends.fmha.flashinfer_trtllm_gen import ( - _get_multi_ctas_kv_counter_size, - _get_workspace_size, - ) + from ..attention.backends.fmha.flashinfer_trtllm_gen import _get_workspace_size + from ..attention.backends.fmha.utils import get_multi_ctas_kv_counter_size return DFlashTrtllmGenOps( append_paged_kv_cache=flashinfer.page.append_paged_kv_cache, batch_context_with_kv_cache=flashinfer.prefill.trtllm_batch_context_with_kv_cache, batch_decode_with_kv_cache=flashinfer.decode.trtllm_batch_decode_with_kv_cache, - get_multi_ctas_kv_counter_size=_get_multi_ctas_kv_counter_size, + get_multi_ctas_kv_counter_size=get_multi_ctas_kv_counter_size, get_workspace_size=_get_workspace_size, ) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 55721b44aceb..dc4258136e08 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -479,7 +479,7 @@ def __init__(self, # clone calls) cannot turn the cleanup restore into AttributeError. self._saved_packed_mask = None self._saved_position_offsets = None - self._saved_position_offsets_cpp = None + self._saved_spec_decoding_query_len = None self._saved_generation_lengths = None @property @@ -503,11 +503,11 @@ def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): if attn_metadata.spec_decoding_position_offsets is not None: self._saved_position_offsets = attn_metadata.spec_decoding_position_offsets.clone( ) - self._saved_position_offsets_cpp = ( - attn_metadata.spec_decoding_position_offsets_cpp) + self._saved_spec_decoding_query_len = ( + attn_metadata.spec_decoding_query_len) else: self._saved_position_offsets = None - self._saved_position_offsets_cpp = None + self._saved_spec_decoding_query_len = None if attn_metadata.spec_decoding_generation_lengths is not None: self._saved_generation_lengths = attn_metadata.spec_decoding_generation_lengths[:batch_size].clone( ) @@ -524,10 +524,10 @@ def _restore_attn_metadata_from_spec_dec(self, attn_metadata): if self._saved_position_offsets is not None: attn_metadata.spec_decoding_position_offsets.copy_( self._saved_position_offsets) - attn_metadata.spec_decoding_position_offsets_cpp = ( - self._saved_position_offsets_cpp) + attn_metadata.spec_decoding_query_len = ( + self._saved_spec_decoding_query_len) self._saved_position_offsets = None - self._saved_position_offsets_cpp = None + self._saved_spec_decoding_query_len = None if self._saved_generation_lengths is not None: batch_size = self._saved_generation_lengths.shape[0] attn_metadata.spec_decoding_generation_lengths[:batch_size].copy_( diff --git a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py index 527c7a240c02..ae7bad83b678 100644 --- a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py @@ -323,7 +323,7 @@ def __init__( def _repack_mask_padded_to_packed(self, mask_buf, n_req, n_tok): """XQA indexes mask flat via cuQSeqLens with per-request stride n_tok * ceil(n_tok/32) (C++ sets mPackedMaskMaxSeqLenQ from - spec_decoding_position_offsets_cpp.sizes()[1]). + the phase position-offset view's size(1)). The padded buffer is [n_req, buf_dim, ceil(buf_dim/32)] with a larger row stride, so when n_tok < buf_dim and n_req > 1 packed reads see misaligned rows. Copy the meaningful @@ -340,9 +340,9 @@ def _repack_mask_padded_to_packed(self, mask_buf, n_req, n_tok): flat[:total_elems] = scratch.view(-1) def _apply_spec_metadata(self, attn_metadata, batch_size, query_len): - """Set spec-dec gen lengths and refresh the C++ position-offset view.""" + """Set the current speculative query lengths and position-offset layout.""" attn_metadata.spec_decoding_generation_lengths[:batch_size] = query_len - attn_metadata.update_position_offsets_for_cpp(query_len) + attn_metadata.spec_decoding_query_len = query_len @nvtx_range("eagle3_dyn._ensure_spec_tree_manager") def _ensure_spec_tree_manager(self, resource_manager): diff --git a/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py index 91bb857fe279..dabdcefbe0bc 100644 --- a/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py @@ -195,10 +195,10 @@ def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): self._saved_packed_mask = None if attn_metadata.spec_decoding_position_offsets is not None: self._saved_position_offsets = attn_metadata.spec_decoding_position_offsets.clone() - self._saved_position_offsets_cpp = attn_metadata.spec_decoding_position_offsets_cpp + self._saved_spec_decoding_query_len = attn_metadata.spec_decoding_query_len else: self._saved_position_offsets = None - self._saved_position_offsets_cpp = None + self._saved_spec_decoding_query_len = None if attn_metadata.spec_decoding_generation_lengths is not None: self._saved_generation_lengths = attn_metadata.spec_decoding_generation_lengths[ :batch_size @@ -225,9 +225,9 @@ def _restore_attn_metadata_from_spec_dec(self, attn_metadata): self._saved_packed_mask = None if self._saved_position_offsets is not None: attn_metadata.spec_decoding_position_offsets.copy_(self._saved_position_offsets) - attn_metadata.spec_decoding_position_offsets_cpp = self._saved_position_offsets_cpp + attn_metadata.spec_decoding_query_len = self._saved_spec_decoding_query_len self._saved_position_offsets = None - self._saved_position_offsets_cpp = None + self._saved_spec_decoding_query_len = None if self._saved_generation_lengths is not None: batch_size = self._saved_generation_lengths.shape[0] attn_metadata.spec_decoding_generation_lengths[:batch_size].copy_( @@ -253,9 +253,9 @@ def _ensure_spec_dec_state_restored(self, attn_metadata, spec_metadata): # Helpers # # ------------------------------------------------------------------ # def _apply_spec_metadata(self, attn_metadata, batch_size, query_len): - """Set spec-dec gen lengths and refresh the C++ position-offset view.""" + """Set the current speculative query lengths and position-offset layout.""" attn_metadata.spec_decoding_generation_lengths[:batch_size] = query_len - attn_metadata.update_position_offsets_for_cpp(query_len) + attn_metadata.spec_decoding_query_len = query_len def _refresh_blackwell_tree_mask_metadata(self, attn_metadata): if not getattr(attn_metadata, "use_spec_decoding", False): diff --git a/tests/microbenchmarks/nvfp4_mla_kv_cache_gather.py b/tests/microbenchmarks/nvfp4_mla_kv_cache_gather.py index a1b51ea6ff60..53128bfe2327 100644 --- a/tests/microbenchmarks/nvfp4_mla_kv_cache_gather.py +++ b/tests/microbenchmarks/nvfp4_mla_kv_cache_gather.py @@ -32,8 +32,8 @@ import torch import tensorrt_llm._torch.custom_ops # noqa: F401 +from tensorrt_llm._torch.attention.backends.fmha.fallback import FallbackFmha from tensorrt_llm._torch.attention.backends.interface import AttentionInputType -from tensorrt_llm.bindings.internal import thop from tensorrt_llm.functional import AttentionMaskType, PositionEmbeddingType from tensorrt_llm.quantization import QuantMode @@ -294,13 +294,13 @@ def call_dsa_for_batch(self, batch_size: int) -> Callable[[], None]: def run() -> None: self.fmha_scheduler_counter.zero_() - thop.attention( + FallbackFmha.attention( q=self.query[:batch_size], k=None, v=None, output=self.attention_output[:batch_size], output_sf=None, - workspace_=self.attention_workspace, + workspace=self.attention_workspace, sequence_length=self.kv_lens_cuda[:batch_size], host_past_key_value_lengths=self.kv_lens_host[:batch_size], host_total_kv_lens=host_total_kv_lens, @@ -366,7 +366,7 @@ def run() -> None: use_spec_decoding=False, is_spec_dec_tree=False, spec_decoding_generation_lengths=None, - spec_decoding_position_offsets_for_cpp=None, + spec_decoding_position_offsets=None, spec_decoding_packed_mask=None, spec_decoding_bl_tree_mask_offset=None, spec_decoding_bl_tree_mask=None, diff --git a/tests/unittest/_torch/attention/fmha_test_utils.py b/tests/unittest/_torch/attention/fmha_test_utils.py index a20a99dbf156..4cdf040d08be 100644 --- a/tests/unittest/_torch/attention/fmha_test_utils.py +++ b/tests/unittest/_torch/attention/fmha_test_utils.py @@ -28,6 +28,8 @@ def __init__(self, local_layer_idx: int = 0) -> None: self.kv_lora_rank = None self.v_head_dim = None self.head_dim = 4 + self.qk_rope_head_dim = 0 + self.rope_append = None self.num_heads = 1 self.num_kv_heads = 1 self.predicted_tokens_per_seq = 1 @@ -35,6 +37,15 @@ def __init__(self, local_layer_idx: int = 0) -> None: self.skip_correction_threshold = 0.0 self.local_layer_idx = local_layer_idx + def out_head_size(self, is_gen_only: bool) -> int: + """Extents the fake layer reports, tolerant of its unset MLA dimensions.""" + if not is_gen_only: + return self.v_head_dim if self.is_mla_enable and self.v_head_dim else self.head_dim + kv_lora_rank = self.kv_lora_rank or 0 + if not (self.is_mla_enable and kv_lora_rank): + return self.head_dim + return kv_lora_rank if self.rope_append else kv_lora_rank + self.qk_rope_head_dim + class FakePhasedFmha(PhasedFmha): def __init__( diff --git a/tests/unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py b/tests/unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py index 019b790a762f..d20bd921a66e 100644 --- a/tests/unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py +++ b/tests/unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py @@ -98,6 +98,7 @@ def test_on_update_kv_lens_rebuilds_stale_map() -> None: # __init__ (bypassed by object.__new__) defaults this to False; # on_update_kv_lens() reads it since #16925. md.in_mtp_draft_loop = False + md._group_remap_batched = {} # Stub collaborators unrelated to the map rebuild (test_dsa_indexer.py style). md.kv_lens_cuda = torch.tensor([100, 200, 300], dtype=torch.int32, device=device) md._compute_kv_lens_row_reorder = Mock() diff --git a/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py b/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py index 40d74b28ede3..69d44e994a9a 100644 --- a/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py +++ b/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py @@ -35,6 +35,7 @@ BlockSparseForwardInputs, SparseRuntimeParams, ) +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.functional import PositionEmbeddingType @@ -120,6 +121,8 @@ def _proxy_reference( class _Attention: + out_head_size = TrtllmAttention.out_head_size + def __init__(self) -> None: self.sparse_params = None self.num_heads = 2 diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py index 6cf7b15fcc38..8d67b5e83658 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py @@ -77,9 +77,8 @@ def test_prepare_sparse_runtime_params_from_predictions() -> None: attention._sparse_kv_offsets = torch.tensor([0, 1], dtype=torch.int32) attention._sparse_attn_indices = torch.tensor([2], dtype=torch.int32) attention._sparse_attn_offsets = None - forward_args = AttentionForwardArgs( - sparse_runtime_params=SparseRuntimeParams(sparse_attn_kv_lens=torch.tensor([3])) - ) + caller_params = SparseRuntimeParams(sparse_attn_kv_lens=torch.tensor([3])) + forward_args = AttentionForwardArgs(sparse_runtime_params=caller_params) runtime_params = prepare_sparse_runtime_params( attention, torch.empty(0), None, None, None, forward_args @@ -90,9 +89,9 @@ def test_prepare_sparse_runtime_params_from_predictions() -> None: assert runtime_params.sparse_attn_indices is attention._sparse_attn_indices assert runtime_params.sparse_attn_offsets is None assert runtime_params.sparse_attn_indices_block_size == 1 - assert ( - runtime_params.sparse_attn_kv_lens is forward_args.sparse_runtime_params.sparse_attn_kv_lens - ) + assert runtime_params.sparse_attn_kv_lens is caller_params.sparse_attn_kv_lens + assert runtime_params is not caller_params + assert caller_params.sparse_kv_indices is None def test_sparse_attn_hook_registration() -> None: @@ -161,11 +160,16 @@ def test_prepare_sparse_runtime_params_without_predictions(sparse_params) -> Non attention = TrtllmAttention.__new__(TrtllmAttention) attention.sparse_params = sparse_params - runtime_params = prepare_sparse_runtime_params( - attention, torch.empty(0), None, None, None, AttentionForwardArgs() + forward_args = AttentionForwardArgs() + q = torch.empty(0) + runtime_params = prepare_sparse_runtime_params(attention, q, None, None, None, forward_args) + next_runtime_params = prepare_sparse_runtime_params( + attention, q, None, None, None, forward_args ) assert runtime_params == SparseRuntimeParams() + assert next_runtime_params == SparseRuntimeParams() + assert next_runtime_params is not runtime_params def test_prepare_sparse_runtime_params_runs_index_hooks_once() -> None: @@ -271,6 +275,28 @@ def test_attention_forward_args_default_to_empty_sparse_runtime_params() -> None assert AttentionForwardArgs().sparse_runtime_params == SparseRuntimeParams() +def test_sparse_prediction_hooks_share_per_call_runtime_params() -> None: + attention = TrtllmAttention.__new__(TrtllmAttention) + attention.sparse_params = None + carriers = [] + + def predict(q, k, metadata, forward_args): + runtime_params = forward_args.sparse_runtime_params + assert isinstance(runtime_params, SparseRuntimeParams) + carriers.append(runtime_params) + runtime_params.aux_kv_cache_pool_ptr = 1234 + return None, None + + attention.sparse_attn_predict = Mock(side_effect=predict) + for forward_args in (AttentionForwardArgs(), AttentionForwardArgs()): + runtime_params = prepare_sparse_runtime_params( + attention, torch.empty(0), None, None, None, forward_args + ) + assert runtime_params is carriers[-1] + assert runtime_params.aux_kv_cache_pool_ptr == 1234 + assert carriers[0] is not carriers[1] + + class _StopAfterShapeValidation(Exception): pass diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mha.py b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py index 9198efc034d2..631616ab9524 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mha.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py @@ -285,7 +285,7 @@ def create_generation_inputs(scenario: MhaGenerationScenario) -> MhaGenerationIn dtype=torch.int32, device=device, ) - metadata.update_position_offsets_for_cpp(scenario.max_query_len) + metadata.spec_decoding_query_len = scenario.max_query_len metadata.spec_decoding_param_prepare_for_blackwell() metadata.prepare() cleanup.pop_all() @@ -340,7 +340,7 @@ def reference_generation_attention( sparse_attn_indices: torch.Tensor, scenario: MhaGenerationScenario, ) -> torch.Tensor: - """Compute page-sparse MHA from equivalent request-local token indices.""" + """Compute page-sparse MHA in FP32 from request-local token indices.""" outputs = [] query_offset = 0 for request_idx in range(scenario.batch_size): @@ -354,7 +354,7 @@ def reference_generation_attention( ), ], dim=0, - ) + ).float() v_full = torch.cat( [ v_history, @@ -363,10 +363,10 @@ def reference_generation_attention( ), ], dim=0, - ) + ).float() for query_idx in range(scenario.query_len): packed_query_idx = query_offset + query_idx - q_token = q[packed_query_idx].view(scenario.num_heads, scenario.head_dim) + q_token = q[packed_query_idx].view(scenario.num_heads, scenario.head_dim).float() head_outputs = [] for head_idx in range(scenario.num_heads): token_indices = sparse_attn_indices[head_idx, packed_query_idx] @@ -378,7 +378,7 @@ def reference_generation_attention( ) attention_probs = torch.nn.functional.softmax( attention_scores, dim=-1, dtype=torch.float32 - ).to(scenario.dtype) + ) head_outputs.append(torch.matmul(attention_probs, v_sparse)) outputs.append(torch.cat(head_outputs, dim=0)) query_offset += scenario.query_len @@ -592,14 +592,13 @@ def _run_page_sparse_mha(scenario: PageSparseMhaScenario) -> None: uses_fp8 = ( attention_scenario.kvcache_dtype == torch.float8_e4m3fn or attention_scenario.fp8_output ) - output_for_comparison = output.float() if uses_fp8 else output - if attention_scenario.fp8_output: - reference_output = reference_output.to(torch.float8_e4m3fn) - reference_for_comparison = reference_output.float() if uses_fp8 else reference_output + # Keep the reference unquantized: rounding it to FP8 can amplify small + # kernel errors across a rounding boundary into a full FP8 step. + output_for_comparison = output.float() assert torch.isfinite(output_for_comparison).all() torch.testing.assert_close( output_for_comparison, - reference_for_comparison, + reference_output, atol=FP8_ATOL if uses_fp8 else ATOL, rtol=FP8_RTOL if uses_fp8 else RTOL, ) diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py index df7cf235fe26..5b260670159c 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -1617,7 +1617,7 @@ def _create_generation_inputs(s: GenerationScenario) -> _GenerationInputs: metadata.spec_decoding_generation_lengths = torch.tensor( [s.query_len] * s.batch_size, dtype=torch.int32, device=device ) - metadata.update_position_offsets_for_cpp(s.max_query_len) + metadata.spec_decoding_query_len = s.max_query_len metadata.spec_decoding_param_prepare_for_blackwell() metadata.prepare() diff --git a/tests/unittest/_torch/attention/test_attention_op_sync.py b/tests/unittest/_torch/attention/test_attention_op_sync.py deleted file mode 100644 index 5ca25019f86c..000000000000 --- a/tests/unittest/_torch/attention/test_attention_op_sync.py +++ /dev/null @@ -1,704 +0,0 @@ -# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. -# -# 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. -"""Static sync test for the fallback ``thop.attention(...)`` call in -``FallbackFmha.forward``. - -That call is the single explicit-kwarg call site for the C++ ``thop.attention`` -binding. This test parses both the call site (Python AST) and the C++ -function declaration in ``attentionOp.h`` (text/regex), and enforces: - -1. Every C++ parameter name appears at the call site (and nothing extra). -2. Every call-site kwarg sourced as ``root.attr[.attr...]`` resolves on - exactly one of ``attn`` / ``metadata`` / ``forward_args``, and its - declared C++ type matches the source attribute's Python type at a - coarse-category level (tensor / int / bool / float / list-of-X). -3. Every dataclass field reachable from ``AttentionForwardArgs`` (including - nested dataclass sub-bags like ``SparseRuntimeParams``) is consumed at - the call site — directly, transitively via a @property of the - containing class, or listed in ``_THOP_EXCLUDED_FIELDS``. -4. Every kwarg passed as a literal constant matches an entry in - ``_THOP_LITERALS`` (both name and value). - -The test is AST-only (no kernel run) so it fails fast and runs without a GPU. -""" - -import ast -import dataclasses -import inspect -import pathlib -import re -import textwrap -import typing -from dataclasses import fields -from types import SimpleNamespace, UnionType - -import pytest -import torch - -from tensorrt_llm._torch.attention.backends.fmha.fallback import ( - _THOP_EXCLUDED_FIELDS, - _THOP_LITERALS, - FallbackFmha, -) -from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention, TrtllmAttentionMetadata - -pytestmark = pytest.mark.cpu_only - - -# Roots used as the LHS of attribute chains at the call site. Match the -# names inside ``FallbackFmha.forward``. -_SOURCE_CLASSES = { - "attn": TrtllmAttention, - "metadata": TrtllmAttentionMetadata, - "forward_args": AttentionForwardArgs, -} - -_THOP_KWARG_SOURCE_ALIASES: dict[str, tuple[str, tuple[str, ...]]] = { - "beam_width": ("metadata", ("effective_beam_width",)), - "context_lengths": ("metadata", ("prompt_lens_cuda_runtime",)), - "head_size": ("attn", ("head_dim",)), - "host_context_lengths": ("metadata", ("prompt_lens_cpu_runtime",)), - "host_past_key_value_lengths": ("metadata", ("kv_lens_runtime",)), - "host_request_types": ("metadata", ("host_request_types_runtime",)), - "sequence_length": ("metadata", ("kv_lens_cuda_runtime",)), - "spec_decoding_target_max_draft_tokens": ( - "metadata", - ("max_total_draft_tokens",), - ), - "skip_softmax_threshold_scale_factor_decode": ( - "forward_args", - ( - "sparse_runtime_params", - "threshold_scale_factor_decode", - ), - ), - "skip_softmax_threshold_scale_factor_prefill": ( - "forward_args", - ( - "sparse_runtime_params", - "threshold_scale_factor_prefill", - ), - ), - "workspace_": ("metadata", ("effective_workspace",)), -} - -# The C++ attention() declaration is the single source of truth for kwarg -# names, ordering, and types. -_HEADER = pathlib.Path(__file__).resolve().parents[4] / ("cpp/tensorrt_llm/thop/attentionOp.h") - - -# ---- C++ declaration parser ------------------------------------------------- - - -_TENSOR_RE = re.compile(r"\btorch::Tensor\b") -_QUANT_MODE_RE = re.compile(r"\bcommon::QuantMode\b") -_INT64_RE = re.compile(r"\bint64_t\b") -_DOUBLE_RE = re.compile(r"\bdouble\b") -_BOOL_RE = re.compile(r"\bbool\b") - - -def _split_top_level(s: str, sep: str = ",") -> list[str]: - """Split ``s`` by ``sep`` at angle-bracket depth 0.""" - out: list[str] = [] - depth = 0 - buf: list[str] = [] - for ch in s: - if ch == "<": - depth += 1 - buf.append(ch) - elif ch == ">": - depth -= 1 - buf.append(ch) - elif ch == sep and depth == 0: - out.append("".join(buf).strip()) - buf = [] - else: - buf.append(ch) - if buf: - out.append("".join(buf).strip()) - return out - - -def _strip_inner_type(cpp_type: str) -> str: - """Strip ``std::optional<...>`` and ``const`` / references to expose the - inner element type. Idempotent.""" - t = cpp_type.replace("const", "").replace("&", "").strip() - m = re.fullmatch(r"std::optional<\s*(.+)\s*>", t) - if m: - t = m.group(1).strip() - return t - - -def _cpp_category(cpp_type: str) -> str: - """Coarse Python-side category for a C++ parameter type. - - Returns one of ``tensor`` / ``int`` / ``bool`` / ``float`` / - ``list_tensor`` / ``list_int`` / ``list_bool`` / ``list_float`` / - ``unknown``. Optional/const wrappers and references are ignored. - """ - bare = _strip_inner_type(cpp_type) - if bare.startswith("std::vector<"): - inner_match = re.fullmatch(r"std::vector<\s*(.+)\s*>", bare) - if inner_match: - inner_cat = _cpp_category(inner_match.group(1)) - return f"list_{inner_cat}" - if _TENSOR_RE.search(bare): - return "tensor" - if _BOOL_RE.fullmatch(bare): - return "bool" - if _INT64_RE.fullmatch(bare) or _QUANT_MODE_RE.fullmatch(bare): - return "int" - if _DOUBLE_RE.fullmatch(bare): - return "float" - return "unknown" - - -def _parse_attention_decl() -> list[tuple[str, str]]: - """Parse the ``void attention(...)`` declaration in ``attentionOp.h`` - and return ``[(name, cpp_type), ...]`` in declaration order.""" - src = _HEADER.read_text() - # Strip line comments to keep the regex tidy. - src = re.sub(r"//[^\n]*", "", src) - m = re.search(r"void\s+attention\s*\(", src) - if not m: - raise AssertionError(f"Could not find void attention(...) in {_HEADER}") - # Walk forward, matching the opening paren at m.end()-1 to its close. - start = m.end() - 1 - depth = 0 - end = None - for i in range(start, len(src)): - if src[i] == "(": - depth += 1 - elif src[i] == ")": - depth -= 1 - if depth == 0: - end = i - break - if end is None: - raise AssertionError("Unbalanced parens in attention() declaration") - body = src[start + 1 : end] - - params: list[tuple[str, str]] = [] - for raw in _split_top_level(body): - # Drop ``= default`` suffix at top level only. - param = _split_top_level(raw, sep="=")[0].strip() - # The name is the trailing identifier. Strip trailing '&' or '*' - # attached to the type, not the name. - m = re.match(r"(.*?)([A-Za-z_]\w*)\s*$", param, re.DOTALL) - if not m: - raise AssertionError(f"Could not parse param: {param!r}") - cpp_type = m.group(1).strip() - name = m.group(2) - params.append((name, cpp_type)) - return params - - -def _binding_kwargs() -> set[str]: - """Set of parameter names declared on ``void attention(...)``.""" - return {name for name, _ in _parse_attention_decl()} - - -def _binding_types() -> dict[str, str]: - """Map parameter name → declared C++ type.""" - return dict(_parse_attention_decl()) - - -# ---- Call-site AST helpers -------------------------------------------------- - - -def _parse_thop_attention_call() -> ast.Call: - """Locate the single ``thop.attention(...)`` call inside - ``FallbackFmha.forward``.""" - src = textwrap.dedent(inspect.getsource(FallbackFmha.forward)) - tree = ast.parse(src) - for node in ast.walk(tree): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) - and node.func.attr == "attention" - and isinstance(node.func.value, ast.Name) - and node.func.value.id == "thop" - ): - return node - raise AssertionError("Could not find thop.attention(...) call in FallbackFmha.forward") - - -def _attribute_path(node: ast.AST) -> tuple[str, tuple[str, ...]] | None: - """If ``node`` is a pure attribute chain ``Name.attr1.attr2...``, return - ``(root_name_id, (attr1, attr2, ...))``. Otherwise return ``None``. - """ - if not isinstance(node, ast.Attribute): - return None - attrs: list[str] = [] - current: ast.AST = node - while isinstance(current, ast.Attribute): - attrs.append(current.attr) - current = current.value - if not isinstance(current, ast.Name): - return None - return current.id, tuple(reversed(attrs)) - - -def _getattr_path(node: ast.AST) -> tuple[str, tuple[str, ...]] | None: - """If ``node`` is ``getattr(Name.attr..., "leaf", )``, return - ``(root_name_id, (attr1, ..., leaf))``. Otherwise return ``None``. - """ - if ( - not isinstance(node, ast.Call) - or not isinstance(node.func, ast.Name) - or node.func.id != "getattr" - or len(node.args) not in (2, 3) - or not isinstance(node.args[1], ast.Constant) - or not isinstance(node.args[1].value, str) - ): - return None - root_path: tuple[str, tuple[str, ...]] - if isinstance(node.args[0], ast.Name): - root_path = (node.args[0].id, ()) - else: - path = _attribute_path(node.args[0]) - if path is None: - return None - root_path = path - root, attrs = root_path - return root, (*attrs, node.args[1].value) - - -def _classify_kwargs() -> tuple[ - dict[str, tuple[str, tuple[str, ...]]], dict[str, object], set[str] -]: - """Split the call site's kwargs into three buckets: - - - ``attr_kwargs``: ``kwarg=source.attr[...]``, - ``kwarg=int(source.attr)``, or ``kwarg=getattr(source, "attr", ...)`` - → ``{kwarg: (root, path)}``. - - ``literal_kwargs``: ``kwarg=`` → ``{kwarg: value}``. - - ``other_kwargs``: kwargs whose value is anything else (e.g. a bare - Name like ``q``). - """ - call = _parse_thop_attention_call() - attr_kwargs: dict[str, tuple[str, tuple[str, ...]]] = {} - literal_kwargs: dict[str, object] = {} - other_kwargs: set[str] = set() - for kw in call.keywords: - v = kw.value - # int(source.attr) → equivalent to source.attr for sync purposes. - if ( - isinstance(v, ast.Call) - and isinstance(v.func, ast.Name) - and v.func.id == "int" - and len(v.args) == 1 - ): - path = _attribute_path(v.args[0]) - if path is not None: - attr_kwargs[kw.arg] = path - continue - other_kwargs.add(kw.arg) - continue - if isinstance(v, ast.Constant): - literal_kwargs[kw.arg] = v.value - continue - path = _getattr_path(v) - if path is not None: - attr_kwargs[kw.arg] = path - continue - path = _attribute_path(v) - if path is not None: - attr_kwargs[kw.arg] = path - continue - other_kwargs.add(kw.arg) - return attr_kwargs, literal_kwargs, other_kwargs - - -# ---- Python-side attribute & type resolution -------------------------------- - - -def _runtime_instance_attrs(cls) -> set[str]: - """Names assigned as ``self. = ...`` anywhere in ``cls`` or any of - its base classes.""" - cache = _runtime_instance_attrs._cache # type: ignore[attr-defined] - if cls in cache: - return cache[cls] - - def _walk_target(tgt: ast.AST, names: set[str]) -> None: - if isinstance(tgt, (ast.Tuple, ast.List)): - for elt in tgt.elts: - _walk_target(elt, names) - elif ( - isinstance(tgt, ast.Attribute) - and isinstance(tgt.value, ast.Name) - and tgt.value.id == "self" - ): - names.add(tgt.attr) - - names: set[str] = set() - for base in cls.__mro__: - if base is object: - continue - try: - src = textwrap.dedent(inspect.getsource(base)) - except (OSError, TypeError): - continue - for node in ast.walk(ast.parse(src)): - if isinstance(node, ast.Assign): - for tgt in node.targets: - _walk_target(tgt, names) - elif isinstance(node, (ast.AnnAssign, ast.AugAssign)): - _walk_target(node.target, names) - cache[cls] = names - return names - - -_runtime_instance_attrs._cache = {} # type: ignore[attr-defined] - - -def _has_attr_or_field(cls, name: str) -> bool: - if hasattr(cls, name): - return True - try: - if name in {f.name for f in fields(cls)}: - return True - except TypeError: - pass - return name in _runtime_instance_attrs(cls) - - -def _dataclass_field_type(cls, name: str): - try: - f = cls.__dataclass_fields__.get(name) # type: ignore[attr-defined] - except AttributeError: - return None - if f is None: - return None - if isinstance(f.type, str): - return None - return _unwrap_optional(f.type) - - -def _unwrap_optional(py_type): - """Return the payload type for ``Optional[T]`` annotations.""" - origin = typing.get_origin(py_type) - if origin in (typing.Union, UnionType): - args = [arg for arg in typing.get_args(py_type) if arg is not type(None)] - if len(args) == 1: - return args[0] - return py_type - - -def _resolve_path(root_cls, path: tuple[str, ...]): - """Walk ``path[:-1]`` from ``root_cls`` and return ``(leaf_cls, - leaf_attr)``. Returns ``(None, leaf_attr)`` if an intermediate step is - not a dataclass field.""" - cls = root_cls - for step in path[:-1]: - nxt = _dataclass_field_type(cls, step) - if nxt is None: - return None, path[-1] - cls = nxt - return cls, path[-1] - - -def _python_category(py_type) -> str: - """Coarse classification of a Python annotation, mirroring - ``_cpp_category``. Returns ``unknown`` for anything we can't classify - confidently (the type check is then skipped for that kwarg).""" - # Unwrap Optional[X] / Union[X, None]. - origin = typing.get_origin(py_type) - if origin in (typing.Union, UnionType): - args = [a for a in typing.get_args(py_type) if a is not type(None)] - if len(args) == 1: - return _python_category(args[0]) - return "unknown" - if origin in (list, typing.List): - args = typing.get_args(py_type) - if args: - return f"list_{_python_category(args[0])}" - return "list_unknown" - if py_type is bool: - return "bool" - if py_type is int: - return "int" - if py_type is float: - return "float" - if isinstance(py_type, type): - if py_type.__name__ == "Tensor": - return "tensor" - return "unknown" - - -# ---- Tests ------------------------------------------------------------------ - - -def test_call_site_kwargs_match_binding_kwargs(): - """The call site's keyword set must equal the C++ binding's keyword set.""" - cpp_kwargs = _binding_kwargs() - attr_kwargs, literal_kwargs, other_kwargs = _classify_kwargs() - call_site_kwargs = set(attr_kwargs) | set(literal_kwargs) | other_kwargs - assert call_site_kwargs == cpp_kwargs, ( - f"missing in call site: {cpp_kwargs - call_site_kwargs}, " - f"unknown to C++ binding: {call_site_kwargs - cpp_kwargs}" - ) - - -def test_each_source_attr_kwarg_resolves_uniquely(): - """``root.attr[.attr...]`` kwargs must resolve along the chain to an - attribute that exists on exactly one of the source classes, and the - leaf's Python-category must match the C++ kwarg's category (where - both are unambiguously known).""" - attr_kwargs, _, _ = _classify_kwargs() - cpp_types = _binding_types() - for thop_kwarg, (root, path) in attr_kwargs.items(): - assert root in _SOURCE_CLASSES, ( - f"thop kwarg `{thop_kwarg}` sources from unknown root `{root}` — " - f"call site must use one of {set(_SOURCE_CLASSES)}." - ) - leaf_cls, leaf_attr = _resolve_path(_SOURCE_CLASSES[root], path) - chain = ".".join((root, *path)) - assert leaf_cls is not None, ( - f"thop kwarg `{thop_kwarg}` reads `{chain}` but an intermediate " - f"step is not a dataclass field of its parent." - ) - assert _has_attr_or_field(leaf_cls, leaf_attr), ( - f"thop kwarg `{thop_kwarg}` reads `{chain}` but `{leaf_attr}` " - f"is not exposed on {leaf_cls.__name__}." - ) - # Coarse type cross-check (skipped when either side is unknown). - cpp_cat = _cpp_category(cpp_types[thop_kwarg]) - leaf_py_type = _dataclass_field_type(leaf_cls, leaf_attr) - if cpp_cat != "unknown" and leaf_py_type is not None: - py_cat = _python_category(leaf_py_type) - if py_cat != "unknown": - assert cpp_cat == py_cat, ( - f"thop kwarg `{thop_kwarg}` (C++ {cpp_types[thop_kwarg]!r}" - f" → {cpp_cat}) bound to `{chain}` " - f"({leaf_py_type!r} → {py_cat})." - ) - - -def test_attr_kwarg_names_match_source_leaf_attrs_except_allowlisted_aliases(): - """Most ``thop.attention`` kwargs should bind to a source attribute with - the same name. Existing aliases must stay explicit so new semantic - mismatches cannot slip in under a broad type-compatible mapping. - """ - attr_kwargs, _, _ = _classify_kwargs() - aliases = {kwarg: source for kwarg, source in attr_kwargs.items() if kwarg != source[1][-1]} - assert aliases == _THOP_KWARG_SOURCE_ALIASES, ( - "Unexpected thop kwarg/source attribute aliases.\n" - f"new or changed aliases: {aliases.items() - _THOP_KWARG_SOURCE_ALIASES.items()}\n" - f"stale allowlist entries: {_THOP_KWARG_SOURCE_ALIASES.items() - aliases.items()}" - ) - - -def test_literal_kwargs_match_allowlist(): - """Every literal-constant kwarg at the call site must appear in - ``_THOP_LITERALS`` with the matching value, and every entry in - ``_THOP_LITERALS`` must be used at the call site.""" - _, literal_kwargs, _ = _classify_kwargs() - unknown = set(literal_kwargs) - set(_THOP_LITERALS) - assert not unknown, ( - f"kwargs passed as literals but not in _THOP_LITERALS: " - f"{sorted(unknown)}. Source from one of {set(_SOURCE_CLASSES)} or " - f"add to the allowlist." - ) - stale = set(_THOP_LITERALS) - set(literal_kwargs) - assert not stale, ( - f"_THOP_LITERALS entries no longer passed as literals: " - f"{sorted(stale)}. Drop them or restore the call-site literal." - ) - for kwarg, value in literal_kwargs.items(): - assert value == _THOP_LITERALS[kwarg], ( - f"thop kwarg `{kwarg}` passed as literal {value!r} but " - f"_THOP_LITERALS expects {_THOP_LITERALS[kwarg]!r}." - ) - - -def _self_attrs_in_property(prop: property) -> set[str]: - """``self.`` names read inside ``prop``'s getter body.""" - src = textwrap.dedent(inspect.getsource(prop.fget)) - return { - node.attr - for node in ast.walk(ast.parse(src)) - if isinstance(node, ast.Attribute) - and isinstance(node.value, ast.Name) - and node.value.id == "self" - } - - -def _collect_chains(root: str) -> set[tuple[str, ...]]: - """All attribute paths in ``FallbackFmha.forward`` that start with - ``Name(root).``.""" - src = textwrap.dedent(inspect.getsource(FallbackFmha.forward)) - chains: set[tuple[str, ...]] = set() - for node in ast.walk(ast.parse(src)): - if not isinstance(node, ast.Attribute): - continue - path: list[str] = [] - cur: ast.AST = node - while isinstance(cur, ast.Attribute): - path.insert(0, cur.attr) - cur = cur.value - if isinstance(cur, ast.Name) and cur.id == root: - chains.add(tuple(path)) - return chains - - -def _verify_consumed(cls, chains: set[tuple[str, ...]], excluded=frozenset()): - """Recursively assert that every field on ``cls`` is consumed by some - chain in ``chains`` (or excluded). For nested dataclass fields, recurse - into the sub-bag with chain tails. Properties on ``cls`` accessed at - the call site transitively consume the fields they read on ``self``. - """ - direct = {p[0] for p in chains if p} - transitive: set[str] = set() - for name in direct: - obj = vars(cls).get(name) - if isinstance(obj, property): - transitive |= _self_attrs_in_property(obj) - consumed = direct | transitive - for f in fields(cls): - if f.name in excluded: - continue - ftype = _dataclass_field_type(cls, f.name) - if ftype is not None and dataclasses.is_dataclass(ftype): - sub = {p[1:] for p in chains if len(p) >= 2 and p[0] == f.name} - assert sub, ( - f"Nested dataclass field `{f.name}` on {cls.__name__} is " - f"declared but `{f.name}.` is never read at the " - f"call site." - ) - _verify_consumed(ftype, sub, excluded=excluded) - else: - assert f.name in consumed, ( - f"Field `{f.name}` on {cls.__name__} not consumed by the " - f"thop call site (directly or via @property). Source it at " - f"the call site or add it to _THOP_EXCLUDED_FIELDS." - ) - - -def test_every_forward_args_field_is_consumed(): - """Recursively check that every dataclass field reachable from - ``AttentionForwardArgs`` (including nested sub-bags such as - ``SparseRuntimeParams``) is consumed at the call site, transitively - via @property where applicable, or listed in ``_THOP_EXCLUDED_FIELDS``. - """ - _verify_consumed( - AttentionForwardArgs, - _collect_chains("forward_args"), - excluded=_THOP_EXCLUDED_FIELDS, - ) - - -def test_no_unexpected_other_kwargs(): - """The only call-site kwargs that aren't ``source.attr`` chains or - allowlisted literals are the ``FallbackFmha.forward`` parameters.""" - _, _, other_kwargs = _classify_kwargs() - expected = {"q", "k", "v"} - unexpected = other_kwargs - expected - assert not unexpected, ( - f"thop kwargs with unexpected expression form: {sorted(unexpected)}. " - f"Allowed: source.attr[.attr...], int(source.attr), literal " - f"constant, or one of {sorted(expected)}." - ) - - -def _all_forward_args_field_names() -> set[str]: - """All dataclass field names reachable from ``AttentionForwardArgs``, - recursively descending into nested dataclass sub-bags.""" - seen: set[str] = set() - - def _walk(cls) -> None: - for f in fields(cls): - seen.add(f.name) - ftype = _dataclass_field_type(cls, f.name) - if ftype is not None and dataclasses.is_dataclass(ftype): - _walk(ftype) - - _walk(AttentionForwardArgs) - return seen - - -def test_excluded_fields_match_real_fields(): - """Every entry in ``_THOP_EXCLUDED_FIELDS`` must name a real field on - ``AttentionForwardArgs`` (or a nested sub-bag). Dead entries (typos, - fields removed in a refactor) silently allow newly-added real fields - to slip past ``test_every_forward_args_field_is_consumed``.""" - stale = set(_THOP_EXCLUDED_FIELDS) - _all_forward_args_field_names() - assert not stale, ( - f"_THOP_EXCLUDED_FIELDS entries that don't match any " - f"AttentionForwardArgs field: {sorted(stale)}. Drop them or " - f"restore the field." - ) - - -def test_no_sequence_kwargs_at_thop_attention_boundary(): - """``thop.attention`` must not accept ``std::vector<...>`` / - ``c10::ArrayRef<...>`` / ``std::array<...>`` parameters. - - Sequence params couple list position to semantic meaning, which the - other sync tests in this file cannot verify element-by-element. Flat - named params let every slot be checked individually (name, type, - source). If a new sequence param creeps back in, flatten it the same - way ``rotary_embedding_scales`` / ``helix_tensor_params`` / - ``spec_decoding_*_params`` were flattened. - """ - sequence_params = [] - for name, cpp_type in _parse_attention_decl(): - bare = cpp_type.strip() - # _strip_inner_type only strips trailing const/&/*; we want to - # detect outer container types regardless of qualifiers, so look - # at the raw type string with leading qualifiers stripped. - outer = re.sub(r"^(const\s+|volatile\s+)+", "", bare) - outer = re.sub(r"\s*(const|&|\*)\s*$", "", outer) - if ( - outer.startswith("std::vector<") - or outer.startswith("c10::ArrayRef<") - or outer.startswith("std::array<") - ): - sequence_params.append((name, cpp_type)) - assert not sequence_params, ( - "thop.attention must not accept Sequence-typed kwargs. Flatten " - "the following params into their named scalar/tensor components:\n" - + "\n".join(f" - {name}: {t}" for name, t in sequence_params) - ) - - -@pytest.mark.parametrize( - ("is_cross", "update_kv_cache", "expected"), - ( - (False, False, False), - (False, True, True), - (True, False, True), - ), -) -def test_fallback_support_matches_thop_kv_update_contract(is_cross, update_kv_cache, expected): - """Do not dispatch requests that the native attention op rejects.""" - fmha = object.__new__(FallbackFmha) - # ``helix_position_offsets`` short-circuits the Helix verify-group check - # that runs first in ``_is_supported``; None is what the real metadata - # carries off the helix path. - metadata = SimpleNamespace(is_cross=is_cross, helix_position_offsets=None) - forward_args = AttentionForwardArgs(update_kv_cache=update_kv_cache) - - assert fmha.is_supported(None, None, None, metadata, forward_args) is expected - - -def test_fallback_rejects_raw_fp8_input(): - """Do not dispatch raw FP8 QKV to the native attention op.""" - fmha = object.__new__(FallbackFmha) - metadata = SimpleNamespace(is_cross=False, helix_position_offsets=None) - forward_args = AttentionForwardArgs(update_kv_cache=True) - q = torch.empty((1, 128), dtype=torch.float8_e4m3fn) - - assert not fmha.is_supported(q, None, None, metadata, forward_args) diff --git a/tests/unittest/_torch/attention/test_combined_fmha.py b/tests/unittest/_torch/attention/test_combined_fmha.py index 87f294744b88..db3bd9853afd 100644 --- a/tests/unittest/_torch/attention/test_combined_fmha.py +++ b/tests/unittest/_torch/attention/test_combined_fmha.py @@ -15,6 +15,7 @@ from types import SimpleNamespace +import pytest import torch from fmha_test_utils import FakeAttention, FakePhasedFmha @@ -27,7 +28,12 @@ from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2, Role -def test_combined_fmha_delegates_phases_and_prepares_max_workspace() -> None: +@pytest.mark.parametrize( + "input_type", [AttentionInputType.mixed, AttentionInputType.generation_only] +) +def test_combined_fmha_delegates_phases_and_prepares_max_workspace( + monkeypatch: pytest.MonkeyPatch, input_type: AttentionInputType +) -> None: events: list[tuple] = [] attn = FakeAttention() context_fmha = FakePhasedFmha( @@ -64,8 +70,8 @@ def test_combined_fmha_delegates_phases_and_prepares_max_workspace() -> None: num_generations=4, kv_lens_cuda_runtime=torch.tensor([2, 1, 5, 5, 5, 5], dtype=torch.int32), kv_lens_runtime=torch.tensor([2, 1, 5, 5, 5, 5], dtype=torch.int32), - prompt_lens_cuda_runtime=torch.tensor([2, 1, 1, 1, 1, 1], dtype=torch.int32), - prompt_lens_cpu_runtime=torch.tensor([2, 1, 1, 1, 1, 1], dtype=torch.int32), + prompt_lens_cuda_runtime=torch.tensor([2, 1, 3, 3, 4, 4], dtype=torch.int32), + prompt_lens_cpu_runtime=torch.tensor([2, 1, 3, 3, 4, 4], dtype=torch.int32), beam_width=2, cache_indirection=None, tokens_per_block=32, @@ -73,20 +79,31 @@ def test_combined_fmha_delegates_phases_and_prepares_max_workspace() -> None: is_cross=False, is_spec_decoding_enabled=False, ) + run_generation = generation_fmha.run_generation + + def check_generation_lengths(params) -> None: + assert params.context_lengths.tolist() == [3, 3, 4, 4] + assert params.context_lengths.data_ptr() == metadata.prompt_lens_cuda_runtime[2:].data_ptr() + run_generation(params) + + monkeypatch.setattr(generation_fmha, "run_generation", check_generation_lengths) + num_tokens = 4 if input_type == AttentionInputType.generation_only else 7 forward_args = AttentionForwardArgs( - output=torch.empty((7, 4)), - attention_input_type=AttentionInputType.mixed, + output=torch.empty((num_tokens, 4)), + attention_input_type=input_type, attention_window_size=8, ) - combined_fmha.forward(torch.empty((7, 4)), None, None, metadata, forward_args) + combined_fmha.forward(torch.empty((num_tokens, 4)), None, None, metadata, forward_args) - assert events == [ + expected_events = [ ("prepare", "context"), ("prepare", "generation"), - ("run", "context", FmhaPhase.CONTEXT, 3, 2, 2), - ("run", "generation", FmhaPhase.GENERATION, 4, 4, 2), ] + if input_type == AttentionInputType.mixed: + expected_events.append(("run", "context", FmhaPhase.CONTEXT, 3, 2, 2)) + expected_events.append(("run", "generation", FmhaPhase.GENERATION, 4, 4, 2)) + assert events == expected_events assert metadata.effective_workspace.numel() == 8 diff --git a/tests/unittest/_torch/attention/test_context_fmha_kernel_presence.py b/tests/unittest/_torch/attention/test_context_fmha_kernel_presence.py index a9e1c62b1e0e..ba1451d02e5e 100644 --- a/tests/unittest/_torch/attention/test_context_fmha_kernel_presence.py +++ b/tests/unittest/_torch/attention/test_context_fmha_kernel_presence.py @@ -11,7 +11,7 @@ path, which attends to the current chunk only and then overwrites the cached prefix. -``get_attention_op`` refuses that combination after initialization. This module +The ``AttentionOp`` constructor refuses that combination after initialization. This module covers the ``fused_context_fmha_kernel_exists`` diagnostic that reports what the running build actually contains, including the case that matters most: that the lookup is able to answer "no". Every test here calls into the native kernel @@ -22,9 +22,11 @@ import pytest import torch +from tensorrt_llm._torch.attention.backends.fmha.interface import StaticAttentionConfig from tensorrt_llm._utils import get_sm_version from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal import thop +from tensorrt_llm.functional import PositionEmbeddingType # The paged-KV context FMHA kernels are generated for 32-token pages. _TOKENS_PER_BLOCK = 32 @@ -46,6 +48,31 @@ def _is_sm100_family() -> bool: return 100 <= get_sm_version() < 110 +@pytest.mark.parametrize( + "head_size,position_embedding_type,error", + [ + (96, PositionEmbeddingType.learned_absolute, "requires a fused context FMHA kernel"), + (64, PositionEmbeddingType.relative, "not supported with relative position embedding"), + ], + ids=["missing-kernel", "relative-position"], +) +def test_attention_op_rejects_unfused_paged_context(head_size, position_embedding_type, error): + if not _is_sm100_family(): + pytest.skip("the unsupported head-size case targets the SM100-family dispatcher") + config = StaticAttentionConfig( + num_heads=1, + num_kv_heads=1, + head_size=head_size, + tokens_per_block=_TOKENS_PER_BLOCK, + type=DataType.BF16, + use_kv_cache=True, + paged_context_fmha=True, + position_embedding_type=position_embedding_type, + ) + with pytest.raises(RuntimeError, match=error): + thop.AttentionOp(config.to_thop_config()) + + @pytest.mark.parametrize( "head_size,tokens_per_block", [(0, 32), (-1, 32), (64, 0), (64, -1)], diff --git a/tests/unittest/_torch/attention/test_flashinfer_trtllm_gen_fmha.py b/tests/unittest/_torch/attention/test_flashinfer_trtllm_gen_fmha.py index bbe1dedf50ca..43cbfcbfe161 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_trtllm_gen_fmha.py +++ b/tests/unittest/_torch/attention/test_flashinfer_trtllm_gen_fmha.py @@ -21,6 +21,7 @@ from tensorrt_llm._torch.attention.backends.fmha.flashinfer_trtllm_gen import ( FlashInferTrtllmGenFmha, + _get_generation_workspace_size, ) from tensorrt_llm._torch.attention.backends.fmha.interface import FmhaPhase from tensorrt_llm._torch.attention.backends.interface import ( @@ -31,6 +32,48 @@ from tensorrt_llm.quantization.mode import QuantMode +@pytest.mark.parametrize("max_num_sequences", [None, 64]) +@pytest.mark.parametrize("warmup_tokens", [1, 64]) +def test_flashinfer_workspace_covers_sequence_capacity( + monkeypatch: pytest.MonkeyPatch, max_num_sequences: int | None, warmup_tokens: int +) -> None: + attn = FakeAttention() + attn.num_heads = attn.num_kv_heads = 16 + attn.head_dim = 64 + attn.rope_dim = 0 + attn.quant_mode = 0 + fmha = FlashInferTrtllmGenFmha(attn) + monkeypatch.setattr(fmha, "_get_multi_processor_count", lambda _: 148) + metadata = SimpleNamespace( + max_num_requests=16, + max_num_sequences=max_num_sequences, + max_context_length=1, + num_ctx_tokens=0, + is_cuda_graph=False, + ) + workspace = torch.empty(0, dtype=torch.uint8) + q = torch.empty((warmup_tokens, 16 * 64), dtype=torch.bfloat16) + forward_args = AttentionForwardArgs( + output=torch.empty_like(q), attention_input_type=AttentionInputType.generation_only + ) + fmha.prepare_workspace(q, None, None, metadata, forward_args, workspace) + + num_sequences = max_num_sequences or metadata.max_num_requests + required = _get_generation_workspace_size( + torch.bfloat16, num_sequences, num_sequences, 16, 64, 16, 0 + ) + assert workspace.nbytes >= required + + # Capture must reuse the allocation reserved by a smaller warmup batch. + q = torch.empty((num_sequences, 16 * 64), dtype=torch.bfloat16) + forward_args.output = torch.empty_like(q) + workspace_ptr = workspace.data_ptr() + metadata.is_cuda_graph = True + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + fmha.prepare_workspace(q, None, None, metadata, forward_args, workspace) + assert workspace.data_ptr() == workspace_ptr + + def test_flashinfer_fp8_mode_remains_implementation_local() -> None: attn = FakeAttention() attn.quant_mode = int(QuantMode.from_description(use_fp8_kv_cache=True)) diff --git a/tests/unittest/_torch/attention/test_fmha_interface.py b/tests/unittest/_torch/attention/test_fmha_interface.py index a265c3a69037..e0f5d4c1bb1f 100644 --- a/tests/unittest/_torch/attention/test_fmha_interface.py +++ b/tests/unittest/_torch/attention/test_fmha_interface.py @@ -13,6 +13,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +import inspect +from types import SimpleNamespace from typing import cast from unittest.mock import Mock, patch @@ -22,6 +24,7 @@ from tensorrt_llm._torch.attention.backends.fmha.fallback import FallbackFmha from tensorrt_llm._torch.attention.backends.fmha.interface import Fmha, FmhaPhase +from tensorrt_llm._torch.attention.backends.fmha.phased import PhasedFmha from tensorrt_llm._torch.attention.backends.fmha.registry import FMHA_LIBS from tensorrt_llm._torch.attention.backends.interface import ( AttentionForwardArgs, @@ -32,6 +35,8 @@ SparseRuntimeParams, ) from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention, TrtllmAttentionMetadata +from tensorrt_llm.bindings.internal import thop +from tensorrt_llm.functional import AttentionMaskType class _MinimalFmha(Fmha): @@ -46,6 +51,502 @@ def forward( pass +@pytest.mark.parametrize( + "input_type,num_contexts", + [ + (AttentionInputType.context_only, 1), + (AttentionInputType.generation_only, 0), + (AttentionInputType.generation_only, 1), + (AttentionInputType.mixed, 0), + (AttentionInputType.mixed, 1), + (None, 0), + (None, 1), + ], +) +@pytest.mark.parametrize( + "is_mla,is_cross,beam_width,fused_epilogue", + [ + (False, False, 1, False), + (True, False, 1, False), + (False, False, 4, False), + (False, True, 4, False), + (True, False, 1, True), + ], + ids=["mha", "mla", "self-beams", "cross-beams", "dsv4-epilogue"], +) +@pytest.mark.parametrize("use_spec_decoding", [False, True]) +@pytest.mark.parametrize("is_spec_decoding_enabled", [False, True]) +def test_legacy_fallback_dispatch_uses_native_phase_params( + monkeypatch: pytest.MonkeyPatch, + input_type: AttentionInputType | None, + num_contexts: int, + is_mla: bool, + is_cross: bool, + beam_width: int, + fused_epilogue: bool, + use_spec_decoding: bool, + is_spec_decoding_enabled: bool, +) -> None: + has_context = num_contexts > 0 and input_type != AttentionInputType.generation_only + num_generations = 0 if input_type == AttentionInputType.context_only else 4 + num_ctx_tokens = 3 if has_context else 0 + num_tokens = num_ctx_tokens + num_generations + prompts = torch.tensor([3] * num_contexts + [9] * num_generations, dtype=torch.int32) + kv_lengths = torch.tensor([5] * num_contexts + [10] * num_generations, dtype=torch.int32) + past_lengths = kv_lengths.clone() + if is_mla: + past_lengths[:num_contexts].zero_() + total_kv_lengths = torch.tensor([5 * num_contexts, 10 * num_generations], dtype=torch.int32) + head_dim = 6 if is_mla else 4 + q = torch.empty((num_tokens, head_dim if is_mla else 3 * head_dim), dtype=torch.bfloat16) + output = torch.empty((num_tokens, 4), dtype=torch.bfloat16) + output_sf = None + if fused_epilogue: + head_dim = 512 + q = torch.empty((num_tokens, 4 * head_dim), dtype=torch.bfloat16) + output = torch.empty((4, num_tokens, head_dim), dtype=torch.float8_e4m3fn) + output_sf = torch.empty( + (4, head_dim // 128, (num_tokens + 3) // 4 * 4), dtype=torch.float32 + ) + counter = torch.zeros(64, dtype=torch.uint8) + spec_lengths = torch.ones(4, dtype=torch.int32) + spec_offsets = torch.zeros((4, 1), dtype=torch.int32) + spec_mask = torch.zeros((4, 1), dtype=torch.int32) + spec_tree_offsets = torch.zeros(4, dtype=torch.int64) + spec_tree_mask = torch.zeros(4, dtype=torch.uint32) + spec_sparse_offsets = torch.zeros(4, dtype=torch.int32) + kv_norm_weight = torch.ones(4, dtype=torch.bfloat16) + calls: list[tuple] = [] + phase_params: list[thop.FmhaParams] = [] + + def record(phase: str, params: thop.FmhaParams) -> int: + assert isinstance(params, thop.FmhaParams) + assert params.max_num_sequences == 8 + assert params.host_context_lengths.data_ptr() == prompts.data_ptr() + assert params.host_past_key_value_lengths.data_ptr() == past_lengths.data_ptr() + assert params.host_total_kv_lens.data_ptr() == total_kv_lengths.data_ptr() + assert params.fwd.kv_norm_weight.data_ptr() == kv_norm_weight.data_ptr() + assert params.fwd.kv_norm_eps == pytest.approx(1e-5) + # The flat compatibility entry point delegates activation to native prepare(). + for native_tensor, source_tensor in ( + (params.spec_decoding_generation_lengths, spec_lengths), + (params.spec_decoding_position_offsets, spec_offsets), + (params.spec_decoding_packed_mask, spec_mask), + (params.spec_decoding_bl_tree_mask_offset, spec_tree_offsets), + (params.spec_decoding_bl_tree_mask, spec_tree_mask), + (params.spec_bl_tree_first_sparse_mask_offset_kv, spec_sparse_offsets), + ): + assert native_tensor.data_ptr() == source_tensor.data_ptr() + if fused_epilogue: + assert params.fwd.enable_dsv4_epilogue_fusion + for actual, expected in ((params.output, output), (params.fwd.output_sf, output_sf)): + assert actual.shape == expected.shape + assert actual.stride() == expected.stride() + assert actual.data_ptr() == expected.data_ptr() + if phase != "workspace": + phase_params.append(params) + offset = 0 if phase == "context" else num_contexts + assert ( + params.context_lengths.tolist() + == prompts[offset : offset + params.num_seqs].tolist() + ) + assert params.qkv_or_q.shape[0] == params.num_tokens + if not fused_epilogue: + assert params.output.shape[0] == params.num_tokens + calls.append( + ( + phase, + params.num_seqs, + params.num_tokens, + params.seq_offset, + params.token_offset, + params.beam_width, + params.num_requests, + ) + ) + return 0 + + op = SimpleNamespace( + get_attention_workspace_size=lambda params, *args: record("workspace", params), + run_context=lambda params: record("context", params), + run_generation=lambda params: record("generation", params), + run_mla_generation=lambda params: record("mla_generation", params), + ) + + def get_op(config, device): + assert config.skip_correction_threshold == pytest.approx(0.25) + assert device == q.device + return op + + monkeypatch.setattr(FallbackFmha, "_compat_attention_ops", SimpleNamespace(get=get_op)) + monkeypatch.setattr( + "tensorrt_llm._torch.attention.backends.fmha.fallback.get_multi_ctas_kv_counter", + lambda *args: counter, + ) + arguments = { + name: param.default if param.default is not inspect.Parameter.empty else None + for name, param in inspect.signature(FallbackFmha.attention).parameters.items() + } + arguments.update( + q=q, + output=output, + output_sf=output_sf, + enable_dsv4_epilogue_fusion=fused_epilogue, + workspace=torch.empty(0, dtype=torch.uint8), + sequence_length=kv_lengths, + context_lengths=prompts, + host_context_lengths=prompts, + host_past_key_value_lengths=past_lengths, + host_total_kv_lens=total_kv_lengths, + num_heads=4 if fused_epilogue else 1, + num_kv_heads=1, + head_size=head_dim, + tokens_per_block=64, + max_num_requests=8, + max_context_length=16, + max_seq_len=16, + attention_window_size=16, + beam_width=beam_width, + mask_type=AttentionMaskType.causal, + quant_mode=0, + q_scaling=1.0, + position_embedding_type=0, + local_layer_idx=0, + rope_dim=0, + rope_base=10000.0, + rope_scale_type=0, + rope_scale=1.0, + rope_short_m_scale=1.0, + rope_long_m_scale=1.0, + rope_max_positions=16, + rope_original_max_positions=16, + predicted_tokens_per_seq=1, + is_spec_decoding_enabled=is_spec_decoding_enabled, + use_spec_decoding=use_spec_decoding, + spec_decoding_generation_lengths=spec_lengths, + spec_decoding_position_offsets=spec_offsets, + spec_decoding_packed_mask=spec_mask, + spec_decoding_bl_tree_mask_offset=spec_tree_offsets, + spec_decoding_bl_tree_mask=spec_tree_mask, + spec_bl_tree_first_sparse_mask_offset_kv=spec_sparse_offsets, + is_fused_qkv=not is_mla, + update_kv_cache=True, + use_paged_context_fmha=False, + is_mla_enable=is_mla, + is_cross=is_cross, + attention_input_type=input_type, + kv_lora_rank=4 if is_mla else None, + qk_rope_head_dim=2 if is_mla else None, + qk_nope_head_dim=4 if is_mla else None, + v_head_dim=4 if is_mla else None, + rope_append=True if is_mla else None, + num_contexts=num_contexts, + num_ctx_tokens=3 * num_contexts, + sparse_attn_indices_block_size=1, + kv_norm_weight=kv_norm_weight, + kv_norm_eps=1e-5, + skip_correction_threshold=0.25, + ) + if fused_epilogue: + arguments.update(kv_lora_rank=448, qk_rope_head_dim=64, v_head_dim=512, rope_append=False) + if has_context and num_generations: + with pytest.raises(ValueError, match="DSv4 fused epilogue requires separate"): + FallbackFmha.attention(**arguments) + assert not calls + return + FallbackFmha.attention(**arguments) + + effective_beams = 1 if is_cross else beam_width + sizing_contexts = num_contexts if has_context else 0 + expected = [ + ("workspace", sizing_contexts, num_ctx_tokens, 0, 0, effective_beams, sizing_contexts) + ] + if has_context: + expected.append( + ("context", num_contexts, num_ctx_tokens, 0, 0, effective_beams, num_contexts) + ) + if num_generations: + phase = "mla_generation" if is_mla else "generation" + expected.append( + (phase, 4, 4, num_contexts, num_ctx_tokens, effective_beams, 4 // effective_beams) + ) + assert calls == expected + if len(phase_params) == 2: + assert phase_params[0] is not phase_params[1] + + +@pytest.mark.parametrize("is_cross", [False, True]) +@pytest.mark.parametrize("beam_width", [1, 4]) +@pytest.mark.parametrize( + "input_type", [AttentionInputType.mixed, AttentionInputType.generation_only] +) +@pytest.mark.parametrize("use_spec_decoding", [False, True]) +@pytest.mark.parametrize("is_spec_decoding_enabled", [False, True]) +@pytest.mark.parametrize("predicted_tokens_per_seq", [1, 4]) +@pytest.mark.parametrize("output_dtype", [torch.bfloat16, torch.uint8]) +def test_fallback_native_generation_lengths_and_beams( + monkeypatch: pytest.MonkeyPatch, + is_cross: bool, + beam_width: int, + input_type: AttentionInputType, + use_spec_decoding: bool, + is_spec_decoding_enabled: bool, + predicted_tokens_per_seq: int, + output_dtype: torch.dtype, +) -> None: + monkeypatch.setattr(TrtllmAttentionMetadata, "_post_init_with_buffers", lambda *args: None) + metadata = TrtllmAttentionMetadata( + max_num_requests=8, + max_num_sequences=8, + max_num_tokens=7, + num_contexts=1, + beam_width=beam_width, + workspace=torch.empty(0, dtype=torch.uint8), + ) + metadata.max_seq_len = 16 + metadata._seq_lens = torch.tensor([3, 1, 1, 1, 1], dtype=torch.int32) + if is_cross: + metadata._seq_lens_kv = torch.tensor([5, 10, 10, 10, 10], dtype=torch.int32) + metadata.num_generations = 4 + metadata._num_ctx_tokens = 3 + metadata.prompt_lens_cuda_runtime = torch.tensor([3, 9, 9, 9, 9], dtype=torch.int32) + metadata.prompt_lens_cpu_runtime = metadata.prompt_lens_cuda_runtime + metadata.kv_lens_runtime = torch.tensor([5, 10, 10, 10, 10], dtype=torch.int32) + metadata.kv_lens_cuda_runtime = metadata.kv_lens_runtime + metadata.is_spec_decoding_enabled = is_spec_decoding_enabled + metadata.use_spec_decoding = use_spec_decoding + metadata.spec_decoding_generation_lengths = torch.full((4,), 4, dtype=torch.int32) + metadata.spec_decoding_position_offsets = torch.zeros((8, 4), dtype=torch.int32) + metadata.spec_decoding_packed_mask = torch.zeros((4,), dtype=torch.int32) + metadata.spec_decoding_bl_tree_mask_offset = torch.zeros((4,), dtype=torch.int64) + metadata.spec_decoding_bl_tree_mask = torch.zeros((4,), dtype=torch.uint32) + metadata.spec_bl_tree_first_sparse_mask_offset_kv = torch.zeros((4,), dtype=torch.int32) + spec_tensors = { + name: getattr(metadata, name) + for name in ( + "spec_decoding_generation_lengths", + "spec_decoding_position_offsets", + "spec_decoding_packed_mask", + "spec_decoding_bl_tree_mask_offset", + "spec_decoding_bl_tree_mask", + "spec_bl_tree_first_sparse_mask_offset_kv", + ) + } + attn = FakeAttention() + attn.predicted_tokens_per_seq = predicted_tokens_per_seq + attn.layer_idx = 0 + attn.get_local_layer_idx = lambda _: 0 + attn.attention_chunk_size = None + attn.rotary_inv_freq = attn.rotary_cos_sin = attn.rope_params = None + fmha = FallbackFmha(attn) + native_calls = [] + context_calls = [] + op = SimpleNamespace(run_context=context_calls.append, run_generation=native_calls.append) + monkeypatch.setattr(fmha, "prepare_workspace", lambda *args: None) + monkeypatch.setattr(fmha, "attention_op", lambda params: op) + spec_active = is_spec_decoding_enabled and use_spec_decoding + num_gen_tokens = 16 if spec_active else 4 + num_tokens = num_gen_tokens + (3 if input_type == AttentionInputType.mixed else 0) + q = torch.empty((num_tokens, 12), dtype=torch.bfloat16) + stored_head_size = 2 if output_dtype == torch.uint8 else 4 + forward_args = AttentionForwardArgs( + output=torch.empty((num_tokens, stored_head_size), dtype=output_dtype), + output_sf=torch.empty(16, dtype=torch.uint8) if output_dtype == torch.uint8 else None, + attention_input_type=input_type, + is_fused_qkv=True, + ) + + fmha.forward(q, None, None, metadata, forward_args) + + assert len(context_calls) == int(input_type == AttentionInputType.mixed) + for context_params in context_calls: + for name in spec_tensors: + assert getattr(context_params, name) is None + assert len(native_calls) == 1 + params = native_calls[0] + assert params.context_lengths.tolist() == [9, 9, 9, 9] + assert params.context_lengths.data_ptr() == metadata.prompt_lens_cuda_runtime[1:].data_ptr() + assert params.beam_width == (1 if is_cross else beam_width) + assert params.num_requests == (4 if is_cross else 4 // beam_width) + assert params.num_seqs == params.num_requests * params.beam_width == 4 + assert metadata.beam_width == beam_width + assert params.output.shape == (num_gen_tokens, 1, stored_head_size) + token_offset = 3 if input_type == AttentionInputType.mixed else 0 + assert params.token_offset == token_offset + assert params.output.data_ptr() == forward_args.output[token_offset:].data_ptr() + if output_dtype == torch.uint8: + assert params.fwd.output_sf.data_ptr() == forward_args.output_sf.data_ptr() + for name, source_tensor in spec_tensors.items(): + native_tensor = getattr(params, name) + if spec_active: + assert native_tensor.data_ptr() == source_tensor.data_ptr() + else: + assert native_tensor is None + assert getattr(metadata, name).data_ptr() == source_tensor.data_ptr() + + +@pytest.mark.parametrize( + "input_type,num_tokens", + [ + (AttentionInputType.context_only, 1), + (AttentionInputType.context_only, 8192), + (AttentionInputType.generation_only, 1), + (AttentionInputType.generation_only, 7), + (AttentionInputType.mixed, 5), + ], +) +def test_fallback_preserves_dsv4_epilogue_output_layout( + monkeypatch: pytest.MonkeyPatch, input_type: AttentionInputType, num_tokens: int +) -> None: + monkeypatch.setattr(TrtllmAttentionMetadata, "_post_init_with_buffers", lambda *args: None) + num_ctx_tokens = num_tokens if input_type == AttentionInputType.context_only else 3 + num_generations = num_tokens if input_type == AttentionInputType.generation_only else 2 + metadata = TrtllmAttentionMetadata( + max_num_requests=8, + max_num_tokens=num_ctx_tokens + num_generations, + num_contexts=1, + workspace=torch.empty(0, dtype=torch.uint8), + ) + metadata.max_seq_len = max(num_ctx_tokens, 16) + metadata._seq_lens = torch.tensor([num_ctx_tokens] + [1] * num_generations, dtype=torch.int32) + metadata._num_ctx_tokens = num_ctx_tokens + metadata.num_generations = num_generations + metadata.prompt_lens_cpu_runtime = torch.tensor( + [num_ctx_tokens] + [9] * num_generations, dtype=torch.int32 + ) + metadata.prompt_lens_cuda_runtime = metadata.prompt_lens_cpu_runtime + metadata.kv_lens_runtime = metadata.prompt_lens_cpu_runtime + metadata.kv_lens_cuda_runtime = metadata.kv_lens_runtime + + attn = FakeAttention() + attn.is_mla_enable = True + attn.num_heads = 8 + attn.head_dim = attn.v_head_dim = 512 + attn.kv_lora_rank = 448 + attn.qk_rope_head_dim = 64 + attn.rope_append = False + attn.layer_idx = 0 + attn.get_local_layer_idx = lambda _: 0 + attn.attention_chunk_size = None + attn.rotary_inv_freq = attn.rotary_cos_sin = attn.rope_params = None + fmha = FallbackFmha(attn) + + groups = 4 + heads_per_group = attn.num_heads // groups + q = torch.empty((num_tokens, attn.num_heads * attn.head_dim), dtype=torch.bfloat16) + output = torch.empty( + (groups, num_tokens, heads_per_group * attn.v_head_dim), dtype=torch.float8_e4m3fn + ) + output_sf = torch.empty( + (groups, heads_per_group * (attn.v_head_dim // 128), (num_tokens + 3) // 4 * 4), + dtype=torch.float32, + ) + forward_args = AttentionForwardArgs( + output=output, + output_sf=output_sf, + attention_input_type=input_type, + enable_dsv4_epilogue_fusion=True, + ) + calls = [] + + def record(phase: str, params: thop.FmhaParams) -> int: + calls.append((phase, params)) + return 0 + + op = SimpleNamespace( + get_attention_workspace_size=lambda params, *args: record("workspace", params), + run_context=lambda params: record("context", params), + run_mla_generation=lambda params: record("generation", params), + ) + monkeypatch.setattr(fmha, "attention_op", lambda params: op) + monkeypatch.setattr( + "tensorrt_llm._torch.attention.backends.fmha.fallback.get_multi_ctas_kv_counter", + lambda *args: torch.zeros(64, dtype=torch.uint8), + ) + + if input_type == AttentionInputType.mixed: + with pytest.raises(ValueError, match="DSv4 fused epilogue requires separate"): + fmha.forward(q, None, None, metadata, forward_args) + assert not calls + return + + fmha.forward(q, None, None, metadata, forward_args) + + phase = "context" if input_type == AttentionInputType.context_only else "generation" + assert [name for name, _ in calls] == ["workspace", phase] + for _, params in calls: + assert params.fwd.enable_dsv4_epilogue_fusion + for actual, expected in ((params.output, output), (params.fwd.output_sf, output_sf)): + assert actual.shape == expected.shape + assert actual.stride() == expected.stride() + assert actual.data_ptr() == expected.data_ptr() + assert calls[-1][1].num_tokens == num_tokens + assert calls[-1][1].seq_offset == int(input_type == AttentionInputType.generation_only) + + +@pytest.mark.parametrize("sm", [90, 100, 103, 120]) +def test_phased_position_offsets_view_tracks_query_width_without_copy( + monkeypatch: pytest.MonkeyPatch, sm: int +) -> None: + monkeypatch.setattr(TrtllmAttentionMetadata, "_post_init_with_buffers", lambda *args: None) + monkeypatch.setattr( + "tensorrt_llm._torch.attention.backends.fmha.phased.get_sm_version", lambda: sm + ) + metadata = TrtllmAttentionMetadata( + max_num_requests=4, + max_num_tokens=32, + num_contexts=0, + workspace=torch.empty(0, dtype=torch.uint8), + ) + metadata.max_seq_len = 32 + metadata.num_generations = 2 + metadata._num_ctx_tokens = 0 + metadata._seq_lens = torch.ones(2, dtype=torch.int32) + metadata.prompt_lens_cuda_runtime = torch.ones(2, dtype=torch.int32) + metadata.prompt_lens_cpu_runtime = metadata.prompt_lens_cuda_runtime + metadata.kv_lens_runtime = torch.full((2,), 16, dtype=torch.int32) + metadata.kv_lens_cuda_runtime = metadata.kv_lens_runtime + metadata.is_spec_decoding_enabled = metadata.use_spec_decoding = True + metadata.spec_decoding_generation_lengths = torch.ones(4, dtype=torch.int32) + offsets_buffer = torch.arange(32, dtype=torch.int32) + metadata.spec_decoding_position_offsets = offsets_buffer + attn = FakeAttention() + fmha = PhasedFmha(attn) + monkeypatch.setattr(fmha, "REQUIRES_PAGED_KV", False) + monkeypatch.setattr(fmha, "NEEDS_BLOCK_EXTENT", False) + views = [] + monkeypatch.setattr( + fmha, "run_generation", lambda params: views.append(params.spec_decoding_position_offsets) + ) + + for query_len in (3, 5, 3): + metadata.spec_decoding_query_len = query_len + metadata.spec_decoding_generation_lengths.fill_(query_len) + metadata._seq_lens.fill_(query_len) + q = torch.empty((2 * query_len, 12), dtype=torch.bfloat16) + fmha.forward( + q, + None, + None, + metadata, + AttentionForwardArgs( + output=torch.empty((2 * query_len, 4), dtype=torch.bfloat16), + attention_input_type=AttentionInputType.generation_only, + is_fused_qkv=True, + ), + ) + + view = views[-1] + expected_width = 8 if sm in (100, 103) else query_len + assert view.shape == (4, expected_width) + assert view.data_ptr() == offsets_buffer.data_ptr() + assert view.untyped_storage().data_ptr() == offsets_buffer.untyped_storage().data_ptr() + offsets_buffer[expected_width] += 100 + assert view[1, 0] == offsets_buffer[expected_width] + assert metadata.spec_decoding_position_offsets is offsets_buffer + assert offsets_buffer.shape == (32,) + + @pytest.mark.parametrize("initial_fused_qkv", [False, True]) @pytest.mark.parametrize( "input_type,sparse", diff --git a/tests/unittest/_torch/attention/test_fmha_manager.py b/tests/unittest/_torch/attention/test_fmha_manager.py index 7cfb59eee093..194076a97b05 100644 --- a/tests/unittest/_torch/attention/test_fmha_manager.py +++ b/tests/unittest/_torch/attention/test_fmha_manager.py @@ -940,6 +940,9 @@ class _OldFmha(FakeFmha): def __init__(self, attn: TrtllmAttention) -> None: super().__init__(attn, "old", events) + def release(self) -> None: + events.append(("release", "old")) + class _NewFmha(FakeFmha): def __init__(self, attn: TrtllmAttention) -> None: quant_states_during_construction.append(attn.has_fp8_kv_cache) @@ -948,6 +951,7 @@ def __init__(self, attn: TrtllmAttention) -> None: attn = TrtllmAttention.__new__(TrtllmAttention) attn.is_mla_enable = False attn.skip_correction_threshold = 0.0 + attn._fmha_manager = None metadata = _make_metadata(num_contexts=0, num_generations=1) forward_args = AttentionForwardArgs(attention_input_type=AttentionInputType.generation_only) q = torch.empty((1, 4)) @@ -973,4 +977,26 @@ def __init__(self, attn: TrtllmAttention) -> None: assert attn._fmha_manager._cache == {} assert isinstance(attn._fmha_manager.fmha_libs[0], _NewFmha) assert quant_states_during_construction == [True] - assert events == [("support", "old", None)] + assert events == [("support", "old", None), ("release", "old")] + + +def test_a_new_manager_does_not_reuse_a_previous_selection() -> None: + """The selection cache belongs to the manager that owns the libs. + + Rebuilding the libs means building a new manager, so a stale entry cannot survive + into a different lib set. A shared or module-level cache would break that. + """ + events: list[tuple] = [] + forward_args = AttentionForwardArgs(attention_input_type=AttentionInputType.generation_only) + q = torch.empty((1, 4)) + metadata = _make_metadata(num_contexts=0, num_generations=1) + + selected = [] + for name in ("old", "new"): + attn, manager = _make_manager() + manager.fmha_libs = [FakeFmha(attn, name, events)] + with patch.object(fmha_manager, "_is_fmha_cache_enabled", return_value=True): + selected.append(manager.select(attn, q, None, None, metadata, forward_args)) + + assert [fmha._name for fmha in selected] == ["old", "new"] + assert events == [("support", "old", None), ("support", "new", None)] diff --git a/tests/unittest/_torch/attention/test_fmha_page_index.py b/tests/unittest/_torch/attention/test_fmha_page_index.py index 9edb2a2e087c..35e68aed7b60 100644 --- a/tests/unittest/_torch/attention/test_fmha_page_index.py +++ b/tests/unittest/_torch/attention/test_fmha_page_index.py @@ -12,9 +12,9 @@ from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import CuteDslMlaFmha from tensorrt_llm._torch.attention.backends.fmha.flashinfer_trtllm_gen import ( FlashInferTrtllmGenFmha, - _get_multi_ctas_kv_counter_size, ) from tensorrt_llm._torch.attention.backends.fmha.phased import FmhaParams +from tensorrt_llm._torch.attention.backends.fmha.utils import get_multi_ctas_kv_counter_size from tensorrt_llm._torch.attention.backends.interface import ( AttentionForwardArgs, AttentionInputType, @@ -102,13 +102,13 @@ def test_multi_ctas_kv_counter_size_covers_beam_expanded_batch() -> None: # product clears the multi-processor floor, so pick a case that does. num_heads, batch, beam, sm_count = 6, 16, 2, 148 needed = num_heads * batch * beam * torch.int32.itemsize - assert _get_multi_ctas_kv_counter_size(num_heads, batch, sm_count) < needed - assert _get_multi_ctas_kv_counter_size(num_heads, batch * beam, sm_count) >= needed + assert get_multi_ctas_kv_counter_size(num_heads, batch, sm_count) < needed + assert get_multi_ctas_kv_counter_size(num_heads, batch * beam, sm_count) >= needed def test_multi_ctas_kv_counter_size_keeps_multi_processor_floor() -> None: num_heads, batch, sm_count = 6, 1, 148 - assert _get_multi_ctas_kv_counter_size(num_heads, batch, sm_count) >= ( + assert get_multi_ctas_kv_counter_size(num_heads, batch, sm_count) >= ( sm_count * torch.int32.itemsize ) @@ -133,7 +133,7 @@ def check_counter_size_args( monkeypatch.setattr( "tensorrt_llm._torch.attention.backends.fmha.flashinfer_trtllm_gen." - "_get_multi_ctas_kv_counter_size", + "get_multi_ctas_kv_counter_size", check_counter_size_args, ) diff --git a/tests/unittest/_torch/attention/test_fmha_params_codegen.py b/tests/unittest/_torch/attention/test_fmha_params_codegen.py new file mode 100644 index 000000000000..b12c2918e019 --- /dev/null +++ b/tests/unittest/_torch/attention/test_fmha_params_codegen.py @@ -0,0 +1,344 @@ +# 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. + +from __future__ import annotations + +import dataclasses +import importlib.util +import sys +import textwrap +from pathlib import Path +from types import ModuleType, SimpleNamespace + +import pytest + +REPO_ROOT = Path(__file__).resolve().parents[4] +GENERATOR_PATH = REPO_ROOT / "scripts" / "generate_fmha_params.py" +CPP_SCHEMA_PATH = REPO_ROOT / "tensorrt_llm" / "_torch" / "attention" / "backends" / "cpp_schema.py" + + +@pytest.fixture(scope="module") +def generator(): + """Load the generator script directly; it depends on nothing but the stdlib.""" + spec = importlib.util.spec_from_file_location("test_fmha_params_generator", GENERATOR_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def cpp_schema(): + """Load field metadata without importing TensorRT-LLM or PyTorch.""" + spec = importlib.util.spec_from_file_location("test_cpp_schema", CPP_SCHEMA_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize("postponed", [False, True]) +def test_schema_types_and_defaults( + generator, cpp_schema, tmp_path: Path, monkeypatch, postponed: bool +) -> None: + fields = """ + flag: bool = False + count: int = 1099511627776 + ratio: float = 1e-6 + optional_flag: Optional[bool] = None + optional_count: Optional[int] = None + optional_ratio: Optional[float] = None + attention_input_type: AttentionInputType = AttentionInputType.mixed + typed: Optional[torch.Tensor] = cpp_metadata(dtype=torch.float32) + indices: Optional[torch.Tensor] = cpp_metadata(dtype=torch.int32) + handwritten: Optional[torch.Tensor] = None + required: torch.Tensor = None + nested: Nested = None + optional_nested: Optional[Nested] = None + private: PythonOnly = None + python_only: Any = None + """ + source = ( + ("from __future__ import annotations\n" if postponed else "") + + "from dataclasses import dataclass\nfrom typing import Any, Optional\n" + + "from enum import IntEnum\nclass AttentionInputType(IntEnum):\n mixed = 0\n" + + "@dataclass\nclass Nested:\n" + + " count: int = 0\n" + + "@dataclass\nclass PythonOnly:\n count: int = 0\n" + + "@dataclass\nclass Example:\n" + + textwrap.indent(textwrap.dedent(fields).strip(), " ") + ) + path = tmp_path / "schema.py" + path.write_text(source) + schema = generator.load_schemas([path], "Example", struct_names=("Nested",)) + struct = schema.structs["Example"] + module = ModuleType("test_scalar_schema") + module.cpp_metadata = cpp_schema.cpp_metadata + module.torch = SimpleNamespace( + Tensor=type("Tensor", (), {"__module__": "torch"}), float32=object(), int32=object() + ) + monkeypatch.setitem(sys.modules, module.__name__, module) + exec(compile(source, "schema.py", "exec", dont_inherit=True), module.__dict__) + assert [f.name for f in struct.fields] == [ + "flag", + "count", + "ratio", + "optional_flag", + "optional_count", + "optional_ratio", + "attention_input_type", + "typed", + "indices", + "handwritten", + "required", + "nested", + "optional_nested", + ] + assert set(schema.structs) == {"Example", "Nested"} + params = module.Example() + assert params.flag is False + assert params.count == 1099511627776 + assert params.ratio == 1e-6 + assert params.optional_count is None + assert params.typed is None + assert params.nested is None + assert params.optional_nested is None + assert params.attention_input_type is module.AttentionInputType.mixed + rendered = generator.render_fields(schema, struct) + assert "TRTLLM_FMHA_PARAM_FIELD(flag, bool)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(count, std::int64_t)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(ratio, double)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(optional_count, std::optional)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(attention_input_type, std::int64_t)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(nested, Nested)" in rendered + assert "TRTLLM_FMHA_PARAM_FIELD(optional_nested, Nested)" in rendered + accessors = generator.render_accessors(schema, struct) + # Scalar widening must not widen a tensor's explicitly specified dtype. + assert "float* getTyped() const" in accessors + assert "std::int32_t* getIndices() const" in accessors + assert "getHandwritten" not in accessors + assert "getRequired" not in accessors + + +def test_legacy_arguments_build_python_params() -> None: + import torch + + from tensorrt_llm._torch.attention.backends.fmha.interface import ( + FmhaParams, + StaticAttentionConfig, + ) + from tensorrt_llm.functional import AttentionMaskType + + arguments = { + "q": torch.empty(0, dtype=torch.bfloat16), + "output": torch.empty(0, dtype=torch.bfloat16), + "num_heads": 8, + "num_seqs": 2, + "mask_type": AttentionMaskType.causal, + "max_attention_window_size": 2048, + "not_an_fmha_field": "ignored", + } + params = FmhaParams._from_arguments(arguments, layer_idx=3) + config = StaticAttentionConfig.from_legacy_arguments(arguments) + + assert config.num_heads == 8 + assert config.mask_type == AttentionMaskType.causal + assert params.num_seqs == 2 + assert params.max_attention_window_size == 2048 + assert params.layer_idx == 3 + assert "num_heads" not in {field.name for field in dataclasses.fields(params)} + + +def _stub_native_holder(monkeypatch): + """Stand in for the native holder, mirroring its nested layout. + + Nested holders and scalar defaults mirror C++ initialization. The sparse-index + sentinel makes a Python None leaving an existing native value untouched observable. + """ + + @dataclasses.dataclass(slots=True) + class NativeSparseRuntimeParams: + sparse_kv_indices: object = "native-sparse-indices" + threshold_scale_factor_prefill: float = 0.0 + threshold_scale_factor_decode: float = 0.0 + + class NativeForwardArgs: + __slots__ = ( + "output", + "output_sf", + "kv_norm_eps", + "update_kv_cache", + "is_fused_qkv", + "attention_window_size", + "attention_input_type", + "sparse_runtime_params", + "sparse_backend_args", + ) + + def __init__(self): + self.attention_window_size = 0 + self.sparse_runtime_params = NativeSparseRuntimeParams() + self.sparse_backend_args = SimpleNamespace() + + class NativeParams: + __slots__ = ( + "fwd", + "output", + "kv_pool", + "beam_width", + "is_cross", + "rotary_embedding_base", + "rotary_embedding_scale", + ) + + def __init__(self): + self.fwd = NativeForwardArgs() + + internal = ModuleType("tensorrt_llm.bindings.internal") + internal.thop = SimpleNamespace(FmhaParams=NativeParams) + monkeypatch.setitem(sys.modules, "tensorrt_llm.bindings.internal", internal) + return NativeParams + + +def _real_schema_classes(): + from tensorrt_llm._torch.attention.backends.fmha.interface import FmhaParams + from tensorrt_llm._torch.attention.backends.interface import ( + AttentionForwardArgs, + PredefinedAttentionMask, + ) + from tensorrt_llm._torch.attention.backends.sparse.params import SparseRuntimeParams + + return FmhaParams, AttentionForwardArgs, SparseRuntimeParams, PredefinedAttentionMask + + +def test_nested_python_args_are_lowered_once(monkeypatch) -> None: + FmhaParams, ForwardArgs, SparseParams, Mask = _real_schema_classes() + _stub_native_holder(monkeypatch) + + native = FmhaParams( + fwd=ForwardArgs( + output="full-output", + output_sf="output-scale", + attention_mask=Mask.CAUSAL, + attention_window_size=2048, + sparse_runtime_params=SparseParams( + sparse_kv_indices="sparse-indices", + threshold_scale_factor_prefill=0.25, + threshold_scale_factor_decode=0.5, + ), + ), + output="phase-output", + kv_pool="full-kv-pool", + beam_width=4, + is_cross=True, + rotary_embedding_base=500000.0, + rotary_embedding_scale=2.0, + ).to_op_params() + + # Keep the phase-local output distinct from the caller-facing full buffer; + # native kernels consume the former and leave the latter and kv_pool unused. + assert native.output == "phase-output" + assert native.kv_pool == "full-kv-pool" + assert native.fwd.output == "full-output" + assert native.fwd.output_sf == "output-scale" + assert native.fwd.kv_norm_eps == 1e-6 + assert native.fwd.update_kv_cache is True + assert native.fwd.is_fused_qkv is False + assert native.fwd.attention_window_size == 2048 + assert native.fwd.sparse_runtime_params.sparse_kv_indices == "sparse-indices" + assert native.fwd.sparse_runtime_params.threshold_scale_factor_prefill == 0.25 + assert native.fwd.sparse_runtime_params.threshold_scale_factor_decode == 0.5 + assert native.beam_width == 4 + assert native.is_cross is True + assert native.rotary_embedding_base == 500000.0 + assert native.rotary_embedding_scale == 2.0 + + +def test_static_config_plain_scalars_reach_native_holder(monkeypatch) -> None: + from tensorrt_llm._torch.attention.backends.fmha.interface import StaticAttentionConfig + + _stub_native_holder(monkeypatch) + + @dataclasses.dataclass(slots=True) + class NativeStaticConfig: + num_heads: int = 0 + q_scaling: float = 0.0 + remove_padding: bool = False + use_kv_cache: bool = False + + sys.modules["tensorrt_llm.bindings.internal"].thop.StaticAttentionConfig = NativeStaticConfig + native = StaticAttentionConfig(num_heads=8, q_scaling=0.125).to_thop_config() + + assert native.num_heads == 8 + assert native.q_scaling == 0.125 + assert native.remove_padding is True + assert native.use_kv_cache is False + + +def test_nested_none_does_not_replace_native_value(monkeypatch) -> None: + FmhaParams, ForwardArgs, SparseParams, Mask = _real_schema_classes() + _stub_native_holder(monkeypatch) + + forward_args = ForwardArgs() + assert forward_args.attention_mask is Mask.CAUSAL + assert isinstance(forward_args.sparse_runtime_params, SparseParams) + assert forward_args.sparse_runtime_params.sparse_kv_indices is None + native = FmhaParams(fwd=forward_args).to_op_params() + + # attention_window_size defaults to None on the Python side, which must leave + # the value-initialized native field alone. + assert native.fwd.attention_window_size == 0 + assert native.fwd.attention_input_type == 0 + assert native.fwd.sparse_runtime_params.sparse_kv_indices == "native-sparse-indices" + assert native.fwd.sparse_runtime_params.threshold_scale_factor_prefill == 0.0 + assert native.fwd.sparse_runtime_params.threshold_scale_factor_decode == 0.0 + + # The cached lowering plan contains field names, never values or None decisions. + forward_args.attention_window_size = 128 + forward_args.sparse_runtime_params = SparseParams(threshold_scale_factor_prefill=0.5) + updated = FmhaParams(fwd=forward_args).to_op_params() + assert updated.fwd.attention_window_size == 128 + assert updated.fwd.sparse_runtime_params.threshold_scale_factor_prefill == 0.5 + assert native.fwd.attention_window_size == 0 + assert native.fwd.sparse_runtime_params.threshold_scale_factor_prefill == 0.0 + + +@pytest.mark.parametrize("phase_lengths", [None, "phase-lengths"]) +def test_phase_values_override_shared_state_when_building_op_params(phase_lengths) -> None: + from tensorrt_llm._torch.attention.backends.fmha.interface import build_op_params + + @dataclasses.dataclass + class SharedParams: + spec_decoding_generation_lengths: object = "retained-lengths" + max_num_requests: int = 8 + + @dataclasses.dataclass + class PhaseParams: + spec_decoding_generation_lengths: object = phase_lengths + + @dataclasses.dataclass(slots=True) + class NativeParams: + spec_decoding_generation_lengths: object = None + max_num_requests: int = 0 + + metadata = SharedParams() + native = NativeParams() + build_op_params(native, metadata, PhaseParams()) + + assert native.spec_decoding_generation_lengths == phase_lengths + assert native.max_num_requests == 8 + assert metadata.spec_decoding_generation_lengths == "retained-lengths" diff --git a/tests/unittest/_torch/attention/test_prims_ts_fmha.py b/tests/unittest/_torch/attention/test_prims_ts_fmha.py index a30f05d25c2e..b0413edbe3c3 100644 --- a/tests/unittest/_torch/attention/test_prims_ts_fmha.py +++ b/tests/unittest/_torch/attention/test_prims_ts_fmha.py @@ -39,6 +39,7 @@ PredefinedAttentionMask, ) from tensorrt_llm._torch.attention.backends.sparse.params import SparseRuntimeParams +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType @@ -70,6 +71,8 @@ def numel(self) -> int: class _Attention: + out_head_size = TrtllmAttention.out_head_size + def __init__( self, *, @@ -86,6 +89,7 @@ def __init__( self.qk_rope_head_dim = 64 if is_mla else None self.qk_nope_head_dim = 128 if is_mla else None self.v_head_dim = 128 if is_mla else None + self.rope_append = True self.predicted_tokens_per_seq = 1 self.sparse_params = None self.skip_correction_threshold = 0.0 @@ -2226,12 +2230,20 @@ def test_workspace_cannot_grow_during_capture(monkeypatch: pytest.MonkeyPatch) - (AttentionInputType.generation_only, True, False, False), ], ) +@pytest.mark.parametrize("output_dtype", [torch.bfloat16, torch.uint8]) +@pytest.mark.parametrize( + "is_spec_decoding_enabled,use_spec_decoding", + [(False, False), (False, True), (True, False), (True, True)], +) def test_phased_forward_routes_query_and_qkv_by_phase( monkeypatch: pytest.MonkeyPatch, input_type: AttentionInputType, is_mla: bool, is_fused_qkv: bool, has_kv: bool, + output_dtype: torch.dtype, + is_spec_decoding_enabled: bool, + use_spec_decoding: bool, ) -> None: attn = _Attention(is_mla=is_mla) fmha = PrimsTSFmha(attn) @@ -2260,7 +2272,8 @@ def test_phased_forward_routes_query_and_qkv_by_phase( if input_type == AttentionInputType.generation_only else fmha.context_out_head_size ) - output = torch.empty((num_tokens, attn.num_heads * out_head_size)) + stored_head_size = out_head_size // 2 if output_dtype == torch.uint8 else out_head_size + output = torch.empty((num_tokens, attn.num_heads * stored_head_size), dtype=output_dtype) metadata = SimpleNamespace( kv_cache_block_offsets=torch.empty(1), effective_workspace=torch.empty(0, dtype=torch.int8), @@ -2274,7 +2287,14 @@ def test_phased_forward_routes_query_and_qkv_by_phase( kv_lens_runtime=torch.tensor([3, 65, 97], dtype=torch.int32), prompt_lens_cuda_runtime=torch.tensor([3, 1, 1], dtype=torch.int32), prompt_lens_cpu_runtime=torch.tensor([3, 1, 1], dtype=torch.int32), - is_spec_decoding_enabled=False, + is_spec_decoding_enabled=is_spec_decoding_enabled, + use_spec_decoding=use_spec_decoding, + spec_decoding_generation_lengths=torch.ones(2, dtype=torch.int32), + spec_decoding_position_offsets=torch.zeros((2, 1), dtype=torch.int32), + spec_decoding_packed_mask=torch.zeros((2, 1, 1), dtype=torch.int32), + spec_decoding_bl_tree_mask_offset=torch.zeros(2, dtype=torch.int64), + spec_decoding_bl_tree_mask=torch.zeros(2, dtype=torch.uint32), + spec_bl_tree_first_sparse_mask_offset_kv=torch.zeros(2, dtype=torch.int32), is_cross=False, kv_cache_manager=None, ) @@ -2288,6 +2308,25 @@ def test_phased_forward_routes_query_and_qkv_by_phase( fmha.forward(q, k, v, metadata, forward_args) + spec_fields = ( + "spec_decoding_generation_lengths", + "spec_decoding_packed_mask", + "spec_decoding_bl_tree_mask_offset", + "spec_decoding_bl_tree_mask", + "spec_bl_tree_first_sparse_mask_offset_kv", + ) + for phase_calls, spec_active in ( + (context_calls, False), + (generation_calls, is_spec_decoding_enabled and use_spec_decoding), + ): + for params in phase_calls: + assert params.use_spec_decoding == spec_active + for name in spec_fields: + source = getattr(metadata, name) + assert getattr(params, name) is (source if spec_active else None) + offsets = metadata.spec_decoding_position_offsets if spec_active else None + assert params.spec_decoding_position_offsets is offsets + if input_type == AttentionInputType.generation_only: run_context.assert_not_called() else: @@ -2328,5 +2367,5 @@ def test_phased_forward_routes_query_and_qkv_by_phase( else: assert params.key_input is None assert params.value_input is None - assert params.output.shape == (params.num_tokens, attn.num_heads, out_head_size) + assert params.output.shape == (params.num_tokens, attn.num_heads, stored_head_size) assert params.output.data_ptr() == output[token_slice].data_ptr() diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py index cb0a1f090d4c..41dae9ac5be1 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py @@ -414,7 +414,7 @@ def call_op( v=None, output=output, output_sf=None, - workspace_=self.workspace, + workspace=self.workspace, sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), host_past_key_value_lengths=torch.tensor(kv_lens, dtype=torch.int32), host_total_kv_lens=torch.tensor([total_ctx_kv, total_gen_kv], dtype=torch.int32), @@ -482,7 +482,7 @@ def call_op( use_spec_decoding=False, is_spec_dec_tree=False, spec_decoding_generation_lengths=None, - spec_decoding_position_offsets_for_cpp=None, + spec_decoding_position_offsets=None, spec_decoding_packed_mask=None, spec_decoding_bl_tree_mask_offset=None, spec_decoding_bl_tree_mask=None, @@ -885,7 +885,7 @@ def call_op( v=None, output=output, output_sf=None, - workspace_=self.workspace, + workspace=self.workspace, sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), host_past_key_value_lengths=torch.tensor(kv_lens, dtype=torch.int32), host_total_kv_lens=torch.tensor( @@ -956,7 +956,7 @@ def call_op( use_spec_decoding=False, is_spec_dec_tree=False, spec_decoding_generation_lengths=None, - spec_decoding_position_offsets_for_cpp=None, + spec_decoding_position_offsets=None, spec_decoding_packed_mask=None, spec_decoding_bl_tree_mask_offset=None, spec_decoding_bl_tree_mask=None, @@ -2993,7 +2993,7 @@ def _call_op( v=v, output=output, output_sf=None, - workspace_=self.workspace, + workspace=self.workspace, sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), host_past_key_value_lengths=torch.tensor( kv_lens if host_past_lens is None else host_past_lens, @@ -3065,7 +3065,7 @@ def _call_op( use_spec_decoding=False, is_spec_dec_tree=False, spec_decoding_generation_lengths=None, - spec_decoding_position_offsets_for_cpp=None, + spec_decoding_position_offsets=None, spec_decoding_packed_mask=None, spec_decoding_bl_tree_mask_offset=None, spec_decoding_bl_tree_mask=None, @@ -5247,8 +5247,7 @@ def test_fp8_mla_mtp_decode_sees_a_torn_pool_h128() -> None: def test_mla_rejects_null_q_lora_rank() -> None: - """q_lora_rank must be an int on the MLA path: the C++ unwraps the - optional unconditionally, so None raises rather than defaulting.""" + """The catalog requires an explicit int rank for MLA, including zero for no q-LoRA.""" torch.manual_seed(309) h = MLA_NUM_HEADS_H32 q, k, v, latent = _random_context_inputs(32, h) @@ -5261,7 +5260,7 @@ def test_mla_rejects_null_q_lora_rank() -> None: env.add_request(0, 32) try: env.call_context([0], [32], q, k, v, latent) - except RuntimeError as exc: - assert "bad optional access" in str(exc), f"unexpected message: {exc}" + except ValueError as exc: + assert "q_lora_rank must be an int for MLA" in str(exc), f"unexpected message: {exc}" else: raise AssertionError("q_lora_rank=None was accepted on the MLA path") diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py index e04a4e01ac4f..ca3d90aa8603 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py @@ -3,7 +3,7 @@ """Every op a target names has to exist in the build it is running against. A target reads private engine surface: custom ops registered under -``torch.ops.trtllm`` and one pybind entry point. With no version pin, nothing +``torch.ops.trtllm`` and the native ``AttentionOp`` class. With no version pin, nothing declares which build it was written for -- so the only thing that can tell you the surface moved is running against it. @@ -62,15 +62,17 @@ def test_every_declared_op_exists(name, module): @pytest.mark.parametrize("name,module", list(_target_modules()), ids=_target_ids()) def test_the_pybind_attention_entry_point_exists(name, module): - """``thop.attention`` is reached through the bindings, not torch.ops. - - Separate from the loop above because a missing pybind symbol fails in a - different way -- an ImportError or an AttributeError on the module object - rather than a missing torch.ops entry -- and both targets go through it. - """ + """FallbackFmha dispatches both targets through the native AttentionOp methods.""" from tensorrt_llm.bindings.internal import thop - assert hasattr(thop, "attention"), ( - f"{name} calls the attention op through tensorrt_llm.bindings.internal" - f".thop.attention, which this build does not expose" - ) + for symbol in ("AttentionOp", "StaticAttentionConfig", "FmhaParams"): + assert hasattr(thop, symbol), f"{name} requires thop.{symbol}" + for method in ( + "get_attention_workspace_size", + "run_context", + "run_generation", + "run_mla_generation", + ): + assert callable(getattr(thop.AttentionOp, method, None)), ( + f"{name} requires thop.AttentionOp.{method}" + ) diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py index 70a6688c4fa2..9dfed0d8eea4 100644 --- a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py @@ -1210,7 +1210,7 @@ def test_mla_dspark_cute_dsl_takes_a_single_page_batch_one(): so CuTe DSL's layout deduction refused to pick a leading dimension even though with both extents 1 the choice cannot change an address. Reachable whenever the draft pool is one page deep and a single request is resident -- a - first-draft-step shape, and since AUTO resolves to CUTEDSL, the default path. + first-draft-step shape when CUTEDSL is selected. The parity bound is the same one the batched test uses; a wrongly deduced axis would show up as a magnitude error, not a last-bit one. @@ -1237,26 +1237,23 @@ def test_mla_dspark_cute_dsl_takes_a_single_page_batch_one(): @needs_cuda -def test_mla_dspark_auto_backend_resolves_to_cutedsl(): - """AUTO is the one-pass kernel, and it must reach the dispatch as such. +def test_mla_dspark_auto_backend_resolves_to_trtllm(): + """AUTO selects TRTLLM and must reach the paged decode dispatch as such. Asserting the dispatch and not just the field: every variant returns the same shape, so a default that resolved but never reached _mla_block_decode_variant would be silent. """ from tensorrt_llm._torch.model_config import ModelConfig - from tensorrt_llm._torch.models.modeling_dspark import ( - MLADSparkForCausalLM, - cute_dsl_mla_decode_unavailability_reason, - ) + from tensorrt_llm._torch.models.modeling_dspark import MLADSparkForCausalLM - reason = cute_dsl_mla_decode_unavailability_reason() + reason = MLADSparkForCausalLM._attention_backend_unavailability_reason("TRTLLM") if reason is not None: - pytest.skip(f"cute-dsl MLA decode unavailable: {reason}") + pytest.skip(f"TRTLLM MLA decode unavailable: {reason}") model_config = ModelConfig(pretrained_config=_tiny_mla_config(), attn_backend="VANILLA") drafter = MLADSparkForCausalLM(model_config, dflash_attention_backend="AUTO") - assert drafter.dflash_attention_backend == "CUTEDSL" - assert drafter._mla_block_decode_variant(paged=True) == "cute_dsl" + assert drafter.dflash_attention_backend == "TRTLLM" + assert drafter._mla_block_decode_variant(paged=True) == "trtllm_gen" @needs_cuda diff --git a/tests/unittest/_torch/speculative/test_eagle3.py b/tests/unittest/_torch/speculative/test_eagle3.py index 655571c9092e..f50ac7a54608 100644 --- a/tests/unittest/_torch/speculative/test_eagle3.py +++ b/tests/unittest/_torch/speculative/test_eagle3.py @@ -54,6 +54,38 @@ KvCacheConfig, MoeConfig, MTPDecodingConfig) +@pytest.mark.parametrize( + "worker_type", [Eagle3OneModelDynamicTreeWorker, MTPEagleDynamicTreeWorker]) +def test_dynamic_tree_restores_position_offsets_query_length( + worker_type) -> None: + worker = object.__new__(worker_type) + offsets = torch.arange(32, dtype=torch.int32) + metadata = SimpleNamespace( + num_seqs=2, + prepare_for_spec_dec=MagicMock(), + restore_from_spec_dec=MagicMock(), + on_update=MagicMock(), + kv_lens_cuda=torch.full((2, ), 16, dtype=torch.int32), + spec_decoding_position_offsets=offsets, + spec_decoding_query_len=3, + spec_decoding_generation_lengths=torch.full((2, ), 3, + dtype=torch.int32), + spec_decoding_packed_mask=torch.zeros((2, 8, 1), dtype=torch.int32), + ) + worker._prepare_attn_metadata_for_spec_dec(metadata) + offsets.fill_(7) + metadata.spec_decoding_query_len = 5 + metadata.spec_decoding_generation_lengths.fill_(5) + worker._restore_attn_metadata_from_spec_dec(metadata) + + assert metadata.spec_decoding_position_offsets is offsets + torch.testing.assert_close(offsets, torch.arange(32, dtype=torch.int32)) + assert metadata.spec_decoding_query_len == 3 + assert metadata.spec_decoding_generation_lengths.tolist() == [3, 3] + metadata.restore_from_spec_dec.assert_called_once() + metadata.on_update.assert_called_once() + + def test_mtp_eagle_refreshes_dsa_metadata_before_draft_forward() -> None: """Refresh DSA mappings after switching to the draft cache.""" events = [] diff --git a/tests/unittest/bindings/test_bindings_ut.py b/tests/unittest/bindings/test_bindings_ut.py index 2764479f0823..92e93089ca50 100644 --- a/tests/unittest/bindings/test_bindings_ut.py +++ b/tests/unittest/bindings/test_bindings_ut.py @@ -11,9 +11,19 @@ import tensorrt_llm.bindings as _tb import tensorrt_llm.bindings.executor as _tbe +from tensorrt_llm.bindings.internal import thop +from tensorrt_llm.functional import PositionEmbeddingType from tensorrt_llm.llmapi.kv_cache_type import KVCacheType +def test_position_embedding_type_deferred() -> None: + config = thop.StaticAttentionConfig() + config.position_embedding_type = PositionEmbeddingType.deferred + + assert config.position_embedding_type == _tb.PositionEmbeddingType.DEFERRED + assert config.position_embedding_type.value == 10 + + def test_quant_mode(): assert _tb.QuantMode.none().value == 0 assert _tb.QuantMode.int4_weights().has_int4_weights