diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh index 658935e..68060e5 100644 --- a/csrc/kernels/attn_dispatchers.cuh +++ b/csrc/kernels/attn_dispatchers.cuh @@ -146,25 +146,13 @@ static inline void launch_paged_decode_mma(PagedAttentionParams& p, int gr int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; constexpr int BC = 16; - // page_size must be >= BC and a multiple of BC so a BC-wide tile never - // straddles two pages (the kernel does one page-table lookup per tile). - bool page_ok = (p.page_size >= BC) && (p.page_size % BC == 0); - if (G >= 1 && page_ok) { - int num_passes = (G + MAX_G - 1) / MAX_G; - int tiles_total = (p.kv_len + BC - 1) / BC; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2); - 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); - } else { - int chunks_total = (p.kv_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); - dim3 grid(p.batch * p.kv_head, 1, p.num_splits); - dim3 block(32, group_size); - paged_attn_decode_split_kv_kernel<<>>(p); - } + int num_passes = (G + MAX_G - 1) / MAX_G; + int tiles_total = (p.kv_len + BC - 1) / BC; + p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2); + 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); } #endif diff --git a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh index bc1eb18..8289531 100644 --- a/csrc/kernels/attn_paged_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv_mma.cuh @@ -55,19 +55,21 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM; // ---- Load tile lambda: paged addressing ---- + // Unified per-element page-table lookup. When page_size >= BC, all + // elements in a tile share the same page, so the lookup is redundant + // but harmless (L1-cached). This avoids a branch on page_size. auto load_tile = [&](int ti, int buf) { int kv0 = ti * Traits::BC; bf16* dK = sK + buf * Traits::BC * Traits::LD; bf16* dV = sV + buf * Traits::BC * Traits::LD; - int logical_page = kv0 / p.page_size; - int phys_page = p.page_table[batch * p.max_pages + logical_page]; - bool page_valid = (phys_page >= 0); #pragma unroll for (int i = lane * Traits::VEC; i < Traits::TOTAL; i += Traits::NUM_THREADS * Traits::VEC) { int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; - bool valid = (kc < p.kv_len) && page_valid; + bool valid = (kc < p.kv_len); + 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; int64_t gmem_base = (int64_t)phys_page * page_stride + (int64_t)page_off * pos_stride