Commit 6ec1a7e95 for llama.cpp
commit 6ec1a7e956cfd5dfc111b6d3fa8e7d2c219106db
Author: lhez <lih@qti.qualcomm.com>
Date: Tue Sep 15 02:29:02 2026 -0700
opencl: add generic ssm_scan (#28881)
* opencl: add generic ssm_scan
* opencl: fix whitespace
diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp
index 39c592e88..5b99f5d00 100644
--- a/ggml/src/ggml-opencl/ggml-opencl.cpp
+++ b/ggml/src/ggml-opencl/ggml-opencl.cpp
@@ -982,6 +982,7 @@ 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 = nullptr;
cl_kernel kernel_ssm_scan_f32_mamba2_d128 = nullptr;
cl_kernel kernel_ssm_scan_f32_mamba2_d256 = nullptr;
@@ -3457,7 +3458,7 @@ 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})
+ // ssm_scan
{
#ifdef GGML_OPENCL_EMBED_KERNELS
const std::string kernel_src {
@@ -3469,8 +3470,34 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
cl_program prog =
build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
+ CL_CHECK((backend_ctx->kernel_ssm_scan_f32 = clCreateKernel(prog, "kernel_ssm_scan_f32", &err), err));
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_kernel * kernels[] = {
+ &backend_ctx->kernel_ssm_scan_f32_mamba2_d128,
+ &backend_ctx->kernel_ssm_scan_f32_mamba2_d256
+ };
+
+ // specialized kernels use subgroups and assume subgroup size is 64,
+ // if device does not support subgroups or subgroup size is not 64,
+ // release these kernels
+ for (int i = 0; i < 2; ++i) {
+ size_t subgroup_size = 0;
+#if CL_TARGET_OPENCL_VERSION >= 210
+ const size_t local_work_size[] = { 64, 1 };
+ const cl_int subgroup_err = clGetKernelSubGroupInfo(*kernels[i], backend_ctx->device, CL_KERNEL_MAX_SUB_GROUP_SIZE_FOR_NDRANGE,
+ sizeof(local_work_size), local_work_size, sizeof(subgroup_size), &subgroup_size, nullptr);
+ if (subgroup_err != CL_SUCCESS) {
+ subgroup_size = 0;
+ }
+#endif
+ // The specialized kernels reduce over one 64-lane subgroup.
+ if (subgroup_size != 64) {
+ CL_CHECK(clReleaseKernel(*kernels[i]));
+ *kernels[i] = nullptr;
+ }
+ }
CL_CHECK(clReleaseProgram(prog));
GGML_LOG_CONT(".");
}
@@ -8734,22 +8761,16 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te
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) {
+ if (op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 ||
+ op->src[1]->type != GGML_TYPE_F32 || op->src[2]->type != GGML_TYPE_F32 ||
+ op->src[3]->type != GGML_TYPE_F32 || op->src[4]->type != GGML_TYPE_F32 ||
+ op->src[5]->type != GGML_TYPE_F32 || op->src[6]->type != GGML_TYPE_I32) {
return false;
}
+
+ const int64_t d_state = op->src[0]->ne[0];
+ return d_state >= 1 && d_state <= 256 && (d_state & (d_state - 1)) == 0;
}
- 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}.
@@ -14043,81 +14064,109 @@ static void ggml_cl_mean(ggml_backend_t backend, const ggml_tensor * src0, const
}
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_ASSERT(dst);
+ GGML_ASSERT(dst->extra);
+ GGML_ASSERT(dst->src[0]);
+ GGML_ASSERT(dst->src[0]->extra);
+ GGML_ASSERT(dst->src[1]);
+ GGML_ASSERT(dst->src[1]->extra);
+ GGML_ASSERT(dst->src[2]);
+ GGML_ASSERT(dst->src[2]->extra);
+ GGML_ASSERT(dst->src[3]);
+ GGML_ASSERT(dst->src[3]->extra);
+ GGML_ASSERT(dst->src[4]);
+ GGML_ASSERT(dst->src[4]->extra);
+ GGML_ASSERT(dst->src[5]);
+ GGML_ASSERT(dst->src[5]->extra);
+ GGML_ASSERT(dst->src[6]);
+ GGML_ASSERT(dst->src[6]->extra);
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);
+ ggml_tensor_extra_cl * extra0 = (ggml_tensor_extra_cl *) dst->src[0]->extra;
+ ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *) dst->src[1]->extra;
+ ggml_tensor_extra_cl * extra2 = (ggml_tensor_extra_cl *) dst->src[2]->extra;
+ ggml_tensor_extra_cl * extra3 = (ggml_tensor_extra_cl *) dst->src[3]->extra;
+ ggml_tensor_extra_cl * extra4 = (ggml_tensor_extra_cl *) dst->src[4]->extra;
+ ggml_tensor_extra_cl * extra5 = (ggml_tensor_extra_cl *) dst->src[5]->extra;
+ ggml_tensor_extra_cl * extra6 = (ggml_tensor_extra_cl *) dst->src[6]->extra;
+ ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *) dst->extra;
+
+ const cl_ulong offset0 = extra0->offset + dst->src[0]->view_offs;
+ const cl_ulong offset1 = extra1->offset + dst->src[1]->view_offs;
+ const cl_ulong offset2 = extra2->offset + dst->src[2]->view_offs;
+ const cl_ulong offset3 = extra3->offset + dst->src[3]->view_offs;
+ const cl_ulong offset4 = extra4->offset + dst->src[4]->view_offs;
+ const cl_ulong offset5 = extra5->offset + dst->src[5]->view_offs;
+ const cl_ulong offset6 = extra6->offset + dst->src[6]->view_offs;
+ const cl_ulong offsetd = extrad->offset + dst->view_offs;
+
+ const ggml_tensor * s = dst->src[0];
+ const ggml_tensor * x = dst->src[1];
+ const ggml_tensor * dt = dst->src[2];
+ const ggml_tensor * A = dst->src[3];
+ const ggml_tensor * B = dst->src[4];
+ const ggml_tensor * C = dst->src[5];
+
+ const cl_ulong s_nb1 = s->nb[1];
+ const cl_ulong s_nb2 = s->nb[2];
+ const cl_ulong s_nb3 = s->nb[3];
+ const cl_ulong x_nb1 = x->nb[1];
+ const cl_ulong x_nb2 = x->nb[2];
+ const cl_ulong x_nb3 = x->nb[3];
+ const cl_ulong dt_nb1 = dt->nb[1];
+ const cl_ulong dt_nb2 = dt->nb[2];
+ const cl_ulong A_nb1 = A->nb[1];
+ const cl_ulong B_nb1 = B->nb[1];
+ const cl_ulong B_nb2 = B->nb[2];
+ const cl_ulong B_nb3 = B->nb[3];
+ const cl_ulong C_nb1 = C->nb[1];
+ const cl_ulong C_nb2 = C->nb[2];
+ const cl_ulong C_nb3 = C->nb[3];
+
+ const cl_uint A_ne0 = A->ne[0];
+ const cl_uint d_state = s->ne[0];
+ const cl_int head_dim = x->ne[0];
+ const cl_int n_head = x->ne[1];
+ const cl_int n_group = B->ne[1];
+ const cl_int n_tokens = x->ne[2];
+ const cl_uint n_seqs = x->ne[3];
+ const cl_uint K = ggml_get_op_params_i32(dst, 0);
+ const cl_ulong s_off_bytes = (cl_ulong) ggml_nelements(x) * sizeof(float);
+
+ cl_kernel kernel = backend_ctx->kernel_ssm_scan_f32;
+ size_t nth = d_state;
+ if (A_ne0 == 1 && K == 1) {
+ cl_kernel kernel_mamba2 = nullptr;
+ if (d_state == 128) {
+ kernel_mamba2 = backend_ctx->kernel_ssm_scan_f32_mamba2_d128;
+ } else if (d_state == 256) {
+ kernel_mamba2 = backend_ctx->kernel_ssm_scan_f32_mamba2_d256;
+ }
+ if (kernel_mamba2 != nullptr) {
+ kernel = kernel_mamba2;
+ nth = 64;
+ }
+ }
- 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, 0, sizeof(cl_mem), &extra0->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_ulong), &offset0));
+ CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra1->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_ulong), &offset1));
+ CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &extra2->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &offset2));
+ CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_mem), &extra3->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_ulong), &offset3));
+ CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_mem), &extra4->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 9, sizeof(cl_ulong), &offset4));
+ CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_mem), &extra5->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_ulong), &offset5));
+ CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_mem), &extra6->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_ulong), &offset6));
+ CL_CHECK(clSetKernelArg(kernel, 14, sizeof(cl_mem), &extrad->data_device));
+ CL_CHECK(clSetKernelArg(kernel, 15, sizeof(cl_ulong), &offsetd));
+ CL_CHECK(clSetKernelArg(kernel, 16, sizeof(cl_ulong), &s_nb2));
+ CL_CHECK(clSetKernelArg(kernel, 17, sizeof(cl_ulong), &s_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));
@@ -14128,15 +14177,30 @@ static void ggml_cl_ssm_scan(ggml_backend_t backend, ggml_tensor * dst) {
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));
+ CL_CHECK(clSetKernelArg(kernel, 28, sizeof(cl_int), &head_dim));
+ CL_CHECK(clSetKernelArg(kernel, 29, sizeof(cl_int), &n_head));
+ CL_CHECK(clSetKernelArg(kernel, 30, sizeof(cl_int), &n_group));
+ CL_CHECK(clSetKernelArg(kernel, 31, sizeof(cl_int), &n_tokens));
+
+ if (kernel == backend_ctx->kernel_ssm_scan_f32) {
+ CL_CHECK(clSetKernelArg(kernel, 32, sizeof(cl_ulong), &s_nb1));
+ CL_CHECK(clSetKernelArg(kernel, 33, sizeof(cl_ulong), &x_nb1));
+ CL_CHECK(clSetKernelArg(kernel, 34, sizeof(cl_ulong), &B_nb1));
+ CL_CHECK(clSetKernelArg(kernel, 35, sizeof(cl_ulong), &C_nb1));
+ CL_CHECK(clSetKernelArg(kernel, 36, sizeof(cl_uint), &A_ne0));
+ CL_CHECK(clSetKernelArg(kernel, 37, sizeof(cl_uint), &d_state));
+ CL_CHECK(clSetKernelArg(kernel, 38, sizeof(cl_uint), &n_seqs));
+ CL_CHECK(clSetKernelArg(kernel, 39, sizeof(cl_uint), &K));
+ CL_CHECK(clSetKernelArg(kernel, 40, d_state * sizeof(float), nullptr));
+ }
- 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 };
+ size_t global_work_size[] = {
+ (size_t) head_dim * (size_t) n_head * nth,
+ (size_t) n_seqs,
+ };
+ size_t local_work_size[] = { nth, 1 };
- backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
+ backend_ctx->enqueue_ndrange_kernel(kernel, 2, 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) {
diff --git a/ggml/src/ggml-opencl/kernels/ssm_scan.cl b/ggml/src/ggml-opencl/kernels/ssm_scan.cl
index 37698d123..1889b74cd 100644
--- a/ggml/src/ggml-opencl/kernels/ssm_scan.cl
+++ b/ggml/src/ggml-opencl/kernels/ssm_scan.cl
@@ -214,3 +214,133 @@ kernel void kernel_ssm_scan_f32_mamba2_d256(
s_warp[tid + 128] = state2;
s_warp[tid + 192] = state3;
}
+
+kernel void kernel_ssm_scan_f32(
+ global const char * s_buf,
+ ulong s_off,
+ global const char * x_buf,
+ ulong x_off,
+ global const char * dt_buf,
+ ulong dt_off,
+ global const char * A_buf,
+ ulong A_off,
+ global const char * B_buf,
+ ulong B_off,
+ global const char * C_buf,
+ ulong C_off,
+ global const char * ids_buf,
+ ulong ids_off,
+ global char * dst_buf,
+ ulong dst_off,
+ ulong s_nb2,
+ ulong s_nb3,
+ ulong x_nb2,
+ ulong x_nb3,
+ ulong dt_nb1,
+ ulong dt_nb2,
+ ulong A_nb1,
+ ulong B_nb2,
+ ulong B_nb3,
+ ulong C_nb2,
+ ulong C_nb3,
+ ulong state_off,
+ int head_dim,
+ int n_head,
+ int n_group,
+ int n_tokens,
+ ulong s_nb1,
+ ulong x_nb1,
+ ulong B_nb1,
+ ulong C_nb1,
+ uint A_ne0,
+ uint d_state,
+ uint n_seqs,
+ uint K,
+ local float * reduce
+) {
+ global const char * s_data = s_buf + s_off;
+ global const char * x_data = x_buf + x_off;
+ global const char * dt_data = dt_buf + dt_off;
+ global const char * A_data = A_buf + A_off;
+ global const char * B_data = B_buf + B_off;
+ global const char * C_data = C_buf + C_off;
+ global const int * ids_data = (global const int *) (ids_buf + ids_off);
+ global float * dst = (global float *) (dst_buf + dst_off);
+ const uint y_elems = state_off / sizeof(float);
+
+ const uint tid = get_local_id(0);
+ const uint inner_idx = get_group_id(0);
+ const uint seq_idx = get_group_id(1);
+ const uint head_idx = inner_idx / head_dim;
+ const uint dim_idx = inner_idx - head_idx * head_dim;
+ const uint group_idx = head_idx / (n_head / n_group);
+ const uint state_slot = (uint) ids_data[seq_idx];
+
+ const ulong s_idx = (ulong) state_slot * s_nb3 +
+ (ulong) head_idx * s_nb2 +
+ (ulong) dim_idx * s_nb1 +
+ (ulong) tid * sizeof(float);
+ float state = *((global const float *) (s_data + s_idx));
+
+ const ulong A_idx = (ulong) head_idx * A_nb1 +
+ (ulong) (tid % A_ne0) * sizeof(float);
+ const float A_value = *((global const float *) (A_data + A_idx));
+
+ for (int token_idx = 0; token_idx < n_tokens; ++token_idx) {
+ const ulong x_idx = (ulong) head_idx * x_nb1 +
+ (ulong) token_idx * x_nb2 +
+ (ulong) seq_idx * x_nb3 +
+ (ulong) dim_idx * sizeof(float);
+ const ulong dt_idx = (ulong) token_idx * dt_nb1 +
+ (ulong) seq_idx * dt_nb2 +
+ (ulong) head_idx * sizeof(float);
+ const ulong B_idx = (ulong) group_idx * B_nb1 +
+ (ulong) token_idx * B_nb2 +
+ (ulong) seq_idx * B_nb3 +
+ (ulong) tid * sizeof(float);
+ const ulong C_idx = (ulong) group_idx * C_nb1 +
+ (ulong) token_idx * C_nb2 +
+ (ulong) seq_idx * C_nb3 +
+ (ulong) tid * sizeof(float);
+
+ const float x_value = *((global const float *) (x_data + x_idx));
+ const float dt_value = *((global const float *) (dt_data + dt_idx));
+ const float B_value = *((global const float *) (B_data + B_idx));
+ const float C_value = *((global const float *) (C_data + C_idx));
+ const float dt_soft_plus = dt_value > 20.0f ? dt_value : log(1.0f + exp(dt_value));
+ const float dA = exp(dt_soft_plus * A_value);
+ const float x_dt = x_value * dt_soft_plus;
+
+ state = mad(state, dA, B_value * x_dt);
+ reduce[tid] = state * C_value;
+ barrier(CLK_LOCAL_MEM_FENCE);
+
+ for (uint stride = d_state / 2; stride > 0; stride >>= 1) {
+ if (tid < stride) {
+ reduce[tid] += reduce[tid + stride];
+ }
+ barrier(CLK_LOCAL_MEM_FENCE);
+ }
+
+ if (tid == 0) {
+ const uint y_idx = dim_idx + head_idx * head_dim +
+ token_idx * n_head * head_dim +
+ seq_idx * n_tokens * n_head * head_dim;
+ dst[y_idx] = reduce[0];
+ }
+
+ const uint snapshot_slot = n_tokens - 1 - token_idx;
+ if (snapshot_slot > 0 && snapshot_slot < K) {
+ const uint snapshot_idx = y_elems + tid + dim_idx * d_state +
+ head_idx * d_state * head_dim +
+ (snapshot_slot * n_seqs + seq_idx) * d_state * head_dim * n_head;
+ dst[snapshot_idx] = state;
+ }
+ barrier(CLK_LOCAL_MEM_FENCE);
+ }
+
+ const uint state_idx = y_elems + tid + dim_idx * d_state +
+ head_idx * d_state * head_dim +
+ seq_idx * d_state * head_dim * n_head;
+ dst[state_idx] = state;
+}