diff --git a/csrc/kernels/attn_kv_source.cuh b/csrc/kernels/attn_kv_source.cuh index 168faed..03ee5fb 100644 --- a/csrc/kernels/attn_kv_source.cuh +++ b/csrc/kernels/attn_kv_source.cuh @@ -183,3 +183,25 @@ struct PagedKV { return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], ok}; } }; + +// ---- Q-tile mapping broadcast ---- +// The flat grid maps a tile index to (batch, local q_tile) once per block +// (thread 0 computes it via KV::map_q_tile), then publishes the result to +// the whole block through shared memory. Tail blocks past the ragged tile +// total get batch == -1 and the whole block exits before any work. +// Usage in a kernel: __shared__ QTileMapper qmap; +template +struct QTileMapper { + int batch; + int q_tile; + + __device__ __forceinline__ bool init(const AttentionParams& p, + int flat_tile, int grid_batch) { + if ((threadIdx.x | threadIdx.y) == 0) { + batch = -1; + KV::template map_q_tile(p, flat_tile, grid_batch, batch, q_tile); + } + __syncthreads(); + return batch >= 0; + } +}; diff --git a/csrc/kernels/attn_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index 945b6b6..a032126 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -36,20 +36,13 @@ template p) { constexpr int DPT = HEAD_DIM / G; - __shared__ int mapped_batch; - __shared__ int mapped_q_tile; - if (threadIdx.x == 0 && threadIdx.y == 0) { - mapped_batch = -1; - KV::template map_q_tile( - p, blockIdx.x, blockIdx.z, mapped_batch, mapped_q_tile); - } - __syncthreads(); - if (mapped_batch < 0) + __shared__ QTileMapper qmap; + if (!qmap.init(p, blockIdx.x, blockIdx.z)) return; - int q_tile = mapped_q_tile; + int q_tile = qmap.q_tile; int q_head = blockIdx.y; - int batch = mapped_batch; + int batch = qmap.batch; int gpos = threadIdx.x; // 0..G-1 (which d-chunk) int row = threadIdx.y; // 0..ROWS-1 int q_row = q_tile * ROWS + row; diff --git a/csrc/kernels/attn_prefill_split_q_mma.cuh b/csrc/kernels/attn_prefill_split_q_mma.cuh index f9140da..3bf9ff9 100644 --- a/csrc/kernels/attn_prefill_split_q_mma.cuh +++ b/csrc/kernels/attn_prefill_split_q_mma.cuh @@ -24,18 +24,11 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { const int tid4 = lane & 3; // 0..3 const int q_head = blockIdx.y; - __shared__ int mapped_batch; - __shared__ int mapped_q_tile; - if (threadIdx.x == 0) { - mapped_batch = -1; - KV::template map_q_tile( - p, blockIdx.x, blockIdx.z, mapped_batch, mapped_q_tile); - } - __syncthreads(); - if (mapped_batch < 0) + __shared__ QTileMapper qmap; + if (!qmap.init(p, blockIdx.x, blockIdx.z)) return; - const int batch = mapped_batch; - const int q_tile = mapped_q_tile; + const int batch = qmap.batch; + const int q_tile = qmap.q_tile; const int kv_head = q_head / (p.q_head / p.kv_head); const int qrow0 = (q_tile * Traits::WARPS + warp) * Traits::BR;