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:
2026-07-08 14:14:11 +08:00
parent 9ebaea840f
commit fd65b9bc23
7 changed files with 168 additions and 102 deletions
+12 -2
View File
@@ -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;
}
+6 -4
View File
@@ -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,
+4 -1
View File
@@ -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;
}
+9 -8
View File
@@ -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]);
}
}