Commit 5e03bdd87 for llama.cpp

commit 5e03bdd8700948b9c41c54dd1b00f28a2aebc03f
Author: Todor Boinovski <tboinovski@gmail.com>
Date:   Mon Oct 5 17:42:40 2026 -0700

    hexagon: ssm-conv updates (#29971)

    * hexagon: ssm-conv double-buffered DMA for prefill and decode restructuring

    * hex-ssm-conv: remove divs from loops and fix trace events

    * hex-dma: improved SSM_CONV dma pipeline and streamlined dma_queue

    ---------

    Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>

diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 1e967086b..1e15e2eb7 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -6080,68 +6080,63 @@ static void ggml_hexagon_precompute_ssm_conv_params(

     const uint32_t raw_rpt = (d_inner + n_threads - 1) / n_threads;
     const uint32_t d_inner_per_thread = hex_round_up(raw_rpt, 32);
-    kparams->d_inner_per_thread = d_inner_per_thread;

-    kparams->src0_row_size_aligned = hex_round_up(ncs * sizeof(float), 128);
-    kparams->src1_row_size_aligned = hex_round_up(d_conv * sizeof(float), 128);
-    kparams->dst_row_size_aligned  = hex_round_up(d_inner * sizeof(float), 128);
+    const uint32_t src1_raw_bytes = hex_round_up(d_inner_per_thread * d_conv * sizeof(float), 128) + 128;
+    const uint32_t src1_T_bytes   = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 128);
+    const uint32_t vtcm_src1_per_thread = src1_raw_bytes + src1_T_bytes;
+
+    uint32_t vtcm_src0_per_thread = 0;
+    uint32_t vtcm_dst_per_thread  = 0;

     if (n_t == 1) {
         kparams->d_inner_tile = d_inner_per_thread;

-        const uint32_t src1_raw_bytes = hex_round_up(d_inner_per_thread * d_conv * sizeof(float), 128) + 128;
-        const uint32_t src1_T_bytes   = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 128);
-        const uint32_t vtcm_src1_per_thread = src1_raw_bytes + src1_T_bytes;
-
-        const uint32_t src0_raw_bytes = hex_round_up(d_inner_per_thread * d_conv * sizeof(float), 128) + 128;
-        const uint32_t src0_T_bytes   = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 128);
-        const uint32_t vtcm_src0_per_thread = src0_raw_bytes + src0_T_bytes;
-
-        const uint32_t vtcm_dst_per_thread = hex_round_up(d_inner_per_thread * sizeof(float), 128);
+        const uint32_t src0_tile_raw_bytes = hex_round_up(d_inner_per_thread * d_conv * sizeof(float), 128);
+        const uint32_t src0_T_bytes        = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 128);
+        vtcm_src0_per_thread = 2 * src0_tile_raw_bytes + src0_T_bytes;

-        kparams->vtcm_src0_size_per_thread = vtcm_src0_per_thread;
-        kparams->vtcm_src1_size_per_thread = vtcm_src1_per_thread;
-        kparams->vtcm_dst_size_per_thread  = vtcm_dst_per_thread;
-
-        kparams->vtcm_src0_size = vtcm_src0_per_thread * n_threads;
-        kparams->vtcm_src1_size = vtcm_src1_per_thread * n_threads;
-        kparams->vtcm_dst_size  = vtcm_dst_per_thread  * n_threads;
-        kparams->vtcm_size      = kparams->vtcm_src0_size + kparams->vtcm_src1_size + kparams->vtcm_dst_size;
+        const uint32_t dst_tile_bytes = hex_round_up(d_inner_per_thread * sizeof(float), 128);
+        vtcm_dst_per_thread = 2 * dst_tile_bytes;
     } else {
-        const uint32_t src1_raw_bytes = hex_round_up(d_inner_per_thread * d_conv * sizeof(float), 128) + 128;
-        const uint32_t src1_T_bytes   = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 128);
-        const uint32_t vtcm_src1_per_thread = src1_raw_bytes + src1_T_bytes;
-
         const size_t vtcm_budget = (sess->vtcm_size > 0 ? sess->vtcm_size / n_threads : (1024 * 1024));
-        const size_t avail_for_src0 = vtcm_budget > vtcm_src1_per_thread ? vtcm_budget - vtcm_src1_per_thread : (128 * 1024);

-        uint32_t d_inner_tile = (uint32_t)((avail_for_src0 / 2) / (ncs * sizeof(float) + n_t * sizeof(float) + 1));
+        // the kernel double-buffers the raw src0 tile and the dst tile, and transposes
+        // one 32-channel block at a time
+        const uint32_t src0_block_T    = hex_round_up(ncs * 32 * sizeof(float), 128);
+        const size_t   fixed_bytes     = vtcm_src1_per_thread + src0_block_T;
+        const size_t   avail_for_tiles = vtcm_budget > fixed_bytes ? vtcm_budget - fixed_bytes : (128 * 1024);
+
+        uint32_t target_max_tile = hex_round_up((d_inner_per_thread + 3) / 4, 32);
+        target_max_tile = (std::max)(target_max_tile, 32u);
+        target_max_tile = (std::min)(target_max_tile, 128u);
+
+        uint32_t d_inner_tile = (uint32_t)(avail_for_tiles / (2 * (ncs + n_t) * sizeof(float)));
         d_inner_tile = (d_inner_tile / 32) * 32;
         if (d_inner_tile == 0) {
             d_inner_tile = 32;
         }
+        if (d_inner_tile > target_max_tile) {
+            d_inner_tile = target_max_tile;
+        }
         if (d_inner_tile > d_inner_per_thread) {
             d_inner_tile = d_inner_per_thread;
         }
         kparams->d_inner_tile = d_inner_tile;

-        const uint32_t src0_tile_raw = hex_round_up(d_inner_tile * ncs * sizeof(float), 128) + 128;
-        const uint32_t src0_tile_T   = hex_round_up(ncs * d_inner_tile * sizeof(float), 128);
-        const uint32_t vtcm_src0_per_thread = src0_tile_raw + src0_tile_T;
-
-        const uint32_t vtcm_dst_per_thread = hex_round_up(d_inner_tile * n_t * sizeof(float), 128);
+        const uint32_t src0_tile_raw = hex_round_up(d_inner_tile * ncs * sizeof(float), 128);
+        vtcm_src0_per_thread = 2 * src0_tile_raw + src0_block_T;

-        kparams->vtcm_src0_size_per_thread = vtcm_src0_per_thread;
-        kparams->vtcm_src1_size_per_thread = vtcm_src1_per_thread;
-        kparams->vtcm_dst_size_per_thread  = vtcm_dst_per_thread;
-
-        kparams->vtcm_src0_size = vtcm_src0_per_thread * n_threads;
-        kparams->vtcm_src1_size = vtcm_src1_per_thread * n_threads;
-        kparams->vtcm_dst_size  = vtcm_dst_per_thread  * n_threads;
-        kparams->vtcm_size      = kparams->vtcm_src0_size + kparams->vtcm_src1_size + kparams->vtcm_dst_size;
+        vtcm_dst_per_thread = 2 * hex_round_up(d_inner_tile * n_t * sizeof(float), 128);
     }

-    kparams->div_n_threads = init_fastdiv_values(n_threads);
+    kparams->vtcm_src0_size_per_thread = vtcm_src0_per_thread;
+    kparams->vtcm_src1_size_per_thread = vtcm_src1_per_thread;
+    kparams->vtcm_dst_size_per_thread  = vtcm_dst_per_thread;
+
+    kparams->vtcm_src0_size = vtcm_src0_per_thread * n_threads;
+    kparams->vtcm_src1_size = vtcm_src1_per_thread * n_threads;
+    kparams->vtcm_dst_size  = vtcm_dst_per_thread  * n_threads;
+    kparams->vtcm_size      = kparams->vtcm_src0_size + kparams->vtcm_src1_size + kparams->vtcm_dst_size;
 }

 static void ggml_hexagon_precompute_gated_delta_net_params(
diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.c b/ggml/src/ggml-hexagon/htp/dma-queue.c
index 464e4b849..ef61d2d24 100644
--- a/ggml/src/ggml-hexagon/htp/dma-queue.c
+++ b/ggml/src/ggml-hexagon/htp/dma-queue.c
@@ -161,9 +161,9 @@ bool dma_queue_push_fallback_contig(dma_queue * q, dma_data ddata, size_t total)
     while (rem_bytes > 0) {
         const uint32_t cur_bytes = MIN(rem_bytes, DMA_SAFE_CHUNK_SIZE);
         dma_data cur_data = dma_make_data(cur_dst, cur_src);
-        if (!dma_ring_push_single_1d(r1, cur_data, cur_bytes)) {
+        if (!dma_ring_push_single_contig(r1, cur_data, cur_bytes)) {
             dma_ring_flush(r1);
-            dma_ring_push_single_1d(r1, cur_data, cur_bytes);
+            dma_ring_push_single_contig(r1, cur_data, cur_bytes);
         }
         cur_dst   += cur_bytes;
         cur_src   += cur_bytes;
diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.h b/ggml/src/ggml-hexagon/htp/dma-queue.h
index a736eb762..9b774bdd9 100644
--- a/ggml/src/ggml-hexagon/htp/dma-queue.h
+++ b/ggml/src/ggml-hexagon/htp/dma-queue.h
@@ -107,6 +107,14 @@ typedef struct {
 #define DMA_MAX_STRIDE_24B     0x00FFFFFFu    // 24-bit HW descriptor limit for strides (16MB - 1)
 #define DMA_SAFE_CHUNK_SIZE    0x00F00000u    // ~15MB safe contiguous chunk size

+#if __HVX_ARCH__ < 75
+#define DMA_MAX_2D_ROW_SIZE    DMA_MAX_SIZE_16B
+#define DMA_MAX_2D_STRIDE      DMA_MAX_STRIDE_16B
+#else
+#define DMA_MAX_2D_ROW_SIZE    DMA_MAX_SIZE_24B
+#define DMA_MAX_2D_STRIDE      DMA_MAX_STRIDE_24B
+#endif
+
 #define DMA_FALLBACK_CAPACITY  16u            // descriptors in secondary fallback ring

 typedef struct dma_ring_s dma_ring;
@@ -216,13 +224,10 @@ static inline bool dma_ring_push_single_1d(dma_ring * r, dma_data ddata, size_t

 static inline bool dma_ring_push_single_2d(dma_ring * r, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
 #if __HVX_ARCH__ > 79
+    assert(!((ddata.src | ddata.dst) >> 40) || nrows == 0);
     const uint32_t src_hi = (uint32_t) (ddata.src >> 32);
     const uint32_t dst_hi = (uint32_t) (ddata.dst >> 32);
     const bool is_ext     = (src_hi | dst_hi) != 0;
-
-    if (is_ext && ((ddata.src >> 40) || (ddata.dst >> 40))) {
-        return false;
-    }
 #endif

     if (((r->push_idx + 1) & r->idx_mask) == r->pop_idx) {
@@ -284,6 +289,16 @@ static inline bool dma_ring_push_single_2d(dma_ring * r, dma_data ddata, size_t
     return true;
 }

+#if __HVX_ARCH__ < 75
+static inline bool dma_ring_push_single_contig(dma_ring * r, dma_data ddata, size_t size) {
+    return dma_ring_push_single_1d(r, ddata, size);
+}
+#else
+static inline bool dma_ring_push_single_contig(dma_ring * r, dma_data ddata, size_t size) {
+    return dma_ring_push_single_2d(r, ddata, size, size, size, 1);
+}
+#endif
+
 static inline dma_data dma_ring_pop(dma_ring * r) {
     dma_data ddata = { 0 };

@@ -374,57 +389,36 @@ static inline uint32_t dma_queue_capacity(dma_queue * q) {
     return dma_ring_capacity(q->ring0);
 }

-#if __HVX_ARCH__ < 75
-
-static inline bool dma_queue_push(dma_queue *q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
-    // Fast path: everything fits in 16 bits
-    if (nrows == 0 || __builtin_expect(
-            nrows      <= DMA_MAX_NROWS &&
-            row_size   <= DMA_MAX_SIZE_16B &&
-            src_stride <= DMA_MAX_STRIDE_16B &&
-            dst_stride <= DMA_MAX_STRIDE_16B, 1)) {
-        return dma_ring_push_single_2d(q->ring0, ddata, dst_stride, src_stride, row_size, nrows);
+static inline bool dma_queue_push(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
+    if (__builtin_expect(nrows == 0, 0)) {
+        return dma_ring_push_single_1d(q->ring0, ddata, 0);
     }

-    // Contiguous block: 1D DMA mode supports up to 24-bit size (16MB)
-    if (nrows == 1 || (row_size == src_stride && row_size == dst_stride)) {
-        size_t total = row_size * nrows;
-        if (total <= DMA_MAX_SIZE_24B) {
-            return dma_ring_push_single_1d(q->ring0, ddata, total);
+    // 1. Hot path: Contiguous or single-row (80-90% of calls)
+    if (nrows == 1 || !((row_size ^ src_stride) | (row_size ^ dst_stride))) {
+        const size_t total = row_size * nrows;
+        if (__builtin_expect(total <= DMA_MAX_SIZE_24B, 1)) {
+            return dma_ring_push_single_contig(q->ring0, ddata, total);
         }
         return dma_queue_push_fallback_contig(q, ddata, total);
     }

-    // Row count overflow with 16-bit strides: chunk 2D descriptors via fallback ring
-    if (row_size <= DMA_MAX_SIZE_16B && src_stride <= DMA_MAX_STRIDE_16B && dst_stride <= DMA_MAX_STRIDE_16B) {
-        return dma_queue_push_fallback_2d(q, ddata, dst_stride, src_stride, row_size, nrows);
-    }
-
-    // Stride or row_size overflow: row-by-row 1D via fallback ring
-    return dma_queue_push_fallback_1d(q, ddata, dst_stride, src_stride, row_size, nrows);
-}
-
-#else // HVX_ARCH >= 75
-
-static inline bool dma_queue_push(dma_queue *q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
-    if (nrows == 0 || __builtin_expect(
-            nrows      <= DMA_MAX_NROWS &&
-            row_size   <= DMA_MAX_SIZE_24B &&
-            src_stride <= DMA_MAX_STRIDE_24B &&
-            dst_stride <= DMA_MAX_STRIDE_24B, 1)) {
+    // 2. Hot path: Standard strided 2D (10-20% of calls)
+    if (__builtin_expect(nrows <= DMA_MAX_NROWS &&
+                         (row_size | src_stride | dst_stride) <= DMA_MAX_2D_ROW_SIZE, 1)) {
         return dma_ring_push_single_2d(q->ring0, ddata, dst_stride, src_stride, row_size, nrows);
     }

-    // Contiguous block exceeding 24 bits
-    if (nrows == 1 || (row_size == src_stride && row_size == dst_stride)) {
-        size_t total = row_size * nrows;
-        return dma_queue_push_fallback_contig(q, ddata, total);
+    // 3. Cold path: Descriptor chunking fallbacks (< 0.1%)
+#if __HVX_ARCH__ < 75
+    if (row_size <= DMA_MAX_SIZE_16B && (src_stride | dst_stride) <= DMA_MAX_STRIDE_16B) {
+        return dma_queue_push_fallback_2d(q, ddata, dst_stride, src_stride, row_size, nrows);
     }
-
+    return dma_queue_push_fallback_1d(q, ddata, dst_stride, src_stride, row_size, nrows);
+#else
     return dma_queue_push_fallback_2d(q, ddata, dst_stride, src_stride, row_size, nrows);
-}
-
 #endif
+}

 static inline void dma_sync_read(dma_queue * dma_q, void * dst, dma_addr_t src, size_t bytes) {
     const uint32_t b = (uint32_t) bytes;
diff --git a/ggml/src/ggml-hexagon/htp/ssm-conv.c b/ggml/src/ggml-hexagon/htp/ssm-conv.c
index 931aa406e..0d28c8990 100644
--- a/ggml/src/ggml-hexagon/htp/ssm-conv.c
+++ b/ggml/src/ggml-hexagon/htp/ssm-conv.c
@@ -135,48 +135,49 @@ static inline void hvx_ssm_conv_unpack_to_T(const float * raw, float * T, uint32
     }
 }

-// HVX 32x32 src0 transpose for prefill: src0 {tile_n, ncs} (VTCM) -> src0_T {ncs, d_inner_tile} (VTCM)
-static inline void transpose_src0_block(const float * src0_block,
-                                        uint32_t      ncs,
-                                        uint32_t      cb_n,
-                                        uint32_t      d_inner_tile,
-                                        float *       src0_T_block_dst,
-                                        uint32_t      cb) {
-    const uint32_t T_TILE = VLEN_FP32;
-
-    HVX_Vector __attribute__((aligned(VLEN))) sub[32];
-
-    for (uint32_t t0 = 0; t0 < ncs; t0 += T_TILE) {
-        const uint32_t t_n = MIN(T_TILE, ncs - t0);
-
-        uint32_t __attribute__((aligned(VLEN))) mask_buf[VLEN_FP32] = { 0 };
-        for (uint32_t k = 0; k < t_n; ++k) {
-            mask_buf[k] = 0xFFFFFFFF;
-        }
-        const HVX_Vector mask = *(const HVX_Vector *) mask_buf;
+// Decode dot product specialization for d_conv == 4: multiply in the raw channel-major layout,
+// then deinterleave the products so each vector holds one tap of 32 channels, and sum.
+// Keeps both operands in DMA layout - no transpose, no scratch.
+static inline void hvx_ssm_conv_decode_4(const float * x, const float * w, float * out, uint32_t n_ch) {
+    for (uint32_t cb = 0; cb < n_ch; cb += VLEN_FP32) {
+        const float * xp = x + cb * 4;
+        const float * wp = w + cb * 4;

-        for (uint32_t r = 0; r < cb_n; ++r) {
-            const float * src_row = src0_block + r * ncs + t0;
-            sub[r] = (t_n == T_TILE) ? *(const HVX_UVector *) src_row : Q6_V_vand_VV(*(const HVX_UVector *) src_row, mask);
-        }
-        for (uint32_t r = cb_n; r < T_TILE; ++r) {
-            sub[r] = hvx_vec_splat_f32(0.0f);
-        }
+        HVX_Vector p0 = Q6_Vqf32_vmpy_VsfVsf(*(const HVX_Vector *)(xp +  0), *(const HVX_Vector *)(wp +  0));
+        HVX_Vector p1 = Q6_Vqf32_vmpy_VsfVsf(*(const HVX_Vector *)(xp + 32), *(const HVX_Vector *)(wp + 32));
+        HVX_Vector p2 = Q6_Vqf32_vmpy_VsfVsf(*(const HVX_Vector *)(xp + 64), *(const HVX_Vector *)(wp + 64));
+        HVX_Vector p3 = Q6_Vqf32_vmpy_VsfVsf(*(const HVX_Vector *)(xp + 96), *(const HVX_Vector *)(wp + 96));

-        hvx_transpose_32x32_f32(sub);
+        HVX_VectorPair p01 = Q6_W_vdeal_VVR(p1, p0, -4);
+        HVX_VectorPair p23 = Q6_W_vdeal_VVR(p3, p2, -4);

-        for (uint32_t r = 0; r < t_n; ++r) {
-            float * dst = src0_T_block_dst + (t0 + r) * d_inner_tile + cb;
-            if (cb_n == T_TILE) {
-                *(HVX_UVector *) dst = sub[r];
-            } else {
-                hvx_vec_store_u(dst, cb_n * sizeof(float), sub[r]);
-            }
-        }
+        HVX_VectorPair q02 = Q6_W_vdeal_VVR(Q6_V_lo_W(p23), Q6_V_lo_W(p01), -4);
+        HVX_VectorPair q13 = Q6_W_vdeal_VVR(Q6_V_hi_W(p23), Q6_V_hi_W(p01), -4);
+
+        HVX_Vector a = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_lo_W(q02), Q6_V_lo_W(q13));
+        HVX_Vector b = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_hi_W(q02), Q6_V_hi_W(q13));
+
+        *(HVX_Vector *)(out + cb) = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vqf32(a, b));
+    }
+}
+
+// Transpose src0 for prefill: one 32-channel block {32, ncs} (VTCM) -> T {ncs, 32} (VTCM).
+// One VTCM gather per output row: lane r picks channel cb+r, the region bound drops
+// the lanes past the channel tail.
+static inline void hvx_ssm_conv_transpose_block(const float * raw_block,
+                                                float *       T,
+                                                uint32_t      ncs,
+                                                uint32_t      cb_n,
+                                                HVX_Vector    vv) {
+    const size_t   base = (size_t) raw_block;
+    const uint32_t mu   = cb_n * ncs * sizeof(float) - 1;
+
+    for (uint32_t t = 0; t < ncs; ++t) {
+        Q6_vgather_ARMVw((HVX_Vector *) (T + (size_t) t * VLEN_FP32), base + t * sizeof(float), mu, vv);
     }
 }

-// Single-row decode worker (n_t == 1)
+// Single-token decode worker (n_t == 1)
 static void ssm_conv_thread_f32_decode(unsigned int nth, unsigned int ith, void * data) {
     struct htp_ssm_conv_context *             scctx   = (struct htp_ssm_conv_context *) data;
     struct htp_ops_context *                  octx    = scctx->octx;
@@ -202,9 +203,11 @@ static void ssm_conv_thread_f32_decode(unsigned int nth, unsigned int ith, void

     const uint32_t d_inner_per_thread = ir1 - ir0;
     const uint32_t d_inner_stride     = hex_round_up(d_inner_per_thread, VLEN_FP32);
+    const uint32_t d_inner_tile       = scctx->d_inner_tile;

-    const size_t src0_stride_seq_bytes = src0->nb[2];
-    const size_t dst_stride_seq_bytes  = dst->nb[2];
+    const size_t src0_stride_inner_bytes = src0->nb[1];
+    const size_t src0_stride_seq_bytes   = src0->nb[2];
+    const size_t dst_stride_seq_bytes    = dst->nb[2];

     uint8_t * src1_spad_base = octx->src1_spad.data + ith * octx->src1_spad.size_per_thread;
     uint8_t * src0_spad_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread;
@@ -216,57 +219,121 @@ static void ssm_conv_thread_f32_decode(unsigned int nth, unsigned int ith, void
     float * src1_raw = (float *) src1_spad_base;
     float * src1_T   = (float *) (src1_spad_base + weight_raw_size);

-    float * src0_raw = (float *) src0_spad_base;
-    float * src0_T   = (float *) (src0_spad_base + weight_raw_size);
+    const size_t src0_tile_raw_bytes = hex_round_up(d_inner_tile * d_conv * sizeof(float), 128);
+    const size_t dst_tile_bytes      = hex_round_up(d_inner_tile * sizeof(float), 128);
+
+    float * src0_tile_raw[2] = { (float *) src0_spad_base, (float *) (src0_spad_base + src0_tile_raw_bytes) };
+    float * src0_T           = (float *) (src0_spad_base + 2 * src0_tile_raw_bytes);

-    float * dst_spad = (float *) dst_spad_base;
+    float * dst_tile[2] = { (float *) dst_spad_base, (float *) (dst_spad_base + dst_tile_bytes) };

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];

-    // 1. Fetch weights src1 from DDR into VTCM via DMA (DMA64-safe)
+    // raw_dot keeps both operands in the DMA layout, so src1 needs no prep pass
+    const bool raw_dot = (d_conv == 4) && (d_inner_per_thread % VLEN_FP32 == 0);
+
+    const uint32_t n_tiles   = (d_inner_per_thread + d_inner_tile - 1) / d_inner_tile;
+    const uint32_t n_chunks  = n_s * n_tiles;
+    const size_t   row_bytes = d_conv * sizeof(float);
+
+    uint32_t s_fetch        = 0;
+    uint32_t tile_off_fetch = 0;
+
+    #define SSM_CONV_DECODE_PUSH_FETCH(c)                                                         \
+        do {                                                                                      \
+            const uint32_t   cur_tile_n = MIN(d_inner_tile, d_inner_per_thread - tile_off_fetch); \
+            const dma_addr_t fetch_ddr  = src0->data + s_fetch * src0_stride_seq_bytes +          \
+                                          (ir0 + tile_off_fetch) * src0_stride_inner_bytes;       \
+            dma_queue_push(dma_q,                                                                 \
+                           dma_make_data((uint8_t *) src0_tile_raw[(c) & 1], fetch_ddr),          \
+                           row_bytes, src0_stride_inner_bytes, row_bytes, cur_tile_n);            \
+            tile_off_fetch += d_inner_tile;                                                       \
+            if (tile_off_fetch >= d_inner_per_thread) {                                           \
+                tile_off_fetch = 0;                                                               \
+                s_fetch++;                                                                        \
+            }                                                                                     \
+        } while (0)
+
+    // Queue weights and initial input tiles together so DDR reads overlap
     const dma_addr_t src1_ddr = src1->data + ir0 * d_conv * sizeof(float);
     dma_queue_push(dma_q, dma_make_data((uint8_t *) src1_raw, src1_ddr), weight_bytes, weight_bytes, weight_bytes, 1);
-    dma_queue_pop(dma_q);

-    // 2. Unpack/transpose src1_raw into src1_T {d_conv, d_inner_stride}
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
-    hvx_ssm_conv_unpack_to_T(src1_raw, src1_T, d_inner_per_thread, d_inner_stride, d_conv);
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
-
-    const size_t input_bytes  = (size_t) d_inner_per_thread * d_conv * sizeof(float);
-    const size_t output_bytes = (size_t) d_inner_per_thread * sizeof(float);
-
-    // 3. Process each sequence
-    for (uint32_t s = 0; s < n_s; ++s) {
-        const dma_addr_t src0_ddr = src0->data + s * src0_stride_seq_bytes + ir0 * d_conv * sizeof(float);
-        dma_queue_push(dma_q, dma_make_data((uint8_t *) src0_raw, src0_ddr), input_bytes, input_bytes, input_bytes, 1);
-        dma_queue_pop(dma_q);
-
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) s);
-        hvx_ssm_conv_unpack_to_T(src0_raw, src0_T, d_inner_per_thread, d_inner_stride, d_conv);
-
-        for (uint32_t cb = 0; cb < d_inner_per_thread; cb += VLEN_FP32) {
-            const uint32_t cb_n = MIN(VLEN_FP32, d_inner_per_thread - cb);
-            HVX_Vector acc = hvx_vec_splat_f32(0.0f);
-            for (uint32_t j = 0; j < d_conv; ++j) {
-                HVX_Vector x = *(const HVX_Vector *)(src0_T + j * d_inner_stride + cb);
-                HVX_Vector w = *(const HVX_Vector *)(src1_T + j * d_inner_stride + cb);
-                acc          = Q6_Vqf32_vadd_Vqf32Vqf32(acc, Q6_Vqf32_vmpy_VsfVsf(x, w));
-            }
-            HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
-            if (cb_n == VLEN_FP32) {
-                *(HVX_Vector *)(dst_spad + cb) = y;
-            } else {
-                hvx_vec_store_u(dst_spad + cb, cb_n * sizeof(float), y);
+    SSM_CONV_DECODE_PUSH_FETCH(0);
+    if (n_chunks > 1) {
+        SSM_CONV_DECODE_PUSH_FETCH(1);
+    }
+
+    dma_queue_pop(dma_q);  // weights
+
+    if (!raw_dot) {
+        // Unpack/transpose src1_raw into src1_T {d_conv, d_inner_stride}
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_W_PREP, (uint16_t) ir0);
+        hvx_ssm_conv_unpack_to_T(src1_raw, src1_T, d_inner_per_thread, d_inner_stride, d_conv);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_W_PREP, (uint16_t) ir0);
+    }
+
+    uint32_t i3       = 0;
+    uint32_t tile_off = 0;
+
+    for (uint32_t c = 0; c < n_chunks; ++c) {
+        const uint32_t tile_n = MIN(d_inner_tile, d_inner_per_thread - tile_off);
+
+        if (c >= 2) {
+            dma_queue_pop(dma_q);  // writeback of chunk c-2, frees dst_tile[c & 1]
+        }
+        dma_queue_pop(dma_q);      // fetch chunk c
+
+        float * restrict out = dst_tile[c & 1];
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i3);
+        if (raw_dot) {
+            const float * xp = src0_tile_raw[c & 1];
+            const float * wp = src1_raw + tile_off * 4;
+            hvx_ssm_conv_decode_4(xp, wp, out, tile_n);
+        } else {
+            const uint32_t tile_stride = hex_round_up(tile_n, VLEN_FP32);
+            hvx_ssm_conv_unpack_to_T(src0_tile_raw[c & 1], src0_T, tile_n, tile_stride, d_conv);
+
+            for (uint32_t cb = 0; cb < tile_n; cb += VLEN_FP32) {
+                const uint32_t cb_n = MIN(VLEN_FP32, tile_n - cb);
+                HVX_Vector acc = hvx_vec_splat_f32(0.0f);
+                for (uint32_t j = 0; j < d_conv; ++j) {
+                    HVX_Vector x = *(const HVX_Vector *)(src0_T + j * tile_stride + cb);
+                    HVX_Vector w = *(const HVX_Vector *)(src1_T + j * d_inner_stride + tile_off + cb);
+                    acc          = Q6_Vqf32_vadd_Vqf32Vqf32(acc, Q6_Vqf32_vmpy_VsfVsf(x, w));
+                }
+                HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
+                if (cb_n == VLEN_FP32) {
+                    *(HVX_Vector *)(out + cb) = y;
+                } else {
+                    hvx_vec_store_u(out + cb, cb_n * sizeof(float), y);
+                }
             }
         }
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) s);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i3);
+
+        const dma_addr_t dst_ddr        = dst->data + i3 * dst_stride_seq_bytes + (ir0 + tile_off) * sizeof(float);
+        const size_t     tile_out_bytes = (size_t) tile_n * sizeof(float);
+        dma_queue_push(dma_q, dma_make_data(dst_ddr, (uint8_t *) out),
+                       tile_out_bytes, tile_out_bytes, tile_out_bytes, 1);

-        const dma_addr_t dst_ddr = dst->data + s * dst_stride_seq_bytes + ir0 * sizeof(float);
-        dma_queue_push(dma_q, dma_make_data(dst_ddr, (uint8_t *) dst_spad), output_bytes, output_bytes, output_bytes, 1);
-        dma_queue_pop(dma_q);
+        if (c + 2 < n_chunks) {
+            SSM_CONV_DECODE_PUSH_FETCH(c + 2);
+        }
+
+        tile_off += d_inner_tile;
+        if (tile_off >= d_inner_per_thread) {
+            tile_off = 0;
+            i3++;
+        }
+    }
+
+    for (uint32_t k = MIN(n_chunks, 2); k > 0; --k) {
+        dma_queue_pop(dma_q);  // drain the last writebacks
     }

+    #undef SSM_CONV_DECODE_PUSH_FETCH
+
     FARF(HIGH, "ssm-conv-f32-decode %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
          src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
@@ -318,82 +385,157 @@ static void ssm_conv_thread_f32_prefill(unsigned int nth, unsigned int ith, void
     float * src1_raw = (float *) src1_spad_base;
     float * src1_T   = (float *) (src1_spad_base + weight_raw_size);

+    // src0 spad holds two raw tiles (fetch of tile n+1 overlaps compute of tile n) plus
+    // the transposed block. dst spad holds two tiles so a writeback can stay in flight.
     const size_t src0_tile_raw_bytes = hex_round_up(d_inner_tile * ncs * sizeof(float), 128);
-    float * src0_tile_raw = (float *) src0_spad_base;
-    float * src0_T        = (float *) (src0_spad_base + src0_tile_raw_bytes);
+    const size_t dst_tile_bytes      = hex_round_up(d_inner_tile * n_t * sizeof(float), 128);

-    float * dst_tile = (float *) dst_spad_base;
+    float * src0_tile_raw[2] = { (float *) src0_spad_base, (float *) (src0_spad_base + src0_tile_raw_bytes) };
+    float * src0_T           = (float *) (src0_spad_base + 2 * src0_tile_raw_bytes);
+
+    float * dst_tile[2] = { (float *) dst_spad_base, (float *) (dst_spad_base + dst_tile_bytes) };

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];

-    // 1. Fetch weights src1 from DDR into VTCM via DMA (DMA64-safe)
+    const uint32_t n_tiles   = (d_inner_per_thread + d_inner_tile - 1) / d_inner_tile;
+    const uint32_t n_chunks  = n_s * n_tiles;
+    const size_t   row_bytes = ncs * sizeof(float);
+
+    uint32_t s_fetch        = 0;
+    uint32_t tile_off_fetch = 0;
+
+    // Chunk c fetches into src0_tile_raw[c & 1] and writes back from dst_tile[c & 1].
+    // Two fetches run ahead, so the queue order is F0 F1 W0 F2 W1 ... and pops follow it.
+    #define SSM_CONV_PUSH_FETCH(c)                                                                \
+        do {                                                                                      \
+            const uint32_t   cur_tile_n = MIN(d_inner_tile, d_inner_per_thread - tile_off_fetch); \
+            const dma_addr_t fetch_ddr  = src0->data + s_fetch * src0_stride_seq_bytes +          \
+                                          (ir0 + tile_off_fetch) * src0_stride_inner_bytes;       \
+            dma_queue_push(dma_q,                                                                 \
+                           dma_make_data((uint8_t *) src0_tile_raw[(c) & 1], fetch_ddr),          \
+                           row_bytes, src0_stride_inner_bytes, row_bytes, cur_tile_n);            \
+            tile_off_fetch += d_inner_tile;                                                       \
+            if (tile_off_fetch >= d_inner_per_thread) {                                           \
+                tile_off_fetch = 0;                                                               \
+                s_fetch++;                                                                        \
+            }                                                                                     \
+        } while (0)
+
+    // Queue weights and initial input tiles together so DDR reads overlap
     const dma_addr_t src1_ddr = src1->data + ir0 * d_conv * sizeof(float);
     dma_queue_push(dma_q, dma_make_data((uint8_t *) src1_raw, src1_ddr), weight_bytes, weight_bytes, weight_bytes, 1);
-    dma_queue_pop(dma_q);

-    // 2. Unpack/transpose src1_raw into src1_T {d_conv, d_inner_stride}
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
+    SSM_CONV_PUSH_FETCH(0);
+    if (n_chunks > 1) {
+        SSM_CONV_PUSH_FETCH(1);
+    }
+
+    dma_queue_pop(dma_q);  // weights
+
+    // Unpack/transpose src1_raw into src1_T {d_conv, d_inner_stride}
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_W_PREP, (uint16_t) ir0);
     hvx_ssm_conv_unpack_to_T(src1_raw, src1_T, d_inner_per_thread, d_inner_stride, d_conv);
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_W_PREP, (uint16_t) ir0);

     const uint32_t C_TILE = VLEN_FP32;

-    for (uint32_t i3 = 0; i3 < n_s; ++i3) {
-        for (uint32_t tile_off = 0; tile_off < d_inner_per_thread; tile_off += d_inner_tile) {
-            const uint32_t tile_n = MIN(d_inner_tile, d_inner_per_thread - tile_off);
+    // gather offsets: lane r reads channel r of the raw tile block
+    uint32_t __attribute__((aligned(VLEN))) gather_off[VLEN_FP32];
+    for (uint32_t r = 0; r < VLEN_FP32; ++r) {
+        gather_off[r] = r * ncs * sizeof(float);
+    }
+    const HVX_Vector vv = *(const HVX_Vector *) gather_off;
+
+    uint32_t i3       = 0;
+    uint32_t tile_off = 0;
+
+    for (uint32_t c = 0; c < n_chunks; ++c) {
+        const uint32_t tile_n = MIN(d_inner_tile, d_inner_per_thread - tile_off);
+
+        const float * restrict raw = src0_tile_raw[c & 1];
+        float * restrict       out = dst_tile[c & 1];
+
+        if (c >= 2) {
+            dma_queue_pop(dma_q);  // writeback of chunk c-2, frees dst_tile[c & 1]
+        }
+        dma_queue_pop(dma_q);      // fetch chunk c
+
+        // Channel block outer, token inner: the taps and the sliding window stay in
+        // registers, so a new output row costs one src0_T load.
+        const uint32_t dst_tile_stride = hex_round_up(tile_n, C_TILE);
+
+        for (uint32_t cb = 0; cb < tile_n; cb += C_TILE) {
+            const uint32_t cb_n = MIN(C_TILE, tile_n - cb);

-            // Fetch src0 chunk from DDR to VTCM via 2D DMA
-            const dma_addr_t src0_tile_ddr = src0->data +
-                i3 * src0_stride_seq_bytes +
-                (ir0 + tile_off) * src0_stride_inner_bytes;
-            const size_t row_bytes = ncs * sizeof(float);
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) tile_off);
+            hvx_ssm_conv_transpose_block(raw + (size_t) cb * ncs, src0_T, ncs, cb_n, vv);
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) tile_off);

-            dma_queue_push(dma_q, dma_make_data((uint8_t *) src0_tile_raw, src0_tile_ddr),
-                           row_bytes, src0_stride_inner_bytes, row_bytes, tile_n);
-            dma_queue_pop(dma_q);
+            const float * restrict wp = src1_T + tile_off + cb;
+            float * restrict       op = out + cb;

-            // Transpose src0 chunk in VTCM into {d_inner_tile, ncs} layout
             htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) tile_off);
-            for (uint32_t cb = 0; cb < tile_n; cb += C_TILE) {
-                const uint32_t cb_n = MIN(C_TILE, tile_n - cb);
-                transpose_src0_block(src0_tile_raw + cb * ncs, ncs, cb_n, d_inner_tile, src0_T, cb);
-            }
+            if (d_conv == 4) {
+                const HVX_Vector w0 = *(const HVX_Vector *) (wp);
+                const HVX_Vector w1 = *(const HVX_Vector *) (wp + d_inner_stride);
+                const HVX_Vector w2 = *(const HVX_Vector *) (wp + 2 * d_inner_stride);
+                const HVX_Vector w3 = *(const HVX_Vector *) (wp + 3 * d_inner_stride);
+
+                HVX_Vector x0 = *(const HVX_Vector *) (src0_T);
+                HVX_Vector x1 = *(const HVX_Vector *) (src0_T + C_TILE);
+                HVX_Vector x2 = *(const HVX_Vector *) (src0_T + 2 * C_TILE);

-            // Compute convolution
-            for (uint32_t t = 0; t < n_t; ++t) {
-                for (uint32_t cb = 0; cb < tile_n; cb += C_TILE) {
-                    const uint32_t cb_n = MIN(C_TILE, tile_n - cb);
+                for (uint32_t t = 0; t < n_t; ++t) {
+                    const HVX_Vector x3 = *(const HVX_Vector *) (src0_T + (t + 3) * C_TILE);

+                    HVX_Vector a = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(x0, w0),
+                                                            Q6_Vqf32_vmpy_VsfVsf(x1, w1));
+                    HVX_Vector b = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(x2, w2),
+                                                            Q6_Vqf32_vmpy_VsfVsf(x3, w3));
+
+                    *(HVX_Vector *) (op + t * dst_tile_stride) = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vqf32(a, b));
+
+                    x0 = x1;
+                    x1 = x2;
+                    x2 = x3;
+                }
+            } else {
+                for (uint32_t t = 0; t < n_t; ++t) {
                     HVX_Vector acc = hvx_vec_splat_f32(0.0f);
                     for (uint32_t j = 0; j < d_conv; ++j) {
-                        HVX_Vector x = *(const HVX_Vector *) (src0_T + (t + j) * d_inner_tile + cb);
-                        HVX_Vector w = *(const HVX_Vector *) (src1_T + j * d_inner_stride + tile_off + cb);
+                        HVX_Vector x = *(const HVX_Vector *) (src0_T + (t + j) * C_TILE);
+                        HVX_Vector w = *(const HVX_Vector *) (wp + j * d_inner_stride);
                         acc          = Q6_Vqf32_vadd_Vqf32Vqf32(acc, Q6_Vqf32_vmpy_VsfVsf(x, w));
                     }
-
-                    HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
-                    float * dst_tile_ptr = dst_tile + t * tile_n + cb;
-                    if (cb_n == C_TILE) {
-                        *(HVX_Vector *) dst_tile_ptr = y;
-                    } else {
-                        hvx_vec_store_u(dst_tile_ptr, cb_n * sizeof(float), y);
-                    }
+                    *(HVX_Vector *) (op + t * dst_tile_stride) = Q6_Vsf_equals_Vqf32(acc);
                 }
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) tile_off);
+        }
+
+        // Writeback dst_tile VTCM -> DDR via 2D DMA
+        dma_queue_push(dma_q,
+                       dma_make_data(dst->data + i3 * dst_stride_seq_bytes + (ir0 + tile_off) * sizeof(float),
+                                     (uint8_t *) out),
+                       dst_stride_token_bytes, dst_tile_stride * sizeof(float), tile_n * sizeof(float), n_t);

-            // Writeback dst_tile from VTCM to DDR via 2D DMA
-            const dma_addr_t dst_tile_ddr = dst->data +
-                i3 * dst_stride_seq_bytes +
-                (ir0 + tile_off) * sizeof(float);
-            const size_t dst_row_bytes = tile_n * sizeof(float);
+        if (c + 2 < n_chunks) {
+            SSM_CONV_PUSH_FETCH(c + 2);
+        }

-            dma_queue_push(dma_q, dma_make_data(dst_tile_ddr, (uint8_t *) dst_tile),
-                           dst_stride_token_bytes, dst_row_bytes, dst_row_bytes, n_t);
-            dma_queue_pop(dma_q);
+        tile_off += d_inner_tile;
+        if (tile_off >= d_inner_per_thread) {
+            tile_off = 0;
+            i3++;
         }
     }

+    for (uint32_t k = MIN(n_chunks, 2); k > 0; --k) {
+        dma_queue_pop(dma_q);  // drain the last writebacks
+    }
+
+    #undef SSM_CONV_PUSH_FETCH
+
     FARF(HIGH, "ssm-conv-f32-prefill %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
          src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
diff --git a/ggml/src/ggml-hexagon/htp/ssm-conv.h b/ggml/src/ggml-hexagon/htp/ssm-conv.h
index be62d7bf5..a071d74b0 100644
--- a/ggml/src/ggml-hexagon/htp/ssm-conv.h
+++ b/ggml/src/ggml-hexagon/htp/ssm-conv.h
@@ -12,13 +12,8 @@ struct htp_ssm_conv_kernel_params {
     uint32_t d_inner;
     uint32_t n_t;
     uint32_t n_s;
-    uint32_t d_inner_per_thread;
     uint32_t d_inner_tile;

-    uint32_t src0_row_size_aligned;
-    uint32_t src1_row_size_aligned;
-    uint32_t dst_row_size_aligned;
-
     uint32_t vtcm_src0_size_per_thread;
     uint32_t vtcm_src1_size_per_thread;
     uint32_t vtcm_dst_size_per_thread;
@@ -27,8 +22,6 @@ struct htp_ssm_conv_kernel_params {
     uint32_t vtcm_src1_size;
     uint32_t vtcm_dst_size;
     uint32_t vtcm_size;
-
-    struct fastdiv_values div_n_threads;
 };

 #if defined(__cplusplus)