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
This commit is contained in:
@@ -1,4 +1,5 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
|
#include <float.h>
|
||||||
#include <torch/extension.h>
|
#include <torch/extension.h>
|
||||||
#include <c10/cuda/CUDAGuard.h>
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
@@ -23,8 +24,8 @@ using bf16 = __nv_bfloat16;
|
|||||||
template<typename P>
|
template<typename P>
|
||||||
inline void alloc_split_partials(P& p) {
|
inline void alloc_split_partials(P& p) {
|
||||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
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 o_part = torch::zeros(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 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.o_part = (float*)o_part.data_ptr();
|
||||||
p.ml_part = (float*)ml_part.data_ptr();
|
p.ml_part = (float*)ml_part.data_ptr();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -67,14 +67,17 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
partial = warp_reduce_sum(partial) * p.scale;
|
partial = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
int kv_idx = chunk_start + s;
|
int kv_idx = chunk_start + s;
|
||||||
|
bool masked = false;
|
||||||
if constexpr (HasMask) {
|
if constexpr (HasMask) {
|
||||||
if (!p.mask[mask_base + kv_idx])
|
if (!p.mask[mask_base + kv_idx])
|
||||||
partial = -FLT_MAX;
|
masked = true;
|
||||||
}
|
}
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
if (kv_idx > p.causal_offset)
|
if (kv_idx > p.causal_offset)
|
||||||
partial = -FLT_MAX;
|
masked = true;
|
||||||
}
|
}
|
||||||
|
if (masked)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
float new_m = fmaxf(m, partial);
|
float new_m = fmaxf(m, partial);
|
||||||
float alpha = __expf(m - new_m);
|
float alpha = __expf(m - new_m);
|
||||||
@@ -85,7 +88,11 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
int logical_page = pos / p.page_size;
|
int logical_page = pos / p.page_size;
|
||||||
int page_offset = pos % p.page_size;
|
int page_offset = pos % p.page_size;
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
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 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)page_offset * p.kv_head * p.head_dim
|
||||||
+ (int64_t)kv_head * p.head_dim;
|
+ (int64_t)kv_head * p.head_dim;
|
||||||
|
|||||||
@@ -31,6 +31,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
__shared__ __align__(16) bf16 sV[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 q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||||
const int qra = gid;
|
const int qra = gid;
|
||||||
const int qrb = gid + 8;
|
const int qrb = gid + 8;
|
||||||
@@ -68,6 +75,9 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = (kc < p.kv_len);
|
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;
|
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
|
||||||
valid = valid && (phys_page >= 0);
|
valid = valid && (phys_page >= 0);
|
||||||
int page_off = kc % p.page_size;
|
int page_off = kc % p.page_size;
|
||||||
|
|||||||
Reference in New Issue
Block a user