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);