From ba7ac35a9ce24947092f859fbf4f6783f00fbf31 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sat, 26 Sep 2026 16:58:53 +0100
Subject: [PATCH 10/18] vulkan : write gated delta net state straight into the
 recurrent cache

Assisted-by: Claude Code (Claude Opus 5.5)
---
 .../ggml-vulkan/ggml-vulkan-push-constants.h  |   1 +
 ggml/src/ggml-vulkan/ggml-vulkan-types.h      |   4 +
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          | 106 +++++++++++++++++-
 .../vulkan-shaders/gated_delta_net.comp       |   9 +-
 4 files changed, 111 insertions(+), 9 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
index f066d1064..48c060e23 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
@@ -699,6 +699,7 @@ struct vk_op_gated_delta_net_push_constants {
     uint32_t neq1, rq3;
     float scale;
     uint32_t K;
+    uint32_t snap_stride;
 };
 
 struct vk_op_ssm_scan_push_constants {
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
index 252359bf8..21cdd044e 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
@@ -74,6 +74,7 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV {
 #include <set>
 
 #include <unordered_map>
+#include <unordered_set>
 
 #include <shared_mutex>
 
@@ -1276,6 +1277,9 @@ struct ggml_backend_vk_context {
     // number of additional consecutive nodes that are being fused with the
     // node currently being processed
     int num_additional_fused_ops {};
+    // gated_delta_net state copies elided because the gdn writes the cache directly (per graph)
+    std::unordered_set<const ggml_tensor *> gdn_elided_cpys;
+    const ggml_tensor * gdn_state_dst {};
     // Bitmask of which fused ops need to write an intermediate value to memory.
     // Bit 'i' means nodes[start_of_fusion + i] writes to memory.
     // If there's no fusion, bit 0 is still set.
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 11485b4e9..ae295b3de 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -3720,7 +3720,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
 
             for (uint32_t kda = 0; kda < 2; kda++) {
                 ggml_vk_create_pipeline(device, device->pipeline_gated_delta_net[si][kda],
-                    gdn_names[si][kda], gdn_len, gdn_data, "main", 7, sizeof(vk_op_gated_delta_net_push_constants),
+                    gdn_names[si][kda], gdn_len, gdn_data, "main", 8, sizeof(vk_op_gated_delta_net_push_constants),
                     wg_denoms, {S_V, kda, device->subgroup_size, lanes_per_column}, 1, true, use_subgroup_ops, device->subgroup_size);
             }
         }
@@ -10176,7 +10176,8 @@ void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx
         pc, {dispatch_x, dispatch_y, 1});
 }
 
-void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) {
+// state_dst: cache view to write the state snapshots to (fused state copy), or nullptr to write them to the tail of dst.
+void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst, const ggml_tensor * state_dst) {
     const ggml_tensor * src_q     = dst->src[0];
     const ggml_tensor * src_v     = dst->src[2];
     const ggml_tensor * src_beta  = dst->src[4];
@@ -10191,7 +10192,8 @@ void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx,
     // K (snapshot slot count) is an op param; state holds s0 only [S_v, S_v, H, n_seqs].
     const uint32_t K = (uint32_t)ggml_get_op_params_i32(dst, 0);
 
-    const uint32_t s_off = S_v * H * n_tokens * n_seqs;
+    uint32_t s_off       = S_v * H * n_tokens * n_seqs;
+    uint32_t snap_stride = S_v * S_v * H * n_seqs;
 
     vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, dst->src[0], dst->src[1], dst->src[2], dst, dst->op);
     GGML_ASSERT(pipeline != nullptr);
@@ -10199,6 +10201,12 @@ void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx,
     ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);
 
     vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst);
+    vk_subbuffer state_buf = dst_buf;
+    if (state_dst != nullptr) {
+        state_buf   = ggml_vk_tensor_subbuffer(ctx, state_dst, true);
+        s_off       = get_misalign_bytes(ctx, state_dst) / sizeof(float);
+        snap_stride = K > 1 ? (uint32_t)(state_dst->nb[2] / sizeof(float)) : 0;
+    }
     vk_subbuffer src_buf[6] = {};
     for (int i = 0; i < 6; i++) {
         src_buf[i] = ggml_vk_tensor_subbuffer(ctx, dst->src[i]);
@@ -10225,11 +10233,12 @@ void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx,
         sb1, sb2, sb3,
         neq1, rq3,
         scale,
-        K
+        K,
+        snap_stride
     };
 
     ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
-        {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], src_buf[5], dst_buf},
+        {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], src_buf[5], dst_buf, state_buf},
         pc, { H, n_seqs, S_v });
 }
 
@@ -12156,6 +12165,9 @@ bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgraph, in
     if ((node->flags & GGML_TENSOR_FLAG_COMPUTE) == 0) {
         return false;
     }
+    if (node->op == GGML_OP_CPY && !ctx->gdn_elided_cpys.empty() && ctx->gdn_elided_cpys.count(node)) {
+        return false; // written directly by the gated_delta_net
+    }
 
     VK_LOG_DEBUG("ggml_vk_build_graph(" << node << ", " << ggml_op_name(node->op) << ")");
     ctx->semaphore_idx = 0;
@@ -12665,7 +12677,7 @@ bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgraph, in
         break;
 
     case GGML_OP_GATED_DELTA_NET:
-        ggml_vk_gated_delta_net(ctx, compute_ctx, node);
+        ggml_vk_gated_delta_net(ctx, compute_ctx, node, ctx->gdn_state_dst);
 
         break;
 
@@ -13951,6 +13963,79 @@ bool ggml_vk_can_fuse_rms_norm_set_rows(ggml_backend_vk_context * ctx, const str
     return true;
 }
 
+bool ggml_vk_tensors_overlap(const ggml_tensor * a, const ggml_tensor * b, bool elementwise);
+
+// Find the cpy that scatters gated_delta_net state snapshots into the recurrent cache, so the shader writes the cache and the cpy is skipped.
+// Graph optimization can put other nodes in between, so this is not a contiguous fusion. Returns the cpy node index, or 0.
+static int ggml_vk_find_gdn_state_cpy(const struct ggml_cgraph * cgraph, int node_idx) {
+    const ggml_tensor * gdn = cgraph->nodes[node_idx];
+    if (gdn->op != GGML_OP_GATED_DELTA_NET || gdn->type != GGML_TYPE_F32 || (gdn->flags & GGML_TENSOR_FLAG_OUTPUT)) {
+        return 0;
+    }
+
+    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);
+    const int64_t       n_written = std::min<int64_t>(n_tokens, K);
+    const size_t        tail_off  = ggml_row_size(GGML_TYPE_F32, S_v * H * n_tokens * n_seqs);
+
+    // the state copy is the first cpy reading this gdn, within a few nodes
+    int j = node_idx + 1;
+    for (; j < cgraph->n_nodes && j <= node_idx + 8; ++j) {
+        const ggml_tensor * n = cgraph->nodes[j];
+        if (n->op == GGML_OP_CPY && n->src[0]->op == GGML_OP_VIEW && n->src[0]->view_src == gdn) {
+            break;
+        }
+    }
+    if (j >= cgraph->n_nodes || j > node_idx + 8) {
+        return 0;
+    }
+    const ggml_tensor * cpy = cgraph->nodes[j];
+    if (cpy->flags & GGML_TENSOR_FLAG_OUTPUT) {
+        return 0;
+    }
+    const ggml_tensor * src = cpy->src[0];
+    const ggml_tensor * dst = cpy->src[1];
+    if (src->op != GGML_OP_VIEW || src->view_src != gdn || src->view_offs != tail_off || !ggml_is_contiguous(src)) {
+        return 0;
+    }
+    const int64_t expected_ne[GGML_MAX_DIMS] = { D, n_seqs, n_written, 1 };
+    if (dst->type != GGML_TYPE_F32 || dst->buffer == nullptr ||
+        !std::equal(expected_ne, expected_ne + GGML_MAX_DIMS, dst->ne) ||
+        dst->nb[0] != sizeof(float) || dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) {
+        return 0;
+    }
+    // the gdn tail is not written any more, so nothing else may read it
+    for (int k = node_idx + 1; k < cgraph->n_nodes; ++k) {
+        if (k == j) {
+            continue;
+        }
+        for (int s = 0; s < GGML_MAX_SRC; ++s) {
+            const ggml_tensor * t = cgraph->nodes[k]->src[s];
+            if (t && t != src && (t == gdn || t->view_src == gdn) && t->view_offs >= tail_off) {
+                return 0;
+            }
+        }
+    }
+    // the cache is now written before the nodes scheduled between gdn and cpy: they must not touch it
+    for (int k = node_idx + 1; k < j; ++k) {
+        const ggml_tensor * n = cgraph->nodes[k];
+        if (ggml_vk_tensors_overlap(n, dst, false)) {
+            return 0;
+        }
+        for (int s = 0; s < GGML_MAX_SRC; ++s) {
+            if (n->src[s] && ggml_vk_tensors_overlap(n->src[s], dst, false)) {
+                return 0;
+            }
+        }
+    }
+    return j;
+}
+
 bool ggml_vk_can_fuse_snake(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx) {
     GGML_UNUSED(ctx);
     if (!ggml_can_fuse(cgraph, node_idx, snake_pattern)) {
@@ -14146,6 +14231,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
     ctx->prealloc_size_add_rms_partials_offset = 0;
     ctx->do_add_rms_partials = false;
     ctx->do_add_rms_partials_offset_calculation = false;
+    ctx->gdn_elided_cpys.clear();
 
     int last_node = cgraph->n_nodes - 1;
 
@@ -14289,6 +14375,14 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
         ctx->fused_topk_qsa = false;
         ctx->fused_rms_norm_mode = RMS_NORM_COUNT;
         const char *fusion_string {};
+        static const bool gdn_state_fusion_disabled = getenv("GGML_VK_DISABLE_GDN_STATE_FUSION") != nullptr;
+        ctx->gdn_state_dst = nullptr;
+        if (!ctx->device->disable_fusion && !gdn_state_fusion_disabled && cgraph->nodes[i]->op == GGML_OP_GATED_DELTA_NET) {
+            if (const int j_cpy = ggml_vk_find_gdn_state_cpy(cgraph, i); j_cpy > 0) {
+                ctx->gdn_state_dst = cgraph->nodes[j_cpy]->src[1];
+                ctx->gdn_elided_cpys.insert(cgraph->nodes[j_cpy]);
+            }
+        }
         if (!ctx->device->disable_fusion) {
             uint32_t num_adds = ggml_vk_fuse_multi_add(ctx, cgraph, i);
             if (num_adds) {
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp b/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp
index 0e384330b..c9d893680 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp
@@ -32,6 +32,7 @@ layout(push_constant) uniform Parameters {
     uint neq1, rq3;
     float scale;
     uint K;
+    uint snap_stride; // floats between snapshot slots in the state output
 };
 
 layout(binding = 0) readonly  buffer QBuf     { FLOAT_TYPE data_q[];     };
@@ -41,6 +42,8 @@ layout(binding = 3) readonly  buffer GBuf     { FLOAT_TYPE data_g[];     };
 layout(binding = 4) readonly  buffer BetaBuf  { FLOAT_TYPE data_beta[];  };
 layout(binding = 5) readonly  buffer StateBuf { FLOAT_TYPE data_state[]; };
 layout(binding = 6)           buffer DstBuf   { FLOAT_TYPE data_dst[];   };
+// Recurrent state output: the tail of dst, or the cache directly when the state copy is fused.
+layout(binding = 7)           buffer StateOutBuf { FLOAT_TYPE data_state_out[]; };
 
 #if !USE_SUBGROUP_ADD && !USE_SUBGROUP_CLUSTERED
 shared FLOAT_TYPE temp[SUBGROUP_SIZE];
@@ -173,9 +176,9 @@ void main() {
         if (K > 1u) {
             const int target_slot = int(n_tokens) - 1 - int(t);
             if (target_slot >= 0 && target_slot < int(K)) {
-                const uint slot_base = s_off + uint(target_slot) * state_size_per_snap + state_out_base;
+                const uint slot_base = s_off + uint(target_slot) * snap_stride + state_out_base;
                 [[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
-                    data_dst[slot_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
+                    data_state_out[slot_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
                 }
             }
         }
@@ -183,7 +186,7 @@ void main() {
 
     if (K == 1u) {
         [[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
-            data_dst[s_off + state_out_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
+            data_state_out[s_off + state_out_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
         }
     }
 }
-- 
2.55.0

