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