From 3439e3104e6c5377b0f0910d58284fd95aea3a24 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 2 Aug 2026 13:18:41 +0800 Subject: [PATCH] perf: launch CUDA kernels on torch's current stream - Thread a cudaStream_t through attn dispatchers onto torch's current stream - Scope the device guard to the entry function so kernels run on tensor device - DISPATCH_HEAD_DIM now forwards varargs so stream reaches each dispatch - Parallelize CPU reference kernels with OpenMP (paged test 31s -> 7s) - Merge decode/prefill standalone tests into attn_test.cu with correctness tables - Drop bench error column (CPU ref too slow at large sizes) - Update cuda_kernels.md for the merged test layout --- csrc/kernels/attn_decode.cu | 5 +- csrc/kernels/attn_dispatchers.cuh | 60 ++--- csrc/kernels/attn_entry_utils.cuh | 14 +- csrc/kernels/attn_paged_decode.cu | 5 +- csrc/kernels/attn_paged_prefill.cu | 5 +- csrc/kernels/attn_prefill.cu | 5 +- csrc/kernels/rotary_emb.cu | 3 +- csrc/tests/attn_decode_test.cu | 185 --------------- csrc/tests/attn_paged_test.cu | 86 +++---- csrc/tests/attn_prefill_test.cu | 169 -------------- csrc/tests/attn_test.cu | 346 +++++++++++++++++++++++++++++ csrc/tests/test_utils.cuh | 32 ++- docs/developer/cuda_kernels.md | 16 +- 13 files changed, 477 insertions(+), 454 deletions(-) delete mode 100644 csrc/tests/attn_decode_test.cu delete mode 100644 csrc/tests/attn_prefill_test.cu create mode 100644 csrc/tests/attn_test.cu diff --git a/csrc/kernels/attn_decode.cu b/csrc/kernels/attn_decode.cu index a303a65..dd4bf53 100644 --- a/csrc/kernels/attn_decode.cu +++ b/csrc/kernels/attn_decode.cu @@ -10,6 +10,9 @@ torch::Tensor attn_decode( double scale, int64_t layout ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); + auto stream = at::cuda::getCurrentCUDAStream(); + AttentionParams p; attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p); TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1"); @@ -20,7 +23,7 @@ torch::Tensor attn_decode( p.o = (bf16*)O_view.data_ptr(); alloc_split_partials(p); - DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p); + DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream); C10_CUDA_CHECK(cudaGetLastError()); return O; } diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh index ac0dda2..c954cc9 100644 --- a/csrc/kernels/attn_dispatchers.cuh +++ b/csrc/kernels/attn_dispatchers.cuh @@ -66,33 +66,33 @@ inline int compute_num_splits(int base_blocks, int tiles_total, #ifndef ASTRAI_NO_MMA template -static inline void launch_prefill_mma(AttentionParams& p) { +static inline void launch_prefill_mma(AttentionParams& p, cudaStream_t stream) { 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); + attn_prefill_split_q_mma_kernel<<>>(p); } #endif template -static inline void launch_prefill_scalar(AttentionParams& p) { +static inline void launch_prefill_scalar(AttentionParams& p, cudaStream_t stream) { 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); + attn_prefill_split_q_kernel_t<<>>(p); } template -static inline void dispatch_prefill(AttentionParams& p) { +static inline void dispatch_prefill(AttentionParams& p, cudaStream_t stream) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); #ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream); #endif } @@ -106,7 +106,7 @@ static inline void dispatch_prefill(AttentionParams& p) { // enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates // the 176-byte spill that STAGES=1+BC=32 suffered. template -static inline void launch_decode_mma(AttentionParams& p, int group_size) { +static inline void launch_decode_mma(AttentionParams& p, int group_size, cudaStream_t stream) { int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; int num_passes = (G + MAX_G - 1) / MAX_G; @@ -116,34 +116,34 @@ static inline void launch_decode_mma(AttentionParams& p, int group_size) { constexpr int STAGES = 2; using Traits = KernelTraits; dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits); - attn_decode_split_kv_mma_kernel<<>>(p); + attn_decode_split_kv_mma_kernel<<>>(p); } #endif template -static inline void launch_decode_scalar(AttentionParams& p, int group_size) { +static inline void launch_decode_scalar(AttentionParams& p, int group_size, cudaStream_t stream) { 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); int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, g); - attn_decode_split_kv_kernel<<>>(p); + attn_decode_split_kv_kernel<<>>(p); } template -static inline void dispatch_decode(AttentionParams& p) { +static inline void dispatch_decode(AttentionParams& p, cudaStream_t stream) { 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 - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream); #endif - attn_decode_combine_kernel<<>>(p); + attn_decode_combine_kernel<<>>(p); } // ====================================================================== @@ -152,7 +152,7 @@ static inline void dispatch_decode(AttentionParams& p) { #ifndef ASTRAI_NO_MMA template -static inline void launch_paged_decode_mma(PagedAttentionParams& p, int) { +static inline void launch_paged_decode_mma(PagedAttentionParams& p, cudaStream_t stream) { int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; constexpr int BC = 16; @@ -162,34 +162,34 @@ static inline void launch_paged_decode_mma(PagedAttentionParams& p, int) { constexpr int STAGES = 2; using Traits = KernelTraits; dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits); - paged_attn_decode_split_kv_mma_kernel <<>>(p); + paged_attn_decode_split_kv_mma_kernel <<>>(p); } #endif template -static inline void launch_paged_decode_scalar(PagedAttentionParams& p, int group_size) { +static inline void launch_paged_decode_scalar(PagedAttentionParams& p, int group_size, cudaStream_t stream) { int chunks_total = (p.max_seq_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); int g = min(group_size, 32); dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, g); - paged_attn_decode_split_kv_kernel<<>>(p); + paged_attn_decode_split_kv_kernel<<>>(p); } template -static inline void dispatch_paged_decode(PagedAttentionParams& p) { +static inline void dispatch_paged_decode(PagedAttentionParams& p, cudaStream_t stream) { 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 - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, 0); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream); #endif - paged_attn_decode_combine_kernel<<>>(p); + paged_attn_decode_combine_kernel<<>>(p); } // ====================================================================== @@ -198,35 +198,35 @@ static inline void dispatch_paged_decode(PagedAttentionParams& p) { #ifndef ASTRAI_NO_MMA template -static inline void launch_paged_prefill_mma(PagedAttentionParams& p) { +static inline void launch_paged_prefill_mma(PagedAttentionParams& p, cudaStream_t stream) { constexpr int WARPS = 4; constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; using Traits = KernelTraits; int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS); dim3 grid(max_q_tiles, p.q_head, p.batch); dim3 block(Traits::NUM_THREADS); - paged_attn_prefill_split_q_mma_kernel<<>>(p); + paged_attn_prefill_split_q_mma_kernel<<>>(p); } #endif template -static inline void launch_paged_prefill_scalar(PagedAttentionParams& p) { +static inline void launch_paged_prefill_scalar(PagedAttentionParams& p, cudaStream_t stream) { constexpr int G = 8, ROWS = 32, P_BC = 32; int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS; dim3 grid(max_q_tiles, p.q_head, p.batch); dim3 block(G, ROWS); paged_attn_prefill_split_q_kernel - <<>>(p); + <<>>(p); } template -static inline void dispatch_paged_prefill(PagedAttentionParams& p) { +static inline void dispatch_paged_prefill(PagedAttentionParams& p, cudaStream_t stream) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); #ifndef ASTRAI_NO_MMA - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p, stream); #else - DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p); + DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream); #endif } diff --git a/csrc/kernels/attn_entry_utils.cuh b/csrc/kernels/attn_entry_utils.cuh index 5a474b4..5575bb1 100644 --- a/csrc/kernels/attn_entry_utils.cuh +++ b/csrc/kernels/attn_entry_utils.cuh @@ -8,14 +8,14 @@ using bf16 = __nv_bfloat16; // 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. -#define DISPATCH_HEAD_DIM(hd, fn, arg) \ +// Usage: DISPATCH_HEAD_DIM(hd, fn, args...) +// Expands to: fn<32>(args...); fn<64>(args...); etc. +#define DISPATCH_HEAD_DIM(hd, fn, ...) \ switch (hd) { \ - case 32: fn<32>(arg); break; \ - case 64: fn<64>(arg); break; \ - case 128: fn<128>(arg); break; \ - case 256: fn<256>(arg); break; \ + case 32: fn<32>(__VA_ARGS__); break; \ + case 64: fn<64>(__VA_ARGS__); break; \ + case 128: fn<128>(__VA_ARGS__); break; \ + case 256: fn<256>(__VA_ARGS__); break; \ default: \ TORCH_CHECK(false, "unsupported head_dim ", hd, \ " (supported: 32, 64, 128, 256)"); \ diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index e8b4f44..97dc44c 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -13,6 +13,9 @@ torch::Tensor attn_paged_decode( int64_t causal_offset, double scale ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); + auto stream = at::cuda::getCurrentCUDAStream(); + PagedAttentionParams p; attn_pack_paged_decode_params(q, k_cache, v_cache, req_to_token, req_pool_indices, kv_indptr, @@ -22,7 +25,7 @@ torch::Tensor attn_paged_decode( p.o = (bf16*)O.data_ptr(); alloc_split_partials(p); - DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p); + DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream); C10_CUDA_CHECK(cudaGetLastError()); return O; } diff --git a/csrc/kernels/attn_paged_prefill.cu b/csrc/kernels/attn_paged_prefill.cu index 7f77fa2..4848ced 100644 --- a/csrc/kernels/attn_paged_prefill.cu +++ b/csrc/kernels/attn_paged_prefill.cu @@ -14,6 +14,9 @@ torch::Tensor attn_paged_prefill( int64_t causal_offset, double scale ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); + auto stream = at::cuda::getCurrentCUDAStream(); + PagedAttentionParams p; attn_pack_paged_prefill_params(q, k_cache, v_cache, req_to_token, req_pool_indices, @@ -23,7 +26,7 @@ torch::Tensor attn_paged_prefill( auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); p.o = (bf16*)O.data_ptr(); - DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p); + DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream); C10_CUDA_CHECK(cudaGetLastError()); return O; } diff --git a/csrc/kernels/attn_prefill.cu b/csrc/kernels/attn_prefill.cu index a0a0cb8..416d680 100644 --- a/csrc/kernels/attn_prefill.cu +++ b/csrc/kernels/attn_prefill.cu @@ -10,6 +10,9 @@ torch::Tensor attn_prefill( double scale, int64_t layout ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); + auto stream = at::cuda::getCurrentCUDAStream(); + AttentionParams p; attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p); TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16"); @@ -18,7 +21,7 @@ torch::Tensor attn_prefill( auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; p.o = (bf16*)O_view.data_ptr(); - DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p); + DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream); C10_CUDA_CHECK(cudaGetLastError()); return O; } diff --git a/csrc/kernels/rotary_emb.cu b/csrc/kernels/rotary_emb.cu index 3e717c9..8e9c76e 100644 --- a/csrc/kernels/rotary_emb.cu +++ b/csrc/kernels/rotary_emb.cu @@ -49,6 +49,7 @@ torch::Tensor rotary_emb( torch::Tensor freqs_cis ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); + auto stream = at::cuda::getCurrentCUDAStream(); TORCH_CHECK(x.is_cuda(), "x must be on CUDA"); TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA"); @@ -77,7 +78,7 @@ torch::Tensor rotary_emb( int block = 256; int grid = std::min((total + block - 1) / block, 1024); - rotary_emb_kernel<<>>( + rotary_emb_kernel<<>>( reinterpret_cast(x.data_ptr()), freqs_cis.data_ptr(), reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), diff --git a/csrc/tests/attn_decode_test.cu b/csrc/tests/attn_decode_test.cu deleted file mode 100644 index e31a4cb..0000000 --- a/csrc/tests/attn_decode_test.cu +++ /dev/null @@ -1,185 +0,0 @@ -/* -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_dispatchers.cuh" - -// Split-K scratch (torch-free) -struct DecodeScratch { - float* o_part = nullptr; - float* ml_part = nullptr; -}; - -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)); -} - -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}, - {1, 32, 4, 1024, 128}, - {1, 32, 4, 2048, 128}, - {1, 32, 4, 4096, 128}, - {16, 32, 4, 2048, 128}, - {32, 32, 4, 1024, 128}, - }; - const int WARMUP = 10, ITERS = 100; - printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); - print_bench_header(); - - for (int ci = 0; ci < 6; ci++) { - int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2]; - int sl = cfgs[ci][3], D = cfgs[ci][4]; - size_t nQ = (size_t)B * Hq * D; - size_t nKV = (size_t)B * Hk * sl * D; - - bf16 *dQ, *dK, *dV, *dO; - cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2); - cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2); - size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big]; - for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf()); - cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice); - for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf()); - cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice); - for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf()); - cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice); - delete[] tmp; - - AttentionParams p; - p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl; - p.head_dim = D; p.use_mask = 0; p.causal_offset = -1; - p.scale = 1.0f / sqrtf((float)D); - set_default_strides(p); - p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO; - - DecodeScratch sc; - setup_scratch(p, sc); - p.o_part = sc.o_part; p.ml_part = sc.ml_part; - - 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); - - char cfg[64]; - snprintf(cfg, sizeof(cfg), - "B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d", - B, Hq, Hk, 1, sl, D, 0); - print_bench_row(cfg, r); - - cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO); - free_scratch(sc); - } -} - -static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) { - int gs = Hq / Hk; - printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n", - B,Hq,Hk,sl,D,gs,causal); - - size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; - float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; - for (size_t i=0;i p; - p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; - p.use_mask=0; p.causal_offset=causal?0:-1; - p.scale=1.0f/sqrtf((float)D); - set_default_strides(p); - p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - - DecodeScratch sc; - setup_scratch(p, sc); - p.o_part = sc.o_part; p.ml_part = sc.ml_part; - - double t0=now_ms(); - dispatch_by_head_dim(D, [&]() { dispatch_decode(p); }); - cudaDeviceSynchronize(); - double kms=now_ms()-t0; - cudaError_t err=cudaGetLastError(); - if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} - - bf16* hOut=new bf16[nQ]; - cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); - - float* ref=new float[nQ]; - cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1); - - float max_abs_err=0, max_rel_err=0; - for (size_t i=0;imax_abs_err) max_abs_err=err; - float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); - if(rel>max_rel_err) max_rel_err=rel; - } - const float atol=0.01f, rtol=0.01f; - bool pass=true; - for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } - } - printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n", - kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL"); - - cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask); - free_scratch(sc); - delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; - - return pass ? 0 : 1; -} - -int main() { - const int configs[][6] = { - {1, 2, 1, 64, 32, 0}, - {1, 32, 4, 512, 128, 0}, - {1, 32, 4, 1024, 128, 0}, - {1, 32, 4, 512, 128, 1}, - }; - int n_cfgs = sizeof(configs) / sizeof(configs[0]); - int fail = 0; - - for (int ci = 0; ci < n_cfgs; ci++) { - int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2]; - int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5]; - fail += run_test(B, Hq, Hk, sl, D, causal); - if (fail) break; - } - - if (fail) { - printf("FAILED\n"); - return fail; - } - printf("All tests passed!\n"); - bench(); - return 0; -} diff --git a/csrc/tests/attn_paged_test.cu b/csrc/tests/attn_paged_test.cu index dc58882..4e1f502 100644 --- a/csrc/tests/attn_paged_test.cu +++ b/csrc/tests/attn_paged_test.cu @@ -1,6 +1,7 @@ // Compile: // nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \ -// --extra-device-vectorization csrc/tests/attn_paged_test.cu \ +// --extra-device-vectorization -Xcompiler -fopenmp \ +// csrc/tests/attn_paged_test.cu \ // -o /tmp/test_paged && /tmp/test_paged #include @@ -24,6 +25,7 @@ static void cpu_paged_decode_ref( for (int b = 0; b < B; b++) { int seq_len = kv_indptr[b + 1] - kv_indptr[b]; int64_t req_idx = req_pool_indices[b]; + #pragma omp parallel for schedule(dynamic) for (int h = 0; h < Hq; h++) { int kv_h = h / n_rep; float mv = -INFINITY, sv = 0.0f; @@ -74,9 +76,10 @@ static void cpu_paged_prefill_ref( int q_len = qo_indptr[b + 1] - qo_indptr[b]; int causal_off = seq_len - q_len; int64_t req_idx = req_pool_indices[b]; + #pragma omp parallel for collapse(2) schedule(dynamic) for (int h = 0; h < Hq; h++) { - int kv_h = h / n_rep; for (int qi = 0; qi < q_len; qi++) { + int kv_h = h / n_rep; float mv = -INFINITY, sv = 0.0f; float accum[256] = {0.0f}; int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len; @@ -106,6 +109,19 @@ static void cpu_paged_prefill_ref( } } +// ---- paged validation table (kernel vs CPU ref, abs error only) ---- +inline void print_paged_header() { + printf("%-58s | %11s | %6s\n", + "config", "max_err", "result"); + printf("----------------------------------------------------------------" + "--------------------------------\n"); +} + +inline void print_paged_row(const char* cfg, float max_err, bool pass) { + printf("%-58s | %11.3e | %s\n", + cfg, max_err, pass ? "PASS" : "FAIL"); +} + // ====================================================================== // DECODE TEST // ====================================================================== @@ -123,10 +139,9 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, int pool_size = B * max_ctx; int num_reqs = B + 4; - printf("DECODE B=%d Hq=%d Hkv=%d D=%d seqs=[", B, Hq, Hkv, HEAD_DIM); - for (int b = 0; b < B; b++) printf("%d%s", seq_lens[b], b < B-1 ? "," : ""); - printf("] causal=%d ... ", causal); - fflush(stdout); + char cfg[80]; + snprintf(cfg, sizeof(cfg), "DECODE B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d", + B, Hq, Hkv, HEAD_DIM, max_sl, causal); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); @@ -211,7 +226,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, p.kv_indptr = d_kvi; p.qo_indptr = nullptr; p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p, 0); }); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -228,8 +243,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; } } - if (pass) printf("PASS (max_err=%.4e)\n", max_err); - else printf("FAIL (max_err=%.4e)\n", max_err); + print_paged_row(cfg, max_err, pass); free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi); free(h_kvi); free(h_q_f); free(h_k_f); free(h_v_f); @@ -254,8 +268,9 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, int pool_size = B * max_ctx; int num_reqs = B + 4; - printf("DECODE-MASK B=%d Hq=%d Hkv=%d D=%d max_sl=%d ... ", B, Hq, Hkv, HEAD_DIM, max_sl); - fflush(stdout); + char cfg[80]; + snprintf(cfg, sizeof(cfg), "DECODE-MASK B=%d Hq=%d Hkv=%d D=%d max_sl=%d", + B, Hq, Hkv, HEAD_DIM, max_sl); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); @@ -346,7 +361,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, p.kv_indptr = d_kvi; p.qo_indptr = nullptr; p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p, 0); }); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -363,8 +378,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; } } - if (pass) printf("PASS (max_err=%.4e)\n", max_err); - else printf("FAIL (max_err=%.4e)\n", max_err); + print_paged_row(cfg, max_err, pass); free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi); free(h_kvi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f); @@ -393,12 +407,9 @@ static int run_prefill_test(int B, int Hq, int Hkv, int pool_size = B * max_ctx; int num_reqs = B + 4; - printf("PREFILL B=%d Hq=%d Hkv=%d D=%d q_lens=[", B, Hq, Hkv, HEAD_DIM); - for (int b = 0; b < B; b++) printf("%d%s", q_lens[b], b < B-1 ? "," : ""); - printf("] kv_lens=["); - for (int b = 0; b < B; b++) printf("%d%s", kv_lens[b], b < B-1 ? "," : ""); - printf("] causal=%d ... ", causal); - fflush(stdout); + char cfg[80]; + snprintf(cfg, sizeof(cfg), "PREFILL B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d", + B, Hq, Hkv, HEAD_DIM, max_sl, causal); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); @@ -486,7 +497,7 @@ static int run_prefill_test(int B, int Hq, int Hkv, p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p, 0); }); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -503,8 +514,7 @@ static int run_prefill_test(int B, int Hq, int Hkv, if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; } } - if (pass) printf("PASS (max_err=%.4e)\n", max_err); - else printf("FAIL (max_err=%.4e)\n", max_err); + print_paged_row(cfg, max_err, pass); free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi); free(h_q_f); free(h_k_f); free(h_v_f); @@ -527,7 +537,9 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { int pool_size = B * max_ctx; int num_reqs = B + 4; - printf("PREFILL-MASK Hq=%d Hkv=%d D=%d q_len=%d ... ", Hq, Hkv, HEAD_DIM, q_len); + char cfg[80]; + snprintf(cfg, sizeof(cfg), "PREFILL-MASK Hq=%d Hkv=%d D=%d q_len=%d", + Hq, Hkv, HEAD_DIM, q_len); fflush(stdout); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); @@ -620,7 +632,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p, 0); }); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -637,8 +649,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; } } - if (pass) printf("PASS (max_err=%.4e)\n", max_err); - else printf("FAIL (max_err=%.4e)\n", max_err); + print_paged_row(cfg, max_err, pass); free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f); @@ -710,15 +721,12 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) { p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; auto launch = [&]() { - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_decode(p, 0); }); }; // Decode: q_len=1, query is the last token → attends to all [0, seq_len). // FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D. double flops = 4.0 * B * Hq * (double)seq_len * HEAD_DIM; - // HBM: K+V read (Q/O negligible for decode). - size_t nKV = (size_t)B * Hkv * seq_len * HEAD_DIM; - double bytes = 2.0 * nKV * sizeof(bf16); - BenchResult r = bench_kernel(launch, 10, 100, flops, bytes); + BenchResult r = bench_kernel(launch, 3, 10, flops); char cfg[64]; snprintf(cfg, sizeof(cfg), "DEC B=%2d Hq=%2d Hk=%d kv=%4d D=%3d", @@ -791,7 +799,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; auto launch = [&]() { - dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p); }); + dispatch_by_head_dim(HEAD_DIM, [&]() { dispatch_paged_prefill(p, 0); }); }; // FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D. // Non-causal: effective = q_len * kv_len. @@ -807,11 +815,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau eff_kv = (double)q_len * kv_len; } double flops = 4.0 * B * Hq * eff_kv * HEAD_DIM; - // HBM: Q read + K read + V read + O write. - size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM; - size_t nQ = (size_t)total_q * Hq * HEAD_DIM; - double bytes = (2.0 * nQ + 2.0 * nKV) * sizeof(bf16); - BenchResult r = bench_kernel(launch, 10, 100, flops, bytes); + BenchResult r = bench_kernel(launch, 3, 10, flops); char cfg[80]; snprintf(cfg, sizeof(cfg), "PRE B=%d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d c=%d", @@ -827,7 +831,8 @@ int main() { int fail = 0; // ===== DECODE TESTS ===== - printf("=== Paged Decode Tests ===\n\n"); + printf("=== Paged Decode Tests ===\n"); + print_paged_header(); fail += run_decode_test<128>(1, 32, 4, 512, 0, 1); fail += run_decode_test<128>(1, 32, 4, 1024, 0, 2); fail += run_decode_test<128>(4, 32, 4, 512, 0, 3); @@ -848,7 +853,8 @@ int main() { if (fail) { printf("\nFAILED decode tests\n"); return fail; } // ===== PREFILL TESTS ===== - printf("\n=== Paged Prefill Tests ===\n\n"); + printf("\n=== Paged Prefill Tests ===\n"); + print_paged_header(); // Single request, pure prefill (q_len == kv_len) { std::vector ql = {512}; diff --git a/csrc/tests/attn_prefill_test.cu b/csrc/tests/attn_prefill_test.cu deleted file mode 100644 index ec137b4..0000000 --- a/csrc/tests/attn_prefill_test.cu +++ /dev/null @@ -1,169 +0,0 @@ -/* -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_dispatchers.cuh" - -// Warmed-up, CUDA-event timed throughput sweep over the production MMA path. -static void bench() { - const int cfgs[][7] = { - {1,32,4,512,512,128,0}, - {1,32,4,1024,1024,128,0}, - {1,32,4,2048,2048,128,0}, - {1,32,4,2048,2048,128,1}, - {4,32,4,2048,2048,128,1}, - {1,32,4,4096,4096,128,1}, - }; - int n = sizeof(cfgs)/sizeof(cfgs[0]); - const int WARMUP = 10, ITERS = 50; - printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); - printf("%-46s | %10s | %10s | %10s\n", - "config", "latency", "bandwidth", "throughput"); - printf("---------------------------------------------------------------" - "----------------------------\n"); - - for (int ci = 0; ci < n; ci++) { - int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2]; - int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6]; - size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D; - - bf16 *dQ,*dK,*dV,*dO,*tmp; - cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2); - cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2); - size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big]; - for (size_t i=0;i p; - p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; - p.use_mask=0; p.causal_offset=causal?0:-1; - set_default_strides(p); - p.scale=1.0f/sqrtf((float)D); - p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - - auto launch = [&]() { dispatch_by_head_dim(D, [&]() { dispatch_prefill(p); }); }; - for (int i=0;i p; - p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; - p.use_mask=0; p.causal_offset=causal?0:-1; - set_default_strides(p); - p.scale=1.0f/sqrtf((float)D); - p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; - - double t0=now_ms(); - dispatch_by_head_dim(D, [&]() { dispatch_prefill(p); }); - cudaDeviceSynchronize(); - double kms=now_ms()-t0; - cudaError_t err=cudaGetLastError(); - if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} - - bf16* hOut=new bf16[nQ]; - cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); - - float* ref=new float[nQ]; - cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); - - float max_abs_err=0, max_rel_err=0; - for (size_t i=0;imax_abs_err) max_abs_err=err; - float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); - if(rel>max_rel_err) max_rel_err=rel; - } - const float atol=0.01f, rtol=0.01f; - bool pass=true; - for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } - } - printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n", - kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL"); - - cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO); - delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; - - return pass ? 0 : 1; -} - -int main() { - const int configs[][7] = { - {1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal - {1,32,4,512,512,128,0}, // standard - {1,32,4,128,256,128,0}, // medium - {1,4,2,256,256,128,1}, // causal - }; - int n_configs = sizeof(configs) / sizeof(configs[0]); - int fail = 0; - - for (int ci = 0; ci < n_configs; ci++) { - int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2]; - int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5]; - int causal=configs[ci][6]; - fail += run_test(B, Hq, Hk, ql, kl, D, causal); - if (fail) break; - } - - if (fail) { - printf("FAILED\n"); - return fail; - } - printf("All tests passed!\n"); - bench(); - return 0; -} diff --git a/csrc/tests/attn_test.cu b/csrc/tests/attn_test.cu new file mode 100644 index 0000000..b1aeeda --- /dev/null +++ b/csrc/tests/attn_test.cu @@ -0,0 +1,346 @@ +/* +Pure-C test — uses shared dispatcher. Combines the decode (split-KV) and +prefill (split-Q) correctness checks + benchmarks into one binary. +nvcc -I csrc -arch=sm_89 -O3 \ + --use_fast_math --ptxas-options=-O3 --extra-device-vectorization \ + -Xcompiler -fopenmp csrc/tests/attn_test.cu -o test && ./test +*/ + +#include "test_utils.cuh" +#include "../kernels/attn_dispatchers.cuh" + +// Split-K scratch (torch-free) +struct DecodeScratch { + float* o_part = nullptr; + float* ml_part = nullptr; +}; + +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)); +} + +static void free_scratch(DecodeScratch& sc) { + cudaFree(sc.o_part); cudaFree(sc.ml_part); +} + +// ====================================================================== +// DECODE +// ====================================================================== + +static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) { + int gs = Hq / Hk; + + size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; + float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; + for (size_t i=0;i p; + p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; + p.use_mask=0; p.causal_offset=causal?0:-1; + p.scale=1.0f/sqrtf((float)D); + set_default_strides(p); + p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; + + DecodeScratch sc; + setup_scratch(p, sc); + p.o_part = sc.o_part; p.ml_part = sc.ml_part; + + double t0=now_ms(); + dispatch_by_head_dim(D, [&]() { dispatch_decode(p, 0); }); + cudaDeviceSynchronize(); + (void)t0; + cudaError_t err=cudaGetLastError(); + if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} + + bf16* hOut=new bf16[nQ]; + cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); + + float* ref=new float[nQ]; + cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1); + + float max_abs_err=0, max_rel_err=0; + for (size_t i=0;imax_abs_err) max_abs_err=err; + float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); + if(rel>max_rel_err) max_rel_err=rel; + } + const float atol=0.01f, rtol=0.01f; + bool pass=true; + for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } + } + char cfg[64]; + snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d seq=%4d D=%3d causal=%d", + B, Hq, Hk, sl, D, causal); + print_test_row(cfg, max_abs_err, max_rel_err, pass); + + cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask); + free_scratch(sc); + delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; + + return pass ? 0 : 1; +} + +static void bench_decode() { + const int cfgs[][5] = { + {1, 32, 4, 512, 128}, + {1, 32, 4, 1024, 128}, + {1, 32, 4, 2048, 128}, + {1, 32, 4, 4096, 128}, + {16, 32, 4, 2048, 128}, + {32, 32, 4, 1024, 128}, + }; + const int WARMUP = 3, ITERS = 10; + printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); + print_bench_header(); + + for (int ci = 0; ci < 6; ci++) { + int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2]; + int sl = cfgs[ci][3], D = cfgs[ci][4]; + size_t nQ = (size_t)B * Hq * D; + size_t nKV = (size_t)B * Hk * sl * D; + + bf16 *dQ, *dK, *dV, *dO; + cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2); + cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2); + size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big]; + for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf()); + cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice); + for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf()); + cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice); + for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf()); + cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice); + delete[] tmp; + + AttentionParams p; + p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl; + p.head_dim = D; p.use_mask = 0; p.causal_offset = -1; + p.scale = 1.0f / sqrtf((float)D); + set_default_strides(p); + p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO; + + DecodeScratch sc; + setup_scratch(p, sc); + p.o_part = sc.o_part; p.ml_part = sc.ml_part; + + auto launch = [&]() { dispatch_by_head_dim(D, [&]() { dispatch_decode(p, 0); }); }; + double flops = 4.0 * B * Hq * (double)sl * D; + BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops); + + char cfg[64]; + snprintf(cfg, sizeof(cfg), + "B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d", + B, Hq, Hk, 1, sl, D, 0); + print_bench_row(cfg, r); + + cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO); + free_scratch(sc); + } +} + +// ====================================================================== +// PREFILL +// ====================================================================== + +static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) { + size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D; + float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; + for (size_t i=0;i p; + p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; + p.use_mask=0; p.causal_offset=causal?0:-1; + set_default_strides(p); + p.scale=1.0f/sqrtf((float)D); + p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; + + double t0=now_ms(); + dispatch_by_head_dim(D, [&]() { dispatch_prefill(p, 0); }); + cudaDeviceSynchronize(); + (void)t0; + cudaError_t err=cudaGetLastError(); + if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;} + + bf16* hOut=new bf16[nQ]; + cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost); + + float* ref=new float[nQ]; + cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); + + float max_abs_err=0, max_rel_err=0; + for (size_t i=0;imax_abs_err) max_abs_err=err; + float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); + if(rel>max_rel_err) max_rel_err=rel; + } + const float atol=0.01f, rtol=0.01f; + bool pass=true; + for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } + } + char cfg[64]; + snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d", + B, Hq, Hk, ql, kl, D, causal); + print_test_row(cfg, max_abs_err, max_rel_err, pass); + + cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO); + delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; + + return pass ? 0 : 1; +} + +static void bench_prefill() { + const int cfgs[][7] = { + {1,32,4,512,512,128,0}, + {1,32,4,1024,1024,128,0}, + {1,32,4,2048,2048,128,0}, + {1,32,4,2048,2048,128,1}, + {4,32,4,2048,2048,128,1}, + {1,32,4,4096,4096,128,1}, + }; + int n = sizeof(cfgs)/sizeof(cfgs[0]); + const int WARMUP = 3, ITERS = 10; + printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); + print_bench_header(); + + for (int ci = 0; ci < n; ci++) { + int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2]; + int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6]; + size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D; + + bf16 *dQ,*dK,*dV,*dO,*tmp; + cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2); + cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2); + size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big]; + for (size_t i=0;i p; + p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; + p.use_mask=0; p.causal_offset=causal?0:-1; + set_default_strides(p); + p.scale=1.0f/sqrtf((float)D); + p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; + + auto launch = [&]() { dispatch_by_head_dim(D, [&]() { dispatch_prefill(p, 0); }); }; + for (int i=0;i BenchResult bench_kernel(Fn launch, int warmup, int iters, - double flops, double bytes) { + double flops) { for (int i = 0; i < warmup; i++) launch(); cudaDeviceSynchronize(); cudaError_t err = cudaGetLastError(); if (err != cudaSuccess) { printf("CUDA error before bench: %s\n", cudaGetErrorString(err)); - return {0, 0, 0}; + return {0, 0}; } cudaEvent_t s, e; @@ -52,19 +51,33 @@ BenchResult bench_kernel(Fn launch, int warmup, int iters, float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters; cudaEventDestroy(s); cudaEventDestroy(e); - return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12}; + return {ms, flops / (ms * 1e-3) / 1e12}; } inline void print_bench_header() { - printf("%-46s | %10s | %10s | %10s\n", - "config", "latency", "bandwidth", "throughput"); + printf("%-46s | %10s | %10s\n", + "config", "latency", "TFLOP/s"); printf("---------------------------------------------------------------" "----------------------------\n"); } inline void print_bench_row(const char* cfg, const BenchResult& r) { - printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n", - cfg, r.ms, r.gbps, r.tflops); + printf("%-46s | %7.4f ms | %6.2f\n", + cfg, r.ms, r.tflops); +} + +// ---- validation table (kernel vs CPU reference) ---- +inline void print_test_header() { + printf("%-46s | %11s | %11s | %6s\n", + "config", "max_abs_err", "max_rel_err", "result"); + printf("----------------------------------------------------------------" + "----------------------------\n"); +} + +inline void print_test_row(const char* cfg, float max_abs_err, + float max_rel_err, bool pass) { + printf("%-46s | %11.3e | %11.3e | %s\n", + cfg, max_abs_err, max_rel_err, pass ? "PASS" : "FAIL"); } template @@ -135,9 +148,10 @@ static void cpu_attention_ref( float scale = 1.0f / sqrtf((float)D); int n_rep = Hq / Hk; for (int b = 0; b < B; b++) { + #pragma omp parallel for collapse(2) schedule(dynamic) for (int h = 0; h < Hq; h++) { - int kv_h = h / n_rep; for (int qi = 0; qi < q_len; qi++) { + int kv_h = h / n_rep; float mv = -INFINITY, sv = 0.0f; float accum[256] = {0.0f}; int lim = kv_len; diff --git a/docs/developer/cuda_kernels.md b/docs/developer/cuda_kernels.md index 6773914..63a5f7b 100644 --- a/docs/developer/cuda_kernels.md +++ b/docs/developer/cuda_kernels.md @@ -117,23 +117,22 @@ Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment ```bash nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \ --ptxas-options=-O3,-v --extra-device-vectorization \ - csrc/tests/attn_decode_test.cu -o /tmp/test && /tmp/test + -Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test ``` Test files: -- `attn_decode_test.cu` — basic decode kernel -- `attn_paged_decode_test.cu` — paged decode kernel -- `attn_prefill_test.cu` — prefill kernel +- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks) +- `attn_paged_test.cu` — paged decode/prefill kernels ## Benchmarks Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86. -Reproduce: +Reproduce (decode + prefill in `attn_test.cu`, paged in `attn_paged_test.cu`): ```bash nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \ --ptxas-options=-O3,-v --extra-device-vectorization \ - csrc/tests/attn__test.cu -o /tmp/test && /tmp/test + -Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test ``` ## Known Optimization Targets @@ -165,9 +164,8 @@ csrc/ │ └── attn_warp_utils.cuh # Warp-level utilities └── tests/ ├── test_utils.cuh # Shared test utilities - ├── attn_decode_test.cu # Decode kernel test - ├── attn_paged_decode_test.cu # Paged decode test - └── attn_prefill_test.cu # Prefill kernel test + ├── attn_test.cu # Decode + prefill kernels + └── attn_paged_test.cu # Paged decode/prefill kernels ``` Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.