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 185649f4ec..e168a077e9 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -204,11 +204,11 @@ void main() { coopmat cache_a; coopmat cache_b; - coopmat int_result[cms_per_row * cms_per_col]; + coopmat int_result; coopmat scales_a; - coopmat scales_b; - coopmat scales; + coopmat scales_b[cms_per_col]; + coopmat scales[cms_per_row * cms_per_col]; coopmat sums[cms_per_row * cms_per_col]; [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col; i++) { @@ -242,13 +242,19 @@ void main() { pos_a_ib += 1; pos_b_ib += 1; - // Calculate quants + // Precompute scales + [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { + coopMatLoad(scales_b[cm_col], buf_b_d, warp_c*WN + cm_col*TN, 0, gl_CooperativeMatrixLayoutRowMajor); + } + [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { + coopMatLoad(scales_a, buf_a_d, warp_r*WM + cm_row*TM, 0, gl_CooperativeMatrixLayoutColumnMajor); [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { - int_result[cm_col * cms_per_row + cm_row] = coopmat(0); + scales[cm_col * cms_per_row + cm_row] = coopMatMulAdd(scales_a, scales_b[cm_col], coopmat(0)); } } + // Calculate quants [[unroll]] for (uint i = 0; i < BK; i += TK) { [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { coopMatLoad(cache_a, buf_a_qs, (warp_r * WM + cm_row * TM) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutRowMajor); @@ -256,21 +262,12 @@ void main() { [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { coopMatLoad(cache_b, buf_b_qs, (warp_c * WN + cm_col * TN) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutColumnMajor); - int_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, int_result[cm_col * cms_per_row + cm_row]); + int_result = coopMatMulAdd(cache_a, cache_b, coopmat(0)); + sums[cm_col * cms_per_row + cm_row] += scales[cm_col * cms_per_row + cm_row] * coopmat(int_result); } } } - // Apply scales - [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { - coopMatLoad(scales_a, buf_a_d, warp_r*WM + cm_row*TM, 0, gl_CooperativeMatrixLayoutColumnMajor); - [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { - coopMatLoad(scales_b, buf_b_d, warp_c*WN + cm_col*TN, 0, gl_CooperativeMatrixLayoutRowMajor); - scales = coopMatMulAdd(scales_a, scales_b, coopmat(0)); - sums[cm_col * cms_per_row + cm_row] += scales * coopmat(int_result[cm_col * cms_per_row + cm_row]); - } - } - barrier(); }