vulkan: fix rms_norm workgroup count overflow (#30145)

This commit is contained in:
Ruben Ortlam
2026-10-09 15:00:14 +02:00
committed by GitHub
parent 609290be6b
commit 5e4878e978
2 changed files with 90 additions and 72 deletions
+14 -2
View File
@@ -9455,6 +9455,8 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 }; elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 };
} else { } else {
elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 }; elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 };
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
} }
break; break;
@@ -10777,7 +10779,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
ggml_vk_tensor_subbuffer(ctx, src0, true), ggml_vk_tensor_subbuffer(ctx, src0, true),
ggml_vk_tensor_subbuffer(ctx, set_rows, true), ggml_vk_tensor_subbuffer(ctx, set_rows, true),
ggml_vk_tensor_subbuffer(ctx, indices), ggml_vk_tensor_subbuffer(ctx, indices),
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] }); }, pc, {
(uint32_t)src0->ne[1],
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
});
ggml_vk_rms_norm_finish(ctx, src0); ggml_vk_rms_norm_finish(ctx, src0);
return; return;
} }
@@ -10824,7 +10830,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
ggml_vk_tensor_subbuffer(ctx, dst, true), ggml_vk_tensor_subbuffer(ctx, dst, true),
ggml_vk_tensor_subbuffer(ctx, residual), ggml_vk_tensor_subbuffer(ctx, residual),
ggml_vk_tensor_subbuffer(ctx, post_scale), ggml_vk_tensor_subbuffer(ctx, post_scale),
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] }); }, pc, {
(uint32_t)src0->ne[1],
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
});
} }
ggml_vk_rms_norm_finish(ctx, src0); ggml_vk_rms_norm_finish(ctx, src0);
return; return;
@@ -10911,6 +10921,8 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
std::array<uint32_t, 3> elements; std::array<uint32_t, 3> elements;
elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] }; elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] };
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
static_assert(max_tensors == 7); static_assert(max_tensors == 7);
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
@@ -53,17 +53,21 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE];
void rms_norm(uint num_iters) { void rms_norm(uint num_iters) {
const uint ncols = p.ne00; const uint ncols = p.ne00;
const uint nrows = gl_NumWorkGroups.x; const uint nrows = gl_NumWorkGroups.x;
const uint nchannels = gl_NumWorkGroups.y; const uint nchannels = p.ne02;
const uint nsamples = p.ne03;
const uint row = gl_WorkGroupID.x; const uint row = gl_WorkGroupID.x;
const uint channel = gl_WorkGroupID.y;
const uint samp = gl_WorkGroupID.z;
const uint tid = gl_LocalInvocationID.x; const uint tid = gl_LocalInvocationID.x;
const uint stride_row = p.nb01; const uint stride_row = p.nb01;
const uint stride_channel = p.nb02; const uint stride_channel = p.nb02;
const uint stride_sample = p.nb03; const uint stride_sample = p.nb03;
// grid.y/z are clamped to the device workgroup limit, iterate over the excess channels/samples
for (uint samp = gl_WorkGroupID.z; samp < nsamples; samp += gl_NumWorkGroups.z) {
for (uint channel = gl_WorkGroupID.y; channel < nchannels; channel += gl_NumWorkGroups.y) {
barrier();
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset(); uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset(); uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
#if RMS_NORM_ROPE_FUSION #if RMS_NORM_ROPE_FUSION
@@ -148,6 +152,8 @@ void rms_norm(uint num_iters) {
} }
} }
#endif #endif
}
}
} }
void main() { void main() {