Skip to content
Open
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
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 12 additions & 2 deletions csrc/api/dense_decode.h
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,8 @@ dense_attn_decode_interface(
const float softmax_scale,
bool is_causal,
std::optional<at::Tensor> &tile_scheduler_metadata, // num_sm_parts x (DecodingSchedMetaSize/4)
std::optional<at::Tensor> &num_splits // batch_size + 1
std::optional<at::Tensor> &num_splits, // batch_size + 1
const int swa_size
) {
// Check arch
Arch arch = Arch();
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -108,7 +114,8 @@ dense_attn_decode_interface(
(DecodingSchedMeta*)tile_scheduler_metadata->data_ptr(),
num_splits->data_ptr<int>(),
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 {
Expand Down Expand Up @@ -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<cutlass::bfloat16_t>(combine_params);
} else if (q_dtype == torch::kHalf) {
Expand All @@ -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});
Expand Down
1 change: 1 addition & 0 deletions csrc/params.h
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,7 @@ struct GetDecodeSchedMetaParams {
int num_sm_parts;

cudaStream_t stream;
int swa_size;
};

struct SparseAttnFwdParams {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 &params) {

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);
Expand Down
6 changes: 4 additions & 2 deletions flash_mla/flash_mla_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down