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)));
}