From f7c8377d393f36ffe6fc72affe5005e51aeebe31 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sat, 26 Sep 2026 14:59:08 +0100
Subject: [PATCH 08/18] cuda : use 8 lanes per state column in gated delta net
 for S_v >= 128

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-cuda/gated_delta_net.cu | 54 +++++++++++++++------------
 1 file changed, 31 insertions(+), 23 deletions(-)

diff --git a/ggml/src/ggml-cuda/gated_delta_net.cu b/ggml/src/ggml-cuda/gated_delta_net.cu
index 1b431a724..0db2ae0c9 100644
--- a/ggml/src/ggml-cuda/gated_delta_net.cu
+++ b/ggml/src/ggml-cuda/gated_delta_net.cu
@@ -1,7 +1,9 @@
 #include "gated_delta_net.cuh"
 #include "ggml-cuda/common.cuh"
 
-template <int S_v, bool KDA, bool keep_rs_t>
+// lpc lanes own each state column (lpc_req == 0: the whole warp), each lane holds S_v/lpc rows.
+// A warp handles warp_size/lpc columns. Fewer lanes per column means shorter reductions.
+template <int S_v, int lpc_req, bool KDA, bool keep_rs_t>
 __global__ void __launch_bounds__((ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v) * 4, 2)
 gated_delta_net_cuda(const float * q,
                                      const float * k,
@@ -30,9 +32,13 @@ gated_delta_net_cuda(const float * q,
                                      int           K) {
     const uint32_t h_idx    = blockIdx.x;
     const uint32_t sequence = blockIdx.y;
-    // each warp owns one column, using warp-level primitives to reduce across rows
-    const int      lane     = threadIdx.x;
-    const int      col      = blockIdx.z * blockDim.y + threadIdx.y;
+    constexpr int warp_size     = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v;
+    constexpr int lpc           = lpc_req == 0 ? warp_size : lpc_req; // 0: whole warp per column
+    constexpr int cols_per_warp = warp_size / lpc;
+    static_assert(warp_size % lpc == 0 && S_v % lpc == 0, "bad lanes per column");
+    // groups of lpc lanes own one column, using warp-level primitives to reduce across rows
+    const int      lane     = threadIdx.x % lpc;
+    const int      col      = (blockIdx.z * blockDim.y + threadIdx.y) * cols_per_warp + threadIdx.x / lpc;
 
     const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic);
     const uint32_t iq3 = fastdiv(sequence, rq3_magic);
@@ -47,16 +53,14 @@ gated_delta_net_cuda(const float * q,
     curr_state += state_in_offset + col * S_v;
     attn_data += (sequence * n_tokens * H + h_idx) * S_v;
 
-    constexpr int warp_size = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v;
-    static_assert(S_v % warp_size == 0, "S_v must be a multiple of warp_size");
-    constexpr int rows_per_lane = (S_v + warp_size - 1) / warp_size;
+    constexpr int rows_per_lane = S_v / lpc;
     float         s_shard[rows_per_lane];
     // state is stored transposed: M[col][i] = S[i][col], row col is contiguous
 
     ggml_cuda_pdl_sync();
 #pragma unroll
     for (int r = 0; r < rows_per_lane; r++) {
-        const int i = r * warp_size + lane;
+        const int i = r * lpc + lane;
         s_shard[r]  = curr_state[i];
     }
 
@@ -76,7 +80,7 @@ gated_delta_net_cuda(const float * q,
         float q_reg[rows_per_lane];
 #pragma unroll
         for (int r = 0; r < rows_per_lane; r++) {
-            const int i = r * warp_size + lane;
+            const int i = r * lpc + lane;
             k_reg[r] = k_t[i];
             q_reg[r] = q_t[i];
         }
@@ -90,7 +94,7 @@ gated_delta_net_cuda(const float * q,
             for (int r = 0; r < rows_per_lane; r++) {
                 kv_shard += s_shard[r] * k_reg[r];
             }
-            float kv_col = warp_reduce_sum<warp_size>(kv_shard);
+            float kv_col = warp_reduce_sum<lpc>(kv_shard);
 
             // delta[col] = (v[col] - g * kv[col]) * beta
             float delta_col = (v_t[col] - g_val * kv_col) * beta_val;
@@ -104,7 +108,7 @@ gated_delta_net_cuda(const float * q,
                 attn_partial += s_shard[r] * q_reg[r];
             }
 
-            float attn_col = warp_reduce_sum<warp_size>(attn_partial);
+            float attn_col = warp_reduce_sum<lpc>(attn_partial);
 
             if (lane == 0) {
                 attn_data[col] = attn_col * scale;
@@ -114,11 +118,11 @@ gated_delta_net_cuda(const float * q,
             float kv_shard = 0.0f;
 #pragma unroll
             for (int r = 0; r < rows_per_lane; r++) {
-                const int i = r * warp_size + lane;
+                const int i = r * lpc + lane;
                 kv_shard += expf(g_t[i]) * s_shard[r] * k_reg[r];
             }
 
-            float kv_col = warp_reduce_sum<warp_size>(kv_shard);
+            float kv_col = warp_reduce_sum<lpc>(kv_shard);
 
             // delta[col] = (v[col] - kv[col]) * beta
             float delta_col = (v_t[col] - kv_col) * beta_val;
@@ -128,12 +132,12 @@ gated_delta_net_cuda(const float * q,
             float attn_partial = 0.0f;
 #pragma unroll
             for (int r = 0; r < rows_per_lane; r++) {
-                const int i = r * warp_size + lane;
+                const int i = r * lpc + lane;
                 s_shard[r]  = expf(g_t[i]) * s_shard[r] + k_reg[r] * delta_col;
                 attn_partial += s_shard[r] * q_reg[r];
             }
 
-            float attn_col = warp_reduce_sum<warp_size>(attn_partial);
+            float attn_col = warp_reduce_sum<lpc>(attn_partial);
 
             if (lane == 0) {
                 attn_data[col] = attn_col * scale;
@@ -150,7 +154,7 @@ gated_delta_net_cuda(const float * q,
                 float * curr_state = state + target_slot * state_slot_stride;
 #pragma unroll
                 for (int r = 0; r < rows_per_lane; r++) {
-                    const int i = r * warp_size + lane;
+                    const int i = r * lpc + lane;
                     curr_state[col * S_v + i] = s_shard[r];
                 }
             }
@@ -160,7 +164,7 @@ gated_delta_net_cuda(const float * q,
     if constexpr (!keep_rs_t) {
 #pragma unroll
         for (int r = 0; r < rows_per_lane; r++) {
-            const int i          = r * warp_size + lane;
+            const int i          = r * lpc + lane;
             state[col * S_v + i] = s_shard[r];
         }
     }
@@ -180,8 +184,12 @@ static void launch_gated_delta_net(
     //TODO: Add chunked kernel for even faster pre-fill
     const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
     const int num_warps = 4;
-    dim3      grid_dims(H, n_seqs, (S_v + num_warps - 1) / num_warps);
-    dim3      block_dims(warp_size <= S_v ? warp_size : S_v, num_warps, 1);
+    const int warp_w    = warp_size <= S_v ? warp_size : S_v;
+    // S_v >= 128: 8 lanes per column (16 rows each), otherwise one column per warp
+    const int lpc       = S_v >= 128 ? 8 : warp_w;
+    const int cols_per_block = num_warps * (warp_w / lpc);
+    dim3      grid_dims(H, n_seqs, (S_v + cols_per_block - 1) / cols_per_block);
+    dim3      block_dims(warp_w, num_warps, 1);
 
     const uint3 neqk1_magic = init_fastdiv_values(neqk1);
     const uint3 rq3_magic   = init_fastdiv_values(rq3);
@@ -189,26 +197,26 @@ static void launch_gated_delta_net(
     const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, stream);
     switch (S_v) {
         case 16:
-            ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t>, launch_params,
+            ggml_cuda_kernel_launch(gated_delta_net_cuda<16, 0, KDA, keep_rs_t>, launch_params,
                 q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
                 n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
                 sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
             break;
         case 32:
-            ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t>, launch_params,
+            ggml_cuda_kernel_launch(gated_delta_net_cuda<32, 0, KDA, keep_rs_t>, launch_params,
                 q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
                 n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
                 sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
             break;
         case 64: {
-            ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t>, launch_params,
+            ggml_cuda_kernel_launch(gated_delta_net_cuda<64, 0, KDA, keep_rs_t>, launch_params,
                 q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
                 n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
                 sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
             break;
         }
         case 128: {
-            ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t>, launch_params,
+            ggml_cuda_kernel_launch(gated_delta_net_cuda<128, 8, KDA, keep_rs_t>, launch_params,
                 q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
                 n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
                 sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
-- 
2.55.0

