Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 37 additions & 11 deletions ggml/src/ggml-opencl/ggml-opencl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -891,6 +891,7 @@ struct ggml_backend_opencl_context {
cl_kernel kernel_gemm_moe_q4_0_q8_1_dp4a = nullptr; // dp4a (int8) q4_0 MoE prefill GEMM
cl_kernel kernel_moe_reorder_b;
cl_kernel kernel_moe_histogram, kernel_moe_scan, kernel_moe_fill, kernel_moe_scatter;
cl_kernel kernel_moe_scatter_stable = nullptr; // deterministic slot assignment
cl_kernel kernel_moe_combine_f32 = nullptr; // fused router-weight mul + cross-expert sum
cl_kernel kernel_mul_mv_id_q4_0_f32_8x_flat;
cl_kernel kernel_mul_mv_id_q8_0_f32, kernel_mul_mv_id_q8_0_f32_flat;
Expand Down Expand Up @@ -4441,6 +4442,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
CL_CHECK((backend_ctx->kernel_moe_scan = clCreateKernel(prog, "kernel_moe_scan", &err), err));
CL_CHECK((backend_ctx->kernel_moe_fill = clCreateKernel(prog, "kernel_moe_fill", &err), err));
CL_CHECK((backend_ctx->kernel_moe_scatter = clCreateKernel(prog, "kernel_moe_scatter", &err), err));
CL_CHECK((backend_ctx->kernel_moe_scatter_stable = clCreateKernel(prog, "kernel_moe_scatter_stable", &err), err));
CL_CHECK(clReleaseProgram(prog));
GGML_LOG_CONT(".");
}
Expand Down Expand Up @@ -20678,18 +20680,42 @@ static void moe_router_reoerder(ggml_backend_t backend, const ggml_tensor * src,
size_t fill_local_size[] = {64, 1, 1};
backend_ctx->enqueue_ndrange_kernel(kernel, 3, fill_global_size, fill_local_size, src);

// Scatter
kernel = backend_ctx->kernel_moe_scatter;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &original_router_buf));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &post_router_buf));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &emap_buf));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &tile_offset_buf));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &slot_counter_buf));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(int), &ne21));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne20));
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(int), &ne02));
// Scatter. The deterministic variant is the default: kernel_moe_scatter derives
// each token's slot from an atomic counter, so the packing inside an expert - and
// with it the output of the ragged prefill GEMM - changes from run to run. Set
// GGML_OPENCL_MOE_STABLE_SCATTER=0 to restore the atomic version.
static const bool stable_scatter = []{
const char * e = getenv("GGML_OPENCL_MOE_STABLE_SCATTER");
return !e || e[0] == '\0' || e[0] != '0';
}();

backend_ctx->enqueue_ndrange_kernel(kernel, 3, histogram_global_size, histogram_local_size, src);
if (stable_scatter) {
kernel = backend_ctx->kernel_moe_scatter_stable;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &original_router_buf));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &post_router_buf));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &emap_buf));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &tile_offset_buf));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(int), &ne21));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(int), &ne20));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne02));

// one workgroup (one wave) per expert; each ranks its own tokens
size_t scatter_global_size[] = {64, (size_t)ne02};
size_t scatter_local_size[] = {64, 1};
backend_ctx->enqueue_ndrange_kernel(kernel, 2, scatter_global_size, scatter_local_size, src);
} else {
kernel = backend_ctx->kernel_moe_scatter;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &original_router_buf));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &post_router_buf));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &emap_buf));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &tile_offset_buf));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &slot_counter_buf));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(int), &ne21));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne20));
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(int), &ne02));

backend_ctx->enqueue_ndrange_kernel(kernel, 3, histogram_global_size, histogram_local_size, src);
}

// [MOE_TILES] env-gated padding probe: read back total_tiles (= Sum_e
// ceil(k_e/n_tile_size)) and compare to the ideal tile count for the real
Expand Down
73 changes: 73 additions & 0 deletions ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,79 @@ __kernel void kernel_moe_scatter(
emap[tile_idx] = val;
}

// Deterministic replacement for kernel_moe_scatter.
//
// kernel_moe_scatter takes each token's slot from atomic_inc(slot_counter[expert]),
// so the token -> slot packing inside an expert depends on which work-item wins the
// atomic and changes from run to run. The ragged prefill GEMM path is sensitive to
// that packing (the non-ragged path is not, since its padded slots alias slot 0 and
// are overwritten last), which makes MoE prompt processing non-reproducible: the same
// binary on the same prompt returns one of several outputs.
//
// Here the slot is the token's rank in flat (n, k) order among the tokens routed to
// the same expert - a fixed function of the routing input. One workgroup per expert
// walks the flat routing list in blocks of 64 and ranks its own tokens with a
// workgroup scan, carrying a running count between blocks. Cost is one pass over the
// routing list per expert; the list is a few KiB and stays in cache.
__kernel void kernel_moe_scatter_stable(
__global const int * input,
__global int * post_router,
__global ushort * emap,
__global const int * tile_offset,
int N,
int topK,
uint n_experts
) {
const int e = get_group_id(1);
const int lid = get_local_id(0);
const int M = N * topK;

__local int scan[64];
__local int running;

if (lid == 0) {
running = 0;
}
barrier(CLK_LOCAL_MEM_FENCE);

for (int base = 0; base < M; base += 64) {
const int j = base + lid;

int pred = 0;
if (j < M) {
const int n = j / topK;
const int k = j - n * topK;
pred = (input[n * (int)n_experts + k] == e) ? 1 : 0;
}

scan[lid] = pred;
barrier(CLK_LOCAL_MEM_FENCE);

// Hillis-Steele inclusive scan over the 64 lanes
for (int off = 1; off < 64; off <<= 1) {
int add = (lid >= off) ? scan[lid - off] : 0;
barrier(CLK_LOCAL_MEM_FENCE);
scan[lid] += add;
barrier(CLK_LOCAL_MEM_FENCE);
}

if (pred) {
const int local_slot = running + (scan[lid] - 1); // exclusive rank
const int tile_idx = tile_offset[e] + (local_slot >> 5);
const int lane = local_slot & 31;

post_router[tile_idx * 32 + lane] = j;
emap[tile_idx] = (ushort)e;
}

barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 63) {
running += scan[63];
}
barrier(CLK_LOCAL_MEM_FENCE);
}
}

__kernel void kernel_moe_fill(
__global int * post_router,
__global int * total_tiles,
Expand Down