From 19df4f2e4dd6de1a458635c7740667febc65f4ac Mon Sep 17 00:00:00 2001 From: zhoutianzi666 <17801055074@163.com> Date: Mon, 15 Jun 2026 22:07:52 +0800 Subject: [PATCH 1/3] support swa --- csrc/api/dense_decode.h | 15 +++++- csrc/params.h | 1 + .../get_decoding_sched_meta.cu | 46 +++++++++++++++++++ flash_mla/flash_mla_interface.py | 6 ++- 4 files changed, 64 insertions(+), 4 deletions(-) diff --git a/csrc/api/dense_decode.h b/csrc/api/dense_decode.h index 7df178a6..20b74bae 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,12 @@ 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 ) { + std::cout << "batch_size" << std::endl; + 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 +115,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 +214,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 +225,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..e7cab0d7 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,53 @@ 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){ + for (int i = 0; i < 100; i++) + std::cout << "niubi" << std::endl; + 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 From ba23741ba7ecf2f05f6c7e236d3c700d604b8fca Mon Sep 17 00:00:00 2001 From: zhoutianzi666 <17801055074@163.com> Date: Mon, 15 Jun 2026 22:08:11 +0800 Subject: [PATCH 2/3] support swa --- csrc/api/dense_decode.h | 1 - 1 file changed, 1 deletion(-) diff --git a/csrc/api/dense_decode.h b/csrc/api/dense_decode.h index 20b74bae..bb52cf4a 100644 --- a/csrc/api/dense_decode.h +++ b/csrc/api/dense_decode.h @@ -80,7 +80,6 @@ dense_attn_decode_interface( if (swa_size > 0 ) { - std::cout << "batch_size" << std::endl; num_sm_parts = batch_size; } From 6fda9bb346f102339a8993b9c2caaa4dc073b2da Mon Sep 17 00:00:00 2001 From: zhoutianzi666 <17801055074@163.com> Date: Tue, 16 Jun 2026 13:55:07 +0800 Subject: [PATCH 3/3] support swa --- .../decode/get_decoding_sched_meta/get_decoding_sched_meta.cu | 2 -- 1 file changed, 2 deletions(-) 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 e7cab0d7..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 @@ -146,8 +146,6 @@ get_mla_metadata_kernel2(__grid_constant__ const GetDecodeSchedMetaParams params void run_get_decoding_sched_meta_kernel(GetDecodeSchedMetaParams ¶ms) { if (params.swa_size > 0){ - for (int i = 0; i < 100; i++) - std::cout << "niubi" << std::endl; get_mla_metadata_kernel2<<<1, 32, 0, params.stream>>>(params); CHECK_CUDA_KERNEL_LAUNCH(); return;