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
+19 -73
View File
@@ -1,15 +1,12 @@
/*
Pure-C test — updated for KernelTraits + IsCausal/HasMask.
Pure-C test — uses shared dispatcher.
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_decode_test.cu -o test && ./test
*/
#include "test_utils.cuh"
#include "../kernels/attn_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_decode_split_kv_mma.cuh"
#endif
#include "../kernels/attn_dispatchers.cuh"
// Split-K scratch (torch-free)
struct DecodeScratch {
@@ -17,71 +14,20 @@ struct DecodeScratch {
float* ml_part = nullptr;
};
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, int BC, bool IsCausal, bool HasMask>
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
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);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
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);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
int gs = 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);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask>
<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
int max_splits = 32;
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
}
template <int HEAD_DIM>
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
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, sc);
else launch_mma_decode<HEAD_DIM, 32, true, false>(p, sc);
} else {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, false, true>(p, sc);
else launch_mma_decode<HEAD_DIM, 32, false, false>(p, sc);
}
return;
}
#endif
if (is_causal) {
if (has_mask) launch_scalar_decode<HEAD_DIM, true, true>(p, sc);
else launch_scalar_decode<HEAD_DIM, true, false>(p, sc);
} else {
if (has_mask) launch_scalar_decode<HEAD_DIM, false, true>(p, sc);
else launch_scalar_decode<HEAD_DIM, false, false>(p, sc);
}
}
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
static void free_scratch(DecodeScratch& sc) {
cudaFree(sc.o_part); cudaFree(sc.ml_part);
}
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
static void bench() {
const int cfgs[][5] = {
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
{1, 32, 4, 512, 128},
{1, 32, 4, 1024, 128},
{1, 32, 4, 2048, 128},
{1, 32, 4, 4096, 128},
@@ -118,10 +64,10 @@ static void bench() {
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
auto launch = [&]() { dispatch_decode(p, sc); };
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); }); };
double flops = 4.0 * B * Hq * (double)sl * D;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
@@ -133,7 +79,7 @@ static void bench() {
print_bench_row(cfg, r);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
cudaFree(sc.o_part); cudaFree(sc.ml_part);
free_scratch(sc);
}
}
@@ -173,11 +119,11 @@ static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
double t0=now_ms();
dispatch_decode(p, sc);
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); });
cudaDeviceSynchronize();
double kms=now_ms()-t0;
cudaError_t err=cudaGetLastError();
@@ -197,7 +143,7 @@ static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
cudaFree(sc.o_part);cudaFree(sc.ml_part);
free_scratch(sc);
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
return (max_err < 0.05f) ? 0 : 1;
@@ -205,10 +151,10 @@ static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
int main() {
const int configs[][6] = {
{1, 2, 1, 64, 32, 0}, // B,Hq,Hk,seq_len,D,causal
{1, 2, 1, 64, 32, 0},
{1, 32, 4, 512, 128, 0},
{1, 32, 4, 1024, 128, 0},
{1, 32, 4, 512, 128, 1}, // causal decode
{1, 32, 4, 512, 128, 1},
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
int fail = 0;