Commit 966baae76 for llama.cpp

commit 966baae76b578d5d3dd4c179a75a35c80ff75577
Author: Jeff Bolz <jbolz@nvidia.com>
Date:   Sun Oct 11 07:02:57 2026 -0500

    vulkan: handle mul_mat_id duplicates in the prepass rather than looping (#29998)

    * vulkan: compute every row of an expert in mul_mm_id when ids repeat

    * vulkan: handle mul_mat_id duplicates in the prepass rather than looping

    Extend the "hoist row ids" optimization to always be enabled and to emit a
    compact list of tile descriptions that need to run, and to emit multiple tiles
    when needed. Then we can launch a tighter upper bound on the number of
    workgroups and the tail can trivially early exit.

    ---------

    Co-authored-by: Anjielon <[EMAIL_REDACTED]>

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
index f1f66a628..fdce443a1 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
@@ -63,7 +63,7 @@ struct vk_mat_mat_id_push_constants {
     uint32_t batch_stride_a; uint32_t batch_stride_b; uint32_t batch_stride_d;
     uint32_t nei0; uint32_t nei1; uint32_t nbi1; uint32_t ne11;
     uint32_t n_experts;
-    uint32_t hoist_row_ids;
+    uint32_t row_ids_offset;
 };

 struct vk_mat_vec_id_push_constants {
@@ -218,9 +218,10 @@ struct vk_op_count_experts_push_constants {
     uint32_t nb01;
     uint32_t a_offset;
     uint32_t n_experts;
-    uint32_t hoist_row_ids;
     uint32_t ne00mp;
     uint32_t ne00L;
+    uint32_t row_tile_size;
+    uint32_t row_ids_offset;
 };

 struct vk_op_glu_push_constants {
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 907ac82ce..80f4d3acf 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1473,9 +1473,8 @@ static bool ggml_vk_matmul_shmem_support(const vk_device& device, const std::vec
     const uint32_t load_bufs = (warptile[1] + warptile[2]) * (warptile[3] + bank_conflict_offset) * type_size;
     const uint32_t mmid_row_ids = mul_mat_id ? (warptile[2] * 2 * sizeof(uint16_t)) : 0;
     const uint32_t coopmat_stage = device->coopmat_support ? warptile[7] * warptile[8] / warps * sizeof(float) : 0;
-    const uint32_t ballots_sh = mul_mat_id ? (warps * 4 * sizeof(uint32_t)) : 0;

-    const uint32_t total_size = load_bufs + mmid_row_ids + coopmat_stage + lut_size + ballots_sh;
+    const uint32_t total_size = load_bufs + mmid_row_ids + coopmat_stage + lut_size;
     const bool supported = total_size <= device->properties.limits.maxComputeSharedMemorySize;

     VK_LOG_DEBUG("ggml_vk_matmul_shmem_support(warptile=(" << warptile[0] << "," << warptile[1] << "," << warptile[2] << "), "
@@ -1540,10 +1539,7 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std:
     const uint32_t buf_b_size = BN * BK_STEP * block_b_size;
     const uint32_t mmid_row_ids = mul_mat_id ? (BN * 2u * (uint32_t)sizeof(uint16_t)) : 0u;

-    const uint32_t warps = warptile[0] / warptile[10];
-    const uint32_t ballots_sh = mul_mat_id ? (warps * 4u * (uint32_t)sizeof(uint32_t)) : 0u;
-
-    const uint32_t total_size = buf_a_size + buf_b_size + mmid_row_ids + ballots_sh + lut_size;
+    const uint32_t total_size = buf_a_size + buf_b_size + mmid_row_ids + lut_size;
     const bool supported = total_size <= device->properties.limits.maxComputeSharedMemorySize;

     VK_LOG_DEBUG("ggml_vk_matmul_int_shmem_support(warptile=(" << warptile[0] << "," << warptile[1] << "," << warptile[2] << "), "
@@ -1573,10 +1569,8 @@ static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const
             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;
@@ -1600,8 +1594,6 @@ static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const
     }
     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;
@@ -3595,7 +3587,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {

     ggml_vk_create_pipeline(device, device->pipeline_count_equal_i32, "count_equal_i32", count_equal_i32_len, count_equal_i32_data, "main", 3, sizeof(vk_op_push_constants), {512, 1, 1}, { device->subgroup_size }, 1);

-    if (device->subgroup_arithmetic && device->subgroup_require_full_support) {
+    if (device->subgroup_arithmetic && device->subgroup_vote && device->subgroup_require_full_support) {
         ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_subgroup_len, count_experts_subgroup_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true, true);
     } else {
         ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_len, count_experts_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true);
@@ -6014,20 +6006,27 @@ static uint32_t ggml_vk_guess_matmul_pipeline_align_map(ggml_backend_vk_context
     return configs[idx].align;
 }

+static uint64_t ggml_vk_mul_mat_id_max_tiles(uint64_t n_experts, uint64_t n_rows, uint32_t bn) {
+    // Each nonempty expert can add at most one partial tile to the combined grid.
+    return std::min(n_rows, CEIL_DIV(n_rows, bn) + std::min(n_experts, n_rows) - 1);
+}
+
 static void ggml_vk_matmul_id(
         ggml_backend_vk_context * ctx, vk_context& subctx, vk_pipeline& pipeline,
         vk_subbuffer&& a, vk_subbuffer&& b, vk_subbuffer&& d, vk_subbuffer&& ids, const vk_subbuffer & expert_count_buf,
         uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d,
         uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d,
         uint32_t n_as, uint32_t nei0, uint32_t nei1, uint32_t nbi1, uint32_t ne11,
-        bool hoist_row_ids) {
+        uint32_t row_ids_offset, uint32_t max_tiles) {
     VK_LOG_DEBUG("ggml_vk_matmul_id(a: (" << a.buffer->buffer << ", " << a.offset << ", " << a.size << "), b: (" << b.buffer->buffer << ", " << b.offset << ", " << b.size << "), d: (" << d.buffer->buffer << ", " << d.offset << ", " << d.size << "), ids: (" << ids.buffer->buffer << ", " << ids.offset << ", " << ids.size << "), expert_count: (" << expert_count_buf.buffer->buffer << ", " << expert_count_buf.offset << ", " << expert_count_buf.size << "), " <<
         "m: " << m << ", n: " << n << ", k: " << k << ", stride_a: " << stride_a << ", stride_b: " << stride_b << ", stride_d: " << stride_d << ", " <<
         "batch_stride_a: " << batch_stride_a << ", batch_stride_b: " << batch_stride_b << ", batch_stride_d: " << batch_stride_d << ", " <<
         "n_as: " << n_as << ", nei0: " << nei0 << ", nei1: " << nei1 << ", nbi1: " << nbi1 << ", ne11: " << ne11 << ")");
     const vk_mat_mat_id_push_constants pc = { m, n, k, stride_a, stride_b, stride_d, batch_stride_a, batch_stride_b, batch_stride_d,
-                                              nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids) };
-    ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, { m, nei1, n_as });
+                                              nei0, nei1, nbi1, ne11, n_as, row_ids_offset };
+    const uint32_t tiles_y = std::min(max_tiles, ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
+    const std::array<uint32_t, 3> elements = { m, tiles_y * pipeline->wg_denoms[1], uint32_t(CEIL_DIV(max_tiles, tiles_y)) };
+    ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, elements);
 }

 bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor) {
@@ -7336,7 +7335,6 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&

     const uint32_t nbi0 = ids->nb[0];
     const uint32_t nbi1 = ids->nb[1];
-    const uint32_t nbi2 = ids->nb[2];

     const uint64_t ne20 = dst->ne[0];
     const uint64_t ne21 = dst->ne[1];
@@ -7344,38 +7342,24 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
     // const uint64_t ne23 = dst->ne[3];

     const uint64_t n_as = ne02;
-    // n_as counts, n_as offsets, one total, then one packed row id per (expert, token).
-    // Hoisting requires 16-bit indices for the packing and a table that fits one binding.
-    const uint64_t hoisted_row_id_words = 2 * n_as + 1 + nei0 * nei1;
-    // 1024 matches MAX_EXPERTS in count_experts.comp and LLAMA_MAX_EXPERTS. It costs
-    // 3 * 1024 * 4 = 12 KiB of shared memory, within the 16 KiB Vulkan guarantees.
-    const bool hoist_row_ids = n_as <= 1024 && nei0 <= 0xffff && nei1 <= 0xffff &&
-                                hoisted_row_id_words * sizeof(uint32_t) <=
-                                    ctx->device->properties.limits.maxStorageBufferRange;

     ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)dst->buffer->context;
     ggml_backend_vk_buffer_context * src0_buf_ctx = (ggml_backend_vk_buffer_context *)src0->buffer->context;
     ggml_backend_vk_buffer_context * src1_buf_ctx = (ggml_backend_vk_buffer_context *)src1->buffer->context;
-    ggml_backend_vk_buffer_context * ids_buf_ctx = (ggml_backend_vk_buffer_context *)ids->buffer->context;

     vk_buffer d_Qx = nullptr;
     size_t qx_buf_offset = 0;
     vk_buffer d_Qy = nullptr;
     size_t qy_buf_offset = 0;
-    vk_buffer d_ids = nullptr;
-    size_t ids_buf_offset = 0;

     bool src0_uma = false;
     bool src1_uma = false;
-    bool ids_uma = false;

     if (ctx->device->uma) {
         ggml_vk_host_get(ctx->device, src0->data, d_Qx, qx_buf_offset);
         ggml_vk_host_get(ctx->device, src1->data, d_Qy, qy_buf_offset);
-        ggml_vk_host_get(ctx->device, ids->data, d_ids, ids_buf_offset);
         src0_uma = d_Qx != nullptr;
         src1_uma = d_Qy != nullptr;
-        ids_uma = d_ids != nullptr;
     }

     // Reformat and convert to fp16 if non-contiguous, or for coopmat2 for better perf
@@ -7445,6 +7429,11 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&

     vk_pipeline pipeline = ggml_vk_guess_matmul_pipeline_map(ctx, *mmp_map, ne01, nei1, aligned, true);

+    const uint32_t bn = pipeline->wg_denoms[1];
+    const uint64_t max_tiles = ggml_vk_mul_mat_id_max_tiles(n_as, nei0 * nei1, bn);
+    const uint64_t row_ids_offset = 1 + 2 * max_tiles;
+    const uint64_t row_map_words = row_ids_offset + max_tiles * bn;
+
     if (ggml_nbytes(src0) > ctx->device->properties.limits.maxStorageBufferRange) {
         pipeline = ggml_vk_get_64b_indexing_pipeline(ctx, pipeline);
     }
@@ -7456,7 +7445,6 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
     const uint64_t qy_sz = ggml_type_size(src1->type) * ggml_nelements(src1) / ggml_blck_size(src1->type);
     const uint64_t x_sz = !qx_needs_dequant ? qx_sz : sizeof(ggml_fp16_t) * x_ne;
     const uint64_t y_sz = quantize_y ? (ggml_vk_align_size(y_ne, 128) * ggml_type_size(GGML_TYPE_Q8_1) / ggml_blck_size(GGML_TYPE_Q8_1)) : (y_f32_kernel ? sizeof(float) * y_ne : sizeof(ggml_fp16_t) * y_ne);
-    const uint64_t ids_sz = nbi2;
     const uint64_t d_sz = sizeof(float) * d_ne;

     vk_pipeline to_fp16_vk_0 = nullptr;
@@ -7498,8 +7486,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
     }
     vk_pipeline count_experts = ctx->device->pipeline_count_experts;

-    const size_t expert_data_size = sizeof(uint32_t) *
-        (hoist_row_ids ? hoisted_row_id_words : n_as);
+    const size_t expert_data_size = sizeof(uint32_t) * row_map_words;

     {
         if (
@@ -7551,11 +7538,6 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
         qy_buf_offset = vk_tensor_offset(src1) + src1->view_offs;
         GGML_ASSERT(d_Qy != nullptr);
     }
-    if (!ids_uma) {
-        d_ids = ids_buf_ctx->dev_buffer;
-        ids_buf_offset = vk_tensor_offset(ids) + ids->view_offs;
-        GGML_ASSERT(d_ids != nullptr);
-    }
     if (qx_needs_dequant) {
         d_X = ctx->prealloc_x;
         GGML_ASSERT(d_X->size >= x_sz);
@@ -7581,6 +7563,8 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
             ggml_vk_sync_buffers(ctx, subctx);
         }
     }
+    vk_subbuffer d_ids = ggml_vk_tensor_subbuffer(ctx, ids, true);
+
     // Count how many times each expert is used
     vk_subbuffer expert_count_buf = { ctx->prealloc_split_k, 0, expert_data_size };
     if (ctx->prealloc_split_k_need_sync) {
@@ -7593,12 +7577,11 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
                                            (uint32_t)(nbi1 / ggml_type_size(ids->type)),
                                            (uint32_t)(get_misalign_bytes(ctx, ids) / ggml_type_size(ids->type)),
                                            (uint32_t)n_as,
-                                           uint32_t(hoist_row_ids),
-                                           0, 0 };
+                                           0, 0, bn, (uint32_t)row_ids_offset };
         init_pushconst_fastdiv(pc);
         ggml_vk_dispatch_pipeline(ctx, subctx, count_experts,
-            { vk_subbuffer{ d_ids, ids_buf_offset, ids_sz }, expert_count_buf }, pc,
-            { hoist_row_ids ? 1u : (uint32_t)n_as, 1, 1});
+            { d_ids, expert_count_buf }, pc,
+            { 1, 1, 1 });
     }

     if (x_non_contig) {
@@ -7671,10 +7654,10 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
     ggml_vk_matmul_id(
         ctx, subctx, pipeline,
         { d_X, x_buf_offset, x_range }, { d_Y, y_buf_offset, y_range },
-        { d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf,
+        { d_D, d_buf_offset, d_sz }, std::move(d_ids), expert_count_buf,
         ne01, ne21, ne10, ne10, stride_b_y, ne01,
         stride_batch_x, stride_batch_y, ne20*ne21,
-        n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids
+        n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, row_ids_offset, max_tiles
     );  // NOLINT

     if (x_non_contig || qx_needs_dequant) {
@@ -7923,11 +7906,14 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte
     }
 }

+static bool ggml_vk_use_mul_mat_vec_id(const ggml_tensor * dst) {
+    const ggml_tensor * src0 = dst->src[0];
+    const ggml_tensor * ids = dst->src[2];
+    return ids->ne[1] <= 8 && (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type));
+}
+
 bool ggml_vk_use_mul_mat_vec_id(const struct ggml_cgraph * cgraph, int node_idx) {
-    ggml_tensor * dst = cgraph->nodes[node_idx];
-    ggml_tensor * src0 = dst->src[0];
-    ggml_tensor * src2 = dst->src[2];
-    return (src2->ne[1] <= 8) && (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type));
+    return ggml_vk_use_mul_mat_vec_id(cgraph->nodes[node_idx]);
 }

 void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx) {
@@ -15431,6 +15417,13 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
             {
                 ggml_type src0_type = op->src[0]->type;
                 if (op->op == GGML_OP_MUL_MAT_ID) {
+                    if (!ggml_vk_use_mul_mat_vec_id(op)) {
+                        const ggml_tensor * ids = op->src[2];
+                        // Shaders store row IDs as uint16_t slot and token indices.
+                        if (ids->ne[0] > int64_t(UINT16_MAX) + 1 || ids->ne[1] > int64_t(UINT16_MAX) + 1) {
+                            return false;
+                        }
+                    }
                     if (!device->mul_mat_id_s[src0_type] && !device->mul_mat_id_m[src0_type] && !device->mul_mat_id_l[src0_type]) {
                         // If there's not enough shared memory for row_ids and the result tile, fallback to CPU
                         return false;
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp
index 06a50181c..e7e295115 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp
@@ -5,6 +5,7 @@
 #ifdef USE_SUBGROUPS
 #extension GL_KHR_shader_subgroup_basic : enable
 #extension GL_KHR_shader_subgroup_arithmetic : enable
+#extension GL_KHR_shader_subgroup_vote : enable
 #endif

 #include "types.glsl"
@@ -18,40 +19,47 @@ layout (push_constant) uniform parameter
     uint32_t nb01;
     uint32_t a_offset;
     uint32_t n_experts;
-    uint32_t hoist_row_ids;
     uint32_t ne00mp;
     uint32_t ne00L;
+    uint32_t row_tile_size;
+    uint32_t row_ids_offset;
 } p;

 #define BLOCK_SIZE 256
+#define EXPERTS_PER_CHUNK 1024

 layout(local_size_x = BLOCK_SIZE, local_size_y = 1, local_size_z = 1) in;

 layout (binding = 0) readonly buffer A {uint data_a[];};
 layout (binding = 1) writeonly buffer D {uint data_d[];};

-// Upper bound on n_experts for the hoisted row-id path. Must match the limit in
-// ggml_vk_mul_mat_id_q_f16 (hoist_row_ids). The non-hoisted reduction below only
-// needs BLOCK_SIZE entries.
-#define MAX_EXPERTS 1024
-
-shared uint vals[MAX_EXPERTS];
-shared uint offsets[MAX_EXPERTS];
-shared uint cursors[MAX_EXPERTS];
+shared uint vals[EXPERTS_PER_CHUNK];
+shared uint offsets[EXPERTS_PER_CHUNK];
+shared uint cursors[EXPERTS_PER_CHUNK];
+
+// data_d[0] is the tile count. Each tile has (expert, valid rows).
+// Packed (token << 16) | slot row ids start at p.row_ids_offset, with each expert on a tile boundary.
+void write_tiles(uint expert, uint count, uint tile_begin, uint n_tiles) {
+    for (uint t = 0; t < n_tiles; ++t) {
+        const uint descriptor = 1 + 2 * (tile_begin + t);
+        data_d[descriptor] = expert;
+        data_d[descriptor + 1] = min(p.row_tile_size, count - t * p.row_tile_size);
+    }
+}

-// data_d layout when p.hoist_row_ids is set:
-//   [0,              n_experts)   per-expert row count
-//   [n_experts,    2*n_experts)   per-expert start offset into the row id region
-//   [2*n_experts]                 total row count
-//   [2*n_experts + 1,         )   row ids grouped by expert, packed as (i01 << 16) | (i00 & 0xffff)
-// Otherwise only data_d[expert_id] is written, holding that expert's row count.
 void main() {
-    const uint expert_id = gl_WorkGroupID.x;
     const uint num_elements = p.ne00 * p.ne01;
     const uint tid = gl_LocalInvocationID.x;
+    const uint tile_shift = findLSB(p.row_tile_size);
+    uint total_tiles = 0;
+#ifdef USE_SUBGROUPS
+    // Use the subgroup that contains invocation 0.
+    const bool prefix_subgroup = subgroupAny(tid == 0);
+#endif

-    if (p.hoist_row_ids != 0) {
-        for (uint e = tid; e < p.n_experts; e += BLOCK_SIZE) {
+    for (uint expert_base = 0; expert_base < p.n_experts;) {
+        const uint n_experts = min(EXPERTS_PER_CHUNK, p.n_experts - expert_base);
+        for (uint e = tid; e < n_experts; e += BLOCK_SIZE) {
             vals[e] = 0;
         }
         barrier();
@@ -59,46 +67,43 @@ void main() {
         for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) {
             const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L);
             const uint i00 = idx - i01 * p.ne00;
-            const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00];
-            if (expert < p.n_experts) {
+            const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00] - expert_base;
+            if (expert < n_experts) {
                 atomicAdd(vals[expert], 1);
             }
         }
         barrier();

 #ifdef USE_SUBGROUPS
-        if (gl_SubgroupID == 0) {
-            // pad the trip count so the subgroup ops stay in uniform control flow
-            const uint n_experts_padded = (p.n_experts + gl_SubgroupSize - 1) & ~(gl_SubgroupSize - 1);
-            uint base = 0;
-            for (uint expert = gl_SubgroupInvocationID; expert < n_experts_padded; expert += gl_SubgroupSize) {
-                const bool in_range = expert < p.n_experts;
-                const uint count = in_range ? vals[expert] : 0;
-                const uint offset = base + subgroupExclusiveAdd(count);
+        if (prefix_subgroup) {
+            const uint n_experts_padded = (n_experts + gl_SubgroupSize - 1) & ~(gl_SubgroupSize - 1);
+            uint base = subgroupAdd(tid == 0 ? total_tiles : 0u);
+            for (uint e = gl_SubgroupInvocationID; e < n_experts_padded; e += gl_SubgroupSize) {
+                const bool in_range = e < n_experts;
+                const uint count = in_range ? vals[e] : 0;
+                const uint n_tiles = (count + p.row_tile_size - 1) >> tile_shift;
+                const uint tile_begin = base + subgroupExclusiveAdd(n_tiles);
                 if (in_range) {
-                    data_d[expert] = count;
-                    data_d[p.n_experts + expert] = offset;
-                    offsets[expert] = offset;
-                    cursors[expert] = 0;
+                    offsets[e] = tile_begin * p.row_tile_size;
+                    cursors[e] = 0;
+                    write_tiles(expert_base + e, count, tile_begin, n_tiles);
                 }
-                base += subgroupAdd(count);
+                base += subgroupAdd(n_tiles);
             }
-            if (subgroupElect()) {
-                data_d[2 * p.n_experts] = base;
+            if (tid == 0) {
+                total_tiles = base;
             }
         }
 #else
         if (tid == 0) {
-            uint offset = 0;
-            for (uint expert = 0; expert < p.n_experts; ++expert) {
-                const uint count = vals[expert];
-                data_d[expert] = count;
-                data_d[p.n_experts + expert] = offset;
-                offsets[expert] = offset;
-                cursors[expert] = 0;
-                offset += count;
+            for (uint e = 0; e < n_experts; ++e) {
+                const uint count = vals[e];
+                const uint n_tiles = (count + p.row_tile_size - 1) >> tile_shift;
+                offsets[e] = total_tiles * p.row_tile_size;
+                cursors[e] = 0;
+                write_tiles(expert_base + e, count, total_tiles, n_tiles);
+                total_tiles += n_tiles;
             }
-            data_d[2 * p.n_experts] = offset;
         }
 #endif
         barrier();
@@ -106,35 +111,16 @@ void main() {
         for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) {
             const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L);
             const uint i00 = idx - i01 * p.ne00;
-            const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00];
-            if (expert < p.n_experts) {
+            const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00] - expert_base;
+            if (expert < n_experts) {
                 const uint row = atomicAdd(cursors[expert], 1);
-                const uint packed_row_id = (i01 << 16) | (i00 & 0xffffu);
-                data_d[2 * p.n_experts + 1 + offsets[expert] + row] = packed_row_id;
+                data_d[p.row_ids_offset + offsets[expert] + row] = (i01 << 16) | i00;
             }
         }
-        return;
+        // The next chunk clears vals. Scatter uses only offsets and cursors.
+        expert_base += n_experts;
     }
-
-    uint count = 0;
-    for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) {
-        const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L);
-        const uint i00 = idx - i01 * p.ne00;
-        const uint a = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00];
-
-        count += uint(a == expert_id);
-    }
-
-    vals[tid] = count;
-    barrier();
-    [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
-        if (tid < s) {
-            vals[tid] += vals[tid + s];
-        }
-        barrier();
-    }
-
     if (tid == 0) {
-        data_d[expert_id] = vals[0];
+        data_d[0] = total_tiles;
     }
 }
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
index 11098ee7b..9f31ef126 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
@@ -124,7 +124,7 @@ 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[];};
+layout (binding = 4) readonly buffer RowMap {uint data_row_map[];};
 #endif

 layout (push_constant) uniform parameter
@@ -146,7 +146,7 @@ layout (push_constant) uniform parameter
     uint nbi1;
     uint ne11;
     uint n_experts;
-    uint hoist_row_ids;
+    uint row_ids_offset;
 #else
     uint base_work_group_z;
     uint num_batches;
@@ -208,13 +208,17 @@ shared ACC_TYPE coopmat_stage[TM * TN * NUM_WARPS];
 #endif

 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]) {
+    const uint tile_idx = gl_WorkGroupID.y + gl_WorkGroupID.z * gl_NumWorkGroups.y;
+    if (tile_idx >= data_row_map[0]) {
         return;
     }
+    const uint expert_idx = data_row_map[1 + 2 * tile_idx];
+    const uint row_begin = tile_idx * BN;
+    _ne1 = data_row_map[2 + 2 * tile_idx];
+    const uint ic = 0;
+#else
+    const uint ic = gl_WorkGroupID.y;
 #endif
 #if defined(NEEDS_INIT_IQ_SHMEM) || defined(MULMAT_QUANT)
     init_iq_shmem(gl_WorkGroupSize);
@@ -285,37 +289,9 @@ void main() {
     const uint loadstride_b = gl_WorkGroupSize.x * LOAD_VEC_B_EFF * LOAD_VEC_BATCH_B / BK;

 #ifdef MUL_MAT_ID
-    if (p.hoist_row_ids != 0) {
-        load_row_ids_hoisted(expert_idx, ic);
-    } else {
-#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;
-
-    uint required_work_items = (_ne1 - ic * BN) * BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B;
-    uint required_warp_c = (_ne1 - ic * BN + WN - 1) / WN;
+    load_row_ids(row_begin);
+    uint required_work_items = _ne1 * BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B;
+    uint required_warp_c = (_ne1 + WN - 1) / WN;
 #endif

 #ifdef MUL_MAT_ID
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp
index 9b59e8bca..e3f1dd5ab 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp
@@ -82,7 +82,7 @@ layout (push_constant) uniform parameter
     uint nbi1;
     uint ne11;
     uint n_experts;
-    uint hoist_row_ids;
+    uint row_ids_offset;
 #else
     uint base_work_group_z;
     uint num_batches;
@@ -182,7 +182,7 @@ f16vec4 mmDecodeA_v(const in decodeBufA bl_in, const in uint blockCoords[2], con

 #ifdef MUL_MAT_ID
 layout (binding = 3) readonly buffer IDS {int data_ids[];};
-layout (binding = 4) readonly buffer Counts {int data_expert_count[];};
+layout (binding = 4) readonly buffer RowMap {uint data_row_map[];};

 shared u16vec4 row_ids[BN];

@@ -191,7 +191,6 @@ layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufB {
 };

 uint _ne1;
-shared uvec4 ballots_sh[BLOCK_SIZE / subgroup_size];

 B_TYPE decodeFuncB(const in decodeBufB bl, const in uint blockCoords[2], const in uint coordInBlock[2])
 {
@@ -238,82 +237,9 @@ D_TYPE perElemOpD(const in uint32_t r, const in uint32_t c, const in D_TYPE elem
     return elem;
 }

-void load_row_ids(uint expert_idx, bool nei0_is_pow2, uint ic) {
-    _ne1 = 0;
-    uint num_elements = p.nei1 * p.nei0;
-    uint nei0shift = findLSB(p.nei0);
-
-    uint ids[16];
-    uint iter = 0;
-
-    uint expert_count = data_expert_count[expert_idx];
-
-    for (uint j = 0; j < num_elements; j += BLOCK_SIZE) {
-        // prefetch up to 16 elements
-        if (iter == 0) {
-            [[unroll]] for (uint k = 0; k < 16; ++k) {
-                uint i = j + gl_LocalInvocationIndex + k*BLOCK_SIZE;
-                bool in_range = i < num_elements;
-                uint ii1;
-                if (nei0_is_pow2) {
-                    ii1 = i >> nei0shift;
-                } else {
-                    ii1 = i / p.nei0;
-                }
-                uint ii0 = i - ii1 * p.nei0;
-                ids[k] = in_range ? data_ids[ii1*p.nbi1 + ii0] : 0;
-            }
-        }
-        uint i = j + gl_LocalInvocationIndex;
-        bool in_range = i < num_elements;
-        uint ii1;
-        if (nei0_is_pow2) {
-            ii1 = i >> nei0shift;
-        } else {
-            ii1 = i / p.nei0;
-        }
-        uint ii0 = i - ii1 * p.nei0;
-        uint id = ids[iter++];
-        uvec4 ballot = subgroupBallot(in_range && id == expert_idx);
-
-        if (gl_SubgroupInvocationID == 0) {
-            ballots_sh[gl_SubgroupID] = ballot;
-        }
-        barrier();
-
-        uint subgroup_base = 0;
-        uint total = 0;
-        for (uint k = 0; k < gl_NumSubgroups; ++k) {
-            if (k == gl_SubgroupID) {
-                subgroup_base = total;
-            }
-            total += subgroupBallotBitCount(ballots_sh[k]);
-        }
-        barrier();
-
-        uint idx = subgroup_base + subgroupBallotExclusiveBitCount(ballot);
-        if (in_range && id == expert_idx && _ne1 + idx >= ic * BN && _ne1 + idx < (ic + 1) * BN) {
-            row_ids[_ne1 + idx - ic * BN] = u16vec4(fastmod(ii0, p.ne11), ii1, ii0, 0);
-        }
-        _ne1 += total;
-        iter &= 15;
-        if (_ne1 >= (ic + 1) * BN || _ne1 == expert_count) {
-            break;
-        }
-    }
-    barrier();
-}
-
-void load_row_ids_hoisted(uint expert_idx, uint ic) {
-    _ne1 = uint(data_expert_count[expert_idx]);
-
-    const uint tile_begin = ic * BN;
-    const uint tile_count = tile_begin < _ne1 ? min(BN, _ne1 - tile_begin) : 0;
-    const uint expert_offset = uint(data_expert_count[p.n_experts + expert_idx]);
-    const uint row_ids_offset = 2 * p.n_experts + 1 + expert_offset + tile_begin;
-
-    for (uint i = gl_LocalInvocationIndex; i < tile_count; i += BLOCK_SIZE) {
-        const uint packed_row_id = uint(data_expert_count[row_ids_offset + i]);
+void load_row_ids(uint row_begin) {
+    for (uint i = gl_LocalInvocationIndex; i < _ne1; i += BLOCK_SIZE) {
+        const uint packed_row_id = data_row_map[p.row_ids_offset + row_begin + i];
         const uint ii0 = packed_row_id & 0xffffu;
         const uint ii1 = packed_row_id >> 16;
         row_ids[i] = u16vec4(fastmod(ii0, p.ne11), ii1, ii0, 0);
@@ -328,13 +254,20 @@ void load_row_ids_hoisted(uint expert_idx, uint ic) {

 void main() {
     const uint tid = gl_LocalInvocationIndex;
-    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]) {
+    const uint tile_idx = gl_WorkGroupID.y + gl_WorkGroupID.z * gl_NumWorkGroups.y;
+    if (tile_idx >= data_row_map[0]) {
         return;
     }
+    const uint expert_idx = data_row_map[1 + 2 * tile_idx];
+    const uint row_begin = tile_idx * BN;
+    _ne1 = data_row_map[2 + 2 * tile_idx];
+    const uint ic = 0;
+#else
+    const uint ic = gl_WorkGroupID.y;
+#endif
+
+#ifdef MUL_MAT_ID
     // initialize to row 0 so we don't need to bounds check
     if (tid < BN) {
         row_ids[tid] = u16vec4(0);
@@ -365,16 +298,7 @@ void main() {
     const uint ik = gl_WorkGroupID.x / blocks_m;

 #ifdef MUL_MAT_ID
-    if (p.hoist_row_ids != 0) {
-        load_row_ids_hoisted(expert_idx, ic);
-    } else if (bitCount(p.nei0) == 1) {
-        load_row_ids(expert_idx, true, ic);
-    } else {
-        load_row_ids(expert_idx, false, ic);
-    }
-
-    // Workgroup has no work
-    if (ic * BN >= _ne1) return;
+    load_row_ids(row_begin);
 #endif

 #ifdef MUL_MAT_ID
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl
index 54ad60b2e..99a644e5a 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl
@@ -2,86 +2,9 @@
 shared u16vec2 row_ids[BN];
 uint _ne1;

-#ifdef MUL_MAT_ID_USE_SUBGROUPS
-shared uvec4 ballots_sh[NUM_WARPS];
-
-void load_row_ids(uint expert_idx, bool nei0_is_pow2, uint ic) {
-    _ne1 = 0;
-    uint num_elements = p.nei1 * p.nei0;
-    uint nei0shift = findLSB(p.nei0);
-
-    uint ids[16];
-    uint iter = 0;
-
-    uint expert_count = data_expert_count[expert_idx];
-
-    for (uint j = 0; j < num_elements; j += BLOCK_SIZE) {
-        // prefetch up to 16 elements
-        if (iter == 0) {
-            [[unroll]] for (uint k = 0; k < 16; ++k) {
-                uint i = j + gl_LocalInvocationIndex + k*BLOCK_SIZE;
-                bool in_range = i < num_elements;
-                uint ii1;
-                if (nei0_is_pow2) {
-                    ii1 = i >> nei0shift;
-                } else {
-                    ii1 = i / p.nei0;
-                }
-                uint ii0 = i - ii1 * p.nei0;
-                ids[k] = in_range ? data_ids[ii1*p.nbi1 + ii0] : 0;
-            }
-        }
-        uint i = j + gl_LocalInvocationIndex;
-        bool in_range = i < num_elements;
-        uint ii1;
-        if (nei0_is_pow2) {
-            ii1 = i >> nei0shift;
-        } else {
-            ii1 = i / p.nei0;
-        }
-        uint ii0 = i - ii1 * p.nei0;
-        uint id = ids[iter++];
-        uvec4 ballot = subgroupBallot(in_range && id == expert_idx);
-
-        if (gl_SubgroupInvocationID == 0) {
-            ballots_sh[gl_SubgroupID] = ballot;
-        }
-        barrier();
-
-        uint subgroup_base = 0;
-        uint total = 0;
-        for (uint k = 0; k < gl_NumSubgroups; ++k) {
-            if (k == gl_SubgroupID) {
-                subgroup_base = total;
-            }
-            total += subgroupBallotBitCount(ballots_sh[k]);
-        }
-        barrier();
-
-        uint idx = subgroup_base + subgroupBallotExclusiveBitCount(ballot);
-        if (in_range && id == expert_idx && _ne1 + idx >= ic * BN && _ne1 + idx < (ic + 1) * BN) {
-            row_ids[_ne1 + idx - ic * BN] = u16vec2(ii0, ii1);
-        }
-        _ne1 += total;
-        iter &= 15;
-        if (_ne1 >= (ic + 1) * BN || _ne1 == expert_count) {
-            break;
-        }
-    }
-    barrier();
-}
-#endif // MUL_MAT_ID_USE_SUBGROUPS
-
-void load_row_ids_hoisted(uint expert_idx, uint ic) {
-    _ne1 = uint(data_expert_count[expert_idx]);
-
-    const uint tile_begin = ic * BN;
-    const uint tile_count = tile_begin < _ne1 ? min(BN, _ne1 - tile_begin) : 0;
-    const uint expert_offset = uint(data_expert_count[p.n_experts + expert_idx]);
-    const uint row_ids_offset = 2 * p.n_experts + 1 + expert_offset + tile_begin;
-
-    for (uint i = gl_LocalInvocationIndex; i < tile_count; i += BLOCK_SIZE) {
-        const uint packed_row_id = uint(data_expert_count[row_ids_offset + i]);
+void load_row_ids(uint row_begin) {
+    for (uint i = gl_LocalInvocationIndex; i < _ne1; i += BLOCK_SIZE) {
+        const uint packed_row_id = data_row_map[p.row_ids_offset + row_begin + i];
         row_ids[i] = u16vec2(packed_row_id & 0xffffu, packed_row_id >> 16);
     }
     barrier();
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp
index 67908e704..2819eb46e 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp
@@ -37,7 +37,7 @@ 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[];};
+layout (binding = 4) readonly buffer RowMap {uint data_row_map[];};
 #endif

 layout (push_constant) uniform parameter
@@ -59,7 +59,7 @@ layout (push_constant) uniform parameter
     uint nbi1;
     uint ne11;
     uint n_experts;
-    uint hoist_row_ids;
+    uint row_ids_offset;
 #else
     uint base_work_group_z;
     uint num_batches;
@@ -105,19 +105,22 @@ block_b_cache cache_b;
 #define LOAD_VEC_A (4 * QUANT_R_MMQ)
 #define LOAD_VEC_B 16

-#define NUM_WARPS (BLOCK_SIZE / WARP)

 #include "mul_mm_id_funcs.glsl"
 #include "mul_mmq_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]) {
+    const uint tile_idx = gl_WorkGroupID.y + gl_WorkGroupID.z * gl_NumWorkGroups.y;
+    if (tile_idx >= data_row_map[0]) {
         return;
     }
+    const uint expert_idx = data_row_map[1 + 2 * tile_idx];
+    const uint row_begin = tile_idx * BN;
+    _ne1 = data_row_map[2 + 2 * tile_idx];
+    const uint ic = 0;
+#else
+    const uint ic = gl_WorkGroupID.y;
 #endif
 #ifdef NEEDS_INIT_IQ_SHMEM
     init_iq_shmem(gl_WorkGroupSize);
@@ -161,34 +164,7 @@ void main() {
     const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK;

 #ifdef MUL_MAT_ID
-    if (p.hoist_row_ids != 0) {
-        load_row_ids_hoisted(expert_idx, ic);
-    } else {
-#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;
+    load_row_ids(row_begin);
 #endif

 #ifdef MUL_MAT_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 7cab9a119..17280d17b 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
@@ -39,7 +39,7 @@ 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[];};
+layout (binding = 4) readonly buffer RowMap {uint data_row_map[];};
 #endif

 layout (push_constant) uniform parameter
@@ -61,7 +61,7 @@ layout (push_constant) uniform parameter
     uint nbi1;
     uint ne11;
     uint n_experts;
-    uint hoist_row_ids;
+    uint row_ids_offset;
 #else
     uint base_work_group_z;
     uint num_batches;
@@ -142,13 +142,22 @@ ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a,
 }

 #ifdef MUL_MAT_ID
-#define NUM_WARPS (BLOCK_SIZE / WARP)
 #include "mul_mm_id_funcs.glsl"
 #endif

 #include "mul_mmq_cm1_funcs.glsl"

 void main() {
+#ifdef MUL_MAT_ID
+    const uint tile_idx = gl_WorkGroupID.y + gl_WorkGroupID.z * gl_NumWorkGroups.y;
+    if (tile_idx >= data_row_map[0]) {
+        return;
+    }
+    const uint expert_idx = data_row_map[1 + 2 * tile_idx];
+    const uint row_begin = tile_idx * BN;
+    _ne1 = data_row_map[2 + 2 * tile_idx];
+    const uint ic = 0;
+#endif
 #if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS)
     if (gl_LocalInvocationIndex < 16u) {
         cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex];
@@ -175,12 +184,7 @@ void main() {
     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;
@@ -231,33 +235,7 @@ void main() {
     const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK;

 #ifdef MUL_MAT_ID
-    if (p.hoist_row_ids != 0) {
-        load_row_ids_hoisted(expert_idx, ic);
-    } else {
-#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
-    }
-
-    if (ic * BN >= _ne1) return;
+    load_row_ids(row_begin);
 #endif

 #ifdef MUL_MAT_ID
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 4e3c6942b..bd599e99e 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -5440,9 +5440,10 @@ struct test_mul_mat_id : public test_case {
     const int64_t k;
     const float amax; // magnitude of src1
     const int64_t m_v; // rows of as in memory, the experts of as are strided for m_v > m, no view for m_v == 0
+    const bool ids_offset;

     std::string vars() override {
-        return VARS_TO_STR10(type_a, type_b, n_mats, n_used, b, m, n, k, amax, m_v);
+        return VARS_TO_STR11(type_a, type_b, n_mats, n_used, b, m, n, k, amax, m_v, ids_offset);
     }

     double max_nmse_err() override {
@@ -5467,9 +5468,9 @@ struct test_mul_mat_id : public test_case {
     test_mul_mat_id(ggml_type type_a = GGML_TYPE_F32, ggml_type type_b = GGML_TYPE_F32,
             int n_mats = 8, int n_used = 2, bool b = false,
             int64_t m = 32, int64_t n = 32, int64_t k = 32,
-            float amax = 1.0f, int64_t m_v = 0)
+            float amax = 1.0f, int64_t m_v = 0, bool ids_offset = false)
         : type_a(type_a), type_b(type_b), n_mats(n_mats), n_used(n_used), b(b),
-            m(m), n(n), k(k), amax(amax), m_v(m_v) {
+            m(m), n(n), k(k), amax(amax), m_v(m_v), ids_offset(ids_offset) {
             GGML_ASSERT(n_used <= n_mats);
             GGML_ASSERT(m_v == 0 || m_v > m);
         }
@@ -5482,10 +5483,10 @@ struct test_mul_mat_id : public test_case {
         }
         ggml_set_name(as, "as");

-        ggml_tensor * ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, n_mats, n);
+        ggml_tensor * ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, n_mats + int(ids_offset), n);
         ggml_set_name(ids, "ids");
-        if (n_used != n_mats) {
-            ids = ggml_view_2d(ctx, ids, n_used, n, ids->nb[1], 0);
+        if (n_used != n_mats || ids_offset) {
+            ids = ggml_view_2d(ctx, ids, n_used, n, ids->nb[1], ids_offset ? sizeof(int32_t) : 0);
             ggml_set_name(ids, "view_of_ids");
         }

@@ -5512,6 +5513,32 @@ struct test_mul_mat_id : public test_case {
     }
 };

+// MUL_MAT_ID with expert ids repeated within a token's row
+struct test_mul_mat_id_dup : public test_mul_mat_id {
+    using test_mul_mat_id::test_mul_mat_id;
+    std::string vars() override { return test_mul_mat_id::vars() + ",dup=1"; }
+    void initialize_tensors(ggml_context * ctx) override {
+        std::default_random_engine rng(1234);
+        for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
+            if (t->type == GGML_TYPE_I32) {
+                if (ggml_is_view_op(t->op)) { continue; }
+                for (int64_t r = 0; r < ggml_nrows(t); r++) {
+                    std::vector<int32_t> data(t->ne[0]);
+                    for (int i = 0; i < t->ne[0]; i++) {
+                        data[i] = (rng() % 4 == 0) ? (int32_t) (rng() % n_mats) : 0;   // repeated ids, mostly expert 0
+                    }
+                    if (r == ggml_nrows(t) - 1) {
+                        data[ids_offset ? 1 : 0] = n_mats - 1;
+                    }
+                    ggml_backend_tensor_set(t, data.data(), r * t->nb[1], t->ne[0] * sizeof(int32_t));
+                }
+            } else {
+                init_tensor_uniform(t);
+            }
+        }
+    }
+};
+
 // FP4 W4A8 path on the MoE path (GGML_PREC_Q8 on src1 disallows 4-bit activations)
 struct test_mul_mat_id_w4a8 : public test_mul_mat_id {
     test_mul_mat_id_w4a8(ggml_type type_a = GGML_TYPE_NVFP4, ggml_type type_b = GGML_TYPE_F32,
@@ -10972,6 +10999,18 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_TQ1_0, GGML_TYPE_F32, 28, 10, false, 1024, 1, 4096));
     test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_TQ1_0, GGML_TYPE_F32, 128, 8, false, 1024, 1, 2048));

+    // repeated expert ids within a row
+    for (ggml_type ta : {GGML_TYPE_F16, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0}) {
+        for (int n : {9, 16, 33, 64}) {
+            test_cases.emplace_back(new test_mul_mat_id_dup(ta, GGML_TYPE_F32, 28, 10, false, 1024, n, 256));
+        }
+    }
+
+    for (ggml_type ta : {GGML_TYPE_F16, GGML_TYPE_Q4_0}) {
+        test_cases.emplace_back(new test_mul_mat_id_dup(ta, GGML_TYPE_F32, 1025, 10, false, 64, 33, 256, 1.0f, 0, true));
+        test_cases.emplace_back(new test_mul_mat_id_dup(ta, GGML_TYPE_F32, 2050, 10, true, 64, 33, 256));
+    }
+
     for (ggml_type type_a : all_types) {
         test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 4, 2, false, 64, 16, 3*ggml_blck_size(type_a)));
     }