diff --git a/csrc/api/dense_decode.h b/csrc/api/dense_decode.h index 7df178a6..bb52cf4a 100644 --- a/csrc/api/dense_decode.h +++ b/csrc/api/dense_decode.h @@ -20,7 +20,8 @@ dense_attn_decode_interface( const float softmax_scale, bool is_causal, std::optional &tile_scheduler_metadata, // num_sm_parts x (DecodingSchedMetaSize/4) - std::optional &num_splits // batch_size + 1 + std::optional &num_splits, // batch_size + 1 + const int swa_size ) { // Check arch Arch arch = Arch(); @@ -77,6 +78,11 @@ dense_attn_decode_interface( .reshape({batch_size, q_seq_per_hk, num_heads, head_size_k}); int num_sm_parts = std::max(arch.num_sms / num_heads_k / cutlass::ceil_div(seqlen_q_ori*num_heads_q/num_heads_k, 64), 1); + + if (swa_size > 0 ) { + num_sm_parts = batch_size; + } + KU_CHECK_SHAPE(q, batch_size, q_seq_per_hk, num_heads, head_size_k); KU_CHECK_SHAPE(kcache, num_blocks, page_block_size, num_heads_k, head_size_k); KU_CHECK_SHAPE(seqlens_k, batch_size); @@ -108,7 +114,8 @@ dense_attn_decode_interface( (DecodingSchedMeta*)tile_scheduler_metadata->data_ptr(), num_splits->data_ptr(), num_sm_parts, - at::cuda::getCurrentCUDAStream().stream() + at::cuda::getCurrentCUDAStream().stream(), + swa_size, }; smxx::decode::run_get_decoding_sched_meta_kernel(get_sched_meta_params); } else { @@ -206,6 +213,8 @@ dense_attn_decode_interface( at::cuda::getCurrentCUDAStream().stream() }; + if (swa_size < 0){ + if (q_dtype == torch::kBFloat16) { smxx::decode::run_flash_mla_combine_kernel(combine_params); } else if (q_dtype == torch::kHalf) { @@ -215,6 +224,7 @@ dense_attn_decode_interface( } else { TORCH_CHECK(false, "Unsupported tensor dtype for query"); } +} out = out.view({batch_size, num_heads_k, seqlen_q_ori, num_q_heads_per_hk, head_size_v}).transpose(1, 2) .reshape({batch_size, seqlen_q_ori, num_heads_q, head_size_v}); diff --git a/csrc/params.h b/csrc/params.h index 4433e8d4..35962822 100644 --- a/csrc/params.h +++ b/csrc/params.h @@ -140,6 +140,7 @@ struct GetDecodeSchedMetaParams { int num_sm_parts; cudaStream_t stream; + int swa_size; }; struct SparseAttnFwdParams { diff --git a/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu b/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu index 083da60c..f705cd40 100644 --- a/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu +++ b/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu @@ -106,7 +106,51 @@ get_mla_metadata_kernel(__grid_constant__ const GetDecodeSchedMetaParams params) } } + +__global__ void __launch_bounds__(32, 1, 1) +get_mla_metadata_kernel2(__grid_constant__ const GetDecodeSchedMetaParams params) { + int *seqlens_k_ptr = params.seqlens_k_ptr; + DecodingSchedMeta *tile_scheduler_metadata_ptr = params.tile_scheduler_metadata_ptr; + int batch_size = params.b; + int block_size_n = params.block_size_n; + int num_sm_parts = params.num_sm_parts; + + if (threadIdx.x == 0) { + for (int i = 0; i < num_sm_parts; ++i) { + DecodingSchedMeta cur_meta; + int seqlen_k = seqlens_k_ptr[i]; + + + cur_meta.begin_req_idx = i; + cur_meta.end_req_idx = i; + + if (seqlen_k >= params.swa_size) { + cur_meta.begin_block_idx = (seqlen_k - params.swa_size) / block_size_n; + } else { + cur_meta.begin_block_idx = 0; + } + + cur_meta.begin_split_idx = 0; + cur_meta.is_first_req_splitted = false; + + cur_meta.end_block_idx = (seqlen_k + block_size_n) / block_size_n; + + cur_meta.is_last_req_splitted = false; + cur_meta.is_first_req_splitted = false; + tile_scheduler_metadata_ptr[i] = cur_meta; + } + } +} + + void run_get_decoding_sched_meta_kernel(GetDecodeSchedMetaParams ¶ms) { + + if (params.swa_size > 0){ + get_mla_metadata_kernel2<<<1, 32, 0, params.stream>>>(params); + CHECK_CUDA_KERNEL_LAUNCH(); + return; + } + int smem_size = sizeof(int) * (params.b*5+1); CHECK_CUDA(cudaFuncSetAttribute(get_mla_metadata_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size)); get_mla_metadata_kernel<<<1, 32, smem_size, params.stream>>>(params); diff --git a/flash_mla/flash_mla_interface.py b/flash_mla/flash_mla_interface.py index a3740b0f..447dd117 100644 --- a/flash_mla/flash_mla_interface.py +++ b/flash_mla/flash_mla_interface.py @@ -66,7 +66,8 @@ def flash_mla_with_kvcache( extra_k_cache: Optional[torch.Tensor] = None, extra_indices_in_kvcache: Optional[torch.Tensor] = None, topk_length: Optional[torch.Tensor] = None, - extra_topk_length: Optional[torch.Tensor] = None + extra_topk_length: Optional[torch.Tensor] = None, + swa_size: int = -1 ) -> Tuple[torch.Tensor, torch.Tensor]: """ Arguments: @@ -166,7 +167,8 @@ def flash_mla_with_kvcache( q, k_cache, head_dim_v, cache_seqlens, block_table, softmax_scale, causal, - sched_meta.tile_scheduler_metadata, sched_meta.num_splits + sched_meta.tile_scheduler_metadata, sched_meta.num_splits, + swa_size ) sched_meta.tile_scheduler_metadata = new_tile_scheduler_metadata sched_meta.num_splits = new_num_splits