From 9ef5996da43b04f273cd2536b2be1c8809d292cb Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 10:38:13 +0200 Subject: [PATCH] merge shmem arrays --- .../vulkan-shaders/mul_mmq_cm1.comp | 41 ++++++++++--------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 12 ++---- 2 files changed, 26 insertions(+), 27 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 1a1186ccce..3b7ae9a6df 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -95,13 +95,16 @@ const uint QPITCH = BK_STEP * (BK / 4) + 4; // Shared memory cache shared uint32_t buf_a_qs[BM * QPITCH]; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) +shared vec2 buf_a_dm[BM * BK_STEP]; // .x = d (== old buf_a_d), .y = m (== old buf_a_m) +#else shared float buf_a_d[BM * BK_STEP * KSCALES]; +#endif shared uint32_t buf_b_qs[BN * QPITCH]; shared float buf_b_d[BN * BK_STEP]; #if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) -shared float buf_a_m[BM * BK_STEP]; shared float buf_b_s[BN * BK_STEP]; #endif @@ -205,8 +208,8 @@ void main() { const uint warp_c = warp_i / (BM / WM); // Probe coopmat element layout to discover row/col mapping - uint elem_row[CM_ELEMS]; - uint elem_col[CM_ELEMS]; + uint8_t elem_row[CM_ELEMS]; + uint8_t elem_col[CM_ELEMS]; for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { cm_layout_probe[i] = int32_t(i); @@ -217,8 +220,8 @@ void main() { coopMatLoad(probe, cm_layout_probe, 0, TN, gl_CooperativeMatrixLayoutRowMajor); [[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) { - elem_row[e] = uint(probe[e]) / TN; - elem_col[e] = uint(probe[e]) % TN; + elem_row[e] = uint8_t(uint(probe[e]) / TN); + elem_col[e] = uint8_t(uint(probe[e]) % TN); } const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); @@ -333,7 +336,7 @@ void main() { float scale_a[CM_ELEMS]; float nbias_a[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + elem_row[e]]; + scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + uint(elem_row[e])]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -343,7 +346,7 @@ void main() { coopmat cache_b; coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); - const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; + const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])]; coopmat acc = coopmat( @@ -385,15 +388,15 @@ void main() { float nbias_a[CM_ELEMS]; float ma[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]]; + vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + uint(elem_row[e])]; + scale_a[e] = dm.x; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } #if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - ma[e] = float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]); + ma[e] = dm.y; #else - ma[e] = float(QUANT_OFFSET) * scale_a[e] - + float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]); + ma[e] = float(QUANT_OFFSET) * scale_a[e] + dm.y; #endif } @@ -407,8 +410,8 @@ void main() { } const uint tile_idx = r * cms_per_col + c; - const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; - const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col[0]]); + const float sb = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])]; + const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + uint(elem_col[0])]); [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), @@ -441,14 +444,14 @@ void main() { float scale_b[cms_per_col]; [[unroll]] for (uint c = 0; c < cms_per_col; c++) { - scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; + scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])]; } [[unroll]] for (uint r = 0; r < cms_per_row; r++) { float scale_a[CM_ELEMS]; float nbias_a[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]]; + scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + uint(elem_row[e])]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -497,10 +500,10 @@ void main() { [[unroll]] for (uint c = 0; c < cms_per_col; c++) { const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - const uint col_i = dc + c * TN + elem_col[e]; + const uint col_i = dc + c * TN + uint(elem_col[e]); if (col_i >= _ne1) continue; - const uint row_g = dr + r * TM + elem_row[e]; + const uint row_g = dr + r * TM + uint(elem_row[e]); if (row_g >= p.M) continue; const u16vec2 row_idx = row_ids[col_i - ic * BN]; @@ -516,8 +519,8 @@ void main() { [[unroll]] for (uint c = 0; c < cms_per_col; c++) { const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - const uint row_g = dr + r * TM + elem_row[e]; - const uint col_g = dc + c * TN + elem_col[e]; + const uint row_g = dr + r * TM + uint(elem_row[e]); + const uint col_g = dc + c * TN + uint(elem_col[e]); if (row_g < p.M && col_g < p.N) { data_d[offsets + col_g * p.stride_d + row_g] = D_TYPE(sums[tile_idx * CM_ELEMS + e]); } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl index 6a8d8a6c36..e2ee2795d2 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl @@ -55,8 +55,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; if (loadr == 0) { - buf_a_d[ks * BM + buf_ib] = float(blk.dm.x); - buf_a_m[ks * BM + buf_ib] = float(blk.dm.y); + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); } } @@ -119,8 +118,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; if (loadr == 0) { - buf_a_d[ks * BM + buf_ib] = float(blk.dm.x); - buf_a_m[ks * BM + buf_ib] = float(blk.dm.y); + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); } } @@ -258,8 +256,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { } vec2 dm = vec2(data_a_packed32[ib_k].dm); float d_scaled = dm.x * float(sc_val); - buf_a_d[ks * BM + buf_ib] = d_scaled; - buf_a_m[ks * BM + buf_ib] = 8.0 * d_scaled - (dm.y * float(mn_val)); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 8.0 * d_scaled - (dm.y * float(mn_val))); } } @@ -316,8 +313,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { } vec2 dm = vec2(data_a_packed32[ib_k].dm); float d_scaled = dm.x * float(sc_val); - buf_a_d[ks * BM + buf_ib] = d_scaled; - buf_a_m[ks * BM + buf_ib] = 16.0 * d_scaled - (dm.y * float(mn_val)); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 16.0 * d_scaled - (dm.y * float(mn_val))); } }