diff --git a/csrc/kernels/attn_decode.cu b/csrc/kernels/attn_decode.cu index 3a8d802..a9fec7a 100644 --- a/csrc/kernels/attn_decode.cu +++ b/csrc/kernels/attn_decode.cu @@ -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 -static void launch_mma_decode_impl(AttentionParams& p) { - using Traits = KernelTraits; - 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<<>>(p); - attn_decode_combine_kernel<<>>(p); -} - -template -static void launch_mma_decode(AttentionParams& p) { - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - launch_mma_decode_impl(p); -} -#endif - -template -static void launch_scalar_decode(AttentionParams& 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<<>>(p); - attn_decode_combine_kernel<<>>(p); -} - -template -static void dispatch_decode(AttentionParams& 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(p); - else launch_mma_decode(p); - } else { - if (has_mask) launch_mma_decode(p); - else launch_mma_decode(p); - } - return; - } -#endif - if (is_causal) { - if (has_mask) launch_scalar_decode(p); - else launch_scalar_decode(p); - } else { - if (has_mask) launch_scalar_decode(p); - else launch_scalar_decode(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; } diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh new file mode 100644 index 0000000..017b95c --- /dev/null +++ b/csrc/kernels/attn_dispatchers.cuh @@ -0,0 +1,196 @@ +#pragma once +// Shared attention dispatchers — used by both production .cu and test .cu. +// No torch dependency; pure CUDA. + +#include +#include +#include "attn_prefill_split_q.cuh" +#include "attn_decode_split_kv.cuh" +#include "attn_paged_decode_split_kv.cuh" +#ifndef ASTRAI_NO_MMA +#include "attn_prefill_split_q_mma.cuh" +#include "attn_decode_split_kv_mma.cuh" +#include "attn_paged_decode_split_kv_mma.cuh" +#endif + +// Split-KV: compute number of splits to fill all SMs for small-batch decode. +inline int compute_num_splits(int base_blocks, int tiles_total) { + int sm_count = 0; + cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0); + int n = (2 * sm_count + base_blocks - 1) / base_blocks; + return std::max(1, std::min(n, std::min(tiles_total, 32))); +} + +// ====================================================================== +// Prefill +// ====================================================================== + +#ifndef ASTRAI_NO_MMA +template +static inline void launch_prefill_mma(AttentionParams& p) { + constexpr int WARPS = 4; + constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; + using Traits = KernelTraits; + dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch); + dim3 block(Traits::NUM_THREADS); + attn_prefill_split_q_mma_kernel<<>>(p); +} +#endif + +template +static inline void launch_prefill_scalar(AttentionParams& p) { + constexpr int G = 8, ROWS = 32, P_BC = 32; + dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); + dim3 block(G, ROWS); + attn_prefill_split_q_kernel_t<<>>(p); +} + +template +static inline void dispatch_prefill(AttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + +#ifndef ASTRAI_NO_MMA + if (is_causal) { + if (has_mask) launch_prefill_mma(p); + else launch_prefill_mma(p); + } else { + if (has_mask) launch_prefill_mma(p); + else launch_prefill_mma(p); + } +#else + if (is_causal) { + if (has_mask) launch_prefill_scalar(p); + else launch_prefill_scalar(p); + } else { + if (has_mask) launch_prefill_scalar(p); + else launch_prefill_scalar(p); + } +#endif +} + +// ====================================================================== +// Decode +// ====================================================================== + +#ifndef ASTRAI_NO_MMA +template +static inline void launch_decode_mma(AttentionParams& p, int group_size) { + int G = p.q_head / p.kv_head; + if (G >= 1 && G <= 16) { + int tiles_total = (p.kv_len + 32 - 1) / 32; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + using Traits = KernelTraits; + dim3 grid(p.kv_head, p.batch, p.num_splits); + attn_decode_split_kv_mma_kernel<<>>(p); + } else { + int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + 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<<>>(p); + } +} +#endif + +template +static inline void launch_decode_scalar(AttentionParams& p, int group_size) { + int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + 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<<>>(p); +} + +template +static inline void dispatch_decode(AttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + int group_size = p.q_head / p.kv_head; + +#ifndef ASTRAI_NO_MMA + if (is_causal) { + if (has_mask) launch_decode_mma(p, group_size); + else launch_decode_mma(p, group_size); + } else { + if (has_mask) launch_decode_mma(p, group_size); + else launch_decode_mma(p, group_size); + } +#else + if (is_causal) { + if (has_mask) launch_decode_scalar(p, group_size); + else launch_decode_scalar(p, group_size); + } else { + if (has_mask) launch_decode_scalar(p, group_size); + else launch_decode_scalar(p, group_size); + } +#endif + + attn_decode_combine_kernel<<>>(p); +} + +// ====================================================================== +// Paged Decode +// ====================================================================== + +#ifndef ASTRAI_NO_MMA +template +static inline void launch_paged_decode_mma(PagedAttentionParams& p, int group_size) { + int G = p.q_head / p.kv_head; + if (G >= 1 && G <= 16 && p.page_size >= 32) { + int tiles_total = (p.kv_len + 32 - 1) / 32; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); + constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; + using Traits = KernelTraits; + dim3 grid(p.kv_head, p.batch, p.num_splits); + paged_attn_decode_split_kv_mma_kernel<<>>(p); + } else { + int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); + dim3 grid(p.batch * p.kv_head, 1, p.num_splits); + dim3 block(32, group_size); + paged_attn_decode_split_kv_kernel<<>>(p); + } +} +#endif + +template +static inline void launch_paged_decode_scalar(PagedAttentionParams& p, int group_size) { + int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; + p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); + size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); + dim3 grid(p.batch * p.kv_head, 1, p.num_splits); + dim3 block(32, group_size); + paged_attn_decode_split_kv_kernel<<>>(p); +} + +template +static inline void dispatch_paged_decode(PagedAttentionParams& p) { + bool is_causal = (p.causal_offset >= 0); + bool has_mask = (p.use_mask && p.mask); + int group_size = p.q_head / p.kv_head; + +#ifndef ASTRAI_NO_MMA + if (is_causal) { + if (has_mask) launch_paged_decode_mma(p, group_size); + else launch_paged_decode_mma(p, group_size); + } else { + if (has_mask) launch_paged_decode_mma(p, group_size); + else launch_paged_decode_mma(p, group_size); + } +#else + if (is_causal) { + if (has_mask) launch_paged_decode_scalar(p, group_size); + else launch_paged_decode_scalar(p, group_size); + } else { + if (has_mask) launch_paged_decode_scalar(p, group_size); + else launch_paged_decode_scalar(p, group_size); + } +#endif + + paged_attn_decode_combine_kernel<<>>(p); +} diff --git a/csrc/kernels/attn_entry_utils.cuh b/csrc/kernels/attn_entry_utils.cuh index 0833bfe..ffb0388 100644 --- a/csrc/kernels/attn_entry_utils.cuh +++ b/csrc/kernels/attn_entry_utils.cuh @@ -5,13 +5,6 @@ using bf16 = __nv_bfloat16; -inline int compute_num_splits(int base_blocks, int tiles_total) { - int sm_count = 0; - cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0); - int n = (2 * sm_count + base_blocks - 1) / base_blocks; - return std::max(1, std::min(n, std::min(tiles_total, 32))); -} - // Dispatch head_dim: shared macro — avoids C++20 lambda template syntax. // Usage: DISPATCH_HEAD_DIM(hd, fn, arg) // Expands to: fn<32>(arg); fn<64>(arg); etc. diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index b68d0e7..b6f8d67 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -1,71 +1,6 @@ -#include "attn_paged_decode_split_kv.cuh" -#ifndef ASTRAI_NO_MMA -#include "attn_paged_decode_split_kv_mma.cuh" -#endif - +#include "attn_dispatchers.cuh" #include "attn_entry_utils.cuh" -#ifndef ASTRAI_NO_MMA -template -static void launch_paged_mma_decode_impl(PagedAttentionParams& p) { - using Traits = KernelTraits; - 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); - - paged_attn_decode_split_kv_mma_kernel - <<>>(p); - paged_attn_decode_combine_kernel<<>>(p); -} - -template -static void launch_paged_mma_decode(PagedAttentionParams& p) { - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - launch_paged_mma_decode_impl(p); -} -#endif - -template -static void launch_paged_scalar_decode(PagedAttentionParams& p) { - int group_size = p.q_head / p.kv_head; - int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; - p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); - alloc_split_partials(p); - - size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); - dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits); - dim3 block = dim3(32, group_size); - paged_attn_decode_split_kv_kernel<<>>(p); - paged_attn_decode_combine_kernel<<>>(p); -} - -template -static void dispatch_paged_decode(PagedAttentionParams& 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 && p.page_size >= 32) { - if (is_causal) { - if (has_mask) launch_paged_mma_decode(p); - else launch_paged_mma_decode(p); - } else { - if (has_mask) launch_paged_mma_decode(p); - else launch_paged_mma_decode(p); - } - return; - } -#endif - if (is_causal) { - if (has_mask) launch_paged_scalar_decode(p); - else launch_paged_scalar_decode(p); - } else { - if (has_mask) launch_paged_scalar_decode(p); - else launch_paged_scalar_decode(p); - } -} - torch::Tensor attn_paged_decode( torch::Tensor q, torch::Tensor page_table, @@ -86,6 +21,7 @@ torch::Tensor attn_paged_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_paged_decode, p); return O; } diff --git a/csrc/kernels/attn_prefill.cu b/csrc/kernels/attn_prefill.cu index c06b73d..d4f1c0a 100644 --- a/csrc/kernels/attn_prefill.cu +++ b/csrc/kernels/attn_prefill.cu @@ -1,53 +1,6 @@ -#include "attn_prefill_split_q.cuh" +#include "attn_dispatchers.cuh" #include "attn_entry_utils.cuh" -#ifndef ASTRAI_NO_MMA -#include "attn_prefill_split_q_mma.cuh" - -template -static void launch_mma_prefill(AttentionParams& p) { - constexpr int WARPS = 4; - constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; - using Traits = KernelTraits; - dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), - p.q_head, p.batch); - dim3 block(Traits::NUM_THREADS, 1, 1); - attn_prefill_split_q_mma_kernel<<>>(p); -} -#endif - -template -static void launch_scalar_prefill(AttentionParams& p) { - constexpr int G = 8, ROWS = 32, P_BC = 32; - dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); - dim3 block(G, ROWS, 1); - attn_prefill_split_q_kernel_t<<>>(p); -} - -template -static void dispatch_prefill(AttentionParams& p) { - bool is_causal = (p.causal_offset >= 0); - bool has_mask = (p.use_mask && p.mask); - -#ifndef ASTRAI_NO_MMA - if (is_causal) { - if (has_mask) launch_mma_prefill(p); - else launch_mma_prefill(p); - } else { - if (has_mask) launch_mma_prefill(p); - else launch_mma_prefill(p); - } -#else - if (is_causal) { - if (has_mask) launch_scalar_prefill(p); - else launch_scalar_prefill(p); - } else { - if (has_mask) launch_scalar_prefill(p); - else launch_scalar_prefill(p); - } -#endif -} - torch::Tensor attn_prefill( torch::Tensor q, torch::Tensor k, diff --git a/csrc/tests/attn_decode_test.cu b/csrc/tests/attn_decode_test.cu index cda0bc3..327f289 100644 --- a/csrc/tests/attn_decode_test.cu +++ b/csrc/tests/attn_decode_test.cu @@ -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 -static void launch_mma_decode(AttentionParams& p, DecodeScratch& sc) { - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - using Traits = KernelTraits; - 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 - <<>>(p); - attn_decode_combine_kernel<<>>(p); -} -#endif - -template -static void launch_scalar_decode(AttentionParams& 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 - <<>>(p); - attn_decode_combine_kernel<<>>(p); +static void setup_scratch(AttentionParams& 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 -static void dispatch_decode_t(AttentionParams& 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(p, sc); - else launch_mma_decode(p, sc); - } else { - if (has_mask) launch_mma_decode(p, sc); - else launch_mma_decode(p, sc); - } - return; - } -#endif - if (is_causal) { - if (has_mask) launch_scalar_decode(p, sc); - else launch_scalar_decode(p, sc); - } else { - if (has_mask) launch_scalar_decode(p, sc); - else launch_scalar_decode(p, sc); - } -} - -static void dispatch_decode(AttentionParams& p, DecodeScratch& sc) { - dispatch_by_head_dim(p.head_dim, [&]() { dispatch_decode_t(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, [&]() { dispatch_decode(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, [&]() { dispatch_decode(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; diff --git a/csrc/tests/attn_paged_decode_test.cu b/csrc/tests/attn_paged_decode_test.cu index 115bf67..5a97559 100644 --- a/csrc/tests/attn_paged_decode_test.cu +++ b/csrc/tests/attn_paged_decode_test.cu @@ -5,12 +5,8 @@ #include #include "test_utils.cuh" -#include "../kernels/attn_paged_decode_split_kv.cuh" -#ifndef ASTRAI_NO_MMA -#include "../kernels/attn_paged_decode_split_kv_mma.cuh" -#endif +#include "../kernels/attn_dispatchers.cuh" -// Copy contiguous K/V from page pool (reference gather) static void gather_kv_cpu( const bf16* h_k_pool, const bf16* h_v_pool, const int64_t* h_pt, int B, int Hkv, int kv_len, @@ -28,7 +24,8 @@ static void gather_kv_cpu( size_t src_base = (size_t)phys * page_stride + (size_t)pg_off * Hkv * head_dim + h * head_dim; - size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim; + size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + + (size_t)pos * head_dim; memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16)); memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16)); } @@ -36,59 +33,6 @@ static void gather_kv_cpu( } } -#ifndef ASTRAI_NO_MMA -template -static void launch_paged_mma_decode(PagedAttentionParams& p) { - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - using Traits = KernelTraits; - int tiles_total = (p.kv_len + 32 - 1) / 32; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - paged_attn_decode_split_kv_mma_kernel - <<>>(p); -} -#endif - -template -static void launch_paged_scalar_decode(PagedAttentionParams& p) { - int group_sz = p.q_head / p.kv_head; - int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK; - p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); - size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16); - paged_attn_decode_split_kv_kernel<<< - dim3(p.batch * p.kv_head, 1, p.num_splits), - dim3(32, group_sz), smem>>>(p); -} - -template -static void launch_paged_decode(PagedAttentionParams& p) { - bool is_causal = (p.causal_offset >= 0); - bool has_mask = (p.use_mask && p.mask); - -#ifndef ASTRAI_NO_MMA - int G_check = p.q_head / p.kv_head; - bool use_mma = G_check >= 1 && G_check <= 16 && p.page_size >= 32; - if (use_mma) { - if (is_causal) { - if (has_mask) launch_paged_mma_decode(p); - else launch_paged_mma_decode(p); - } else { - if (has_mask) launch_paged_mma_decode(p); - else launch_paged_mma_decode(p); - } - } else -#endif - { - if (is_causal) { - if (has_mask) launch_paged_scalar_decode(p); - else launch_paged_scalar_decode(p); - } else { - if (has_mask) launch_paged_scalar_decode(p); - else launch_paged_scalar_decode(p); - } - } - paged_attn_decode_combine_kernel<<>>(p); -} - template static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) { printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ", @@ -97,23 +41,22 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causa int max_pages = (kv_len + page_size - 1) / page_size; int n_phys_pages = B * max_pages; + int max_splits = 32; size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16); size_t sz_o = sz_q; size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t); - int max_splits = 32; size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float); size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float); - bf16 *d_q, *d_o_paged, *d_o_ref; + bf16 *d_q, *d_o_paged; bf16 *d_k_pool, *d_v_pool; int64_t* d_pt; float *d_op, *d_ml; cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o_paged, sz_o); - cudaMalloc(&d_o_ref, sz_o); cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_pt, sz_pt); @@ -136,7 +79,8 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causa for (int h = 0; h < Hkv; h++) { for (int d = 0; d < HEAD_DIM; d++) { float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d)); - size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d; + size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + + h * HEAD_DIM + d; h_k_pool[idx] = __float2bfloat16(v); h_v_pool[idx] = __float2bfloat16(v * 0.3f); } @@ -170,20 +114,19 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causa cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, causal ? 0 : -1); - float scale_val = 1.0f / sqrtf((float)HEAD_DIM); - PagedAttentionParams p; + PagedAttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1; p.kv_len = kv_len; p.head_dim = HEAD_DIM; p.use_mask = 0; p.causal_offset = causal ? 0 : -1; set_default_paged_strides(p); - p.num_splits = 1; p.scale = scale_val; + p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.page_size = page_size; p.max_pages = max_pages; p.page_table = d_pt; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q = d_q; p.mask = nullptr; p.o = d_o_paged; p.o_part = d_op; p.ml_part = d_ml; - launch_paged_decode(p); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p); }); cudaDeviceSynchronize(); bf16* h_o_bf16 = (bf16*)malloc(sz_o); @@ -222,7 +165,7 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causa free(h_k_cont); free(h_v_cont); free(h_q_f); free(h_k_f); free(h_v_f); free(h_o_ref); free(h_o_bf16); free(h_o_paged); - cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref); + cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt); cudaFree(d_op); cudaFree(d_ml); @@ -248,28 +191,26 @@ static const TestCase TESTS[] = { {128, 2, 32, 8, 512, 128, 0, 11}, {128, 1, 16, 2, 256, 128, 0, 12}, {128, 2, 32, 4, 512, 128, 0, 13}, - {128, 2, 8, 2, 128, 128, 1, 14}, // causal paged decode + {128, 2, 8, 2, 128, 128, 1, 14}, // causal }; static int dispatch_test(const TestCase& tc) { - bool matched = false; int r = 0; dispatch_by_head_dim(tc.head_dim, [&]() { - matched = true; r = run_test(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed); }); - return matched ? r : 1; + return r; } template static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) { int max_pages = (kv_len + page_size - 1) / page_size; int n_phys_pages = B * max_pages; + int max_splits = 32; size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t); - int max_splits = 32; size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float); size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float); @@ -296,13 +237,12 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) { cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice); free(h_pt); - float scale_val = 1.0f / sqrtf((float)HEAD_DIM); - PagedAttentionParams pa; + PagedAttentionParams pa; pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1; pa.kv_len = kv_len; pa.head_dim = HEAD_DIM; pa.use_mask = 0; pa.causal_offset = -1; set_default_paged_strides(pa); - pa.num_splits = 1; pa.scale = scale_val; + pa.scale = 1.0f / sqrtf((float)HEAD_DIM); pa.page_size = page_size; pa.max_pages = max_pages; pa.page_table = d_pt; pa.k_cache = d_k_pool; pa.v_cache = d_v_pool; @@ -310,7 +250,9 @@ static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) { pa.o_part = d_op; pa.ml_part = d_ml; const int WARMUP = 10, ITERS = 100; - auto launch = [&]() { launch_paged_decode(pa); }; + auto launch = [&]() { + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(pa); }); + }; double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM; size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM; double bytes = 2.0 * (2.0 * nKV * sizeof(bf16)); diff --git a/csrc/tests/attn_prefill_test.cu b/csrc/tests/attn_prefill_test.cu index 8fd3dab..78aebc4 100644 --- a/csrc/tests/attn_prefill_test.cu +++ b/csrc/tests/attn_prefill_test.cu @@ -1,68 +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_prefill_test.cu -o test && ./test */ #include "test_utils.cuh" -#include "../kernels/attn_prefill_split_q.cuh" -#ifndef ASTRAI_NO_MMA -#include "../kernels/attn_prefill_split_q_mma.cuh" -#endif - -#ifndef ASTRAI_NO_MMA -template -static void launch_mma_prefill(AttentionParams& p) { - constexpr int WARPS = 4; - constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; - using Traits = KernelTraits; - dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), - p.q_head, p.batch); - dim3 block(Traits::NUM_THREADS, 1, 1); - attn_prefill_split_q_mma_kernel<<>>(p); -} -#endif - -template -static void launch_scalar_prefill(AttentionParams& p) { - constexpr int G = 8, ROWS = 32, P_BC = 32; - dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); - dim3 block(G, ROWS, 1); - attn_prefill_split_q_kernel_t<<>>(p); -} - -template -static void launch_prefill_dispatch(AttentionParams& p) { - bool is_causal = (p.causal_offset >= 0); - bool has_mask = (p.use_mask && p.mask); -#ifndef ASTRAI_NO_MMA - if (is_causal) { - if (has_mask) launch_mma_prefill(p); - else launch_mma_prefill(p); - } else { - if (has_mask) launch_mma_prefill(p); - else launch_mma_prefill(p); - } -#else - if (is_causal) { - if (has_mask) launch_scalar_prefill(p); - else launch_scalar_prefill(p); - } else { - if (has_mask) launch_scalar_prefill(p); - else launch_scalar_prefill(p); - } -#endif -} - -static void dispatch_prefill(AttentionParams& p) { - switch (p.head_dim) { - case 64: launch_prefill_dispatch<64>(p); break; - case 128: launch_prefill_dispatch<128>(p); break; - default: printf("bench: unsupported D=%d\n", p.head_dim); - } -} +#include "../kernels/attn_dispatchers.cuh" // Warmed-up, CUDA-event timed throughput sweep over the production MMA path. static void bench() { @@ -105,14 +49,15 @@ static void bench() { p.scale=1.0f/sqrtf((float)D); p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - for (int i=0;i() { dispatch_prefill(p); }); }; + for (int i=0;i() { dispatch_prefill(p); }); cudaDeviceSynchronize(); double kms=now_ms()-t0; cudaError_t err=cudaGetLastError(); diff --git a/csrc/tests/test_utils.cuh b/csrc/tests/test_utils.cuh index 3a5c466..21d3e44 100644 --- a/csrc/tests/test_utils.cuh +++ b/csrc/tests/test_utils.cuh @@ -18,16 +18,6 @@ inline double now_ms() { return duration_cast(steady_clock::now().time_since_epoch()).count(); } -inline int compute_num_splits(int base_blocks, int tiles_total) { - int sm_count = 0; - cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0); - int n = (2 * sm_count + base_blocks - 1) / base_blocks; - if (n > tiles_total) n = tiles_total; - if (n > 32) n = 32; - if (n < 1) n = 1; - return n; -} - #define CUDA_CHECK(call) \ do { \ cudaError_t _e = (call); \