From 2e29ed45d3a9ee8609903eda77a2ff991d361b14 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 27 Jul 2026 22:41:11 +0800 Subject: [PATCH] perf: shrink decode tile to BC=16 for higher occupancy MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - BC=32→16 halves smem (32KB→16KB for D=128), doubling blocks/SM (3→6) - D=256 now fits STAGES=2 double-buffer in 32KB, eliminating 176-byte spill - min_tiles_per_split=2 avoids excessive split overhead on small kv - paged decode: require page_size multiple of BC so tiles stay page-aligned Benchmark (L20 sm_89, D=128): - B=1 kv=4096: 0.0134→0.0122ms (+9% BW) - B=16 kv=2048: 0.0434→0.0352ms (+23% BW) - B=32 kv=1024: 0.0343→0.0282ms (+22% BW) --- csrc/kernels/attn_dispatchers.cuh | 34 +++++++++++++++++++++---------- 1 file changed, 23 insertions(+), 11 deletions(-) diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh index 52d2a8d..658935e 100644 --- a/csrc/kernels/attn_dispatchers.cuh +++ b/csrc/kernels/attn_dispatchers.cuh @@ -15,11 +15,15 @@ #endif // Split-KV: compute number of splits to fill all SMs for small-batch decode. -inline int compute_num_splits(int base_blocks, int tiles_total) { +// Caps splits so each split processes at least `min_tiles_per_split` tiles, +// avoiding excessive loop/prologue overhead when tiles are small. +inline int compute_num_splits(int base_blocks, int tiles_total, + int min_tiles_per_split = 1) { int sm_count = 0; cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0); int n = (2 * sm_count + base_blocks - 1) / base_blocks; - return std::max(1, std::min(n, std::min(tiles_total, MAX_SPLITS))); + int max_by_work = tiles_total / min_tiles_per_split; + return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS))); } // ====================================================================== @@ -75,15 +79,20 @@ static inline void dispatch_prefill(AttentionParams& p) { // ====================================================================== #ifndef ASTRAI_NO_MMA +// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM). +// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags), +// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates +// the 176-byte spill that STAGES=1+BC=32 suffered. template static inline void launch_decode_mma(AttentionParams& p, int group_size) { int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; int num_passes = (G + MAX_G - 1) / MAX_G; - int tiles_total = (p.kv_len + 32 - 1) / 32; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - using Traits = KernelTraits; + constexpr int BC = 16; + 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); attn_decode_split_kv_mma_kernel<<>>(p); } @@ -136,13 +145,16 @@ template static inline void launch_paged_decode_mma(PagedAttentionParams& p, int group_size) { int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; - bool page_ok = (p.page_size >= 32); + 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 + 32 - 1) / 32; - p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); - constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; - using Traits = KernelTraits; + 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 {