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];