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;