diff --git a/csrc/tests/attn_paged_test.cu b/csrc/tests/attn_paged_test.cu index 9ec27ad..d796d3d 100644 --- a/csrc/tests/attn_paged_test.cu +++ b/csrc/tests/attn_paged_test.cu @@ -9,6 +9,9 @@ #include "test_utils.cuh" #include "../kernels/attn_dispatchers.cuh" +struct PagedDecodeDispatch { AttentionParams& p; template void operator()() { dispatch_paged_decode(p, 0); } }; +struct PagedPrefillDispatch { AttentionParams& p; template void operator()() { dispatch_paged_prefill(p, 0); } }; + // ---- CPU reference: paged decode with variable seq_lens ---- // Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D] // req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B] @@ -227,7 +230,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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -235,7 +238,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float)); for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); - const float atol = 0.02f, rtol = 0.02f; + const float atol = 0.01f, rtol = 0.01f; bool pass = true; float max_err = 0.0f; for (int i = 0; i < B * Hq * HEAD_DIM; i++) { @@ -362,7 +365,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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -370,7 +373,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float)); for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); - const float atol = 0.02f, rtol = 0.02f; + const float atol = 0.01f, rtol = 0.01f; bool pass = true; float max_err = 0.0f; for (int i = 0; i < B * Hq * HEAD_DIM; i++) { @@ -498,7 +501,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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -506,7 +509,7 @@ static int run_prefill_test(int B, int Hq, int Hkv, float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float)); for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); - const float atol = 0.02f, rtol = 0.02f; + const float atol = 0.01f, rtol = 0.01f; bool pass = true; float max_err = 0.0f; for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) { @@ -633,7 +636,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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); cudaDeviceSynchronize(); bf16* h_o_bf = (bf16*)malloc(sz_q); @@ -641,7 +644,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float)); for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); - const float atol = 0.02f, rtol = 0.02f; + const float atol = 0.01f, rtol = 0.01f; bool pass = true; float max_err = 0.0f; for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) { @@ -722,7 +725,7 @@ 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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); }; // 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. @@ -800,7 +803,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, 0); }); + dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); }; // FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D. // Non-causal: effective = q_len * kv_len. diff --git a/csrc/tests/attn_test.cu b/csrc/tests/attn_test.cu index 029b129..c09bf29 100644 --- a/csrc/tests/attn_test.cu +++ b/csrc/tests/attn_test.cu @@ -9,6 +9,9 @@ nvcc -I csrc -arch=sm_89 -O3 \ #include "test_utils.cuh" #include "../kernels/attn_dispatchers.cuh" +struct DecodeDispatch { AttentionParams& p; template void operator()() { dispatch_decode(p, 0); } }; +struct PrefillDispatch { AttentionParams& p; template void operator()() { dispatch_prefill(p, 0); } }; + // Split-K scratch (torch-free) struct DecodeScratch { float* o_part = nullptr; @@ -30,8 +33,6 @@ static void free_scratch(DecodeScratch& sc) { // ====================================================================== 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() { dispatch_decode(p, 0); }); + dispatch_by_head_dim(D, DecodeDispatch{p}); cudaDeviceSynchronize(); (void)t0; cudaError_t err=cudaGetLastError(); @@ -146,7 +147,7 @@ static void bench_decode() { 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); }); }; + auto launch = [&]() { dispatch_by_head_dim(D, DecodeDispatch{p}); }; double flops = 4.0 * B * Hq * (double)sl * D; BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops); @@ -190,7 +191,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca 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); }); + dispatch_by_head_dim(D, PrefillDispatch{p}); cudaDeviceSynchronize(); (void)t0; cudaError_t err=cudaGetLastError(); @@ -263,7 +264,7 @@ static void bench_prefill() { 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); }); }; + auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); }; for (int i=0;i