refactor: adopt FA2-style KernelTraits + compile-time causal/mask dispatch

- Introduce KernelTraits<HEAD_DIM, BC, WARPS, STAGES> compile-time config bundle, replacing scattered <KD, NC8, KT2, ...> template params
- Template all MMA and scalar kernels on IsCausal/HasMask bools to eliminate inner-loop runtime branches
- Dispatch to 4-path IsCausal/HasMask kernel variants at entry points based on p.causal_offset and p.use_mask
- Update standalone test files with new kernel signatures, add causal test cases
- Fix duplicate using bf16 in MMA kernels that include attn_mma_utils.cuh
This commit is contained in:
2026-07-21 21:52:46 +08:00
parent ccf728a1b7
commit a01e8bbe98
13 changed files with 682 additions and 611 deletions
+39 -22
View File
@@ -3,9 +3,26 @@
#ifndef ASTRAI_NO_MMA
#include "attn_decode_split_kv_mma.cuh"
template <int HEAD_DIM, int BC, int STAGES, bool IsCausal, bool HasMask>
static void launch_mma_decode_impl(AttentionParams<bf16>& p) {
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM, int BC, bool IsCausal, bool HasMask>
static void launch_mma_decode(AttentionParams<bf16>& p) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
launch_mma_decode_impl<HEAD_DIM, BC, STAGES, IsCausal, HasMask>(p);
}
#endif
// Scalar fallback: one warp per query head, split-KV across grid.z.
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_decode(AttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
@@ -13,37 +30,38 @@ static void launch_scalar_decode(AttentionParams<bf16>& p) {
alloc_split_partials(p);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, group_size);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
// to preserve occupancy).
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_mma_decode(AttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_decode(AttentionParams<bf16>& p) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) {
launch_mma_decode<HEAD_DIM, 32>(p);
if (is_causal) {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, true, true>(p);
else launch_mma_decode<HEAD_DIM, 32, true, false>(p);
} else {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, false, true>(p);
else launch_mma_decode<HEAD_DIM, 32, false, false>(p);
}
return;
}
#endif
launch_scalar_decode(p);
if (is_causal) {
if (has_mask) launch_scalar_decode<HEAD_DIM, true, true>(p);
else launch_scalar_decode<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_scalar_decode<HEAD_DIM, false, true>(p);
else launch_scalar_decode<HEAD_DIM, false, false>(p);
}
}
torch::Tensor attn_decode(
@@ -60,7 +78,6 @@ torch::Tensor attn_decode(
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
// O matches Q's original layout
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
+15 -11
View File
@@ -12,6 +12,7 @@ __device__ inline float warp_reduce_sum(float val) {
return val;
}
template <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head;
@@ -48,7 +49,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// Load K into shared memory (gather from strided global)
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int kv_idx = chunk_start + s;
@@ -60,24 +62,30 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
for (int s = 0; s < this_chunk; s++) {
float partial = 0.0f;
for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial += q_reg[i] * __bfloat162float(
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
}
if constexpr (IsCausal) {
if (kv_idx > p.causal_offset)
partial = -FLT_MAX;
}
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
float beta = expf(partial - new_m);
d = d * alpha + beta;
// V: stride-based read
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d;
int v_off = kv_base + kv_idx * p.kv_stride_l
+ lane * hd_per_thread * p.kv_stride_d;
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
acc_reg[i] = acc_reg[i] * alpha
+ __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
m = new_m;
}
__syncthreads();
@@ -97,9 +105,6 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
}
}
// Reduce split-K partials into the final bf16 output. One block per (batch,
// q_head); each thread folds across all splits with a single-pass
// online-rescale reduction (expf + FMA counts halved vs 3-pass original).
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int bh = blockIdx.x;
int d = threadIdx.x;
@@ -126,7 +131,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
// Stride-based output write (q_len=1 for decode, so stride_l not needed)
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv);
}
+45 -66
View File
@@ -4,34 +4,17 @@
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
// GEMM that reuses each loaded K/V tile across all G heads.
//
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core
// work on its own. But GQA gives us G = q_head / kv_head query heads that all
// share one kv_head. We pack those G heads into the M=16 rows of
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that
// reuses each loaded K/V tile across all G heads (K/V load is the decode
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
// partitioned across gridDim.z blocks so that a decode with only
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
// slice; the combine kernel below reduces across splits. Fixes the "grid too
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
// small-batch decode.
template <int HEAD_DIM, int BC, int STAGES = 2>
// IsCausal and HasMask are compile-time bools — no runtime branch in the
// inner compute loop.
//
// Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
template <typename Traits, bool IsCausal, bool HasMask>
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x;
const int gid = lane >> 2;
const int tid4 = lane & 3;
@@ -42,46 +25,44 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct
// from global to registers).
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
// Double-buffered shared memory for K/V (no sQ needed)
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// ---- Load Q directly from global into mma A-operand registers ----
// Load Q directly from global into mma A-operand registers.
// stride_row = p.q_stride_h for decode (q_len=1).
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int qra = gid;
const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4];
float Oacc[Traits::DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// ---- Load tile lambda: predicated cp.async, unified full/partial ----
// ---- Load tile lambda: predicated cp.async ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
bf16* dK = sK + buf * BC * LD;
bf16* dV = sV + buf * BC * LD;
int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
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;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
@@ -89,50 +70,48 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
// Prologue
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK;
// Wait for current tile, then issue next tile's prefetch (overlaps
// with this tile's compute). Single syncwarp covers both hazards.
// When STAGES==1, no prefetch — load happens at end of prior iter.
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
// Decode: q_len=1, so qrow0=qrow1=0
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0,
p.mask_b_stride, 0,
batch,
p.mask, has_mask,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp();
if constexpr (STAGES == 1) {
if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
}
@@ -144,18 +123,18 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
return bh * p.num_splits + split;
};
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1];
}
if (r1 < G) {
int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3];
}
+75 -77
View File
@@ -3,10 +3,41 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// Shared MMA utilities for tensor-core GQA kernels.
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing.
// ============================================================================
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
//
// Bundles all dimension-dependent constants so device functions only need a
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
// ============================================================================
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
struct KernelTraits {
static constexpr int HEAD_DIM = HEAD_DIM_;
static constexpr int BC = BC_; // K/V tile size along seq dim
static constexpr int WARPS = WARPS_; // warps per block
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
static constexpr int BR = 16; // Q rows per warp (mma M=16)
// Derived: mma.sync.m16n8k16 tile counts
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
static constexpr int LD = HEAD_DIM; // smem leading dim
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
// mask = log2(LD/8) bits, clamped to stay within LD.
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
static constexpr int NUM_THREADS = WARPS * 32;
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
};
// ---- PTX wrappers ----
using bf16 = __nv_bfloat16;
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
const unsigned* b, const float* c) {
asm volatile(
@@ -37,9 +68,7 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
}
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the
// scalar per-thread fragment packing, cutting shared-load instructions and bank
// conflicts. Each lane supplies the shared address of one 8-wide row.
// 16x16 / 16x8 tile) with the exact register layout mma expects.
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
@@ -60,29 +89,19 @@ __device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
}
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
// existing HEAD_DIM>=64 call sites working unchanged.
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
}
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly,
// bypassing registers. Eliminates shared-store bank conflicts and cuts
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
// Requires sm_80+.
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
:: "r"(smem_addr), "l"(gmem_ptr));
}
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the
// destination (src-size operand = 0 → no bytes read from src, so an
// out-of-bounds src address is never dereferenced). Lets full and partial
// tiles share one uniform async load path — no scalar fallback branch.
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr,
bool pred) {
@@ -100,9 +119,6 @@ __device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
// Wait until at most N commit groups are still in flight. Used for
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
// continue while ensuring the current tile's data is ready.
template <int N>
__device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N));
@@ -139,78 +155,65 @@ __device__ inline void load_q_mma_frags(
}
// ---------------------------------------------------------------------------
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
// between the two kernels; only the per-row causal/mask bounds differ.
// ---------------------------------------------------------------------------
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
// caller to avoid bf16 precision loss).
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as
// runtime ints lets the compiler fold them while keeping the signature clean.
template <int KD, int NC8>
// Traits provides KD, NC8, LD, and SWIZ_MASK.
// ---------------------------------------------------------------------------
template <typename Traits>
__device__ inline void mma_compute_scores(
const unsigned Qa[KD][4],
const unsigned Qa[Traits::KD][4],
const bf16* __restrict__ sK,
int LD,
int SWIZ_MASK,
int lane,
float Sacc[NC8][4])
float Sacc[Traits::NC8][4])
{
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++) {
for (int n8 = 0; n8 < Traits::NC8; n8++) {
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
int krow_l = n8 * 8 + (lane & 7);
int kcol_h = (lane & 8) ? 8 : 0;
#pragma unroll
for (int kt = 0; kt < KD; kt++) {
for (int kt = 0; kt < Traits::KD; kt++) {
unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
ldmatrix_x2(b, &sK[krow_l * Traits::LD
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
}
}
}
// ---------------------------------------------------------------------------
// Online softmax + Oacc rescale for one K/V tile.
// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits;
// decode: same value for both rows since q_len==1).
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0).
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len).
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm),
// rescales Oacc by exp(m_old - nm), and updates m/l — all in place.
template <int NC8, int DN8>
//
// HasMask is a compile-time template bool: when false, the mask branch is
// entirely dead-code-eliminated from the inner unrolled loop.
// ---------------------------------------------------------------------------
template <typename Traits, bool HasMask>
__device__ inline void mma_softmax_tile(
int kv0,
int maxc0,
int maxc1,
int qrow0,
int qrow1,
int mask_b_stride,
int mask_q_stride,
int maxc0, int maxc1,
int qrow0, int qrow1,
int mask_b_stride, int mask_q_stride,
int mask_batch,
const bool* __restrict__ mask,
bool has_mask,
float Sacc[NC8][4],
float Oacc[DN8][4],
float Sacc[Traits::NC8][4],
float Oacc[Traits::DN8][4],
float& m0, float& m1,
float& l0, float& l1,
int lane)
{
int tid4 = lane & 3;
// Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream
// without per-element sentinel checks. Compute tile-local row maxima.
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++) {
for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4;
int c1 = cc + 1;
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]);
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]);
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]);
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]);
bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
@@ -220,29 +223,20 @@ __device__ inline void mma_softmax_tile(
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
}
// Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2).
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
// nm = max(running max m, tile-local max rmax) — updated running maximum.
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
// corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm ==
// -FLT_MAX), exp(0) = 1 — correct, no guard needed.
float corr0 = __expf(m0 - nm0);
float corr1 = __expf(m1 - nm1);
// pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm)
// gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8
// per-element comparisons.
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
// P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give
// exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge.
float rsum0 = 0.0f, rsum1 = 0.0f;
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++) {
for (int n8 = 0; n8 < Traits::NC8; n8++) {
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
@@ -261,22 +255,25 @@ __device__ inline void mma_softmax_tile(
m0 = nm0; m1 = nm1;
#pragma unroll
for (int j = 0; j < DN8; j++) {
for (int j = 0; j < Traits::DN8; j++) {
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
}
}
// ---------------------------------------------------------------------------
// O += P @ V (Sacc must contain P = attention weights after softmax).
template <int DN8, int KT2>
// Traits provides DN8, KT2, LD, and SWIZ_MASK.
// ---------------------------------------------------------------------------
template <typename Traits>
__device__ inline void mma_pv_accumulate(
float Sacc[][4],
const bf16* __restrict__ sV,
int LD, int SWIZ_MASK, int lane,
float Oacc[DN8][4])
int lane,
float Oacc[Traits::DN8][4])
{
#pragma unroll
for (int kt2 = 0; kt2 < KT2; kt2++) {
for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
unsigned Pa[4];
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
@@ -284,9 +281,10 @@ __device__ inline void mma_pv_accumulate(
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
int vrow_l = kt2 * 16 + (lane & 15);
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
}
}
+39 -15
View File
@@ -5,6 +5,27 @@
#include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, int BC, int STAGES, bool IsCausal, bool HasMask>
static void launch_paged_mma_decode_impl(PagedAttentionParams<bf16>& p) {
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM, int BC, bool IsCausal, bool HasMask>
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
launch_paged_mma_decode_impl<HEAD_DIM, BC, STAGES, IsCausal, HasMask>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
@@ -14,32 +35,35 @@ static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
dim3 block = dim3(32, group_size);
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16 && p.page_size >= 32) {
launch_paged_mma_decode<HEAD_DIM, 32>(p);
if (is_causal) {
if (has_mask) launch_paged_mma_decode<HEAD_DIM, 32, true, true>(p);
else launch_paged_mma_decode<HEAD_DIM, 32, true, false>(p);
} else {
if (has_mask) launch_paged_mma_decode<HEAD_DIM, 32, false, true>(p);
else launch_paged_mma_decode<HEAD_DIM, 32, false, false>(p);
}
return;
}
#endif
launch_paged_scalar_decode(p);
if (is_causal) {
if (has_mask) launch_paged_scalar_decode<HEAD_DIM, true, true>(p);
else launch_paged_scalar_decode<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_paged_scalar_decode<HEAD_DIM, false, true>(p);
else launch_paged_scalar_decode<HEAD_DIM, false, false>(p);
}
}
torch::Tensor attn_paged_decode(
+13 -7
View File
@@ -12,7 +12,7 @@ __device__ inline float paged_warp_reduce_sum(float val) {
return val;
}
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
template <int HEAD_DIM, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head;
@@ -22,7 +22,6 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32;
// Q: stride-based [batch, q_head, q_len=1, head_dim]
float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ lane * hd_per_thread * p.q_stride_d;
@@ -46,7 +45,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int pos = chunk_start + s;
@@ -69,14 +69,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
float partial = 0.0f;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial += q_reg[i] * __bfloat162float(
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = paged_warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
}
if constexpr (IsCausal) {
if (kv_idx > p.causal_offset)
partial = -FLT_MAX;
}
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
@@ -93,7 +98,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
+ (int64_t)kv_head * p.head_dim;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
acc_reg[i] = acc_reg[i] * alpha
+ __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
} else {
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
+41 -56
View File
@@ -4,25 +4,14 @@
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Paged split-KV tensor-core decode via GQA head-packing.
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
// directly from the page pool through a page table, eliminating the gather
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
// the page-table lookup happens once per tile for cp.async.
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
// Reads K/V directly from the page pool through a page table — one tile
// (BC=32) fits within a single page (page_size >= 32), so the page-table
// lookup happens once per tile for cp.async.
//
// IsCausal and HasMask are compile-time bools.
template <typename Traits, bool IsCausal, bool HasMask>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x;
const int gid = lane >> 2;
const int tid4 = lane & 3;
@@ -33,123 +22,119 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// ---- Load Q directly from global into mma A-operand registers ----
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int qra = gid;
const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4];
float Oacc[Traits::DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// Paged strides (constant for the block)
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
// ---- Load tile lambda: predicated cp.async, paged addressing ----
// ---- Load tile lambda: paged addressing ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
bf16* dK = sK + buf * BC * LD;
bf16* dV = sV + buf * BC * LD;
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 * VEC; i < TOTAL; i += 32 * VEC) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
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;
int page_off = kc % p.page_size;
int64_t gmem_base = (int64_t)phys_page * page_stride
+ (int64_t)page_off * pos_stride
+ head_off;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
}
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if constexpr (Traits::STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0,
p.mask_b_stride, 0,
batch,
p.mask, has_mask,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp();
if constexpr (STAGES == 1) {
if constexpr (Traits::STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
}
}
// ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split;
};
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1];
}
if (r1 < G) {
int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3];
}
+37 -19
View File
@@ -3,30 +3,48 @@
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_mma_prefill(AttentionParams<bf16>& p) {
constexpr int WARPS = 4;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS),
p.q_head, p.batch);
dim3 block(Traits::NUM_THREADS, 1, 1);
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
// KV tile: bigger tiles amortize the per-tile cp.async wait + barrier +
// loop overhead over more tensor-core work (this kernel is latency-bound,
// not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for
// D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem,
// over the 48KB static cap.
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
// Static shared memory — double-buffered K/V only (no sQ: Q goes direct
// to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes.
// Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB),
// D=256→1 (32KB, BC=16).
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_prefill(AttentionParams<bf16>& 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, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
}
template <int HEAD_DIM>
static void dispatch_prefill(AttentionParams<bf16>& 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_mma_prefill<HEAD_DIM, true, true>(p);
else launch_mma_prefill<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_mma_prefill<HEAD_DIM, false, true>(p);
else launch_mma_prefill<HEAD_DIM, false, false>(p);
}
#else
if (is_causal) {
if (has_mask) launch_scalar_prefill<HEAD_DIM, true, true>(p);
else launch_scalar_prefill<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_scalar_prefill<HEAD_DIM, false, true>(p);
else launch_scalar_prefill<HEAD_DIM, false, false>(p);
}
#endif
}
+11 -14
View File
@@ -6,12 +6,9 @@
using bf16 = __nv_bfloat16;
// v9: group-split register blocking. G threads cooperate on one query row,
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps
// occupancy high; the S dot product is reduced across the G-lane group with a
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce.
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long.
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
// G*ROWS a multiple of 32 with groups warp-aligned.
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
// compile-time bools — the compiler eliminates dead branches.
// Templated on <HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>.
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
@@ -21,8 +18,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
return v;
}
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
float4 raw = *reinterpret_cast<const float4*>(p);
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
@@ -34,7 +30,7 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
}
}
template <int HEAD_DIM, int G, int ROWS, int P_BC>
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
constexpr int DPT = HEAD_DIM / G;
@@ -73,8 +69,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
int tt = G * ROWS;
int lid = row * G + gpos;
// per-group shuffle mask: only the G lanes of this row's group participate,
// so causal masking (differing loop bounds across rows in a warp) is safe.
int lane_in_warp = lid & 31;
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
@@ -95,13 +89,15 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
__syncthreads();
int lim = tlen;
if (p.causal_offset >= 0 && q_row < p.q_len) {
if constexpr (IsCausal) {
if (q_row < p.q_len) {
int ep = q_row + p.causal_offset + 1;
if (kv0 >= ep)
lim = 0;
else if (kv0 + tlen > ep)
lim = ep - kv0;
}
}
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
for (int s = 0; s < lim; s++) {
@@ -118,8 +114,10 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
float dot = group_reduce_sum<G>(part, gmask);
int kv_idx = kv0 + s;
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx])
if constexpr (HasMask) {
if (!p.mask[mask_row_base + kv_idx])
dot = -FLT_MAX;
}
float nm = fmaxf(m, dot);
float al = __expf(m - nm);
@@ -141,7 +139,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
}
if (q_row < p.q_len) {
// O: stride-based write
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
+53 -103
View File
@@ -4,121 +4,76 @@
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Tensor-core prefill flash attention (raw mma.sync PTX).
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once
// straight from global into the mma A-operand layout (no smem staging) and
// kept resident in registers across the tile loop. S, O, and the online-softmax
// stats (m, l) also live in registers.
// Shared memory is statically sized via template parameters — no dynamic
// allocation. The mma fragment layout is used directly: the S accumulator
// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so
// softmax needs no shuffle repack; row reductions fold across the 4-lane
// thread group. Templated on <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
// cores via mma.sync.m16n8k16 (f32 accumulate).
//
// Software pipeline: K/V are double-buffered and loaded via cp.async one tile
// ahead, so the next tile streams from global memory while the current tile's
// tensor-core math runs — hiding load latency (long_scoreboard). A single
// __syncthreads per tile both publishes the freshly loaded tile cross-warp and
// (because it runs before the next prefetch) guards the buffer being refilled,
// so no second barrier is needed. Predicated cp.async (cp_async_16_pred)
// zero-fills rows past kv_len, unifying full and partial tiles on one path.
// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more
// tensor-core work — this kernel is latency-bound (low occupancy from high
// register pressure), so fewer, larger tiles beat many tiny ones.
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
// dead branches in the inner compute loop (FA2-style).
//
// Optimizations: load Q fragments directly from global in mma A-operand layout
// (no sQ staging, no prologue barriers); post-multiply scale in float after
// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores;
// causal tile skipping (block-level prefetch bound + warp-level compute skip);
// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD
// padding (LD=HEAD_DIM).
template <int HEAD_DIM, int WARPS, int BC>
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
template <typename Traits, bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
constexpr int BR = 16;
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32;
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
const int gid = lane >> 2; // 0..7
const int tid4 = lane & 3; // 0..3
const int nthreads = WARPS * 32;
const int q_head = blockIdx.y;
const int batch = blockIdx.z;
const int kv_head = q_head / (p.q_head / p.kv_head);
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
// ---- Static shared memory: double-buffered K/V ----
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs
// while the current tile's tensor-core math executes, hiding global-load
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in.
constexpr int STAGES = 2;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
// to registers in mma A-operand layout).
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// Load Q fragments straight from global into mma A-operand layout.
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len).
// See attn_mma_utils.cuh for the shared template.
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
const int qra = qrow0 + gid;
const int qrb = qrow0 + gid + 8;
const bool va = qra < p.q_len, vb = qrb < p.q_len;
unsigned Qa[KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4];
float Oacc[Traits::DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles = (p.kv_len + BC - 1) / BC;
const int qr0 = qrow0 + gid; // row for c0/c1
const int qr1 = qrow0 + gid + 8; // row for c2/c3
const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC;
const int qr0 = qrow0 + gid;
const int qr1 = qrow0 + gid + 8;
// Causal tile-skip bounds (no-op when causal_offset < 0)
const int use_skip = (p.causal_offset >= 0) ? 1 : 0;
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
// Causal tile-skip bounds (dead code when IsCausal == false)
const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset;
const int block_max_kv =
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
const int has_mask = p.use_mask && p.mask;
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
+ p.causal_offset;
// Last active tile: block-level causal bound (all warps in the block share
// the K/V load, so the prefetch range is the block max, not per-warp).
int t_end = tiles - 1;
if (use_skip) {
int bt = block_max_kv / BC;
if constexpr (IsCausal) {
int bt = block_max_kv / Traits::BC;
if (bt < t_end) t_end = bt;
}
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
constexpr int TOTAL = BC * HEAD_DIM;
// ---- Load tile lambda: predicated cp.async ----
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
bf16* dK = sK + buf * BC * LD;
bf16* dV = sV + buf * BC * LD;
int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
for (int i = threadIdx.x * 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;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
@@ -132,65 +87,60 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
for (int ti = 0; ti <= t_end; ti++) {
int buf = ti & 1;
// Wait for the current tile's async copies, then a single barrier: it
// both publishes this tile's data cross-warp AND guarantees the prior
// compute on the buffer we are about to refill has finished. Issuing
// the next tile's load *after* this barrier lets one barrier cover both
// hazards (vs two), while the load still overlaps this tile's math.
// Wait for current tile, then publish cross-warp + guard buffer reuse.
cp_async_wait_group<0>();
__syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * Traits::BC;
// Warp-level causal skip
if (!use_skip || kv0 <= max_kv) {
// Warp-level causal skip (dead branch eliminated when IsCausal == false)
if (!IsCausal || kv0 <= max_kv) {
// S = Q @ K^T + scale + online softmax + O += P @ V
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
float Sacc[Traits::NC8][4];
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
// post-multiply scale in float (no bf16 precision loss from pre-scaling Q)
// Post-multiply scale in float (no bf16 precision loss)
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1)
int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
: p.kv_len;
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1)
int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
: p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1,
p.mask_b_stride, p.mask_q_stride,
batch,
p.mask, has_mask,
p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
} // if active (warp-level causal skip)
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
}
}
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair,
// halves store count and removes the uncoalesced scalar-store penalty)
// ---- write output: packed bf16x2 stores ----
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
// O: stride-based write
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
if (qr0 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
Oacc[dn8][1] * rl0);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
*reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
}
if (qr1 < p.q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
*reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
}
}
}
+60 -31
View File
@@ -1,5 +1,5 @@
/*
Pure-C test:
Pure-C test — updated for KernelTraits + IsCausal/HasMask.
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_decode_test.cu -o test && ./test
@@ -11,35 +11,29 @@ nvcc -I csrc -arch=sm_89 -O3 \
#include "../kernels/attn_decode_split_kv_mma.cuh"
#endif
// Split-K scratch (torch-free): the production launcher allocates these from
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
// Split-K scratch (torch-free)
struct DecodeScratch {
float* o_part = nullptr;
float* ml_part = nullptr;
};
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
#ifndef ASTRAI_NO_MMA
static bool decode_use_mma(const AttentionParams<bf16>& p) {
int G = p.q_head / p.kv_head;
return !p.use_mask && G > 1 && G <= 16;
}
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
template <int HEAD_DIM, int BC, bool IsCausal, bool HasMask>
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
int gs = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
@@ -48,16 +42,36 @@ static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
p.ml_part = sc.ml_part;
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask>
<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) {
if (is_causal) {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, true, true>(p, sc);
else launch_mma_decode<HEAD_DIM, 32, true, false>(p, sc);
} else {
if (has_mask) launch_mma_decode<HEAD_DIM, 32, false, true>(p, sc);
else launch_mma_decode<HEAD_DIM, 32, false, false>(p, sc);
}
return;
}
#endif
launch_scalar_decode(p, sc);
if (is_causal) {
if (has_mask) launch_scalar_decode<HEAD_DIM, true, true>(p, sc);
else launch_scalar_decode<HEAD_DIM, true, false>(p, sc);
} else {
if (has_mask) launch_scalar_decode<HEAD_DIM, false, true>(p, sc);
else launch_scalar_decode<HEAD_DIM, false, false>(p, sc);
}
}
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
@@ -123,18 +137,10 @@ static void bench() {
}
}
int main() {
const int configs[][5] = {
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
{1, 32, 4, 512, 128},
{1, 32, 4, 1024, 128},
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_cfgs; ci++) {
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
int gs = Hq / Hk;
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
B,Hq,Hk,sl,D,gs,causal);
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
@@ -161,12 +167,11 @@ int main() {
AttentionParams<bf16> p;
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.causal_offset=-1;
p.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D);
set_default_strides(p);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
// Split-K scratch (max 32 splits), sized for the production MMA path.
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
@@ -182,7 +187,7 @@ int main() {
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1);
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
float max_err=0;
for (size_t i=0;i<nQ;i++){
@@ -194,6 +199,30 @@ int main() {
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
cudaFree(sc.o_part);cudaFree(sc.ml_part);
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
return (max_err < 0.05f) ? 0 : 1;
}
int main() {
const int configs[][6] = {
{1, 2, 1, 64, 32, 0}, // B,Hq,Hk,seq_len,D,causal
{1, 32, 4, 512, 128, 0},
{1, 32, 4, 1024, 128, 0},
{1, 32, 4, 512, 128, 1}, // causal decode
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
int fail = 0;
for (int ci = 0; ci < n_cfgs; ci++) {
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
fail += run_test(B, Hq, Hk, sl, D, causal);
if (fail) break;
}
if (fail) {
printf("FAILED\n");
return fail;
}
printf("All tests passed!\n");
bench();
+59 -32
View File
@@ -36,34 +36,63 @@ static void gather_kv_cpu(
}
}
template <int HEAD_DIM>
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
#ifndef ASTRAI_NO_MMA
int G_check = p.q_head / p.kv_head;
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
if (use_mma) {
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_paged_mma_decode(PagedAttentionParams<bf16, float>& p) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
} else
}
#endif
{
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_paged_scalar_decode(PagedAttentionParams<bf16, float>& p) {
int group_sz = p.q_head / p.kv_head;
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);
paged_attn_decode_split_kv_kernel<<<
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<
dim3(p.batch * p.kv_head, 1, p.num_splits),
dim3(32, group_sz), smem>>>(p);
}
template <int HEAD_DIM>
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
int G_check = p.q_head / p.kv_head;
bool use_mma = G_check >= 1 && G_check <= 16 && p.page_size >= 32;
if (use_mma) {
if (is_causal) {
if (has_mask) launch_paged_mma_decode<HEAD_DIM, true, true>(p);
else launch_paged_mma_decode<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_paged_mma_decode<HEAD_DIM, false, true>(p);
else launch_paged_mma_decode<HEAD_DIM, false, false>(p);
}
} else
#endif
{
if (is_causal) {
if (has_mask) launch_paged_scalar_decode<HEAD_DIM, true, true>(p);
else launch_paged_scalar_decode<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_paged_scalar_decode<HEAD_DIM, false, true>(p);
else launch_paged_scalar_decode<HEAD_DIM, false, false>(p);
}
}
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size;
@@ -138,13 +167,14 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1);
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv,
1, kv_len, HEAD_DIM, causal ? 0 : -1);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
PagedAttentionParams<bf16, float> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
p.use_mask = 0; p.causal_offset = -1;
p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
set_default_paged_strides(p);
p.num_splits = 1; p.scale = scale_val;
p.page_size = page_size; p.max_pages = max_pages;
@@ -201,25 +231,24 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed)
struct TestCase {
int head_dim;
int B, Hq, Hkv, kv_len, page_size, seed;
int B, Hq, Hkv, kv_len, page_size, causal, seed;
};
static const TestCase TESTS[] = {
{128, 1, 1, 1, 8, 128, 1},
{128, 1, 4, 4, 128, 128, 2},
{128, 2, 4, 4, 256, 128, 3},
{128, 1, 4, 1, 64, 64, 4},
{128, 1, 8, 2, 64, 128, 5},
{128, 2, 16, 4, 128, 128, 6},
{64, 1, 4, 2, 32, 128, 7},
{256, 1, 2, 1, 16, 128, 8},
{32, 1, 4, 2, 32, 64, 9},
{128, 3, 8, 2, 256, 128, 10},
{128, 2, 32, 8, 512, 128, 11},
#ifndef ASTRAI_NO_MMA
{128, 1, 16, 2, 256, 128, 12},
{128, 2, 32, 4, 512, 128, 13},
#endif
{128, 1, 1, 1, 8, 128, 0, 1},
{128, 1, 4, 4, 128, 128, 0, 2},
{128, 2, 4, 4, 256, 128, 0, 3},
{128, 1, 4, 1, 64, 64, 0, 4},
{128, 1, 8, 2, 64, 128, 0, 5},
{128, 2, 16, 4, 128, 128, 0, 6},
{64, 1, 4, 2, 32, 128, 0, 7},
{256, 1, 2, 1, 16, 128, 0, 8},
{32, 1, 4, 2, 32, 64, 0, 9},
{128, 3, 8, 2, 256, 128, 0, 10},
{128, 2, 32, 8, 512, 128, 0, 11},
{128, 1, 16, 2, 256, 128, 0, 12},
{128, 2, 32, 4, 512, 128, 0, 13},
{128, 2, 8, 2, 128, 128, 1, 14}, // causal paged decode
};
static int dispatch_test(const TestCase& tc) {
@@ -227,13 +256,11 @@ static int dispatch_test(const TestCase& tc) {
int r = 0;
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
matched = true;
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
});
return matched ? r : 1;
}
// Warmed-up, CUDA-event timed sweep over paged decode configs.
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
template <int HEAD_DIM>
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
int max_pages = (kv_len + page_size - 1) / page_size;
+66 -29
View File
@@ -1,5 +1,5 @@
/*
Pure-C test:
Pure-C test — updated for KernelTraits + IsCausal/HasMask.
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_prefill_test.cu -o test && ./test
@@ -11,35 +11,60 @@ nvcc -I csrc -arch=sm_89 -O3 \
#include "../kernels/attn_prefill_split_q_mma.cuh"
#endif
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
template <int HEAD_DIM>
static void launch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_mma_prefill(AttentionParams<bf16>& p) {
constexpr int WARPS = 4;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS),
p.q_head, p.batch);
dim3 block(Traits::NUM_THREADS, 1, 1);
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
}
#endif
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch_scalar_prefill(AttentionParams<bf16>& 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, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC,
IsCausal, HasMask><<<grid, block>>>(p);
}
template <int HEAD_DIM>
static void launch_prefill_dispatch(AttentionParams<bf16>& 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_mma_prefill<HEAD_DIM, true, true>(p);
else launch_mma_prefill<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_mma_prefill<HEAD_DIM, false, true>(p);
else launch_mma_prefill<HEAD_DIM, false, false>(p);
}
#else
if (is_causal) {
if (has_mask) launch_scalar_prefill<HEAD_DIM, true, true>(p);
else launch_scalar_prefill<HEAD_DIM, true, false>(p);
} else {
if (has_mask) launch_scalar_prefill<HEAD_DIM, false, true>(p);
else launch_scalar_prefill<HEAD_DIM, false, false>(p);
}
#endif
}
static void dispatch_prefill(AttentionParams<bf16>& p) {
switch (p.head_dim) {
case 64: launch_prefill<64>(p); break;
case 128: launch_prefill<128>(p); break;
case 64: launch_prefill_dispatch<64>(p); break;
case 128: launch_prefill_dispatch<128>(p); break;
default: printf("bench: unsupported D=%d\n", p.head_dim);
}
}
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
static void bench() {
const int cfgs[][7] = {
{1,32,4,512,512,128,0},
@@ -94,7 +119,6 @@ static void bench() {
double flops = 4.0*B*Hq*(double)ql*kl*D;
if (causal) flops *= 0.5;
double tflops = flops/(ms*1e-3)/1e12;
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
double gbps = bytes/(ms*1e-3)/1e9;
@@ -110,19 +134,7 @@ static void bench() {
}
}
int main() {
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_configs; ci++) {
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
int causal=configs[ci][6];
static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
B,Hq,Hk,ql,kl,D,causal);
@@ -171,6 +183,31 @@ int main() {
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
return (max_err < 0.05f) ? 0 : 1;
}
int main() {
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[0]);
int fail = 0;
for (int ci = 0; ci < n_configs; ci++) {
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
int causal=configs[ci][6];
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
if (fail) break;
}
if (fail) {
printf("FAILED\n");
return fail;
}
printf("All tests passed!\n");
bench();