From 21ddead23888e0f2e05d4fdc49965943b8a68b26 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 31 Jul 2026 21:00:01 +0800 Subject: [PATCH] fix: stabilize paged decode attention kernels - zero-fill split partials so combine skips unwritten splits deterministically - skip loading masked KV in paged decode kernels to avoid 0*NaN output poisoning - zero-fill shared memory tile buffers to prevent stale NaN leaking into softmax --- csrc/kernels/attn_entry_utils.cuh | 5 +++-- csrc/kernels/attn_paged_decode_split_kv.cuh | 13 ++++++++++--- csrc/kernels/attn_paged_decode_split_kv_mma.cuh | 10 ++++++++++ 3 files changed, 23 insertions(+), 5 deletions(-) diff --git a/csrc/kernels/attn_entry_utils.cuh b/csrc/kernels/attn_entry_utils.cuh index b35730b..44d6498 100644 --- a/csrc/kernels/attn_entry_utils.cuh +++ b/csrc/kernels/attn_entry_utils.cuh @@ -1,4 +1,5 @@ #pragma once +#include #include #include #include "attn_common.h" @@ -23,8 +24,8 @@ using bf16 = __nv_bfloat16; template inline void alloc_split_partials(P& p) { auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); - auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt); - auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt); + auto o_part = torch::zeros(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt); + auto ml_part = torch::full(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, -FLT_MAX, fopt); p.o_part = (float*)o_part.data_ptr(); p.ml_part = (float*)ml_part.data_ptr(); } diff --git a/csrc/kernels/attn_paged_decode_split_kv.cuh b/csrc/kernels/attn_paged_decode_split_kv.cuh index b3ee857..a2461a1 100644 --- a/csrc/kernels/attn_paged_decode_split_kv.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv.cuh @@ -67,14 +67,17 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) partial = warp_reduce_sum(partial) * p.scale; int kv_idx = chunk_start + s; + bool masked = false; if constexpr (HasMask) { if (!p.mask[mask_base + kv_idx]) - partial = -FLT_MAX; + masked = true; } if constexpr (IsCausal) { if (kv_idx > p.causal_offset) - partial = -FLT_MAX; + masked = true; } + if (masked) + partial = -FLT_MAX; float new_m = fmaxf(m, partial); float alpha = __expf(m - new_m); @@ -85,7 +88,11 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) int logical_page = pos / p.page_size; int page_offset = pos % p.page_size; int phys_page = p.page_table[batch * p.max_pages + logical_page]; - if (phys_page >= 0) { + if (masked) { + #pragma unroll + for (int i = 0; i < hd_per_thread; i++) + acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f); + } else if (phys_page >= 0) { int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim + (int64_t)page_offset * p.kv_head * p.head_dim + (int64_t)kv_head * p.head_dim; diff --git a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh index 8289531..92ae734 100644 --- a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh @@ -31,6 +31,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams __shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD]; __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; + #pragma unroll + for (int i = lane; i < Traits::STAGES * Traits::BC * Traits::LD; i += 32) { + sK[i] = __float2bfloat16(0.0f); + sV[i] = __float2bfloat16(0.0f); + } + __syncwarp(); + const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int qra = gid; const int qrb = gid + 8; @@ -68,6 +75,9 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = (kc < p.kv_len); + if constexpr (HasMask) { + valid = valid && p.mask[batch * p.mask_b_stride + kc]; + } int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0; valid = valid && (phys_page >= 0); int page_off = kc % p.page_size;