opencl: make the MoE expert scatter deterministic (#26464)
This commit is contained in:
@@ -895,6 +895,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_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_reorder_b;
|
||||||
cl_kernel kernel_moe_histogram, kernel_moe_scan, kernel_moe_fill, kernel_moe_scatter;
|
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_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_q4_0_f32_8x_flat;
|
||||||
cl_kernel kernel_mul_mv_id_q8_0_f32, kernel_mul_mv_id_q8_0_f32_flat;
|
cl_kernel kernel_mul_mv_id_q8_0_f32, kernel_mul_mv_id_q8_0_f32_flat;
|
||||||
@@ -4463,6 +4464,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_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_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 = 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));
|
CL_CHECK(clReleaseProgram(prog));
|
||||||
GGML_LOG_CONT(".");
|
GGML_LOG_CONT(".");
|
||||||
}
|
}
|
||||||
@@ -20863,18 +20865,42 @@ static void moe_router_reoerder(ggml_backend_t backend, const ggml_tensor * src,
|
|||||||
size_t fill_local_size[] = {64, 1, 1};
|
size_t fill_local_size[] = {64, 1, 1};
|
||||||
backend_ctx->enqueue_ndrange_kernel(kernel, 3, fill_global_size, fill_local_size, src);
|
backend_ctx->enqueue_ndrange_kernel(kernel, 3, fill_global_size, fill_local_size, src);
|
||||||
|
|
||||||
// Scatter
|
// Scatter. The deterministic variant is the default: kernel_moe_scatter derives
|
||||||
kernel = backend_ctx->kernel_moe_scatter;
|
// each token's slot from an atomic counter, so the packing inside an expert - and
|
||||||
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &original_router_buf));
|
// with it the output of the ragged prefill GEMM - changes from run to run. Set
|
||||||
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &post_router_buf));
|
// GGML_OPENCL_MOE_STABLE_SCATTER=0 to restore the atomic version.
|
||||||
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &emap_buf));
|
static const bool stable_scatter = []{
|
||||||
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &tile_offset_buf));
|
const char * e = getenv("GGML_OPENCL_MOE_STABLE_SCATTER");
|
||||||
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &slot_counter_buf));
|
return !e || e[0] == '\0' || e[0] != '0';
|
||||||
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);
|
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
|
// [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
|
// ceil(k_e/n_tile_size)) and compare to the ideal tile count for the real
|
||||||
|
|||||||
@@ -68,6 +68,79 @@ __kernel void kernel_moe_scatter(
|
|||||||
emap[tile_idx] = val;
|
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(
|
__kernel void kernel_moe_fill(
|
||||||
__global int * post_router,
|
__global int * post_router,
|
||||||
__global int * total_tiles,
|
__global int * total_tiles,
|
||||||
|
|||||||
Reference in New Issue
Block a user