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)