opencl: port fused ssm_scan kernel (Mamba-2, d_state in {128, 256}) to GPU (#26439)

* opencl: port fused ssm_scan kernel (Mamba-2, d_state in {128, 256})

Fold the fused per-token SSM_SCAN recurrent step from opencl/gdn-qwen36-35b
onto the unified base. Previously SSM_SCAN fell back to CPU here; now scalar-A
Mamba-2 with d_state in {128,256}, all-f32, runs on GPU. Other shapes (incl.
Mamba-1 element-wise A) still fall back. test-backend-ops -o SSM_SCAN passes on
Adreno X2-90. opt-out via GGML_OPENCL_DISABLE_SSM_SCAN=1.

* opencl: cleanup

* opencl: require K == 1

---------

Co-authored-by: Li He <lih@qti.qualcomm.com>
This commit is contained in:
Hongqiang Wang
2026-08-19 13:35:17 -07:00
committed by GitHub
co-authored by Li He
parent cd644c3954
commit b062ba735e
3 changed files with 360 additions and 0 deletions
+143
View File
@@ -866,6 +866,9 @@ struct ggml_backend_opencl_context {
// [size_idx][kda][tgpp] where size_idx: 0=S_V=16, 1=32, 2=64, 3=128; kda: 0 or 1.
// tgpp 0 = TG variant (COLS_PER_LANE_GROUP=1), tgpp 1 = prefill variant (COLS_PER_LANE_GROUP=4).
cl_kernel kernel_gated_delta_net_f32[4][2][2] = {};
cl_kernel kernel_ssm_scan_f32_mamba2_d128 = nullptr;
cl_kernel kernel_ssm_scan_f32_mamba2_d256 = nullptr;
cl_kernel kernel_timestep_embedding;
cl_kernel kernel_gemv_moe_q4_0_f32_ns, kernel_gemm_moe_q4_0_f32_ns, kernel_gemm_moe_q4_0_f32_ns_bin;
cl_kernel kernel_gemm_moe_q8_0_f32_ns;
@@ -3154,6 +3157,24 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
GGML_LOG_CONT(".");
}
// ssm_scan (Mamba-2 fused per-token recurrent step; d_state in {128, 256})
{
#ifdef GGML_OPENCL_EMBED_KERNELS
const std::string kernel_src {
#include "ssm_scan.cl.h"
};
#else
const std::string kernel_src = read_file("ssm_scan.cl");
#endif
cl_program prog =
build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
CL_CHECK((backend_ctx->kernel_ssm_scan_f32_mamba2_d128 = clCreateKernel(prog, "kernel_ssm_scan_f32_mamba2_d128", &err), err));
CL_CHECK((backend_ctx->kernel_ssm_scan_f32_mamba2_d256 = clCreateKernel(prog, "kernel_ssm_scan_f32_mamba2_d256", &err), err));
CL_CHECK(clReleaseProgram(prog));
GGML_LOG_CONT(".");
}
// gated_delta_net: one kernel per (S_V, KDA, tgpp) triple.
{
#ifdef GGML_OPENCL_EMBED_KERNELS
@@ -7301,6 +7322,23 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te
(op->src[0]->type == GGML_TYPE_F16 && op->src[1]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32);
case GGML_OP_SSM_CONV:
return (op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32);
case GGML_OP_SSM_SCAN: {
// Mamba-2 fused per-token scan. Requires src3->ne[0] == 1 (scalar
// A per head); d_state in {128, 256}; all sources f32. Falls back
// to CPU otherwise (incl. Mamba-1 element-wise A).
for (int i = 0; i < 6; ++i) {
if (op->src[i]->type != GGML_TYPE_F32) {
return false;
}
}
if (op->type != GGML_TYPE_F32) {
return false;
}
const int K = ggml_get_op_params_i32(op, 0);
const int d_state = (int) op->src[0]->ne[0];
const bool is_mamba2 = (op->src[3]->ne[0] == 1);
return is_mamba2 && (d_state == 128 || d_state == 256) && (K == 1);
}
case GGML_OP_GATED_DELTA_NET:
{
// Match the Vulkan backend: only F32 -> F32, S_v in {16, 32, 64, 128}.
@@ -12260,6 +12298,103 @@ static void ggml_cl_mean(ggml_backend_t backend, const ggml_tensor * src0, const
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
}
static void ggml_cl_ssm_scan(ggml_backend_t backend, ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0]; // s
const ggml_tensor * src1 = dst->src[1]; // x
const ggml_tensor * src2 = dst->src[2]; // dt
const ggml_tensor * src3 = dst->src[3]; // A
const ggml_tensor * src4 = dst->src[4]; // B
const ggml_tensor * src5 = dst->src[5]; // C
const ggml_tensor * src6 = dst->src[6]; // ids
GGML_ASSERT(src0 && src1 && src2 && src3 && src4 && src5 && src6 && dst);
ggml_backend_opencl_context * backend_ctx = (ggml_backend_opencl_context *) backend->context;
ggml_tensor_extra_cl * e0 = (ggml_tensor_extra_cl *) src0->extra;
ggml_tensor_extra_cl * e1 = (ggml_tensor_extra_cl *) src1->extra;
ggml_tensor_extra_cl * e2 = (ggml_tensor_extra_cl *) src2->extra;
ggml_tensor_extra_cl * e3 = (ggml_tensor_extra_cl *) src3->extra;
ggml_tensor_extra_cl * e4 = (ggml_tensor_extra_cl *) src4->extra;
ggml_tensor_extra_cl * e5 = (ggml_tensor_extra_cl *) src5->extra;
ggml_tensor_extra_cl * e6 = (ggml_tensor_extra_cl *) src6->extra;
ggml_tensor_extra_cl * ed = (ggml_tensor_extra_cl *) dst->extra;
cl_ulong o0 = e0->offset + src0->view_offs;
cl_ulong o1 = e1->offset + src1->view_offs;
cl_ulong o2 = e2->offset + src2->view_offs;
cl_ulong o3 = e3->offset + src3->view_offs;
cl_ulong o4 = e4->offset + src4->view_offs;
cl_ulong o5 = e5->offset + src5->view_offs;
cl_ulong o6 = e6->offset + src6->view_offs;
cl_ulong od = ed->offset + dst->view_offs;
const int d_state = (int) src0->ne[0];
const int head_dim = (int) src0->ne[1];
const int n_head = (int) src1->ne[1];
const int n_group = (int) src4->ne[1];
const int n_tokens = (int) src1->ne[2];
const int n_seqs = (int) src1->ne[3];
// Mirror CPU ref: s_off = ggml_nelements(src1) * sizeof(float)
const cl_ulong s_off_bytes = (cl_ulong) ggml_nelements(src1) * sizeof(float);
cl_kernel kernel = (d_state == 128)
? backend_ctx->kernel_ssm_scan_f32_mamba2_d128
: backend_ctx->kernel_ssm_scan_f32_mamba2_d256;
GGML_ASSERT(kernel != nullptr);
cl_ulong s0_nb2 = src0->nb[2];
cl_ulong s0_nb3 = src0->nb[3];
cl_ulong x_nb2 = src1->nb[2];
cl_ulong x_nb3 = src1->nb[3];
cl_ulong dt_nb1 = src2->nb[1];
cl_ulong dt_nb2 = src2->nb[2];
cl_ulong A_nb1 = src3->nb[1];
cl_ulong B_nb2 = src4->nb[2];
cl_ulong B_nb3 = src4->nb[3];
cl_ulong C_nb2 = src5->nb[2];
cl_ulong C_nb3 = src5->nb[3];
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &e0->data_device));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_ulong), &o0));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &e1->data_device));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_ulong), &o1));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &e2->data_device));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &o2));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_mem), &e3->data_device));
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_ulong), &o3));
CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_mem), &e4->data_device));
CL_CHECK(clSetKernelArg(kernel, 9, sizeof(cl_ulong), &o4));
CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_mem), &e5->data_device));
CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_ulong), &o5));
CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_mem), &e6->data_device));
CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_ulong), &o6));
CL_CHECK(clSetKernelArg(kernel, 14, sizeof(cl_mem), &ed->data_device));
CL_CHECK(clSetKernelArg(kernel, 15, sizeof(cl_ulong), &od));
CL_CHECK(clSetKernelArg(kernel, 16, sizeof(cl_ulong), &s0_nb2));
CL_CHECK(clSetKernelArg(kernel, 17, sizeof(cl_ulong), &s0_nb3));
CL_CHECK(clSetKernelArg(kernel, 18, sizeof(cl_ulong), &x_nb2));
CL_CHECK(clSetKernelArg(kernel, 19, sizeof(cl_ulong), &x_nb3));
CL_CHECK(clSetKernelArg(kernel, 20, sizeof(cl_ulong), &dt_nb1));
CL_CHECK(clSetKernelArg(kernel, 21, sizeof(cl_ulong), &dt_nb2));
CL_CHECK(clSetKernelArg(kernel, 22, sizeof(cl_ulong), &A_nb1));
CL_CHECK(clSetKernelArg(kernel, 23, sizeof(cl_ulong), &B_nb2));
CL_CHECK(clSetKernelArg(kernel, 24, sizeof(cl_ulong), &B_nb3));
CL_CHECK(clSetKernelArg(kernel, 25, sizeof(cl_ulong), &C_nb2));
CL_CHECK(clSetKernelArg(kernel, 26, sizeof(cl_ulong), &C_nb3));
CL_CHECK(clSetKernelArg(kernel, 27, sizeof(cl_ulong), &s_off_bytes));
CL_CHECK(clSetKernelArg(kernel, 28, sizeof(int), &head_dim));
CL_CHECK(clSetKernelArg(kernel, 29, sizeof(int), &n_head));
CL_CHECK(clSetKernelArg(kernel, 30, sizeof(int), &n_group));
CL_CHECK(clSetKernelArg(kernel, 31, sizeof(int), &n_tokens));
size_t global_work_size[] = { (size_t)n_head * head_dim * 64, (size_t)n_seqs, 1 };
size_t local_work_size[] = { 64, 1, 1 };
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
}
static void ggml_cl_ssm_conv(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
GGML_ASSERT(src0);
GGML_ASSERT(src0->extra);
@@ -24746,6 +24881,14 @@ bool ggml_cl_compute_forward(ggml_backend_t backend, struct ggml_tensor * tensor
}
func = ggml_cl_ssm_conv;
break;
case GGML_OP_SSM_SCAN:
if (!any_on_device) {
return false;
}
// SSM_SCAN has 7 source tensors, so it cannot use the standard
// (src0, src1, dst) func signature. Dispatch directly and return.
ggml_cl_ssm_scan(backend, tensor);
return true;
case GGML_OP_GATED_DELTA_NET:
if (!any_on_device) {
return false;