#pragma once // Shared attention dispatchers — used by both production .cu and test .cu. // No torch dependency; pure CUDA. #include #include #include "attn_warp_utils.cuh" #include "attn_prefill_split_q.cuh" #include "attn_decode_split_kv.cuh" #include "attn_paged_decode_split_kv.cuh" #ifndef ASTRAI_NO_MMA #include "attn_prefill_split_q_mma.cuh" #include "attn_decode_split_kv_mma.cuh" #include "attn_paged_decode_split_kv_mma.cuh" #endif // Split-KV: compute number of splits to fill all SMs for small-batch decode. // 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; int max_by_work = tiles_total / min_tiles_per_split; return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS))); } // ====================================================================== // Prefill // ====================================================================== #ifndef ASTRAI_NO_MMA template static inline void launch_prefill_mma(AttentionParams& p) { constexpr int WARPS = 4; constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; using Traits = KernelTraits; dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch); dim3 block(Traits::NUM_THREADS); attn_prefill_split_q_mma_kernel<<>>(p); } #endif template static inline void launch_prefill_scalar(AttentionParams& p) { constexpr int G = 8, ROWS = 32, P_BC = 32; dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch); dim3 block(G, ROWS); attn_prefill_split_q_kernel_t<<>>(p); } template static inline void dispatch_prefill(AttentionParams& p) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); #ifndef ASTRAI_NO_MMA if (is_causal) { if (has_mask) launch_prefill_mma(p); else launch_prefill_mma(p); } else { if (has_mask) launch_prefill_mma(p); else launch_prefill_mma(p); } #else if (is_causal) { if (has_mask) launch_prefill_scalar(p); else launch_prefill_scalar(p); } else { if (has_mask) launch_prefill_scalar(p); else launch_prefill_scalar(p); } #endif } // ====================================================================== // Decode // ====================================================================== #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; 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); } #endif template static inline void launch_decode_scalar(AttentionParams& p, int group_size) { int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, g); attn_decode_split_kv_kernel<<>>(p); } template static inline void dispatch_decode(AttentionParams& p) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA if (is_causal) { if (has_mask) launch_decode_mma(p, group_size); else launch_decode_mma(p, group_size); } else { if (has_mask) launch_decode_mma(p, group_size); else launch_decode_mma(p, group_size); } #else if (is_causal) { if (has_mask) launch_decode_scalar(p, group_size); else launch_decode_scalar(p, group_size); } else { if (has_mask) launch_decode_scalar(p, group_size); else launch_decode_scalar(p, group_size); } #endif attn_decode_combine_kernel<<>>(p); } // ====================================================================== // Paged Decode // ====================================================================== #ifndef ASTRAI_NO_MMA 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; 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); } } #endif template static inline void launch_paged_decode_scalar(PagedAttentionParams& p, int group_size) { 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); int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, g); paged_attn_decode_split_kv_kernel<<>>(p); } template static inline void dispatch_paged_decode(PagedAttentionParams& p) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA if (is_causal) { if (has_mask) launch_paged_decode_mma(p, group_size); else launch_paged_decode_mma(p, group_size); } else { if (has_mask) launch_paged_decode_mma(p, group_size); else launch_paged_decode_mma(p, group_size); } #else if (is_causal) { if (has_mask) launch_paged_decode_scalar(p, group_size); else launch_paged_decode_scalar(p, group_size); } else { if (has_mask) launch_paged_decode_scalar(p, group_size); else launch_paged_decode_scalar(p, group_size); } #endif paged_attn_decode_combine_kernel<<>>(p); }