refactor: streamline Q block mapping
- bypass shared mapping for contiguous attention - centralize paged Q tile broadcast in KV policy helpers
This commit is contained in:
@@ -184,24 +184,29 @@ struct PagedKV {
|
||||
}
|
||||
};
|
||||
|
||||
// ---- 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<ROWS, KV> qmap;
|
||||
// ---- Q-block mapping ----
|
||||
// Contiguous grids map directly to (batch, q_tile). Paged grids flatten the
|
||||
// ragged Q tiles, so one thread resolves the request and broadcasts it.
|
||||
template <int ROWS, typename KV>
|
||||
struct QTileMapper {
|
||||
int batch;
|
||||
int q_tile;
|
||||
__device__ __forceinline__ bool map_q_block(
|
||||
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
|
||||
if constexpr (!KV::kPaged) {
|
||||
batch = blockIdx.z;
|
||||
q_tile = blockIdx.x;
|
||||
return true;
|
||||
} else {
|
||||
__shared__ int mapped_batch;
|
||||
__shared__ int mapped_q_tile;
|
||||
|
||||
__device__ __forceinline__ bool init(const AttentionParams<bf16>& p,
|
||||
int flat_tile, int grid_batch) {
|
||||
if ((threadIdx.x | threadIdx.y) == 0) {
|
||||
batch = -1;
|
||||
KV::template map_q_tile<ROWS>(p, flat_tile, grid_batch, batch, q_tile);
|
||||
mapped_batch = -1;
|
||||
KV::template map_q_tile<ROWS>(
|
||||
p, blockIdx.x, blockIdx.z, mapped_batch, mapped_q_tile);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
batch = mapped_batch;
|
||||
q_tile = mapped_q_tile;
|
||||
return batch >= 0;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -36,13 +36,11 @@ template <int HEAD_DIM, typename KV, int G, int ROWS, int P_BC, bool IsCausal, b
|
||||
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
constexpr int DPT = HEAD_DIM / G;
|
||||
|
||||
__shared__ QTileMapper<ROWS, KV> qmap;
|
||||
if (!qmap.init(p, blockIdx.x, blockIdx.z))
|
||||
int batch, q_tile;
|
||||
if (!map_q_block<ROWS, KV>(p, batch, q_tile))
|
||||
return;
|
||||
|
||||
int q_tile = qmap.q_tile;
|
||||
int q_head = blockIdx.y;
|
||||
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;
|
||||
|
||||
@@ -24,11 +24,9 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
const int tid4 = lane & 3; // 0..3
|
||||
|
||||
const int q_head = blockIdx.y;
|
||||
__shared__ QTileMapper<Traits::BR * Traits::WARPS, KV> qmap;
|
||||
if (!qmap.init(p, blockIdx.x, blockIdx.z))
|
||||
int batch, q_tile;
|
||||
if (!map_q_block<Traits::BR * Traits::WARPS, KV>(p, batch, q_tile))
|
||||
return;
|
||||
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;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user