From 464754143b18a476841406d17fae0368e88dd809 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 13:48:58 +0200 Subject: [PATCH] skip computation for inactive tiles --- .../vulkan-shaders/mul_mmq_cm1.comp | 20 +++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp index db97784599..007ea3f51f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -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++) {