From ab4e443f0f29b2b7e834da138838e5986aabc2a3 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:37:12 +0200 Subject: [PATCH] preload scales --- .../vulkan-shaders/mul_mmq_cm1.comp | 21 +++++++++++++++---- 1 file changed, 17 insertions(+), 4 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 4dc091a52d..c0fcb885aa 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -360,14 +360,27 @@ void main() { } } - // Apply scales directly from coopmat elements + // Pre-load scales into registers + ACC_TYPE scale_a[cms_per_row * CM_ELEMS]; + ACC_TYPE 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] = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]]); + } + } + [[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] = ACC_TYPE(buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]]); + } + } + + // Apply scales from registers [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { const uint tile_idx = cm_col * cms_per_row + cm_row; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - const ACC_TYPE da = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + cm_row * TM + elem_row[e]]); - const ACC_TYPE db = ACC_TYPE(buf_b_d[ks * BN + warp_c * WN + cm_col * TN + elem_col[e]]); - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db; + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) + * scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e]; } } }