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:
+39
-22
@@ -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();
|
||||
|
||||
@@ -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])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
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);
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
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];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
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;
|
||||
#pragma unroll
|
||||
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
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 * 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,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
// 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,
|
||||
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);
|
||||
}
|
||||
@@ -143,19 +122,19 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
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++) {
|
||||
#pragma unroll
|
||||
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];
|
||||
}
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
if constexpr (HasMask) {
|
||||
if (!p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
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++)
|
||||
|
||||
@@ -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,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
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];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
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;
|
||||
#pragma unroll
|
||||
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++)
|
||||
#pragma unroll
|
||||
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,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
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,
|
||||
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++) {
|
||||
#pragma unroll
|
||||
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];
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,12 +89,14 @@ __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) {
|
||||
int ep = q_row + p.causal_offset + 1;
|
||||
if (kv0 >= ep)
|
||||
lim = 0;
|
||||
else if (kv0 + tlen > ep)
|
||||
lim = ep - kv0;
|
||||
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;
|
||||
@@ -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])
|
||||
dot = -FLT_MAX;
|
||||
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;
|
||||
|
||||
@@ -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,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
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];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
float Oacc[Traits::DN8][4];
|
||||
#pragma unroll
|
||||
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;
|
||||
#pragma unroll
|
||||
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
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 * 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)
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
// Post-multiply scale in float (no bf16 precision loss)
|
||||
#pragma unroll
|
||||
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)
|
||||
: p.kv_len;
|
||||
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
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++) {
|
||||
#pragma unroll
|
||||
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;
|
||||
Oacc[dn8][1] * rl0);
|
||||
*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;
|
||||
Oacc[dn8][3] * rl1);
|
||||
*reinterpret_cast<__nv_bfloat162*>(
|
||||
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+109
-80
@@ -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,77 +137,92 @@ static void bench() {
|
||||
}
|
||||
}
|
||||
|
||||
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];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
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=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;
|
||||
|
||||
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));
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_decode(p, sc);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
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, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
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[][5] = {
|
||||
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
||||
{1, 32, 4, 512, 128},
|
||||
{1, 32, 4, 1024, 128},
|
||||
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], gs = Hq / Hk;
|
||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
||||
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;
|
||||
}
|
||||
|
||||
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];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
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.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));
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_decode(p, sc);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
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);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
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;
|
||||
if (fail) {
|
||||
printf("FAILED\n");
|
||||
return fail;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
|
||||
@@ -36,34 +36,63 @@ static void gather_kv_cpu(
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_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<Traits, IsCausal, HasMask>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
}
|
||||
#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<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 = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||
bool use_mma = G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||
if (use_mma) {
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
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>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
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
|
||||
{
|
||||
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<<<
|
||||
dim3(p.batch * p.kv_head, 1, p.num_splits),
|
||||
dim3(32, group_sz), smem>>>(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);
|
||||
}
|
||||
}
|
||||
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;
|
||||
|
||||
+100
-63
@@ -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,6 +134,59 @@ static void bench() {
|
||||
}
|
||||
}
|
||||
|
||||
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);
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
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
|
||||
@@ -118,59 +195,19 @@ int main() {
|
||||
{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];
|
||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
|
||||
if (fail) break;
|
||||
}
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
if (fail) {
|
||||
printf("FAILED\n");
|
||||
return fail;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
|
||||
Reference in New Issue
Block a user