From 1a97c2c54d13a4611c6bdc450386fe4fe24d6d6e Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 09:33:15 +0200 Subject: [PATCH] add BK_STEP to shader, default to 2 --- .../vulkan-shaders/mul_mmq_cm1.comp | 108 ++++++++++-------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 27 +++-- 2 files changed, 71 insertions(+), 64 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 860d9fea23..9497e1f0e0 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -80,15 +80,16 @@ layout (constant_id = 9) const uint TK = 16; layout (constant_id = 10) const uint WARP = 32; #define BK 32 +#define BK_STEP 2 -const uint shmem_stride = (BK / 4) + 4; +const uint QPITCH = BK_STEP * (BK / 4) + 4; // Shared memory cache -shared uint32_t buf_a_qs[BM * shmem_stride]; -shared float16_t buf_a_d[BM]; +shared uint32_t buf_a_qs[BM * QPITCH]; +shared float16_t buf_a_d[BM * BK_STEP]; -shared uint32_t buf_b_qs[BN * shmem_stride]; -shared float16_t buf_b_d[BN]; +shared uint32_t buf_b_qs[BN * QPITCH]; +shared float16_t buf_b_d[BN * BK_STEP]; #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -228,58 +229,65 @@ void main() { sums[i] = ACC_TYPE(0.0); } - for (uint block = start_k; block < end_k; block += BK) { - [[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) { - const uint buf_ib = loadc_a + l; - const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; - const uint iqs = loadr_a; - - block_a_to_shmem(buf_ib, ib, iqs); - } - [[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) { - const uint buf_ib = loadc_b + l; - + for (uint block = start_k; block < end_k; block += BK * BK_STEP) { + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { + const bool k_in_bounds = block + ks * BK < end_k; + [[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) { + const uint buf_ib = loadc_a + l; + const uint ib = pos_a_ib + buf_ib * p.stride_a / BK + ks; + if (k_in_bounds) { + block_a_to_shmem(buf_ib, ib, loadr_a, ks); + } else if (loadr_a == 0) { + buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(0.0); + } + } + [[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) { + const uint buf_ib = loadc_b + l; #ifdef MUL_MAT_ID - const u16vec2 row_idx = row_ids[buf_ib]; - const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK; + const u16vec2 row_idx = row_ids[buf_ib]; + const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK + ks; #else - const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; + const uint ib = pos_b_ib + buf_ib * p.stride_b / BK + ks; #endif - const uint iqs = loadr_b; - - block_b_to_shmem(buf_ib, ib, iqs); - } - - barrier(); - - pos_a_ib += 1; - pos_b_ib += 1; - - [[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) { - cm_result[idx] = 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); - - [[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); - - cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]); + if (k_in_bounds) { + block_b_to_shmem(buf_ib, ib, loadr_b, ks); + } else if (loadr_b == 0) { + buf_b_d[ks * BN + buf_ib] = FLOAT_TYPE(0.0); } } } - // Apply scales directly from coopmat elements - [[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[warp_r * WM + cm_row * TM + elem_row[e]]); - const ACC_TYPE db = ACC_TYPE(buf_b_d[warp_c * WN + cm_col * TN + elem_col[e]]); - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db; + barrier(); + + pos_a_ib += BK_STEP; + pos_b_ib += BK_STEP; + + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { + [[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) { + cm_result[idx] = coopmat(0); + } + + [[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) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutRowMajor); + + [[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) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + + cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]); + } + } + } + + // Apply scales directly from coopmat elements + [[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; + } } } } 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 4e092b9bd4..3f9a1c440e 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 @@ -9,7 +9,7 @@ #if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1) // 2-byte loads for Q4_0 blocks (18 bytes) // 4-byte loads for Q4_1 blocks (20 bytes) -void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { +void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) { #ifdef DATA_A_Q4_0 const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2], data_a_packed16[ib].qs[iqs * 2 + 1])); @@ -24,12 +24,12 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; - buf_a_qs[buf_ib * shmem_stride + iqs ] = lo4; - buf_a_qs[buf_ib * shmem_stride + iqs + 4] = hi4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs + 4] = hi4; if (iqs == 0) { #ifdef DATA_A_Q4_0 - buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); + buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); #else // DATA_A_Q4_1 #endif } @@ -44,14 +44,14 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { #if defined(DATA_A_Q8_0) // 2-byte loads for Q8_0 blocks (34 bytes) -void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { +void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) { const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2], data_a_packed16[ib].qs[iqs * 2 + 1])); - buf_a_qs[buf_ib * shmem_stride + iqs] = vui; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs] = vui; if (iqs == 0) { - buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); + buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); } } #endif @@ -78,18 +78,17 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { // 2-byte loads for Q6_K blocks (210 bytes) #endif -void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { +void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) { const uint ib_outer = ib / 4; const uint ib_inner = ib % 4; if (iqs == 0) { - // Divide by TK for matmul scale application - buf_b_d[buf_ib] = data_b[ib_outer].ds[ib_inner].x; + buf_b_d[ks * BN + buf_ib] = data_b[ib_outer].ds[ib_inner].x; } const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs]; - buf_b_qs[buf_ib * shmem_stride + iqs * 4 ] = values.x; - buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 1] = values.y; - buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 2] = values.z; - buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 3] = values.w; + buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 ] = values.x; + buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 1] = values.y; + buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 2] = values.z; + buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 3] = values.w; }