From 817e5f83eb68a2cf5111ef194ae5b37579363ea7 Mon Sep 17 00:00:00 2001 From: Titaniumtown Date: Wed, 16 Sep 2026 23:51:50 -0700 Subject: [PATCH] sycl: ssm_conv: fuse the SiLU epilogue into the ssm_conv kernel (#28929) --- ggml/src/ggml-sycl/fusion.cpp | 48 +++++ ggml/src/ggml-sycl/ggml-sycl.cpp | 14 ++ ggml/src/ggml-sycl/ssm_conv.cpp | 303 +++++++++++++++++++++++++++---- ggml/src/ggml-sycl/ssm_conv.hpp | 1 + 4 files changed, 327 insertions(+), 39 deletions(-) diff --git a/ggml/src/ggml-sycl/fusion.cpp b/ggml/src/ggml-sycl/fusion.cpp index b5e79bea54..6b1f55f2fb 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 beaba8a4ae..46b1f21590 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 3eafa1a680..a871435186 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 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 +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(idx % n_t); + const int channel = static_cast((idx / n_t) % d_inner); + const int seq = static_cast(idx / (static_cast(n_t) * static_cast(d_inner))); + + const float *s = src_data + + static_cast(seq) * static_cast(src_stride_seq) + + static_cast(channel) * static_cast(src_stride_inner) + + static_cast(token); + + const float *c = weights + static_cast(channel) * static_cast(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(seq) * static_cast(dst_stride_seq) + + static_cast(token) * static_cast(dst_stride_token) + + static_cast(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 +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(d_inner) * static_cast(n_t) * static_cast(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(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(idx % n_t); - const int channel = static_cast((idx / n_t) % d_inner); - const int seq = static_cast(idx / (static_cast(n_t) * static_cast(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(seq) * static_cast(src_stride_seq) - + static_cast(channel) * static_cast(src_stride_inner) - + static_cast(token); - - const float *c = weights + static_cast(channel) * static_cast(d_conv); - - float sumf = 0.0f; - for (int i0 = 0; i0 < d_conv; ++i0) { - sumf += s[i0] * c[i0]; + ssm_conv_element(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 size_t dst_idx = - static_cast(seq) * static_cast(dst_stride_seq) + - static_cast(token) * static_cast(dst_stride_token) + - static_cast(channel); - - dst_data[dst_idx] = sumf; - } - ); - }); + ); + }); + } } -inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { +// 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 +static __dpct_inline__ void ssm_conv_tile( + nd_item<1> it, local_accessor 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(it.get_local_id(0)); + const size_t g = it.get_group(0); + const int tt = static_cast(g % nt_tiles); + const int ct = static_cast((g / nt_tiles) % nc_tiles); + const int seq = static_cast(g / (static_cast(nt_tiles) * nc_tiles)); + const int t0 = tt * TT, c0 = ct * TC; + + 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(seq) * src_stride_seq + + static_cast(c0 + c) * src_stride_inner + tok; + const float *cw = weights + static_cast(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(seq) * dst_stride_seq + + static_cast(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 +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(nt_tiles) * nc_tiles * n_s; + + if constexpr (FUSED) { + q.submit([&](handler &h) { + local_accessor 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(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); + + q.submit([&](handler &h) { + local_accessor 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(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, 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(src0->data); const float *weights = static_cast(src1->data); - float *dst_data = static_cast(dst->data); + const bool apply_silu = silu_dst != nullptr; + float *dst_data = static_cast((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(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 1a8ad05f0c..72c9066232 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);