From a33ca04f6022b1b88fb0381b0ca3a69b6e1e2063 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 9 Aug 2026 00:31:47 +0800 Subject: [PATCH] fix: synchronize final decode async copy - wait for the final split-KV tile before reading shared memory - cover long decode with production context capacity --- csrc/kernels/attn_decode_split_kv_mma.cuh | 5 ++++- csrc/tests/attn_paged_test.cu | 10 +++++++--- 2 files changed, 11 insertions(+), 4 deletions(-) diff --git a/csrc/kernels/attn_decode_split_kv_mma.cuh b/csrc/kernels/attn_decode_split_kv_mma.cuh index b9ca271..853b485 100644 --- a/csrc/kernels/attn_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_decode_split_kv_mma.cuh @@ -121,7 +121,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { load_tile(ti_begin + i, i); for (int it = 0; it < ntiles; it++) { - cp_async_wait_group(); + if (it + 1 == ntiles) + cp_async_wait_group<0>(); + else + cp_async_wait_group(); __syncwarp(); process_tile(it, it & (STAGES - 1)); __syncwarp(); diff --git a/csrc/tests/attn_paged_test.cu b/csrc/tests/attn_paged_test.cu index 8c79441..9ec27ad 100644 --- a/csrc/tests/attn_paged_test.cu +++ b/csrc/tests/attn_paged_test.cu @@ -127,14 +127,15 @@ inline void print_paged_row(const char* cfg, float max_err, bool pass) { // ====================================================================== template static int run_decode_test(int B, int Hq, int Hkv, int max_seq, - int causal, int seed) { + int causal, int seed, int context_capacity = 0, + int fixed_seq_len = 0) { // Variable seq_lens per request srand(seed); std::vector seq_lens(B); for (int b = 0; b < B; b++) - seq_lens[b] = 8 + rand() % (max_seq - 8); + seq_lens[b] = fixed_seq_len ? fixed_seq_len : 8 + rand() % (max_seq - 8); int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end()); - int max_ctx = max_sl + 16; + int max_ctx = context_capacity ? context_capacity : max_sl + 16; int pool_size = B * max_ctx; int num_reqs = B + 4; @@ -844,6 +845,9 @@ int main() { fail += run_decode_test<256>(1, 2, 1, 256, 0, 9); fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10); fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11); + // Production keeps a fixed 32768-wide request table. This forces 32 + // splits, so seq_len > 512 gives each split multiple cp.async tiles. + fail += run_decode_test<64>(1, 24, 4, 1100, 0, 12, 32768, 1100); // Decode with 2D mask (regression: mixed seq_lens + HasMask) fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);