feat: support HEAD_DIM=32 and split extension into loader/ops
- add case 32 to decode/prefill dispatch switch - fix swiz_col out-of-bounds for HEAD_DIM=32: XOR mask now limited to chunk count (3 for 32, 7 for >=64) instead of always 7, which produced column offsets >= LD=32 and corrupted shared memory - restructure decode dispatch to #ifndef/#else/#endif matching prefill - split astrai/extension/__init__.py into loader.py (kernel .so discovery) and ops.py (wrapper functions + torch SDPA fallback); __init__.py now re-exports the public API
This commit is contained in:
@@ -21,13 +21,20 @@ static void dispatch_decode(GQAParams& p) {
|
||||
gqa_decode_attn_mma_kernel<HEAD_DIM, BC><<<grid, block, smem>>>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
// scalar fallback (per-KV-head, one warp per query head)
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 block(32, group_size);
|
||||
dim3 grid(p.batch * p.kv_head);
|
||||
gqa_decode_attn_kernel<<<grid, block, smem>>>(p);
|
||||
#else
|
||||
// scalar fallback (per-KV-head, one warp per query head)
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 block(32, group_size);
|
||||
dim3 grid(p.batch * p.kv_head);
|
||||
gqa_decode_attn_kernel<<<grid, block, smem>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
torch::Tensor gqa_decode_attn(
|
||||
@@ -74,6 +81,9 @@ torch::Tensor gqa_decode_attn(
|
||||
p.o = (bf16*)O.data_ptr();
|
||||
|
||||
switch (p.head_dim) {
|
||||
case 32:
|
||||
dispatch_decode<32>(p);
|
||||
break;
|
||||
case 64:
|
||||
dispatch_decode<64>(p);
|
||||
break;
|
||||
@@ -85,7 +95,7 @@ torch::Tensor gqa_decode_attn(
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "decode: unsupported head_dim ", p.head_dim,
|
||||
" (supported: 64, 128, 256)");
|
||||
" (supported: 32, 64, 128, 256)");
|
||||
}
|
||||
return O;
|
||||
}
|
||||
|
||||
@@ -58,10 +58,12 @@ __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) = ((d>>3)^(r&7))<<3 | (d&7).
|
||||
// Works for any d; aligned (d%8==0) simplifies to d ^ ((r&7)<<3).
|
||||
__device__ __forceinline__ int swiz_col(int d, int r) {
|
||||
return ((d >> 3) ^ (r & 7)) << 3 | (d & 7);
|
||||
// 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,
|
||||
|
||||
@@ -68,6 +68,9 @@ torch::Tensor gqa_prefill_attn(
|
||||
p.o = (bf16*)O.data_ptr();
|
||||
|
||||
switch (p.head_dim) {
|
||||
case 32:
|
||||
dispatch_prefill<32>(p);
|
||||
break;
|
||||
case 64:
|
||||
dispatch_prefill<64>(p);
|
||||
break;
|
||||
@@ -79,7 +82,7 @@ torch::Tensor gqa_prefill_attn(
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "prefill: unsupported head_dim ", p.head_dim,
|
||||
" (supported: 64,128,256)");
|
||||
" (supported: 32,64,128,256)");
|
||||
}
|
||||
return O;
|
||||
}
|
||||
|
||||
@@ -30,6 +30,7 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
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;
|
||||
@@ -61,12 +62,12 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
int qr = qrow0 + r;
|
||||
bf16 qv = (qr < p.q_len) ? p.q[q_base + qr * HEAD_DIM + d]
|
||||
: __float2bfloat16(0.0f);
|
||||
sQ[r * LD + swiz_col(d, r)] = __hmul(qv, scale_bf16);
|
||||
sQ[r * LD + swiz_col(d, r, SWIZ_MASK)] = __hmul(qv, scale_bf16);
|
||||
}
|
||||
__syncwarp();
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++)
|
||||
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l)]);
|
||||
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l, SWIZ_MASK)]);
|
||||
}
|
||||
__syncthreads(); // prevent next warp from overwriting sQ prematurely
|
||||
}
|
||||
@@ -106,8 +107,8 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
int r = i / HEAD_DIM;
|
||||
int d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
cp_async_16(&sK[r * LD + swiz_col(d, r)], &p.k[kv_base + kc * HEAD_DIM + d]);
|
||||
cp_async_16(&sV[r * LD + swiz_col(d, r)], &p.v[kv_base + kc * HEAD_DIM + d]);
|
||||
cp_async_16(&sK[r * LD + swiz_col(d, r, SWIZ_MASK)], &p.k[kv_base + kc * HEAD_DIM + d]);
|
||||
cp_async_16(&sV[r * LD + swiz_col(d, r, SWIZ_MASK)], &p.v[kv_base + kc * HEAD_DIM + d]);
|
||||
}
|
||||
cp_async_commit();
|
||||
cp_async_wait_all();
|
||||
@@ -116,9 +117,9 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bf16 z = __float2bfloat16(0.0f);
|
||||
sK[r * LD + swiz_col(d, r)] = (kc < p.kv_len)
|
||||
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] = (kc < p.kv_len)
|
||||
? p.k[kv_base + kc * HEAD_DIM + d] : z;
|
||||
sV[r * LD + swiz_col(d, r)] = (kc < p.kv_len)
|
||||
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] = (kc < p.kv_len)
|
||||
? p.v[kv_base + kc * HEAD_DIM + d] : z;
|
||||
}
|
||||
}
|
||||
@@ -137,7 +138,7 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l)]);
|
||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||
}
|
||||
}
|
||||
@@ -218,7 +219,7 @@ __global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l)]);
|
||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user