Commit 817e5f83e for llama.cpp

commit 817e5f83eb68a2cf5111ef194ae5b37579363ea7
Author: Titaniumtown <titaniumtown@proton.me>
Date:   Wed Sep 16 23:51:50 2026 -0700

    sycl: ssm_conv: fuse the SiLU epilogue into the ssm_conv kernel (#28929)

diff --git a/ggml/src/ggml-sycl/fusion.cpp b/ggml/src/ggml-sycl/fusion.cpp
index b5e79bea5..6b1f55f2f 100644
--- a/ggml/src/ggml-sycl/fusion.cpp
+++ b/ggml/src/ggml-sycl/fusion.cpp
@@ -208,5 +208,53 @@ bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializ
         return true;
     }

+    if (ops.size() == 2 && ops.begin()[0] == GGML_OP_SSM_CONV && ops.begin()[1] == GGML_OP_UNARY &&
+        unary_ops.size() == 1 && unary_ops.begin()[0] == GGML_UNARY_OP_SILU) {
+        const ggml_tensor * ssm_conv = cgraph->nodes[node_idx];
+        const ggml_tensor * silu     = cgraph->nodes[node_idx + 1];
+
+        if (ggml_get_unary_op(silu) != unary_ops.begin()[0]) {
+            return false;
+        }
+        if (ssm_conv->type != GGML_TYPE_F32 || silu->type != GGML_TYPE_F32) {
+            return false;
+        }
+        // the fused kernel writes the SiLU output with dense strides, so it must be contiguous
+        if (!ggml_is_contiguous(silu)) {
+            return false;
+        }
+
+        return true;
+    }
+
+    if (ops.size() == 3 && ops.begin()[0] == GGML_OP_SSM_CONV && ops.begin()[1] == GGML_OP_ADD &&
+        ops.begin()[2] == GGML_OP_UNARY && unary_ops.size() == 1 && unary_ops.begin()[0] == GGML_UNARY_OP_SILU) {
+        const ggml_tensor * ssm_conv = cgraph->nodes[node_idx];
+        const ggml_tensor * add      = cgraph->nodes[node_idx + 1];
+        const ggml_tensor * silu     = cgraph->nodes[node_idx + 2];
+
+        if (ggml_get_unary_op(silu) != unary_ops.begin()[0]) {
+            return false;
+        }
+        if (ssm_conv->type != GGML_TYPE_F32 || add->type != GGML_TYPE_F32 || silu->type != GGML_TYPE_F32) {
+            return false;
+        }
+        // the fused kernel writes the SiLU output with dense strides, so it must be contiguous
+        if (!ggml_is_contiguous(silu)) {
+            return false;
+        }
+
+        // ADD must consume ssm_conv's output and broadcast a 1-D channel-wise bias
+        const ggml_tensor * bias = (add->src[0] == ssm_conv) ? add->src[1] : add->src[0];
+        if (bias->type != GGML_TYPE_F32 || !ggml_is_contiguous(bias)) {
+            return false;
+        }
+        if (ggml_nelements(bias) != ssm_conv->ne[0] || bias->ne[0] != ssm_conv->ne[0]) {
+            return false;
+        }
+
+        return true;
+    }
+
     return false;
 }
diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp
index beaba8a4a..46b1f2159 100644
--- a/ggml/src/ggml-sycl/ggml-sycl.cpp
+++ b/ggml/src/ggml-sycl/ggml-sycl.cpp
@@ -6034,6 +6034,20 @@ static void ggml_backend_sycl_graph_compute_impl(ggml_backend_sycl_context * syc
             }
         }

+        if (node->op == GGML_OP_SSM_CONV &&
+            ggml_sycl_can_fuse(cgraph, i, { GGML_OP_SSM_CONV, GGML_OP_ADD, GGML_OP_UNARY }, { GGML_UNARY_OP_SILU })) {
+            ggml_sycl_ssm_conv_fused(*sycl_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]);
+            i += 2;
+            continue;
+        }
+
+        if (node->op == GGML_OP_SSM_CONV &&
+            ggml_sycl_can_fuse(cgraph, i, { GGML_OP_SSM_CONV, GGML_OP_UNARY }, { GGML_UNARY_OP_SILU })) {
+            ggml_sycl_ssm_conv_fused(*sycl_ctx, node, nullptr, cgraph->nodes[i + 1]);
+            i++;
+            continue;
+        }
+
         if (node->op == GGML_OP_MUL_MAT && ggml_sycl_mul_mat_glu_mmvq_fused(*sycl_ctx, cgraph, i)) {
             i += 2;
             continue;
diff --git a/ggml/src/ggml-sycl/ssm_conv.cpp b/ggml/src/ggml-sycl/ssm_conv.cpp
index 3eafa1a68..a87143518 100644
--- a/ggml/src/ggml-sycl/ssm_conv.cpp
+++ b/ggml/src/ggml-sycl/ssm_conv.cpp
@@ -1,11 +1,71 @@
 #include "ssm_conv.hpp"
 #include "common.hpp"
+#include "element_wise.hpp"

 #include <cstdio>

 using namespace sycl;

-static void kernel_ssm_conv(
+// One output element of the conv. DC is d_conv as a compile-time constant (0 keeps the
+// runtime loop); unfused callers pass literal false/nullptr so the epilogue folds away.
+template <int DC>
+static __dpct_inline__ void ssm_conv_element(
+    size_t idx,
+    const float *src_data,
+    const float *weights,
+    float *dst_data,
+    int d_conv,
+    int d_inner,
+    int n_t,
+    int src_stride_inner,
+    int src_stride_seq,
+    int dst_stride_token,
+    int dst_stride_seq,
+    bool apply_silu,
+    const float *bias
+) {
+    // src is token-contiguous per channel, dst is channel-contiguous per token,
+    // so indexing token-fastest coalesces the d_conv loads.
+    const int token   = static_cast<int>(idx % n_t);
+    const int channel = static_cast<int>((idx / n_t) % d_inner);
+    const int seq     = static_cast<int>(idx / (static_cast<size_t>(n_t) * static_cast<size_t>(d_inner)));
+
+    const float *s = src_data
+        + static_cast<size_t>(seq) * static_cast<size_t>(src_stride_seq)
+        + static_cast<size_t>(channel) * static_cast<size_t>(src_stride_inner)
+        + static_cast<size_t>(token);
+
+    const float *c = weights + static_cast<size_t>(channel) * static_cast<size_t>(d_conv);
+
+    float sumf = 0.0f;
+    if constexpr (DC > 0) {
+#pragma unroll
+        for (int i0 = 0; i0 < DC; ++i0) {
+            sumf += s[i0] * c[i0];
+        }
+    } else {
+        for (int i0 = 0; i0 < d_conv; ++i0) {
+            sumf += s[i0] * c[i0];
+        }
+    }
+
+    // fused bias add: the ADD node broadcasts a 1-D channel bias over tokens
+    if (bias != nullptr) {
+        sumf += bias[channel];
+    }
+
+    const size_t dst_idx =
+        static_cast<size_t>(seq) * static_cast<size_t>(dst_stride_seq) +
+        static_cast<size_t>(token) * static_cast<size_t>(dst_stride_token) +
+        static_cast<size_t>(channel);
+
+    dst_data[dst_idx] = apply_silu ? op_silu(sumf) : sumf;
+}
+
+// FUSED=false keeps apply_silu/bias out of the kernel capture list, so the unfused launch
+// takes the pre-fusion argument list; matters at n_t == 1, where the op is launch-bound.
+template <int DC, bool FUSED>
+static void kernel_ssm_conv_impl(
     queue &q,
     const float *src_data,
     const float *weights,
@@ -18,7 +78,9 @@ static void kernel_ssm_conv(
     int src_stride_inner,
     int src_stride_seq,
     int dst_stride_token,
-    int dst_stride_seq
+    int dst_stride_seq,
+    bool apply_silu,
+    const float *bias
 ) {
     const size_t total_work = static_cast<size_t>(d_inner) * static_cast<size_t>(n_t) * static_cast<size_t>(n_s);
     const size_t work_group_size = 256;
@@ -27,53 +89,199 @@ static void kernel_ssm_conv(
     const range<1> global_range(num_work_groups * work_group_size);
     const range<1> local_range(work_group_size);

-    q.submit([&](handler &h) {
-        h.parallel_for(
-            nd_range<1>(global_range, local_range),
-            [=](nd_item<1> item) {
-                const size_t idx = item.get_global_id(0);
-                if (idx >= total_work) {
-                    return;
+    if constexpr (FUSED) {
+        q.submit([&](handler &h) {
+            h.parallel_for(
+                nd_range<1>(global_range, local_range),
+                [=](nd_item<1> item) {
+                    const size_t idx = item.get_global_id(0);
+                    if (idx >= total_work) {
+                        return;
+                    }
+
+                    ssm_conv_element<DC>(idx, src_data, weights, dst_data, d_conv, d_inner, n_t,
+                                         src_stride_inner, src_stride_seq, dst_stride_token,
+                                         dst_stride_seq, apply_silu, bias);
                 }
+            );
+        });
+    } else {
+        GGML_UNUSED(apply_silu);
+        GGML_UNUSED(bias);

-                // src has the tokens of one channel contiguous, dst has the channels of one
-                // token contiguous, so either the loads or the store must be strided. Indexing
-                // token-fastest coalesces the d_conv loads, which measured faster except for
-                // short, cache-resident rows.
-                const int token   = static_cast<int>(idx % n_t);
-                const int channel = static_cast<int>((idx / n_t) % d_inner);
-                const int seq     = static_cast<int>(idx / (static_cast<size_t>(n_t) * static_cast<size_t>(d_inner)));
+        q.submit([&](handler &h) {
+            h.parallel_for(
+                nd_range<1>(global_range, local_range),
+                [=](nd_item<1> item) {
+                    const size_t idx = item.get_global_id(0);
+                    if (idx >= total_work) {
+                        return;
+                    }

-                const float *s = src_data
-                    + static_cast<size_t>(seq) * static_cast<size_t>(src_stride_seq)
-                    + static_cast<size_t>(channel) * static_cast<size_t>(src_stride_inner)
-                    + static_cast<size_t>(token);
+                    ssm_conv_element<DC>(idx, src_data, weights, dst_data, d_conv, d_inner, n_t,
+                                         src_stride_inner, src_stride_seq, dst_stride_token,
+                                         dst_stride_seq, false, nullptr);
+                }
+            );
+        });
+    }
+}

-                const float *c = weights + static_cast<size_t>(channel) * static_cast<size_t>(d_conv);
+// SLM transpose tile: coalesces both the loads and the stores. The +1 pad makes the row
+// stride 33, coprime with 32 banks, so both phases are bank-conflict-free.
+template <int DC, int TT, int TC, int WG>
+static __dpct_inline__ void ssm_conv_tile(
+    nd_item<1> it, local_accessor<float, 1> tile, const float *src_data, const float *weights,
+    float *dst_data, int n_t, int nt_tiles, int nc_tiles, int src_stride_inner,
+    int src_stride_seq, int dst_stride_token, int dst_stride_seq, bool apply_silu,
+    const float *bias
+) {
+    const int    lid = static_cast<int>(it.get_local_id(0));
+    const size_t g   = it.get_group(0);
+    const int    tt  = static_cast<int>(g % nt_tiles);
+    const int    ct  = static_cast<int>((g / nt_tiles) % nc_tiles);
+    const int    seq = static_cast<int>(g / (static_cast<size_t>(nt_tiles) * nc_tiles));
+    const int    t0 = tt * TT, c0 = ct * TC;

-                float sumf = 0.0f;
-                for (int i0 = 0; i0 < d_conv; ++i0) {
-                    sumf += s[i0] * c[i0];
-                }
+    const int ti = lid % TT;
+    const int cj = lid / TT;
+#pragma unroll
+    for (int r = 0; r < TC / (WG / TT); ++r) {
+        const int c   = cj + r * (WG / TT);
+        const int tok = t0 + ti;
+        float sumf = 0.0f;
+        if (tok < n_t) {
+            const float *s = src_data + static_cast<size_t>(seq) * src_stride_seq
+                           + static_cast<size_t>(c0 + c) * src_stride_inner + tok;
+            const float *cw = weights + static_cast<size_t>(c0 + c) * DC;
+#pragma unroll
+            for (int i = 0; i < DC; ++i) sumf += s[i] * cw[i];
+            if (bias != nullptr) sumf += bias[c0 + c];
+            if (apply_silu) sumf = op_silu(sumf);
+        }
+        tile[c * (TT + 1) + ti] = sumf;
+    }
+    it.barrier(access::fence_space::local_space);
+
+    const int cc = lid % TC;
+    const int tj = lid / TC;
+#pragma unroll
+    for (int r = 0; r < TT / (WG / TC); ++r) {
+        const int t   = tj + r * (WG / TC);
+        const int tok = t0 + t;
+        if (tok < n_t) {
+            dst_data[static_cast<size_t>(seq) * dst_stride_seq
+                     + static_cast<size_t>(tok) * dst_stride_token + c0 + cc]
+                = tile[cc * (TT + 1) + t];
+        }
+    }
+}
+
+// Same FUSED split as kernel_ssm_conv_impl. The fused instantiation keeps the runtime
+// apply_silu/bias branches: at n_t >= 32 they are amortized over the whole tile.
+template <int DC, bool FUSED>
+static void kernel_ssm_conv_tiled(
+    queue &q, const float *src_data, const float *weights, float *dst_data,
+    int d_inner, int n_t, int n_s, int src_stride_inner, int src_stride_seq,
+    int dst_stride_token, int dst_stride_seq, bool apply_silu, const float *bias
+) {
+    constexpr int TT = 32, TC = 32, WG = 256;
+    const int nt_tiles = (n_t + TT - 1) / TT;
+    const int nc_tiles = d_inner / TC;
+    const size_t groups = static_cast<size_t>(nt_tiles) * nc_tiles * n_s;

-                const size_t dst_idx =
-                    static_cast<size_t>(seq) * static_cast<size_t>(dst_stride_seq) +
-                    static_cast<size_t>(token) * static_cast<size_t>(dst_stride_token) +
-                    static_cast<size_t>(channel);
+    if constexpr (FUSED) {
+        q.submit([&](handler &h) {
+            local_accessor<float, 1> tile(range<1>(TC * (TT + 1)), h);
+            h.parallel_for(nd_range<1>(range<1>(groups * WG), range<1>(WG)), [=](nd_item<1> it) {
+                ssm_conv_tile<DC, TT, TC, WG>(it, tile, src_data, weights, dst_data, n_t, nt_tiles,
+                                              nc_tiles, src_stride_inner, src_stride_seq,
+                                              dst_stride_token, dst_stride_seq, apply_silu, bias);
+            });
+        });
+    } else {
+        GGML_UNUSED(apply_silu);
+        GGML_UNUSED(bias);

-                dst_data[dst_idx] = sumf;
-            }
-        );
-    });
+        q.submit([&](handler &h) {
+            local_accessor<float, 1> tile(range<1>(TC * (TT + 1)), h);
+            h.parallel_for(nd_range<1>(range<1>(groups * WG), range<1>(WG)), [=](nd_item<1> it) {
+                ssm_conv_tile<DC, TT, TC, WG>(it, tile, src_data, weights, dst_data, n_t, nt_tiles,
+                                              nc_tiles, src_stride_inner, src_stride_seq,
+                                              dst_stride_token, dst_stride_seq, false, nullptr);
+            });
+        });
+    }
+}
+
+static void kernel_ssm_conv(
+    queue &q,
+    const float *src_data,
+    const float *weights,
+    float *dst_data,
+    int d_conv,
+    int d_inner,
+    int n_t,
+    int n_s,
+    int ncs,
+    int src_stride_inner,
+    int src_stride_seq,
+    int dst_stride_token,
+    int dst_stride_seq,
+    bool apply_silu,
+    const float *bias
+) {
+    // Only the fused instantiations carry apply_silu/bias as kernel arguments; the plain
+    // ssm_conv launch keeps the argument list it had before the fusion landed.
+    const bool fused = apply_silu || bias != nullptr;
+
+    // d_inner must be a multiple of 32 so the channel tiles are exact; the transpose is only
+    // worth it for n_t >= 32. d_conv == 4 is the only window with a DC-specialized kernel.
+    if (d_conv == 4 && n_t >= 32 && (d_inner % 32) == 0) {
+        if (fused) {
+            kernel_ssm_conv_tiled<4, true>(q, src_data, weights, dst_data, d_inner, n_t, n_s,
+                                           src_stride_inner, src_stride_seq, dst_stride_token,
+                                           dst_stride_seq, apply_silu, bias);
+        } else {
+            kernel_ssm_conv_tiled<4, false>(q, src_data, weights, dst_data, d_inner, n_t, n_s,
+                                            src_stride_inner, src_stride_seq, dst_stride_token,
+                                            dst_stride_seq, apply_silu, bias);
+        }
+        return;
+    }
+
+    if (d_conv == 4) {
+        if (fused) {
+            kernel_ssm_conv_impl<4, true>(q, src_data, weights, dst_data, d_conv, d_inner, n_t, n_s,
+                                          ncs, src_stride_inner, src_stride_seq, dst_stride_token,
+                                          dst_stride_seq, apply_silu, bias);
+        } else {
+            kernel_ssm_conv_impl<4, false>(q, src_data, weights, dst_data, d_conv, d_inner, n_t, n_s,
+                                           ncs, src_stride_inner, src_stride_seq, dst_stride_token,
+                                           dst_stride_seq, apply_silu, bias);
+        }
+        return;
+    }
+
+    if (fused) {
+        kernel_ssm_conv_impl<0, true>(q, src_data, weights, dst_data, d_conv, d_inner, n_t, n_s,
+                                      ncs, src_stride_inner, src_stride_seq, dst_stride_token,
+                                      dst_stride_seq, apply_silu, bias);
+    } else {
+        kernel_ssm_conv_impl<0, false>(q, src_data, weights, dst_data, d_conv, d_inner, n_t, n_s,
+                                       ncs, src_stride_inner, src_stride_seq, dst_stride_token,
+                                       dst_stride_seq, apply_silu, bias);
+    }
 }

-inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
+inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst, ggml_tensor * silu_dst = nullptr, const float * bias = nullptr) {
     ggml_tensor * src0 = dst->src[0];
     ggml_tensor * src1 = dst->src[1];

     GGML_ASSERT(src0->type == GGML_TYPE_F32);
     GGML_ASSERT(src1->type == GGML_TYPE_F32);
     GGML_ASSERT(dst->type  == GGML_TYPE_F32);
+    GGML_ASSERT(bias == nullptr || silu_dst != nullptr);

     const int d_conv   = src1->ne[0];
     const int ncs      = src0->ne[0];
@@ -104,7 +312,8 @@ inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor *

         const float *src_data = static_cast<const float *>(src0->data);
         const float *weights  = static_cast<const float *>(src1->data);
-        float *dst_data       = static_cast<float *>(dst->data);
+        const bool apply_silu = silu_dst != nullptr;
+        float *dst_data       = static_cast<float *>((silu_dst ? silu_dst : dst)->data);

         GGML_ASSERT(src_data && weights && dst_data);

@@ -121,7 +330,9 @@ inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor *
             src_stride_inner,
             src_stride_seq,
             dst_stride_token,
-            dst_stride_seq
+            dst_stride_seq,
+            apply_silu,
+            bias
         );

     } catch (const std::exception &e) {
@@ -134,3 +345,17 @@ void ggml_sycl_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
     scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
     ggml_sycl_op_ssm_conv(ctx, dst);
 }
+
+// Fused ssm_conv + ADD + SiLU: write silu(conv(x) + b) straight into silu_dst, eliding the
+// standalone SiLU launch and its HBM round-trip of the conv output.
+void ggml_sycl_ssm_conv_fused(ggml_backend_sycl_context & ctx, ggml_tensor * dst, ggml_tensor * add, ggml_tensor * silu_dst) {
+    scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
+    GGML_ASSERT(silu_dst && ggml_are_same_shape(dst, silu_dst) && silu_dst->type == GGML_TYPE_F32);
+    // the fused kernel reads only the ADD's bias operand; the ADD result is never written
+    const float * bias = nullptr;
+    if (add != nullptr) {
+        const ggml_tensor * bias_t = (add->src[0] == dst) ? add->src[1] : add->src[0];
+        bias = static_cast<const float *>(bias_t->data);
+    }
+    ggml_sycl_op_ssm_conv(ctx, dst, silu_dst, bias);
+}
diff --git a/ggml/src/ggml-sycl/ssm_conv.hpp b/ggml/src/ggml-sycl/ssm_conv.hpp
index 1a8ad05f0..72c906623 100644
--- a/ggml/src/ggml-sycl/ssm_conv.hpp
+++ b/ggml/src/ggml-sycl/ssm_conv.hpp
@@ -3,3 +3,4 @@
 #include "common.hpp"

 void ggml_sycl_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
+void ggml_sycl_ssm_conv_fused(ggml_backend_sycl_context & ctx, ggml_tensor * dst, ggml_tensor * add, ggml_tensor * silu_dst);