perf: reduce MMA kernel registers, switch to static smem
- Move Qa[KD][4] into tile loop (reload from sQ per tile) cutting ~32 resident registers for HEAD_DIM=128 - Replace extern __shared__ with static template-sized smem (no cudaFuncSetAttribute or dynamic allocation needed) - Add __launch_bounds__ with MIN_BLOCKS param, dispatch by HEAD_DIM (hd=128→4, hd=64→6, hd=32→6) - Remove dynamic smem from scalar kernel and C test - Result: hd=128 168→128 regs, 25%→33% occupancy
This commit is contained in:
@@ -97,14 +97,13 @@ int main() {
|
||||
constexpr int G=8, ROWS=32, P_BC=32;
|
||||
dim3 grid((ql+ROWS-1)/ROWS, Hq, B);
|
||||
dim3 block(G, ROWS, 1);
|
||||
size_t smem=2*P_BC*D*sizeof(bf16);
|
||||
printf("grid=(%d,%d,%d) block=(%d,%d,%d) smem=%zu\n",
|
||||
grid.x,grid.y,grid.z, block.x,block.y,block.z, smem);
|
||||
printf("grid=(%d,%d,%d) block=(%d,%d,%d)\n",
|
||||
grid.x,grid.y,grid.z, block.x,block.y,block.z);
|
||||
|
||||
double t0=now_ms();
|
||||
switch (D) {
|
||||
case 64: gqa_prefill_attn_kernel_t<64, G,ROWS,P_BC><<<grid,block,smem>>>(p); break;
|
||||
case 128: gqa_prefill_attn_kernel_t<128,G,ROWS,P_BC><<<grid,block,smem>>>(p); break;
|
||||
case 64: gqa_prefill_attn_kernel_t<64, G,ROWS,P_BC><<<grid,block>>>(p); break;
|
||||
case 128: gqa_prefill_attn_kernel_t<128,G,ROWS,P_BC><<<grid,block>>>(p); break;
|
||||
default: printf("unsupported D=%d\n",D); return 1;
|
||||
}
|
||||
cudaDeviceSynchronize();
|
||||
|
||||
Reference in New Issue
Block a user