webgpu: fix SSM_SCAN binding aliasing (#29750)

This commit is contained in:
Masashi Yoshimura
2026-10-01 11:11:48 +09:00
committed by GitHub
parent f7b384c1e5
commit 0c1e57098b
3 changed files with 109 additions and 146 deletions
+16 -21
View File
@@ -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;
+66 -73
View File
@@ -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<ggml_webgpu_ssm_scan_shader_decisions *>(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<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));
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<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;
+27 -52
View File
@@ -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];