From 13f6d8e3891962563ed58cbdf657d90484f1879e Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 7 Mar 2026 14:42:41 +0100 Subject: [PATCH 01/50] vulkan: add int8 coopmat quantized matmul shader --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 38 +- .../vulkan-shaders/mul_mmq_cm1.comp | 329 ++++++++++++++++++ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 85 +++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 4 + 4 files changed, 451 insertions(+), 5 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 39b4cd359803..1e62923f5a12 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4308,6 +4308,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t tk_m = device->coopmat_support ? device->coopmat_k : 1; const uint32_t tk_s = device->coopmat_support ? device->coopmat_k : 1; + const uint32_t itm_l = device->coopmat_int_support ? device->coopmat_int_m : 4; + const uint32_t itm_m = device->coopmat_int_support ? device->coopmat_int_m : 4; + const uint32_t itm_s = device->coopmat_int_support ? device->coopmat_int_m : 2; + const uint32_t itn_l = device->coopmat_int_support ? device->coopmat_int_n : 4; + const uint32_t itn_m = device->coopmat_int_support ? device->coopmat_int_n : 2; + const uint32_t itn_s = device->coopmat_int_support ? device->coopmat_int_n : 1; + const uint32_t itk_l = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t itk_m = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t itk_s = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; @@ -4319,9 +4329,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; // Integer MMQ has a smaller shared memory profile, but heavier register use - l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; - m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, 2, 2, 1, mm_warp_8 }; - s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, itm_l, itn_l, itk_l, mm_warp_8 }; + m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, itm_m, itn_m, itk_m, mm_warp_8 }; + s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; @@ -4783,6 +4793,14 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (device->mul_mat ## ID ## _s[TYPE]) \ ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true); \ +#define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + if (device->mul_mat ## ID ## _l[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true); \ + if (device->mul_mat ## ID ## _m[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true); \ + if (device->mul_mat ## ID ## _s[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true); \ + // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->coopmat_acc_f16_support) { \ @@ -4792,6 +4810,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM(TYPE, PIPELINE_NAME . f32acc, NAMELC, , WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ } \ +#define CREATE_MMQ2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MMQ(TYPE, PIPELINE_NAME . f16acc, NAMELC, _f16acc, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MMQ(TYPE, PIPELINE_NAME . f32acc, NAMELC, , WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MM(GGML_TYPE_F32, pipeline_matmul_f32, matmul_f32_f32, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_F32, pipeline_matmul_f32_f16, matmul_f32_f16, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_f16, matmul_f16, wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); @@ -4833,10 +4855,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } else #endif { - CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_MXFP4], matmul_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } + if (device->coopmat_int_support) { + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_int, vk_mat_mat_push_constants, 3, ); + } + GGML_ASSERT(device->subgroup_ballot); CREATE_MM(GGML_TYPE_F32, pipeline_matmul_id_f32, matmul_id_subgroup_f32_f32, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4880,6 +4905,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } +#undef CREATE_MMQ2 +#undef CREATE_MMQ #undef CREATE_MM2 #undef CREATE_MM } else @@ -9294,7 +9321,8 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub const bool y_f32_kernel = src1->type == GGML_TYPE_F32 && !y_non_contig; - bool quantize_y = ctx->device->integer_dot_product && src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && !y_non_contig && (ne11 * ne10) % 4 == 0; + bool quantize_y = (ctx->device->integer_dot_product || ctx->device->coopmat_int_support) && + src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && !y_non_contig && (ne11 * ne10) % 4 == 0; // Check for mmq first vk_matmul_pipeline mmp = quantize_y ? ggml_vk_get_mul_mat_mat_pipeline(ctx, src0->type, GGML_TYPE_Q8_1, (ggml_prec)dst->op_params[0]) : nullptr; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp new file mode 100644 index 000000000000..185649f4ec9b --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -0,0 +1,329 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +#extension GL_KHR_shader_subgroup_basic : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : enable + +#if defined(MUL_MAT_ID_USE_SUBGROUPS) +#extension GL_KHR_shader_subgroup_basic : enable +#extension GL_KHR_shader_subgroup_ballot : enable +#endif + +#ifdef MUL_MAT_ID +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#endif + +#include "types.glsl" + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; +#if defined(A_TYPE_PACKED16) +layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];}; +#endif +#if defined(A_TYPE_PACKED32) +layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32[];}; +#endif +layout (binding = 1) readonly buffer B {block_q8_1_x4_packed128 data_b[];}; +layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; + +#ifdef MUL_MAT_ID +layout (binding = 3) readonly buffer IDS {int data_ids[];}; +layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +#endif + +layout (push_constant) uniform parameter +{ + uint M; + uint N; + uint K; + uint stride_a; + uint stride_b; + uint stride_d; + + uint batch_stride_a; + uint batch_stride_b; + uint batch_stride_d; + +#ifdef MUL_MAT_ID + uint nei0; + uint nei1; + uint nbi1; + uint ne11; +#else + uint base_work_group_z; + uint num_batches; + uint k_split; + uint ne02; + uint ne12; + uint broadcast2; + uint broadcast3; +#endif +} p; + +layout (constant_id = 0) const uint BLOCK_SIZE = 64; +layout (constant_id = 1) const uint BM = 64; +layout (constant_id = 2) const uint BN = 64; +// layout (constant_id = 3) const uint BK = 32; +layout (constant_id = 4) const uint WM = 32; +layout (constant_id = 5) const uint WN = 32; +layout (constant_id = 6) const uint WMITER = 2; +layout (constant_id = 7) const uint TM = 16; +layout (constant_id = 8) const uint TN = 16; +layout (constant_id = 9) const uint TK = 16; +layout (constant_id = 10) const uint WARP = 32; + +#define BK 32 + +const uint shmem_stride = (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_b_qs[BN * shmem_stride]; +shared float16_t buf_b_d[BN]; + +#define LOAD_VEC_A (4 * QUANT_R) +#define LOAD_VEC_B 16 + +#define NUM_WARPS (BLOCK_SIZE / WARP) + +shared ACC_TYPE coopmat_stage[TM * TN * NUM_WARPS]; + +#include "mul_mm_id_funcs.glsl" +#include "mul_mmq_cm1_funcs.glsl" + +void main() { + const uint ic = gl_WorkGroupID.y; + +#ifdef MUL_MAT_ID + const uint expert_idx = gl_WorkGroupID.z; + if (ic * BN >= data_expert_count[expert_idx]) { + return; + } +#endif +#ifdef NEEDS_INIT_IQ_SHMEM + init_iq_shmem(gl_WorkGroupSize); +#endif + +#ifndef MUL_MAT_ID + const uint batch_idx = gl_WorkGroupID.z + p.base_work_group_z; + + const uint i13 = batch_idx / p.ne12; + const uint i12 = batch_idx % p.ne12; + + const uint i03 = i13 / p.broadcast3; + const uint i02 = i12 / p.broadcast2; + + const uint batch_idx_a = i03 * p.ne02 + i02; +#endif + + const uint blocks_m = (p.M + BM - 1) / BM; + const uint ir = gl_WorkGroupID.x % blocks_m; + const uint ik = gl_WorkGroupID.x / blocks_m; + + const uint WNITER = (WM * WN) / (WARP * TM * TN * WMITER); + const uint WSUBM = WM / WMITER; + const uint WSUBN = WN / WNITER; + + const uint warp_i = gl_SubgroupID; + + const uint tiw = gl_SubgroupInvocationID; + + const uint cms_per_row = WM / TM; + const uint cms_per_col = WN / TN; + + const uint storestride = WARP / TM; + const uint store_r = tiw % TM; + const uint store_c = tiw / TM; + + const uint warp_r = warp_i % (BM / WM); + const uint warp_c = warp_i / (BM / WM); + + const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); + const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A); + const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B); + const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B); + + const uint loadstride_a = BLOCK_SIZE * LOAD_VEC_A / BK; + const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK; + +#ifdef MUL_MAT_ID +#ifdef MUL_MAT_ID_USE_SUBGROUPS + if (bitCount(p.nei0) == 1) { + load_row_ids(expert_idx, true, ic); + } else { + load_row_ids(expert_idx, false, ic); + } +#else + _ne1 = 0; + for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { + for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { + if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { + if (_ne1 >= ic * BN) { + row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + } + _ne1++; + } + } + } + + barrier(); +#endif + + // Workgroup has no work + if (ic * BN >= _ne1) return; +#endif + +#ifdef MUL_MAT_ID + const uint start_k = 0; + const uint end_k = p.K; +#else + const uint start_k = ik * p.k_split; + const uint end_k = min(p.K, (ik + 1) * p.k_split); +#endif + + uint pos_a_ib = +#ifdef MUL_MAT_ID + expert_idx * (p.batch_stride_a / BK) + +#else + batch_idx_a * (p.batch_stride_a / BK) + +#endif + (ir * BM * p.stride_a + start_k) / BK; +#ifdef MUL_MAT_ID + uint pos_b_ib = 0; +#else + uint pos_b_ib = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / BK; +#endif + + coopmat cache_a; + coopmat cache_b; + coopmat int_result[cms_per_row * cms_per_col]; + + coopmat scales_a; + coopmat scales_b; + coopmat scales; + coopmat sums[cms_per_row * cms_per_col]; + + [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col; i++) { + sums[i] = coopmat(0.0f); + } + + 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; + +#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; +#else + const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; +#endif + const uint iqs = loadr_b; + + block_b_to_shmem(buf_ib, ib, iqs); + } + + barrier(); + + pos_a_ib += 1; + pos_b_ib += 1; + + // Calculate quants + [[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++) { + int_result[cm_col * cms_per_row + cm_row] = 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) * 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); + + int_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, int_result[cm_col * cms_per_row + cm_row]); + } + } + } + + // 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(); + } + + const uint dr = ir * BM + warp_r * WM; + const uint dc = ic * BN + warp_c * WN; + +#ifdef MUL_MAT_ID + [[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++) { + coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); + + [[unroll]] for (uint col = 0; col < TN; col += storestride) { + const uint row_i = dc + cm_col * TN + col + store_c; + if (row_i >= _ne1) break; + + const u16vec2 row_idx = row_ids[row_i - ic * BN]; + + if (dr + cm_row * TM + store_r < p.M) { + data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + } + } + } + } +#else + const uint offsets = batch_idx * p.batch_stride_d + ik * p.batch_stride_d * p.num_batches; + const bool is_aligned = p.stride_d % 4 == 0; // Assumption: D_TYPE == float + + [[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 bool is_in_bounds = dr + (cm_row + 1) * TM <= p.M && dc + (cm_col + 1) * TN <= p.N; + + if (is_aligned && is_in_bounds) { + // Full coopMat is within bounds and stride_d is aligned with 16B + coopmat cm_dtype = coopmat(sums[cm_col * cms_per_row + cm_row]); + coopMatStore(cm_dtype, data_d, offsets + (dc + cm_col * TN) * p.stride_d + dr + cm_row * TM, p.stride_d, gl_CooperativeMatrixLayoutColumnMajor); + } else if (is_in_bounds) { + // Full coopMat is within bounds, but stride_d is not aligned + coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); + + [[unroll]] for (uint col = 0; col < TN; col += storestride) { + data_d[offsets + (dc + cm_col * TN + col + store_c) * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + } + } else if (dr + cm_row * TM < p.M && dc + cm_col * TN < p.N) { + // Partial coopMat is within bounds + coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); + + [[unroll]] for (uint col = 0; col < TN; col += storestride) { + if (dr + cm_row * TM + store_r < p.M && dc + cm_col * TN + col + store_c < p.N) { + data_d[offsets + (dc + cm_col * TN + col + store_c) * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + } + } + } + } + } +#endif // MUL_MAT_ID +} 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 new file mode 100644 index 000000000000..415eb40f3d79 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl @@ -0,0 +1,85 @@ +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require + +#include "types.glsl" + +// Each iqs value maps to a 32-bit integer + +#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) { +#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])); +#else // DATA_A_Q4_1 + const uint32_t vui = data_a_packed32[ib].qs[iqs]; +#endif + + uint32_t lo4 = vui & 0x0F0F0F0F; + uint32_t hi4 = (vui >> 4) & 0x0F0F0F0F; + + // subtract 8 from each byte + 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; + + if (iqs == 0) { +#ifdef DATA_A_Q4_0 + buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); +#else // DATA_A_Q4_1 +#endif + } +} +#endif + +#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) +// 2-byte loads for Q5_0 blocks (22 bytes) +// 4-byte loads for Q5_1 blocks (24 bytes) +} +#endif + +#if defined(DATA_A_Q8_0) +// 2-byte loads for Q8_0 blocks (34 bytes) +#endif + +#if defined(DATA_A_MXFP4) +// 1-byte loads for mxfp4 blocks (17 bytes) +#endif + +// For k-quants, ib and iqs still assume 32-wide blocks, but k-quants are 256-wide +// iqs still refers to a 32-bit integer, meaning 0..7 for 32-wide quants +#if defined(DATA_A_Q2_K) +// 4-byte loads for Q2_K blocks (84 bytes) +#endif + +#if defined(DATA_A_Q3_K) +// 2-byte loads for Q3_K blocks (110 bytes) +#endif + +#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) +// 4-byte loads for Q4_K blocks (144 bytes) and Q5_K blocks (176 bytes) +#endif + +#if defined(DATA_A_Q6_K) +// 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) { + 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 / float16_t(TK); + } + + 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; +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index d375c2d12771..cd0066f8aa35 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -627,6 +627,10 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"},}), fp16, coopmat, coopmat2, f16acc); } #endif + + if (coopmat && tname == "q4_0") { + string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"},}), fp16, coopmat, coopmat2, f16acc); + } } } From e086b9ca0552cc1b59d9b57b3d36207f4dc6c303 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 7 Mar 2026 14:56:25 +0100 Subject: [PATCH 02/50] apply scales inline --- .../vulkan-shaders/mul_mmq_cm1.comp | 29 +++++++++---------- 1 file changed, 13 insertions(+), 16 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 185649f4ec9b..e168a077e90f 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(); } From 5f034d6086bd3b5413c4c3b8e26cb89170584700 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 7 Mar 2026 22:11:40 +0100 Subject: [PATCH 03/50] use scalar sums --- .../vulkan-shaders/mul_mmq_cm1.comp | 160 ++++++++++++------ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 2 +- .../vulkan-shaders/vulkan-shaders-gen.cpp | 7 +- 3 files changed, 117 insertions(+), 52 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 e168a077e90f..e5f0733c128c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -31,6 +31,7 @@ layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32 #endif layout (binding = 1) readonly buffer B {block_q8_1_x4_packed128 data_b[];}; layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; +layout (binding = 2) writeonly buffer D4 {D_TYPE_VEC4 data_dv4[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; @@ -94,7 +95,7 @@ shared float16_t buf_b_d[BN]; #define NUM_WARPS (BLOCK_SIZE / WARP) -shared ACC_TYPE coopmat_stage[TM * TN * NUM_WARPS]; +shared ivec4 coopmat_stage[TM * TN * NUM_WARPS / 4]; #include "mul_mm_id_funcs.glsl" #include "mul_mmq_cm1_funcs.glsl" @@ -204,17 +205,17 @@ void main() { coopmat cache_a; coopmat cache_b; - coopmat int_result; + coopmat cm_result[cms_per_row * cms_per_col]; - coopmat scales_a; - coopmat scales_b[cms_per_col]; - coopmat scales[cms_per_row * cms_per_col]; - coopmat sums[cms_per_row * cms_per_col]; + const uint accs_per_thread = (WM * WN) / WARP / 4; + ACC_TYPE_VEC4 sums[accs_per_thread]; - [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col; i++) { - sums[i] = coopmat(0.0f); + [[unroll]] for (uint i = 0; i < accs_per_thread; i++) { + sums[i] = ACC_TYPE_VEC4(0.0f); } + const uint chunks_per_thread_per_tile = (TM * TN) / (WARP * 4); + 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; @@ -242,16 +243,8 @@ void main() { pos_a_ib += 1; pos_b_ib += 1; - // 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++) { - scales[cm_col * cms_per_row + cm_row] = coopMatMulAdd(scales_a, scales_b[cm_col], coopmat(0)); - } + [[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) { + cm_result[idx] = coopmat(0); } // Calculate quants @@ -262,8 +255,37 @@ 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 = 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); + cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]); + } + } + } + + // Store to shmem + const uint subgroup_vec_stride = (TM * TN) / 4; + const uint subgroup_offset = warp_i * subgroup_vec_stride; + + [[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; + coopMatStore(cm_result[tile_idx], coopmat_stage, subgroup_offset, TM / 4, gl_CooperativeMatrixLayoutColumnMajor); + + controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); + + // Each thread grabs chunks and applies the scales + [[unroll]] for (uint chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { + const uint local_chunk = chunk * WARP + tiw; + const uint col_local = local_chunk / (TM / 4); + const uint row_group = local_chunk % (TM / 4); + const uint row0_local = row_group * 4; + const ivec4 qs = coopmat_stage[subgroup_offset + col_local * (TM / 4) + row_group]; + + const uint a_row0 = warp_r * WM + cm_row * TM + row0_local; + const uint b_col = warp_c * WN + cm_col * TN + col_local; + + const ACC_TYPE_VEC4 da = ACC_TYPE_VEC4(buf_a_d[a_row0], buf_a_d[a_row0+1], buf_a_d[a_row0+2], buf_a_d[a_row0+3]); + const ACC_TYPE db = ACC_TYPE(buf_b_d[b_col]); + + sums[tile_idx * chunks_per_thread_per_tile + chunk] += ACC_TYPE_VEC4(qs) * da * db; } } } @@ -274,50 +296,92 @@ void main() { const uint dr = ir * BM + warp_r * WM; const uint dc = ic * BN + warp_c * WN; + const bool is_aligned = p.stride_d % 4 == 0; + #ifdef MUL_MAT_ID [[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++) { - coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); + const uint tile_idx = cm_col * cms_per_row + cm_row; + [[unroll]] for (uint chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { + const uint local_chunk = chunk * WARP + tiw; + const uint col_local = local_chunk / (TM / 4); + const uint row_group = local_chunk % (TM / 4); + const uint row0_local = row_group * 4; + + const uint row_i = dc + cm_col * TN + col_local; - [[unroll]] for (uint col = 0; col < TN; col += storestride) { - const uint row_i = dc + cm_col * TN + col + store_c; if (row_i >= _ne1) break; + const uint row0_g = dr + cm_row * TM + row0_local; const u16vec2 row_idx = row_ids[row_i - ic * BN]; - - if (dr + cm_row * TM + store_r < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + const uint store_offset = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + row0_g; + const uint acc_idx = tile_idx * chunks_per_thread_per_tile + chunk; + + if (row0_g + 3 < p.M && is_aligned && (store_offset % 4) == 0) { + data_dv4[store_offset / 4] = D_TYPE_VEC4(sums[acc_idx]); + } else if (row0_g + 3 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + data_d[store_offset + 2] = D_TYPE(vals.z); + data_d[store_offset + 3] = D_TYPE(vals.w); + } else if (row0_g + 2 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + data_d[store_offset + 2] = D_TYPE(vals.z); + } else if (row0_g + 1 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + } else if (row0_g < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset] = D_TYPE(vals.x); } } } } #else const uint offsets = batch_idx * p.batch_stride_d + ik * p.batch_stride_d * p.num_batches; - const bool is_aligned = p.stride_d % 4 == 0; // Assumption: D_TYPE == float [[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 bool is_in_bounds = dr + (cm_row + 1) * TM <= p.M && dc + (cm_col + 1) * TN <= p.N; - - if (is_aligned && is_in_bounds) { - // Full coopMat is within bounds and stride_d is aligned with 16B - coopmat cm_dtype = coopmat(sums[cm_col * cms_per_row + cm_row]); - coopMatStore(cm_dtype, data_d, offsets + (dc + cm_col * TN) * p.stride_d + dr + cm_row * TM, p.stride_d, gl_CooperativeMatrixLayoutColumnMajor); - } else if (is_in_bounds) { - // Full coopMat is within bounds, but stride_d is not aligned - coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); - - [[unroll]] for (uint col = 0; col < TN; col += storestride) { - data_d[offsets + (dc + cm_col * TN + col + store_c) * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); - } - } else if (dr + cm_row * TM < p.M && dc + cm_col * TN < p.N) { - // Partial coopMat is within bounds - coopMatStore(sums[cm_col * cms_per_row + cm_row], coopmat_stage, warp_i * TM * TN, TM, gl_CooperativeMatrixLayoutColumnMajor); - - [[unroll]] for (uint col = 0; col < TN; col += storestride) { - if (dr + cm_row * TM + store_r < p.M && dc + cm_col * TN + col + store_c < p.N) { - data_d[offsets + (dc + cm_col * TN + col + store_c) * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); - } + const uint tile_idx = cm_col * cms_per_row + cm_row; + [[unroll]] for (uint chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { + const uint local_chunk = chunk * WARP + tiw; + const uint col_local = local_chunk / (TM / 4); + const uint row_group = local_chunk % (TM / 4); + const uint row0_local = row_group * 4; + + const uint col_g = dc + cm_col * TN + col_local; + + if (col_g >= p.N) break; + + const uint row0_g = dr + cm_row * TM + row0_local; + + const uint store_offset = offsets + col_g * p.stride_d + row0_g; + const uint acc_idx = tile_idx * chunks_per_thread_per_tile + chunk; + + if (row0_g + 3 < p.M && is_aligned && (store_offset % 4) == 0) { + data_dv4[store_offset / 4] = D_TYPE_VEC4(sums[acc_idx]); + } else if (row0_g + 3 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + data_d[store_offset + 2] = D_TYPE(vals.z); + data_d[store_offset + 3] = D_TYPE(vals.w); + } else if (row0_g + 2 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + data_d[store_offset + 2] = D_TYPE(vals.z); + } else if (row0_g + 1 < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset ] = D_TYPE(vals.x); + data_d[store_offset + 1] = D_TYPE(vals.y); + } else if (row0_g < p.M) { + const ACC_TYPE_VEC4 vals = sums[acc_idx]; + data_d[store_offset] = D_TYPE(vals.x); } } } 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 415eb40f3d79..99a674b960e6 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 @@ -74,7 +74,7 @@ void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs) { if (iqs == 0) { // Divide by TK for matmul scale application - buf_b_d[buf_ib] = data_b[ib_outer].ds[ib_inner].x / float16_t(TK); + buf_b_d[buf_ib] = data_b[ib_outer].ds[ib_inner].x; } const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index cd0066f8aa35..e0a3f4c9f70b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -468,8 +468,9 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c base_dict["FLOAT16"] = "1"; } - base_dict["ACC_TYPE" ] = f16acc ? "float16_t" : "float"; - base_dict["ACC_TYPEV2"] = f16acc ? "f16vec2" : "vec2"; + base_dict["ACC_TYPE" ] = f16acc ? "float16_t" : "float"; + base_dict["ACC_TYPEV2" ] = f16acc ? "f16vec2" : "vec2"; + base_dict["ACC_TYPE_VEC4"] = f16acc ? "f16vec4" : "vec4"; if (f16acc) { base_dict["ACC_TYPE_MAX"] = "float16_t(65504.0)"; } @@ -629,7 +630,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c #endif if (coopmat && tname == "q4_0") { - string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"},}), fp16, coopmat, coopmat2, f16acc); + string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); } } } From 162560885a6fa0714d869907c32cfce3620387d0 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Fri, 21 Aug 2026 14:58:56 +0200 Subject: [PATCH 04/50] probe and directly access coopmat values instead of going through shmem Co-authored-by: Piotr Wilkin (ilintar) --- .../vulkan-shaders/mul_mmq_cm1.comp | 152 +++++------------- 1 file changed, 44 insertions(+), 108 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 e5f0733c128c..860d9fea23f6 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -94,8 +94,9 @@ shared float16_t buf_b_d[BN]; #define LOAD_VEC_B 16 #define NUM_WARPS (BLOCK_SIZE / WARP) +const uint CM_ELEMS = (TM * TN) / WARP; -shared ivec4 coopmat_stage[TM * TN * NUM_WARPS / 4]; +shared int32_t cm_layout_probe[TM * TN]; #include "mul_mm_id_funcs.glsl" #include "mul_mmq_cm1_funcs.glsl" @@ -140,13 +141,28 @@ void main() { const uint cms_per_row = WM / TM; const uint cms_per_col = WN / TN; - const uint storestride = WARP / TM; - const uint store_r = tiw % TM; - const uint store_c = tiw / TM; - const uint warp_r = warp_i % (BM / WM); 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]; + + for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { + cm_layout_probe[i] = int32_t(i); + } + barrier(); + + coopmat probe; + 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; + } + + barrier(); + const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A); const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B); @@ -207,15 +223,11 @@ void main() { coopmat cache_b; coopmat cm_result[cms_per_row * cms_per_col]; - const uint accs_per_thread = (WM * WN) / WARP / 4; - ACC_TYPE_VEC4 sums[accs_per_thread]; - - [[unroll]] for (uint i = 0; i < accs_per_thread; i++) { - sums[i] = ACC_TYPE_VEC4(0.0f); + ACC_TYPE sums[cms_per_row * cms_per_col * CM_ELEMS]; + [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col * CM_ELEMS; i++) { + sums[i] = ACC_TYPE(0.0); } - const uint chunks_per_thread_per_tile = (TM * TN) / (WARP * 4); - 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; @@ -260,32 +272,14 @@ void main() { } } - // Store to shmem - const uint subgroup_vec_stride = (TM * TN) / 4; - const uint subgroup_offset = warp_i * subgroup_vec_stride; - + // 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; - coopMatStore(cm_result[tile_idx], coopmat_stage, subgroup_offset, TM / 4, gl_CooperativeMatrixLayoutColumnMajor); - - controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); - - // Each thread grabs chunks and applies the scales - [[unroll]] for (uint chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { - const uint local_chunk = chunk * WARP + tiw; - const uint col_local = local_chunk / (TM / 4); - const uint row_group = local_chunk % (TM / 4); - const uint row0_local = row_group * 4; - const ivec4 qs = coopmat_stage[subgroup_offset + col_local * (TM / 4) + row_group]; - - const uint a_row0 = warp_r * WM + cm_row * TM + row0_local; - const uint b_col = warp_c * WN + cm_col * TN + col_local; - - const ACC_TYPE_VEC4 da = ACC_TYPE_VEC4(buf_a_d[a_row0], buf_a_d[a_row0+1], buf_a_d[a_row0+2], buf_a_d[a_row0+3]); - const ACC_TYPE db = ACC_TYPE(buf_b_d[b_col]); - - sums[tile_idx * chunks_per_thread_per_tile + chunk] += ACC_TYPE_VEC4(qs) * da * db; + [[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; } } } @@ -296,48 +290,20 @@ void main() { const uint dr = ir * BM + warp_r * WM; const uint dc = ic * BN + warp_c * WN; - const bool is_aligned = p.stride_d % 4 == 0; - #ifdef MUL_MAT_ID [[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 chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { - const uint local_chunk = chunk * WARP + tiw; - const uint col_local = local_chunk / (TM / 4); - const uint row_group = local_chunk % (TM / 4); - const uint row0_local = row_group * 4; - - const uint row_i = dc + cm_col * TN + col_local; - - if (row_i >= _ne1) break; - - const uint row0_g = dr + cm_row * TM + row0_local; - const u16vec2 row_idx = row_ids[row_i - ic * BN]; - const uint store_offset = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + row0_g; - const uint acc_idx = tile_idx * chunks_per_thread_per_tile + chunk; - - if (row0_g + 3 < p.M && is_aligned && (store_offset % 4) == 0) { - data_dv4[store_offset / 4] = D_TYPE_VEC4(sums[acc_idx]); - } else if (row0_g + 3 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - data_d[store_offset + 2] = D_TYPE(vals.z); - data_d[store_offset + 3] = D_TYPE(vals.w); - } else if (row0_g + 2 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - data_d[store_offset + 2] = D_TYPE(vals.z); - } else if (row0_g + 1 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - } else if (row0_g < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset] = D_TYPE(vals.x); - } + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + const uint col_i = dc + cm_col * TN + elem_col[e]; + if (col_i >= _ne1) continue; + + const uint row_g = dr + cm_row * TM + elem_row[e]; + if (row_g >= p.M) continue; + + const u16vec2 row_idx = row_ids[col_i - ic * BN]; + const uint store_offset = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + row_g; + data_d[store_offset] = D_TYPE(sums[tile_idx * CM_ELEMS + e]); } } } @@ -347,41 +313,11 @@ void main() { [[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 chunk = 0; chunk < chunks_per_thread_per_tile; chunk++) { - const uint local_chunk = chunk * WARP + tiw; - const uint col_local = local_chunk / (TM / 4); - const uint row_group = local_chunk % (TM / 4); - const uint row0_local = row_group * 4; - - const uint col_g = dc + cm_col * TN + col_local; - - if (col_g >= p.N) break; - - const uint row0_g = dr + cm_row * TM + row0_local; - - const uint store_offset = offsets + col_g * p.stride_d + row0_g; - const uint acc_idx = tile_idx * chunks_per_thread_per_tile + chunk; - - if (row0_g + 3 < p.M && is_aligned && (store_offset % 4) == 0) { - data_dv4[store_offset / 4] = D_TYPE_VEC4(sums[acc_idx]); - } else if (row0_g + 3 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - data_d[store_offset + 2] = D_TYPE(vals.z); - data_d[store_offset + 3] = D_TYPE(vals.w); - } else if (row0_g + 2 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - data_d[store_offset + 2] = D_TYPE(vals.z); - } else if (row0_g + 1 < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset ] = D_TYPE(vals.x); - data_d[store_offset + 1] = D_TYPE(vals.y); - } else if (row0_g < p.M) { - const ACC_TYPE_VEC4 vals = sums[acc_idx]; - data_d[store_offset] = D_TYPE(vals.x); + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + const uint row_g = dr + cm_row * TM + elem_row[e]; + const uint col_g = dc + cm_col * TN + 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]); } } } From 513f43600ce6f2af06276bac9d2e2f30ed2394b0 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 09:23:19 +0200 Subject: [PATCH 05/50] add q8_0 support --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 1 + .../ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl | 10 ++++++++++ .../ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 3 files changed, 12 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 1e62923f5a12..cd95cbb9b204 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4860,6 +4860,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (device->coopmat_int_support) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_int, vk_mat_mat_push_constants, 3, ); } GGML_ASSERT(device->subgroup_ballot); 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 99a674b960e6..4e092b9bd413 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 @@ -44,6 +44,16 @@ 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) { + 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; + + if (iqs == 0) { + buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); + } +} #endif #if defined(DATA_A_MXFP4) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index e0a3f4c9f70b..5c863dee4dea 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -629,7 +629,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && tname == "q4_0") { + if (coopmat && (tname == "q4_0" || tname == "q8_0")) { string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); } } From b121ef3b3381134ffac8628796c576037ccd7199 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 09:33:15 +0200 Subject: [PATCH 06/50] add BK_STEP to shader, default to 2 Co-authored-by: Piotr Wilkin (ilintar) --- .../vulkan-shaders/mul_mmq_cm1.comp | 94 ++++++++++--------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 27 +++--- 2 files changed, 64 insertions(+), 57 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 860d9fea23f6..9497e1f0e0b6 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); + 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); + } + } } barrier(); - pos_a_ib += 1; - pos_b_ib += 1; + pos_a_ib += BK_STEP; + pos_b_ib += BK_STEP; - [[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) { - cm_result[idx] = coopmat(0); - } + [[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); + } - // 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 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) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutColumnMajor); + [[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]); + 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[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; + // 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 4e092b9bd413..3f9a1c440e76 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; } From 2a878d9ea998acccbddb7c2bf96a482025ef1f2e Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 09:57:40 +0200 Subject: [PATCH 07/50] use larger workgroups Co-authored-by: Piotr Wilkin (ilintar) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 10 ++++++++-- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 8 ++++---- 2 files changed, 12 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index cd95cbb9b204..a18d759d6525 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4239,6 +4239,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_warptile_id, m_warptile_id, s_warptile_id, l_warptile_mmq, m_warptile_mmq, s_warptile_mmq, l_warptile_mmq_int, m_warptile_mmq_int, s_warptile_mmq_int, + l_warptile_mmq_cm1_int, m_warptile_mmq_cm1_int, s_warptile_mmq_cm1_int, l_warptile_mmq_int_k, m_warptile_mmq_int_k, s_warptile_mmq_int_k, l_warptile_mmq_k, m_warptile_mmq_k, s_warptile_mmq_k, l_warptile_mmqid, m_warptile_mmqid, s_warptile_mmqid, @@ -4333,6 +4334,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, itm_m, itn_m, itk_m, mm_warp_8 }; s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; + // Coopmat int8 cm1 shader uses larger workgroups for better occupancy + l_warptile_mmq_cm1_int = { subgroup_size_8 * 8, 128, 128, 32, 64, 32, 2, itm_l, itn_l, itk_l, subgroup_size_8 }; + m_warptile_mmq_cm1_int = { subgroup_size_8 * 4, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, subgroup_size_8 }; + s_warptile_mmq_cm1_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; + // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; m_warptile_mmq_int_k = { 128, 64, 64, 32, mm_warp_8, 32, 1, 2, 2, 1, mm_warp_8 }; @@ -4859,8 +4865,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } if (device->coopmat_int_support) { - CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } GGML_ASSERT(device->subgroup_ballot); 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 9497e1f0e0b6..1ea5bae70e5c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -67,11 +67,11 @@ layout (push_constant) uniform parameter #endif } p; -layout (constant_id = 0) const uint BLOCK_SIZE = 64; -layout (constant_id = 1) const uint BM = 64; -layout (constant_id = 2) const uint BN = 64; +layout (constant_id = 0) const uint BLOCK_SIZE = 256; +layout (constant_id = 1) const uint BM = 128; +layout (constant_id = 2) const uint BN = 128; // layout (constant_id = 3) const uint BK = 32; -layout (constant_id = 4) const uint WM = 32; +layout (constant_id = 4) const uint WM = 64; layout (constant_id = 5) const uint WN = 32; layout (constant_id = 6) const uint WMITER = 2; layout (constant_id = 7) const uint TM = 16; From 5553b8991202bf7fc28c7b8c79d04511f310aeac Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:30:55 +0200 Subject: [PATCH 08/50] double buffering Co-authored-by: Piotr Wilkin (ilintar) --- .../vulkan-shaders/mul_mmq_cm1.comp | 134 ++++++++++++++---- 1 file changed, 110 insertions(+), 24 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 1ea5bae70e5c..4dc091a52dcc 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -229,39 +229,120 @@ void main() { sums[i] = ACC_TYPE(0.0); } - 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; + // Double-buffering: prefetch registers + const uint A_LOADS = (BM + loadstride_a - 1) / loadstride_a; + const uint B_LOADS = (BN + loadstride_b - 1) / loadstride_b; + + uint32_t pre_a_qs[A_LOADS * BK_STEP]; + float16_t pre_a_d [A_LOADS * BK_STEP]; + ivec4 pre_b_qs[B_LOADS * BK_STEP]; + float16_t pre_b_d [B_LOADS * BK_STEP]; + +// Prefetch: global memory → registers +#define PREFETCH_BLOCK(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + pre_a_qs[li * BK_STEP + ks] = \ + pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ + data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ + pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + B_IB_CALC \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \ + const uint ib_outer = ib_k / 4; \ + const uint ib_inner = ib_k % 4; \ + pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ + pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ + } \ + } \ + } + #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 + ks; +#define B_IB_CALC \ + 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; #else - const uint ib = pos_b_ib + buf_ib * p.stride_b / BK + ks; +#define B_IB_CALC \ + const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; #endif - 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); - } - } - } + +// Store: registers → shared memory (with quant-specific unpacking) +#define STORE_BLOCK_TO_LDS(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const uint idx = li * BK_STEP + ks; \ + STORE_A_QS(buf_ib, ks, idx) \ + if (loadr_a == 0) { \ + buf_a_d[ks * BM + buf_ib] = pre_a_d[idx]; \ + } \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const bool in_bounds = (blk) + ks * BK < end_k; \ + const uint idx = li * BK_STEP + ks; \ + const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \ + const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \ + buf_b_qs[base ] = v.x; \ + buf_b_qs[base + 1] = v.y; \ + buf_b_qs[base + 2] = v.z; \ + buf_b_qs[base + 3] = v.w; \ + if (loadr_b == 0) { \ + buf_b_d[ks * BN + buf_ib] = in_bounds ? pre_b_d[idx] : float16_t(0.0); \ + } \ + } \ + } \ + } + +#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1) +#define STORE_A_QS(buf_ib, ks, idx) \ + uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ + uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ + lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ + hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; +#elif defined(DATA_A_Q8_0) +#define STORE_A_QS(buf_ib, ks, idx) \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; +#endif + + // Prefetch first block + if (start_k < end_k) { + PREFETCH_BLOCK(start_k) + } + + for (uint block = start_k; block < end_k; block += BK * BK_STEP) { + // Store prefetched data to shmem + STORE_BLOCK_TO_LDS(block) barrier(); pos_a_ib += BK_STEP; pos_b_ib += BK_STEP; + // Prefetch next block (overlaps with compute) + const uint next_block = block + BK * BK_STEP; + if (next_block < end_k) { + PREFETCH_BLOCK(next_block) + } + + // Compute from shmem [[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); @@ -295,6 +376,11 @@ void main() { barrier(); } +#undef STORE_A_QS +#undef STORE_BLOCK_TO_LDS +#undef B_IB_CALC +#undef PREFETCH_BLOCK + const uint dr = ir * BM + warp_r * WM; const uint dc = ic * BN + warp_c * WN; From b453be38a3369236da27ea9fc93057a219585b08 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:37:12 +0200 Subject: [PATCH 09/50] 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 4dc091a52dcc..c0fcb885aa68 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]; } } } From 14a6080dd59f459412e1d79ad8343f31d201f7da Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:39:25 +0200 Subject: [PATCH 10/50] coopmat load first, then wmma --- .../vulkan-shaders/mul_mmq_cm1.comp | 25 +++++++++++++------ 1 file changed, 17 insertions(+), 8 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 c0fcb885aa68..38f38a2b26ad 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -220,8 +220,6 @@ void main() { uint pos_b_ib = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / BK; #endif - coopmat cache_a; - coopmat cache_b; coopmat cm_result[cms_per_row * cms_per_col]; ACC_TYPE sums[cms_per_row * cms_per_col * CM_ELEMS]; @@ -348,14 +346,25 @@ void main() { 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); + const uint K_SUB = BK / TK; + coopmat all_a[cms_per_row * K_SUB]; + coopmat all_b[cms_per_col * K_SUB]; - [[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); + [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(all_a[cm_row * K_SUB + h], buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + } + } + [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(all_b[cm_col * K_SUB + h], buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + h * (TK / 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]); + [[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++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(all_a[cm_row * K_SUB + h], all_b[cm_col * K_SUB + h], cm_result[cm_col * cms_per_row + cm_row]); } } } From 95d14c99ead21c9d7ddf9a61318b98ccc0c7f523 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:41:33 +0200 Subject: [PATCH 11/50] use float for scales --- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 8 ++++---- .../src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl | 6 +++--- 2 files changed, 7 insertions(+), 7 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 38f38a2b26ad..fb63acc06c1f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -86,10 +86,10 @@ const uint QPITCH = BK_STEP * (BK / 4) + 4; // Shared memory cache shared uint32_t buf_a_qs[BM * QPITCH]; -shared float16_t buf_a_d[BM * BK_STEP]; +shared float buf_a_d[BM * BK_STEP]; shared uint32_t buf_b_qs[BN * QPITCH]; -shared float16_t buf_b_d[BN * BK_STEP]; +shared float buf_b_d[BN * BK_STEP]; #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -283,7 +283,7 @@ void main() { const uint idx = li * BK_STEP + ks; \ STORE_A_QS(buf_ib, ks, idx) \ if (loadr_a == 0) { \ - buf_a_d[ks * BM + buf_ib] = pre_a_d[idx]; \ + buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ } \ } \ } \ @@ -301,7 +301,7 @@ void main() { buf_b_qs[base + 2] = v.z; \ buf_b_qs[base + 3] = v.w; \ if (loadr_b == 0) { \ - buf_b_d[ks * BN + buf_ib] = in_bounds ? pre_b_d[idx] : float16_t(0.0); \ + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ } \ } \ } \ 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 3f9a1c440e76..14cb6bfd87d1 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 @@ -29,7 +29,7 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const ui if (iqs == 0) { #ifdef DATA_A_Q4_0 - buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); + buf_a_d[ks * BM + buf_ib] = float(data_a_packed16[ib].d); #else // DATA_A_Q4_1 #endif } @@ -51,7 +51,7 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const ui buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs] = vui; if (iqs == 0) { - buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d); + buf_a_d[ks * BM + buf_ib] = float(data_a_packed16[ib].d); } } #endif @@ -83,7 +83,7 @@ void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const ui const uint ib_inner = ib % 4; if (iqs == 0) { - buf_b_d[ks * BN + buf_ib] = data_b[ib_outer].ds[ib_inner].x; + buf_b_d[ks * BN + buf_ib] = float(data_b[ib_outer].ds[ib_inner].x); } const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs]; From ad015d6a44b99b89c682486c01bb5fab6323deac Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:45:00 +0200 Subject: [PATCH 12/50] add faster RDNA int->float conversion --- .../vulkan-shaders/mul_mmq_cm1.comp | 30 ++++++++++++++----- 1 file changed, 23 insertions(+), 7 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 fb63acc06c1f..7185a3a68690 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -96,6 +96,9 @@ shared float buf_b_d[BN * BK_STEP]; #define NUM_WARPS (BLOCK_SIZE / WARP) const uint CM_ELEMS = (TM * TN) / WARP; +#define ACC_BIAS_BITS 0x4B400000 +#define ACC_BIAS_F 12582912.0f +const bool USE_MAGIC_BIAS = WARP != 32; shared int32_t cm_layout_probe[TM * TN]; @@ -343,7 +346,8 @@ void main() { // Compute from shmem [[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); + cm_result[idx] = coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); } const uint K_SUB = BK / TK; @@ -370,16 +374,20 @@ void main() { } // Pre-load scales into registers - ACC_TYPE scale_a[cms_per_row * CM_ELEMS]; - ACC_TYPE scale_b[cms_per_col * CM_ELEMS]; + float scale_a[cms_per_row * CM_ELEMS]; + float nbias_a[cms_per_row * CM_ELEMS]; + 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] = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]]); + scale_a[r * CM_ELEMS + e] = buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]]; + if (USE_MAGIC_BIAS) { + nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + 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]]); + scale_b[c * CM_ELEMS + e] = buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]]; } } @@ -388,8 +396,16 @@ void main() { [[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++) { - 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]; + if (USE_MAGIC_BIAS) { + const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), + scale_a[cm_row * CM_ELEMS + e], + nbias_a[cm_row * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[cm_col * CM_ELEMS + e], + float(sums[tile_idx * CM_ELEMS + e]))); + } else { + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) + * scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e]); + } } } } From 31b9bc683eb710df6d57c97bc5353037cabb687d Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:46:57 +0200 Subject: [PATCH 13/50] revert load reordering and scale pre-loading --- .../vulkan-shaders/mul_mmq_cm1.comp | 56 +++++-------------- 1 file changed, 14 insertions(+), 42 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 7185a3a68690..44fadf33de6c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -350,61 +350,33 @@ void main() { USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); } - const uint K_SUB = BK / TK; - coopmat all_a[cms_per_row * K_SUB]; - coopmat all_b[cms_per_col * K_SUB]; + coopmat cache_a; + coopmat cache_b; - [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(all_a[cm_row * K_SUB + h], buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); - } - } - [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(all_b[cm_col * K_SUB + h], buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); - } - } + [[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_row = 0; cm_row < cms_per_row; cm_row++) { - [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(all_a[cm_row * K_SUB + h], all_b[cm_col * K_SUB + h], cm_result[cm_col * cms_per_row + cm_row]); - } - } - } + [[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); - // Pre-load scales into registers - float scale_a[cms_per_row * CM_ELEMS]; - float nbias_a[cms_per_row * CM_ELEMS]; - 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]]; - if (USE_MAGIC_BIAS) { - nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + e]; + cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]); } } } - [[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]]; - } - } - // Apply scales from registers + // 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 float da = buf_a_d[ks * BM + warp_r * WM + cm_row * TM + elem_row[e]]; + const float db = buf_b_d[ks * BN + warp_c * WN + cm_col * TN + elem_col[e]]; if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), - scale_a[cm_row * CM_ELEMS + e], - nbias_a[cm_row * CM_ELEMS + e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[cm_col * CM_ELEMS + e], - float(sums[tile_idx * CM_ELEMS + e]))); + const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), da, -ACC_BIAS_F * da); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, db, float(sums[tile_idx * CM_ELEMS + e]))); } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) - * scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) * da * db); } } } From 23356762ca5de775e3aa853864d17aceb0e45f7a Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 10:54:21 +0200 Subject: [PATCH 14/50] workgroup scheduling for cache proximity --- .../vulkan-shaders/mul_mmq_cm1.comp | 24 +++++++++++++++---- 1 file changed, 19 insertions(+), 5 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 44fadf33de6c..3d3b30020811 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -81,6 +81,7 @@ layout (constant_id = 10) const uint WARP = 32; #define BK 32 #define BK_STEP 2 +#define GROUP_A_BUDGET (16u * 1024u * 1024u) const uint QPITCH = BK_STEP * (BK / 4) + 4; @@ -106,14 +107,31 @@ shared int32_t cm_layout_probe[TM * TN]; #include "mul_mmq_cm1_funcs.glsl" void main() { - const uint ic = gl_WorkGroupID.y; + const uint blocks_m = (p.M + BM - 1) / BM; + const uint ik = gl_WorkGroupID.x / blocks_m; #ifdef MUL_MAT_ID + const uint ic = gl_WorkGroupID.y; + const uint ir = gl_WorkGroupID.x % blocks_m; const uint expert_idx = gl_WorkGroupID.z; if (ic * BN >= data_expert_count[expert_idx]) { return; } +#else + // L2-friendly workgroup scheduling + const uint blocks_n = (p.N + BN - 1) / BN; + const uint a_panel_bytes = BM * p.K + (BM * p.K) / 16; + const uint group_m = clamp(GROUP_A_BUDGET / max(a_panel_bytes, 1u), 1u, min(blocks_m, 32u)); + const uint tiles_per_group = group_m * blocks_n; + const uint lin = gl_WorkGroupID.y * blocks_m + (gl_WorkGroupID.x % blocks_m); + const uint group_id = lin / tiles_per_group; + const uint first_m = group_id * group_m; + const uint gsize = min(blocks_m - first_m, group_m); + const uint in_group = lin - group_id * tiles_per_group; + const uint ir = first_m + in_group % gsize; + const uint ic = in_group / gsize; #endif + #ifdef NEEDS_INIT_IQ_SHMEM init_iq_shmem(gl_WorkGroupSize); #endif @@ -130,10 +148,6 @@ void main() { const uint batch_idx_a = i03 * p.ne02 + i02; #endif - const uint blocks_m = (p.M + BM - 1) / BM; - const uint ir = gl_WorkGroupID.x % blocks_m; - const uint ik = gl_WorkGroupID.x / blocks_m; - const uint WNITER = (WM * WN) / (WARP * TM * TN * WMITER); const uint WSUBM = WM / WMITER; const uint WSUBN = WN / WNITER; From 9f069f42e82f11c8e59b5f52f2c064f4285a2f8b Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 11:24:02 +0200 Subject: [PATCH 15/50] clean up --- .../ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 17 +++-------------- 1 file changed, 3 insertions(+), 14 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 3d3b30020811..64951ae8a90b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -31,7 +31,6 @@ layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32 #endif layout (binding = 1) readonly buffer B {block_q8_1_x4_packed128 data_b[];}; layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; -layout (binding = 2) writeonly buffer D4 {D_TYPE_VEC4 data_dv4[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; @@ -73,7 +72,6 @@ layout (constant_id = 2) const uint BN = 128; // layout (constant_id = 3) const uint BK = 32; layout (constant_id = 4) const uint WM = 64; layout (constant_id = 5) const uint WN = 32; -layout (constant_id = 6) const uint WMITER = 2; layout (constant_id = 7) const uint TM = 16; layout (constant_id = 8) const uint TN = 16; layout (constant_id = 9) const uint TK = 16; @@ -95,7 +93,6 @@ shared float buf_b_d[BN * BK_STEP]; #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 -#define NUM_WARPS (BLOCK_SIZE / WARP) const uint CM_ELEMS = (TM * TN) / WARP; #define ACC_BIAS_BITS 0x4B400000 #define ACC_BIAS_F 12582912.0f @@ -103,8 +100,10 @@ const bool USE_MAGIC_BIAS = WARP != 32; shared int32_t cm_layout_probe[TM * TN]; +#ifdef MUL_MAT_ID +#define NUM_WARPS (BLOCK_SIZE / WARP) #include "mul_mm_id_funcs.glsl" -#include "mul_mmq_cm1_funcs.glsl" +#endif void main() { const uint blocks_m = (p.M + BM - 1) / BM; @@ -132,10 +131,6 @@ void main() { const uint ic = in_group / gsize; #endif -#ifdef NEEDS_INIT_IQ_SHMEM - init_iq_shmem(gl_WorkGroupSize); -#endif - #ifndef MUL_MAT_ID const uint batch_idx = gl_WorkGroupID.z + p.base_work_group_z; @@ -148,14 +143,8 @@ void main() { const uint batch_idx_a = i03 * p.ne02 + i02; #endif - const uint WNITER = (WM * WN) / (WARP * TM * TN * WMITER); - const uint WSUBM = WM / WMITER; - const uint WSUBN = WN / WNITER; - const uint warp_i = gl_SubgroupID; - const uint tiw = gl_SubgroupInvocationID; - const uint cms_per_row = WM / TM; const uint cms_per_col = WN / TN; From 15b3ce0de34cc6e35cb3f2bf37e6b59438321aa1 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 11:24:38 +0200 Subject: [PATCH 16/50] Revert "revert load reordering and scale pre-loading" This reverts commit fbaefe0eaaccbb6ab9d733ad84580428f18db410. --- .../vulkan-shaders/mul_mmq_cm1.comp | 56 ++++++++++++++----- 1 file changed, 42 insertions(+), 14 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 64951ae8a90b..4f0a364248e7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -353,33 +353,61 @@ void main() { USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); } - coopmat cache_a; - coopmat cache_b; + const uint K_SUB = BK / TK; + coopmat all_a[cms_per_row * K_SUB]; + coopmat all_b[cms_per_col * K_SUB]; - [[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_row = 0; cm_row < cms_per_row; cm_row++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(all_a[cm_row * K_SUB + h], buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + } + } + [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(all_b[cm_col * K_SUB + h], buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + } + } - [[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); + [[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++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(all_a[cm_row * K_SUB + h], all_b[cm_col * K_SUB + h], cm_result[cm_col * cms_per_row + cm_row]); + } + } + } - cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]); + // Pre-load scales into registers + float scale_a[cms_per_row * CM_ELEMS]; + float nbias_a[cms_per_row * CM_ELEMS]; + 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]]; + if (USE_MAGIC_BIAS) { + nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + 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] = buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]]; + } + } - // Apply scales directly from coopmat elements + // 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 float da = buf_a_d[ks * BM + warp_r * WM + cm_row * TM + elem_row[e]]; - const float db = buf_b_d[ks * BN + warp_c * WN + cm_col * TN + elem_col[e]]; if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), da, -ACC_BIAS_F * da); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, db, float(sums[tile_idx * CM_ELEMS + e]))); + const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), + scale_a[cm_row * CM_ELEMS + e], + nbias_a[cm_row * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[cm_col * CM_ELEMS + e], + float(sums[tile_idx * CM_ELEMS + e]))); } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) * da * db); + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) + * scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e]); } } } From 57b42767849084ea0ef7ad67fab9b749b7a6bcdb Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 12:24:21 +0200 Subject: [PATCH 17/50] use wave32 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a18d759d6525..83047be06c17 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4335,9 +4335,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; // Coopmat int8 cm1 shader uses larger workgroups for better occupancy - l_warptile_mmq_cm1_int = { subgroup_size_8 * 8, 128, 128, 32, 64, 32, 2, itm_l, itn_l, itk_l, subgroup_size_8 }; - m_warptile_mmq_cm1_int = { subgroup_size_8 * 4, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, subgroup_size_8 }; - s_warptile_mmq_cm1_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; + // Force wave32 for VOPD on RDNA3+ + l_warptile_mmq_cm1_int = { 512, 128, 128, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; + m_warptile_mmq_cm1_int = { 128, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; + s_warptile_mmq_cm1_int = { 32, 32, 32, 32, 32, 32, 2, itm_s, itn_s, itk_s, 32 }; // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; @@ -4801,11 +4802,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, 32); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, 32); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, 32); \ // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ From a9e0e89c0b94b4fe716249a3968de384f6c0ff97 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 12:41:33 +0200 Subject: [PATCH 18/50] increase large tile size --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 83047be06c17..12153af8e1b0 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4248,6 +4248,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { std::array l_wg_denoms, m_wg_denoms, s_wg_denoms, l_mmq_wg_denoms, m_mmq_wg_denoms, s_mmq_wg_denoms, l_mmq_wg_denoms_k, m_mmq_wg_denoms_k, s_mmq_wg_denoms_k, + l_mmq_cm1_wg_denoms, m_mmq_cm1_wg_denoms, s_mmq_cm1_wg_denoms, l_mmqid_wg_denoms, m_mmqid_wg_denoms, s_mmqid_wg_denoms; uint32_t l_align, m_align, s_align; @@ -4336,8 +4337,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // Coopmat int8 cm1 shader uses larger workgroups for better occupancy // Force wave32 for VOPD on RDNA3+ - l_warptile_mmq_cm1_int = { 512, 128, 128, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; - m_warptile_mmq_cm1_int = { 128, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; + l_warptile_mmq_cm1_int = { 640, 128, 160, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; + m_warptile_mmq_cm1_int = { 256, 128, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; s_warptile_mmq_cm1_int = { 32, 32, 32, 32, 32, 32, 2, itm_s, itn_s, itk_s, 32 }; // K-quants use even more registers, mitigate by setting WMITER to 1 @@ -4379,6 +4380,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_mmq_wg_denoms = l_wg_denoms = {128, 128, 1 }; m_mmq_wg_denoms = m_wg_denoms = { 64, 64, 1 }; s_mmq_wg_denoms = s_wg_denoms = { 32, 32, 1 }; + l_mmq_cm1_wg_denoms = {128, 160, 1 }; + m_mmq_cm1_wg_denoms = {128, 64, 1 }; + s_mmq_cm1_wg_denoms = { 32, 32, 1 }; l_align = 128; m_align = 64; s_align = 32; @@ -4866,8 +4870,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } if (device->coopmat_int_support) { - CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_cm1_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_cm1_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } GGML_ASSERT(device->subgroup_ballot); From ef5c62992fba765c3017216bacd63ea2866adda3 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 13:04:28 +0200 Subject: [PATCH 19/50] Revert "increase large tile size" This reverts commit 7fabc25c5e0c9047a7a25b23bba5cad017de23ec. --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 12153af8e1b0..83047be06c17 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4248,7 +4248,6 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { std::array l_wg_denoms, m_wg_denoms, s_wg_denoms, l_mmq_wg_denoms, m_mmq_wg_denoms, s_mmq_wg_denoms, l_mmq_wg_denoms_k, m_mmq_wg_denoms_k, s_mmq_wg_denoms_k, - l_mmq_cm1_wg_denoms, m_mmq_cm1_wg_denoms, s_mmq_cm1_wg_denoms, l_mmqid_wg_denoms, m_mmqid_wg_denoms, s_mmqid_wg_denoms; uint32_t l_align, m_align, s_align; @@ -4337,8 +4336,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // Coopmat int8 cm1 shader uses larger workgroups for better occupancy // Force wave32 for VOPD on RDNA3+ - l_warptile_mmq_cm1_int = { 640, 128, 160, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; - m_warptile_mmq_cm1_int = { 256, 128, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; + l_warptile_mmq_cm1_int = { 512, 128, 128, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; + m_warptile_mmq_cm1_int = { 128, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; s_warptile_mmq_cm1_int = { 32, 32, 32, 32, 32, 32, 2, itm_s, itn_s, itk_s, 32 }; // K-quants use even more registers, mitigate by setting WMITER to 1 @@ -4380,9 +4379,6 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_mmq_wg_denoms = l_wg_denoms = {128, 128, 1 }; m_mmq_wg_denoms = m_wg_denoms = { 64, 64, 1 }; s_mmq_wg_denoms = s_wg_denoms = { 32, 32, 1 }; - l_mmq_cm1_wg_denoms = {128, 160, 1 }; - m_mmq_cm1_wg_denoms = {128, 64, 1 }; - s_mmq_cm1_wg_denoms = { 32, 32, 1 }; l_align = 128; m_align = 64; s_align = 32; @@ -4870,8 +4866,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } if (device->coopmat_int_support) { - CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_cm1_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_cm1_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } GGML_ASSERT(device->subgroup_ballot); From 2b9c53381e232ca9dc809dd1b226bf8192f3e519 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 13:07:06 +0200 Subject: [PATCH 20/50] restructure for vgpr use --- .../vulkan-shaders/mul_mmq_cm1.comp | 110 +++++++++--------- 1 file changed, 57 insertions(+), 53 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 4f0a364248e7..db977845995f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -155,21 +155,26 @@ void main() { uint elem_row[CM_ELEMS]; uint elem_col[CM_ELEMS]; - for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { - cm_layout_probe[i] = int32_t(i); - } - barrier(); + if (WARP == 32) { + [[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) { + elem_row[e] = gl_SubgroupInvocationID / TN + 2 * e; + elem_col[e] = gl_SubgroupInvocationID % TN; + } + } else { + for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { + cm_layout_probe[i] = int32_t(i); + } + barrier(); - coopmat probe; - coopMatLoad(probe, cm_layout_probe, 0, TN, gl_CooperativeMatrixLayoutRowMajor); + coopmat probe; + 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; + [[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) { + elem_row[e] = uint(probe[e]) / TN; + elem_col[e] = uint(probe[e]) % TN; + } } - barrier(); - const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A); const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B); @@ -226,8 +231,6 @@ void main() { uint pos_b_ib = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / BK; #endif - coopmat cm_result[cms_per_row * cms_per_col]; - ACC_TYPE sums[cms_per_row * cms_per_col * CM_ELEMS]; [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col * CM_ELEMS; i++) { sums[i] = ACC_TYPE(0.0); @@ -348,35 +351,21 @@ void main() { // Compute from shmem [[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( - USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); - } - const uint K_SUB = BK / TK; - coopmat all_a[cms_per_row * K_SUB]; - coopmat all_b[cms_per_col * K_SUB]; + coopmat cache_a[cms_per_row * K_SUB]; + coopmat cache_b[cms_per_col * K_SUB]; - [[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) { + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(all_a[cm_row * K_SUB + h], buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + 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); } } - [[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) { + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(all_b[cm_col * K_SUB + h], buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + 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); } } - [[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++) { - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(all_a[cm_row * K_SUB + h], all_b[cm_col * K_SUB + h], cm_result[cm_col * cms_per_row + cm_row]); - } - } - } - - // Pre-load scales into registers float scale_a[cms_per_row * CM_ELEMS]; float nbias_a[cms_per_row * CM_ELEMS]; float scale_b[cms_per_col * CM_ELEMS]; @@ -394,20 +383,35 @@ void main() { } } - // 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; + coopmat accs[cms_per_row * cms_per_col]; + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); + } + + accs[r * cms_per_col + c] = acc; + } + } + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[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++) { if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(cm_result[tile_idx][e])), - scale_a[cm_row * CM_ELEMS + e], - nbias_a[cm_row * CM_ELEMS + e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[cm_col * CM_ELEMS + e], + const float t = fma(intBitsToFloat(int(accs[tile_idx][e])), + scale_a[r * CM_ELEMS + e], + nbias_a[r * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c * CM_ELEMS + e], float(sums[tile_idx * CM_ELEMS + e]))); } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(cm_result[tile_idx][e]) - * scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(accs[tile_idx][e]) + * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); } } } @@ -426,14 +430,14 @@ void main() { const uint dc = ic * BN + warp_c * WN; #ifdef MUL_MAT_ID - [[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 r = 0; r < cms_per_row; r++) { + [[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 + cm_col * TN + elem_col[e]; + const uint col_i = dc + c * TN + elem_col[e]; if (col_i >= _ne1) continue; - const uint row_g = dr + cm_row * TM + elem_row[e]; + const uint row_g = dr + r * TM + elem_row[e]; if (row_g >= p.M) continue; const u16vec2 row_idx = row_ids[col_i - ic * BN]; @@ -445,12 +449,12 @@ void main() { #else const uint offsets = batch_idx * p.batch_stride_d + ik * p.batch_stride_d * p.num_batches; - [[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 r = 0; r < cms_per_row; r++) { + [[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 + cm_row * TM + elem_row[e]; - const uint col_g = dc + cm_col * TN + elem_col[e]; + const uint row_g = dr + r * TM + elem_row[e]; + const uint col_g = dc + c * TN + 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]); } From 57dd3f8bd3e1e84c5cc59c7ee5e35c16a0ada065 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 13:48:58 +0200 Subject: [PATCH 21/50] 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 db977845995f..007ea3f51f40 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++) { From ab6e84bbddeacac7395275a9cc1d2d0094351ad2 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 14:36:41 +0200 Subject: [PATCH 22/50] only force subgroup size 32 on AMD RDNA --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 83047be06c17..e5b7278b5678 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4334,11 +4334,18 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, itm_m, itn_m, itk_m, mm_warp_8 }; s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; - // Coopmat int8 cm1 shader uses larger workgroups for better occupancy - // Force wave32 for VOPD on RDNA3+ - l_warptile_mmq_cm1_int = { 512, 128, 128, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; - m_warptile_mmq_cm1_int = { 128, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; - s_warptile_mmq_cm1_int = { 32, 32, 32, 32, 32, 32, 2, itm_s, itn_s, itk_s, 32 }; + // RDNA3.5 preferred wave32 here + const bool cm1_use_wave32 = device->vendor_id == VK_VENDOR_ID_AMD && + device->subgroup_size_control && + device->subgroup_min_size <= 32 && device->subgroup_max_size >= 32; + const uint32_t cm1_sg = cm1_use_wave32 ? 32 : device->subgroup_size; + const auto cm1_bs = [cm1_sg](uint32_t bm, uint32_t bn) { + return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32); + }; + + l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm_l, itn_l, itk_l, cm1_sg }; + m_warptile_mmq_cm1_int = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; + s_warptile_mmq_cm1_int = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; @@ -4802,11 +4809,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, cm1_sg); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, cm1_sg); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, cm1_sg); \ // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ From 42a6269bdaac91f03e19c7b6c8950fbe2d456955 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 14:37:05 +0200 Subject: [PATCH 23/50] use BK_STEP 4 --- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 007ea3f51f40..bdb50d5ff3fa 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -78,7 +78,7 @@ layout (constant_id = 9) const uint TK = 16; layout (constant_id = 10) const uint WARP = 32; #define BK 32 -#define BK_STEP 2 +#define BK_STEP 4 #define GROUP_A_BUDGET (16u * 1024u * 1024u) const uint QPITCH = BK_STEP * (BK / 4) + 4; From 1825270eb46b1be4909d48f7f149653b2bfc4529 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 14:50:46 +0200 Subject: [PATCH 24/50] fix compilation --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index e5b7278b5678..75bb4426a4af 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4252,6 +4252,12 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { uint32_t l_align, m_align, s_align; + // RDNA3.5 preferred wave32 here + const bool cm1_use_wave32 = device->vendor_id == VK_VENDOR_ID_AMD && + device->subgroup_size_control && + device->subgroup_min_size <= 32 && device->subgroup_max_size >= 32; + const uint32_t cm1_sg = cm1_use_wave32 ? 32 : device->subgroup_size; + vk_pipeline wait_pipeline; CompileTask claimed_task {}; bool has_claimed_task = false; @@ -4334,11 +4340,6 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, itm_m, itn_m, itk_m, mm_warp_8 }; s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; - // RDNA3.5 preferred wave32 here - const bool cm1_use_wave32 = device->vendor_id == VK_VENDOR_ID_AMD && - device->subgroup_size_control && - device->subgroup_min_size <= 32 && device->subgroup_max_size >= 32; - const uint32_t cm1_sg = cm1_use_wave32 ? 32 : device->subgroup_size; const auto cm1_bs = [cm1_sg](uint32_t bm, uint32_t bn) { return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32); }; From eb46b8cf0a2efb5059cd77ac9601bd5d57b4b1e3 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 15:12:28 +0200 Subject: [PATCH 25/50] move quant-specific prefetch function out of main file --- .../vulkan-shaders/mul_mmq_cm1.comp | 84 +-------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 162 ++++++++---------- 2 files changed, 77 insertions(+), 169 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 bdb50d5ff3fa..9a4b0cc5beab 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -245,89 +245,7 @@ void main() { ivec4 pre_b_qs[B_LOADS * BK_STEP]; float16_t pre_b_d [B_LOADS * BK_STEP]; -// Prefetch: global memory → registers -#define PREFETCH_BLOCK(blk) \ - [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ - const uint buf_ib = loadc_a + li * loadstride_a; \ - if (buf_ib < BM) { \ - const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \ - [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - pre_a_qs[li * BK_STEP + ks] = \ - pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ - data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ - pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ - } \ - } \ - } \ - [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ - const uint buf_ib = loadc_b + li * loadstride_b; \ - if (buf_ib < BN) { \ - B_IB_CALC \ - [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \ - const uint ib_outer = ib_k / 4; \ - const uint ib_inner = ib_k % 4; \ - pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ - pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ - } \ - } \ - } - -#ifdef MUL_MAT_ID -#define B_IB_CALC \ - 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; -#else -#define B_IB_CALC \ - const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; -#endif - -// Store: registers → shared memory (with quant-specific unpacking) -#define STORE_BLOCK_TO_LDS(blk) \ - [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ - const uint buf_ib = loadc_a + li * loadstride_a; \ - if (buf_ib < BM) { \ - [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - const uint idx = li * BK_STEP + ks; \ - STORE_A_QS(buf_ib, ks, idx) \ - if (loadr_a == 0) { \ - buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ - } \ - } \ - } \ - } \ - [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ - const uint buf_ib = loadc_b + li * loadstride_b; \ - if (buf_ib < BN) { \ - [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - const bool in_bounds = (blk) + ks * BK < end_k; \ - const uint idx = li * BK_STEP + ks; \ - const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \ - const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \ - buf_b_qs[base ] = v.x; \ - buf_b_qs[base + 1] = v.y; \ - buf_b_qs[base + 2] = v.z; \ - buf_b_qs[base + 3] = v.w; \ - if (loadr_b == 0) { \ - buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ - } \ - } \ - } \ - } - -#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1) -#define STORE_A_QS(buf_ib, ks, idx) \ - uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ - uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ - lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ - hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; -#elif defined(DATA_A_Q8_0) -#define STORE_A_QS(buf_ib, ks, idx) \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; -#endif +#include "mul_mmq_cm1_funcs.glsl" // Prefetch first block if (start_k < end_k) { 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 14cb6bfd87d1..fcb914a44351 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 @@ -1,94 +1,84 @@ -#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require -#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require -#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require - -#include "types.glsl" - -// Each iqs value maps to a 32-bit integer - +// Quant-specific A-side unpacking: registers → shared memory #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, 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])); -#else // DATA_A_Q4_1 - const uint32_t vui = data_a_packed32[ib].qs[iqs]; +#define STORE_A_QS(buf_ib, ks, idx) \ + uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ + uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ + lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ + hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; +#elif defined(DATA_A_Q8_0) +#define STORE_A_QS(buf_ib, ks, idx) \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; #endif - uint32_t lo4 = vui & 0x0F0F0F0F; - uint32_t hi4 = (vui >> 4) & 0x0F0F0F0F; - - // subtract 8 from each byte - lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; - hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; - - 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[ks * BM + buf_ib] = float(data_a_packed16[ib].d); -#else // DATA_A_Q4_1 -#endif - } -} -#endif - -#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) -// 2-byte loads for Q5_0 blocks (22 bytes) -// 4-byte loads for Q5_1 blocks (24 bytes) -} +#ifdef MUL_MAT_ID +#define B_IB_CALC \ + 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; +#else +#define B_IB_CALC \ + const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; #endif -#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, 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 * QPITCH + ks * (BK / 4) + iqs] = vui; - - if (iqs == 0) { - buf_a_d[ks * BM + buf_ib] = float(data_a_packed16[ib].d); +// Prefetch: global memory → registers +#define PREFETCH_BLOCK(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + pre_a_qs[li * BK_STEP + ks] = \ + pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ + data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ + pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + B_IB_CALC \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \ + const uint ib_outer = ib_k / 4; \ + const uint ib_inner = ib_k % 4; \ + pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ + pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ + } \ + } \ } -} -#endif -#if defined(DATA_A_MXFP4) -// 1-byte loads for mxfp4 blocks (17 bytes) -#endif - -// For k-quants, ib and iqs still assume 32-wide blocks, but k-quants are 256-wide -// iqs still refers to a 32-bit integer, meaning 0..7 for 32-wide quants -#if defined(DATA_A_Q2_K) -// 4-byte loads for Q2_K blocks (84 bytes) -#endif - -#if defined(DATA_A_Q3_K) -// 2-byte loads for Q3_K blocks (110 bytes) -#endif - -#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) -// 4-byte loads for Q4_K blocks (144 bytes) and Q5_K blocks (176 bytes) -#endif - -#if defined(DATA_A_Q6_K) -// 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, const uint ks) { - const uint ib_outer = ib / 4; - const uint ib_inner = ib % 4; - - if (iqs == 0) { - buf_b_d[ks * BN + buf_ib] = float(data_b[ib_outer].ds[ib_inner].x); +// Store: registers → shared memory (with quant-specific unpacking) +#define STORE_BLOCK_TO_LDS(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const uint idx = li * BK_STEP + ks; \ + STORE_A_QS(buf_ib, ks, idx) \ + if (loadr_a == 0) { \ + buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ + } \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const bool in_bounds = (blk) + ks * BK < end_k; \ + const uint idx = li * BK_STEP + ks; \ + const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \ + const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \ + buf_b_qs[base ] = v.x; \ + buf_b_qs[base + 1] = v.y; \ + buf_b_qs[base + 2] = v.z; \ + buf_b_qs[base + 3] = v.w; \ + if (loadr_b == 0) { \ + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ + } \ + } \ + } \ } - - const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs]; - 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; -} From d2a5ac69c9d67624d08aa1ae6200656885284e0d Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 15:29:39 +0200 Subject: [PATCH 26/50] add q4_1, q5_0, q5_1 support --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 3 ++ .../vulkan-shaders/mul_mmq_cm1.comp | 44 +++++++++++++++++ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 49 ++++++++++++++++++- .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 4 files changed, 95 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 75bb4426a4af..d127fd578fbf 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4875,6 +4875,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (device->coopmat_int_support) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } 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 9a4b0cc5beab..4673dbce13a8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -90,6 +90,17 @@ shared float buf_a_d[BM * BK_STEP]; 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) +shared float buf_a_m[BM * BK_STEP]; +shared float buf_b_s[BN * BK_STEP]; +#endif + +#if defined(DATA_A_Q4_1) +#define QUANT_OFFSET 8.0 +#elif defined(DATA_A_Q5_1) +#define QUANT_OFFSET 16.0 +#endif + #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -245,6 +256,14 @@ void main() { ivec4 pre_b_qs[B_LOADS * BK_STEP]; float16_t pre_b_d [B_LOADS * BK_STEP]; +#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) + uint32_t pre_a_qh[A_LOADS * BK_STEP]; +#endif +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + float16_t pre_a_m [A_LOADS * BK_STEP]; + float16_t pre_b_s [B_LOADS * BK_STEP]; +#endif + #include "mul_mmq_cm1_funcs.glsl" // Prefetch first block @@ -308,6 +327,22 @@ void main() { } } +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + float corr_a[cms_per_row * CM_ELEMS]; + float sum_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++) { + corr_a[r * CM_ELEMS + e] = float(QUANT_OFFSET) * scale_a[r * CM_ELEMS + e] + + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; + } + } + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + sum_b[c * CM_ELEMS + e] = buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]; + } + } +#endif + coopmat accs[cms_per_row * cms_per_col]; [[unroll]] for (uint r = 0; r < cms_per_row; r++) { @@ -338,6 +373,10 @@ void main() { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(accs[tile_idx][e]) * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); } +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( + corr_a[r * CM_ELEMS + e] * sum_b[c * CM_ELEMS + e]); +#endif } } } @@ -351,6 +390,11 @@ void main() { #undef STORE_BLOCK_TO_LDS #undef B_IB_CALC #undef PREFETCH_BLOCK +#undef PREFETCH_A_QH +#undef PREFETCH_A_M +#undef PREFETCH_B_S +#undef STORE_A_M +#undef STORE_B_S const uint dr = ir * BM + a_row0; const uint dc = ic * BN + b_col0; 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 fcb914a44351..453190fe470c 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 @@ -7,11 +7,51 @@ hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; +#elif defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) +#define STORE_A_QS(buf_ib, ks, idx) \ + uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ + uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ + const uint32_t qh = pre_a_qh[idx]; \ + lo4 |= ((qh >> (4u * loadr_a )) & 0xFu) * 0x02040810u & 0x10101010u; \ + hi4 |= ((qh >> (4u * loadr_a + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; \ + lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ + hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; #elif defined(DATA_A_Q8_0) #define STORE_A_QS(buf_ib, ks, idx) \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; #endif +// Quant-specific extra prefetch helpers (no-ops for types that don't need them) +#if defined(DATA_A_Q5_0) +#define PREFETCH_A_QH(li, ks, ib) \ + pre_a_qh[(li) * BK_STEP + (ks)] = \ + pack32(u16vec2(data_a_packed16[(ib) + (ks)].qh[0], \ + data_a_packed16[(ib) + (ks)].qh[1])); +#elif defined(DATA_A_Q5_1) +#define PREFETCH_A_QH(li, ks, ib) \ + pre_a_qh[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].qh; +#else +#define PREFETCH_A_QH(li, ks, ib) +#endif + +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#define PREFETCH_A_M(li, ks, ib) \ + pre_a_m[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].m; +#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ + pre_b_s[(li) * BK_STEP + (ks)] = data_b[(ib_outer)].ds[(ib_inner)].y; +#define STORE_A_M(buf_ib, ks, idx) \ + buf_a_m[(ks) * BM + (buf_ib)] = float(pre_a_m[idx]); +#define STORE_B_S(in_bounds, buf_ib, ks, idx) \ + buf_b_s[(ks) * BN + (buf_ib)] = (in_bounds) ? float(pre_b_s[idx]) : 0.0f; +#else +#define PREFETCH_A_M(li, ks, ib) +#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) +#define STORE_A_M(buf_ib, ks, idx) +#define STORE_B_S(in_bounds, buf_ib, ks, idx) +#endif + #ifdef MUL_MAT_ID #define B_IB_CALC \ const u16vec2 row_idx = row_ids[buf_ib]; \ @@ -33,6 +73,8 @@ pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ + PREFETCH_A_QH(li, ks, ib) \ + PREFETCH_A_M(li, ks, ib) \ } \ } \ } \ @@ -46,6 +88,7 @@ const uint ib_inner = ib_k % 4; \ pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ + PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ } \ } \ } @@ -59,7 +102,8 @@ const uint idx = li * BK_STEP + ks; \ STORE_A_QS(buf_ib, ks, idx) \ if (loadr_a == 0) { \ - buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ + buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ + STORE_A_M(buf_ib, ks, idx) \ } \ } \ } \ @@ -77,7 +121,8 @@ buf_b_qs[base + 2] = v.z; \ buf_b_qs[base + 3] = v.w; \ if (loadr_b == 0) { \ - buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ + STORE_B_S(in_bounds, buf_ib, ks, idx) \ } \ } \ } \ diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 5c863dee4dea..45512fc86f1a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -629,7 +629,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && (tname == "q4_0" || tname == "q8_0")) { + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0")) { string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); } } From ed294e8724a011e36a7c837c3db84d9c0230c835 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 15:49:48 +0200 Subject: [PATCH 27/50] restructure mmq cm1 functions --- .../vulkan-shaders/mul_mmq_cm1.comp | 26 +- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 266 +++++++++++++----- 2 files changed, 195 insertions(+), 97 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 4673dbce13a8..0e30f2b6a1ef 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -116,6 +116,8 @@ shared int32_t cm_layout_probe[TM * TN]; #include "mul_mm_id_funcs.glsl" #endif +#include "mul_mmq_cm1_funcs.glsl" + void main() { const uint blocks_m = (p.M + BM - 1) / BM; const uint ik = gl_WorkGroupID.x / blocks_m; @@ -251,20 +253,8 @@ void main() { const uint A_LOADS = (BM + loadstride_a - 1) / loadstride_a; const uint B_LOADS = (BN + loadstride_b - 1) / loadstride_b; - uint32_t pre_a_qs[A_LOADS * BK_STEP]; - float16_t pre_a_d [A_LOADS * BK_STEP]; - ivec4 pre_b_qs[B_LOADS * BK_STEP]; - float16_t pre_b_d [B_LOADS * BK_STEP]; - -#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) - uint32_t pre_a_qh[A_LOADS * BK_STEP]; -#endif -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) - float16_t pre_a_m [A_LOADS * BK_STEP]; - float16_t pre_b_s [B_LOADS * BK_STEP]; -#endif - -#include "mul_mmq_cm1_funcs.glsl" + block_a_prefetch pre_a[A_LOADS * BK_STEP]; + block_b_prefetch pre_b[B_LOADS * BK_STEP]; // Prefetch first block if (start_k < end_k) { @@ -386,15 +376,9 @@ void main() { barrier(); } -#undef STORE_A_QS +#undef PREFETCH_BLOCK #undef STORE_BLOCK_TO_LDS #undef B_IB_CALC -#undef PREFETCH_BLOCK -#undef PREFETCH_A_QH -#undef PREFETCH_A_M -#undef PREFETCH_B_S -#undef STORE_A_M -#undef STORE_B_S const uint dr = ir * BM + a_row0; const uint dc = ic * BN + b_col0; 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 453190fe470c..330c36f384a3 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 @@ -1,56 +1,196 @@ -// Quant-specific A-side unpacking: registers → shared memory -#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1) -#define STORE_A_QS(buf_ib, ks, idx) \ - uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ - uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ - lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ - hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; -#elif defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) -#define STORE_A_QS(buf_ib, ks, idx) \ - uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ - uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ - const uint32_t qh = pre_a_qh[idx]; \ - lo4 |= ((qh >> (4u * loadr_a )) & 0xFu) * 0x02040810u & 0x10101010u; \ - hi4 |= ((qh >> (4u * loadr_a + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; \ - lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ - hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; +// Per-quant-type data structures and functions for the cm1 int8 coopmat path. +// Each quant type defines: +// struct block_a_prefetch — register data for one A-block per thread +// block_a_load() — load from global memory into a block_a_prefetch +// block_a_to_shmem() — unpack and write to shared memory + +#if defined(DATA_A_Q4_0) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_Q4_1) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; + float16_t m; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + blk.m = data_a_packed16[ib].m; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + buf_a_m[ks * BM + buf_ib] = float(blk.m); + } +} + +#elif defined(DATA_A_Q5_0) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; + uint32_t qh; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + blk.qh = pack32(u16vec2(data_a_packed16[ib].qh[0], data_a_packed16[ib].qh[1])); + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 |= ((blk.qh >> (4u * loadr )) & 0xFu) * 0x02040810u & 0x10101010u; + hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; + lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_Q5_1) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; + float16_t m; + uint32_t qh; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + blk.m = data_a_packed16[ib].m; + blk.qh = data_a_packed16[ib].qh; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 |= ((blk.qh >> (4u * loadr )) & 0xFu) * 0x02040810u & 0x10101010u; + hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; + lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + buf_a_m[ks * BM + buf_ib] = float(blk.m); + } +} + #elif defined(DATA_A_Q8_0) -#define STORE_A_QS(buf_ib, ks, idx) \ - buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +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] = blk.qs; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + #endif -// Quant-specific extra prefetch helpers (no-ops for types that don't need them) -#if defined(DATA_A_Q5_0) -#define PREFETCH_A_QH(li, ks, ib) \ - pre_a_qh[(li) * BK_STEP + (ks)] = \ - pack32(u16vec2(data_a_packed16[(ib) + (ks)].qh[0], \ - data_a_packed16[(ib) + (ks)].qh[1])); -#elif defined(DATA_A_Q5_1) -#define PREFETCH_A_QH(li, ks, ib) \ - pre_a_qh[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].qh; -#else -#define PREFETCH_A_QH(li, ks, ib) +// ===== B-side: load and store ===== + +struct block_b_prefetch { + ivec4 qs; + float16_t d; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + float16_t s; #endif +}; +block_b_prefetch block_b_load(uint ib_outer, uint ib_inner, uint loadr) { + block_b_prefetch blk; + blk.qs = data_b[ib_outer].qs[ib_inner * 2 + loadr]; + blk.d = data_b[ib_outer].ds[ib_inner].x; #if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) -#define PREFETCH_A_M(li, ks, ib) \ - pre_a_m[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].m; -#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ - pre_b_s[(li) * BK_STEP + (ks)] = data_b[(ib_outer)].ds[(ib_inner)].y; -#define STORE_A_M(buf_ib, ks, idx) \ - buf_a_m[(ks) * BM + (buf_ib)] = float(pre_a_m[idx]); -#define STORE_B_S(in_bounds, buf_ib, ks, idx) \ - buf_b_s[(ks) * BN + (buf_ib)] = (in_bounds) ? float(pre_b_s[idx]) : 0.0f; -#else -#define PREFETCH_A_M(li, ks, ib) -#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) -#define STORE_A_M(buf_ib, ks, idx) -#define STORE_B_S(in_bounds, buf_ib, ks, idx) + blk.s = data_b[ib_outer].ds[ib_inner].y; #endif + return blk; +} + +void block_b_to_shmem(block_b_prefetch blk, uint buf_ib, uint ks, uint loadr, bool in_bounds) { + const ivec4 v = in_bounds ? blk.qs : ivec4(0); + const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr * 4; + buf_b_qs[base ] = v.x; + buf_b_qs[base + 1] = v.y; + buf_b_qs[base + 2] = v.z; + buf_b_qs[base + 3] = v.w; + if (loadr == 0) { + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(blk.d) : 0.0f; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + buf_b_s[ks * BN + buf_ib] = in_bounds ? float(blk.s) : 0.0f; +#endif + } +} + +// ===== Framework macros ===== #ifdef MUL_MAT_ID #define B_IB_CALC \ @@ -62,19 +202,13 @@ const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; #endif -// Prefetch: global memory → registers #define PREFETCH_BLOCK(blk) \ [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ const uint buf_ib = loadc_a + li * loadstride_a; \ if (buf_ib < BM) { \ const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \ [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - pre_a_qs[li * BK_STEP + ks] = \ - pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ - data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ - pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ - PREFETCH_A_QH(li, ks, ib) \ - PREFETCH_A_M(li, ks, ib) \ + pre_a[li * BK_STEP + ks] = block_a_load(ib + ks, loadr_a); \ } \ } \ } \ @@ -84,27 +218,17 @@ B_IB_CALC \ [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \ - const uint ib_outer = ib_k / 4; \ - const uint ib_inner = ib_k % 4; \ - pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ - pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ - PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ + pre_b[li * BK_STEP + ks] = block_b_load(ib_k / 4, ib_k % 4, loadr_b); \ } \ } \ } -// Store: registers → shared memory (with quant-specific unpacking) #define STORE_BLOCK_TO_LDS(blk) \ [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ const uint buf_ib = loadc_a + li * loadstride_a; \ if (buf_ib < BM) { \ [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ - const uint idx = li * BK_STEP + ks; \ - STORE_A_QS(buf_ib, ks, idx) \ - if (loadr_a == 0) { \ - buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ - STORE_A_M(buf_ib, ks, idx) \ - } \ + block_a_to_shmem(pre_a[li * BK_STEP + ks], buf_ib, ks, loadr_a); \ } \ } \ } \ @@ -113,17 +237,7 @@ if (buf_ib < BN) { \ [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ const bool in_bounds = (blk) + ks * BK < end_k; \ - const uint idx = li * BK_STEP + ks; \ - const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \ - const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \ - buf_b_qs[base ] = v.x; \ - buf_b_qs[base + 1] = v.y; \ - buf_b_qs[base + 2] = v.z; \ - buf_b_qs[base + 3] = v.w; \ - if (loadr_b == 0) { \ - buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ - STORE_B_S(in_bounds, buf_ib, ks, idx) \ - } \ + block_b_to_shmem(pre_b[li * BK_STEP + ks], buf_ib, ks, loadr_b, in_bounds); \ } \ } \ } From fa37db93ced375cf7cd2f7f3793ddefb9ac533c2 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 15:59:55 +0200 Subject: [PATCH 28/50] enable mul_mat_id support --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index d127fd578fbf..623a82fcbe64 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4879,6 +4879,12 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } GGML_ASSERT(device->subgroup_ballot); From d8ed37209ac53b201fc2e73fa9889fe6cdfff4fc Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 16:20:59 +0200 Subject: [PATCH 29/50] fix segfault --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 623a82fcbe64..85ef6e01f167 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -7883,10 +7883,19 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte assert(src1_type == GGML_TYPE_F16); return prec == GGML_PREC_DEFAULT ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f32acc; } + + vk_matmul_pipeline pipelines; if (ctx->device->coopmat_support) { - return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + pipelines = (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + } else { + pipelines = (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + } + + if (pipelines->is_empty()) { + return nullptr; } - return (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + + return pipelines; } static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context * ctx, ggml_type a_type, ggml_type b_type, uint32_t num_cols, uint32_t m, uint32_t k) { From 3cd76432d76d60c6aca4b75fb7d642a4f15ad26f Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 16:39:39 +0200 Subject: [PATCH 30/50] fix mul_mat_id bug --- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 4 ++++ 1 file changed, 4 insertions(+) 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 0e30f2b6a1ef..94b6325cc633 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -263,7 +263,11 @@ void main() { const uint a_row0 = warp_r * WM; const uint b_col0 = warp_c * WN; +#ifdef MUL_MAT_ID + const bool active_col_tile = ic * BN + b_col0 < _ne1; +#else const bool active_col_tile = ic * BN + b_col0 < p.N; +#endif barrier(); From dc08fa91015cf28995ca89a7328410fe0c69b00c Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 17:10:17 +0200 Subject: [PATCH 31/50] support iq4_nl and mxfp4 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 26 ++++---- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 62 +++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 3 files changed, 78 insertions(+), 12 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 85ef6e01f167..29adfcc3d94e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4874,17 +4874,21 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } if (device->coopmat_int_support) { - CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - - CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } GGML_ASSERT(device->subgroup_ballot); 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 330c36f384a3..3fe104cf98a6 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 @@ -153,6 +153,68 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { } } +#elif defined(DATA_A_IQ4_NL) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = + pack32(i8vec4(kvalues_iq4nl_const[lo_idx.x], kvalues_iq4nl_const[lo_idx.y], + kvalues_iq4nl_const[lo_idx.z], kvalues_iq4nl_const[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = + pack32(i8vec4(kvalues_iq4nl_const[hi_idx.x], kvalues_iq4nl_const[hi_idx.y], + kvalues_iq4nl_const[hi_idx.z], kvalues_iq4nl_const[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_MXFP4) + +struct block_a_prefetch { + uint32_t qs; + uint8_t e; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u8vec4(data_a[ib].qs[loadr * 4], + data_a[ib].qs[loadr * 4 + 1], + data_a[ib].qs[loadr * 4 + 2], + data_a[ib].qs[loadr * 4 + 3])); + blk.e = data_a[ib].e; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = + pack32(i8vec4(kvalues_mxfp4_const[lo_idx.x], kvalues_mxfp4_const[lo_idx.y], + kvalues_mxfp4_const[lo_idx.z], kvalues_mxfp4_const[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = + pack32(i8vec4(kvalues_mxfp4_const[hi_idx.x], kvalues_mxfp4_const[hi_idx.y], + kvalues_mxfp4_const[hi_idx.z], kvalues_mxfp4_const[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5; + } +} + #endif // ===== B-side: load and store ===== diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 45512fc86f1a..d1c778242339 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -629,7 +629,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0")) { + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4")) { string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); } } From ba8402e9f6a48da89c478b48aa3fc30450b9d71c Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 17:23:15 +0200 Subject: [PATCH 32/50] remove elem row/col fast path, invalid for RDNA4 --- .../vulkan-shaders/mul_mmq_cm1.comp | 25 +++++++------------ 1 file changed, 9 insertions(+), 16 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 94b6325cc633..7e79b0d80fa1 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -168,24 +168,17 @@ void main() { uint elem_row[CM_ELEMS]; uint elem_col[CM_ELEMS]; - if (WARP == 32) { - [[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) { - elem_row[e] = gl_SubgroupInvocationID / TN + 2 * e; - elem_col[e] = gl_SubgroupInvocationID % TN; - } - } else { - for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { - cm_layout_probe[i] = int32_t(i); - } - barrier(); + for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { + cm_layout_probe[i] = int32_t(i); + } + barrier(); - coopmat probe; - coopMatLoad(probe, cm_layout_probe, 0, TN, gl_CooperativeMatrixLayoutRowMajor); + coopmat probe; + 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; - } + [[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) { + elem_row[e] = uint(probe[e]) / TN; + elem_col[e] = uint(probe[e]) % TN; } const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); From 5037f027aa3928e7be9914ba052cb904ca3ce1c6 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 17:37:24 +0200 Subject: [PATCH 33/50] use shmem arrays for LUTs --- .../ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 16 ++++++++++++++++ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 16 ++++++++-------- 2 files changed, 24 insertions(+), 8 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 7e79b0d80fa1..e9ffd608e598 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -101,6 +101,10 @@ shared float buf_b_s[BN * BK_STEP]; #define QUANT_OFFSET 16.0 #endif +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) +shared int8_t cm1_kvalues[16]; +#endif + #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -119,6 +123,18 @@ shared int32_t cm_layout_probe[TM * TN]; #include "mul_mmq_cm1_funcs.glsl" void main() { +#if defined(DATA_A_IQ4_NL) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex]; + } + barrier(); +#elif defined(DATA_A_MXFP4) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; + } + barrier(); +#endif + const uint blocks_m = (p.M + BM - 1) / BM; const uint ik = gl_WorkGroupID.x / blocks_m; 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 3fe104cf98a6..294dd75a2a67 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 @@ -172,11 +172,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = - pack32(i8vec4(kvalues_iq4nl_const[lo_idx.x], kvalues_iq4nl_const[lo_idx.y], - kvalues_iq4nl_const[lo_idx.z], kvalues_iq4nl_const[lo_idx.w])); + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = - pack32(i8vec4(kvalues_iq4nl_const[hi_idx.x], kvalues_iq4nl_const[hi_idx.y], - kvalues_iq4nl_const[hi_idx.z], kvalues_iq4nl_const[hi_idx.w])); + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); if (loadr == 0) { buf_a_d[ks * BM + buf_ib] = float(blk.d); @@ -204,11 +204,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = - pack32(i8vec4(kvalues_mxfp4_const[lo_idx.x], kvalues_mxfp4_const[lo_idx.y], - kvalues_mxfp4_const[lo_idx.z], kvalues_mxfp4_const[lo_idx.w])); + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = - pack32(i8vec4(kvalues_mxfp4_const[hi_idx.x], kvalues_mxfp4_const[hi_idx.y], - kvalues_mxfp4_const[hi_idx.z], kvalues_mxfp4_const[hi_idx.w])); + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); if (loadr == 0) { buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5; From ab94ffdd3054f80c8bb0299b40c2470b0e06dcc1 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 17:44:44 +0200 Subject: [PATCH 34/50] use 4-byte loads where possible --- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 28 ++++++++----------- 1 file changed, 11 insertions(+), 17 deletions(-) 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 294dd75a2a67..7b2d8a188088 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 @@ -36,16 +36,13 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { struct block_a_prefetch { uint32_t qs; - float16_t d; - float16_t m; + f16vec2 dm; }; block_a_prefetch block_a_load(uint ib, uint loadr) { block_a_prefetch blk; - blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], - data_a_packed16[ib].qs[loadr * 2 + 1])); - blk.d = data_a_packed16[ib].d; - blk.m = data_a_packed16[ib].m; + blk.qs = data_a_packed32[ib].qs[loadr]; + blk.dm = data_a_packed32[ib].dm; return blk; } @@ -58,8 +55,8 @@ 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.d); - buf_a_m[ks * BM + buf_ib] = float(blk.m); + buf_a_d[ks * BM + buf_ib] = float(blk.dm.x); + buf_a_m[ks * BM + buf_ib] = float(blk.dm.y); } } @@ -99,18 +96,15 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { struct block_a_prefetch { uint32_t qs; - float16_t d; - float16_t m; + f16vec2 dm; uint32_t qh; }; block_a_prefetch block_a_load(uint ib, uint loadr) { block_a_prefetch blk; - blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], - data_a_packed16[ib].qs[loadr * 2 + 1])); - blk.d = data_a_packed16[ib].d; - blk.m = data_a_packed16[ib].m; - blk.qh = data_a_packed16[ib].qh; + blk.qs = data_a_packed32[ib].qs[loadr]; + blk.dm = data_a_packed32[ib].dm; + blk.qh = data_a_packed32[ib].qh; return blk; } @@ -125,8 +119,8 @@ 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.d); - buf_a_m[ks * BM + buf_ib] = float(blk.m); + buf_a_d[ks * BM + buf_ib] = float(blk.dm.x); + buf_a_m[ks * BM + buf_ib] = float(blk.dm.y); } } From 4695f108bb89d5d9199b26dd3870d62518259d3d Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 11:46:11 +0200 Subject: [PATCH 35/50] add q3_k, q4_k, q5_k, q6_k and nvfp4 support --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 15 + .../vulkan-shaders/mul_mmq_cm1.comp | 85 +++++- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 275 +++++++++++++++++- .../vulkan-shaders/vulkan-shaders-gen.cpp | 3 +- 4 files changed, 361 insertions(+), 17 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 29adfcc3d94e..cc95a6c85553 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4240,6 +4240,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_warptile_mmq, m_warptile_mmq, s_warptile_mmq, l_warptile_mmq_int, m_warptile_mmq_int, s_warptile_mmq_int, l_warptile_mmq_cm1_int, m_warptile_mmq_cm1_int, s_warptile_mmq_cm1_int, + l_warptile_mmq_cm1_int_k, m_warptile_mmq_cm1_int_k, s_warptile_mmq_cm1_int_k, l_warptile_mmq_int_k, m_warptile_mmq_int_k, s_warptile_mmq_int_k, l_warptile_mmq_k, m_warptile_mmq_k, s_warptile_mmq_k, l_warptile_mmqid, m_warptile_mmqid, s_warptile_mmqid, @@ -4348,6 +4349,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_cm1_int = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; s_warptile_mmq_cm1_int = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; + l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128), 64, 128, 32, std::min(cm1_sg, 64u), 32, 2, itm_l, itn_l, itk_l, cm1_sg }; + m_warptile_mmq_cm1_int_k = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; + s_warptile_mmq_cm1_int_k = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; + // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; m_warptile_mmq_int_k = { 128, 64, 64, 32, mm_warp_8, 32, 1, 2, 2, 1, mm_warp_8 }; @@ -4881,6 +4886,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4889,6 +4899,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } GGML_ASSERT(device->subgroup_ballot); 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 e9ffd608e598..22eff4093ce3 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -20,6 +20,12 @@ #include "types.glsl" +#if defined(DATA_A_Q3_K) || defined(DATA_A_Q6_K) || defined(DATA_A_NVFP4) +#define KSCALES 2 +#else +#define KSCALES 1 +#endif + layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; @@ -85,27 +91,31 @@ const uint QPITCH = BK_STEP * (BK / 4) + 4; // Shared memory cache shared uint32_t buf_a_qs[BM * QPITCH]; -shared float buf_a_d[BM * BK_STEP]; +shared float buf_a_d[BM * BK_STEP * KSCALES]; 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) +#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 -#if defined(DATA_A_Q4_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q4_K) #define QUANT_OFFSET 8.0 -#elif defined(DATA_A_Q5_1) +#elif defined(DATA_A_Q5_1) || defined(DATA_A_Q5_K) #define QUANT_OFFSET 16.0 #endif -#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) shared int8_t cm1_kvalues[16]; #endif +#if defined(DATA_A_QUANT_K) || defined(DATA_A_NVFP4) +#define LOAD_VEC_A 8 +#else #define LOAD_VEC_A (4 * QUANT_R) +#endif #define LOAD_VEC_B 16 const uint CM_ELEMS = (TM * TN) / WARP; @@ -133,6 +143,16 @@ void main() { cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; } barrier(); +#elif defined(DATA_A_NVFP4) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; + } +#if !defined(USE_OCP_FP4) + for (uint i = gl_LocalInvocationIndex; i < 128u; i += BLOCK_SIZE) { + ue4m3_fp32_lut[i] = ue4m3_to_fp32_build(i); + } +#endif + barrier(); #endif const uint blocks_m = (p.M + BM - 1) / BM; @@ -313,9 +333,52 @@ void main() { } } + float scale_b[cms_per_col * CM_ELEMS]; + [[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 + b_col0 + c * TN + elem_col[e]]; + } + } + +#if KSCALES == 2 + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + float scale_a[cms_per_row * CM_ELEMS]; + float nbias_a[cms_per_row * 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 * KSCALES + h) * 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]; + } + } + } + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); + + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + if (USE_MAGIC_BIAS) { + const float t = fma(intBitsToFloat(int(acc[e])), + scale_a[r * CM_ELEMS + e], + nbias_a[r * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c * CM_ELEMS + e], + float(sums[tile_idx * CM_ELEMS + e]))); + } else { + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) + * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); + } + } + } + } + } +#else float scale_a[cms_per_row * CM_ELEMS]; float nbias_a[cms_per_row * CM_ELEMS]; - 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 + a_row0 + r * TM + elem_row[e]]; @@ -324,13 +387,8 @@ 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 + b_col0 + c * TN + elem_col[e]]; - } - } -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) float corr_a[cms_per_row * CM_ELEMS]; float sum_b[cms_per_col * CM_ELEMS]; [[unroll]] for (uint r = 0; r < cms_per_row; r++) { @@ -376,13 +434,14 @@ void main() { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(accs[tile_idx][e]) * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); } -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( corr_a[r * CM_ELEMS + e] * sum_b[c * CM_ELEMS + e]); #endif } } } +#endif // KSCALES } } 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 7b2d8a188088..7536df4b0b2b 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 @@ -209,6 +209,275 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { } } +// LOAD_VEC_A=8 for k-quants and NVFP4: loadr has 4 positions, each writes 2 uint32 + +#elif defined(DATA_A_Q4_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint qs_base = (sub >> 1) * 8; + + uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2]; + uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1]; + if ((sub & 1u) != 0u) { + blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F; + blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F; + } else { + blk.qs0 = raw0 & 0x0F0F0F0F; + blk.qs1 = raw1 & 0x0F0F0F0F; + } + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x08080808) ^ 0x80808080; + uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x08080808) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + uint sc_val, mn_val; + if (sub < 4) { + sc_val = uint(data_a[ib_k].scales[sub]) & 0x3Fu; + mn_val = uint(data_a[ib_k].scales[sub + 4]) & 0x3Fu; + } else { + sc_val = (uint(data_a[ib_k].scales[sub + 4]) & 0xFu) | ((uint(data_a[ib_k].scales[sub - 4]) & 0xC0u) >> 2); + mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); + } + vec2 dm = vec2(data_a_packed32[ib_k].dm); + buf_a_d[ks * BM + buf_ib] = dm.x * float(sc_val); + buf_a_m[ks * BM + buf_ib] = -(dm.y * float(mn_val)); + } +} + +#elif defined(DATA_A_Q5_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint32_t qh0; + uint32_t qh1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint qs_base = (sub >> 1) * 8; + + uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2]; + uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1]; + if ((sub & 1u) != 0u) { + blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F; + blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F; + } else { + blk.qs0 = raw0 & 0x0F0F0F0F; + blk.qs1 = raw1 & 0x0F0F0F0F; + } + blk.qh0 = ((data_a_packed32[ib_k].qh[loadr * 2 ] >> sub) & 0x01010101) << 4; + blk.qh1 = ((data_a_packed32[ib_k].qh[loadr * 2 + 1] >> sub) & 0x01010101) << 4; + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = blk.qs0 | blk.qh0; + uint32_t v1 = blk.qs1 | blk.qh1; + v0 = ((v0 | 0x80808080) - 0x10101010) ^ 0x80808080; + v1 = ((v1 | 0x80808080) - 0x10101010) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + uint sc_val, mn_val; + if (sub < 4) { + sc_val = uint(data_a[ib_k].scales[sub]) & 0x3Fu; + mn_val = uint(data_a[ib_k].scales[sub + 4]) & 0x3Fu; + } else { + sc_val = (uint(data_a[ib_k].scales[sub + 4]) & 0xFu) | ((uint(data_a[ib_k].scales[sub - 4]) & 0xC0u) >> 2); + mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); + } + vec2 dm = vec2(data_a_packed32[ib_k].dm); + buf_a_d[ks * BM + buf_ib] = dm.x * float(sc_val); + buf_a_m[ks * BM + buf_ib] = -(dm.y * float(mn_val)); + } +} + +#elif defined(DATA_A_Q6_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint g = sub / 4; + const uint j = sub % 4; + + const uint ql_u16 = g * 32 + (j & 1) * 16 + loadr * 4; + const uint qh_u16 = g * 16 + loadr * 4; + const uint qh_shift = j * 2; + + uint32_t ql0 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16 ], + data_a_packed16[ib_k].ql[ql_u16 + 1])); + uint32_t ql1 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16 + 2], + data_a_packed16[ib_k].ql[ql_u16 + 3])); + if (j >= 2) { + ql0 = (ql0 >> 4) & 0x0F0F0F0F; + ql1 = (ql1 >> 4) & 0x0F0F0F0F; + } else { + ql0 = ql0 & 0x0F0F0F0F; + ql1 = ql1 & 0x0F0F0F0F; + } + + uint32_t qh0 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16 ], + data_a_packed16[ib_k].qh[qh_u16 + 1])); + uint32_t qh1 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16 + 2], + data_a_packed16[ib_k].qh[qh_u16 + 3])); + + blk.qs0 = ql0 | (((qh0 >> qh_shift) & 0x03030303) << 4); + blk.qs1 = ql1 | (((qh1 >> qh_shift) & 0x03030303) << 4); + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x20202020) ^ 0x80808080; + uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x20202020) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + i8vec2 sc = unpack8(int32_t(int16_t(data_a_packed16[ib_k].scales[sub]))).xy; + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.x); + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.y); + } +} + +#elif defined(DATA_A_Q3_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint g = sub / 4; + const uint j = sub % 4; + const uint qs_shift = j * 2; + const uint hm_bit = j + g * 4; + + const uint qs_u16 = g * 16 + loadr * 4; + uint32_t qs0 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16 ], + data_a_packed16[ib_k].qs[qs_u16 + 1])); + uint32_t qs1 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16 + 2], + data_a_packed16[ib_k].qs[qs_u16 + 3])); + + const uint hm_u16 = loadr * 4; + uint32_t hm0 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16 ], + data_a_packed16[ib_k].hmask[hm_u16 + 1])); + uint32_t hm1 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16 + 2], + data_a_packed16[ib_k].hmask[hm_u16 + 3])); + + blk.qs0 = ((qs0 >> qs_shift) & 0x03030303) | (((hm0 >> hm_bit) & 0x01010101) << 2); + blk.qs1 = ((qs1 >> qs_shift) & 0x03030303) | (((hm1 >> hm_bit) & 0x01010101) << 2); + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x04040404) ^ 0x80808080; + uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x04040404) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + const uint is = sub * 2; + uint lo = uint(data_a_packed16[ib_k].scales[(is % 8) / 2]); + lo = (lo >> (4 * (is / 8))) & 0x0F0Fu; + uint hi = uint(data_a_packed16[ib_k].scales[(8 + (is % 4)) / 2]); + hi = (hi >> (2 * (is / 4))) & 0x0303u; + uint combined = lo | (hi << 4); + i8vec2 sc = unpack8(int32_t(combined)).xy; + float d = float(data_a_packed16[ib_k].d); + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = d * float(int(sc.x) - 32); + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = d * float(int(sc.y) - 32); + } +} + +#elif defined(DATA_A_NVFP4) + +struct block_a_prefetch { + uint32_t qs; + uint8_t d0; + uint8_t d1; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 2; + const uint ihalf = ib % 2; + const uint sub = ihalf * 2 + (loadr >> 1); + const uint byte_group = loadr & 1u; + + blk.qs = pack32(u8vec4(data_a[ib_k].qs[sub * 8 + byte_group * 4], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 1], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 2], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 3])); + blk.d0 = data_a[ib_k].d[ihalf * 2]; + blk.d1 = data_a[ib_k].d[ihalf * 2 + 1]; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + const uint sub_base = (loadr >> 1) * 4; + const uint byte_group = loadr & 1u; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + byte_group] = + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + 2 + byte_group] = + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = ue4m3_to_fp32(blk.d0) * 0.5; + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = ue4m3_to_fp32(blk.d1) * 0.5; + } +} + #endif // ===== B-side: load and store ===== @@ -216,7 +485,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { struct block_b_prefetch { ivec4 qs; float16_t d; -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) float16_t s; #endif }; @@ -225,7 +494,7 @@ block_b_prefetch block_b_load(uint ib_outer, uint ib_inner, uint loadr) { block_b_prefetch blk; blk.qs = data_b[ib_outer].qs[ib_inner * 2 + loadr]; blk.d = data_b[ib_outer].ds[ib_inner].x; -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) blk.s = data_b[ib_outer].ds[ib_inner].y; #endif return blk; @@ -240,7 +509,7 @@ void block_b_to_shmem(block_b_prefetch blk, uint buf_ib, uint ks, uint loadr, bo buf_b_qs[base + 3] = v.w; if (loadr == 0) { buf_b_d[ks * BN + buf_ib] = in_bounds ? float(blk.d) : 0.0f; -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) buf_b_s[ks * BN + buf_ib] = in_bounds ? float(blk.s) : 0.0f; #endif } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index d1c778242339..27f56b3764c8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -629,7 +629,8 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4")) { + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4" + || tname == "q3_k" || tname == "q4_k" || tname == "q5_k" || tname == "q6_k" || tname == "nvfp4")) { string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); } } From 22f588c10d79c88364fdf25b728e5ecb341f7933 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 12:49:49 +0200 Subject: [PATCH 36/50] fix l warptile --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 25 +++++++++++++++---------- 1 file changed, 15 insertions(+), 10 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index cc95a6c85553..78f318247e2d 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4249,6 +4249,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { std::array l_wg_denoms, m_wg_denoms, s_wg_denoms, l_mmq_wg_denoms, m_mmq_wg_denoms, s_mmq_wg_denoms, l_mmq_wg_denoms_k, m_mmq_wg_denoms_k, s_mmq_wg_denoms_k, + l_mmq_cm1_wg_denoms_k, m_mmq_cm1_wg_denoms_k, s_mmq_cm1_wg_denoms_k, l_mmqid_wg_denoms, m_mmqid_wg_denoms, s_mmqid_wg_denoms; uint32_t l_align, m_align, s_align; @@ -4353,6 +4354,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_cm1_int_k = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; s_warptile_mmq_cm1_int_k = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; + l_mmq_cm1_wg_denoms_k = { l_warptile_mmq_cm1_int_k[1], l_warptile_mmq_cm1_int_k[2], 1 }; + m_mmq_cm1_wg_denoms_k = { m_warptile_mmq_cm1_int_k[1], m_warptile_mmq_cm1_int_k[2], 1 }; + s_mmq_cm1_wg_denoms_k = { s_warptile_mmq_cm1_int_k[1], s_warptile_mmq_cm1_int_k[2], 1 }; + // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; m_warptile_mmq_int_k = { 128, 64, 64, 32, mm_warp_8, 32, 1, 2, 2, 1, mm_warp_8 }; @@ -4886,11 +4891,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4899,11 +4904,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } GGML_ASSERT(device->subgroup_ballot); From 3985f6467019506a03bfecc2c254b00363a58db7 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 13:55:03 +0200 Subject: [PATCH 37/50] improve performance --- .../vulkan-shaders/mul_mmq_cm1.comp | 168 +++++++++++------- 1 file changed, 107 insertions(+), 61 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 22eff4093ce3..337003683ded 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -84,7 +84,11 @@ layout (constant_id = 9) const uint TK = 16; layout (constant_id = 10) const uint WARP = 32; #define BK 32 +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) +#define BK_STEP 2 +#else #define BK_STEP 4 +#endif #define GROUP_A_BUDGET (16u * 1024u * 1024u) const uint QPITCH = BK_STEP * (BK / 4) + 4; @@ -319,94 +323,146 @@ void main() { // Compute from shmem [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { const uint K_SUB = BK / TK; - coopmat cache_a[cms_per_row * K_SUB]; - coopmat cache_b[cms_per_col * K_SUB]; - - [[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, (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, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); - } - } - - float scale_b[cms_per_col * CM_ELEMS]; - [[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 + b_col0 + c * TN + elem_col[e]]; - } - } #if KSCALES == 2 [[unroll]] for (uint h = 0; h < K_SUB; h++) { - float scale_a[cms_per_row * CM_ELEMS]; - float nbias_a[cms_per_row * CM_ELEMS]; [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + coopmat cache_a; + coopMatLoad(cache_a, buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + + float scale_a[CM_ELEMS]; + float nbias_a[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_a[r * CM_ELEMS + 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 + elem_row[e]]; if (USE_MAGIC_BIAS) { - nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + e]; + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } } - } - [[unroll]] for (uint r = 0; r < cms_per_row; r++) { [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat cache_b; + coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + + float scale_b[CM_ELEMS]; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + scale_b[e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; + } + coopmat acc = coopmat( USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); - acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); + acc = coopMatMulAdd(cache_a, cache_b, acc); const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), - scale_a[r * CM_ELEMS + e], - nbias_a[r * CM_ELEMS + e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c * CM_ELEMS + e], + scale_a[e], + nbias_a[e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[e], float(sums[tile_idx * CM_ELEMS + e]))); } else { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); + * scale_a[e] * scale_b[e]); } } } } } -#else - float scale_a[cms_per_row * CM_ELEMS]; - float nbias_a[cms_per_row * CM_ELEMS]; +#elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + // Correction types: interleaved cache loading to reduce register pressure [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + coopmat cache_a[K_SUB]; + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_a[h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + } + + float scale_a[CM_ELEMS]; + float nbias_a[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_a[r * CM_ELEMS + 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 + elem_row[e]]; if (USE_MAGIC_BIAS) { - nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + e]; + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } } - } -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - float corr_a[cms_per_row * CM_ELEMS]; - float sum_b[cms_per_col * CM_ELEMS]; - [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + float corr_a[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - corr_a[r * CM_ELEMS + e] = float(QUANT_OFFSET) * scale_a[r * CM_ELEMS + e] + corr_a[e] = float(QUANT_OFFSET) * scale_a[e] + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; } + + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat cache_b[K_SUB]; + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_b[h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + } + + float scale_b[CM_ELEMS]; + float sum_b[CM_ELEMS]; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + scale_b[e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; + sum_b[e] = buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]; + } + + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + acc = coopMatMulAdd(cache_a[h], cache_b[h], acc); + } + + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + if (USE_MAGIC_BIAS) { + const float t = fma(intBitsToFloat(int(acc[e])), + scale_a[e], + nbias_a[e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[e], + float(sums[tile_idx * CM_ELEMS + e]))); + } else { + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) + * scale_a[e] * scale_b[e]); + } + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( + corr_a[e] * sum_b[e]); + } + } + } +#else + // Simple types: pre-load all caches for better data reuse + coopmat cache_a[cms_per_row * K_SUB]; + coopmat cache_b[cms_per_col * K_SUB]; + + [[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, (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 e = 0; e < CM_ELEMS; e++) { - sum_b[c * CM_ELEMS + e] = buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]; + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); } } -#endif - coopmat accs[cms_per_row * cms_per_col]; + float scale_b[cms_per_col * CM_ELEMS]; + [[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 + b_col0 + c * TN + elem_col[e]]; + } + } [[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]]; + if (USE_MAGIC_BIAS) { + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; + } + } + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { coopmat acc = coopmat( @@ -416,28 +472,18 @@ void main() { acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); } - accs[r * cms_per_col + c] = acc; - } - } - - [[unroll]] for (uint r = 0; r < cms_per_row; r++) { - [[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++) { if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(accs[tile_idx][e])), - scale_a[r * CM_ELEMS + e], - nbias_a[r * CM_ELEMS + e]); + const float t = fma(intBitsToFloat(int(acc[e])), + scale_a[e], + nbias_a[e]); sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c * CM_ELEMS + e], float(sums[tile_idx * CM_ELEMS + e]))); } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(accs[tile_idx][e]) - * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) + * scale_a[e] * scale_b[c * CM_ELEMS + e]); } -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( - corr_a[r * CM_ELEMS + e] * sum_b[c * CM_ELEMS + e]); -#endif } } } From 092661df981f2158a9a5d9dd765bc4b4122fc1ca Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 15:45:24 +0200 Subject: [PATCH 38/50] improve performance --- .../vulkan-shaders/mul_mmq_cm1.comp | 37 +++++++++---------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 10 +++-- 2 files changed, 23 insertions(+), 24 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 337003683ded..3f5d5a94197b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -370,7 +370,6 @@ void main() { } } #elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - // Correction types: interleaved cache loading to reduce register pressure [[unroll]] for (uint r = 0; r < cms_per_row; r++) { coopmat cache_a[K_SUB]; [[unroll]] for (uint h = 0; h < K_SUB; h++) { @@ -386,25 +385,12 @@ void main() { } } - float corr_a[CM_ELEMS]; - [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - corr_a[e] = float(QUANT_OFFSET) * scale_a[e] - + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; - } - [[unroll]] for (uint c = 0; c < cms_per_col; c++) { coopmat cache_b[K_SUB]; [[unroll]] for (uint h = 0; h < K_SUB; h++) { coopMatLoad(cache_b[h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); } - float scale_b[CM_ELEMS]; - float sum_b[CM_ELEMS]; - [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_b[e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; - sum_b[e] = buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]; - } - coopmat acc = coopmat( USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); @@ -415,18 +401,29 @@ void main() { const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), - scale_a[e], - nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[e], + scale_a[e], nbias_a[e]); + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, sb, float(sums[tile_idx * CM_ELEMS + e]))); } else { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * scale_b[e]); + * scale_a[e] * sb); } - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( - corr_a[e] * sum_b[e]); + } + } + + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { +#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + const float ma = buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; +#else + const float ma = float(QUANT_OFFSET) * scale_a[e] + + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; +#endif + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + sums[(r * cms_per_col + c) * CM_ELEMS + e] += ACC_TYPE( + ma * buf_b_s[ks * BN + b_col0 + c * TN + elem_col[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 7536df4b0b2b..6a8d8a6c367c 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 @@ -257,8 +257,9 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); } vec2 dm = vec2(data_a_packed32[ib_k].dm); - buf_a_d[ks * BM + buf_ib] = dm.x * float(sc_val); - buf_a_m[ks * BM + buf_ib] = -(dm.y * float(mn_val)); + 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)); } } @@ -314,8 +315,9 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); } vec2 dm = vec2(data_a_packed32[ib_k].dm); - buf_a_d[ks * BM + buf_ib] = dm.x * float(sc_val); - buf_a_m[ks * BM + buf_ib] = -(dm.y * float(mn_val)); + 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)); } } From eeeacedd4e3ce86d7195130c34d5e58578681bda Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 08:57:10 +0200 Subject: [PATCH 39/50] improvements --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 8 ++-- .../vulkan-shaders/mul_mmq_cm1.comp | 43 +++++++++---------- 2 files changed, 25 insertions(+), 26 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 78f318247e2d..2755d442289f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4892,8 +4892,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); @@ -4905,8 +4905,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } 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 3f5d5a94197b..fdc7d89862f6 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -370,6 +370,14 @@ void main() { } } #elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + // Preload all B fragments for this ks to avoid re-loading them per r (halves B LDS reads). + coopmat cache_b[cms_per_col * K_SUB]; + [[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, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + } + } + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { coopmat cache_a[K_SUB]; [[unroll]] for (uint h = 0; h < K_SUB; h++) { @@ -378,54 +386,45 @@ void main() { float scale_a[CM_ELEMS]; 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]]; 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]]); +#else + ma[e] = float(QUANT_OFFSET) * scale_a[e] + + float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]); +#endif } [[unroll]] for (uint c = 0; c < cms_per_col; c++) { - coopmat cache_b[K_SUB]; - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(cache_b[h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); - } - coopmat acc = coopmat( USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); [[unroll]] for (uint h = 0; h < K_SUB; h++) { - acc = coopMatMulAdd(cache_a[h], cache_b[h], acc); + acc = coopMatMulAdd(cache_a[h], cache_b[c * K_SUB + h], acc); } const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; + const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]); if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), scale_a[e], nbias_a[e]); sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, sb, - float(sums[tile_idx * CM_ELEMS + e]))); + fma(ma[e], bs, + float(sums[tile_idx * CM_ELEMS + e])))); } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * sb); + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( + fma(float(acc[e]) * scale_a[e], sb, ma[e] * bs)); } } } - - [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { -#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - const float ma = buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; -#else - const float ma = float(QUANT_OFFSET) * scale_a[e] - + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; -#endif - [[unroll]] for (uint c = 0; c < cms_per_col; c++) { - sums[(r * cms_per_col + c) * CM_ELEMS + e] += ACC_TYPE( - ma * buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]); - } - } } #else // Simple types: pre-load all caches for better data reuse From 2ba1d225e16e200fbcb935b2f134abdd82e58c02 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 09:29:14 +0200 Subject: [PATCH 40/50] dedup b scales --- .../vulkan-shaders/mul_mmq_cm1.comp | 23 ++++++++----------- 1 file changed, 9 insertions(+), 14 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 fdc7d89862f6..1a1186ccce1e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -343,10 +343,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); - float scale_b[CM_ELEMS]; - [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - scale_b[e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; - } + const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; coopmat acc = coopmat( @@ -359,11 +356,11 @@ void main() { const float t = fma(intBitsToFloat(int(acc[e])), scale_a[e], nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[e], + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b_v, float(sums[tile_idx * CM_ELEMS + e]))); } else { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * scale_b[e]); + * scale_a[e] * scale_b_v); } } } @@ -410,9 +407,9 @@ 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]]); [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]]; - const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]); if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), scale_a[e], nbias_a[e]); @@ -442,11 +439,9 @@ void main() { } } - float scale_b[cms_per_col * CM_ELEMS]; + float scale_b[cms_per_col]; [[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 + b_col0 + c * TN + elem_col[e]]; - } + scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; } [[unroll]] for (uint r = 0; r < cms_per_row; r++) { @@ -474,11 +469,11 @@ void main() { const float t = fma(intBitsToFloat(int(acc[e])), scale_a[e], nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c * CM_ELEMS + e], + sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c], float(sums[tile_idx * CM_ELEMS + e]))); } else { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * scale_b[c * CM_ELEMS + e]); + * scale_a[e] * scale_b[c]); } } } From 9ef5996da43b04f273cd2536b2be1c8809d292cb Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 10:38:13 +0200 Subject: [PATCH 41/50] 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 1a1186ccce1e..3b7ae9a6df73 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 6a8d8a6c367c..e2ee2795d2ff 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))); } } From f252045d8928c1f7c618f46b043df6e6b77c5fb4 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 11:31:03 +0200 Subject: [PATCH 42/50] undo uint8_t, gate to RDNA3/4 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 3 +- .../vulkan-shaders/mul_mmq_cm1.comp | 30 +++++++++---------- 2 files changed, 17 insertions(+), 16 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 2755d442289f..22cf2eeb9055 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4883,7 +4883,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } - if (device->coopmat_int_support) { + // cm1 int8 MMQ assumes the RDNA wave32 WMMA accumulator layout (RDNA4 reports as AMD_RDNA3) + if (device->coopmat_int_support && device->architecture == AMD_RDNA3) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); 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 3b7ae9a6df73..80fa147aa057 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -208,8 +208,8 @@ void main() { const uint warp_c = warp_i / (BM / WM); // Probe coopmat element layout to discover row/col mapping - uint8_t elem_row[CM_ELEMS]; - uint8_t elem_col[CM_ELEMS]; + uint elem_row[CM_ELEMS]; + uint elem_col[CM_ELEMS]; for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { cm_layout_probe[i] = int32_t(i); @@ -220,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] = uint8_t(uint(probe[e]) / TN); - elem_col[e] = uint8_t(uint(probe[e]) % TN); + elem_row[e] = uint(probe[e]) / TN; + elem_col[e] = uint(probe[e]) % TN; } const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); @@ -336,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 + uint(elem_row[e])]; + scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + elem_row[e]]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -346,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 + uint(elem_col[0])]; + const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]]; coopmat acc = coopmat( @@ -388,7 +388,7 @@ void main() { float nbias_a[CM_ELEMS]; float ma[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + uint(elem_row[e])]; + vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + elem_row[e]]; scale_a[e] = dm.x; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; @@ -410,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 + uint(elem_col[0])]; - const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + uint(elem_col[0])]); + 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]]); [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), @@ -444,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 + uint(elem_col[0])]; + scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + 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 + uint(elem_row[e])]; + scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -500,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 + uint(elem_col[e]); + const uint col_i = dc + c * TN + elem_col[e]; if (col_i >= _ne1) continue; - const uint row_g = dr + r * TM + uint(elem_row[e]); + const uint row_g = dr + r * TM + elem_row[e]; if (row_g >= p.M) continue; const u16vec2 row_idx = row_ids[col_i - ic * BN]; @@ -519,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 + uint(elem_row[e]); - const uint col_g = dc + c * TN + uint(elem_col[e]); + const uint row_g = dr + r * TM + elem_row[e]; + const uint col_g = dc + c * TN + 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]); } From 0eab519531405f41791bba9298651255bfe31a4d Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 14:05:20 +0200 Subject: [PATCH 43/50] add RDNA4 architecture, use for hardcoded coopmat elem thread access, set BK_STEP back to 4 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 26 ++++++---- .../vulkan-shaders/mul_mmq_cm1.comp | 52 +++++++------------ 2 files changed, 37 insertions(+), 41 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 22cf2eeb9055..cf76db9a0954 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -389,12 +389,14 @@ static void ggml_vk_synchronize(ggml_backend_vk_context * ctx); static constexpr uint32_t mul_mat_vec_max_cols = 8; static constexpr uint32_t p021_max_gqa_ratio = 8; +// Values are passed to mul_mmq_cm1.comp as spec constant CM_ARCH; keep CM_ARCH_AMD_RDNA4 in sync. enum vk_device_architecture { OTHER, AMD_GCN, AMD_RDNA1, AMD_RDNA2, AMD_RDNA3, + AMD_RDNA4, INTEL_XE1, INTEL_XE2, NVIDIA_PRE_TURING, @@ -410,6 +412,7 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& bool amd_shader_core_properties = false; bool integer_dot_product = false; bool subgroup_size_control = false; + bool shader_float8 = false; // RDNA4-only (fp8 WMMA), distinguishes it from RDNA3 for (const auto& properties : ext_props) { if (strcmp("VK_AMD_shader_core_properties", properties.extensionName) == 0) { @@ -418,6 +421,8 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& integer_dot_product = true; } else if (strcmp("VK_EXT_subgroup_size_control", properties.extensionName) == 0) { subgroup_size_control = true; + } else if (strcmp("VK_EXT_shader_float8", properties.extensionName) == 0) { + shader_float8 = true; } } @@ -444,6 +449,9 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& if (shader_core_props_amd.wavefrontsPerSimd == 20) { return vk_device_architecture::AMD_RDNA1; } + if (shader_float8) { + return vk_device_architecture::AMD_RDNA4; + } if (integer_dot_props.integerDotProduct4x8BitPackedMixedSignednessAccelerated) { return vk_device_architecture::AMD_RDNA3; } @@ -4346,13 +4354,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32); }; - l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm_l, itn_l, itk_l, cm1_sg }; - m_warptile_mmq_cm1_int = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; - s_warptile_mmq_cm1_int = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; + l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm_l, itn_l, itk_l, cm1_sg, (uint32_t)device->architecture }; + m_warptile_mmq_cm1_int = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg, (uint32_t)device->architecture }; + s_warptile_mmq_cm1_int = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg, (uint32_t)device->architecture }; - l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128), 64, 128, 32, std::min(cm1_sg, 64u), 32, 2, itm_l, itn_l, itk_l, cm1_sg }; - m_warptile_mmq_cm1_int_k = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; - s_warptile_mmq_cm1_int_k = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; + l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128), 64, 128, 32, std::min(cm1_sg, 64u), 32, 2, itm_l, itn_l, itk_l, cm1_sg, (uint32_t)device->architecture }; + m_warptile_mmq_cm1_int_k = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg, (uint32_t)device->architecture }; + s_warptile_mmq_cm1_int_k = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg, (uint32_t)device->architecture }; l_mmq_cm1_wg_denoms_k = { l_warptile_mmq_cm1_int_k[1], l_warptile_mmq_cm1_int_k[2], 1 }; m_mmq_cm1_wg_denoms_k = { m_warptile_mmq_cm1_int_k[1], m_warptile_mmq_cm1_int_k[2], 1 }; @@ -4883,8 +4891,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } - // cm1 int8 MMQ assumes the RDNA wave32 WMMA accumulator layout (RDNA4 reports as AMD_RDNA3) - if (device->coopmat_int_support && device->architecture == AMD_RDNA3) { + // cm1 int8 MMQ assumes the RDNA wave32 WMMA accumulator layout (row order set per-arch via CM_ARCH) + if (device->coopmat_int_support && (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4)) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); @@ -19360,7 +19368,7 @@ static bool ggml_vk_khr_cooperative_matrix_support(const vk::PhysicalDevicePrope case VK_VENDOR_ID_AMD: if (driver_props.driverID == vk::DriverId::eAmdProprietary || driver_props.driverID == vk::DriverId::eAmdOpenSource) { // Workaround for AMD proprietary driver reporting support on all GPUs - return arch == vk_device_architecture::AMD_RDNA3; + return arch == vk_device_architecture::AMD_RDNA3 || arch == vk_device_architecture::AMD_RDNA4; } return true; default: 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 80fa147aa057..df31bef5d6ea 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -82,13 +82,11 @@ layout (constant_id = 7) const uint TM = 16; layout (constant_id = 8) const uint TN = 16; layout (constant_id = 9) const uint TK = 16; layout (constant_id = 10) const uint WARP = 32; +layout (constant_id = 11) const uint CM_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp) +#define CM_ARCH_AMD_RDNA4 5u // must match enum vk_device_architecture #define BK 32 -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) -#define BK_STEP 2 -#else #define BK_STEP 4 -#endif #define GROUP_A_BUDGET (16u * 1024u * 1024u) const uint QPITCH = BK_STEP * (BK / 4) + 4; @@ -130,7 +128,12 @@ const uint CM_ELEMS = (TM * TN) / WARP; #define ACC_BIAS_F 12582912.0f const bool USE_MAGIC_BIAS = WARP != 32; -shared int32_t cm_layout_probe[TM * TN]; +// Coopmat accumulator layout (RDNA): col constant per lane, 8 distinct rows. +// Row order differs by arch: RDNA4 blocked, RDNA3/3.5 interleaved. +uint cm_elem_row(uint e) { + const uint row_half = gl_SubgroupInvocationID / TN; + return (CM_ARCH == CM_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e); +} #ifdef MUL_MAT_ID #define NUM_WARPS (BLOCK_SIZE / WARP) @@ -207,22 +210,7 @@ void main() { const uint warp_r = warp_i % (BM / WM); 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]; - - for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) { - cm_layout_probe[i] = int32_t(i); - } - barrier(); - - coopmat probe; - 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; - } + const uint elem_col0 = gl_SubgroupInvocationID % TN; // constant per lane const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A); @@ -336,7 +324,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 + cm_elem_row(e)]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -346,7 +334,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 + elem_col0]; coopmat acc = coopmat( @@ -388,7 +376,7 @@ void main() { float nbias_a[CM_ELEMS]; float ma[CM_ELEMS]; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + elem_row[e]]; + vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + cm_elem_row(e)]; scale_a[e] = dm.x; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; @@ -410,8 +398,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 + elem_col0]; + const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col0]); [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { if (USE_MAGIC_BIAS) { const float t = fma(intBitsToFloat(int(acc[e])), @@ -444,14 +432,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 + elem_col0]; } [[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 + cm_elem_row(e)]; if (USE_MAGIC_BIAS) { nbias_a[e] = -ACC_BIAS_F * scale_a[e]; } @@ -500,10 +488,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 + elem_col0; if (col_i >= _ne1) continue; - const uint row_g = dr + r * TM + elem_row[e]; + const uint row_g = dr + r * TM + cm_elem_row(e); if (row_g >= p.M) continue; const u16vec2 row_idx = row_ids[col_i - ic * BN]; @@ -519,8 +507,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 + cm_elem_row(e); + const uint col_g = dc + c * TN + elem_col0; 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]); } From c30a8740f2c600aa2226c4c1fea235962539fe96 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Thu, 27 Aug 2026 16:57:05 +0200 Subject: [PATCH 44/50] improve offset application --- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 10 ---------- .../ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl | 6 ++++-- 2 files changed, 4 insertions(+), 12 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 df31bef5d6ea..bd41b1c18685 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -106,12 +106,6 @@ shared float buf_b_d[BN * BK_STEP]; shared float buf_b_s[BN * BK_STEP]; #endif -#if defined(DATA_A_Q4_1) || defined(DATA_A_Q4_K) -#define QUANT_OFFSET 8.0 -#elif defined(DATA_A_Q5_1) || defined(DATA_A_Q5_K) -#define QUANT_OFFSET 16.0 -#endif - #if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) shared int8_t cm1_kvalues[16]; #endif @@ -381,11 +375,7 @@ void main() { 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] = dm.y; -#else - ma[e] = float(QUANT_OFFSET) * scale_a[e] + dm.y; -#endif } [[unroll]] for (uint c = 0; c < cms_per_col; c++) { 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 e2ee2795d2ff..c9d2da0fef51 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,7 +55,8 @@ 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_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); + const float d = float(blk.dm.x); + buf_a_dm[ks * BM + buf_ib] = vec2(d, 8.0 * d + float(blk.dm.y)); } } @@ -118,7 +119,8 @@ 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_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); + const float d = float(blk.dm.x); + buf_a_dm[ks * BM + buf_ib] = vec2(d, 16.0 * d + float(blk.dm.y)); } } From 6421fbe532a927e2ad080e1860d578ae043fc3af Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Fri, 28 Aug 2026 16:02:01 +0200 Subject: [PATCH 45/50] clean up --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 1 + .../vulkan-shaders/mul_mmq_cm1.comp | 87 ++++++++----------- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 54 +++++------- 3 files changed, 60 insertions(+), 82 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index cf76db9a0954..8534e9e2324d 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4888,6 +4888,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } else #endif { + CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_MXFP4], matmul_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } 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 bd41b1c18685..8c195e0f1269 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -10,7 +10,6 @@ #extension GL_KHR_memory_scope_semantics : enable #if defined(MUL_MAT_ID_USE_SUBGROUPS) -#extension GL_KHR_shader_subgroup_basic : enable #extension GL_KHR_shader_subgroup_ballot : enable #endif @@ -91,10 +90,9 @@ layout (constant_id = 11) const uint CM_ARCH = 0; // vk_device_architecture (gg 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) +shared vec2 buf_a_dm[BM * BK_STEP]; // .x = d, .y = m #else shared float buf_a_d[BM * BK_STEP * KSCALES]; #endif @@ -122,13 +120,21 @@ const uint CM_ELEMS = (TM * TN) / WARP; #define ACC_BIAS_F 12582912.0f const bool USE_MAGIC_BIAS = WARP != 32; -// Coopmat accumulator layout (RDNA): col constant per lane, 8 distinct rows. -// Row order differs by arch: RDNA4 blocked, RDNA3/3.5 interleaved. +// Accumulator row for element e: RDNA4 blocked, RDNA3/3.5 interleaved. uint cm_elem_row(uint e) { const uint row_half = gl_SubgroupInvocationID / TN; return (CM_ARCH == CM_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e); } +// min_term = asymmetric-quant min*b_sum correction (0 for symmetric types). +ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a, float scale_b, float min_term) { + if (USE_MAGIC_BIAS) { + const float t = fma(intBitsToFloat(acc_e), scale_a, nbias_a); + return ACC_TYPE(fma(t, scale_b, float(prev) + min_term)); + } + return prev + ACC_TYPE(fma(float(acc_e) * scale_a, scale_b, min_term)); +} + #ifdef MUL_MAT_ID #define NUM_WARPS (BLOCK_SIZE / WARP) #include "mul_mm_id_funcs.glsl" @@ -237,7 +243,6 @@ void main() { barrier(); #endif - // Workgroup has no work if (ic * BN >= _ne1) return; #endif @@ -274,7 +279,6 @@ void main() { block_a_prefetch pre_a[A_LOADS * BK_STEP]; block_b_prefetch pre_b[B_LOADS * BK_STEP]; - // Prefetch first block if (start_k < end_k) { PREFETCH_BLOCK(start_k) } @@ -290,7 +294,6 @@ void main() { barrier(); for (uint block = start_k; block < end_k; block += BK * BK_STEP) { - // Store prefetched data to shmem STORE_BLOCK_TO_LDS(block) barrier(); @@ -298,14 +301,12 @@ void main() { pos_a_ib += BK_STEP; pos_b_ib += BK_STEP; - // Prefetch next block (overlaps with compute) const uint next_block = block + BK * BK_STEP; if (next_block < end_k) { 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; @@ -337,35 +338,37 @@ void main() { const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(acc[e])), - scale_a[e], - nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b_v, - float(sums[tile_idx * CM_ELEMS + e]))); - } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * scale_b_v); - } + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b_v, 0.0); } } } } #elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) - // Preload all B fragments for this ks to avoid re-loading them per r (halves B LDS reads). + // Preload all A/B fragments up front (ILP). + coopmat cache_a[cms_per_row * K_SUB]; coopmat cache_b[cms_per_col * K_SUB]; + + [[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, (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, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); } } - [[unroll]] for (uint r = 0; r < cms_per_row; r++) { - coopmat cache_a[K_SUB]; - [[unroll]] for (uint h = 0; h < K_SUB; h++) { - coopMatLoad(cache_a[h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); - } + float scale_b[cms_per_col]; + float bs[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_col0]; + bs[c] = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col0]); + } + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { float scale_a[CM_ELEMS]; float nbias_a[CM_ELEMS]; float ma[CM_ELEMS]; @@ -384,28 +387,19 @@ void main() { USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); [[unroll]] for (uint h = 0; h < K_SUB; h++) { - acc = coopMatMulAdd(cache_a[h], cache_b[c * K_SUB + h], acc); + acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); } const uint tile_idx = r * cms_per_col + c; - const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0]; - const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col0]); [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(acc[e])), - scale_a[e], nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, sb, - fma(ma[e], bs, - float(sums[tile_idx * CM_ELEMS + e])))); - } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( - fma(float(acc[e]) * scale_a[e], sb, ma[e] * bs)); - } + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b[c], ma[e] * bs[c]); } } } #else - // Simple types: pre-load all caches for better data reuse + // Preload all A/B fragments up front (ILP). coopmat cache_a[cms_per_row * K_SUB]; coopmat cache_b[cms_per_col * K_SUB]; @@ -446,16 +440,9 @@ void main() { const uint tile_idx = r * cms_per_col + c; [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { - if (USE_MAGIC_BIAS) { - const float t = fma(intBitsToFloat(int(acc[e])), - scale_a[e], - nbias_a[e]); - sums[tile_idx * CM_ELEMS + e] = ACC_TYPE(fma(t, scale_b[c], - float(sums[tile_idx * CM_ELEMS + e]))); - } else { - sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(acc[e]) - * scale_a[e] * scale_b[c]); - } + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b[c], 0.0); } } } 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 c9d2da0fef51..3b07934f9898 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 @@ -47,16 +47,14 @@ block_a_prefetch block_a_load(uint ib, uint loadr) { } void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned nibbles; the -8 offset is absorbed by the min term. uint32_t lo4 = blk.qs & 0x0F0F0F0F; uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; - lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; - hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; if (loadr == 0) { - const float d = float(blk.dm.x); - buf_a_dm[ks * BM + buf_ib] = vec2(d, 8.0 * d + float(blk.dm.y)); + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); } } @@ -109,18 +107,16 @@ block_a_prefetch block_a_load(uint ib, uint loadr) { } void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned 5-bit values; the -16 offset is absorbed by the min term. uint32_t lo4 = blk.qs & 0x0F0F0F0F; uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; lo4 |= ((blk.qh >> (4u * loadr )) & 0xFu) * 0x02040810u & 0x10101010u; hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; - lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; - hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; if (loadr == 0) { - const float d = float(blk.dm.x); - buf_a_dm[ks * BM + buf_ib] = vec2(d, 16.0 * d + float(blk.dm.y)); + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); } } @@ -240,25 +236,22 @@ block_a_prefetch block_a_load(uint ib, uint loadr) { } void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { - uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x08080808) ^ 0x80808080; - uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x08080808) ^ 0x80808080; - buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; - buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + // Store raw unsigned nibbles (blk.qs already masked); no -8 recentering needed. + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = blk.qs0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = blk.qs1; if (loadr == 0) { const uint ib_k = blk.ib / 8; const uint sub = blk.ib % 8; - uint sc_val, mn_val; - if (sub < 4) { - sc_val = uint(data_a[ib_k].scales[sub]) & 0x3Fu; - mn_val = uint(data_a[ib_k].scales[sub + 4]) & 0x3Fu; - } else { - sc_val = (uint(data_a[ib_k].scales[sub + 4]) & 0xFu) | ((uint(data_a[ib_k].scales[sub - 4]) & 0xC0u) >> 2); - mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); - } + const uint j = sub & 3u; + const uint s_j = uint(data_a[ib_k].scales[j]); + const uint s_j4 = uint(data_a[ib_k].scales[j + 4]); + const uint s_j8 = uint(data_a[ib_k].scales[j + 8]); + const uint sc_val = (sub < 4) ? (s_j & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j >> 6) << 4)); + const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4) | ((s_j4 >> 6) << 4)); vec2 dm = vec2(data_a_packed32[ib_k].dm); float d_scaled = dm.x * float(sc_val); - buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 8.0 * d_scaled - (dm.y * float(mn_val))); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val))); } } @@ -295,27 +288,24 @@ block_a_prefetch block_a_load(uint ib, uint loadr) { } void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned 5-bit values (qs nibble | qh bit); no -16 recentering needed. uint32_t v0 = blk.qs0 | blk.qh0; uint32_t v1 = blk.qs1 | blk.qh1; - v0 = ((v0 | 0x80808080) - 0x10101010) ^ 0x80808080; - v1 = ((v1 | 0x80808080) - 0x10101010) ^ 0x80808080; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; if (loadr == 0) { const uint ib_k = blk.ib / 8; const uint sub = blk.ib % 8; - uint sc_val, mn_val; - if (sub < 4) { - sc_val = uint(data_a[ib_k].scales[sub]) & 0x3Fu; - mn_val = uint(data_a[ib_k].scales[sub + 4]) & 0x3Fu; - } else { - sc_val = (uint(data_a[ib_k].scales[sub + 4]) & 0xFu) | ((uint(data_a[ib_k].scales[sub - 4]) & 0xC0u) >> 2); - mn_val = (uint(data_a[ib_k].scales[sub + 4]) >> 4) | ((uint(data_a[ib_k].scales[sub]) & 0xC0u) >> 2); - } + const uint j = sub & 3u; + const uint s_j = uint(data_a[ib_k].scales[j]); + const uint s_j4 = uint(data_a[ib_k].scales[j + 4]); + const uint s_j8 = uint(data_a[ib_k].scales[j + 8]); + const uint sc_val = (sub < 4) ? (s_j & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j >> 6) << 4)); + const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4) | ((s_j4 >> 6) << 4)); vec2 dm = vec2(data_a_packed32[ib_k].dm); float d_scaled = dm.x * float(sc_val); - buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 16.0 * d_scaled - (dm.y * float(mn_val))); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val))); } } From 04d7ef7b247b89a61060837492f3778396a1687a Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Fri, 28 Aug 2026 18:38:12 +0200 Subject: [PATCH 46/50] fix iq4_nl and nvfp4 performance --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 2 ++ 1 file changed, 2 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8534e9e2324d..c1f649eda583 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4078,6 +4078,8 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std: case GGML_TYPE_Q5_1: block_a_size = std430_size({{16, 4}, {4, 4}, {fp2_size, fp2_align}}); break; // qs[16/4] + qh + dm(vec2) case GGML_TYPE_Q8_0: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + dm case GGML_TYPE_MXFP4: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + d + case GGML_TYPE_IQ4_NL: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + d + case GGML_TYPE_NVFP4: block_a_size = std430_size({{32, 4}, {fp2_size, fp2_align}}); break; // qs[8] + d_scales(vec2) case GGML_TYPE_Q2_K: block_a_size = std430_size({{ 8, 4}, {2, 2}, {fp2_size, fp2_align}}); break; // qs[2] + scales(u8vec2) + dm(vec2) case GGML_TYPE_Q3_K: block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}}); break; // qs[4] + d_scales(vec2) case GGML_TYPE_Q4_K: block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}}); break; // qs[4] + dm(vec2) From 584ae05d5f95df2ed79a163b2996dfcdac78644c Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 29 Aug 2026 07:36:13 +0200 Subject: [PATCH 47/50] rdna4 tuning --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index c1f649eda583..9ad7cde1a22e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4896,18 +4896,21 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // cm1 int8 MMQ assumes the RDNA wave32 WMMA accumulator layout (row order set per-arch via CM_ARCH) if (device->coopmat_int_support && (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4)) { + // RDNA4: skip quants that regress on the int8 MMQ path, dispatch falls back to dequant + const bool rdna4 = device->architecture == AMD_RDNA4; + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); - CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); } CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4920,7 +4923,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); - CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } } GGML_ASSERT(device->subgroup_ballot); From e85ff8433461c1fce7d37fcad1aff4c2ee0c0024 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 29 Aug 2026 09:42:10 +0200 Subject: [PATCH 48/50] use BK_STEP 2 on MUL_MAT_ID --- ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 4 ++++ 1 file changed, 4 insertions(+) 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 8c195e0f1269..50f067b0f31e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -85,7 +85,11 @@ layout (constant_id = 11) const uint CM_ARCH = 0; // vk_device_architecture (gg #define CM_ARCH_AMD_RDNA4 5u // must match enum vk_device_architecture #define BK 32 +#ifdef MUL_MAT_ID +#define BK_STEP 2 +#else #define BK_STEP 4 +#endif #define GROUP_A_BUDGET (16u * 1024u * 1024u) const uint QPITCH = BK_STEP * (BK / 4) + 4; From 8b194a15eb47f5ab5f48341e86b386f4e6485549 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 29 Aug 2026 11:42:29 +0200 Subject: [PATCH 49/50] adapt to upstream changes --- .../vulkan-shaders/mul_mmq_cm1.comp | 34 +++++++++++-------- 1 file changed, 20 insertions(+), 14 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 50f067b0f31e..8c2ce006b33a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -60,6 +60,8 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; + uint n_experts; + uint hoist_row_ids; #else uint base_work_group_z; uint num_batches; @@ -225,27 +227,31 @@ void main() { const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK; #ifdef MUL_MAT_ID -#ifdef MUL_MAT_ID_USE_SUBGROUPS - if (bitCount(p.nei0) == 1) { - load_row_ids(expert_idx, true, ic); + if (p.hoist_row_ids != 0) { + load_row_ids_hoisted(expert_idx, ic); } else { - load_row_ids(expert_idx, false, ic); - } +#ifdef MUL_MAT_ID_USE_SUBGROUPS + if (bitCount(p.nei0) == 1) { + load_row_ids(expert_idx, true, ic); + } else { + load_row_ids(expert_idx, false, ic); + } #else - _ne1 = 0; - for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { - for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { - if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { - if (_ne1 >= ic * BN) { - row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + _ne1 = 0; + for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { + for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { + if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { + if (_ne1 >= ic * BN) { + row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + } + _ne1++; } - _ne1++; } } - } - barrier(); + barrier(); #endif + } if (ic * BN >= _ne1) return; #endif From 965e57103fce2c4329cdfc2b8300f8f7ed57c9fe Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Sat, 29 Aug 2026 13:58:45 +0200 Subject: [PATCH 50/50] fix shmem support function, clean up comments --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 124 ++++++++++++++---- .../vulkan-shaders/mul_mmq_cm1.comp | 8 +- 2 files changed, 103 insertions(+), 29 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 9ad7cde1a22e..c96a12e42310 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -389,7 +389,6 @@ static void ggml_vk_synchronize(ggml_backend_vk_context * ctx); static constexpr uint32_t mul_mat_vec_max_cols = 8; static constexpr uint32_t p021_max_gqa_ratio = 8; -// Values are passed to mul_mmq_cm1.comp as spec constant CM_ARCH; keep CM_ARCH_AMD_RDNA4 in sync. enum vk_device_architecture { OTHER, AMD_GCN, @@ -412,7 +411,7 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& bool amd_shader_core_properties = false; bool integer_dot_product = false; bool subgroup_size_control = false; - bool shader_float8 = false; // RDNA4-only (fp8 WMMA), distinguishes it from RDNA3 + bool shader_float8 = false; for (const auto& properties : ext_props) { if (strcmp("VK_AMD_shader_core_properties", properties.extensionName) == 0) { @@ -4113,6 +4112,66 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std: return supported; } +static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const std::vector& warptile, bool mul_mat_id, ggml_type src0_type) { + + bool kscales2 = false; // two scale sets per block + bool has_dm = false; // d+m as vec2 + b-side sum + bool has_kvalues = false; + switch (src0_type) { + case GGML_TYPE_Q4_0: case GGML_TYPE_Q5_0: case GGML_TYPE_Q8_0: + break; + case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_1: + case GGML_TYPE_Q4_K: case GGML_TYPE_Q5_K: + has_dm = true; break; + case GGML_TYPE_IQ4_NL: case GGML_TYPE_MXFP4: + has_kvalues = true; break; + case GGML_TYPE_Q3_K: case GGML_TYPE_Q6_K: + kscales2 = true; break; + case GGML_TYPE_NVFP4: + kscales2 = true; has_kvalues = true; break; + default: + return false; + } + + const uint32_t BLOCK_SIZE = warptile[0]; + const uint32_t BM = warptile[1]; + const uint32_t BN = warptile[2]; + const uint32_t WARP = warptile[10]; + + const uint32_t BK = 32; + const uint32_t BK_STEP = mul_mat_id ? 2u : 4u; + const uint32_t QPITCH = BK_STEP * (BK / 4u) + 4u; + const uint32_t KSCALES = kscales2 ? 2u : 1u; + + uint32_t total = 0; + total += BM * QPITCH * (uint32_t)sizeof(uint32_t); // buf_a_qs + total += BN * QPITCH * (uint32_t)sizeof(uint32_t); // buf_b_qs + total += has_dm ? (BM * BK_STEP * 2u * (uint32_t)sizeof(float)) // buf_a_dm (vec2) + : (BM * BK_STEP * KSCALES * (uint32_t)sizeof(float)); // buf_a_d + total += BN * BK_STEP * (uint32_t)sizeof(float); // buf_b_d + if (has_dm) { + total += BN * BK_STEP * (uint32_t)sizeof(float); // buf_b_s + } + if (has_kvalues) { + total += 16u * (uint32_t)sizeof(int8_t); // cm1_kvalues[16] + } + if (src0_type == GGML_TYPE_NVFP4 && !device->ocp_fp4) { + total += 128u * (uint32_t)sizeof(float); // ue4m3_fp32_lut[128] + } + if (mul_mat_id) { + total += BN * 2u * (uint32_t)sizeof(uint16_t); // row_ids[BN] (u16vec2) + const uint32_t num_warps = BLOCK_SIZE / std::max(WARP, 1u); + total += num_warps * 4u * (uint32_t)sizeof(uint32_t); // ballots_sh[NUM_WARPS] (uvec4) + } + + const bool supported = total <= device->properties.limits.maxComputeSharedMemorySize; + + VK_LOG_DEBUG("ggml_vk_matmul_cm1_int_shmem_support(warptile=(" << warptile[0] << "," << warptile[1] << "," << warptile[2] << "), " + "mul_mat_id=" << mul_mat_id << ", src0_type=" << ggml_type_name(src0_type) << ", total=" << total << ", supported=" << supported); + + return supported; +} + struct GpuPipelineConfig { // GPU architecture identifier. // Example: vk_device_architecture::AMD_GCN @@ -4411,6 +4470,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_align = 64; s_align = 32; + const bool use_cm1_int = device->coopmat_int_support && + (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4); + for (uint32_t i = 0; i < GGML_TYPE_COUNT; ++i) { ggml_type t = (ggml_type)i; // Disable medium and large matrix multiplication if not enough shared memory is available @@ -4438,37 +4500,50 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { device->mul_mat_id_l[i] = false; } - // The q8_1 mmq path has its own (larger) shmem layout, check it separately. - // K-quants use the _int_k warptiles, others use _int. + // cm1 splits k-tiles on the KSCALES==2 types and shares tiles between dense/id. const bool is_k_quant = (t == GGML_TYPE_Q2_K || t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q4_K || t == GGML_TYPE_Q5_K || t == GGML_TYPE_Q6_K); - const auto & s_int = is_k_quant ? s_warptile_mmq_int_k : s_warptile_mmq_int; - const auto & m_int = is_k_quant ? m_warptile_mmq_int_k : m_warptile_mmq_int; - const auto & l_int = is_k_quant ? l_warptile_mmq_int_k : l_warptile_mmq_int; - const auto & s_intid = is_k_quant ? s_warptile_mmqid_int_k : s_warptile_mmqid_int; - const auto & m_intid = is_k_quant ? m_warptile_mmqid_int_k : m_warptile_mmqid_int; - const auto & l_intid = is_k_quant ? l_warptile_mmqid_int_k : l_warptile_mmqid_int; - - if (!ggml_vk_matmul_int_shmem_support(device, s_int, false, t)) { + const bool cm1_k_tile = (t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q6_K || + t == GGML_TYPE_NVFP4); + + const auto & s_int = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int) + : (is_k_quant ? s_warptile_mmq_int_k : s_warptile_mmq_int); + const auto & m_int = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int) + : (is_k_quant ? m_warptile_mmq_int_k : m_warptile_mmq_int); + const auto & l_int = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int) + : (is_k_quant ? l_warptile_mmq_int_k : l_warptile_mmq_int); + const auto & s_intid = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int) + : (is_k_quant ? s_warptile_mmqid_int_k : s_warptile_mmqid_int); + const auto & m_intid = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int) + : (is_k_quant ? m_warptile_mmqid_int_k : m_warptile_mmqid_int); + const auto & l_intid = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int) + : (is_k_quant ? l_warptile_mmqid_int_k : l_warptile_mmqid_int); + + const auto int_shmem_support = [&](const std::vector& wt, bool id) { + return use_cm1_int ? ggml_vk_matmul_cm1_int_shmem_support(device, wt, id, t) + : ggml_vk_matmul_int_shmem_support(device, wt, id, t); + }; + + if (!int_shmem_support(s_int, false)) { device->mul_mat_s_int[i] = false; device->mul_mat_m_int[i] = false; device->mul_mat_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, m_int, false, t)) { + } else if (!int_shmem_support(m_int, false)) { device->mul_mat_m_int[i] = false; device->mul_mat_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, l_int, false, t)) { + } else if (!int_shmem_support(l_int, false)) { device->mul_mat_l_int[i] = false; } - if (!ggml_vk_matmul_int_shmem_support(device, s_intid, true, t)) { + if (!int_shmem_support(s_intid, true)) { device->mul_mat_id_s_int[i] = false; device->mul_mat_id_m_int[i] = false; device->mul_mat_id_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, m_intid, true, t)) { + } else if (!int_shmem_support(m_intid, true)) { device->mul_mat_id_m_int[i] = false; device->mul_mat_id_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, l_intid, true, t)) { + } else if (!int_shmem_support(l_intid, true)) { device->mul_mat_id_l_int[i] = false; } } @@ -4829,11 +4904,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true); \ #define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ - if (device->mul_mat ## ID ## _l[TYPE]) \ + if (device->mul_mat ## ID ## _l_int[TYPE]) \ ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, cm1_sg); \ - if (device->mul_mat ## ID ## _m[TYPE]) \ + if (device->mul_mat ## ID ## _m_int[TYPE]) \ ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, cm1_sg); \ - if (device->mul_mat ## ID ## _s[TYPE]) \ + if (device->mul_mat ## ID ## _s_int[TYPE]) \ ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, cm1_sg); \ // Create 2 variants, {f16,f32} accumulator @@ -4894,11 +4969,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } - // cm1 int8 MMQ assumes the RDNA wave32 WMMA accumulator layout (row order set per-arch via CM_ARCH) - if (device->coopmat_int_support && (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4)) { - // RDNA4: skip quants that regress on the int8 MMQ path, dispatch falls back to dequant - const bool rdna4 = device->architecture == AMD_RDNA4; - + // Some quants are not performant on RDNA4, those fall back to FP16 matmul + const bool rdna3 = device->architecture == AMD_RDNA3; + const bool rdna4 = device->architecture == AMD_RDNA4; + if (device->coopmat_int_support && (rdna3 || rdna4)) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); 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 8c2ce006b33a..eccd69869b79 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -83,8 +83,8 @@ layout (constant_id = 7) const uint TM = 16; layout (constant_id = 8) const uint TN = 16; layout (constant_id = 9) const uint TK = 16; layout (constant_id = 10) const uint WARP = 32; -layout (constant_id = 11) const uint CM_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp) -#define CM_ARCH_AMD_RDNA4 5u // must match enum vk_device_architecture +layout (constant_id = 11) const uint DEVICE_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp) +#define VK_ARCH_AMD_RDNA4 5u #define BK 32 #ifdef MUL_MAT_ID @@ -129,7 +129,7 @@ const bool USE_MAGIC_BIAS = WARP != 32; // Accumulator row for element e: RDNA4 blocked, RDNA3/3.5 interleaved. uint cm_elem_row(uint e) { const uint row_half = gl_SubgroupInvocationID / TN; - return (CM_ARCH == CM_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e); + return (DEVICE_ARCH == VK_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e); } // min_term = asymmetric-quant min*b_sum correction (0 for symmetric types). @@ -216,7 +216,7 @@ void main() { const uint warp_r = warp_i % (BM / WM); const uint warp_c = warp_i / (BM / WM); - const uint elem_col0 = gl_SubgroupInvocationID % TN; // constant per lane + const uint elem_col0 = gl_SubgroupInvocationID % TN; const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A);