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;
+}