Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
The table of contents is too big for display.
Diff view
Diff view
  •  
  •  
  •  
The diff you're trying to view is too large. We only load the first 3000 changed files.
36 changes: 22 additions & 14 deletions cpp/tensorrt_llm/common/attentionOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -881,15 +881,16 @@ size_t AttentionOp::getWorkspaceSizeForContext(tensorrt_llm::DataType type, int3
else if (useSageAttnSeparateQkv)
{
fp8_q_buf_size = max_num_tokens * static_cast<size_t>(local_hidden_units_qo);
fp8_k_buf_size = max_num_tokens * static_cast<size_t>(local_hidden_units_kv);
fp8_v_buf_size = max_num_tokens * static_cast<size_t>(local_hidden_units_kv);
fp8_k_buf_size = total_kv_len * static_cast<size_t>(local_hidden_units_kv);
fp8_v_buf_size = total_kv_len * static_cast<size_t>(local_hidden_units_kv);
}

int32_t const q_max_n_blk = mSageAttnNumEltsPerBlkQ > 0 ? tc::divUp(input_seq_length, mSageAttnNumEltsPerBlkQ) : 0;
int32_t const k_max_n_blk = mSageAttnNumEltsPerBlkK > 0 ? tc::divUp(kv_seq_length, mSageAttnNumEltsPerBlkK) : 0;
size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * batch_size * static_cast<size_t>(q_max_n_blk);
size_t const sage_k_sfs_buffer_size
= sizeof(float) * mNumAttnKVHeads * batch_size * static_cast<size_t>(k_max_n_blk);
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<size_t>(q_max_n_blk);
size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast<size_t>(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;
Expand Down Expand Up @@ -1577,15 +1578,16 @@ int AttentionOp::enqueueContext(EnqueueContextParams<T> const& params, cudaStrea
fp8_v_buf_size = params.total_kv_len * static_cast<size_t>(local_hidden_units_kv);
}

int32_t const q_max_n_blk
= mSageAttnNumEltsPerBlkQ > 0 ? tc::divUp(params.input_seq_length, mSageAttnNumEltsPerBlkQ) : 0;
int32_t const k_max_n_blk = mSageAttnNumEltsPerBlkK > 0 ? tc::divUp(kv_seq_length, mSageAttnNumEltsPerBlkK) : 0;
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 * params.batch_size * static_cast<size_t>(q_max_n_blk);
size_t const sage_k_sfs_buffer_size
= sizeof(float) * mNumAttnKVHeads * params.batch_size * static_cast<size_t>(k_max_n_blk);
size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast<size_t>(q_max_n_blk);
size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast<size_t>(k_max_n_blk);
size_t const sage_v_sfs_buffer_size = sizeof(float) * v_max_n_blk;

size_t const padding_offset_size
Expand Down Expand Up @@ -1905,6 +1907,8 @@ int AttentionOp::enqueueContext(EnqueueContextParams<T> const& params, cudaStrea
{
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<int>(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.");
Expand All @@ -1928,8 +1932,10 @@ int AttentionOp::enqueueContext(EnqueueContextParams<T> const& params, cudaStrea

// 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;
Expand All @@ -1938,8 +1944,10 @@ int AttentionOp::enqueueContext(EnqueueContextParams<T> const& params, cudaStrea

// 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;
Comment thread
coderabbitai[bot] marked this conversation as resolved.
sageQuantParams.ptrQk = params.k_ptr;
sageQuantParams.ptrQkQuant = workspaceViews.fp8KBuf;
sageQuantParams.ptrQkScale = workspaceViews.sageKScale;
Expand Down
Loading
Loading