refactor: extract shared dispatcher header, unify MMA/scalar dispatch format

- Merge 3 duplicated dispatch blocks into single attn_dispatchers.cuh
- Merge compute_num_splits from attn_utils.cuh into dispatcher header
- All dim3 grid/block declarations and <<<>>> launches are single-line
- Production .cu files (35-42 loc) only handle torch wrapping + pybind11
- Test files include dispatcher header directly, removing all #ifndef ASTRAI_NO_MMA duplication
This commit is contained in:
2026-07-21 22:21:39 +08:00
parent a01e8bbe98
commit f7a16efc9d
9 changed files with 245 additions and 406 deletions
+2 -64
View File
@@ -1,69 +1,6 @@
#include "attn_decode_split_kv.cuh"
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_decode_split_kv_mma.cuh"
template <int HEAD_DIM, int BC, int STAGES, bool IsCausal, bool HasMask>
static void launch_mma_decode_impl(AttentionParams<bf16>& p) {
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM, int BC, bool IsCausal, bool HasMask>
static void launch_mma_decode(AttentionParams<bf16>& p) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
launch_mma_decode_impl<HEAD_DIM, BC, STAGES, IsCausal, HasMask>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_decode(AttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
alloc_split_partials(p);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static void dispatch_decode(AttentionParams<bf16>& p) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) {
if (is_causal) {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, true, true>(p);
else launch_mma_decode<HEAD_DIM, 32, true, false>(p);
} else {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, false, true>(p);
else launch_mma_decode<HEAD_DIM, 32, false, false>(p);
}
return;
}
#endif
if (is_causal) {
if (has_mask) launch_scalar_decode<HEAD_DIM, true, true>(p);
else launch_scalar_decode<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_scalar_decode<HEAD_DIM, false, true>(p);
else launch_scalar_decode<HEAD_DIM, false, false>(p);
}
}
torch::Tensor attn_decode(
torch::Tensor q,
torch::Tensor k,
@@ -82,6 +19,7 @@ torch::Tensor attn_decode(
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
return O;
}