diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index d1cf78083..47a266d7d 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -81,6 +81,7 @@ struct ggml_webgpu_shader_lib_context { ggml_tensor * src4; ggml_tensor * src5; ggml_tensor * dst; + ggml_tensor * dst_fuse; uint32_t max_wg_size; size_t wg_mem_limit_bytes = 0; @@ -412,12 +413,13 @@ struct ggml_webgpu_im2col_pipeline_key_hash { /** Gated Delta Net **/ struct ggml_webgpu_gated_delta_net_pipeline_key { - int type; - int s_v; - int kda; + int type; + int s_v; + int kda; + bool fused_cache; bool operator==(const ggml_webgpu_gated_delta_net_pipeline_key & other) const { - return type == other.type && s_v == other.s_v && kda == other.kda; + return type == other.type && s_v == other.s_v && kda == other.kda && fused_cache == other.fused_cache; } }; @@ -1865,6 +1867,7 @@ class ggml_webgpu_shader_lib { key.type = context.dst->type; key.s_v = (int) context.src2->ne[0]; key.kda = context.src3->ne[0] == context.src2->ne[0]; + key.fused_cache = context.dst_fuse != nullptr; auto it = gated_delta_net_pipelines.find(key); if (it != gated_delta_net_pipelines.end()) { @@ -1887,6 +1890,11 @@ class ggml_webgpu_shader_lib { variant += "_kda"; } + if (key.fused_cache) { + defines.push_back("FUSED_CACHE"); + variant += "_fused_cache"; + } + defines.push_back("S_V=" + std::to_string(key.s_v) + "u"); defines.push_back("WG_SIZE=" + std::to_string(key.s_v) + "u"); diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 9b494d421..86f0e958a 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -1383,7 +1383,8 @@ static webgpu_encoded_op ggml_webgpu_gated_delta_net(webgpu_context & ctx, ggml_tensor * src3, ggml_tensor * src4, ggml_tensor * src5, - ggml_tensor * dst) { + ggml_tensor * dst, + ggml_tensor * dst_fuse) { ggml_webgpu_shader_lib_context shader_lib_ctx = {}; shader_lib_ctx.src0 = src0; shader_lib_ctx.src1 = src1; @@ -1391,6 +1392,7 @@ static webgpu_encoded_op ggml_webgpu_gated_delta_net(webgpu_context & ctx, shader_lib_ctx.src3 = src3; shader_lib_ctx.src4 = src4; shader_lib_ctx.dst = dst; + shader_lib_ctx.dst_fuse = dst_fuse; shader_lib_ctx.max_wg_size = ctx->global_ctx->capabilities.limits.maxComputeInvocationsPerWorkgroup; webgpu_pipeline pipeline = ctx->shader_lib->get_gated_delta_net_pipeline(shader_lib_ctx); @@ -1426,6 +1428,8 @@ static webgpu_encoded_op ggml_webgpu_gated_delta_net(webgpu_context & ctx, (uint32_t) (src2->ne[3] / src0->ne[3]), K, scale_u32, + dst_fuse ? (uint32_t) (dst_fuse->nb[2] / ggml_type_size(dst_fuse->type)) : 0, + dst_fuse ? (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, dst_fuse) / ggml_type_size(dst_fuse->type)) : 0, }; std::vector entries = { @@ -1435,6 +1439,10 @@ static webgpu_encoded_op ggml_webgpu_gated_delta_net(webgpu_context & ctx, ggml_webgpu_make_tensor_bind_group_entry(ctx, 6, dst), }; + if (dst_fuse) { + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 7, dst_fuse)); + } + return ggml_backend_webgpu_build(ctx, pipeline, params, entries, h, n_seqs); } @@ -3220,6 +3228,67 @@ static bool ggml_webgpu_can_fuse_rms_norm_mul(const struct ggml_cgraph * cgraph, return true; } +static bool ggml_webgpu_can_fuse_gdn_cache(const struct ggml_cgraph * cgraph, int node_idx, int & num_encoded_ops) { + const ggml_tensor * gdn = cgraph->nodes[node_idx]; + + // the kernel skips the snapshot tail, so the gdn output must not be a graph output + if (gdn->op != GGML_OP_GATED_DELTA_NET || gdn->type != GGML_TYPE_F32 || (gdn->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return false; + } + + const ggml_tensor * src_v = gdn->src[2]; + const int64_t S_v = src_v->ne[0]; + const int64_t H = src_v->ne[1]; + const int64_t n_tokens = src_v->ne[2]; + const int64_t n_seqs = src_v->ne[3]; + const int64_t D = S_v * S_v * H; + const int64_t K = ggml_get_op_params_i32(gdn, 0); // snapshot slot count + const int64_t n_written = std::min(n_tokens, K); // newest n_written slots are written + + // snapshot tail starts right after the attention scores + const size_t tail_off = ggml_row_size(GGML_TYPE_F32, S_v * H * n_tokens * n_seqs); + + // snapshot cpy is the first real node after the gdn (skip views/no-ops) + const ggml_tensor * cpy = nullptr; + int cpy_idx = 0; + for (int j = node_idx + 1; j < cgraph->n_nodes && cpy == nullptr; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + if (ggml_op_is_empty(n->op) || ggml_is_empty(n)) { + continue; + } + if (n->op != GGML_OP_CPY || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return false; + } + cpy = n; + cpy_idx = j; + } + if (cpy == nullptr) { + return false; + } + + const ggml_tensor * cpy_src = cpy->src[0]; // view of the gdn snapshot tail + const ggml_tensor * cpy_dst = cpy->src[1]; // cache view the kernel writes to + + // src must be this gdn's snapshot tail (contiguous, at the tail offset) + if (cpy_src->op != GGML_OP_VIEW || cpy_src->view_src != gdn || cpy_src->view_offs != tail_off || + !ggml_is_contiguous(cpy_src)) { + return false; + } + + // dst is the [D, n_seqs, n_written] cache view; require nb[1] == D (the per-seq stride the kernel + // assumes). ggml_cpy pins src to the same element count. + const std::array expected_ne = { D, n_seqs, n_written, 1 }; + if (cpy_dst->op != GGML_OP_VIEW || cpy_dst->type != GGML_TYPE_F32 || cpy_dst->data == nullptr || + !std::equal(expected_ne.begin(), expected_ne.end(), cpy_dst->ne) || + cpy_dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || cpy_dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) { + return false; + } + + num_encoded_ops = cpy_idx - node_idx + 1; + + return true; +} + static webgpu_encoded_op ggml_webgpu_upscale(webgpu_context ctx, ggml_tensor * src, ggml_tensor * dst) { const uint32_t mode_flags = (uint32_t) ggml_get_op_params_i32(dst, 0); std::vector params = { (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src) / ggml_type_size(src->type)), @@ -3358,7 +3427,14 @@ static std::optional ggml_webgpu_encode(webgpu_context ctx, return ggml_webgpu_ssm_scan(ctx, src0, src1, src2, node->src[3], node->src[4], node->src[5], node->src[6], node); case GGML_OP_GATED_DELTA_NET: - return ggml_webgpu_gated_delta_net(ctx, src0, src1, src2, node->src[3], node->src[4], node->src[5], node); + if (ggml_webgpu_can_fuse_gdn_cache(cgraph, node_idx, num_encoded_ops)) { + ggml_tensor * dst_fuse = cgraph->nodes[node_idx + num_encoded_ops - 1]->src[1]; + return ggml_webgpu_gated_delta_net(ctx, src0, src1, src2, node->src[3], node->src[4], node->src[5], + node, dst_fuse); + } else { + return ggml_webgpu_gated_delta_net(ctx, src0, src1, src2, node->src[3], node->src[4], node->src[5], + node, nullptr); + } case GGML_OP_PAD: return ggml_webgpu_pad(ctx, src0, node); case GGML_OP_ARGMAX: diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl index 7d7b34755..6f4b5a31c 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl @@ -19,6 +19,16 @@ var src_state: array; @group(0) @binding(6) var dst: array; +#ifdef FUSED_CACHE +@group(0) @binding(7) +var dst_fuse: array; +#define DST_SNAP dst_fuse +#define PARAMS_BINDING 8 +#else +#define DST_SNAP dst +#define PARAMS_BINDING 7 +#endif + struct Params { h: u32, n_tokens: u32, @@ -41,9 +51,11 @@ struct Params { rq3: u32, K: u32, scale: f32, + dst_fuse_nb2: u32, + dst_fuse_off: u32, }; -@group(0) @binding(7) +@group(0) @binding(PARAMS_BINDING) var params: Params; var sh_k: array; @@ -66,7 +78,14 @@ fn main( // input state holds s0 only [S_v, S_v, H, n_seqs]: per-seq stride is H*D. let state_in_base = (seq_id * params.h + head_id) * state_size; let state_out_base = (seq_id * params.h + head_id) * state_size; + +#ifdef FUSED_CACHE + let state_size_per_snap = params.dst_fuse_nb2; + let snap_off = params.dst_fuse_off; +#else let state_size_per_snap = state_size * params.h * params.n_seqs; + let snap_off = params.s_off; +#endif var state: array; for (var i = 0u; i < S_V; i++) { @@ -131,9 +150,9 @@ fn main( // snapshot slot mapping: slot 0 = most recent state, slot s = s tokens back. let target_slot = i32(params.n_tokens) - 1 - i32(t); if (target_slot >= 0 && target_slot < i32(params.K)) { - let slot_base = params.s_off + u32(target_slot) * state_size_per_snap + state_out_base; + let slot_base = snap_off + u32(target_slot) * state_size_per_snap + state_out_base; for (var i = 0u; i < S_V; i++) { - dst[slot_base + col * S_V + i] = state[i]; + DST_SNAP[slot_base + col * S_V + i] = state[i]; } } } @@ -143,7 +162,7 @@ fn main( if (params.K == 1u) { for (var i = 0u; i < S_V; i++) { - dst[params.s_off + state_out_base + col * S_V + i] = state[i]; + DST_SNAP[snap_off + state_out_base + col * S_V + i] = state[i]; } } }