Commit 0c1e57098 for llama.cpp
commit 0c1e57098bba43ac29e6e3b677cdceebdd22334f
Author: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>
Date: Thu Oct 1 11:11:48 2026 +0900
webgpu: fix SSM_SCAN binding aliasing (#29750)
diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
index 47a266d7d..778a1b4bf 100644
--- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
+++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
@@ -136,11 +136,11 @@ struct ggml_webgpu_ssm_conv_shader_decisions {
};
struct ggml_webgpu_ssm_scan_pipeline_key {
- int type;
- int d_state;
- bool xbc_overlap;
- bool a_overlap;
- bool ids_overlap;
+ int type;
+ int d_state;
+ uint8_t xbc_overlap;
+ bool a_overlap;
+ bool ids_overlap;
bool operator==(const ggml_webgpu_ssm_scan_pipeline_key & other) const {
return type == other.type && d_state == other.d_state && xbc_overlap == other.xbc_overlap &&
@@ -163,7 +163,7 @@ struct ggml_webgpu_ssm_scan_pipeline_key_hash {
struct ggml_webgpu_ssm_scan_shader_decisions {
uint32_t wg_size;
uint32_t tokens_per_tile;
- bool xbc_overlap = false;
+ uint8_t xbc_overlap = 0;
bool a_overlap = false;
bool ids_overlap = false;
};
@@ -1797,16 +1797,11 @@ class ggml_webgpu_shader_lib {
return ssm_conv_pipelines[key];
}
- webgpu_pipeline get_ssm_scan_pipeline(const ggml_webgpu_shader_lib_context & context,
- bool xbc_overlap,
- bool a_overlap,
- bool ids_overlap) {
+ webgpu_pipeline get_ssm_scan_pipeline(const ggml_webgpu_shader_lib_context & context, uint8_t xbc_overlap) {
ggml_webgpu_ssm_scan_pipeline_key key = {};
key.type = context.dst->type;
key.d_state = (int) context.src0->ne[0];
key.xbc_overlap = xbc_overlap;
- key.a_overlap = a_overlap;
- key.ids_overlap = ids_overlap;
auto it = ssm_scan_pipelines.find(key);
if (it != ssm_scan_pipelines.end()) {
@@ -1838,15 +1833,17 @@ class ggml_webgpu_shader_lib {
variant += "_wg_reduce";
}
- if (key.xbc_overlap) {
+ if (key.xbc_overlap == 0b110) { // x/B
+ defines.push_back("XB_OVERLAP");
+ variant += "_xb_overlap";
+ } else if (key.xbc_overlap == 0b011) { // B/C
+ defines.push_back("BC_OVERLAP");
+ variant += "_bc_overlap";
+ } else if (key.xbc_overlap == 0b111) { // x/B/C
defines.push_back("XBC_OVERLAP");
+ variant += "_xbc_overlap";
}
- if (key.a_overlap) {
- defines.push_back("A_OVERLAP");
- }
- if (key.ids_overlap) {
- defines.push_back("IDS_OVERLAP");
- }
+
variant += "_d" + std::to_string(key.d_state);
auto processed = preprocessor.preprocess(wgsl_ssm_scan, defines);
@@ -1854,8 +1851,6 @@ class ggml_webgpu_shader_lib {
decisions->wg_size = wg_size;
decisions->tokens_per_tile = tokens_per_tile;
decisions->xbc_overlap = key.xbc_overlap;
- decisions->a_overlap = key.a_overlap;
- decisions->ids_overlap = key.ids_overlap;
webgpu_pipeline pipeline = ggml_webgpu_create_pipeline(device, processed, variant);
pipeline.context = decisions;
ssm_scan_pipelines[key] = pipeline;
diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp
index 1ebff43f3..dd806ab99 100644
--- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp
+++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp
@@ -1242,59 +1242,53 @@ static webgpu_encoded_op ggml_webgpu_ssm_scan(webgpu_context & ctx,
shader_lib_ctx.dst = dst;
shader_lib_ctx.max_wg_size = ctx->global_ctx->capabilities.limits.maxComputeInvocationsPerWorkgroup;
shader_lib_ctx.supports_subgroups = ctx->global_ctx->capabilities.supports_subgroups;
- bool xbc_overlap = ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src1, src2) ||
- ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src1, src4) ||
- ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src1, src5) ||
- ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src2, src4) ||
- ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src2, src5) ||
- ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src4, src5);
- bool a_overlap = false;
- bool ids_overlap = false;
+
+ uint8_t xbc_overlap = 0;
+ if (ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src1, src4)) {
+ xbc_overlap |= 0b110; // x/B
+ }
+ if (ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src1, src5)) {
+ xbc_overlap |= 0b101; // x/C
+ }
+ if (ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src4, src5)) {
+ xbc_overlap |= 0b011; // B/C
+ }
+
+ webgpu_pipeline pipeline = ctx->shader_lib->get_ssm_scan_pipeline(shader_lib_ctx, xbc_overlap);
+ auto * decisions = static_cast<ggml_webgpu_ssm_scan_shader_decisions *>(pipeline.context.get());
+ xbc_overlap = decisions->xbc_overlap;
+
+ uint32_t offset_x = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src1) / ggml_type_size(src1->type));
+ uint32_t offset_dt = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src2) / ggml_type_size(src2->type));
+ uint32_t offset_A = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src3) / ggml_type_size(src3->type));
+ uint32_t offset_B = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src4) / ggml_type_size(src4->type));
+ uint32_t offset_C = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src5) / ggml_type_size(src5->type));
+ uint32_t offset_ids = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src6) / ggml_type_size(src6->type));
+
ggml_webgpu_merged_binding_range xbc_merged_range = {};
- if (xbc_overlap) {
- xbc_merged_range = ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src2, src4, src5 });
- a_overlap = ggml_webgpu_tensor_binding_overlap_range(ctx->global_ctx, src3, src1->buffer,
- xbc_merged_range.offset, xbc_merged_range.size);
- if (a_overlap) {
- xbc_merged_range = ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src2, src3, src4, src5 });
- }
- ids_overlap = ggml_webgpu_tensor_binding_overlap_range(ctx->global_ctx, src6, src1->buffer,
- xbc_merged_range.offset, xbc_merged_range.size);
- if (ids_overlap) {
- xbc_merged_range =
- a_overlap ? ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src2, src3, src4, src5, src6 }) :
- ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src2, src4, src5, src6 });
- }
+
+ if (xbc_overlap == 0b110) { // x/B
+ xbc_merged_range = ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src4 });
+ offset_x = ggml_webgpu_tensor_merged_element_offset(src1, xbc_merged_range);
+ offset_B = ggml_webgpu_tensor_merged_element_offset(src4, xbc_merged_range);
+ } else if (xbc_overlap == 0b011) { // B/C
+ xbc_merged_range = ggml_webgpu_tensor_merged_binding_range(ctx, { src4, src5 });
+ offset_B = ggml_webgpu_tensor_merged_element_offset(src4, xbc_merged_range);
+ offset_C = ggml_webgpu_tensor_merged_element_offset(src5, xbc_merged_range);
+ } else if (xbc_overlap == 0b111) { // x/B/C
+ xbc_merged_range = ggml_webgpu_tensor_merged_binding_range(ctx, { src1, src4, src5 });
+ offset_x = ggml_webgpu_tensor_merged_element_offset(src1, xbc_merged_range);
+ offset_B = ggml_webgpu_tensor_merged_element_offset(src4, xbc_merged_range);
+ offset_C = ggml_webgpu_tensor_merged_element_offset(src5, xbc_merged_range);
}
- webgpu_pipeline pipeline =
- ctx->shader_lib->get_ssm_scan_pipeline(shader_lib_ctx, xbc_overlap, a_overlap, ids_overlap);
- auto * decisions = static_cast<ggml_webgpu_ssm_scan_shader_decisions *>(pipeline.context.get());
- xbc_overlap = decisions->xbc_overlap;
- a_overlap = decisions->a_overlap;
- ids_overlap = decisions->ids_overlap;
-
- uint32_t offset_x = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src1) / ggml_type_size(src1->type));
- uint32_t offset_dt = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src2) / ggml_type_size(src2->type));
- uint32_t offset_A = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src3) / ggml_type_size(src3->type));
- uint32_t offset_B = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src4) / ggml_type_size(src4->type));
- uint32_t offset_C = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src5) / ggml_type_size(src5->type));
- uint32_t offset_ids = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src6) / ggml_type_size(src6->type));
- size_t xbc_bind_offset = 0;
- size_t xbc_bind_size = 0;
- if (xbc_overlap) {
+ GGML_ASSERT(xbc_overlap == 0 || xbc_overlap == 0b110 || xbc_overlap == 0b011 || xbc_overlap == 0b111);
+
+ size_t xbc_bind_offset = 0;
+ size_t xbc_bind_size = 0;
+ if (xbc_overlap > 0) {
xbc_bind_offset = xbc_merged_range.offset;
xbc_bind_size = xbc_merged_range.size;
- offset_x = ggml_webgpu_tensor_merged_element_offset(src1, xbc_merged_range);
- offset_dt = ggml_webgpu_tensor_merged_element_offset(src2, xbc_merged_range);
- if (a_overlap) {
- offset_A = ggml_webgpu_tensor_merged_element_offset(src3, xbc_merged_range);
- }
- offset_B = ggml_webgpu_tensor_merged_element_offset(src4, xbc_merged_range);
- offset_C = ggml_webgpu_tensor_merged_element_offset(src5, xbc_merged_range);
- if (ids_overlap) {
- offset_ids = ggml_webgpu_tensor_merged_element_offset(src6, xbc_merged_range);
- }
}
std::vector<uint32_t> params = {
@@ -1338,34 +1332,33 @@ static webgpu_encoded_op ggml_webgpu_ssm_scan(webgpu_context & ctx,
(uint32_t) ggml_get_op_params_i32(dst, 0),
};
+ uint32_t binding_num = 0;
+
std::vector<wgpu::BindGroupEntry> entries = {
- ggml_webgpu_make_tensor_bind_group_entry(ctx, 0, src0),
+ ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src0),
};
- if (xbc_overlap) {
- entries.push_back(
- ggml_webgpu_make_bind_group_entry(1, ggml_webgpu_tensor_buf(src1), xbc_bind_offset, xbc_bind_size));
- if (ids_overlap) {
- if (!a_overlap) {
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, src3));
- }
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, a_overlap ? 2 : 3, dst));
- } else if (a_overlap) {
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, src6));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 3, dst));
- } else {
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, src3));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 3, src6));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 4, dst));
- }
- } else {
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, src1));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, src2));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 3, src3));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 4, src4));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 5, src5));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 6, src6));
- entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 7, dst));
+ // xbc_merged binding
+ if (xbc_overlap > 0) {
+ entries.push_back(ggml_webgpu_make_bind_group_entry(binding_num++, ggml_webgpu_tensor_buf(src1),
+ xbc_bind_offset, xbc_bind_size));
+ }
+ // x
+ if (!(xbc_overlap & 0b100)) {
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src1));
}
+ // B
+ if (!(xbc_overlap & 0b010)) {
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src4));
+ }
+ // C
+ if (!(xbc_overlap & 0b001)) {
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src5));
+ }
+
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src2));
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src3));
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, src6));
+ entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num++, dst));
const uint32_t total_wg = (uint32_t) (src0->ne[1] * src0->ne[2] * src1->ne[3]);
const uint32_t max_wg_per_dim = ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension;
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl
index 57f012ad0..d031a985f 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl
@@ -45,41 +45,32 @@ struct Params {
};
@group(0) @binding(0) var<storage, read_write> s_in: array<f32>;
+
+// binding for x/B/C merged status
#ifdef XBC_OVERLAP
-#ifdef IDS_OVERLAP
-@group(0) @binding(1) var<storage, read_write> x_dt_B_C_ids_merged: array<u32>;
-#ifdef A_OVERLAP
-@group(0) @binding(2) var<storage, read_write> dst: array<f32>;
-@group(0) @binding(3) var<uniform> params: Params;
-#else
-@group(0) @binding(2) var<storage, read_write> A: array<f32>;
-@group(0) @binding(3) var<storage, read_write> dst: array<f32>;
-@group(0) @binding(4) var<uniform> params: Params;
-#endif
-#else
-@group(0) @binding(1) var<storage, read_write> x_dt_B_C_merged: array<f32>;
-#ifdef A_OVERLAP
-@group(0) @binding(2) var<storage, read_write> ids: array<i32>;
-@group(0) @binding(3) var<storage, read_write> dst: array<f32>;
-@group(0) @binding(4) var<uniform> params: Params;
-#else
-@group(0) @binding(2) var<storage, read_write> A: array<f32>;
-@group(0) @binding(3) var<storage, read_write> ids: array<i32>;
-@group(0) @binding(4) var<storage, read_write> dst: array<f32>;
-@group(0) @binding(5) var<uniform> params: Params;
-#endif
-#endif
+@group(0) @binding(1) var<storage, read_write> merged: array<f32>;
+#define BIND_DT 2
+#elif defined(BC_OVERLAP)
+@group(0) @binding(1) var<storage, read_write> merged: array<f32>;
+@group(0) @binding(2) var<storage, read_write> x: array<f32>;
+#define BIND_DT 3
+#elif defined(XB_OVERLAP)
+@group(0) @binding(1) var<storage, read_write> merged: array<f32>;
+@group(0) @binding(2) var<storage, read_write> C: array<f32>;
+#define BIND_DT 3
#else
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
-@group(0) @binding(2) var<storage, read_write> dt: array<f32>;
-@group(0) @binding(3) var<storage, read_write> A: array<f32>;
-@group(0) @binding(4) var<storage, read_write> B: array<f32>;
-@group(0) @binding(5) var<storage, read_write> C: array<f32>;
-@group(0) @binding(6) var<storage, read_write> ids: array<i32>;
-@group(0) @binding(7) var<storage, read_write> dst: array<f32>;
-@group(0) @binding(8) var<uniform> params: Params;
+@group(0) @binding(2) var<storage, read_write> B: array<f32>;
+@group(0) @binding(3) var<storage, read_write> C: array<f32>;
+#define BIND_DT 4
#endif
+@group(0) @binding(BIND_DT) var<storage, read_write> dt: array<f32>;
+@group(0) @binding(BIND_DT + 1) var<storage, read_write> A: array<f32>;
+@group(0) @binding(BIND_DT + 2) var<storage, read_write> ids: array<i32>;
+@group(0) @binding(BIND_DT + 3) var<storage, read_write> dst: array<f32>;
+@group(0) @binding(BIND_DT + 4) var<uniform> params: Params;
+
var<workgroup> shared_x_dt: array<f32, TOKENS_PER_TILE>;
var<workgroup> shared_dtsp: array<f32, TOKENS_PER_TILE>;
var<workgroup> shared_reduce: array<f32, TOKENS_PER_TILE * WG_SIZE>;
@@ -88,22 +79,14 @@ fn reduce_base(token_in_tile: u32) -> u32 {
return token_in_tile * WG_SIZE;
}
-#ifdef XBC_OVERLAP
+#if defined(XBC_OVERLAP) || defined(XB_OVERLAP) || defined(BC_OVERLAP)
fn read_merged_f32(idx: u32) -> f32 {
-#ifdef IDS_OVERLAP
- return bitcast<f32>(x_dt_B_C_ids_merged[idx]);
-#else
- return x_dt_B_C_merged[idx];
-#endif
+ return merged[idx];
}
#endif
fn read_state_slot(i3: u32) -> u32 {
-#ifdef IDS_OVERLAP
- return x_dt_B_C_ids_merged[params.offset_ids + i3];
-#else
return u32(ids[params.offset_ids + i3]);
-#endif
}
@compute @workgroup_size(WG_SIZE)
@@ -133,11 +116,7 @@ fn main(
var s_prev = s_in[s_idx];
let a_idx = params.offset_A + (tid % params.a_ne0) + ir * params.stride_A1;
-#ifdef A_OVERLAP
- let A0 = read_merged_f32(a_idx);
-#else
let A0 = A[a_idx];
-#endif
for (var token_base = 0u; token_base < params.n_seq_tokens; token_base += TOKENS_PER_TILE) {
if (tid < TOKENS_PER_TILE) {
@@ -145,14 +124,10 @@ fn main(
if (token < params.n_seq_tokens) {
let x_idx = params.offset_x + i1 + ir * params.stride_x1 + token * params.stride_x2 + i3 * params.stride_x3;
let dt_idx = params.offset_dt + ir + token * params.stride_dt1 + i3 * params.stride_dt2;
-#ifdef XBC_OVERLAP
- let dt0 = read_merged_f32(dt_idx);
-#else
let dt0 = dt[dt_idx];
-#endif
let dtsp = select(log(1.0 + exp(dt0)), dt0, dt0 > 20.0);
shared_dtsp[tid] = dtsp;
-#ifdef XBC_OVERLAP
+#if defined(XBC_OVERLAP) || defined(XB_OVERLAP)
shared_x_dt[tid] = read_merged_f32(x_idx) * dtsp;
#else
shared_x_dt[tid] = x[x_idx] * dtsp;
@@ -174,7 +149,7 @@ fn main(
let b_idx = params.offset_B + tid + g * params.stride_B1 + token * params.stride_B2 + i3 * params.stride_B3;
let c_idx = params.offset_C + tid + g * params.stride_C1 + token * params.stride_C2 + i3 * params.stride_C3;
-#ifdef XBC_OVERLAP
+#if defined(XBC_OVERLAP) || defined(BC_OVERLAP) || defined(XB_OVERLAP)
let s = s_prev * dA + read_merged_f32(b_idx) * x_dt;
#else
let s = s_prev * dA + B[b_idx] * x_dt;
@@ -191,7 +166,7 @@ fn main(
}
#ifdef USE_SUBGROUP_REDUCTION
-#ifdef XBC_OVERLAP
+#if defined(XBC_OVERLAP) || defined(BC_OVERLAP)
let subgroup_partial = subgroupAdd(s * read_merged_f32(c_idx));
#else
let subgroup_partial = subgroupAdd(s * C[c_idx]);
@@ -200,7 +175,7 @@ fn main(
shared_reduce[reduce_idx - tid + subgroup_id] = subgroup_partial;
}
#else
-#ifdef XBC_OVERLAP
+#if defined(XBC_OVERLAP) || defined(BC_OVERLAP)
shared_reduce[reduce_idx] = s * read_merged_f32(c_idx);
#else
shared_reduce[reduce_idx] = s * C[c_idx];