From 0c1e57098bba43ac29e6e3b677cdceebdd22334f Mon Sep 17 00:00:00 2001 From: Masashi Yoshimura Date: Thu, 1 Oct 2026 11:11:48 +0900 Subject: [PATCH] webgpu: fix SSM_SCAN binding aliasing (#29750) --- .../ggml-webgpu/ggml-webgpu-shader-lib.hpp | 37 ++--- ggml/src/ggml-webgpu/ggml-webgpu.cpp | 139 +++++++++--------- .../ggml-webgpu/wgsl-shaders/ssm_scan.wgsl | 79 ++++------ 3 files changed, 109 insertions(+), 146 deletions(-) diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 47a266d7de..778a1b4bf7 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 1ebff43f38..dd806ab99b 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; - 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 }); - } + + 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, a_overlap, ids_overlap); - auto * decisions = static_cast(pipeline.context.get()); - xbc_overlap = decisions->xbc_overlap; - a_overlap = decisions->a_overlap; - ids_overlap = decisions->ids_overlap; + webgpu_pipeline pipeline = ctx->shader_lib->get_ssm_scan_pipeline(shader_lib_ctx, xbc_overlap); + auto * decisions = static_cast(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)); - size_t xbc_bind_offset = 0; - size_t xbc_bind_size = 0; - if (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 == 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); + } + + 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 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 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 57f012ad0f..d031a985f7 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 s_in: array; + +// binding for x/B/C merged status #ifdef XBC_OVERLAP -#ifdef IDS_OVERLAP -@group(0) @binding(1) var x_dt_B_C_ids_merged: array; -#ifdef A_OVERLAP -@group(0) @binding(2) var dst: array; -@group(0) @binding(3) var params: Params; -#else -@group(0) @binding(2) var A: array; -@group(0) @binding(3) var dst: array; -@group(0) @binding(4) var params: Params; -#endif -#else -@group(0) @binding(1) var x_dt_B_C_merged: array; -#ifdef A_OVERLAP -@group(0) @binding(2) var ids: array; -@group(0) @binding(3) var dst: array; -@group(0) @binding(4) var params: Params; -#else -@group(0) @binding(2) var A: array; -@group(0) @binding(3) var ids: array; -@group(0) @binding(4) var dst: array; -@group(0) @binding(5) var params: Params; -#endif -#endif +@group(0) @binding(1) var merged: array; +#define BIND_DT 2 +#elif defined(BC_OVERLAP) +@group(0) @binding(1) var merged: array; +@group(0) @binding(2) var x: array; +#define BIND_DT 3 +#elif defined(XB_OVERLAP) +@group(0) @binding(1) var merged: array; +@group(0) @binding(2) var C: array; +#define BIND_DT 3 #else @group(0) @binding(1) var x: array; -@group(0) @binding(2) var dt: array; -@group(0) @binding(3) var A: array; -@group(0) @binding(4) var B: array; -@group(0) @binding(5) var C: array; -@group(0) @binding(6) var ids: array; -@group(0) @binding(7) var dst: array; -@group(0) @binding(8) var params: Params; +@group(0) @binding(2) var B: array; +@group(0) @binding(3) var C: array; +#define BIND_DT 4 #endif +@group(0) @binding(BIND_DT) var dt: array; +@group(0) @binding(BIND_DT + 1) var A: array; +@group(0) @binding(BIND_DT + 2) var ids: array; +@group(0) @binding(BIND_DT + 3) var dst: array; +@group(0) @binding(BIND_DT + 4) var params: Params; + var shared_x_dt: array; var shared_dtsp: array; var shared_reduce: array; @@ -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(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];