Commit ba6439a6b for llama.cpp

commit ba6439a6b595224c480d71572c0e4ac91d3e6afd
Author: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>
Date:   Fri Oct 9 22:00:37 2026 +0900

    webgpu: use 2D workgroup dispatch for all the ops which use 1D dispatch (e.g., rms_norm) (#30219)

    * remove 1D workgroups dispatching

    * formatting

diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp
index b14a1fe2b..3fbac9675 100644
--- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp
+++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp
@@ -847,8 +847,11 @@ static webgpu_encoded_op ggml_webgpu_set(webgpu_context & ctx,
     entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index, src1));
     entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index + 1, dst));

-    uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
@@ -897,8 +900,11 @@ static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src
         ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst),
     };

-    uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_solve_tri(webgpu_context & ctx,
@@ -1495,8 +1501,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_set_rows(webgpu_context & ct
     } else {
         threads = src->ne[0] * src->ne[1] * src->ne[2] * src->ne[3];
     }
-    uint32_t wg_x = CEIL_DIV(threads, decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, 1);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(threads, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 // Workgroup size is a common constant
@@ -1557,9 +1566,11 @@ static webgpu_encoded_op ggml_webgpu_get_rows(webgpu_context & ctx,
     uint32_t blocks_per_row = (uint32_t) (dst->ne[0] / (decisions->vectorized ? 4 : 1));
     uint32_t total_rows     = (uint32_t) (dst->ne[1] * dst->ne[2] * dst->ne[3]);
     uint32_t total_threads  = float_parallel ? blocks_per_row * total_rows : total_rows;
-    uint32_t wg_x           = CEIL_DIV(total_threads, decisions->wg_size);
-
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(total_threads, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static void ggml_webgpu_quantize_q8_dispatch(webgpu_context &                    ctx,
@@ -2347,7 +2358,8 @@ static webgpu_encoded_op ggml_webgpu_unary_op(webgpu_context & ctx, ggml_tensor
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
     }

-    uint32_t wg_x, wg_y;
+    uint32_t wg_x;
+    uint32_t wg_y;
     uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
     compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
     return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
@@ -2425,7 +2437,8 @@ static webgpu_encoded_op ggml_webgpu_binary_op(webgpu_context & ctx,
         }
     }

-    uint32_t wg_x, wg_y;
+    uint32_t wg_x;
+    uint32_t wg_y;
     uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
     compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
     return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
@@ -2556,8 +2569,11 @@ static webgpu_encoded_op ggml_webgpu_concat(webgpu_context & ctx,
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
     }

-    uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor * src0, ggml_tensor * dst) {
@@ -2591,8 +2607,12 @@ static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor *

     webgpu_pipeline pipeline  = ctx->shader_lib->get_repeat_pipeline(shader_lib_ctx);
     auto *          decisions = static_cast<ggml_webgpu_generic_shader_decisions *>(pipeline.context.get());
-    uint32_t        wg_x      = CEIL_DIV(ne, decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context & ctx,
@@ -2637,6 +2657,7 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
         (uint32_t) dst->ne[0],
         (uint32_t) dst->ne[1],
         (uint32_t) dst->ne[2],
+        (uint32_t) dst->ne[3],
         ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(rn_dst, 0))  // epsilon, treated as f32 in the shader
     };

@@ -2676,7 +2697,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
     }

-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
+    uint32_t wg_x;
+    uint32_t wg_y;
+    compute_2d_workgroups(ggml_nrows(dst), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
+                          wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
@@ -2692,6 +2717,7 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
         (uint32_t) src->ne[0],
         (uint32_t) src->ne[1],
         (uint32_t) src->ne[2],
+        (uint32_t) src->ne[3],
         ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 0))  // epsilon, treated as f32 in the shader
     };

@@ -2707,7 +2733,12 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
     if (!decisions->inplace) {
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
     }
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(src));
+
+    uint32_t wg_x;
+    uint32_t wg_y;
+    compute_2d_workgroups(ggml_nrows(src), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
+                          wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
@@ -2796,8 +2827,11 @@ static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));
     }

-    uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
@@ -2870,8 +2904,11 @@ static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
     }
     entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));

-    uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
@@ -2909,7 +2946,8 @@ static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * s
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
     }

-    uint32_t wg_x, wg_y;
+    uint32_t wg_x;
+    uint32_t wg_y;
     uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
     compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
     return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
@@ -2982,7 +3020,11 @@ static webgpu_encoded_op ggml_webgpu_soft_max(webgpu_context & ctx,
         entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num, dst));
     }

-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = ggml_nrows(dst);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_argmax(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
@@ -3165,8 +3207,11 @@ static webgpu_encoded_op ggml_webgpu_cumsum(webgpu_context & ctx, ggml_tensor *
     shader_lib_ctx.max_wg_size = ctx->global_ctx->capabilities.limits.maxComputeInvocationsPerWorkgroup;

     webgpu_pipeline pipeline = ctx->shader_lib->get_cumsum_pipeline(shader_lib_ctx);
-    uint32_t        wg_x     = ggml_nrows(dst);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t        wg_x;
+    uint32_t        wg_y;
+    uint32_t        total_wg = ggml_nrows(dst);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
@@ -3190,8 +3235,11 @@ static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor

     webgpu_pipeline pipeline = ctx->shader_lib->get_sum_rows_pipeline(shader_lib_ctx);

-    uint32_t wg_x = total_sum ? 1 : ggml_nrows(dst);
-    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
+    uint32_t wg_x;
+    uint32_t wg_y;
+    uint32_t total_wg = total_sum ? 1 : ggml_nrows(dst);
+    compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
+    return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
 }

 static bool ggml_webgpu_can_fuse_rms_norm_mul(const struct ggml_cgraph * cgraph, int node_idx) {
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl
index 7ccad73f4..de6d6e6c7 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl
@@ -53,10 +53,12 @@ var<storage, read_write> dst: array<DataType>;
 var<uniform> params: Params;
 #endif
 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {

-    if (gid.x < params.ne) {
-        var i = gid.x;
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i < params.ne) {
+        var i = gid_i;
         let i3 = i / (params.ne2 * params.ne1 * params.ne0);
         i = i % (params.ne2 * params.ne1 * params.ne0);
         let i2 = i / (params.ne1 * params.ne0);
@@ -72,9 +74,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
                              ni[2] * params.stride_src0_2 +
                              ni[3] * params.stride_src0_3;
 #ifdef SRC_OVERLAP
-            dst[params.offset_dst + gid.x] = merged_src[params.offset_src0 + src_i];
+            dst[params.offset_dst + gid_i] = merged_src[params.offset_src0 + src_i];
 #else
-            dst[params.offset_dst + gid.x] = src0[params.offset_src0 + src_i];
+            dst[params.offset_dst + gid_i] = src0[params.offset_src0 + src_i];
 #endif
         } else {
             ni[params.dim] -= params.src0_nedim;
@@ -83,9 +85,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
                              ni[2] * params.stride_src1_2 +
                              ni[3] * params.stride_src1_3;
 #ifdef SRC_OVERLAP
-            dst[params.offset_dst + gid.x] = merged_src[params.offset_src1 + src_i];
+            dst[params.offset_dst + gid_i] = merged_src[params.offset_src1 + src_i];
 #else
-            dst[params.offset_dst + gid.x] = src1[params.offset_src1 + src_i];
+            dst[params.offset_dst + gid_i] = src1[params.offset_src1 + src_i];
 #endif
         }
     }
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/cumsum.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/cumsum.wgsl
index e622552c4..58c703dcf 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/cumsum.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/cumsum.wgsl
@@ -17,8 +17,11 @@ var<workgroup> shared_sum: array<f32, WG_SIZE>;

 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wid: vec3<u32>,
+        @builtin(num_workgroups) num_wg: vec3<u32>,
         @builtin(local_invocation_id) lid: vec3<u32>) {
-    let row_idx = params.offset_src + wid.x * params.ne0;
+
+    let wid_i = wid.x + wid.y * num_wg.x;
+    let row_idx = params.offset_src + wid_i * params.ne0;
     let elems = (params.ne0 + WG_SIZE - 1) / WG_SIZE;
     var local_sum: f32 = 0.0;
     for (var col = lid.x * elems; col < (lid.x + 1) * elems && col < params.ne0; col ++) {
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl
index 6bbed5d3b..6f1c367d5 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl
@@ -157,12 +157,15 @@ fn b_value(base: u32) -> DataType {
 #endif

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
-    if (gid.x >= params.ne) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i >= params.ne) {
         return;
     }

-    var i = gid.x;
+    var i = gid_i;
     let i3 = i / (params.ne2 * params.ne1 * params.ne0);
     i = i % (params.ne2 * params.ne1 * params.ne0);
     let i2 = i / (params.ne1 * params.ne0);
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/pad.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/pad.wgsl
index ea63b9a73..2cb999bef 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/pad.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/pad.wgsl
@@ -45,12 +45,15 @@ fn wrap_around(idx: i32, n: u32) -> u32 {
 }

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
-    if (gid.x >= params.ne) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i >= params.ne) {
         return;
     }

-    var i = gid.x;
+    var i = gid_i;
     let dst_plane = params.dst_ne2 * params.dst_ne1 * params.dst_ne0;
     let i3 = i / dst_plane;
     i = i % dst_plane;
@@ -82,5 +85,5 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
     }
 #endif

-    dst[params.offset_dst + gid.x] = value;
+    dst[params.offset_dst + gid_i] = value;
 }
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/repeat.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/repeat.wgsl
index 43b883e67..9debcd9c5 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/repeat.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/repeat.wgsl
@@ -45,9 +45,12 @@ var<storage, read_write> dst: array<DataType>;
 var<uniform> params: Params;

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
-    if (gid.x < params.ne) {
-        var i = gid.x;
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i < params.ne) {
+        var i = gid_i;
         let i3 = i / (params.ne2 * params.ne1 * params.ne0);
         i = i % (params.ne2 * params.ne1 * params.ne0);
         let i2 = i / (params.ne1 * params.ne0);
@@ -65,6 +68,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
                            a_i2 * params.stride_src0_2 +
                            a_i3 * params.stride_src0_3;

-        dst[params.offset_dst + gid.x] = src0[params.offset_src0 + a_index];
+        dst[params.offset_dst + gid_i] = src0[params.offset_src0 + a_index];
     }
 }
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl
index c9e424ffc..46a208029 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl
@@ -88,6 +88,7 @@ struct Params {
     ne0: u32,
     ne1: u32,
     ne2: u32,
+    ne3: u32,

     eps: f32
 };
@@ -96,10 +97,14 @@ var<workgroup> scratch: array<f32, WG_SIZE>;

 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wid: vec3<u32>,
+        @builtin(num_workgroups) num_wg: vec3<u32>,
         @builtin(local_invocation_id) lid: vec3<u32>) {

-    // one thread per row
-    var i = wid.x;
+    // one workgroup per row
+    var i = wid.x + wid.y * num_wg.x;
+    if (i >= params.ne1 * params.ne2 * params.ne3) {
+        return;
+    }
     let i3 = i / (params.ne2 * params.ne1);
     i = i % (params.ne2 * params.ne1);
     let i2 = i / params.ne1;
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/rope.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/rope.wgsl
index 6ff53088c..31a73c613 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/rope.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/rope.wgsl
@@ -145,9 +145,12 @@ fn pair_offset(is_neox: bool, is_mrope: bool, is_vision: bool) -> u32 {
 }

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
     // two elements per n_threads
-    if (gid.x >= params.n_threads) {
+    if (gid_i >= params.n_threads) {
         return;
     }

@@ -156,7 +159,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
     let is_imrope = params.mode == 40;
     let is_vision = params.mode == 24;

-    var i = gid.x * 2; // start index for this thread
+    var i = gid_i * 2; // start index for this thread
     let i3 = i / (params.ne2 * params.ne1 * params.ne0);
     i = i % (params.ne2 * params.ne1 * params.ne0);
     let i2 = i / (params.ne1 * params.ne0);
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl
index 7629bf5b4..a87bf897a 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl
@@ -31,6 +31,7 @@ struct Params {
     ne0: u32,
     ne1: u32,
     ne2: u32,
+    ne3: u32,

     eps: f32
 };
@@ -53,10 +54,14 @@ var<workgroup> scratch: array<f32, WG_SIZE * 2u>;

 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wid: vec3<u32>,
+        @builtin(num_workgroups) num_wg: vec3<u32>,
         @builtin(local_invocation_id) lid: vec3<u32>) {

-    // one thread per row
-    var i = wid.x;
+    // one workgroup per row
+    var i = wid.x + wid.y * num_wg.x;
+    if (i >= params.ne1 * params.ne2 * params.ne3) {
+        return;
+    }
     let i3 = i / (params.ne2 * params.ne1);
     i = i % (params.ne2 * params.ne1);
     let i2 = i / params.ne1;
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl
index 0a7ae9bdb..919db34c3 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl
@@ -84,26 +84,29 @@ fn in_set_view(rel: u32, coords: vec4<u32>) -> bool {
 }

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
-    if (gid.x >= params.ne) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i >= params.ne) {
         return;
     }

 #ifdef INPLACE
-    let coords = decode_src1_coords(gid.x);
+    let coords = decode_src1_coords(gid_i);

     let src1_idx = params.offset_src1 + src1_idx_from_coords(coords);
     let dst_idx = params.offset_view + view_rel_from_coords(coords);

     dst[dst_idx] = src1[src1_idx];
 #else
-    let rel = select(params.ne, gid.x - params.offset_view, gid.x >= params.offset_view);
+    let rel = select(params.ne, gid_i - params.offset_view, gid_i >= params.offset_view);
     let coords = decode_view_coords(rel);

     if (rel < params.stride_dst13 * params.src1_ne3 && in_set_view(rel, coords)) {
-        dst[gid.x] = src1[params.offset_src1 + src1_idx_from_coords(coords)];
+        dst[gid_i] = src1[params.offset_src1 + src1_idx_from_coords(coords)];
     } else {
-        dst[gid.x] = src0[params.offset_src0 + gid.x];
+        dst[gid_i] = src0[params.offset_src0 + gid_i];
     }
 #endif
 }
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl
index 91c3d9c74..2e1b0111e 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl
@@ -70,13 +70,16 @@ struct Params {
 var<uniform> params: Params;

 @compute @workgroup_size(WG_SIZE)
-fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
-    if (gid.x >= (params.ne3 * params.ne2 * params.n_rows * params.ne0) / VEC_SIZE) {
+fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
+        @builtin(global_invocation_id) gid: vec3<u32>) {
+
+    let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
+    if (gid_i >= (params.ne3 * params.ne2 * params.n_rows * params.ne0) / VEC_SIZE) {
         return;
     }

     let elems_per_row = params.ne0 / VEC_SIZE;
-    var i = gid.x / elems_per_row;
+    var i = gid_i / elems_per_row;

     let i_src3 = i / (params.ne2 * params.n_rows);

@@ -107,6 +110,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
     let i_dst_row = params.offset_dst + idx_val * params.stride_dst1 + i_src2 * params.stride_dst2 + i_src3 * params.stride_dst3;
     let i_src_row = params.offset_src + i_src1 * params.stride_src1 + i_src2 * params.stride_src2 + i_src3 * params.stride_src3;

-    let col_idx = gid.x % elems_per_row;
+    let col_idx = gid_i % elems_per_row;
     dst[i_dst_row / VEC_SIZE + col_idx] = DST_TYPE(src[i_src_row / VEC_SIZE + col_idx]);
 }
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl
index 1c29a9221..8f92a04c1 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl
@@ -124,9 +124,10 @@ var<workgroup> scratch: array<f32, WG_SIZE>;

 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wid: vec3<u32>,
+        @builtin(num_workgroups) num_wg: vec3<u32>,
         @builtin(local_invocation_id) lid: vec3<u32>) {

-    var i = wid.x;
+    var i = wid.x + wid.y * num_wg.x;
     let i3 = i / (params.ne2 * params.ne1);
     i = i % (params.ne2 * params.ne1);
     let i2 = i / params.ne1;
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/sum_rows.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/sum_rows.wgsl
index 6ea2de9b7..1ead3ac3f 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/sum_rows.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/sum_rows.wgsl
@@ -25,9 +25,10 @@ var<workgroup> shared_sum: array<f32, WG_SIZE>;

 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wid: vec3<u32>,
+        @builtin(num_workgroups) num_wg: vec3<u32>,
         @builtin(local_invocation_id) lid: vec3<u32>) {

-    var i = wid.x;
+    var i = wid.x + wid.y * num_wg.x;
     let i3 = i / (params.ne2 * params.ne1);
     i = i % (params.ne2 * params.ne1);
     let i2 = i / params.ne1;