skip computation for inactive tiles

This commit is contained in:
Ruben Ortlam
2026-08-24 13:48:58 +02:00
parent 2b6c3aa6f8
commit 464754143b
@@ -334,6 +334,12 @@ void main() {
PREFETCH_BLOCK(start_k)
}
const uint a_row0 = warp_r * WM;
const uint b_col0 = warp_c * WN;
const bool active_col_tile = ic * BN + b_col0 < p.N;
barrier();
for (uint block = start_k; block < end_k; block += BK * BK_STEP) {
// Store prefetched data to shmem
STORE_BLOCK_TO_LDS(block)
@@ -349,6 +355,7 @@ void main() {
PREFETCH_BLOCK(next_block)
}
if (active_col_tile) {
// Compute from shmem
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
const uint K_SUB = BK / TK;
@@ -357,12 +364,12 @@ void main() {
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
[[unroll]] for (uint h = 0; h < K_SUB; h++) {
coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (warp_r * WM + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
}
}
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
[[unroll]] for (uint h = 0; h < K_SUB; h++) {
coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (warp_c * WN + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
}
}
@@ -371,7 +378,7 @@ void main() {
float scale_b[cms_per_col * CM_ELEMS];
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
scale_a[r * CM_ELEMS + e] = buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]];
scale_a[r * CM_ELEMS + e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
if (USE_MAGIC_BIAS) {
nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + e];
}
@@ -379,7 +386,7 @@ void main() {
}
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
scale_b[c * CM_ELEMS + e] = buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]];
scale_b[c * CM_ELEMS + e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]];
}
}
@@ -417,6 +424,7 @@ void main() {
}
}
}
}
barrier();
}
@@ -426,8 +434,8 @@ void main() {
#undef B_IB_CALC
#undef PREFETCH_BLOCK
const uint dr = ir * BM + warp_r * WM;
const uint dc = ic * BN + warp_c * WN;
const uint dr = ir * BM + a_row0;
const uint dc = ic * BN + b_col0;
#ifdef MUL_MAT_ID
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {