Commit 5e4878e97 for llama.cpp

commit 5e4878e9787fff9bcc5c12478c47618b73268610
Author: Ruben Ortlam <rortlam@redhat.com>
Date:   Fri Oct 9 15:00:14 2026 +0200

    vulkan: fix rms_norm workgroup count overflow (#30145)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 998cc693c..42e407ebf 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -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 };
         } else {
             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;

@@ -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, set_rows, true),
                 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);
         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, residual),
                     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);
         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;
         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);
         ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
index ee813842c..314c8e565 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
@@ -53,101 +53,107 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE];
 void rms_norm(uint num_iters) {
     const uint ncols     = p.ne00;
     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 channel   = gl_WorkGroupID.y;
-    const uint samp      = gl_WorkGroupID.z;
     const uint tid       = gl_LocalInvocationID.x;

     const uint stride_row       = p.nb01;
     const uint stride_channel   = p.nb02;
     const uint stride_sample    = p.nb03;

-    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();
+    // 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 b_offset = src1_idx(0, row, channel, samp) + get_boffset();
 #if RMS_NORM_ROPE_FUSION
-    // Per-row offset in shared memory
-    uint32_t d_offset = 0;
+            // Per-row offset in shared memory
+            uint32_t d_offset = 0;
 #elif RMS_NORM_SET_ROWS_FUSION
-    uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
+            uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
 #else
-    uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
+            uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
 #endif
-    FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
+            FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp

-    [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
-        FLOAT_TYPE xi = FLOAT_TYPE(0);
-        if (col < ncols) {
-            xi = FLOAT_TYPE(data_a[a_offset + col]);
-        }
-        sum += xi * xi;
-    }
+            [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
+                FLOAT_TYPE xi = FLOAT_TYPE(0);
+                if (col < ncols) {
+                    xi = FLOAT_TYPE(data_a[a_offset + col]);
+                }
+                sum += xi * xi;
+            }

-    sumsh[tid] = sum;
-    // sum up partial sums and write back result
-    barrier();
-    [[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
-        if (tid < s) {
-            sum += sumsh[tid + s];
             sumsh[tid] = sum;
-        }
-        barrier();
-    }
-    sum = sumsh[0];
-
-    const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
-    const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
-
-    if (do_multiply) {
-        if (ncols > p.ne10) {
-            [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
-                if (col >= ncols) {
-                    continue;
+            // sum up partial sums and write back result
+            barrier();
+            [[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
+                if (tid < s) {
+                    sum += sumsh[tid + s];
+                    sumsh[tid] = sum;
                 }
-                FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
+                barrier();
+            }
+            sum = sumsh[0];
+
+            const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
+            const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
+
+            if (do_multiply) {
+                if (ncols > p.ne10) {
+                    [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
+                        if (col >= ncols) {
+                            continue;
+                        }
+                        FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
 #if RMS_NORM_ADD_FUSION
-                value += FLOAT_TYPE(data_c[d_offset + col]);
-                if (do_post_multiply) {
-                    value *= FLOAT_TYPE(data_e[0]);
-                }
+                        value += FLOAT_TYPE(data_c[d_offset + col]);
+                        if (do_post_multiply) {
+                            value *= FLOAT_TYPE(data_e[0]);
+                        }
 #endif
-                data_d[d_offset + col] = D_TYPE(value);
-            }
-        } else {
-            [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
-                if (col >= ncols) {
-                    continue;
-                }
-                FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
+                        data_d[d_offset + col] = D_TYPE(value);
+                    }
+                } else {
+                    [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
+                        if (col >= ncols) {
+                            continue;
+                        }
+                        FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
 #if RMS_NORM_ADD_FUSION
-                value += FLOAT_TYPE(data_c[d_offset + col]);
-                if (do_post_multiply) {
-                    value *= FLOAT_TYPE(data_e[0]);
-                }
+                        value += FLOAT_TYPE(data_c[d_offset + col]);
+                        if (do_post_multiply) {
+                            value *= FLOAT_TYPE(data_e[0]);
+                        }
 #endif
-                data_d[d_offset + col] = D_TYPE(value);
-            }
-        }
-    } else {
-        [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
-            if (col >= ncols) {
-                continue;
+                        data_d[d_offset + col] = D_TYPE(value);
+                    }
+                }
+            } else {
+                [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
+                    if (col >= ncols) {
+                        continue;
+                    }
+                    data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
+                }
             }
-            data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
-        }
-    }
 #if RMS_NORM_ROPE_FUSION
-    barrier();
-    rope_params rp = p.rope;
-    for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
-        if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
-            rope_neox(t, row, channel, samp, rp);
-        } else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
-            rope_norm(t, row, channel, samp, rp);
+            barrier();
+            rope_params rp = p.rope;
+            for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
+                if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
+                    rope_neox(t, row, channel, samp, rp);
+                } else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
+                    rope_norm(t, row, channel, samp, rp);
+                }
+            }
+#endif
         }
     }
-#endif
 }

 void main() {