Commit 0c3626ec0 for llama.cpp

commit 0c3626ec06e913e2aa5ee62256d63ed0d95d969e
Author: Max Krasnyansky <maxk@qti.qualcomm.com>
Date:   Mon Sep 21 01:00:28 2026 -0700

    hexagon: overhaul of buffer and DMA handling to support 64bit mappings + improvements (#29197)

    * hex-dma64: enable support extended buffer mappings and 64bit dma

    hex-dma64: expand binary ops to support more DMA scenarios

    hex-dma64: add binary-ops.h

    hex-dma64: add --hex-dma64 to run.py and fix minor issues

    hex-dma64: update SSM_CONV to use dma with proper support for 64bit

    hex-ops: remove obsolete gate for % 128 in binary ops

    hex-l2: dont check weight tensors against dirty ranges

    hex-dma64: most binary ops now support dma

    hex-dma: use dma_addr_t instead of plain uint64_t to avoid overhead on older targets

    hex-dma: update all dma users to use dma_data (instead of pointers)

    hex-dma64: simplify lazy buffer mapping and clonning

    hex-fusion: factor out try_fuse_common that checks for dma64 buffers

    hex-bufs: minor cleanup for mmaping logic

    hex-bufs: simplify buffer clonning

    hex-ssm-conv: tighten gating checks and check vtcm size in kparams

    hex-binary: fix incorred mod/wrap in scalar ops

    hex-binary: make sure to call precompute kparams in support checks

    hex-dma64: update addr handling in mm,concat,binary

    hex-dma64: fixing up leftover of dma_addr_t conversion

    hex-binary: redo the kernel selection again and fix regressions in MOEs

    hex-binary: specialize per-type/per-op

    hex-binary: vtcm-layout and per-src dma-queue

    hex-dma64: update dma_push to transparently handle 64bit/extended

    * hex-cpy: fix improper rebase with the fixes for cont. tensors

    * hex-dma-cpy: update CPY to use safe dma rows/size limits

    * hex-mmap: bump number of mmaps to 64 to allow avoid eviction in larger models

    * hex-dma: add support for the secondary ring as a fallback for too-large transactions

    * hex-rope: fix freq_factors access with 64bit dma

    * hex-dma: audit all ops for proper use/gards for 64bit addresses

    * hex-dma64: uninline glu-compute funcs to avoid register pressure due to 64bit addr math

    * hex-dma64: refactor binary ops to separate dma loops

    * hex-devel: add inspect script to help with dbg and analysis

    * hex-dma: refactor dma-pipelines in unary-ops

    * hex-dma: rewrite softmax to use dma

    * hex-dma: rewrite GDN dma loops and improve HVX register usage

    * hex-gdn: fuse GDN+CPY

    * hex-mm: factor out HVX solver

    * hex-mm: remove hvx-flat kernels, the chunked version now handles vtcm limits much better

    * hex-buffs: reject huge buffer allocations that we cannot memory map

    * hex-inspect: add logic to look for float promo calls

    * hex-mm: reduce HVX register spills in HVX prompt kernels

    * hex-bufs: do not double count buffers from tensors in the same op

    * hex-roll: fix merge conflict

    * hex-dma: reroute all matmul ddr kernels to new chunked dma/vtcm kernels

    * hex-dev: update developer docs to include inspection for register spils and float promos

    * hex-ops: forgot to add new headers

    * hex-softmax: fix gpt-oss dims

    * hex-dma64: cleanup dma_addr_t casts

    * hex-dma64: add support for dma/vtcm for flash-atten with sinks

    * hex-mm-add: fix MUL_MAT+ADD fusion with bias.weights in extended bufs

    * hex-add-id: add support for dma for src1 (exp. table)

    * hex-dma: imrpove v73 fallback paths

    * hex-bufs: do not drop extended mappings during va defrag

    * hex-scripts: fix flake8 warnings

    * hex-docs: fix editor-config warnings

    * hex-inspect: fix warnings from ty

diff --git a/docs/backend/snapdragon/README.md b/docs/backend/snapdragon/README.md
index 5d32a5877..ca79ca885 100644
--- a/docs/backend/snapdragon/README.md
+++ b/docs/backend/snapdragon/README.md
@@ -327,6 +327,9 @@ on 4 physical NPUs, or `--devices 'HTP0[0-1:0],HTP1[0-1:1]'` on 2 physical NPUs
 - `GGML_HEXAGON_HOSTBUF=1` (default: 0, disabled)
   Enables allocating host buffers for debugging. By default, host buffers are disabled.

+- `GGML_HEXAGON_DMA64=0` (default: enabled on v81+)
+  Disables 64-bit DMA for model weights. Set to `1` to enable it explicitly on a supported architecture.
+
 - `GGML_HEXAGON_VERBOSE=1`
   Enables verbose logging of Ops from the backend. Example output:

diff --git a/docs/backend/snapdragon/developer.md b/docs/backend/snapdragon/developer.md
index 633643c16..378d47653 100644
--- a/docs/backend/snapdragon/developer.md
+++ b/docs/backend/snapdragon/developer.md
@@ -146,6 +146,28 @@ Writing high-performance operators for Hexagon requires following specific guide
   python3 scripts/snapdragon/ggml-hexagon-align-macros.py --fix ggml/src/ggml-hexagon/htp/
   ```

+### Binary Inspection and Spill Analysis
+
+Use [`scripts/snapdragon/ggml-hexagon-inspect.py`](../../../scripts/snapdragon/ggml-hexagon-inspect.py) to audit Hexagon binaries for register
+spills, unexpected float promotions, or disassembly:
+
+- Always verify that compute kernels have zero in-loop vector spills (`--spills --strict`) and no float promotions (`--promotions`).
+- Avoid excessive loop unrolling (`#pragma unroll`), which increases register pressure and causes spills.
+
+```bash
+# Check for vector and scalar register spills
+python3 scripts/snapdragon/ggml-hexagon-inspect.py --spills --strict --func "^compute_"
+
+# Check for float promotions
+python3 scripts/snapdragon/ggml-hexagon-inspect.py --promotions --func "^compute_"
+
+# Disassemble with annotated loops and spill markers
+python3 scripts/snapdragon/ggml-hexagon-inspect.py --disasm compute_same_shape_div_f32
+
+# Resolve crash addresses to function symbols and lines
+python3 scripts/snapdragon/ggml-hexagon-inspect.py --addr2line 0x51a30 0x5ba54
+```
+
 ## Multi-Device Partitioning (mdev)

 Multi-device (mdev) mode enables row-level tensor parallel execution across multiple physical NPU cores or virtual NPU
diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 352434b6a..ec5a4aeb6 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -22,6 +22,7 @@
 #include <deque>
 #include <algorithm>
 #include <cmath>
+#include <initializer_list>

 #ifdef _WIN32
 #    define WIN32_LEAN_AND_MEAN
@@ -53,11 +54,15 @@
 #include "htp-opnode.h"
 #include "htp-ops.h"
 #include "htp/matmul-ops.h"
+#include "htp/binary-ops.h"
 #include "htp/flash-attn-ops.h"
 #include "htp/unary-ops.h"
 #include "htp/get-rows-ops.h"
 #include "htp/set-rows-ops.h"
+#include "htp/softmax-ops.h"
 #include "htp/rope-ops.h"
+#include "htp/ssm-conv.h"
+#include "htp/gated-delta-net-ops.h"
 #include "htp_iface.h"
 #include "htp-drv.h"

@@ -91,8 +96,9 @@ static int    opt_etm     = 0;
 static int    opt_verbose = 0;
 static int    opt_profile = 0; // profiling mode (0-disabled, 1-basic, 2-pmu)
 static bool   opt_hostbuf = false;
+static bool   opt_dma64   = false;

-static int    opt_mm_select = 3; // 3 = HMX -> Tiled -> Flat -> CPU, 2 = Tiled -> Flat -> CPU, 1 = Flat -> CPU
+static int    opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
 static int    opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
 static int    opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = unfused ALLREDUCE (DMA), 0 = fallback to CPY+FENCE

@@ -113,6 +119,7 @@ enum ggml_hexagon_fusion_flags {
     GGML_HEXAGON_FUSE_MUL_MAT_ADD   = (1 << 3), // 8
     GGML_HEXAGON_FUSE_MUL_MAT_NX    = (1 << 4), // 16
     GGML_HEXAGON_FUSE_MUL_MAT_ID_NX = (1 << 5), // 32
+    GGML_HEXAGON_FUSE_GDN_CPY       = (1 << 6), // 64
 };

 static inline bool ggml_hexagon_is_fusion_enabled(int flag) {
@@ -299,6 +306,15 @@ static void ggml_hexagon_precompute_unary_params(
     struct htp_unary_kernel_params * kparams
 );

+static bool ggml_hexagon_precompute_binary_params(
+    const struct ggml_hexagon_session * sess,
+    uint32_t op,
+    const struct ggml_tensor * src0,
+    const struct ggml_tensor * src1,
+    const struct ggml_tensor * dst,
+    struct htp_binary_kernel_params * kparams
+);
+
 static void ggml_hexagon_precompute_get_rows_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
@@ -315,12 +331,32 @@ static void ggml_hexagon_precompute_set_rows_params(
     struct htp_set_rows_kernel_params * kparams
 );

+static void ggml_hexagon_precompute_softmax_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * op,
+    struct htp_softmax_kernel_params * kparams
+);
+
 static void ggml_hexagon_precompute_rope_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * op,
     struct htp_rope_kernel_params * kparams
 );

+static void ggml_hexagon_precompute_ssm_conv_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * src0,
+    const struct ggml_tensor * src1,
+    const struct ggml_tensor * dst,
+    struct htp_ssm_conv_kernel_params * kparams
+);
+
+static void ggml_hexagon_precompute_gated_delta_net_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * op,
+    struct htp_gdn_kernel_params * kparams
+);
+
 static void ggml_hexagon_precompute_fused_mmnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
@@ -349,6 +385,7 @@ static bool ggml_hexagon_precompute_allreduce_params(
 );

 static bool mm_is_hmx_eligible(const ggml_tensor * t);
+static htp_op_code op_remap_to_htp(const ggml_tensor * t);
 static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams);
 static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams);
 static bool is_mergeable_mul_mat(const ggml_tensor * t);
@@ -474,12 +511,14 @@ struct ggml_hexagon_session {
                            const std::vector<const ggml_tensor *> & sync_tensors, uint32_t rank, uint32_t n_ranks,
                            uint32_t fence_seq_entry = 0, uint32_t fence_seq_exit = 0);

+    void start_batch();
     void flush_sync(bool all = true);
     void flush_async();
     void flush_batch(size_t min_ops = 1);
     void flush_peers();
     void flush_pending(bool all = true);

+    ggml_hexagon_shared_buffer * mmap_tensor(const ggml_tensor * t);
     bool clone_buffer(const ggml_hexagon_shared_buffer*);
     void release_buffer(const ggml_hexagon_shared_buffer*);
     void unclone_buffer(const ggml_hexagon_shared_buffer*);
@@ -515,6 +554,9 @@ struct ggml_backend_hexagon_device_context {
     ggml_hexagon_session * session() {
         if (!sess) {
             sess = std::make_unique<ggml_hexagon_session>(config, dev);
+            if (max_bufsize > sess->max_vmem) {
+                max_bufsize = sess->max_vmem;
+            }
         }
         return sess.get();
     }
@@ -563,15 +605,21 @@ struct ggml_hexagon_shared_buffer {
     std::vector<ggml_hexagon_tensor_extra *>   tensor_extra;
     bool     mapped;
     bool     pinned;
+    bool     extended;

     const char * c_name() const { return sess->c_name(); }
     uint8_t *    base()   const { return mem ? mem->base : nullptr; }
     size_t       size()   const { return mem ? mem->size : 0;  }
     int          fd()     const { return mem ? mem->fd   : -1; }

-    void mmap() {
-        if (!this->mem) return;
-        fastrpc_map_flags flags = this->pinned ? FASTRPC_MAP_FD : FASTRPC_MAP_FD_DELAYED;
+    void mmap(bool extended = false) {
+        if (!this->mem)   return;
+        if (this->mapped) return;
+
+        GGML_ASSERT(!this->pinned || !extended);
+
+        this->extended = extended;
+        fastrpc_map_flags flags = this->pinned ? FASTRPC_MAP_FD : (extended ? FASTRPC_MAP_FD_DELAYED_EXTENDED : FASTRPC_MAP_FD_DELAYED);

         int err = fastrpc_mmap(sess->domain_id, fd(), (void *) base(), 0, size(), flags);
         if (err != 0) {
@@ -580,8 +628,8 @@ struct ggml_hexagon_shared_buffer {
             throw std::runtime_error("ggml-hex: fastrpc_mmap failed (see log for details)");
         }

-        HEX_VERBOSE("ggml-hex: %s mapped buffer: base %p size %zu fd %d pinned %u\n",
-                sess->c_name(), (void *) base(), size(), fd(), pinned);
+        HEX_VERBOSE("ggml-hex: %s mapped buffer: base %p size %zu fd %d pinned %u extended %u\n",
+                sess->c_name(), (void *) base(), size(), fd(), pinned, extended);

         this->mapped = true;
     }
@@ -611,7 +659,9 @@ struct ggml_hexagon_shared_buffer {

         HEX_VERBOSE("ggml-hex: %s allocated buffer: base %p size %zu fd %d pinned %d\n", sess->c_name(),
                     (void *) base(), this->size(), fd(), (int) pinned);
-        mmap();
+        if (this->pinned) {
+            mmap();
+        }
     }

     void free() {
@@ -623,13 +673,18 @@ struct ggml_hexagon_shared_buffer {
     }

     ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, size_t size, bool pinned = false) {
-        this->sess   = sess;
-        this->mapped = false;
-        this->pinned = pinned;
+        this->sess     = sess;
+        this->mapped   = false;
+        this->pinned   = pinned;
+        this->extended = false;

         // Size adjustment inside the buffer class: 4K aligned data size + 4K guard page
         size_t guard_offset = (size + 4095) & ~4095;
         size_t total_size   = guard_offset + 4096;
+        if (!pinned && opt_dma64) {
+            constexpr size_t extended_align = 2 * 1024 * 1024;
+            total_size = (total_size + extended_align - 1) & ~(extended_align - 1);
+        }

         alloc(total_size);
     }
@@ -640,6 +695,7 @@ struct ggml_hexagon_shared_buffer {
         this->mem    = other.mem;
         this->mapped = false;
         this->pinned = other.pinned;
+        this->extended = other.extended;
     }

     ~ggml_hexagon_shared_buffer() {
@@ -663,6 +719,7 @@ struct ggml_hexagon_fence_buffer : public ggml_hexagon_shared_buffer {
         backend_buffer.buft    = buft;
         backend_buffer.context = static_cast<ggml_hexagon_shared_buffer *>(this);
         backend_buffer.size    = size;
+        mmap(false);
     }

     uint8_t * alloc_slot(uint32_t n_slots = 1) {
@@ -703,11 +760,6 @@ inline void ggml_hexagon_session::free_fence(void * ptr, uint32_t n_slots) {
     }
 }

-static ggml_hexagon_session * ggml_backend_hexagon_buffer_get_sess(ggml_backend_buffer_t buffer) {
-    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(buffer->context);
-    return sbuf->sess;
-}
-
 static void ggml_backend_hexagon_buffer_free_buffer(ggml_backend_buffer_t buffer) {
     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(buffer->context);
     sbuf->sess->unclone_buffer(sbuf);
@@ -2024,7 +2076,17 @@ static const char * ggml_backend_hexagon_buffer_type_name(ggml_backend_buffer_ty
 static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
             ggml_backend_buffer_type_t buffer_type, size_t size) {
     auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
+    if (size > dev_ctx->max_bufsize) {
+        GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_bufsize %zu\n",
+                       dev_ctx->c_name(), size, dev_ctx->max_bufsize);
+        return nullptr;
+    }
     auto sess    = dev_ctx->session();
+    if (sess && sess->max_vmem && size > sess->max_vmem) {
+        GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_vmem %zu\n",
+                       dev_ctx->c_name(), size, sess->max_vmem);
+        return nullptr;
+    }
     try {
         ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false);
         return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_buffer_interface, sbuf, size);
@@ -2037,7 +2099,17 @@ static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
 static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer(
             ggml_backend_buffer_type_t buffer_type, size_t size) {
     auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
+    if (size > dev_ctx->max_bufsize) {
+        GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_bufsize %zu\n",
+                       dev_ctx->c_name(), size, dev_ctx->max_bufsize);
+        return nullptr;
+    }
     auto sess    = dev_ctx->session();
+    if (sess && sess->max_vmem && size > sess->max_vmem) {
+        GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_vmem %zu\n",
+                       dev_ctx->c_name(), size, sess->max_vmem);
+        return nullptr;
+    }
     try {
         ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false);
         return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_host_buffer_interface, sbuf, size);
@@ -2067,7 +2139,9 @@ static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffe

 static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) {
     auto * context = static_cast<ggml_backend_hexagon_buffer_type_context *>(buft->context);
-    return context->dev_ctx->max_bufsize;
+    auto dev_ctx = context->dev_ctx;
+    dev_ctx->session();
+    return dev_ctx->max_bufsize;
 }

 static bool ggml_backend_hexagon_buffer_type_is_host(ggml_backend_buffer_type_t buft) {
@@ -2139,7 +2213,7 @@ struct ggml_hexagon_opbatch {
     unsigned int n_bufs;     // num buffers in the batch
     unsigned int n_tens;     // num tensors ...
     unsigned int n_ops;      // num ops ...
-    size_t       b_vmem;     // sum of all buffer sizes
+    size_t       b_vmem;     // sum of non-extended buffer sizes

     unsigned int n_bufs_max;
     unsigned int n_tens_max;
@@ -2198,11 +2272,14 @@ struct ggml_hexagon_opbatch {
         b_map.insert({sbuf->fd(), bi});

         htp_buf_desc &b = h_bufs[bi];
-        b.base = (uint64_t) sbuf->base();
-        b.fd   = sbuf->fd();
-        b.size = sbuf->size();
+        b.base  = (uint64_t) sbuf->base();
+        b.fd    = sbuf->fd();
+        b.size  = sbuf->size();
+        b.flags = sbuf->extended ? HTP_BUF_EXTENDED : 0;

-        b_vmem += b.size;
+        if (!sbuf->extended) {
+            b_vmem += b.size;
+        }

         HEX_VERBOSE("ggml-hex: %s add-buffer #%u : fd %d base %p size %zu : vmem %zu\n", sess->c_name(), bi, b.fd, (void*) sbuf->base(), (size_t) b.size, b_vmem);

@@ -2298,21 +2375,33 @@ struct ggml_hexagon_opbatch {
     }

     bool fit_op(const htp_opnode & node) const {
-        if (n_ops >= n_ops_max ) return false;
+        if (n_ops >= n_ops_max) return false;

         // check how much extras we will need
         size_t extra_bufs = 0;
         size_t extra_vmem = 0;
         size_t extra_tens = 0;

+        int seen_bufs[HTP_OP_MAX_BUFS];
+        int n_seen_bufs = 0;
+
         auto fit_tensor = [&](const ggml_tensor *t) {
             if (!t) return;
             if (!t_map.count(t)) {
                 extra_tens++;

                 auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                if (!b_map.count(sbuf->fd())) {
-                    extra_vmem += sbuf->size();
+                int fd = sbuf->fd();
+                if (!b_map.count(fd)) {
+                    for (int i = 0; i < n_seen_bufs; i++) {
+                        if (seen_bufs[i] == fd) return;
+                    }
+                    if (n_seen_bufs < HTP_OP_MAX_BUFS) {
+                        seen_bufs[n_seen_bufs++] = fd;
+                    }
+                    if (!sbuf->extended) {
+                        extra_vmem += sbuf->size();
+                    }
                     extra_bufs += 1;
                 }
             }
@@ -2402,6 +2491,50 @@ struct ggml_hexagon_opbatch {
         }
     }

+    bool try_fuse_common(std::initializer_list<const ggml_tensor *> tensors) const {
+        size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
+
+        int seen_bufs[HTP_OP_MAX_BUFS];
+        int n_seen_bufs = 0;
+
+        for (const auto * t : tensors) {
+            if (!t || t_map.count(t)) {
+                continue;
+            }
+            extra_tens++;
+            auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
+            int fd = sbuf->fd();
+            if (!b_map.count(fd)) {
+                bool found = false;
+                for (int i = 0; i < n_seen_bufs; i++) {
+                    if (seen_bufs[i] == fd) {
+                        found = true;
+                        break;
+                    }
+                }
+                if (!found) {
+                    if (n_seen_bufs < HTP_OP_MAX_BUFS) {
+                        seen_bufs[n_seen_bufs++] = fd;
+                    }
+                    if (!sbuf->extended) {
+                        extra_vmem += sbuf->size();
+                    }
+                    extra_bufs += 1;
+                }
+            }
+        }
+
+        if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+            return false;
+        }
+
+        return true;
+    }
+
+    bool try_fuse_common(const ggml_tensor * t1, const ggml_tensor * t2) const {
+        return try_fuse_common({t1, t2});
+    }
+
     bool try_fuse_allreduce_add(const htp_opnode & node) {
         if (n_ops == 0 || opt_ar_select != 2) return false;
         if (node.opcode != HTP_OP_ADD) return false;
@@ -2466,20 +2599,7 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-        auto fit_t = [&](const ggml_tensor * t) {
-            if (!t_map.count(t)) {
-                extra_tens++;
-                auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                if (!b_map.count(sbuf->fd())) {
-                    extra_vmem += sbuf->size();
-                    extra_bufs += 1;
-                }
-            }
-        };
-        fit_t(res_tensor);
-        fit_t(add_dst);
-        if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+        if (!try_fuse_common(res_tensor, add_dst)) {
             return false;
         }

@@ -2556,20 +2676,7 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-        auto fit_t = [&](const ggml_tensor * t) {
-            if (!t_map.count(t)) {
-                extra_tens++;
-                auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                if (!b_map.count(sbuf->fd())) {
-                    extra_vmem += sbuf->size();
-                    extra_bufs += 1;
-                }
-            }
-        };
-        fit_t(weight);
-        fit_t(node.dst());
-        if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+        if (!try_fuse_common(weight, node.dst())) {
             return false;
         }

@@ -2626,8 +2733,15 @@ struct ggml_hexagon_opbatch {
         const ggml_tensor * src0 = last_node.src0();
         const ggml_tensor * src1 = last_node.src1();

+        if (src2->type != GGML_TYPE_F32) return false;
+
+        const struct htp_mm_kernel_params * orig_kparams = (const struct htp_mm_kernel_params *) last_node.kernel_params;
         struct htp_mm_kernel_params kparams;
         ggml_hexagon_precompute_fused_matmul_add_params(sess, src0, src1, src2, node.dst(), &kparams);
+        if (kparams.kernel_type == HTP_MM_KERNEL_UNSUPPORTED) {
+            return false;
+        }
+
         const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3];
         const bool can_fuse = (kparams.n_hmx > 0) || (src1_nrows == 1);
         if (!can_fuse) return false;
@@ -2638,20 +2752,19 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-        auto fit_t = [&](const ggml_tensor * t) {
-            if (!t_map.count(t)) {
-                extra_tens++;
-                auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                if (!b_map.count(sbuf->fd())) {
-                    extra_vmem += sbuf->size();
-                    extra_bufs += 1;
-                }
+        if (kparams.n_hmx > 0 && orig_kparams->n_hmx > 0) {
+            if (kparams.m_chunk < orig_kparams->m_chunk ||
+                kparams.n_chunk < orig_kparams->n_chunk ||
+                kparams.n_act_threads < orig_kparams->n_act_threads) {
+                HEX_VERBOSE("ggml-hex: %s skip MUL_MAT_ADD fusion: HMX efficiency reduced (m %d->%d, n %d->%d, th %d->%d)\n",
+                            sess->c_name(), orig_kparams->m_chunk, kparams.m_chunk,
+                            orig_kparams->n_chunk, kparams.n_chunk,
+                            orig_kparams->n_act_threads, kparams.n_act_threads);
+                return false;
             }
-        };
-        fit_t(src2);
-        fit_t(node.dst());
-        if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+        }
+
+        if (!try_fuse_common(src2, node.dst())) {
             return false;
         }

@@ -2723,20 +2836,7 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

-            size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-            auto fit_t = [&](const ggml_tensor * t) {
-                if (!t_map.count(t)) {
-                    extra_tens++;
-                    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                    if (!b_map.count(sbuf->fd())) {
-                        extra_vmem += sbuf->size();
-                        extra_bufs += 1;
-                    }
-                }
-            };
-            fit_t(w_in);
-            fit_t(d_in);
-            if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+            if (!try_fuse_common(w_in, d_in)) {
                 return false;
             }

@@ -2787,20 +2887,7 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

-            size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-            auto fit_t = [&](const ggml_tensor * t) {
-                if (!t_map.count(t)) {
-                    extra_tens++;
-                    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                    if (!b_map.count(sbuf->fd())) {
-                        extra_vmem += sbuf->size();
-                        extra_bufs += 1;
-                    }
-                }
-            };
-            fit_t(w1);
-            fit_t(node.dst());
-            if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+            if (!try_fuse_common(w1, node.dst())) {
                 return false;
             }

@@ -2882,20 +2969,7 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

-            size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-            auto fit_t = [&](const ggml_tensor * t) {
-                if (!t_map.count(t)) {
-                    extra_tens++;
-                    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                    if (!b_map.count(sbuf->fd())) {
-                        extra_vmem += sbuf->size();
-                        extra_bufs += 1;
-                    }
-                }
-            };
-            fit_t(w_in);
-            fit_t(d_in);
-            if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+            if (!try_fuse_common(w_in, d_in)) {
                 return false;
             }

@@ -2948,20 +3022,7 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

-            size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
-            auto fit_t = [&](const ggml_tensor * t) {
-                if (!t_map.count(t)) {
-                    extra_tens++;
-                    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
-                    if (!b_map.count(sbuf->fd())) {
-                        extra_vmem += sbuf->size();
-                        extra_bufs += 1;
-                    }
-                }
-            };
-            fit_t(w1);
-            fit_t(node.dst());
-            if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
+            if (!try_fuse_common(w1, node.dst())) {
                 return false;
             }

@@ -3005,6 +3066,70 @@ struct ggml_hexagon_opbatch {
         return false;
     }

+    bool try_fuse_gdn_cpy(const htp_opnode & node) {
+        if (n_ops == 0 || node.opcode != HTP_OP_CPY) return false;
+
+        htp_opnode & last_node = ops[n_ops - 1];
+        if (last_node.opcode != HTP_OP_GATED_DELTA_NET) return false;
+        if (last_node.outputs.size() != 1) return false;
+
+        const ggml_tensor * gdn_out  = last_node.dst();
+        const ggml_tensor * cpy_node = node.node;
+        const ggml_tensor * cpy_src  = node.src0();
+        const ggml_tensor * cpy_dst  = node.dst();
+
+        if (!cpy_src || !cpy_dst || !cpy_dst->data) return false;
+        if (gdn_out->type != GGML_TYPE_F32 || cpy_src->type != GGML_TYPE_F32 || cpy_dst->type != GGML_TYPE_F32) return false;
+        if ((gdn_out->flags & GGML_TENSOR_FLAG_OUTPUT) || (cpy_node->flags & GGML_TENSOR_FLAG_OUTPUT)) return false;
+
+        const ggml_tensor * v = last_node.node->src[2];
+        if (!v) return false;
+
+        const int64_t S_v      = v->ne[0];
+        const int64_t H        = v->ne[1];
+        const int64_t n_tokens = v->ne[2];
+        const int64_t n_seqs   = v->ne[3];
+        const int64_t K        = ggml_get_op_params_i32(last_node.node, 0);
+        const size_t  tail_off = (size_t) S_v * H * n_tokens * n_seqs * sizeof(float);
+
+        const int64_t D         = S_v * S_v * H;
+        const int64_t n_written = std::min<int64_t>(n_tokens, K);
+
+        if (cpy_src->op != GGML_OP_VIEW || (cpy_src->view_src != gdn_out && cpy_src->view_src->data != gdn_out->data) ||
+            cpy_src->view_offs != tail_off || !ggml_is_contiguous(cpy_src)) {
+            return false;
+        }
+
+        if (cpy_dst->ne[0] != D || cpy_dst->ne[1] != n_seqs || cpy_dst->nb[0] != sizeof(float)) {
+            return false;
+        }
+        if (n_seqs > 1 && cpy_dst->nb[1] != (size_t) D * sizeof(float)) {
+            return false;
+        }
+        if (n_written > 1) {
+            if (cpy_dst->ne[2] != n_written || cpy_dst->nb[2] != (size_t) D * n_seqs * sizeof(float)) {
+                return false;
+            }
+        }
+
+        if (!try_fuse_common({cpy_dst})) {
+            return false;
+        }
+
+        last_node.name += "+CPY";
+        last_node.outputs.push_back(cpy_dst);
+        last_node.fused.push_back(node.node);
+
+        htp_op_desc & o = h_ops[n_ops - 1];
+        o.dst[1] = add_tensor(cpy_dst);
+        for (uint32_t d = 2; d < HTP_OP_MAX_OUTPUTS; d++) {
+            o.dst[d] = 0xffff;
+        }
+
+        HEX_VERBOSE("ggml-hex: %s fused GATED_DELTA_NET+CPY (#%u)\n", sess->c_name(), n_ops - 1);
+        return true;
+    }
+
     bool try_fuse(const htp_opnode & node) {
         if (!opt_opfusion) return false;
         if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_ALLREDUCE_ADD) && try_fuse_allreduce_add(node)) return true;
@@ -3012,6 +3137,7 @@ struct ggml_hexagon_opbatch {
         if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ADD)   && try_fuse_mul_mat_add(node))   return true;
         if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_NX)    && try_fuse_mul_mat_nx(node))    return true;
         if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ID_NX) && try_fuse_mul_mat_id_nx(node)) return true;
+        if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_GDN_CPY)       && try_fuse_gdn_cpy(node))       return true;
         return false;
     }
 };
@@ -3130,8 +3256,8 @@ struct ggml_hexagon_opqueue {
             }
             htp_tensor *t = (htp_tensor*) t_ptr;
             for (unsigned int i=0; i < req.n_tensors; i++) {
-                GGML_LOG_DEBUG("ggml-hex: %s htp-tensor #%u : bi %u offset %u size %u : %zu:%zu:%zu:%zu\n",
-                            shm_buf->sess->c_name(), i, t[i].bi, t[i].data, t[i].size,
+                GGML_LOG_DEBUG("ggml-hex: %s htp-tensor #%u : bi %u offset %llu size %u : %zu:%zu:%zu:%zu\n",
+                            shm_buf->sess->c_name(), i, t[i].bi, (unsigned long long) t[i].data, t[i].size,
                             (size_t) t[i].ne[0], (size_t) t[i].ne[1], (size_t) t[i].ne[2], (size_t) t[i].ne[3]);
             }
         }
@@ -3269,6 +3395,12 @@ void ggml_hexagon_session::flush_sync(bool all) {
     flush_pending(all);
 }

+void ggml_hexagon_session::start_batch() {
+    if (this->mdev.count > 1) {
+        enqueue_mdev_group();
+    }
+}
+
 void ggml_hexagon_session::flush_batch(size_t min_ops) {
     if (op_batch->n_ops < min_ops) { return; }

@@ -3318,14 +3450,13 @@ void ggml_hexagon_session::flush_batch(size_t min_ops) {

 void ggml_hexagon_session::enqueue_op(const htp_opnode & node) {
     auto clone_tensor_buffer = [this](const ggml_tensor * t) {
-        if (t && t->buffer && ggml_backend_buffer_is_hexagon(t->buffer)) {
-            auto sbuf = static_cast<const ggml_hexagon_shared_buffer *>(t->buffer->context);
-            if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != this) {
-                this->clone_buffer(sbuf);
-            }
-            for (auto & sub : this->mdev.sessions) {
-                sub->clone_buffer(sbuf);
-            }
+        auto sbuf = this->mmap_tensor(t);
+        if (!sbuf) return;
+        if (sbuf->sess != this) {
+            this->clone_buffer(sbuf);
+        }
+        for (auto & sub : this->mdev.sessions) {
+            sub->clone_buffer(sbuf);
         }
     };

@@ -3344,8 +3475,13 @@ void ggml_hexagon_session::enqueue_op(const htp_opnode & node) {
         flush_async();
     }

-    if (this->mdev.count > 1 && op_batch->n_ops == 0) {
-        enqueue_mdev_group();
+    if (op_batch->empty()) {
+        start_batch();
+    }
+
+    if (!op_batch->fit_op(node)) {
+        GGML_ABORT("ggml-hex: %s op does not fit into empty batch (vmem/tensor/buffer limit exceeded)\n",
+                   c_name());
     }

     op_batch->add_op(node);
@@ -3358,19 +3494,19 @@ void ggml_hexagon_session::enqueue_mdev_group() {

     static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };
     ggml_tensor dummy_t {};
-    dummy_t.buffer = &this->fence_buf->backend_buffer;
-    dummy_t.extra  = &fence_extra;
-    dummy_t.data   = (void *) fence_slot;
-    dummy_t.type   = GGML_TYPE_I8;
-    dummy_t.ne[0]  = HTP_FENCE_SLOT_SIZE;
-    dummy_t.ne[1]  = (int64_t) this->mdev.count;
-    dummy_t.ne[2]  = 1;
-    dummy_t.ne[3]  = 1;
-    dummy_t.nb[0]  = 1;
-    dummy_t.nb[1]  = HTP_FENCE_SLOT_SIZE;
-    dummy_t.nb[2]  = dummy_t.nb[1] * dummy_t.ne[1];
-    dummy_t.nb[3]  = dummy_t.nb[2];
-    dummy_t.op     = GGML_OP_NONE;
+    dummy_t.buffer       = &this->fence_buf->backend_buffer;
+    dummy_t.extra        = &fence_extra;
+    dummy_t.data         = (void *) fence_slot;
+    dummy_t.type         = GGML_TYPE_I8;
+    dummy_t.ne[0]        = HTP_FENCE_SLOT_SIZE;
+    dummy_t.ne[1]        = (int64_t) this->mdev.count;
+    dummy_t.ne[2]        = 1;
+    dummy_t.ne[3]        = 1;
+    dummy_t.nb[0]        = 1;
+    dummy_t.nb[1]        = HTP_FENCE_SLOT_SIZE;
+    dummy_t.nb[2]        = dummy_t.nb[1] * dummy_t.ne[1];
+    dummy_t.nb[3]        = dummy_t.nb[2];
+    dummy_t.op           = GGML_OP_NONE;
     dummy_t.op_params[0] = (int32_t) this->mdev.idx;

     ggml_tensor * node = group_node.add_dummy(dummy_t);
@@ -3379,9 +3515,6 @@ void ggml_hexagon_session::enqueue_mdev_group() {
     group_node.outputs.clear();
     group_node.name = "MDEV_GROUP";

-    if (this->fence_buf->sess != this) {
-        this->clone_buffer(this->fence_buf);
-    }
     for (auto & sub : this->mdev.sessions) {
         sub->clone_buffer(this->fence_buf);
     }
@@ -3588,7 +3721,19 @@ void ggml_hexagon_session::enqueue_allreduce(
     this->enqueue_op(ar_node);
 }

-bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer *sbuf)
+ggml_hexagon_shared_buffer * ggml_hexagon_session::mmap_tensor(const ggml_tensor * t) {
+    if (!t) return nullptr;
+
+    auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
+    if (!sbuf->mapped) {
+        const bool is_weight = ggml_backend_buffer_get_usage(t->buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS;
+        const bool extended  = opt_dma64 && is_weight;
+        sbuf->mmap(extended);
+    }
+    return sbuf;
+}
+
+bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer * sbuf)
 {
     GGML_ASSERT(sbuf && sbuf->mem);
     if (sbuf->sess == this) return true;
@@ -3600,12 +3745,14 @@ bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer *sbuf)

     if (this->cloned_buffers.find(fd) != this->cloned_buffers.end()) return true;

+    GGML_ASSERT(sbuf->mapped);
+
     HEX_VERBOSE("ggml-hex: %s clone-buffer: %s base %p size %zu fd %d\n", this->name.c_str(),
                 sbuf->c_name(), sbuf->base(), sbuf->size(), fd);

     auto clone = std::make_unique<ggml_hexagon_shared_buffer>(this, *sbuf);
     try {
-        clone->mmap();
+        clone->mmap(sbuf->extended);
     } catch (const std::exception & exc) {
         GGML_LOG_ERROR("ggml-hex: %s lazy mapping of buffer context failed: %s\n", this->c_name(), exc.what());
         return false;
@@ -4081,7 +4228,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(
         const uint32_t DK_pad = hex_round_up(DK, 64);
         const uint32_t DV_pad = hex_round_up(DV, 64);
         size_t Br = 0, Bc = 0;
-        int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0);
+        int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0, sinks != nullptr, n_head);
         if (ret == 0) {
             kparams->kernel_type = HTP_FA_KERNEL_HMX;
             kparams->Br = Br;
@@ -4091,7 +4238,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(

             kparams->u.hmx.g_br = hex_align_up(G * Br, 32);
             kparams->u.hmx.pipeline = (kparams->n_kv_blocks >= 3 && sess->n_threads >= 2) ? 1 : 0;
-            kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0);
+            kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0, sinks != nullptr, n_head);

             const size_t row_vec_bytes = hex_align_up(Bc * sizeof(uint16_t), 256);
             kparams->u.hmx.row_buf_stride = row_vec_bytes / 128; // HVX vector is 128 bytes
@@ -4122,7 +4269,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(
     const size_t size_k_row_padded = hex_round_up(k->ne[0] * 2, 128);
     const size_t size_v_row_padded = hex_round_up(v->ne[0] * 2, 128);

-    kparams->vtcm_size = hvx_fa_compute_vtcm_usage(DK, DV, kparams->is_q_fp32 != 0, mask != nullptr, sess->n_threads);
+    kparams->vtcm_size = hvx_fa_compute_vtcm_usage(DK, DV, kparams->is_q_fp32 != 0, mask != nullptr, sinks != nullptr, n_head, sess->n_threads);

     kparams->u.hvx.size_q_row_padded = size_q_row_padded;
     kparams->u.hvx.size_k_row_padded = size_k_row_padded;
@@ -4238,9 +4385,15 @@ static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_ses
         return false;
     }

-    return true;
+    const uint32_t total_rows = (uint32_t) (H * n_seqs);
+    const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, total_rows);
+    struct htp_gdn_vtcm_layout layout;
+    htp_gdn_vtcm_layout_build(&layout, (uint32_t) S_v, n_threads ? n_threads : 1);
+    if (layout.total_bytes > sess->vtcm_size) {
+        return false;
+    }

-    GGML_UNUSED(sess);
+    return true;
 }

 static bool ggml_hexagon_matmul_is_hmx_eligible(
@@ -4311,6 +4464,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     int ne11_padded,
     bool is_matmul_id,
     bool is_batched,
+    size_t src2_size,
     size_t vtcm_budget,
     struct htp_mm_kernel_params * kparams
 ) {
@@ -4331,7 +4485,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
         // Try grouped path first
         const bool use_dma_activation = (src1->nb[1]/sizeof(float) > (size_t)ne00_padded);
-        if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, use_dma_activation, n_threads, pipeline, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+        if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, use_dma_activation, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
             use_grouped = true;
         }
     }
@@ -4339,7 +4493,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     if (!use_grouped) {
         // Fallback to simple 2D path (group_size = 1)
         const int m_id_rows = (dst && is_matmul_id) ? (int) ((size_t) dst->ne[1] * dst->ne[2]) : 0;
-        if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, ne11_padded, ne11, n_threads, pipeline, is_matmul_id, aligned_tile_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+        if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, ne11_padded, ne11, n_threads, pipeline, is_matmul_id, aligned_tile_size, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
             return false;
         }
     }
@@ -4358,6 +4512,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     kparams->div_n_act_threads = init_fastdiv_values(act_threads_selected);
     kparams->div_ne00_padded   = init_fastdiv_values(ne00_padded);
     kparams->vtcm_src1_size = 0;
+    kparams->vtcm_src2_size = (int32_t) src2_size;
     kparams->vtcm_dst_size = 0;

     if (is_batched && !is_matmul_id) {
@@ -4387,6 +4542,11 @@ static void ggml_hexagon_precompute_hvx_mm_params(
     size_t vtcm_budget,
     struct htp_mm_kernel_params * kparams
 ) {
+    if (opt_mm_select < 1) {
+        kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+        return;
+    }
+
     kparams->n_hmx = 0;
     kparams->n_threads = sess->n_threads;

@@ -4410,29 +4570,30 @@ static void ggml_hexagon_precompute_hvx_mm_params(
             for (uint32_t d = max_prefetch; d >= 2; d /= 2) {
                 htp_mm_hvx_vtcm_layout_build(
                     &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                    0, src0->nb[1], 0, src2_row_size, d, true, false
+                    0, src0->nb[1], kparams->src1_row_size, 0, d, true, false
                 );
                 if (L.total_bytes <= vtcm_budget) {
                     best_n_prefetch = d;
                     break;
                 }
             }
-            if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) {
-                htp_mm_hvx_vtcm_layout_build(
-                    &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                    0, src0->nb[1], 0, src2_row_size, 2, true, false
-                );
+            if (L.total_bytes > vtcm_budget) {
+                kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+                return;
             }
-            kparams->n_prefetch = best_n_prefetch;
+            kparams->n_prefetch     = best_n_prefetch;
             kparams->vtcm_size      = L.total_bytes;
             kparams->vtcm_src0_size = L.src0_bytes;
             kparams->vtcm_src1_size = L.src1_bytes;
             kparams->vtcm_dst_size  = L.dst_bytes;
+            goto done_quant;
         } else {
-            bool try_tiled = (k_align && opt_mm_select >= 2);
+            bool try_tiled = (k_align && opt_mm_select >= 1);
             if (try_tiled) {
-                kparams->src1_row_size = (wtype == GGML_TYPE_Q4_1 || wtype == GGML_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
-                if (src1_nrows < (int)sess->n_threads) {
+                kparams->src1_row_size = (wtype == GGML_TYPE_Q4_1 || wtype == GGML_TYPE_Q4_K)
+                                       ? htp_mm_q8_1_tiled_row_size(ne10)
+                                       : htp_mm_q8_0_tiled_row_size(ne10);
+                if (src1_nrows < (int) sess->n_threads) {
                     kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_BLOCK;
                 } else {
                     kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
@@ -4451,113 +4612,72 @@ static void ggml_hexagon_precompute_hvx_mm_params(
                         break;
                     }
                 }
-                if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) {
-                    htp_mm_hvx_vtcm_layout_build(
-                        &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                        dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 2, false, false
-                    );
-                }
-
-                kparams->n_prefetch = best_n_prefetch;

-                if (L.total_bytes <= vtcm_budget) {
-                    kparams->vtcm_size = L.total_bytes;
+                uint32_t m_chunk = 0;
+                if (htp_mm_hvx_solve_vtcm_params(
+                        kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
+                        dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, best_n_prefetch, vtcm_budget,
+                        &L, &m_chunk)) {
+                    kparams->n_prefetch     = best_n_prefetch;
+                    kparams->m_chunk        = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
+                    kparams->vtcm_size      = L.total_bytes;
                     kparams->vtcm_src0_size = L.src0_bytes;
                     kparams->vtcm_src1_size = L.src1_bytes;
-                    kparams->vtcm_dst_size = L.dst_bytes;
+                    kparams->vtcm_src2_size = L.src2_bytes;
+                    kparams->vtcm_dst_size  = L.dst_bytes;
                     goto done_quant;
                 }
-                HEX_VERBOSE("ggml-hex: %s HVX tiled path VTCM size needed (%zu) > budget (%zu), falling back to HVX flat\n", sess->name.c_str(), L.total_bytes, vtcm_budget);
             }

-            // Flat HVX fallback
-            {
-                kparams->src1_row_size = (wtype == GGML_TYPE_Q4_1 || wtype == GGML_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
-                kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
-
-                struct htp_mm_hvx_vtcm_layout L;
-                htp_mm_hvx_vtcm_layout_build(
-                    &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                    dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false
-                );
-
-                kparams->n_prefetch = 16;
-                kparams->vtcm_size = L.total_bytes;
-                kparams->vtcm_src0_size = L.src0_bytes;
-                kparams->vtcm_src1_size = L.src1_bytes;
-                kparams->vtcm_dst_size = L.dst_bytes;
-            }
+            kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+            return;
         }

     done_quant:;
     } else if (wtype == GGML_TYPE_F16) {
         // F16 HVX
-        const bool is_batched  = (ne02 > 1) || (ne03 > 1);
-        const bool is_permuted = ggml_is_permuted(src0) || ggml_is_permuted(src1);
-
         struct htp_mm_hvx_vtcm_layout L;
-        htp_mm_hvx_vtcm_layout_build(
-            &L, HTP_MM_KERNEL_HVX_F16_F16_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
-            dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false
-        );
-
-        if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) {
+        uint32_t m_chunk = 0;
+        if (htp_mm_hvx_solve_vtcm_params(
+                HTP_MM_KERNEL_HVX_F16_F16_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
+                dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, vtcm_budget,
+                &L, &m_chunk)) {
             kparams->kernel_type = HTP_MM_KERNEL_HVX_F16_F16_VTCM;
+            kparams->m_chunk = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
             kparams->src1_row_size = hex_round_up(ne10 * 2, 128);
             kparams->vtcm_size = L.total_bytes;
             kparams->vtcm_src0_size = L.src0_bytes;
             kparams->vtcm_src1_size = L.src1_bytes;
+            kparams->vtcm_src2_size = L.src2_bytes;
             kparams->vtcm_dst_size = L.dst_bytes;
             kparams->n_prefetch = 16;
-        } else {
-            if (src1->type == GGML_TYPE_F32) {
-                kparams->kernel_type = HTP_MM_KERNEL_HVX_F16_F32_DDR;
-            } else {
-                kparams->kernel_type = HTP_MM_KERNEL_HVX_F16_F16_DDR;
-            }
-            kparams->src1_row_size = src1->nb[1];
-            htp_mm_hvx_vtcm_layout_build(
-                &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false
-            );
-            kparams->vtcm_size = L.total_bytes;
-            kparams->vtcm_src0_size = L.src0_bytes;
-            kparams->vtcm_src1_size = L.src1_bytes;
-            kparams->vtcm_dst_size = L.dst_bytes;
-            kparams->n_prefetch = 16;
+            return;
         }
+
+        kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+        return;
     } else {
         // F32 HVX
-        const bool is_batched  = (ne02 > 1) || (ne03 > 1);
-        const bool is_permuted = ggml_is_permuted(src0) || ggml_is_permuted(src1);
-
         struct htp_mm_hvx_vtcm_layout L;
-        htp_mm_hvx_vtcm_layout_build(
-            &L, HTP_MM_KERNEL_HVX_F32_F32_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
-            dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false
-        );
-
-        if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) {
+        uint32_t m_chunk = 0;
+        if (htp_mm_hvx_solve_vtcm_params(
+                HTP_MM_KERNEL_HVX_F32_F32_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
+                dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, vtcm_budget,
+                &L, &m_chunk)) {
             kparams->kernel_type = HTP_MM_KERNEL_HVX_F32_F32_VTCM;
+            kparams->m_chunk = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
             kparams->src1_row_size = hex_round_up(ne10 * 4, 128);
             kparams->vtcm_size = L.total_bytes;
             kparams->vtcm_src0_size = L.src0_bytes;
             kparams->vtcm_src1_size = L.src1_bytes;
+            kparams->vtcm_src2_size = L.src2_bytes;
             kparams->vtcm_dst_size = L.dst_bytes;
             kparams->n_prefetch = 16;
-        } else {
-            kparams->kernel_type = HTP_MM_KERNEL_HVX_F32_F32_DDR;
-            kparams->src1_row_size = src1->nb[1];
-            htp_mm_hvx_vtcm_layout_build(
-                &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
-                dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false
-            );
-            kparams->vtcm_size = L.total_bytes;
-            kparams->vtcm_src0_size = L.src0_bytes;
-            kparams->vtcm_src1_size = L.src1_bytes;
-            kparams->vtcm_dst_size = L.dst_bytes;
-            kparams->n_prefetch = 16;
+            return;
         }
+
+        kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+        return;
     }
 }

@@ -4567,6 +4687,7 @@ static void ggml_hexagon_precompute_matmul_params_impl(
     const struct ggml_tensor * src1,
     const struct ggml_tensor * dst,
     const size_t src2_row_size,
+    const size_t src2_size,
     struct htp_mm_kernel_params * kparams
 ) {
     memset(kparams, 0, sizeof(*kparams));
@@ -4593,9 +4714,9 @@ static void ggml_hexagon_precompute_matmul_params_impl(
     const size_t vtcm_budget = sess->vtcm_size;

     // Check HMX eligibility and try precomputing HMX parameters
-    bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 3);
+    bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 2);
     if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, src1, dst, ne01_padded, is_matmul_id, is_batched)) {
-        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, dst, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, is_matmul_id, is_batched, vtcm_budget, kparams)) {
+        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, dst, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, is_matmul_id, is_batched, src2_size, vtcm_budget, kparams)) {
             goto finalize;
         }
     }
@@ -4608,7 +4729,7 @@ finalize:
     kparams->div_ne1      = init_fastdiv_values(ne11);
     kparams->div_r2       = init_fastdiv_values(ne02 > 0 ? ne12 / ne02 : 1);
     kparams->div_r3       = init_fastdiv_values(ne03 > 0 ? ne13 / ne03 : 1);
-    kparams->div_ne11     = init_fastdiv_values(ne11);
+    kparams->div_ne12     = init_fastdiv_values(ne12);
 }

 static void ggml_hexagon_precompute_matmul_params(
@@ -4618,7 +4739,7 @@ static void ggml_hexagon_precompute_matmul_params(
     const struct ggml_tensor * dst,
     struct htp_mm_kernel_params * kparams
 ) {
-    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, kparams);
+    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
 }

 static void ggml_hexagon_precompute_fused_matmul_add_params(
@@ -4629,7 +4750,83 @@ static void ggml_hexagon_precompute_fused_matmul_add_params(
     const struct ggml_tensor * dst,
     struct htp_mm_kernel_params * kparams
 ) {
-    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, src2->nb[1], kparams);
+    const size_t src2_size = src2 ? hex_round_up(ggml_nbytes(src2), 128) : 0;
+    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, src2 ? src2->nb[1] : 0, src2_size, kparams);
+}
+
+static bool ggml_hexagon_precompute_binary_params(
+    const struct ggml_hexagon_session * sess,
+    uint32_t op,
+    const struct ggml_tensor * src0,
+    const struct ggml_tensor * src1,
+    const struct ggml_tensor * dst,
+    struct htp_binary_kernel_params * kparams
+) {
+    memset(kparams, 0, sizeof(*kparams));
+
+    const size_t elem_size = ggml_type_size(src0->type);
+    const size_t src0_row_size = src0->ne[0] * elem_size;
+    const size_t src1_row_size = src1->ne[0] * elem_size;
+    const size_t dst_row_size  = dst->ne[0]  * elem_size;
+
+    const size_t src0_row_size_aligned = hex_round_up(src0_row_size, 128);
+    const size_t src1_row_size_aligned = hex_round_up(src1_row_size, 128);
+    const size_t dst_row_size_aligned  = hex_round_up(dst_row_size,  128);
+
+    const bool is_add_id = op == HTP_OP_ADD_ID;
+    const bool is_scalar = !is_add_id && src1->ne[0] == 1;
+    const bool is_transposed = src0->nb[1] < src0_row_size || src1->nb[1] < src1_row_size || dst->nb[1] < dst_row_size;
+    const bool is_same_shape = !is_add_id && !is_scalar && !is_transposed &&
+        src1->ne[0] == src0->ne[0] &&
+        (src1->ne[1] == src0->ne[1] || src1->ne[1] == 1) &&
+        (src1->ne[2] == src0->ne[2] || src1->ne[2] == 1) &&
+        (src1->ne[3] == src0->ne[3] || src1->ne[3] == 1);
+    const bool is_row_bcast = is_same_shape && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
+    const bool is_complex   = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] == src0->ne[0]);
+
+    enum htp_binary_kernel_type kernel_type;
+    size_t src1_size = 0;
+
+    if (is_add_id) {
+        kernel_type = HTP_BINARY_KERNEL_ADD_ID;
+        src1_size = hex_round_up(src1->ne[1] * src1_row_size_aligned, 128);
+    } else if (is_row_bcast) {
+        kernel_type = HTP_BINARY_KERNEL_ROW_BCAST;
+        src1_size = src1_row_size_aligned;
+    } else if (is_scalar) {
+        const bool is_scalar_static = (src1->ne[2] == 1 && src1->ne[3] == 1) &&
+            (src1->ne[1] == 1 || src1->nb[1] == elem_size);
+        if (is_scalar_static) {
+            kernel_type = HTP_BINARY_KERNEL_SCALAR_DMA;
+            src1_size = hex_round_up(src1->ne[1] * elem_size, 128);
+        } else {
+            kernel_type = HTP_BINARY_KERNEL_SCALAR;
+        }
+    } else if (is_same_shape) {
+        kernel_type = HTP_BINARY_KERNEL_SAME_SHAPE;
+    } else if (is_complex) {
+        kernel_type = HTP_BINARY_KERNEL_COMPLEX;
+    } else {
+        kernel_type = HTP_BINARY_KERNEL_REPEAT;
+    }
+
+    kparams->kernel_type = kernel_type;
+    kparams->n_threads = sess->n_threads;
+    kparams->src0_row_size_aligned = src0_row_size_aligned;
+    kparams->src1_row_size_aligned = src1_row_size_aligned;
+    kparams->dst_row_size_aligned  = dst_row_size_aligned;
+    kparams->src1_size             = src1_size;
+
+    struct htp_binary_vtcm_layout L;
+    htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
+    if (L.rows_per_buffer == 0 || L.total_bytes > sess->vtcm_size) {
+        return false;
+    }
+
+    kparams->rows_per_buffer = L.rows_per_buffer;
+    kparams->vtcm_size       = L.total_bytes;
+
+    return true;
 }

 static void ggml_hexagon_precompute_unary_params(
@@ -4790,6 +4987,78 @@ static void ggml_hexagon_precompute_set_rows_params(
     kparams->vtcm_size = vtcm_layout.total_bytes;
 }

+static void ggml_hexagon_precompute_softmax_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * op,
+    struct htp_softmax_kernel_params * kparams
+) {
+    memset(kparams, 0, sizeof(*kparams));
+
+    const struct ggml_tensor * src0 = op->src[0];
+    const struct ggml_tensor * src1 = op->src[1];
+
+    const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
+    const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, src0_nrows);
+
+    float scale = 1.0f;
+    float max_bias = 0.0f;
+    memcpy(&scale,    &op->op_params[0], sizeof(float));
+    memcpy(&max_bias, &op->op_params[1], sizeof(float));
+
+    kparams->scale    = scale;
+    kparams->max_bias = max_bias;
+
+    const uint32_t n_head = src0->ne[2];
+    const uint32_t n_head_log2 = 1u << (uint32_t) floor(log2(n_head));
+    kparams->n_head      = n_head;
+    kparams->n_head_log2 = n_head_log2;
+
+    if (max_bias > 0.0f && n_head_log2 > 0) {
+        kparams->m0 = powf(2.0f, -(max_bias) / n_head_log2);
+        kparams->m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2);
+    } else {
+        kparams->m0 = 1.0f;
+        kparams->m1 = 1.0f;
+    }
+
+    kparams->use_src1 = (src1 != nullptr) ? 1 : 0;
+    kparams->use_f16  = (src1 != nullptr && src1->type == GGML_TYPE_F16) ? 1 : 0;
+
+    const uint32_t ne00 = src0->ne[0];
+    const uint32_t ne10 = src1 ? src1->ne[0] : 1;
+
+    struct htp_softmax_vtcm_layout layout;
+    htp_softmax_vtcm_layout_build(&layout, ne00, ne10, kparams->use_src1 != 0, kparams->use_f16 != 0, n_threads);
+
+    kparams->n_threads                 = n_threads;
+    kparams->src0_nrows                = src0_nrows;
+    kparams->src0_nrows_per_thread     = (src0_nrows + n_threads - 1) / n_threads;
+    kparams->vtcm_size                 = (uint32_t) layout.total_bytes;
+    kparams->vtcm_src0_size_per_thread = (uint32_t) layout.src0_bytes_per_thread;
+    kparams->vtcm_src1_size_per_thread = (uint32_t) layout.src1_bytes_per_thread;
+    kparams->vtcm_dst_size_per_thread  = (uint32_t) layout.dst_bytes_per_thread;
+    kparams->src0_row_size_aligned     = (uint32_t) layout.src0_spad_half_size;
+    kparams->src1_row_size_aligned     = (uint32_t) layout.src1_spad_half_size;
+    kparams->dst_row_size_aligned      = (uint32_t) layout.dst_spad_half_size;
+    kparams->src0_spad_half_size       = (uint32_t) layout.src0_spad_half_size;
+    kparams->src1_spad_half_size       = (uint32_t) layout.src1_spad_half_size;
+    kparams->dst_spad_half_size        = (uint32_t) layout.dst_spad_half_size;
+    if (!kparams->use_src1) {
+        kparams->kernel_id = HTP_SOFTMAX_KERNEL_NOMASK;
+    } else if (kparams->use_f16) {
+        kparams->kernel_id = HTP_SOFTMAX_KERNEL_MASK_F16;
+    } else {
+        kparams->kernel_id = HTP_SOFTMAX_KERNEL_MASK_F32;
+    }
+
+    if (src0->ne[1] > 0) kparams->div_ne01 = init_fastdiv_values(src0->ne[1]);
+    if (src0->ne[2] > 0) kparams->div_ne02 = init_fastdiv_values(src0->ne[2]);
+    const uint32_t ne12 = src1 ? src1->ne[2] : 1;
+    const uint32_t ne13 = src1 ? src1->ne[3] : 1;
+    if (ne12 > 0) kparams->div_ne12 = init_fastdiv_values(ne12);
+    if (ne13 > 0) kparams->div_ne13 = init_fastdiv_values(ne13);
+}
+
 static void ggml_hexagon_precompute_rope_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * op,
@@ -4798,13 +5067,15 @@ static void ggml_hexagon_precompute_rope_params(
     memset(kparams, 0, sizeof(*kparams));

     const struct ggml_tensor * src0 = op->src[0];
+    const struct ggml_tensor * src2 = op->src[2];
     const struct ggml_tensor * dst  = op;

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
     const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, src0_nrows);
+    const uint32_t n_freq_factors = src2 ? (uint32_t) src2->ne[0] : 0;

     struct htp_rope_vtcm_layout layout;
-    htp_rope_vtcm_layout_build(&layout, src0->ne[0], n_threads);
+    htp_rope_vtcm_layout_build(&layout, src0->ne[0], n_threads, n_freq_factors);

     kparams->n_threads              = n_threads;
     kparams->src0_nrows             = src0_nrows;
@@ -4813,6 +5084,8 @@ static void ggml_hexagon_precompute_rope_params(
     kparams->spad_per_thread        = (uint32_t) layout.bytes_per_thread;
     kparams->theta_cache_offset     = (uint32_t) layout.theta_cache_size_aligned;
     kparams->src0_row_size_aligned  = (uint32_t) layout.src0_row_size_aligned;
+    kparams->freq_factors_offset    = (uint32_t) (layout.bytes_per_thread * n_threads);
+    kparams->freq_factors_size      = (uint32_t) layout.freq_factors_size_aligned;

     if (src0_nrows > 0) {
         kparams->div_ne2_ne1 = init_fastdiv_values(dst->ne[2] * dst->ne[1]);
@@ -4820,6 +5093,146 @@ static void ggml_hexagon_precompute_rope_params(
     }
 }

+static void ggml_hexagon_precompute_ssm_conv_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * src0,
+    const struct ggml_tensor * src1,
+    const struct ggml_tensor * dst,
+    struct htp_ssm_conv_kernel_params * kparams
+) {
+    memset(kparams, 0, sizeof(*kparams));
+
+    const uint32_t d_conv  = (uint32_t) src1->ne[0];
+    const uint32_t d_inner = (uint32_t) src0->ne[1];
+    const uint32_t n_t     = (uint32_t) dst->ne[1];
+    const uint32_t n_s     = (uint32_t) dst->ne[2];
+    const uint32_t ncs     = (uint32_t) src0->ne[0];
+
+    const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, (d_inner + 31) / 32);
+
+    kparams->n_threads = n_threads;
+    kparams->d_conv    = d_conv;
+    kparams->d_inner   = d_inner;
+    kparams->n_t       = n_t;
+    kparams->n_s       = n_s;
+
+    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);
+
+    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);
+
+        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;
+    } 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));
+        d_inner_tile = (d_inner_tile / 32) * 32;
+        if (d_inner_tile == 0) {
+            d_inner_tile = 32;
+        }
+        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);
+
+        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;
+    }
+
+    kparams->div_n_threads = init_fastdiv_values(n_threads);
+}
+
+static void ggml_hexagon_precompute_gated_delta_net_params(
+    const struct ggml_hexagon_session * sess,
+    const struct ggml_tensor * op,
+    struct htp_gdn_kernel_params * kparams
+) {
+    memset(kparams, 0, sizeof(*kparams));
+
+    const struct ggml_tensor * q     = op->src[0];
+    const struct ggml_tensor * k     = op->src[1];
+    const struct ggml_tensor * v     = op->src[2];
+    const struct ggml_tensor * g     = op->src[3];
+    const struct ggml_tensor * state = op->src[5];
+
+    const uint32_t S_v      = (uint32_t) v->ne[0];
+    const uint32_t H        = (uint32_t) v->ne[1];
+    const uint32_t n_tokens = (uint32_t) v->ne[2];
+    const uint32_t n_seqs   = (uint32_t) v->ne[3];
+    const uint32_t K        = (uint32_t) ggml_get_op_params_i32(op, 0);
+
+    const uint32_t rq3 = (uint32_t) (n_seqs / q->ne[3]);
+    const uint32_t rk3 = (uint32_t) (n_seqs / k->ne[3]);
+    const uint32_t total_rows = H * n_seqs;
+    const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, total_rows);
+
+    struct htp_gdn_vtcm_layout layout;
+    htp_gdn_vtcm_layout_build(&layout, S_v, n_threads ? n_threads : 1);
+
+    kparams->n_threads           = n_threads ? n_threads : 1;
+    kparams->S_v                 = S_v;
+    kparams->H                   = H;
+    kparams->n_tokens            = n_tokens;
+    kparams->n_seqs              = n_seqs;
+    kparams->K                   = K;
+    kparams->total_rows          = total_rows;
+    kparams->rows_per_thread     = (total_rows + kparams->n_threads - 1) / kparams->n_threads;
+    kparams->kda                 = (g->ne[0] == S_v) ? 1 : 0;
+    kparams->state_aligned       = (uint32_t) layout.state_aligned;
+    kparams->vtcm_per_thread     = (uint32_t) layout.bytes_per_thread;
+    kparams->vtcm_size           = (uint32_t) layout.total_bytes;
+    kparams->state_seq_stride    = (uint32_t) (state->nb[3] / sizeof(float));
+    kparams->state_size_per_snap = S_v * S_v * H * n_seqs;
+    kparams->scale               = 1.0f / sqrtf((float) S_v);
+
+    if (H > 0)                  kparams->div_H         = init_fastdiv_values(H);
+    if (q->ne[1] > 0)           kparams->div_q1        = init_fastdiv_values((uint32_t) q->ne[1]);
+    if (k->ne[1] > 0)           kparams->div_k1        = init_fastdiv_values((uint32_t) k->ne[1]);
+    if (rq3 > 0)                kparams->div_rq3       = init_fastdiv_values(rq3);
+    if (rk3 > 0)                kparams->div_rk3       = init_fastdiv_values(rk3);
+    if (kparams->n_threads > 0) kparams->div_n_threads = init_fastdiv_values(kparams->n_threads);
+}
+
 static void ggml_hexagon_precompute_fused_mmnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0, // W0
@@ -4849,9 +5262,9 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
     const size_t vtcm_budget = sess->vtcm_size;
     const bool is_batched = (ne02 * ne03 > 1 || ne12 * ne13 > 1);

-    bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 3);
+    bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 2);
     if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, src1, nullptr, ne01_padded, false, is_batched)) {
-        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, nullptr, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, false, is_batched, vtcm_budget, kparams)) {
+        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, nullptr, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, false, is_batched, 0, vtcm_budget, kparams)) {
             kparams->n_weights = n_weights;
             goto finalize;
         }
@@ -4886,7 +5299,7 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
         }

         struct htp_mm_hvx_vtcm_layout L;
-        bool try_tiled = (opt_mm_select >= 2);
+        bool try_tiled = (opt_mm_select >= 1);

         // Test tiled first
         htp_mm_hvx_vtcm_layout_build(
@@ -4903,19 +5316,8 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
             kparams->n_prefetch     = best_n_prefetch;
             kparams->n_weights      = n_weights;
         } else {
-            kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
-            size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1 || wtype == GGML_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
-
-            htp_mm_hvx_vtcm_layout_build(
-                &L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads,
-                0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true
-            );
-            kparams->vtcm_src0_size = L.src0_bytes;
-            kparams->vtcm_src1_size = L.src1_bytes;
-            kparams->vtcm_dst_size  = L.dst_bytes;
-            kparams->vtcm_size      = L.total_bytes;
-            kparams->n_prefetch     = best_n_prefetch;
-            kparams->n_weights      = n_weights;
+            kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
+            return;
         }
     }

@@ -4924,7 +5326,7 @@ finalize:
     kparams->div_ne1      = init_fastdiv_values(ne11);
     kparams->div_r2       = init_fastdiv_values(ne02 > 0 ? ne12 / ne02 : 1);
     kparams->div_r3       = init_fastdiv_values(ne03 > 0 ? ne13 / ne03 : 1);
-    kparams->div_ne11     = init_fastdiv_values(ne11);
+    kparams->div_ne12     = init_fastdiv_values(ne12);
 }

 static void ggml_hexagon_precompute_fused_mmidnx_params(
@@ -4935,7 +5337,7 @@ static void ggml_hexagon_precompute_fused_mmidnx_params(
     int32_t n_weights,
     struct htp_mm_kernel_params * kparams
 ) {
-    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, kparams);
+    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
     kparams->n_weights = n_weights;
 }

@@ -4969,12 +5371,19 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
         case GGML_TYPE_MXFP4:
         case GGML_TYPE_Q4_K:
         case GGML_TYPE_Q6_K:
+            if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
+                return false;
+            }
+
             if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
                 return false;
             }

-            if (src1->ne[2] != 1 || src1->ne[3] != 1) {
-                return false;  // no broadcasting (for now)
+            if (src1->ne[2] < src0->ne[2] || src1->ne[3] < src0->ne[3]) {
+                return false;
+            }
+            if (src1->ne[2] % src0->ne[2] != 0 || src1->ne[3] % src0->ne[3] != 0) {
+                return false;
             }

             if (!src0->buffer) {
@@ -4989,6 +5398,9 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
             if (src1->ne[2] < src0->ne[2] || src1->ne[3] < src0->ne[3]) {
                 return false;
             }
+            if (src1->ne[2] % src0->ne[2] != 0 || src1->ne[3] % src0->ne[3] != 0) {
+                return false;
+            }
             break;

         case GGML_TYPE_F32:
@@ -5001,6 +5413,9 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
             if (src1->ne[2] < src0->ne[2] || src1->ne[3] < src0->ne[3]) {
                 return false;
             }
+            if (src1->ne[2] % src0->ne[2] != 0 || src1->ne[3] % src0->ne[3] != 0) {
+                return false;
+            }
             break;

         default:
@@ -5009,7 +5424,7 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s

     struct htp_mm_kernel_params kparams;
     ggml_hexagon_precompute_matmul_params(sess, src0, src1, dst, &kparams);
-    if ((size_t)kparams.vtcm_size > sess->vtcm_size) {
+    if (kparams.kernel_type == HTP_MM_KERNEL_UNSUPPORTED || (size_t) kparams.vtcm_size > sess->vtcm_size) {
         HEX_VERBOSE("ggml-hex: %s supported MUL_MAT VTCM size needed (%d) > budget (%zu)\n", sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
         return false;
     }
@@ -5035,6 +5450,10 @@ static bool ggml_hexagon_supported_mul_mat_id(const struct ggml_hexagon_session
         case GGML_TYPE_MXFP4:
         case GGML_TYPE_Q4_K:
         case GGML_TYPE_Q6_K:
+            if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
+                return false;
+            }
+
             if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
                 return false;
             }
@@ -5050,7 +5469,7 @@ static bool ggml_hexagon_supported_mul_mat_id(const struct ggml_hexagon_session

     struct htp_mm_kernel_params kparams;
     ggml_hexagon_precompute_matmul_params(sess, src0, src1, dst, &kparams);
-    if ((size_t)kparams.vtcm_size > sess->vtcm_size) {
+    if (kparams.kernel_type == HTP_MM_KERNEL_UNSUPPORTED || (size_t) kparams.vtcm_size > sess->vtcm_size) {
         HEX_VERBOSE("ggml-hex: %s supported MUL_MAT_ID VTCM size needed (%d) > budget (%zu)\n", sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
         return false;
     }
@@ -5093,37 +5512,42 @@ static bool ggml_hexagon_supported_binary(const struct ggml_hexagon_session * se
         return false;
     }

-    return true;
-
-    GGML_UNUSED(sess);
+    struct htp_binary_kernel_params kparams;
+    return ggml_hexagon_precompute_binary_params(sess, op_remap_to_htp(op), src0, src1, dst, &kparams);
 }

 static bool ggml_hexagon_supported_add_id(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
     const struct ggml_tensor * src0 = op->src[0];
     const struct ggml_tensor * src1 = op->src[1];
+    const struct ggml_tensor * src2 = op->src[2];
     const struct ggml_tensor * dst  = op;

-    if (src0->type != GGML_TYPE_F32) {
+    if (!src2) {
         return false;
     }
-    if (src1->type != GGML_TYPE_F32) {
+    if (src0->type != GGML_TYPE_F32 || src1->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32 || src2->type != GGML_TYPE_I32) {
         return false;
     }
-    if (dst->type != GGML_TYPE_F32) {
+    if (!ggml_are_same_shape(src0, dst)) {
         return false;
     }
-    if (!ggml_are_same_shape(src0, dst)) {
+    if (src1->ne[0] != src0->ne[0] || src1->ne[2] != 1 || src1->ne[3] != 1) {
+        return false;
+    }
+    if (src2->ne[0] != src0->ne[1] || src2->ne[1] != src0->ne[2]) {
+        return false;
+    }
+    if (src0->nb[0] != sizeof(float) || src1->nb[0] != sizeof(float) || dst->nb[0] != sizeof(float) || src2->nb[0] != sizeof(int32_t)) {
         return false;
     }

-    // REVISIT: add support for non-contigiuos tensors
+    // REVISIT: add support for non-contiguous tensors
     if (!ggml_is_contiguous(src0) || !ggml_is_contiguous(src1) || !ggml_is_contiguous(dst)) {
         return false;
     }

-    return true;
-
-    GGML_UNUSED(sess);
+    struct htp_binary_kernel_params kparams;
+    return ggml_hexagon_precompute_binary_params(sess, HTP_OP_ADD_ID, src0, src1, dst, &kparams);
 }

 static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
@@ -5251,6 +5675,10 @@ static bool ggml_hexagon_supported_softmax(const struct ggml_hexagon_session * s
         return false;
     }

+    if (src0->ne[2] > 512) {
+        return false;
+    }
+
     if (src1) {
         if (src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_F16) {
             return false;
@@ -5296,6 +5724,14 @@ static bool ggml_hexagon_supported_softmax(const struct ggml_hexagon_session * s
         return false;
     }

+    const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
+    const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, src0_nrows);
+    struct htp_softmax_vtcm_layout layout;
+    htp_softmax_vtcm_layout_build(&layout, src0->ne[0], src1 ? src1->ne[0] : 1, src1 != nullptr, src1 && src1->type == GGML_TYPE_F16, n_threads);
+    if (layout.total_bytes > sess->vtcm_size) {
+        return false;
+    }
+
     return true;

     GGML_UNUSED(sess);
@@ -5505,9 +5941,10 @@ static bool ggml_hexagon_supported_rope(const struct ggml_hexagon_session * sess
     }

     const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, src0_nrows);
+    const uint32_t n_freq_factors = src2 ? (uint32_t) src2->ne[0] : 0;

     struct htp_rope_vtcm_layout layout;
-    htp_rope_vtcm_layout_build(&layout, src0->ne[0], n_threads);
+    htp_rope_vtcm_layout_build(&layout, src0->ne[0], n_threads, n_freq_factors);
     if (layout.total_bytes > sess->vtcm_size) {
         return false;
     }
@@ -5530,11 +5967,14 @@ static bool ggml_hexagon_supported_ssm_conv(const struct ggml_hexagon_session *
         return false; // src0 should be effectively 3D
     }

-    const int d_conv = src1->ne[0];
+    const int d_conv  = src1->ne[0];
     const int d_inner = src0->ne[1];
     const int n_t = dst->ne[1];
     const int n_s = dst->ne[2];

+    if (d_conv == 0 || d_conv > 32 || d_inner == 0) {
+        return false;
+    }
     if (src0->ne[0] != d_conv - 1 + n_t || src0->ne[1] != d_inner || src0->ne[2] != n_s) {
         return false;
     }
@@ -5551,9 +5991,13 @@ static bool ggml_hexagon_supported_ssm_conv(const struct ggml_hexagon_session *
         return false;
     }

-    return true;
+    struct htp_ssm_conv_kernel_params kparams;
+    ggml_hexagon_precompute_ssm_conv_params(sess, src0, src1, dst, &kparams);
+    if ((size_t) kparams.vtcm_size > sess->vtcm_size) {
+        return false;
+    }

-    GGML_UNUSED(sess);
+    return true;
 }

 static bool ggml_hexagon_supported_im2col(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
@@ -5796,7 +6240,7 @@ static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struc
         return false;  // Q6_K has no fused HVX kernel
     }

-    return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
+    return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW;
 }

 static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams) {
@@ -5926,6 +6370,11 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
                     node.node->src[0], node.node->src[1], node.node,
                     (struct htp_mm_kernel_params *)node.kernel_params
                 );
+            } else if (node.opcode == HTP_OP_MUL || node.opcode == HTP_OP_ADD || node.opcode == HTP_OP_ADD_ID || node.opcode == HTP_OP_SUB || node.opcode == HTP_OP_DIV) {
+                const ggml_tensor * src1 = node.node->src[1];
+                GGML_ASSERT(ggml_hexagon_precompute_binary_params(sess,
+                    node.opcode, node.node->src[0], src1, node.node,
+                    (struct htp_binary_kernel_params *) node.kernel_params));
             } else if (node.opcode == HTP_OP_FLASH_ATTN_EXT) {
                 ggml_hexagon_precompute_flash_attn_params(sess,
                     node.node,
@@ -5954,6 +6403,21 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
                     node.node,
                     (struct htp_rope_kernel_params *)node.kernel_params
                 );
+            } else if (node.opcode == HTP_OP_SSM_CONV) {
+                ggml_hexagon_precompute_ssm_conv_params(sess,
+                    node.node->src[0], node.node->src[1], node.dst(),
+                    (struct htp_ssm_conv_kernel_params *)node.kernel_params
+                );
+            } else if (node.opcode == HTP_OP_SOFTMAX) {
+                ggml_hexagon_precompute_softmax_params(sess,
+                    node.node,
+                    (struct htp_softmax_kernel_params *)node.kernel_params
+                );
+            } else if (node.opcode == HTP_OP_GATED_DELTA_NET) {
+                ggml_hexagon_precompute_gated_delta_net_params(sess,
+                    node.node,
+                    (struct htp_gdn_kernel_params *)node.kernel_params
+                );
             }
             computed_nodes.push_back(std::move(node));
         }
@@ -6230,7 +6694,9 @@ static uint64_t ggml_hexagon_session_key(const ggml_hexagon_session * sess) {
 static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) {
     auto sess_src = static_cast<ggml_hexagon_session *>(backend_src->context);
     auto sess_dst = static_cast<ggml_hexagon_session *>(backend_dst->context);
-    auto sbuf_dst = (ggml_hexagon_shared_buffer *) dst->buffer->context;
+
+    sess_src->mmap_tensor(src);
+    auto sbuf_dst = sess_dst->mmap_tensor(dst);

     if (!sess_src->clone_buffer(sbuf_dst)) { return false; }

@@ -6277,7 +6743,9 @@ static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_
 static bool ggml_hexagon_cpy_tensor_async_virt(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) {
     auto sess_src = static_cast<ggml_hexagon_session *>(backend_src->context);
     auto sess_dst = static_cast<ggml_hexagon_session *>(backend_dst->context);
-    auto sbuf_src = (ggml_hexagon_shared_buffer *) src->buffer->context;
+
+    auto sbuf_src = sess_src->mmap_tensor(src);
+    sess_dst->mmap_tensor(dst);

     if (!sess_dst->clone_buffer(sbuf_src)) { return false; }

@@ -6905,7 +7373,7 @@ static const struct ggml_backend_device_i ggml_backend_hexagon_device_i = {
 ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) {
     GGML_LOG_INFO("ggml-hex: Hexagon backend (experimental) : allocating new registry : ndev %zu\n", opt_ndev);

-    GGML_LOG_INFO("ggml-hex: Hexagon Arch version v%d\n", opt_arch);
+    GGML_LOG_INFO("ggml-hex: Hexagon Arch version v%d, DMA64 %s\n", opt_arch, opt_dma64 ? "enabled" : "disabled");

     // Create devices
     for (size_t i = 0; i < opt_ndev; i++) {
@@ -7270,6 +7738,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
     const char * str_mbuf     = getenv("GGML_HEXAGON_MBUF");
     const char * str_optrace  = getenv("GGML_HEXAGON_OPTRACE");
     const char * str_hostbuf  = getenv("GGML_HEXAGON_HOSTBUF");
+    const char * str_dma64    = getenv("GGML_HEXAGON_DMA64");

     // Init Arch first since it affects other defaults
     if (!str_arch) {
@@ -7297,6 +7766,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {

     // Update vmem default
     opt_vmem = opt_arch >= 75 ? HTP_OP_MAX_VMEM_DEFAULT : 3000 * MiB;
+    opt_dma64 = opt_arch > 79 && (!str_dma64 || atoi(str_dma64) != 0);

     auto RE_ICASE = std::regex_constants::icase;

diff --git a/ggml/src/ggml-hexagon/htp-opnode.h b/ggml/src/ggml-hexagon/htp-opnode.h
index ef7b5184f..0716a8d21 100644
--- a/ggml/src/ggml-hexagon/htp-opnode.h
+++ b/ggml/src/ggml-hexagon/htp-opnode.h
@@ -14,7 +14,11 @@
 #include "htp/matmul-ops.h"
 #include "htp/flash-attn-ops.h"
 #include "htp/unary-ops.h"
+#include "htp/binary-ops.h"
 #include "htp/allreduce-ops.h"
+#include "htp/ssm-conv.h"
+#include "htp/gated-delta-net-ops.h"
+#include "htp/softmax-ops.h"

 struct htp_opnode {
     ggml_tensor * node   { nullptr };
@@ -325,10 +329,6 @@ struct htp_opformat {
             } else if (type == HTP_MM_KERNEL_HVX_F16_F16_VTCM || type == HTP_MM_KERNEL_HVX_F32_F32_VTCM ||
                        type == HTP_MM_KERNEL_HVX_QUANT_ROW    || type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
                 path = "hvx-tiled";
-            } else if (type == HTP_MM_KERNEL_HVX_F16_F16_DDR  || type == HTP_MM_KERNEL_HVX_F16_F32_DDR ||
-                       type == HTP_MM_KERNEL_HVX_F32_F32_DDR  || type == HTP_MM_KERNEL_HVX_F32_F16_DDR ||
-                       type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) {
-                path = "hvx-flat";
             }
             snprintf(str, max_size, "%s vtcm %d", path, (int) kparams->vtcm_size);
         } else if (node.opcode == HTP_OP_FLASH_ATTN_EXT) {
@@ -350,6 +350,21 @@ struct htp_opformat {
             snprintf(str, max_size, "seq 0x%x", (uint32_t) node.node->op_params[0]);
         } else if (node.opcode == HTP_OP_ALLREDUCE && node.node) {
             snprintf(str, max_size, "seq 0x%x -> 0x%x", (uint32_t) node.node->op_params[0], (uint32_t) node.node->op_params[1]);
+        } else if (node.opcode == HTP_OP_SSM_CONV) {
+            const auto * kparams = (const struct htp_ssm_conv_kernel_params *) node.kernel_params;
+            snprintf(str, max_size, "%s vtcm %d", kparams->n_t == 1 ? "decode" : "prefill", (int) kparams->vtcm_size);
+        } else if (node.opcode == HTP_OP_SOFTMAX) {
+            const auto * kparams = (const struct htp_softmax_kernel_params *) node.kernel_params;
+            snprintf(str, max_size, "k%d nth %d vtcm %d", (int) kparams->kernel_id, (int) kparams->n_threads, (int) kparams->vtcm_size);
+        } else if (node.opcode == HTP_OP_GATED_DELTA_NET) {
+            const auto * kparams = (const struct htp_gdn_kernel_params *) node.kernel_params;
+            snprintf(str, max_size, "%s vtcm %u",
+                     kparams->kda ? "kda" : "scalar",
+                     (unsigned int) (kparams->vtcm_size ? kparams->vtcm_size : kparams->vtcm_per_thread * kparams->n_threads));
+        } else if (node.opcode == HTP_OP_MUL || node.opcode == HTP_OP_ADD || node.opcode == HTP_OP_ADD_ID ||
+                   node.opcode == HTP_OP_SUB || node.opcode == HTP_OP_DIV) {
+            const auto * kparams = (const struct htp_binary_kernel_params *) node.kernel_params;
+            snprintf(str, max_size, "vtcm %u", (unsigned int) kparams->vtcm_size);
         } else {
             snprintf(str, max_size, "----");
         }
diff --git a/ggml/src/ggml-hexagon/htp/act-ops.c b/ggml/src/ggml-hexagon/htp/act-ops.c
index 5911c0890..d59ac0770 100644
--- a/ggml/src/ggml-hexagon/htp/act-ops.c
+++ b/ggml/src/ggml-hexagon/htp/act-ops.c
@@ -7,7 +7,7 @@
 #include <math.h>
 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"

 #define GGML_COMMON_DECL_C
@@ -53,13 +53,24 @@
     const uint32_t nb2 = dst->nb[2];                     \
     const uint32_t nb3 = dst->nb[3];

+struct htp_act_context;
+
+typedef void (*glu_compute_fn_t)(const float * restrict src0,
+                                 const float * restrict src1,
+                                 float * restrict dst,
+                                 const uint32_t num_rows,
+                                 const struct htp_act_context * actx);
+
 struct htp_act_context {
     struct htp_ops_context * octx;

+    glu_compute_fn_t         compute;
+    const char *             op_str;
+
     // Precomputed values
-    const uint8_t *          data_src0;
-    const uint8_t *          data_src1;
-    uint8_t *                data_dst;
+    dma_addr_t               data_src0;
+    dma_addr_t               data_src1;
+    dma_addr_t               data_dst;

     size_t                   src0_row_size;
     size_t                   src1_row_size;
@@ -134,10 +145,10 @@ static inline void htp_act_vtcm_layout_build(struct htp_act_vtcm_layout * L,

 // swiglu(x) = x1 * sigmoid(x0)
 static void swiglu_f32(const float * restrict src0,
-                       const float * restrict src1,
-                       float * restrict dst,
-                       const uint32_t num_rows,
-                       const struct htp_act_context * actx) {
+                                                 const float * restrict src1,
+                                                 float * restrict dst,
+                                                 const uint32_t num_rows,
+                                                 const struct htp_act_context * actx) {
     htp_glu_op_preamble;

     for (uint32_t ib = 0; ib < num_rows; ib++) {
@@ -152,10 +163,10 @@ static void swiglu_f32(const float * restrict src0,

 // out = x * sigmoid(alpha * x) * (clamp(y, -limit, limit) + 1.f)
 static void swiglu_oai_f32(const float * restrict src0,
-                           const float * restrict src1,
-                           float * restrict dst,
-                           const uint32_t num_rows,
-                           const struct htp_act_context * actx) {
+                                                     const float * restrict src1,
+                                                     float * restrict dst,
+                                                     const uint32_t num_rows,
+                                                     const struct htp_act_context * actx) {
     htp_glu_op_preamble;
     const float alpha = ((const float *) (actx->octx->op_params))[2];
     const float limit = ((const float *) (actx->octx->op_params))[3];
@@ -181,10 +192,10 @@ static void swiglu_oai_f32(const float * restrict src0,
 }

 static void swiglu_clamp_f32(const float * restrict src0,
-                             const float * restrict src1,
-                             float * restrict dst,
-                             const uint32_t                 num_rows,
-                             const struct htp_act_context * actx) {
+                                                       const float * restrict src1,
+                                                       float * restrict dst,
+                                                       const uint32_t                 num_rows,
+                                                       const struct htp_act_context * actx) {
     htp_glu_op_preamble;
     const float limit = ((const float *) (actx->octx->op_params))[3];

@@ -353,10 +364,10 @@ static inline void hvx_geglu_quick_f32_aa(uint8_t * restrict dst, const uint8_t

 // geglu(x, g) = gelu(x) * g
 static void geglu_f32(const float * restrict src0,
-                      const float * restrict src1,
-                      float * restrict dst,
-                      const uint32_t num_rows,
-                      const struct htp_act_context * actx) {
+                                                const float * restrict src1,
+                                                float * restrict dst,
+                                                const uint32_t num_rows,
+                                                const struct htp_act_context * actx) {
     htp_glu_op_preamble;

     for (uint32_t ib = 0; ib < num_rows; ib++) {
@@ -385,111 +396,100 @@ static void geglu_quick_f32(const float * restrict src0,
     }
 }

-#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR)                                                                   \
-    static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) {                           \
-        struct htp_act_context * actx = (struct htp_act_context *) data;                                                 \
-        htp_act_preamble;                                                                                                \
-                                                                                                                         \
-        struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL;                            \
-                                                                                                                         \
-        size_t src0_row_size = actx->src0_row_size;                                                                      \
-        size_t src1_row_size = actx->src1_row_size;                                                                      \
-        size_t dst_row_size  = actx->dst_row_size;                                                                       \
-                                                                                                                         \
-        size_t src0_row_stride = actx->src0_row_stride;                                                                  \
-        size_t src1_row_stride = actx->src1_row_stride;                                                                  \
-                                                                                                                         \
-        const uint32_t src0_nrows            = actx->src0_nrows;                                                         \
-        const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread;                                              \
-                                                                                                                         \
-        const uint32_t src0_start_row = actx->row_start + src0_nrows_per_thread * ith;                                   \
-        const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, actx->row_start + src0_nrows);       \
-                                                                                                                         \
-        /* no work for this thread */                                                                                    \
-        if (src0_start_row >= src0_end_row) {                                                                            \
-            return;                                                                                                      \
-        }                                                                                                                \
-                                                                                                                         \
-        const uint8_t * restrict data_src0 = actx->data_src0;                                                            \
-        const uint8_t * restrict data_src1 = actx->data_src1;                                                            \
-        uint8_t * restrict data_dst        = actx->data_dst;                                                             \
-                                                                                                                         \
-        const size_t src0_row_size_aligned = actx->src0_row_size_aligned;                                                \
-        const size_t src1_row_size_aligned = actx->src1_row_size_aligned;                                                \
-        const size_t dst_row_size_aligned  = actx->dst_row_size_aligned;                                                 \
-                                                                                                                         \
-        uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread);                   \
-        uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread);                   \
-        uint8_t * restrict dst_spad_data  = actx->vtcm_dst  + (ith * actx->vtcm_dst_size_per_thread);                    \
-                                                                                                                         \
-        size_t src0_spad_half_size = actx->src0_spad_half_size;                                                          \
-        size_t src1_spad_half_size = actx->src1_spad_half_size;                                                          \
-        size_t dst_spad_half_size  = actx->dst_spad_half_size;                                                           \
-                                                                                                                         \
-        const int BLOCK = actx->block;                                                                                   \
-        if (BLOCK == 0) {                                                                                                \
-            FARF(ERROR,                                                                                                  \
-                 OP_STR                                                                                                  \
-                 " : current VTCM reservation %zu is too small for even 1 row per thread, needed at least %zu\n",        \
-                 actx->vtcm_src0_size_per_thread, src0_row_size_aligned);                                                \
-            return;                                                                                                      \
-        }                                                                                                                \
-                                                                                                                         \
-        dma_queue * dma_queue = actx->octx->ctx->dma[ith];                                                               \
-                                                                                                                         \
-        /* See discussion: https://github.com/ggml-org/llama.cpp/pull/18151#issuecomment-3678235379 */                   \
-        for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) {   \
-            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                   \
-                                                                                                                         \
-            /* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */                                     \
-            dma_queue_push_vtcm_to_ddr(dma_queue,                                                                        \
-                                       dma_make_ptr(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)),          \
-                                       dst_row_size, dst_row_size_aligned, 0);                                           \
-                                                                                                                         \
-            dma_queue_push(                                                                                              \
-                dma_queue,                                                                                               \
-                dma_make_ptr(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)),     \
-                src0_row_size_aligned, src0_row_stride, src0_row_size, block_size);                                      \
-            dma_queue_push(                                                                                              \
-                dma_queue,                                                                                               \
-                dma_make_ptr(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)),     \
-                src1_row_size_aligned, src1_row_stride, src1_row_size, block_size);                                      \
-        }                                                                                                                \
-                                                                                                                         \
-        for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) {                                             \
-            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                   \
-                                                                                                                         \
-            float * dst_spad  = (float *) dma_queue_pop(dma_queue).src;                                                  \
-            float * src0_spad = (float *) dma_queue_pop(dma_queue).dst;                                                  \
-            float * src1_spad = (float *) dma_queue_pop(dma_queue).dst;                                                  \
-                                                                                                                         \
-            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                       \
-            CORE_EXPR;                                                                                                   \
-            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                        \
-                                                                                                                         \
-            dma_queue_push_vtcm_to_ddr(dma_queue, dma_make_ptr(data_dst + (ir * dst_row_size), dst_spad),                \
-                                       dst_row_size, dst_row_size_aligned, block_size);                                  \
-                                                                                                                         \
-            /* prefetch N+2 loop iteration if any */                                                                     \
-            const uint32_t pref_block = (ir + BLOCK * 2);                                                                \
-            if (pref_block < src0_end_row) {                                                                             \
-                const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block);                                  \
-                dma_queue_push(dma_queue, dma_make_ptr(src0_spad, data_src0 + (pref_block * src0_row_stride)),           \
-                               src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size);                  \
-                dma_queue_push(dma_queue, dma_make_ptr(src1_spad, data_src1 + (pref_block * src1_row_stride)),           \
-                               src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size);                  \
-            }                                                                                                            \
-        }                                                                                                                \
-                                                                                                                         \
-        dma_queue_flush(dma_queue);                                                                                      \
-                                                                                                                         \
+static void glu_f32_per_thread(unsigned int nth, unsigned int ith, void * data) {
+    struct htp_act_context * actx = (struct htp_act_context *) data;
+    htp_act_preamble;
+
+    struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL;
+
+    size_t src0_row_size = actx->src0_row_size;
+    size_t src1_row_size = actx->src1_row_size;
+    size_t dst_row_size  = actx->dst_row_size;
+
+    size_t src0_row_stride = actx->src0_row_stride;
+    size_t src1_row_stride = actx->src1_row_stride;
+
+    const uint32_t src0_nrows            = actx->src0_nrows;
+    const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread;
+
+    const uint32_t src0_start_row = actx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, actx->row_start + src0_nrows);
+
+    /* no work for this thread */
+    if (src0_start_row >= src0_end_row) {
+        return;
     }

-DEFINE_GLU_PER_THREAD(swiglu, "swiglu-f32", swiglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
-DEFINE_GLU_PER_THREAD(swiglu_oai, "swiglu-oai-f32", swiglu_oai_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
-DEFINE_GLU_PER_THREAD(swiglu_clamp, "swiglu-clamp-f32", swiglu_clamp_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
-DEFINE_GLU_PER_THREAD(geglu, "geglu-f32", geglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
-DEFINE_GLU_PER_THREAD(geglu_quick, "geglu-quick-f32", geglu_quick_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
+    const dma_addr_t data_src0 = actx->data_src0;
+    const dma_addr_t data_src1 = actx->data_src1;
+    const dma_addr_t data_dst  = actx->data_dst;
+
+    const size_t src0_row_size_aligned = actx->src0_row_size_aligned;
+    const size_t src1_row_size_aligned = actx->src1_row_size_aligned;
+    const size_t dst_row_size_aligned  = actx->dst_row_size_aligned;
+
+    uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread);
+    uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread);
+    uint8_t * restrict dst_spad_data  = actx->vtcm_dst  + (ith * actx->vtcm_dst_size_per_thread);
+
+    size_t src0_spad_half_size = actx->src0_spad_half_size;
+    size_t src1_spad_half_size = actx->src1_spad_half_size;
+    size_t dst_spad_half_size  = actx->dst_spad_half_size;
+
+    const int BLOCK = actx->block;
+    if (BLOCK == 0) {
+        FARF(ERROR, "%s : VTCM reservation %zu is too small, needed %zu\n",
+             actx->op_str, actx->vtcm_src0_size_per_thread, src0_row_size_aligned);
+        return;
+    }
+
+    dma_queue * dma_q = actx->octx->ctx->dma[ith];
+    glu_compute_fn_t compute = actx->compute;
+
+    for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) {
+        const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);
+
+        /* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */
+        dma_queue_push(dma_q,
+                       dma_make_data(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)),
+                       dst_row_size, dst_row_size_aligned, dst_row_size, 0);
+
+        dma_queue_push(dma_q,
+            dma_make_data(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)),
+            src0_row_size_aligned, src0_row_stride, src0_row_size, block_size);
+
+        dma_queue_push(dma_q,
+            dma_make_data(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)),
+            src1_row_size_aligned, src1_row_stride, src1_row_size, block_size);
+    }
+
+    for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) {
+        const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);
+
+        float * dst_spad  = (float *) dma_queue_pop(dma_q).src;
+        float * src0_spad = (float *) dma_queue_pop(dma_q).dst;
+        float * src1_spad = (float *) dma_queue_pop(dma_q).dst;
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+        compute(src0_spad, src1_spad, dst_spad, block_size, actx);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+
+        dma_queue_push(dma_q, dma_make_data(data_dst + (ir * dst_row_size), dst_spad),
+                       dst_row_size, dst_row_size_aligned, dst_row_size, block_size);
+
+        /* prefetch N+2 loop iteration if any */
+        const uint32_t pref_block = (ir + BLOCK * 2);
+        if (pref_block < src0_end_row) {
+            const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block);
+            dma_queue_push(dma_q, dma_make_data(src0_spad, data_src0 + (pref_block * src0_row_stride)),
+                           src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size);
+            dma_queue_push(dma_q, dma_make_data(src1_spad, data_src1 + (pref_block * src1_row_stride)),
+                           src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size);
+        }
+    }
+
+    dma_queue_flush(dma_q);
+}

 static int execute_op_activations_f32(struct htp_ops_context * octx) {
     const struct htp_tensor * src0 = octx->src[0];
@@ -501,33 +501,33 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    worker_callback_t act_op_func;
-    const char *      op_type = NULL;
+    glu_compute_fn_t compute_fn = NULL;
+    const char *      op_type    = NULL;

     switch (octx->op) {
         case HTP_OP_GLU_SWIGLU:
-            act_op_func = (worker_callback_t)glu_swiglu_f32_per_thread;
-            op_type     = "swiglu-f32";
+            compute_fn = swiglu_f32;
+            op_type    = "swiglu-f32";
             break;

         case HTP_OP_GLU_SWIGLU_OAI:
-            act_op_func = (worker_callback_t)glu_swiglu_oai_f32_per_thread;
-            op_type     = "swiglu-oai-f32";
+            compute_fn = swiglu_oai_f32;
+            op_type    = "swiglu-oai-f32";
             break;

         case HTP_OP_GLU_SWIGLU_CLAMP:
-            act_op_func = (worker_callback_t) glu_swiglu_clamp_f32_per_thread;
-            op_type     = "swiglu-clamp-f32";
+            compute_fn = swiglu_clamp_f32;
+            op_type    = "swiglu-clamp-f32";
             break;

         case HTP_OP_GLU_GEGLU:
-            act_op_func = (worker_callback_t)glu_geglu_f32_per_thread;
-            op_type     = "geglu-f32";
+            compute_fn = geglu_f32;
+            op_type    = "geglu-f32";
             break;

         case HTP_OP_GLU_GEGLU_QUICK:
-            act_op_func = (worker_callback_t)glu_geglu_quick_f32_per_thread;
-            op_type     = "geglu-quick-f32";
+            compute_fn = geglu_quick_f32;
+            op_type    = "geglu-quick-f32";
             break;
         default:
             FARF(ERROR, "Unsupported activations Op %u\n", octx->op);
@@ -588,13 +588,11 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
              L.src0_bytes_per_thread * n_threads, L.src1_bytes_per_thread * n_threads, L.dst_bytes_per_thread * n_threads);
     }

-    if ((octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        return HTP_STATUS_OK;
-    }
-
     // Prepare context
     struct htp_act_context actx;
-    actx.octx = octx;
+    actx.octx    = octx;
+    actx.compute = compute_fn;
+    actx.op_str  = op_type;

     actx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);

@@ -628,9 +626,9 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {

     actx.nc = dst->ne[0];

-    // Pointers and GLU logic
-    const uint8_t * data_src0 = (const uint8_t *) src0->data;
-    const uint8_t * data_src1 = src1 ? (const uint8_t *) src1->data : NULL;
+    // Addresses and GLU logic
+    dma_addr_t data_src0 = src0->data;
+    dma_addr_t data_src1 = src1 ? src1->data : 0;

     if (!src1 && (octx->op == HTP_OP_GLU_SWIGLU ||
                   octx->op == HTP_OP_GLU_SWIGLU_OAI ||
@@ -651,9 +649,9 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {

     actx.data_src0 = data_src0;
     actx.data_src1 = data_src1;
-    actx.data_dst  = (uint8_t *) dst->data;
+    actx.data_dst  = dst->data;

-    work_queue_run(octx->ctx->work_queue, act_op_func, &actx, n_threads);
+    work_queue_run(octx->ctx->work_queue, (worker_callback_t)glu_f32_per_thread, &actx, n_threads);
     return HTP_STATUS_OK;
 }

diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.c b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
index d6e7f0d10..7b577befb 100644
--- a/ggml/src/ggml-hexagon/htp/allreduce-ops.c
+++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
@@ -14,7 +14,7 @@
 #include "htp-ops.h"
 #include "hvx-utils.h"
 #include "htp-tensor.h"
-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hex-profile.h"
 #include "allreduce-ops.h"
 #include "htp-fence.h"
@@ -38,97 +38,97 @@ struct htp_allreduce_context {
     uint8_t * res_spad_base;
 };

-#define DEFINE_ALLREDUCE_THREAD_DMA_1D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD)                                             \
-static void allreduce_thread_dma_1d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                       \
-    struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data;                                      \
-    struct htp_ops_context * octx = actx->octx;                                                                       \
-                                                                                                                      \
-    const uint32_t n_ranks     = actx->n_ranks;                                                                       \
-    const uint32_t n_dsts      = actx->n_dsts;                                                                        \
-    const uint32_t block_elems = actx->block_elems;                                                                   \
-                                                                                                                      \
-    const uint32_t dr  = actx->elems_per_thread;                                                                      \
-    const uint32_t ir0 = actx->rank_elem_start + dr * ith;                                                            \
-    const uint32_t ir1 = MIN(ir0 + dr, actx->rank_elem_start + actx->rank_nelem);                                     \
-    if (ir0 >= ir1) return;                                                                                           \
-                                                                                                                      \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                            \
-    dma_queue * q = octx->ctx->dma[ith];                                                                              \
-                                                                                                                      \
-    uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS];                                                                 \
-    for (uint32_t s = 0; s < n_ranks; s++) {                                                                          \
-        src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread);                               \
-    }                                                                                                                 \
-    uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread);                               \
-    uint8_t * res_spad_base = HAS_ADD ? (actx->res_spad_base + (ith * actx->vtcm_size_per_thread)) : NULL;            \
-                                                                                                                      \
-    const size_t spad_half = actx->vtcm_size_per_thread / 2;                                                          \
-    uint32_t ir_prefetch = ir0;                                                                                       \
-    int spad_idx = 0;                                                                                                 \
-                                                                                                                      \
-    for (int k = 0; k < 2 && ir_prefetch < ir1; k++) {                                                                \
-        uint32_t cur_elems = MIN(block_elems, ir1 - ir_prefetch);                                                     \
-        size_t   cur_bytes = cur_elems * sizeof(TYPE);                                                                \
-        uint8_t * d_spad = dst_spad_base + spad_idx * spad_half;                                                      \
-        for (uint32_t d = 0; d < n_dsts; d++) {                                                                       \
-            uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + ir_prefetch * sizeof(TYPE);                           \
-            dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 0);                       \
-        }                                                                                                             \
-        for (uint32_t s = 0; s < n_ranks; s++) {                                                                      \
-            uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half;                                               \
-            const uint8_t * s_ddr = (const uint8_t *) octx->src[s]->data + ir_prefetch * sizeof(TYPE);                \
-            dma_queue_push(q, dma_make_ptr(s_spad, s_ddr), cur_bytes, cur_bytes, cur_bytes, 1);                       \
-        }                                                                                                             \
-        if (HAS_ADD) {                                                                                                \
-            uint8_t * r_spad = res_spad_base + spad_idx * spad_half;                                                  \
-            const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE);      \
-            dma_queue_push(q, dma_make_ptr(r_spad, r_ddr), cur_bytes, cur_bytes, cur_bytes, 1);                       \
-        }                                                                                                             \
-        ir_prefetch += cur_elems;                                                                                     \
-        spad_idx ^= 1;                                                                                                \
-    }                                                                                                                 \
-                                                                                                                      \
-    for (uint32_t ir = ir0; ir < ir1; ) {                                                                             \
-        uint32_t cur_elems = MIN(block_elems, ir1 - ir);                                                              \
-        size_t   cur_bytes = cur_elems * sizeof(TYPE);                                                                \
-        uint8_t * d_spad = NULL;                                                                                      \
-        for (uint32_t d = 0; d < n_dsts; d++) {                                                                       \
-            d_spad = (uint8_t *) dma_queue_pop(q).src;                                                                \
-        }                                                                                                             \
-        uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS];                                                                    \
-        for (uint32_t s = 0; s < n_ranks; s++) {                                                                      \
-            s_spad[s] = (uint8_t *) dma_queue_pop(q).dst;                                                             \
-        }                                                                                                             \
-        uint8_t * r_spad = HAS_ADD ? (uint8_t *) dma_queue_pop(q).dst : NULL;                                         \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);                                             \
-        HVX_ADD_FN(d_spad, s_spad[0], s_spad[1], cur_elems);                                                          \
-        for (uint32_t s = 2; s < n_ranks; s++) {                                                                      \
-            HVX_ADD_FN(d_spad, d_spad, s_spad[s], cur_elems);                                                         \
-        }                                                                                                             \
-        if (HAS_ADD) {                                                                                                \
-            HVX_ADD_FN(d_spad, d_spad, r_spad, cur_elems);                                                            \
-        }                                                                                                             \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);                                              \
-        for (uint32_t d = 0; d < n_dsts; d++) {                                                                       \
-            uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + ir * sizeof(TYPE);                                    \
-            dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 1);                       \
-        }                                                                                                             \
-        if (ir_prefetch < ir1) {                                                                                      \
-            uint32_t next_elems = MIN(block_elems, ir1 - ir_prefetch);                                                \
-            size_t   next_bytes = next_elems * sizeof(TYPE);                                                          \
-            for (uint32_t s = 0; s < n_ranks; s++) {                                                                  \
-                const uint8_t * s_next = (const uint8_t *) octx->src[s]->data + ir_prefetch * sizeof(TYPE);           \
-                dma_queue_push(q, dma_make_ptr(s_spad[s], s_next), next_bytes, next_bytes, next_bytes, 1);            \
-            }                                                                                                         \
-            if (HAS_ADD) {                                                                                            \
-                const uint8_t * r_next = (const uint8_t *) octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \
-                dma_queue_push(q, dma_make_ptr(r_spad, r_next), next_bytes, next_bytes, next_bytes, 1);               \
-            }                                                                                                         \
-            ir_prefetch += next_elems;                                                                                \
-        }                                                                                                             \
-        ir += cur_elems;                                                                                              \
-    }                                                                                                                 \
-    dma_queue_flush(q);                                                                                               \
+#define DEFINE_ALLREDUCE_THREAD_DMA_1D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD)                                       \
+static void allreduce_thread_dma_1d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                 \
+    struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data;                                \
+    struct htp_ops_context * octx = actx->octx;                                                                 \
+                                                                                                                \
+    const uint32_t n_ranks     = actx->n_ranks;                                                                 \
+    const uint32_t n_dsts      = actx->n_dsts;                                                                  \
+    const uint32_t block_elems = actx->block_elems;                                                             \
+                                                                                                                \
+    const uint32_t dr  = actx->elems_per_thread;                                                                \
+    const uint32_t ir0 = actx->rank_elem_start + dr * ith;                                                      \
+    const uint32_t ir1 = MIN(ir0 + dr, actx->rank_elem_start + actx->rank_nelem);                               \
+    if (ir0 >= ir1) return;                                                                                     \
+                                                                                                                \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                      \
+    dma_queue * dma_q = octx->ctx->dma[ith];                                                                    \
+                                                                                                                \
+    uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS];                                                           \
+    for (uint32_t s = 0; s < n_ranks; s++) {                                                                    \
+        src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread);                         \
+    }                                                                                                           \
+    uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread);                         \
+    uint8_t * res_spad_base = HAS_ADD ? (actx->res_spad_base + (ith * actx->vtcm_size_per_thread)) : NULL;      \
+                                                                                                                \
+    const size_t spad_half = actx->vtcm_size_per_thread / 2;                                                    \
+    uint32_t ir_prefetch = ir0;                                                                                 \
+    int spad_idx = 0;                                                                                           \
+                                                                                                                \
+    for (int k = 0; k < 2 && ir_prefetch < ir1; k++) {                                                          \
+        uint32_t cur_elems = MIN(block_elems, ir1 - ir_prefetch);                                               \
+        size_t   cur_bytes = cur_elems * sizeof(TYPE);                                                          \
+        uint8_t * d_spad = dst_spad_base + spad_idx * spad_half;                                                \
+        for (uint32_t d = 0; d < n_dsts; d++) {                                                                 \
+            dma_addr_t d_ddr = octx->dsts[d]->data + ir_prefetch * sizeof(TYPE);                                \
+            dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 0);            \
+        }                                                                                                       \
+        for (uint32_t s = 0; s < n_ranks; s++) {                                                                \
+            uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half;                                         \
+            const dma_addr_t s_ddr = octx->src[s]->data + ir_prefetch * sizeof(TYPE);                           \
+            dma_queue_push(dma_q, dma_make_data(s_spad, s_ddr), cur_bytes, cur_bytes, cur_bytes, 1);            \
+        }                                                                                                       \
+        if (HAS_ADD) {                                                                                          \
+            uint8_t * r_spad = res_spad_base + spad_idx * spad_half;                                            \
+            const dma_addr_t r_ddr = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE);                 \
+            dma_queue_push(dma_q, dma_make_data(r_spad, r_ddr), cur_bytes, cur_bytes, cur_bytes, 1);            \
+        }                                                                                                       \
+        ir_prefetch += cur_elems;                                                                               \
+        spad_idx ^= 1;                                                                                          \
+    }                                                                                                           \
+                                                                                                                \
+    for (uint32_t ir = ir0; ir < ir1; ) {                                                                       \
+        uint32_t cur_elems = MIN(block_elems, ir1 - ir);                                                        \
+        size_t   cur_bytes = cur_elems * sizeof(TYPE);                                                          \
+        uint8_t * d_spad = NULL;                                                                                \
+        for (uint32_t d = 0; d < n_dsts; d++) {                                                                 \
+            d_spad = (uint8_t *) dma_queue_pop(dma_q).src;                                                      \
+        }                                                                                                       \
+        uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS];                                                              \
+        for (uint32_t s = 0; s < n_ranks; s++) {                                                                \
+            s_spad[s] = (uint8_t *) dma_queue_pop(dma_q).dst;                                                   \
+        }                                                                                                       \
+        uint8_t * r_spad = HAS_ADD ? (uint8_t *) dma_queue_pop(dma_q).dst : NULL;                               \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);                                       \
+        HVX_ADD_FN(d_spad, s_spad[0], s_spad[1], cur_elems);                                                    \
+        for (uint32_t s = 2; s < n_ranks; s++) {                                                                \
+            HVX_ADD_FN(d_spad, d_spad, s_spad[s], cur_elems);                                                   \
+        }                                                                                                       \
+        if (HAS_ADD) {                                                                                          \
+            HVX_ADD_FN(d_spad, d_spad, r_spad, cur_elems);                                                      \
+        }                                                                                                       \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);                                        \
+        for (uint32_t d = 0; d < n_dsts; d++) {                                                                 \
+            dma_addr_t d_ddr = octx->dsts[d]->data + ir * sizeof(TYPE);                                         \
+            dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 1);            \
+        }                                                                                                       \
+        if (ir_prefetch < ir1) {                                                                                \
+            uint32_t next_elems = MIN(block_elems, ir1 - ir_prefetch);                                          \
+            size_t   next_bytes = next_elems * sizeof(TYPE);                                                    \
+            for (uint32_t s = 0; s < n_ranks; s++) {                                                            \
+                const dma_addr_t s_next = octx->src[s]->data + ir_prefetch * sizeof(TYPE);                      \
+                dma_queue_push(dma_q, dma_make_data(s_spad[s], s_next), next_bytes, next_bytes, next_bytes, 1); \
+            }                                                                                                   \
+            if (HAS_ADD) {                                                                                      \
+                const dma_addr_t r_next = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE);            \
+                dma_queue_push(dma_q, dma_make_data(r_spad, r_next), next_bytes, next_bytes, next_bytes, 1);    \
+            }                                                                                                   \
+            ir_prefetch += next_elems;                                                                          \
+        }                                                                                                       \
+        ir += cur_elems;                                                                                        \
+    }                                                                                                           \
+    dma_queue_flush(dma_q);                                                                                     \
 }

 DEFINE_ALLREDUCE_THREAD_DMA_1D(f16,     __fp16, hvx_add_f16_aaa, 0)
@@ -154,7 +154,7 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
     if (r0 >= r1) return;                                                                                                                         \
                                                                                                                                                   \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                                        \
-    dma_queue * q = octx->ctx->dma[ith];                                                                                                          \
+    dma_queue * dma_q = octx->ctx->dma[ith];                                                                                                      \
                                                                                                                                                   \
     uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS];                                                                                             \
     for (uint32_t s = 0; s < n_ranks; s++) {                                                                                                      \
@@ -171,18 +171,18 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
         uint32_t cur_rows = MIN(block_rows, r1 - r_prefetch);                                                                                     \
         uint8_t * d_spad = dst_spad_base + spad_idx * spad_half;                                                                                  \
         for (uint32_t d = 0; d < n_dsts; d++) {                                                                                                   \
-            uint8_t * d_ddr  = (uint8_t *) octx->dsts[d]->data + r_prefetch * octx->dsts[d]->nb[1];                                               \
-            dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, 0);                                 \
+            dma_addr_t d_ddr = octx->dsts[d]->data + r_prefetch * octx->dsts[d]->nb[1];                                                           \
+            dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, 0);                            \
         }                                                                                                                                         \
         for (uint32_t s = 0; s < n_ranks; s++) {                                                                                                  \
             uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half;                                                                           \
-            const uint8_t * s_ddr = (const uint8_t *) octx->src[s]->data + r_prefetch * octx->src[s]->nb[1];                                      \
-            dma_queue_push(q, dma_make_ptr(s_spad, s_ddr), row_size_aligned, octx->src[s]->nb[1], row_bytes, cur_rows);                           \
+            const dma_addr_t s_ddr = octx->src[s]->data + r_prefetch * octx->src[s]->nb[1];                                                       \
+            dma_queue_push(dma_q, dma_make_data(s_spad, s_ddr), row_size_aligned, octx->src[s]->nb[1], row_bytes, cur_rows);                      \
         }                                                                                                                                         \
         if (HAS_ADD && !IS_ROW_BCAST) {                                                                                                           \
             uint8_t * r_spad = res_spad_base + spad_idx * spad_half;                                                                              \
-            const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1];                  \
-            dma_queue_push(q, dma_make_ptr(r_spad, r_ddr), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, cur_rows);                 \
+            const dma_addr_t r_ddr = octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1];                                   \
+            dma_queue_push(dma_q, dma_make_data(r_spad, r_ddr), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, cur_rows);            \
         }                                                                                                                                         \
         r_prefetch += cur_rows;                                                                                                                   \
         spad_idx ^= 1;                                                                                                                            \
@@ -192,13 +192,13 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
         uint32_t cur_rows = MIN(block_rows, r1 - r);                                                                                              \
         uint8_t * d_spad = NULL;                                                                                                                  \
         for (uint32_t d = 0; d < n_dsts; d++) {                                                                                                   \
-            d_spad = (uint8_t *) dma_queue_pop(q).src;                                                                                            \
+            d_spad = (uint8_t *) dma_queue_pop(dma_q).src;                                                                                        \
         }                                                                                                                                         \
         uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS];                                                                                                \
         for (uint32_t s = 0; s < n_ranks; s++) {                                                                                                  \
-            s_spad[s] = (uint8_t *) dma_queue_pop(q).dst;                                                                                         \
+            s_spad[s] = (uint8_t *) dma_queue_pop(dma_q).dst;                                                                                     \
         }                                                                                                                                         \
-        uint8_t * r_spad = (HAS_ADD && !IS_ROW_BCAST) ? (uint8_t *) dma_queue_pop(q).dst : NULL;                                                  \
+        uint8_t * r_spad = (HAS_ADD && !IS_ROW_BCAST) ? (uint8_t *) dma_queue_pop(dma_q).dst : NULL;                                              \
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r);                                                                          \
         for (uint32_t row = 0; row < cur_rows; row++) {                                                                                           \
             uint8_t * d_row = d_spad + row * row_size_aligned;                                                                                    \
@@ -216,24 +216,24 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
         }                                                                                                                                         \
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r);                                                                           \
         for (uint32_t d = 0; d < n_dsts; d++) {                                                                                                   \
-            uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + r * octx->dsts[d]->nb[1];                                                         \
-            dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, cur_rows);                          \
+            dma_addr_t d_ddr = octx->dsts[d]->data + r * octx->dsts[d]->nb[1];                                                                    \
+            dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, cur_rows);                     \
         }                                                                                                                                         \
         if (r_prefetch < r1) {                                                                                                                    \
             uint32_t next_rows = MIN(block_rows, r1 - r_prefetch);                                                                                \
             for (uint32_t s = 0; s < n_ranks; s++) {                                                                                              \
-                const uint8_t * s_next = (const uint8_t *) octx->src[s]->data + r_prefetch * octx->src[s]->nb[1];                                 \
-                dma_queue_push(q, dma_make_ptr(s_spad[s], s_next), row_size_aligned, octx->src[s]->nb[1], row_bytes, next_rows);                  \
+                const dma_addr_t s_next = octx->src[s]->data + r_prefetch * octx->src[s]->nb[1];                                                  \
+                dma_queue_push(dma_q, dma_make_data(s_spad[s], s_next), row_size_aligned, octx->src[s]->nb[1], row_bytes, next_rows);             \
             }                                                                                                                                     \
             if (HAS_ADD && !IS_ROW_BCAST) {                                                                                                       \
-                const uint8_t * r_next = (const uint8_t *) octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1];             \
-                dma_queue_push(q, dma_make_ptr(r_spad, r_next), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, next_rows);           \
+                const dma_addr_t r_next = octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1];                              \
+                dma_queue_push(dma_q, dma_make_data(r_spad, r_next), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, next_rows);      \
             }                                                                                                                                     \
             r_prefetch += next_rows;                                                                                                              \
         }                                                                                                                                         \
         r += cur_rows;                                                                                                                            \
     }                                                                                                                                             \
-    dma_queue_flush(q);                                                                                                                           \
+    dma_queue_flush(dma_q);                                                                                                                       \
 }

 DEFINE_ALLREDUCE_THREAD_DMA_2D(f16,           __fp16, hvx_add_f16_aaa, 0, 0)
@@ -406,11 +406,11 @@ int op_allreduce(struct htp_ops_context * octx) {
         }

         if (has_add && actx.is_row_bcast) {
-            const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data;
+            const dma_addr_t r_ddr = octx->src[2 * n_ranks]->data;
             const uint32_t row_bytes = actx.ne0 * (dst->type == HTP_TYPE_F16 ? sizeof(__fp16) : sizeof(float));
-            dma_queue * q = octx->ctx->dma[0];
-            dma_queue_push(q, dma_make_ptr(actx.res_spad_base, r_ddr), actx.row_size_aligned, 0, row_bytes, 1);
-            dma_queue_pop(q);
+            dma_queue * dma_q = octx->ctx->dma[0];
+            dma_queue_push(dma_q, dma_make_data(actx.res_spad_base, r_ddr), actx.row_size_aligned, 0, row_bytes, 1);
+            dma_queue_pop(dma_q);
         }

         work_queue_run(octx->ctx->work_queue, reduce_fun, &actx, n_threads);
diff --git a/ggml/src/ggml-hexagon/htp/argsort-ops.c b/ggml/src/ggml-hexagon/htp/argsort-ops.c
index 6ee614d3d..9e9e74651 100644
--- a/ggml/src/ggml-hexagon/htp/argsort-ops.c
+++ b/ggml/src/ggml-hexagon/htp/argsort-ops.c
@@ -9,7 +9,7 @@
 #include "ggml.h"

 #include "hvx-utils.h"
-#include "hex-dma.h"
+#include "dma-queue.h"

 #include "hex-common.h"
 #include "htp-ctx.h"
@@ -591,6 +591,10 @@ int op_argsort(struct htp_ops_context * octx) {
     const struct htp_tensor * src0 = octx->src[0];
     const struct htp_tensor * dst  = octx->dst;

+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     const uint32_t total_rows  = src0->ne[1] * src0->ne[2] * src0->ne[3];
     const size_t dst_row_size  = dst->ne[0]  * sizeof(int32_t);

diff --git a/ggml/src/ggml-hexagon/htp/binary-ops.c b/ggml/src/ggml-hexagon/htp/binary-ops.c
index bfa849e0e..155e85237 100644
--- a/ggml/src/ggml-hexagon/htp/binary-ops.c
+++ b/ggml/src/ggml-hexagon/htp/binary-ops.c
@@ -8,13 +8,14 @@
 #include <math.h>
 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "hex-common.h"
 #include "hex-profile.h"
+#include "binary-ops.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
@@ -26,6 +27,8 @@
 // Context for binary operations
 struct htp_binary_context {
     struct htp_ops_context * octx;
+    struct htp_binary_vtcm_layout vtcm_layout;
+    uint8_t * vtcm_base;

     struct fastdiv_values src0_dim1_div; // ne01
     struct fastdiv_values src0_dim2_div; // ne02
@@ -42,9 +45,12 @@ struct htp_binary_context {
     size_t   src0_row_size_aligned;
     size_t   src1_row_size_aligned;
     size_t   dst_row_size_aligned;
+    size_t   row_size_bytes;

     bool split_at_ne01;
     bool split_at_ne02;
+
+    void * compute;
 };

 #define htp_binary_preamble                        \
@@ -95,114 +101,302 @@ static inline uint32_t calc_block_size(struct htp_binary_context * bctx, uint32_
     return MIN(bctx->block_max, block_limit);
 }

-// Macro for scalar op switch
-#define COMPUTE_SCALAR_OP(DST, SRC, VAL, TYPE, N)                                               \
-    if(TYPE == HTP_TYPE_F32) {                                                                  \
-        switch (octx->op) {                                                                     \
-            case HTP_OP_ADD: hvx_add_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
-            case HTP_OP_SUB: hvx_sub_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
-            case HTP_OP_MUL: hvx_mul_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
-            case HTP_OP_DIV: hvx_mul_scalar_f32_aa(DST, SRC, 1.0f / (*(float *)VAL), N); break; \
-            default: break;                                                                     \
-        }                                                                                       \
-    }                                                                                           \
-    else {                                                                                      \
-        switch (octx->op) {                                                                     \
-            case HTP_OP_ADD: hvx_add_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
-            case HTP_OP_SUB: hvx_sub_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
-            case HTP_OP_MUL: hvx_mul_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
-            case HTP_OP_DIV: hvx_div_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
-            default: break;                                                                     \
-        }                                                                                       \
-    }
+// Out-of-line compute micro-kernels
+
+typedef void (*compute_scalar_dma_t)(
+    uint8_t * dst, const uint8_t * src0, const void * s1_table,
+    uint32_t cur_i11, uint32_t ne11, uint32_t n_rows,
+    size_t dst_stride, size_t src0_stride, uint32_t ne00);
+
+#define DEFINE_COMPUTE_SCALAR_DMA(NAME, TYPE, HVX_STMT)         \
+static void compute_scalar_dma_##NAME(                          \
+    uint8_t * dst, const uint8_t * src0, const void * s1_table, \
+    uint32_t cur_i11, uint32_t ne11, uint32_t n_rows,           \
+    size_t dst_stride, size_t src0_stride, uint32_t ne00) {     \
+    const TYPE * table = (const TYPE *) s1_table;               \
+    for (uint32_t r = 0; r < n_rows; r++) {                     \
+        uint8_t * r_dst = dst + r * dst_stride;                 \
+        const uint8_t * r_src0 = src0 + r * src0_stride;        \
+        TYPE val = table[cur_i11];                              \
+        HVX_STMT;                                               \
+        if (ne11 > 1 && ++cur_i11 == ne11) {                    \
+            cur_i11 = 0;                                        \
+        }                                                       \
+    }                                                           \
+}
+
+DEFINE_COMPUTE_SCALAR_DMA(add_f32, float,    hvx_add_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(add_f16, _Float16, hvx_add_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(sub_f32, float,    hvx_sub_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(sub_f16, _Float16, hvx_sub_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(mul_f32, float,    hvx_mul_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(mul_f16, _Float16, hvx_mul_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR_DMA(div_f32, float,    hvx_mul_scalar_f32_aa(r_dst, r_src0, 1.0f / (val), ne00))
+DEFINE_COMPUTE_SCALAR_DMA(div_f16, _Float16, hvx_div_scalar_f16_aa(r_dst, r_src0, val, ne00))
+
+typedef void (*compute_scalar_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_ptr, uint32_t s1_stride,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00);
+
+#define DEFINE_COMPUTE_SCALAR(NAME, TYPE, HVX_STMT)                                          \
+static void compute_scalar_##NAME(                                                           \
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_ptr, uint32_t s1_stride,       \
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00) {                 \
+    for (uint32_t r = 0; r < n_rows; r++) {                                                  \
+        uint8_t * r_dst = dst + r * dst_stride;                                              \
+        const uint8_t * r_src0 = src0 + r * src0_stride;                                     \
+        TYPE val = *(const TYPE *)(src1_ptr + r * s1_stride);                                \
+        HVX_STMT;                                                                            \
+    }                                                                                        \
+}
+
+DEFINE_COMPUTE_SCALAR(add_f32, float,    hvx_add_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(add_f16, _Float16, hvx_add_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(sub_f32, float,    hvx_sub_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(sub_f16, _Float16, hvx_sub_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(mul_f32, float,    hvx_mul_scalar_f32_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(mul_f16, _Float16, hvx_mul_scalar_f16_aa(r_dst, r_src0, val, ne00))
+DEFINE_COMPUTE_SCALAR(div_f32, float,    hvx_mul_scalar_f32_aa(r_dst, r_src0, 1.0f / (val), ne00))
+DEFINE_COMPUTE_SCALAR(div_f16, _Float16, hvx_div_scalar_f16_aa(r_dst, r_src0, val, ne00))
+
+typedef void (*compute_same_shape_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, size_t src1_stride, uint32_t ne00);
+
+#define DEFINE_COMPUTE_SAME_SHAPE(NAME, HVX_FN)                                                  \
+static void compute_same_shape_##NAME(                                                           \
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1,                                   \
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, size_t src1_stride, uint32_t ne00) { \
+    for (uint32_t r = 0; r < n_rows; r++) {                                                      \
+        HVX_FN(dst + r * dst_stride, src0 + r * src0_stride, src1 + r * src1_stride, ne00);      \
+    }                                                                                            \
+}
+
+DEFINE_COMPUTE_SAME_SHAPE(add_f32, hvx_add_f32_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(add_f16, hvx_add_f16_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(sub_f32, hvx_sub_f32_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(sub_f16, hvx_sub_f16_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(mul_f32, hvx_mul_f32_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(mul_f16, hvx_mul_f16_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(div_f32, hvx_div_f32_aaa)
+DEFINE_COMPUTE_SAME_SHAPE(div_f16, hvx_div_f16_aaa)
+
+typedef void (*compute_row_bcast_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00);
+
+#define DEFINE_COMPUTE_ROW_BCAST(NAME, HVX_FN)                                               \
+static void compute_row_bcast_##NAME(                                                        \
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1,                               \
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00) {                 \
+    for (uint32_t r = 0; r < n_rows; r++) {                                                  \
+        HVX_FN(dst + r * dst_stride, src0 + r * src0_stride, src1, ne00);                    \
+    }                                                                                        \
+}

-// Macro for vector op switch (All Aligned)
-#define COMPUTE_VECTOR_OP_AAA(DST, SRC0, SRC1, TYPE, N)                  \
-    if(TYPE == HTP_TYPE_F32) {                                           \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f32_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f32_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f32_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f32_aaa(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
-    }                                                                    \
-    else {                                                               \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f16_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f16_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f16_aaa(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f16_aaa(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
+DEFINE_COMPUTE_ROW_BCAST(add_f32, hvx_add_f32_aaa)
+DEFINE_COMPUTE_ROW_BCAST(add_f16, hvx_add_f16_aaa)
+DEFINE_COMPUTE_ROW_BCAST(sub_f32, hvx_sub_f32_aaa)
+DEFINE_COMPUTE_ROW_BCAST(sub_f16, hvx_sub_f16_aaa)
+DEFINE_COMPUTE_ROW_BCAST(mul_f32, hvx_mul_f32_aaa)
+DEFINE_COMPUTE_ROW_BCAST(mul_f16, hvx_mul_f16_aaa)
+DEFINE_COMPUTE_ROW_BCAST(div_f32, hvx_div_f32_aaa)
+DEFINE_COMPUTE_ROW_BCAST(div_f16, hvx_div_f16_aaa)
+
+typedef void (*compute_complex_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_plane,
+    uint32_t i01, uint32_t ne11, const struct fastdiv_values * div11, uint32_t nb11,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00);
+
+#define DEFINE_COMPUTE_COMPLEX(NAME, HVX_FN)                                                 \
+static void compute_complex_##NAME(                                                          \
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_plane,                         \
+    uint32_t i01, uint32_t ne11, const struct fastdiv_values * div11, uint32_t nb11,         \
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00) {                 \
+    for (uint32_t r = 0; r < n_rows; r++) {                                                  \
+        uint32_t i11 = fastmodulo(i01 + r, ne11, div11);                                     \
+        const uint8_t * r_src1 = src1_plane + i11 * nb11;                                    \
+        HVX_FN(dst + r * dst_stride, src0 + r * src0_stride, r_src1, ne00);                  \
+    }                                                                                        \
+}
+
+DEFINE_COMPUTE_COMPLEX(add_f32, hvx_add_f32_aau)
+DEFINE_COMPUTE_COMPLEX(add_f16, hvx_add_f16_aau)
+DEFINE_COMPUTE_COMPLEX(sub_f32, hvx_sub_f32_aau)
+DEFINE_COMPUTE_COMPLEX(sub_f16, hvx_sub_f16_aau)
+DEFINE_COMPUTE_COMPLEX(mul_f32, hvx_mul_f32_aau)
+DEFINE_COMPUTE_COMPLEX(mul_f16, hvx_mul_f16_aau)
+DEFINE_COMPUTE_COMPLEX(div_f32, hvx_div_f32_aau)
+DEFINE_COMPUTE_COMPLEX(div_f16, hvx_div_f16_aau)
+
+typedef void (*compute_repeat_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_plane,
+    uint32_t i01, uint32_t ne11, const struct fastdiv_values * div11, uint32_t nb11,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00, uint32_t ne10);
+
+#define DEFINE_COMPUTE_REPEAT(NAME, TYPE, HVX_FN)                                            \
+static void compute_repeat_##NAME(                                                           \
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_plane,                         \
+    uint32_t i01, uint32_t ne11, const struct fastdiv_values * div11, uint32_t nb11,         \
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00, uint32_t ne10) {  \
+    for (uint32_t r = 0; r < n_rows; r++) {                                                  \
+        uint32_t i11 = fastmodulo(i01 + r, ne11, div11);                                     \
+        const uint8_t * r_src1_row = src1_plane + i11 * nb11;                                \
+        uint8_t * r_dst = dst + r * dst_stride;                                              \
+        const uint8_t * r_src0 = src0 + r * src0_stride;                                     \
+        for (uint32_t c = 0; c < ne00; c += ne10) {                                          \
+            uint32_t len = MIN(ne10, ne00 - c);                                              \
+            HVX_FN(r_dst + c * sizeof(TYPE), r_src0 + c * sizeof(TYPE), r_src1_row, len);    \
+        }                                                                                    \
+    }                                                                                        \
+}
+
+DEFINE_COMPUTE_REPEAT(add_f32, float,    hvx_add_f32_uuu)
+DEFINE_COMPUTE_REPEAT(add_f16, _Float16, hvx_add_f16_uuu)
+DEFINE_COMPUTE_REPEAT(sub_f32, float,    hvx_sub_f32_uuu)
+DEFINE_COMPUTE_REPEAT(sub_f16, _Float16, hvx_sub_f16_uuu)
+DEFINE_COMPUTE_REPEAT(mul_f32, float,    hvx_mul_f32_uuu)
+DEFINE_COMPUTE_REPEAT(mul_f16, _Float16, hvx_mul_f16_uuu)
+DEFINE_COMPUTE_REPEAT(div_f32, float,    hvx_div_f32_uuu)
+DEFINE_COMPUTE_REPEAT(div_f16, _Float16, hvx_div_f16_uuu)
+
+typedef void (*compute_add_id_t)(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_data, const char * src2_data,
+    uint32_t i01, uint32_t i02, uint32_t nb20, uint32_t nb21, uint32_t src1_stride,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00);
+
+static void compute_add_id_f32(
+    uint8_t * dst, const uint8_t * src0, const uint8_t * src1_data, const char * src2_data,
+    uint32_t i01, uint32_t i02, uint32_t nb20, uint32_t nb21, uint32_t src1_stride,
+    uint32_t n_rows, size_t dst_stride, size_t src0_stride, uint32_t ne00) {
+    for (uint32_t r = 0; r < n_rows; r++) {
+        uint32_t r_i01 = i01 + r;
+        const int32_t idx = *(const int32_t *)(src2_data + r_i01 * nb20 + i02 * nb21);
+        if (idx < 0) {
+            memcpy(dst + r * dst_stride, src0 + r * src0_stride, ne00 * sizeof(float));
+            continue;
+        }
+        const uint8_t * r_src1 = src1_data + idx * src1_stride;
+        const uint8_t * r_src0 = src0 + r * src0_stride;
+        uint8_t * r_dst = dst + r * dst_stride;
+        hvx_add_f32_aaa(r_dst, r_src0, r_src1, ne00);
     }
+}
+
+// 1a. Scalar src1 in VTCM via DMA (ne10 == 1, ne12 == 1, ne13 == 1)
+static void binary_thread_scalar_dma(unsigned int nth, unsigned int ith, void * data) {
+    struct htp_binary_context * bctx = (struct htp_binary_context *) data;
+    struct htp_ops_context * octx = bctx->octx;
+    htp_binary_preamble;
+
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
+    const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
+    if (start_row >= end_row) return;
+
+    FARF(HIGH, "binary-scalar-dma: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);
+    size_t src0_spad_half    = layout->src0_spad_half_size;
+    size_t dst_spad_half     = layout->dst_spad_half_size;
+    const void * s1_table    = VTCM_LAYOUT_PTR(const void, bctx->vtcm_base, layout->off_src1);
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+    uint32_t ir_prefetch = start_row;
+    int spad_idx = 0;
+
+    for (int k = 0; k < 2 && ir_prefetch < end_row; k++) {
+        uint32_t current_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+        uint32_t i03, i02, i01, rem;
+        i03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+        rem = ir_prefetch - i03 * (ne02 * ne01);
+        i02 = fastdiv(rem, &bctx->src0_dim1_div);
+        i01 = rem - i02 * ne01;
+
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

-// Macro for vector op switch (Dst Aligned, Src0 Aligned, Src1 Unaligned)
-#define COMPUTE_VECTOR_OP_AAU(DST, SRC0, SRC1, TYPE, N)                  \
-    if(TYPE == HTP_TYPE_F32) {                                           \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f32_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f32_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f32_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f32_aau(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
-    }                                                                    \
-    else {                                                               \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f16_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f16_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f16_aau(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f16_aau(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
+        uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
+        uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;
+
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes,
+                                                             current_block_size);
+        ir_prefetch += current_block_size;
+        spad_idx ^= 1;
     }

-// Macro for vector op switch (All Unaligned - generic loop used in element repeat)
-#define COMPUTE_VECTOR_OP_UUU(DST, SRC0, SRC1, TYPE, N)                  \
-    if(TYPE == HTP_TYPE_F32) {                                           \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f32_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f32_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f32_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f32_uuu(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
-    }                                                                    \
-    else {                                                               \
-        switch (octx->op) {                                              \
-            case HTP_OP_ADD: hvx_add_f16_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_SUB: hvx_sub_f16_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_MUL: hvx_mul_f16_uuu(DST, SRC0, SRC1, N); break; \
-            case HTP_OP_DIV: hvx_div_f16_uuu(DST, SRC0, SRC1, N); break; \
-            default: break;                                              \
-        }                                                                \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_scalar_dma_t compute = (compute_scalar_dma_t) bctx->compute;
+
+    for (uint32_t ir = start_row; ir < end_row; ) {
+        uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
+
+        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
+
+        uint32_t i03, i02, i01, rem;
+        i03 = fastdiv(ir, &bctx->src0_dim12_div);
+        rem = ir - i03 * (ne02 * ne01);
+        i02 = fastdiv(rem, &bctx->src0_dim1_div);
+        i01 = rem - i02 * ne01;
+
+        uint32_t cur_i11 = fastmodulo(i01, ne11, &bctx->src1_dim1_div);
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
+        compute(d_spad, s0_spad, s1_table, cur_i11, ne11, current_block_size,
+                bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
+
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes,
+                                                                current_block_size);
+
+        if (ir_prefetch < end_row) {
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02 = fastdiv(prem, &bctx->src0_dim1_div);
+            p01 = prem - p02 * ne01;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes,
+                                                               next_block_size);
+            ir_prefetch += next_block_size;
+        }
+        ir += current_block_size;
     }

-// 1. Scalar src1 (ne10 == 1)
-static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
+    dma_queue_flush(dma_q);
+}
+
+// 1b. Scalar src1 dynamic / pointer (ne10 == 1)
+static void binary_thread_scalar(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;
     htp_binary_preamble;

-    const uint32_t src0_type = octx->src[0]->type;
-    const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    FARF(HIGH, "binary-scalar: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+    FARF(HIGH, "binary-scalar: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
-    size_t src0_spad_half    = octx->src0_spad.size_per_thread / 2;
-    size_t dst_spad_half     = octx->dst_spad.size_per_thread  / 2;
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);
+    size_t src0_spad_half    = layout->src0_spad_half_size;
+    size_t dst_spad_half     = layout->dst_spad_half_size;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

-    // Preamble
     for (int k = 0; k < 2 && ir_prefetch < end_row; k++) {
         uint32_t current_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
         uint32_t i03, i02, i01, rem;
@@ -211,26 +405,26 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
         i02 = fastdiv(rem, &bctx->src0_dim1_div);
         i01 = rem - i02 * ne01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }

-    // Main loop
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_scalar_t compute = (compute_scalar_t) bctx->compute;

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);

-        uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i03, i02, i01, rem;
         i03 = fastdiv(ir, &bctx->src0_dim12_div);
@@ -238,67 +432,62 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
         i02 = fastdiv(rem, &bctx->src0_dim1_div);
         i01 = rem - i02 * ne01;

-        // src1 indices (broadcast/repeat)
         uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
         uint32_t i12 = fastmodulo(i02, ne12, &bctx->src1_dim2_div);
         uint32_t i11 = fastmodulo(i01, ne11, &bctx->src1_dim1_div);

-        uint8_t * src1_ptr = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
+        const uint8_t * src1_ptr = (const uint8_t *)(uintptr_t) src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
         uint32_t s1_stride = (ne11 == 1) ? 0 : nb11;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
-            COMPUTE_SCALAR_OP(r_dst, r_src0, src1_ptr, src0_type, ne00);
-            src1_ptr += s1_stride;
-        }
+        compute(d_spad, s0_spad, src1_ptr, s1_stride, current_block_size,
+                bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
-             uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03, p02, p01, prem;
-             p03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             prem = ir_prefetch - p03 * (ne02 * ne01);
-             p02 = fastdiv(prem, &bctx->src0_dim1_div);
-             p01 = prem - p02 * ne01;
-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-
-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
-             ir_prefetch += next_block_size;
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02 = fastdiv(prem, &bctx->src0_dim1_div);
+            p01 = prem - p02 * ne01;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
+            ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 // 2. Vector Same Shape (ne1x == ne0x) or Simple Broadcast
-static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, void * data) {
+static void binary_thread_vector_same_shape(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;
     htp_binary_preamble;

-    const uint32_t src0_type = octx->src[0]->type;
-    const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    FARF(HIGH, "binary-same-shape: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+    FARF(HIGH, "binary-same-shape: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * src1_spad_base = octx->src1_spad.data + (ith * octx->src1_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * src1_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src1) + (ith * layout->src1_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);

-    size_t src0_spad_half = octx->src0_spad.size_per_thread / 2;
-    size_t src1_spad_half = octx->src1_spad.size_per_thread / 2;
-    size_t dst_spad_half  = octx->dst_spad.size_per_thread  / 2;
+    size_t src0_spad_half = layout->src0_spad_half_size;
+    size_t src1_spad_half = layout->src1_spad_half_size;
+    size_t dst_spad_half  = layout->dst_spad_half_size;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

@@ -314,36 +503,33 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
         uint32_t i12 = (ne12 == 1) ? 0 : i02;
         uint32_t i11 = (ne11 == 1) ? 0 : i01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * src1_curr = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t src1_curr = src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * s1_spad = src1_spad_base + spad_idx * src1_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
-        dma_queue_push(q, dma_make_ptr(s1_spad, src1_curr), bctx->src1_row_size_aligned, nb11, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(s1_spad, src1_curr), bctx->src1_row_size_aligned, nb11, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_same_shape_t compute = (compute_same_shape_t) bctx->compute;

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
-        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
-        uint8_t * s1_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
+        uint8_t * s1_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_src1 = s1_spad + r * bctx->src1_row_size_aligned;
-            uint8_t * r_dst  = d_spad  + r * bctx->dst_row_size_aligned;
-            COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
-        }
+        compute(d_spad, s0_spad, s1_spad, current_block_size,
+                bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, bctx->src1_row_size_aligned, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint32_t i03, i02, i01, rem;
@@ -351,61 +537,61 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
         rem = ir - i03 * (ne02 * ne01);
         i02 = fastdiv(rem, &bctx->src0_dim1_div);
         i01 = rem - i02 * ne01;
-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
-             uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03, p02, p01, prem;
-             p03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             prem = ir_prefetch - p03 * (ne02 * ne01);
-             p02 = fastdiv(prem, &bctx->src0_dim1_div);
-             p01 = prem - p02 * ne01;
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03 = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02 = fastdiv(prem, &bctx->src0_dim1_div);
+            p01 = prem - p02 * ne01;

-             uint32_t p13 = (ne13 == 1) ? 0 : p03;
-             uint32_t p12 = (ne12 == 1) ? 0 : p02;
-             uint32_t p11 = (ne11 == 1) ? 0 : p01;
+            uint32_t p13 = (ne13 == 1) ? 0 : p03;
+            uint32_t p12 = (ne12 == 1) ? 0 : p02;
+            uint32_t p11 = (ne11 == 1) ? 0 : p01;

-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-             uint8_t * s1_next = (uint8_t *)src1->data + p13 * nb13 + p12 * nb12 + p11 * nb11;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_addr_t s1_next = src1->data + p13 * nb13 + p12 * nb12 + p11 * nb11;

-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
-             dma_queue_push(q, dma_make_ptr(s1_spad, s1_next), bctx->src1_row_size_aligned, nb11, row_size_bytes, next_block_size);
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
+            dma_queue_push(dma_q, dma_make_data(s1_spad, s1_next), bctx->src1_row_size_aligned, nb11, row_size_bytes, next_block_size);

-             ir_prefetch += next_block_size;
+            ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 // 3. Row Broadcast (ne11 == 1, ne12 == 1, single row src1)
-static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith, void * data) {
+static void binary_thread_vector_row_broadcast(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;
     htp_binary_preamble;

-    const uint32_t src0_type  = octx->src[0]->type;
-    const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    FARF(HIGH, "binary-row-bcast: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+    FARF(HIGH, "binary-row-bcast: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * src1_spad_base = octx->src1_spad.data + (ith * octx->src1_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);

-    size_t src0_spad_half = octx->src0_spad.size_per_thread / 2;
-    size_t dst_spad_half  = octx->dst_spad.size_per_thread  / 2;
+    size_t src0_spad_half = layout->src0_spad_half_size;
+    size_t dst_spad_half  = layout->dst_spad_half_size;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

-    void * s1_ptr = (void *) src1_spad_base;
+    void * s1_ptr = VTCM_LAYOUT_PTR(void, bctx->vtcm_base, layout->off_src1);

     for (int k = 0; k < 2 && ir_prefetch < end_row; k++) {
         uint32_t current_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
@@ -414,77 +600,76 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_row_bcast_t compute = (compute_row_bcast_t) bctx->compute;

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
-        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_src1 = (uint8_t *)s1_ptr; // Constant
-            uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
-            COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
-        }
+        compute(d_spad, s0_spad, (const uint8_t *)s1_ptr, current_block_size,
+                bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
         uint32_t rem = ir - i03 * (ne02 * ne01);
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;
-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
-             uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             uint32_t prem = ir_prefetch - p03 * (ne02 * ne01);
-             uint32_t p02  = fastdiv(prem, &bctx->src0_dim1_div);
-             uint32_t p01  = prem - p02 * ne01;
-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
-             ir_prefetch += next_block_size;
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02  = fastdiv(prem, &bctx->src0_dim1_div);
+            p01  = prem - p02 * ne01;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
+            ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 // 4. Vector Complex (ne10 == ne00, complex broadcast)
-static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void * data) {
+static void binary_thread_vector_complex(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;
     htp_binary_preamble;

-    const uint32_t src0_type = octx->src[0]->type;
-    const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    FARF(HIGH, "binary-complex: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+    FARF(HIGH, "binary-complex: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
-    size_t src0_spad_half    = octx->src0_spad.size_per_thread / 2;
-    size_t dst_spad_half     = octx->dst_spad.size_per_thread  / 2;
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);
+    size_t src0_spad_half    = layout->src0_spad_half_size;
+    size_t dst_spad_half     = layout->dst_spad_half_size;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

@@ -495,86 +680,81 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_complex_t compute = (compute_complex_t) bctx->compute;

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
-        uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
         uint32_t rem = ir - i03 * (ne02 * ne01);
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

+        uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
+        uint32_t i12 = fastmodulo(i02, ne12, &bctx->src1_dim2_div);
+        const uint8_t * src1_plane = (const uint8_t *)(uintptr_t) src1->data + i13 * nb13 + i12 * nb12;
+
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint32_t r_i01 = i01 + r;
-            uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
-            uint32_t i12 = fastmodulo(i02, ne12, &bctx->src1_dim2_div);
-            uint32_t i11 = fastmodulo(r_i01, ne11, &bctx->src1_dim1_div);
-
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_src1 = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
-            uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
-
-            // Read src1 from DDR (unaligned)
-            COMPUTE_VECTOR_OP_AAU(r_dst, r_src0, r_src1, src0_type, ne00);
-        }
+        compute(d_spad, s0_spad, src1_plane, i01, ne11, &bctx->src1_dim1_div, nb11,
+                current_block_size, bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
-             uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             uint32_t prem = ir_prefetch - p03 * (ne02 * ne01);
-             uint32_t p02  = fastdiv(prem, &bctx->src0_dim1_div);
-             uint32_t p01  = prem - p02 * ne01;
-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
-             ir_prefetch += next_block_size;
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02  = fastdiv(prem, &bctx->src0_dim1_div);
+            p01  = prem - p02 * ne01;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
+            ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 // 5. Element Repeat (ne10 != ne00)
-static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void * data) {
+static void binary_thread_element_repeat(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;
     htp_binary_preamble;

-    const uint32_t src0_type = octx->src[0]->type;
-    const uint32_t elem_size_bytes = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
-    const uint32_t row_size_bytes = ne00 * elem_size_bytes;;
+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
-    size_t src0_spad_half    = octx->src0_spad.size_per_thread / 2;
-    size_t dst_spad_half     = octx->dst_spad.size_per_thread  / 2;
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);
+    size_t src0_spad_half    = layout->src0_spad_half_size;
+    size_t dst_spad_half     = layout->dst_spad_half_size;

-    FARF(HIGH, "binary-repeat: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
+    FARF(HIGH, "binary-repeat: %d/%d (%u:%u) row-size %u (%u)",
+         ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

@@ -585,71 +765,62 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    compute_repeat_t compute = (compute_repeat_t) bctx->compute;

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
-        uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad  = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
         uint32_t rem = ir - i03 * (ne02 * ne01);
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

+        uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
+        uint32_t i12 = fastmodulo(i02, ne12, &bctx->src1_dim2_div);
+        const uint8_t * src1_plane = (const uint8_t *)(uintptr_t) src1->data + i13 * nb13 + i12 * nb12;
+
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint32_t r_i01 = i01 + r;
-            uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
-            uint32_t i12 = fastmodulo(i02, ne12, &bctx->src1_dim2_div);
-            uint32_t i11 = fastmodulo(r_i01, ne11, &bctx->src1_dim1_div);
-
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_src1_row = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
-            uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
-
-            // Repeat src1 row
-            for (uint32_t c = 0; c < ne00; c += ne10) {
-                uint32_t len = MIN(ne10, ne00 - c);
-                // Use UUU for speed and simplicity
-                COMPUTE_VECTOR_OP_UUU(r_dst + c * elem_size_bytes, r_src0 + c * elem_size_bytes, r_src1_row, src0_type, len);
-            }
-        }
+        compute(d_spad, s0_spad, src1_plane, i01, ne11, &bctx->src1_dim1_div, nb11,
+                current_block_size, bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00, ne10);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
-             uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             uint32_t prem = ir_prefetch - p03 * (ne02 * ne01);
-             uint32_t p02  = fastdiv(prem, &bctx->src0_dim1_div);
-             uint32_t p01  = prem - p02 * ne01;
-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
-             ir_prefetch += next_block_size;
+            uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
+            uint32_t p03, p02, p01, prem;
+            p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+            prem = ir_prefetch - p03 * (ne02 * ne01);
+            p02  = fastdiv(prem, &bctx->src0_dim1_div);
+            p01  = prem - p02 * ne01;
+            dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+            dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
+            ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 // 6. ADD_ID (src1 gathered via src2 indices)
-static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
+static void binary_thread_add_id_f32(unsigned int nth, unsigned int ith, void * data) {
     struct htp_binary_context * bctx = (struct htp_binary_context *) data;
     struct htp_ops_context * octx = bctx->octx;

@@ -662,27 +833,29 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t ne01 = src0->ne[1];
     const uint32_t ne02 = src0->ne[2];
     const uint32_t ne03 = src0->ne[3];
-    const uint32_t ne11 = src1->ne[1]; // for bounds check

     const uint32_t nb01 = src0->nb[1];
     const uint32_t nb02 = src0->nb[2];
     const uint32_t nb03 = src0->nb[3];
-    const uint32_t nb11 = src1->nb[1]; // src1 row stride
+    const uint32_t src1_stride = bctx->src1_row_size_aligned;

     const uint32_t nb1 = dst->nb[1];
     const uint32_t nb2 = dst->nb[2];
     const uint32_t nb3 = dst->nb[3];

+    const uint32_t row_size_bytes = bctx->row_size_bytes;
     const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
     const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

-    uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * dst_spad_base  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
-    size_t src0_spad_half    = octx->src0_spad.size_per_thread / 2;
-    size_t dst_spad_half     = octx->dst_spad.size_per_thread  / 2;
+    const struct htp_binary_vtcm_layout * layout = &bctx->vtcm_layout;
+    uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
+    uint8_t * dst_spad_base  = VTCM_LAYOUT_PTR(uint8_t, bctx->vtcm_base, layout->off_dst)  + (ith * layout->dst_bytes_per_thread);
+    const uint8_t * vtcm_src1 = VTCM_LAYOUT_PTR(const uint8_t, bctx->vtcm_base, layout->off_src1);
+    size_t src0_spad_half    = layout->src0_spad_half_size;
+    size_t dst_spad_half     = layout->dst_spad_half_size;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     uint32_t ir_prefetch = start_row;
     int spad_idx = 0;

@@ -693,14 +866,14 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

-        uint8_t * src0_curr = (uint8_t *)src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
-        uint8_t * dst_curr  = (uint8_t *)dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;
+        dma_addr_t src0_curr = src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01;
+        dma_addr_t dst_curr  = dst->data  + i03 * nb3  + i02 * nb2  + i01 * nb1;

         uint8_t * s0_spad = src0_spad_base + spad_idx * src0_spad_half;
         uint8_t * d_spad  = dst_spad_base  + spad_idx * dst_spad_half;

-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, ne00 * sizeof(float), 0);
-        dma_queue_push(q, dma_make_ptr(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, ne00 * sizeof(float), current_block_size);
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, 0);
+        dma_queue_push(dma_q, dma_make_data(s0_spad, src0_curr), bctx->src0_row_size_aligned, nb01, row_size_bytes, current_block_size);
         ir_prefetch += current_block_size;
         spad_idx ^= 1;
     }
@@ -709,8 +882,8 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {

     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
-        uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
-        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
+        uint8_t * d_spad = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
         uint32_t rem = ir - i03 * (ne02 * ne01);
@@ -718,42 +891,37 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
         uint32_t i01 = rem - i02 * ne01;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
-        for (uint32_t r = 0; r < current_block_size; r++) {
-            uint32_t r_i01 = i01 + r; // linear within block since we split at ne01
-
-            const int32_t idx = *(int32_t *)((char *)src2->data + r_i01 * src2->nb[0] + i02 * src2->nb[1]);
-
-            uint8_t * r_src1 = (uint8_t *)src1->data + idx * nb11;
-            uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
-            uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
-
-            hvx_add_f32_aau(r_dst, r_src0, r_src1, ne00);
-        }
+        compute_add_id_t compute = (compute_add_id_t) bctx->compute;
+        compute(d_spad, s0_spad, vtcm_src1, (const char *)(uintptr_t)src2->data,
+                i01, i02, src2->nb[0], src2->nb[1], src1_stride,
+                current_block_size, bctx->dst_row_size_aligned, bctx->src0_row_size_aligned, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
-        dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, ne00 * sizeof(float), current_block_size);
+        dma_addr_t dst_curr = dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);

         if (ir_prefetch < end_row) {
              uint32_t next_block_size = calc_block_size(bctx, ir_prefetch, end_row, ne01, ne02);
-             uint32_t p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
-             uint32_t prem = ir_prefetch - p03 * (ne02 * ne01);
-             uint32_t p02  = fastdiv(prem, &bctx->src0_dim1_div);
-             uint32_t p01  = prem - p02 * ne01;
-             uint8_t * s0_next = (uint8_t *)src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
-             dma_queue_push(q, dma_make_ptr(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, ne00 * sizeof(float), next_block_size);
+             uint32_t p03, p02, p01, prem;
+             p03  = fastdiv(ir_prefetch, &bctx->src0_dim12_div);
+             prem = ir_prefetch - p03 * (ne02 * ne01);
+             p02  = fastdiv(prem, &bctx->src0_dim1_div);
+             p01  = prem - p02 * ne01;
+             dma_addr_t s0_next = src0->data + p03 * nb03 + p02 * nb02 + p01 * nb01;
+             dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), bctx->src0_row_size_aligned, nb01, row_size_bytes, next_block_size);
              ir_prefetch += next_block_size;
         }
         ir += current_block_size;
     }

-    dma_queue_flush(q);
+    dma_queue_flush(dma_q);
 }

 static int execute_op_binary(struct htp_ops_context * octx) {
     const struct htp_tensor * src0 = octx->src[0];
     const struct htp_tensor * src1 = octx->src[1];
     const struct htp_tensor * dst  = octx->dst;
+    const struct htp_binary_kernel_params * kparams = (const struct htp_binary_kernel_params *) octx->kernel_params;

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];

@@ -770,7 +938,8 @@ static int execute_op_binary(struct htp_ops_context * octx) {
     if (octx->ctx->mdev.count > 1) {
         uint32_t rows_per_chunk = 0;
         htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk,
+                                                       octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
         row_start = range.start;
         nrows     = range.count;
     }
@@ -779,92 +948,57 @@ static int execute_op_binary(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

-    const uint32_t n_threads = octx->n_threads;
-
-    size_t src0_row_size_aligned = hex_round_up(src0_row_size, VLEN);
-    size_t src1_row_size_aligned = hex_round_up(src1_row_size, VLEN);
-    size_t dst_row_size_aligned  = hex_round_up(dst_row_size,  VLEN);
-
-    bool is_add_id = (octx->op == HTP_OP_ADD_ID);
-    bool is_scalar = !is_add_id && (src1->ne[0] == 1);
-
-    bool is_transposed = (src0->nb[1] < src0_row_size || src1->nb[1] < src1_row_size || dst->nb[1] < dst_row_size);
-
-    bool is_same_shape = !is_add_id && !is_scalar && !is_transposed &&
-               (src1->ne[0] == src0->ne[0] && src0->ne[0] % VLEN == 0) &&
-               (src1->ne[1] == src0->ne[1] || src1->ne[1] == 1) &&
-               (src1->ne[2] == src0->ne[2] || src1->ne[2] == 1) &&
-               (src1->ne[3] == src0->ne[3] || src1->ne[3] == 1);
-
-    bool is_row_bcast = is_same_shape && (src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1);
-    bool is_complex   = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] == src0->ne[0]);
-    bool is_repeat    = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] != src0->ne[0]);
-
-    size_t spad_row_total;
-    if (is_same_shape) {
-        spad_row_total = 2 * (src0_row_size_aligned + src1_row_size_aligned + dst_row_size_aligned);
-    } else {
-        spad_row_total = 2 * (src0_row_size_aligned + dst_row_size_aligned);
-    }
-
-    size_t rows_per_buffer = octx->ctx->vtcm_size / (n_threads * spad_row_total);
-
-    // Adjust for static src1 in row_bcast case
-    if (is_row_bcast) {
-        size_t needed_static = src1_row_size_aligned;
-        if (octx->ctx->vtcm_size < needed_static) return HTP_STATUS_VTCM_TOO_SMALL;
-        size_t avail = octx->ctx->vtcm_size - needed_static;
-        rows_per_buffer = avail / (n_threads * spad_row_total);
+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
     }

-    if (rows_per_buffer < 1) {
-        FARF(ERROR, "binary: VTCM too small\n");
-        return HTP_STATUS_VTCM_TOO_SMALL;
+    const uint32_t n_threads = octx->n_threads;
+    const size_t src0_row_size_aligned = kparams->src0_row_size_aligned;
+    const size_t src1_row_size_aligned = kparams->src1_row_size_aligned;
+    const size_t dst_row_size_aligned  = kparams->dst_row_size_aligned;
+
+    if (htp_tensor_is_extended(src1)) {
+        if (kparams->kernel_type != HTP_BINARY_KERNEL_SAME_SHAPE &&
+            kparams->kernel_type != HTP_BINARY_KERNEL_ROW_BCAST &&
+            kparams->kernel_type != HTP_BINARY_KERNEL_SCALAR_DMA &&
+            kparams->kernel_type != HTP_BINARY_KERNEL_ADD_ID) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
     }

-    octx->src0_spad.size_per_thread = rows_per_buffer * 2 * src0_row_size_aligned;
-    octx->dst_spad.size_per_thread  = rows_per_buffer * 2 * dst_row_size_aligned;
-
-    if (is_add_id || is_scalar || is_complex || is_repeat || is_row_bcast) {
-        octx->src1_spad.size_per_thread = 0;
-    } else {
-        octx->src1_spad.size_per_thread = rows_per_buffer * 2 * src1_row_size_aligned;
+    if (octx->op == HTP_OP_ADD_ID && htp_tensor_is_extended(octx->src[2])) {
+        return HTP_STATUS_NO_SUPPORT;
     }

-    octx->dst_spad.size  = n_threads * octx->dst_spad.size_per_thread;
-    octx->src0_spad.size = n_threads * octx->src0_spad.size_per_thread;
-    if (is_row_bcast) {
-        octx->src1_spad.size = src1_row_size_aligned;
-    } else {
-        octx->src1_spad.size = n_threads * octx->src1_spad.size_per_thread;
-    }
+    struct htp_binary_context bctx;
+    bctx.vtcm_base = (uint8_t *) octx->ctx->vtcm_base;
+    htp_binary_vtcm_layout_build(&bctx.vtcm_layout, kparams, octx->ctx->vtcm_size);

-    if (octx->ctx->vtcm_size < (octx->src0_spad.size + octx->src1_spad.size + octx->dst_spad.size)) {
+    if (bctx.vtcm_layout.rows_per_buffer == 0 || bctx.vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
         return HTP_STATUS_VTCM_TOO_SMALL;
     }

-    octx->src0_spad.data = octx->ctx->vtcm_base;                        octx->src0_spad.src = NULL;
-    octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size; octx->src1_spad.src = NULL;
-    octx->dst_spad.data  = octx->src1_spad.data + octx->src1_spad.size; octx->dst_spad.src  = NULL;
-
-    if ((octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        return HTP_STATUS_OK;
-    }
-
-    dma_queue * q = octx->ctx->dma[0];
-    if (is_row_bcast) {
-        dma_queue_push(q, dma_make_ptr(octx->src1_spad.data, (const void *) src1->data), src1_row_size_aligned, 0, src1->ne[0] * elem_size, 1);
+    dma_queue * dma_q = octx->ctx->dma[0];
+    uint8_t * vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, bctx.vtcm_base, bctx.vtcm_layout.off_src1);
+    if (kparams->kernel_type == HTP_BINARY_KERNEL_ROW_BCAST) {
+        dma_queue_push(dma_q, dma_make_data(vtcm_src1, src1->data), bctx.vtcm_layout.src1_size, 0, src1->ne[0] * elem_size, 1);
+    } else if (kparams->kernel_type == HTP_BINARY_KERNEL_SCALAR_DMA) {
+        dma_queue_push(dma_q, dma_make_data(vtcm_src1, src1->data), bctx.vtcm_layout.src1_size, 0, src1->ne[1] * elem_size, 1);
+    } else if (kparams->kernel_type == HTP_BINARY_KERNEL_ADD_ID) {
+        dma_queue_push(dma_q, dma_make_data(vtcm_src1, src1->data),
+                       kparams->src1_row_size_aligned, src1->nb[1],
+                       src1->ne[0] * elem_size, src1->ne[1]);
     }

-    struct htp_binary_context bctx;
     bctx.octx                  = octx;
     bctx.nrows_per_thread      = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
     bctx.total_rows            = nrows;
     bctx.row_start             = row_start;
-    bctx.block_max             = rows_per_buffer;
+    bctx.block_max             = bctx.vtcm_layout.rows_per_buffer;
     bctx.src0_row_size_aligned = src0_row_size_aligned;
     bctx.src1_row_size_aligned = src1_row_size_aligned;
     bctx.dst_row_size_aligned  = dst_row_size_aligned;
+    bctx.row_size_bytes        = src0_row_size;

     bctx.src0_dim1_div  = init_fastdiv_values(src0->ne[1]);
     bctx.src0_dim2_div  = init_fastdiv_values(src0->ne[2]);
@@ -880,19 +1014,153 @@ static int execute_op_binary(struct htp_ops_context * octx) {
     bool src0_contig_dim2 = (src0->nb[3] == src0->ne[2] * src0->nb[2]);
     bool dst_contig_dim2  = (dst->nb[3]  == src0->ne[2] * dst->nb[2]);

-    bctx.split_at_ne01 = (src0->ne[2] > 1) && ((src1->ne[1] > 1) || (src1->ne[2] > 1) || !src0_contig_dim1 || !dst_contig_dim1);
+    bctx.split_at_ne01 = (octx->op == HTP_OP_ADD_ID) ||
+        ((src0->ne[2] > 1) && ((src1->ne[1] > 1) || (src1->ne[2] > 1) || !src0_contig_dim1 || !dst_contig_dim1));
     bctx.split_at_ne02 = (src0->ne[3] > 1) && ((src1->ne[2] > 1) || (src1->ne[3] > 1) || !src0_contig_dim2 || !dst_contig_dim2);

-    worker_callback_t worker_func;
-    if (is_add_id)          worker_func = binary_job_add_id;
-    else if (is_scalar)     worker_func = binary_job_scalar;
-    else if (is_row_bcast)  worker_func = binary_job_vector_row_broadcast;
-    else if (is_same_shape) worker_func = binary_job_vector_same_shape;
-    else if (is_complex)    worker_func = binary_job_vector_complex;
-    else                    worker_func = binary_job_element_repeat;
+    worker_callback_t worker_func = NULL;
+    void * compute_func = NULL;
+
+    switch (kparams->kernel_type) {
+        case HTP_BINARY_KERNEL_SAME_SHAPE:
+            worker_func = binary_thread_vector_same_shape;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_same_shape_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_same_shape_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_same_shape_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_same_shape_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_same_shape_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_same_shape_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_same_shape_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_same_shape_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_ROW_BCAST:
+            worker_func = binary_thread_vector_row_broadcast;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_row_bcast_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_row_bcast_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_row_bcast_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_row_bcast_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_row_bcast_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_row_bcast_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_row_bcast_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_row_bcast_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_SCALAR_DMA:
+            worker_func = binary_thread_scalar_dma;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_scalar_dma_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_scalar_dma_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_scalar_dma_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_scalar_dma_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_scalar_dma_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_scalar_dma_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_scalar_dma_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_scalar_dma_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_SCALAR:
+            worker_func = binary_thread_scalar;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_scalar_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_scalar_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_scalar_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_scalar_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_scalar_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_scalar_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_scalar_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_scalar_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_COMPLEX:
+            worker_func = binary_thread_vector_complex;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_complex_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_complex_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_complex_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_complex_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_complex_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_complex_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_complex_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_complex_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_REPEAT:
+            worker_func = binary_thread_element_repeat;
+            if (src0_type == HTP_TYPE_F32) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_repeat_add_f32; break;
+                    case HTP_OP_SUB: compute_func = compute_repeat_sub_f32; break;
+                    case HTP_OP_MUL: compute_func = compute_repeat_mul_f32; break;
+                    case HTP_OP_DIV: compute_func = compute_repeat_div_f32; break;
+                    default: break;
+                }
+            } else if (src0_type == HTP_TYPE_F16) {
+                switch (octx->op) {
+                    case HTP_OP_ADD: compute_func = compute_repeat_add_f16; break;
+                    case HTP_OP_SUB: compute_func = compute_repeat_sub_f16; break;
+                    case HTP_OP_MUL: compute_func = compute_repeat_mul_f16; break;
+                    case HTP_OP_DIV: compute_func = compute_repeat_div_f16; break;
+                    default: break;
+                }
+            }
+            break;
+        case HTP_BINARY_KERNEL_ADD_ID:
+            if (octx->op == HTP_OP_ADD_ID && src0_type == HTP_TYPE_F32) {
+                worker_func = binary_thread_add_id_f32;
+                compute_func = (void *) compute_add_id_f32;
+            }
+            break;
+        default: break;
+    }
+
+    if (!worker_func || !compute_func) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
+    bctx.compute = compute_func;

-    if (is_row_bcast) {
-        dma_queue_pop(q);
+    if (kparams->kernel_type == HTP_BINARY_KERNEL_ROW_BCAST ||
+        kparams->kernel_type == HTP_BINARY_KERNEL_SCALAR_DMA ||
+        kparams->kernel_type == HTP_BINARY_KERNEL_ADD_ID) {
+        dma_queue_pop(dma_q);
     }

     work_queue_run(octx->ctx->work_queue, worker_func, &bctx, n_threads);
diff --git a/ggml/src/ggml-hexagon/htp/binary-ops.h b/ggml/src/ggml-hexagon/htp/binary-ops.h
new file mode 100644
index 000000000..b99f2ad64
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/binary-ops.h
@@ -0,0 +1,111 @@
+#ifndef HTP_BINARY_OPS_H
+#define HTP_BINARY_OPS_H
+
+#include <stddef.h>
+#include <stdint.h>
+#include <string.h>
+
+#include "hex-common.h"
+#include "htp-ops.h"
+#include "htp-vtcm.h"
+
+enum htp_binary_kernel_type {
+    HTP_BINARY_KERNEL_SAME_SHAPE = 0,
+    HTP_BINARY_KERNEL_ROW_BCAST,
+    HTP_BINARY_KERNEL_SCALAR_DMA,
+    HTP_BINARY_KERNEL_SCALAR,
+    HTP_BINARY_KERNEL_ADD_ID,
+    HTP_BINARY_KERNEL_COMPLEX,
+    HTP_BINARY_KERNEL_REPEAT,
+};
+
+struct htp_binary_kernel_params {
+    uint32_t kernel_type;
+    uint32_t n_threads;
+    uint32_t rows_per_buffer;
+
+    uint32_t src0_row_size_aligned;
+    uint32_t src1_row_size_aligned;
+    uint32_t dst_row_size_aligned;
+
+    uint32_t src1_size;
+    uint32_t vtcm_size;
+};
+
+#if defined(__cplusplus)
+static_assert(sizeof(struct htp_binary_kernel_params) <= 128, "htp_binary_kernel_params is too large for kernel_params blob");
+#else
+_Static_assert(sizeof(struct htp_binary_kernel_params) <= 128, "htp_binary_kernel_params is too large for kernel_params blob");
+#endif
+
+struct htp_binary_vtcm_layout {
+    size_t total_bytes;
+    size_t off_src0;
+    size_t off_src1;
+    size_t off_dst;
+
+    size_t src0_bytes_per_thread;
+    size_t src1_bytes_per_thread;
+    size_t dst_bytes_per_thread;
+
+    size_t src0_spad_half_size;
+    size_t src1_spad_half_size;
+    size_t dst_spad_half_size;
+
+    size_t src1_size;
+    uint32_t rows_per_buffer;
+};
+
+static inline void htp_binary_vtcm_layout_build(
+    struct htp_binary_vtcm_layout * L,
+    const struct htp_binary_kernel_params * kparams,
+    size_t vtcm_size
+) {
+    memset(L, 0, sizeof(*L));
+
+    const uint32_t n_threads = kparams->n_threads;
+    if (n_threads == 0) {
+        return;
+    }
+
+    const size_t spad_row_total = (kparams->kernel_type == HTP_BINARY_KERNEL_SAME_SHAPE)
+        ? 2 * (kparams->src0_row_size_aligned + kparams->src1_row_size_aligned + kparams->dst_row_size_aligned)
+        : 2 * (kparams->src0_row_size_aligned + kparams->dst_row_size_aligned);
+
+    if (spad_row_total == 0 || vtcm_size < kparams->src1_size) {
+        return;
+    }
+
+    const size_t rows_per_buffer = (vtcm_size - kparams->src1_size) / (n_threads * spad_row_total);
+    if (rows_per_buffer == 0) {
+        return;
+    }
+
+    L->rows_per_buffer = (uint32_t) rows_per_buffer;
+    L->src1_size = kparams->src1_size;
+
+    L->src0_bytes_per_thread = rows_per_buffer * 2 * kparams->src0_row_size_aligned;
+    L->dst_bytes_per_thread  = rows_per_buffer * 2 * kparams->dst_row_size_aligned;
+    L->src1_bytes_per_thread = (kparams->kernel_type == HTP_BINARY_KERNEL_SAME_SHAPE)
+        ? rows_per_buffer * 2 * kparams->src1_row_size_aligned
+        : 0;
+
+    L->src0_spad_half_size = L->src0_bytes_per_thread / 2;
+    L->src1_spad_half_size = L->src1_bytes_per_thread / 2;
+    L->dst_spad_half_size  = L->dst_bytes_per_thread / 2;
+
+    const size_t src0_total = n_threads * L->src0_bytes_per_thread;
+    const size_t src1_total = (kparams->src1_size > 0)
+        ? kparams->src1_size
+        : n_threads * L->src1_bytes_per_thread;
+    const size_t dst_total  = n_threads * L->dst_bytes_per_thread;
+
+    size_t off = 0;
+    VTCM_LAYOUT_ALLOC(off, off_src0, src0_total);
+    VTCM_LAYOUT_ALLOC(off, off_src1, src1_total);
+    VTCM_LAYOUT_ALLOC(off, off_dst,  dst_total);
+
+    L->total_bytes = off;
+}
+
+#endif
diff --git a/ggml/src/ggml-hexagon/htp/concat-ops.c b/ggml/src/ggml-hexagon/htp/concat-ops.c
index 966e867b3..1fa6ec1bd 100644
--- a/ggml/src/ggml-hexagon/htp/concat-ops.c
+++ b/ggml/src/ggml-hexagon/htp/concat-ops.c
@@ -6,7 +6,7 @@
 #include "hexagon_types.h"
 #include "hexagon_protos.h"
 #include "hvx_hexagon_protos.h"
-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "htp-vtcm.h"
 #include "hvx-utils.h"
 #include "hex-fastdiv.h"
@@ -41,7 +41,7 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
     const uint32_t end_i   = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
     if (start_i >= end_i) return;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

     uint8_t * spad0_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread;
     uint8_t * spad1_base = octx->src1_spad.data + ith * octx->src1_spad.size_per_thread;
@@ -64,14 +64,14 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
         uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;

         uint32_t src1_width_bytes = current_block_i * sizeof(float);
-        uint8_t * src1_ptr = (uint8_t *)src1->data + i * src1->nb[1];
-        dma_queue_push(q, dma_make_ptr(spad1_base, src1_ptr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
+        const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
+        dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);

         uint32_t src0_row_bytes = src0_ne0 * sizeof(float);
-        uint8_t * src0_ptr = (uint8_t *)src0->data + i * src0->nb[1];
-        dma_queue_push(q, dma_make_ptr(spad0_base, src0_ptr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
+        const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
+        dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);

-        dma_queue_pop(q); // src1
+        dma_queue_pop(dma_q); // src1

         HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);

@@ -87,12 +87,12 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);

-        dma_queue_pop(q); // src0
+        dma_queue_pop(dma_q); // src0

-        uint8_t * dst_ptr = (uint8_t *)dst->data + i * dst->nb[1];
-        dma_queue_push(q, dma_make_ptr(dst_ptr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);
+        const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
+        dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);

-        dma_queue_pop(q);
+        dma_queue_pop(dma_q);
     }
 }

@@ -112,7 +112,7 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
     const uint32_t end_i   = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
     if (start_i >= end_i) return;

-    dma_queue * q = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

     uint8_t * spad0_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread;
     uint8_t * spad1_base = octx->src1_spad.data + ith * octx->src1_spad.size_per_thread;
@@ -135,14 +135,14 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
         uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;

         uint32_t src1_width_bytes = current_block_i * sizeof(__fp16);
-        uint8_t * src1_ptr = (uint8_t *)src1->data + i * src1->nb[1];
-        dma_queue_push(q, dma_make_ptr(spad1_base, src1_ptr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
+        const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
+        dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);

         uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16);
-        uint8_t * src0_ptr = (uint8_t *)src0->data + i * src0->nb[1];
-        dma_queue_push(q, dma_make_ptr(spad0_base, src0_ptr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
+        const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
+        dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);

-        dma_queue_pop(q); // src1
+        dma_queue_pop(dma_q); // src1

         HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);

@@ -158,12 +158,12 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);

-        dma_queue_pop(q); // src0
+        dma_queue_pop(dma_q); // src0

-        uint8_t * dst_ptr = (uint8_t *)dst->data + i * dst->nb[1];
-        dma_queue_push(q, dma_make_ptr(dst_ptr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);
+        const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
+        dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);

-        dma_queue_pop(q);
+        dma_queue_pop(dma_q);
     }
 }

@@ -304,6 +304,10 @@ int op_concat(struct htp_ops_context * octx) {
             worker_func = concat_2d_f16_transposed;
         }
     } else {
+        if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(src1) || htp_tensor_is_extended(dst)) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
+
         const uint32_t total_elements = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
         uint32_t elem_start = 0;
         uint32_t nelems     = total_elements;
diff --git a/ggml/src/ggml-hexagon/htp/cpy-ops.c b/ggml/src/ggml-hexagon/htp/cpy-ops.c
index 490efd687..4453dda3a 100644
--- a/ggml/src/ggml-hexagon/htp/cpy-ops.c
+++ b/ggml/src/ggml-hexagon/htp/cpy-ops.c
@@ -49,6 +49,30 @@ struct htp_copy_context {
     struct fastdiv_values div_ne02_ne01_ne00;
 };

+static inline void cpy_dma_sametype_reshape_contig(
+    dma_queue * dma_q,
+    dma_addr_t dst,
+    dma_addr_t src0,
+    uint32_t total_bytes
+) {
+    if (total_bytes == 0) {
+        return;
+    }
+
+    const uint32_t max_chunk = DMA_SAFE_CHUNK_SIZE;
+    while (total_bytes > 0) {
+        const uint32_t chunk = MIN(total_bytes, max_chunk);
+        if (!dma_queue_push(dma_q, dma_make_data(dst, src0), chunk, chunk, chunk, /*nrows=*/ 1)) {
+            dma_queue_flush(dma_q);
+            dma_queue_push(dma_q, dma_make_data(dst, src0), chunk, chunk, chunk, /*nrows=*/ 1);
+        }
+        dst         += chunk;
+        src0        += chunk;
+        total_bytes -= chunk;
+    }
+    dma_queue_flush(dma_q);
+}
+
 #define cpy_preamble                              \
     const struct htp_tensor *src0 = octx->src[0]; \
     const struct htp_tensor *dst  = octx->dst;    \
@@ -73,129 +97,131 @@ struct htp_copy_context {
     const uint32_t  nb2 = dst->nb[2];             \
     const uint32_t  nb3 = dst->nb[3];

-#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                       \
-static void cpy_thread_##NAME##_sameshape(unsigned int nth, unsigned int ith, void * data) {   \
-    struct htp_copy_context * ct = (struct htp_copy_context *) data;                           \
-    struct htp_ops_context * octx = ct->octx;                                                  \
-    cpy_preamble;                                                                              \
-    const uint32_t dr  = ct->src0_nrows_per_thread;                                            \
-    const uint32_t ir0 = ct->row_start + dr * ith;                                             \
-    const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);                             \
-    if (ir0 >= ir1) return;                                                                    \
-    const bool contiguous = (nb01 == ne00 * ELEM_SIZE) && (nb1 == nb01) &&                     \
-                            (nb02 == ne01 * nb01)      && (nb2 == nb02) &&                     \
-                            (nb03 == ne02 * nb02)      && (nb3 == nb03);                       \
-    const uint32_t ne02_ne01 = ne02 * ne01;                                                    \
-    uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);                                           \
-    uint32_t rem = ir0 - i03 * ne02_ne01;                                                      \
-    uint32_t i02 = fastdiv(rem, &ct->div_ne01);                                                \
-    uint32_t i01 = rem - i02 * ne01;                                                           \
-    uint8_t * dst_ptr  = (uint8_t *) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;               \
-    uint8_t * src0_ptr = (uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;              \
-    if (contiguous) {                                                                          \
-        hvx_copy_uu(dst_ptr, src0_ptr, (ir1 - ir0) * ne00, ELEM_SIZE);                         \
-        return;                                                                                \
-    }                                                                                          \
-    for (uint32_t r = ir0; r < ir1; r++) {                                                     \
-        hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2);                                      \
-        hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE);                                       \
-        dst_ptr  += nb1;                                                                       \
-        src0_ptr += nb01;                                                                      \
-        if (++i01 == ne01) {                                                                   \
-            i01 = 0;                                                                           \
-            if (++i02 == ne02) {                                                               \
-                i02 = 0;                                                                       \
-                i03++;                                                                         \
-            }                                                                                  \
-            dst_ptr  = (uint8_t *) dst->data  + i02*nb2  + i03*nb3;                            \
-            src0_ptr = (uint8_t *) src0->data + i02*nb02 + i03*nb03;                           \
-        }                                                                                      \
-    }                                                                                          \
+#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                                           \
+static void cpy_thread_##NAME##_sameshape(unsigned int nth, unsigned int ith, void * data) {                       \
+    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                               \
+    struct htp_ops_context * octx = ct->octx;                                                                      \
+    cpy_preamble;                                                                                                  \
+    const uint32_t dr  = ct->src0_nrows_per_thread;                                                                \
+    const uint32_t ir0 = ct->row_start + dr * ith;                                                                 \
+    const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);                                                 \
+    if (ir0 >= ir1) return;                                                                                        \
+    const bool contiguous = htp_tensor_is_contiguous(src0, ELEM_SIZE) && htp_tensor_is_contiguous(dst, ELEM_SIZE); \
+    if (contiguous) {                                                                                              \
+        dma_queue * dma_q = octx->ctx->dma[ith];                                                                   \
+        dma_addr_t dst_addr  = dst->data  + ir0 * ne00 * ELEM_SIZE;                                                \
+        dma_addr_t src0_addr = src0->data + ir0 * ne00 * ELEM_SIZE;                                                \
+        cpy_dma_sametype_reshape_contig(dma_q, dst_addr, src0_addr, (ir1 - ir0) * ne00 * ELEM_SIZE);               \
+        return;                                                                                                    \
+    }                                                                                                              \
+    const uint32_t ne02_ne01 = ne02 * ne01;                                                                        \
+    uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);                                                               \
+    uint32_t rem = ir0 - i03 * ne02_ne01;                                                                          \
+    uint32_t i02 = fastdiv(rem, &ct->div_ne01);                                                                    \
+    uint32_t i01 = rem - i02 * ne01;                                                                               \
+    uint8_t * dst_ptr  = (uint8_t *) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;                                   \
+    uint8_t * src0_ptr = (uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;                                  \
+    for (uint32_t r = ir0; r < ir1; r++) {                                                                         \
+        hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2);                                                          \
+        hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE);                                                           \
+        dst_ptr  += nb1;                                                                                           \
+        src0_ptr += nb01;                                                                                          \
+        if (++i01 == ne01) {                                                                                       \
+            i01 = 0;                                                                                               \
+            if (++i02 == ne02) {                                                                                   \
+                i02 = 0;                                                                                           \
+                i03++;                                                                                             \
+            }                                                                                                      \
+            dst_ptr  = (uint8_t *) dst->data  + i02*nb2  + i03*nb3;                                                \
+            src0_ptr = (uint8_t *) src0->data + i02*nb02 + i03*nb03;                                               \
+        }                                                                                                          \
+    }                                                                                                              \
 }

 DEFINE_CPY_SAMESHAPE(f32,  float, 4)
 DEFINE_CPY_SAMESHAPE(f16, __fp16, 2)

-#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                               \
-static void cpy_thread_##NAME##_reshape(unsigned int nth, unsigned int ith, void * data) {           \
-    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                 \
-    struct htp_ops_context * octx = ct->octx;                                                        \
-    cpy_preamble;                                                                                    \
-    const uint32_t th_nelem = ct->elem_per_thread;                                                   \
-    const uint32_t th_start = ct->elem_start + ith * th_nelem;                                       \
-    const uint32_t th_end   = MIN(th_start + th_nelem, ct->elem_start + ct->nelem);                  \
-    if (th_start >= th_end) return;                                                                  \
-                                                                                                     \
-    if (htp_tensor_is_contiguous(src0, ELEM_SIZE) && htp_tensor_is_contiguous(dst, ELEM_SIZE)) {     \
-        hvx_copy_uu((uint8_t *) dst->data + (size_t) th_start * ELEM_SIZE,                           \
-                    (const uint8_t *) src0->data + (size_t) th_start * ELEM_SIZE,                    \
-                    th_end - th_start, ELEM_SIZE);                                                   \
-        return;                                                                                      \
-    }                                                                                                \
-                                                                                                     \
-    const uint32_t ne01_ne00      = ne01 * ne00;                                                     \
-    const uint32_t ne02_ne01_ne00 = ne02 * ne01_ne00;                                                \
-    const uint32_t ne1_ne0        = ne1 * ne0;                                                       \
-    const uint32_t ne2_ne1_ne0    = ne2 * ne1_ne0;                                                   \
-                                                                                                     \
-    uint32_t e = th_start;                                                                           \
-    uint32_t i13 = fastdiv(e, &ct->div_ne2_ne1_ne0);                                                 \
-    uint32_t rem = e - i13 * ne2_ne1_ne0;                                                            \
-    uint32_t i12 = fastdiv(rem, &ct->div_ne1_ne0);                                                   \
-    uint32_t rem2 = rem - i12 * ne1_ne0;                                                             \
-    uint32_t i11 = fastdiv(rem2, &ct->div_ne0);                                                      \
-    uint32_t i10 = rem2 - i11 * ne0;                                                                 \
-                                                                                                     \
-    uint32_t i03 = fastdiv(e, &ct->div_ne02_ne01_ne00);                                              \
-    uint32_t rem_s = e - i03 * ne02_ne01_ne00;                                                       \
-    uint32_t i02 = fastdiv(rem_s, &ct->div_ne01_ne00);                                               \
-    uint32_t rem2_s = rem_s - i02 * ne01_ne00;                                                       \
-    uint32_t i01 = fastdiv(rem2_s, &ct->div_ne00);                                                   \
-    uint32_t i00 = rem2_s - i01 * ne00;                                                              \
-                                                                                                     \
-    char * dst_ptr        = (char *)       dst->data  + i10*nb0  + i11*nb1  + i12*nb2  + i13*nb3;    \
-    const char * src0_ptr = (const char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03;   \
-                                                                                                     \
-    const bool rows_contig = (nb00 == ELEM_SIZE) && (nb0 == ELEM_SIZE);                              \
-                                                                                                     \
-    while (e < th_end) {                                                                             \
-        uint32_t run = 1;                                                                            \
-        if (rows_contig) {                                                                           \
-            run = MIN(MIN(ne00 - i00, ne0 - i10), th_end - e);                                       \
-            hvx_copy_uu((uint8_t *) dst_ptr, (const uint8_t *) src0_ptr, run, ELEM_SIZE);            \
-        } else {                                                                                     \
-            *((ELEM_TYPE *) dst_ptr) = *((const ELEM_TYPE *) src0_ptr);                              \
-        }                                                                                            \
-        e += run;                                                                                    \
-                                                                                                     \
-        dst_ptr += run * nb0;                                                                        \
-        i10     += run;                                                                              \
-        if (i10 == ne0) {                                                                            \
-            i10 = 0;                                                                                 \
-            if (++i11 == ne1) {                                                                      \
-                i11 = 0;                                                                             \
-                if (++i12 == ne2) {                                                                  \
-                    i12 = 0;                                                                         \
-                    i13++;                                                                           \
-                }                                                                                    \
-            }                                                                                        \
-            dst_ptr = (char *) dst->data + i11*nb1 + i12*nb2 + i13*nb3;                              \
-        }                                                                                            \
-                                                                                                     \
-        src0_ptr += run * nb00;                                                                      \
-        i00      += run;                                                                             \
-        if (i00 == ne00) {                                                                           \
-            i00 = 0;                                                                                 \
-            if (++i01 == ne01) {                                                                     \
-                i01 = 0;                                                                             \
-                if (++i02 == ne02) {                                                                 \
-                    i02 = 0;                                                                         \
-                    i03++;                                                                           \
-                }                                                                                    \
-            }                                                                                        \
-            src0_ptr = (const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;                   \
-        }                                                                                            \
-    }                                                                                                \
+#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                                \
+static void cpy_thread_##NAME##_reshape(unsigned int nth, unsigned int ith, void * data) {            \
+    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                  \
+    struct htp_ops_context * octx = ct->octx;                                                         \
+    cpy_preamble;                                                                                     \
+    const uint32_t th_nelem = ct->elem_per_thread;                                                    \
+    const uint32_t th_start = ct->elem_start + ith * th_nelem;                                        \
+    const uint32_t th_end   = MIN(th_start + th_nelem, ct->elem_start + ct->nelem);                   \
+    if (th_start >= th_end) return;                                                                   \
+                                                                                                      \
+    if (htp_tensor_is_contiguous(src0, ELEM_SIZE) && htp_tensor_is_contiguous(dst, ELEM_SIZE)) {      \
+        dma_queue * dma_q = octx->ctx->dma[ith];                                                      \
+        dma_addr_t dst_addr  = dst->data  + th_start * ELEM_SIZE;                                     \
+        dma_addr_t src0_addr = src0->data + th_start * ELEM_SIZE;                                     \
+        cpy_dma_sametype_reshape_contig(dma_q, dst_addr, src0_addr, (th_end - th_start) * ELEM_SIZE); \
+        return;                                                                                       \
+    }                                                                                                 \
+                                                                                                      \
+    const uint32_t ne01_ne00      = ne01 * ne00;                                                      \
+    const uint32_t ne02_ne01_ne00 = ne02 * ne01_ne00;                                                 \
+    const uint32_t ne1_ne0        = ne1 * ne0;                                                        \
+    const uint32_t ne2_ne1_ne0    = ne2 * ne1_ne0;                                                    \
+                                                                                                      \
+    uint32_t e = th_start;                                                                            \
+    uint32_t i13 = fastdiv(e, &ct->div_ne2_ne1_ne0);                                                  \
+    uint32_t rem = e - i13 * ne2_ne1_ne0;                                                             \
+    uint32_t i12 = fastdiv(rem, &ct->div_ne1_ne0);                                                    \
+    uint32_t rem2 = rem - i12 * ne1_ne0;                                                              \
+    uint32_t i11 = fastdiv(rem2, &ct->div_ne0);                                                       \
+    uint32_t i10 = rem2 - i11 * ne0;                                                                  \
+                                                                                                      \
+    uint32_t i03 = fastdiv(e, &ct->div_ne02_ne01_ne00);                                               \
+    uint32_t rem_s = e - i03 * ne02_ne01_ne00;                                                        \
+    uint32_t i02 = fastdiv(rem_s, &ct->div_ne01_ne00);                                                \
+    uint32_t rem2_s = rem_s - i02 * ne01_ne00;                                                        \
+    uint32_t i01 = fastdiv(rem2_s, &ct->div_ne00);                                                    \
+    uint32_t i00 = rem2_s - i01 * ne00;                                                               \
+                                                                                                      \
+    char * dst_ptr        = (char *)       dst->data  + i10*nb0  + i11*nb1  + i12*nb2  + i13*nb3;     \
+    const char * src0_ptr = (const char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03;    \
+                                                                                                      \
+    const bool rows_contig = (nb00 == ELEM_SIZE) && (nb0 == ELEM_SIZE);                               \
+                                                                                                      \
+    while (e < th_end) {                                                                              \
+        uint32_t run = 1;                                                                             \
+        if (rows_contig) {                                                                            \
+            run = MIN(MIN(ne00 - i00, ne0 - i10), th_end - e);                                        \
+            hvx_copy_uu((uint8_t *) dst_ptr, (const uint8_t *) src0_ptr, run, ELEM_SIZE);             \
+        } else {                                                                                      \
+            *((ELEM_TYPE *) dst_ptr) = *((const ELEM_TYPE *) src0_ptr);                               \
+        }                                                                                             \
+        e += run;                                                                                     \
+                                                                                                      \
+        dst_ptr += run * nb0;                                                                         \
+        i10     += run;                                                                               \
+        if (i10 == ne0) {                                                                             \
+            i10 = 0;                                                                                  \
+            if (++i11 == ne1) {                                                                       \
+                i11 = 0;                                                                              \
+                if (++i12 == ne2) {                                                                   \
+                    i12 = 0;                                                                          \
+                    i13++;                                                                            \
+                }                                                                                     \
+            }                                                                                         \
+            dst_ptr = (char *) dst->data + i11*nb1 + i12*nb2 + i13*nb3;                               \
+        }                                                                                             \
+                                                                                                      \
+        src0_ptr += run * nb00;                                                                       \
+        i00      += run;                                                                              \
+        if (i00 == ne00) {                                                                            \
+            i00 = 0;                                                                                  \
+            if (++i01 == ne01) {                                                                      \
+                i01 = 0;                                                                              \
+                if (++i02 == ne02) {                                                                  \
+                    i02 = 0;                                                                          \
+                    i03++;                                                                            \
+                }                                                                                     \
+            }                                                                                         \
+            src0_ptr = (const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;                    \
+        }                                                                                             \
+    }                                                                                                 \
 }

 DEFINE_CPY_RESHAPE(f32,  float, 4)
@@ -273,6 +299,27 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi
     }
 }

+static inline void cpy_dma_push_2d_chunked(
+    dma_queue * dma_q,
+    dma_addr_t  dst,
+    dma_addr_t  src,
+    size_t      dst_stride,
+    size_t      src_stride,
+    size_t      row_size,
+    uint32_t    nrows
+) {
+    while (nrows > 0) {
+        const uint32_t cur_rows = MIN(nrows, DMA_MAX_NROWS);
+        if (!dma_queue_push(dma_q, dma_make_data(dst, src), dst_stride, src_stride, row_size, cur_rows)) {
+            dma_queue_flush(dma_q);
+            dma_queue_push(dma_q, dma_make_data(dst, src), dst_stride, src_stride, row_size, cur_rows);
+        }
+        dst   += cur_rows * dst_stride;
+        src   += cur_rows * src_stride;
+        nrows -= cur_rows;
+    }
+}
+
 static inline void cpy_dma_sametype_sameshape(
     struct htp_ops_context * octx,
     const struct htp_tensor * dst,
@@ -282,46 +329,35 @@ static inline void cpy_dma_sametype_sameshape(
     uint32_t nb01, uint32_t nb02, uint32_t nb03,
     uint32_t  nb1, uint32_t  nb2, uint32_t nb3
 ) {
+    const bool contiguous = htp_tensor_is_contiguous(src0, elem_size) && htp_tensor_is_contiguous(dst, elem_size);
+
+    dma_queue * dma_q = octx->ctx->dma[0];
+
+    if (contiguous) {
+        cpy_dma_sametype_reshape_contig(dma_q, dst->data, src0->data, ne00 * elem_size * ne01 * ne02 * ne03);
+        return;
+    }
+
     const bool contiguous_outer =
         (ne02 == 1 || (nb02 == ne01 * nb01 && nb2 == ne01 * nb1)) &&
         (ne03 == 1 || (nb03 == ne02 * nb02 && nb3 == ne02 * nb2));

-    dma_queue * q = octx->ctx->dma[0];
-
     if (contiguous_outer) {
-        if (!dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03)) {
-            dma_queue_flush(q);
-            dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03);
-        }
-        dma_queue_flush(q);
+        uint32_t total_rows = ne01 * ne02 * ne03;
+        cpy_dma_push_2d_chunked(dma_q, dst->data, src0->data, nb1, nb01, ne00 * elem_size, total_rows);
+        dma_queue_flush(dma_q);
         return;
     }

     for (uint32_t i03 = 0; i03 < ne03; i03++) {
         for (uint32_t i02 = 0; i02 < ne02; i02++) {
-            uint8_t * dst_ptr  = (uint8_t *) dst->data  + i02 * nb2  + i03 * nb3;
-            uint8_t * src0_ptr = (uint8_t *) src0->data + i02 * nb02 + i03 * nb03;
-
-            if (!dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01)) {
-                dma_queue_flush(q);
-                dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01);
-            }
+            dma_addr_t dst_data  = dst->data  + i02 * nb2  + i03 * nb3;
+            dma_addr_t src0_data = src0->data + i02 * nb02 + i03 * nb03;
+            cpy_dma_push_2d_chunked(dma_q, dst_data, src0_data, nb1, nb01, ne00 * elem_size, ne01);
         }
     }

-    dma_queue_flush(q);
-}
-
-static inline void cpy_dma_sametype_reshape_contig(
-    struct htp_ops_context * octx,
-    const struct htp_tensor * dst,
-    const struct htp_tensor * src0,
-    uint32_t total_bytes
-) {
-    dma_queue * q = octx->ctx->dma[0];
-    dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data),
-                   total_bytes, total_bytes, total_bytes, /*nrows=*/ 1);
-    dma_queue_pop(q);
+    dma_queue_flush(dma_q);
 }

 static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
@@ -345,10 +381,6 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const bool sametype   = (src0->type == dst->type);
     const bool transposed = (nb00 > nb01) || (nb0 > nb1) ||
                             (nb00 != ct.src0_type_size) || (nb0 != ct.dst_type_size) ||
@@ -360,6 +392,15 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
     const bool src_is_contiguous = htp_tensor_is_contiguous(src0, ct.src0_type_size);
     const bool dst_is_contiguous = htp_tensor_is_contiguous(dst, ct.dst_type_size);

+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+        if (!sametype) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
+        if (!sameshape && !(src_is_contiguous && dst_is_contiguous && octx->ctx->mdev.count <= 1)) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
+    }
+
     if (sameshape) {
         const uint32_t total_rows = ne01 * ne02 * ne03;
         const uint32_t row_size   = ne00 * ct.dst_type_size;
@@ -373,7 +414,8 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
         if (octx->ctx->mdev.count > 1) {
             const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
             const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
-            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, can_split ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, can_split ? rows_per_chunk : 0,
+                                                               octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
             row_start = range.start;
             nrows     = range.count;
         }
@@ -386,9 +428,11 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
         ct.nrows     = nrows;
         ct.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);

-        if (sametype && octx->ctx->mdev.count <= 1) {
-            *use_dma = true;
-            cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
+        if (sametype && (octx->ctx->mdev.count <= 1 || htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst))) {
+            if (octx->ctx->mdev.idx == 0) {
+                *use_dma = true;
+                cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
+            }
         } else {
             work_queue_func_t copy_fun = NULL;
             if (sametype) {
@@ -408,7 +452,7 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {

         if (octx->ctx->mdev.count <= 1 && dst_is_contiguous && src_is_contiguous) {
             *use_dma = true;
-            cpy_dma_sametype_reshape_contig(octx, dst, src0, total_elems * ct.dst_type_size);
+            cpy_dma_sametype_reshape_contig(octx->ctx->dma[0], dst->data, src0->data, total_elems * ct.dst_type_size);
             return HTP_STATUS_OK;
         }

@@ -424,7 +468,8 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {

         if (octx->ctx->mdev.count > 1) {
             const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
-            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_elems, can_split ? elems_per_line : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_elems, can_split ? elems_per_line : 0,
+                                                               octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
             elem_start = range.start;
             nelem      = range.count;
         }
@@ -461,6 +506,9 @@ int op_cpy(struct htp_ops_context * octx) {

         if (octx->ctx->mdev.idx == 0) {
             const struct htp_tensor * sync = octx->src[1];
+            if (htp_tensor_is_extended(sync)) {
+                return HTP_STATUS_NO_SUPPORT;
+            }
             const uint32_t seq = (uint32_t) octx->op_params[0];
             atomic_uint * sync_fence = (atomic_uint *) (uintptr_t) sync->data;
             htp_fence_write(sync_fence, seq, octx->status);
diff --git a/ggml/src/ggml-hexagon/htp/cumsum-ops.c b/ggml/src/ggml-hexagon/htp/cumsum-ops.c
index 971fa3bcc..eaab7d7e5 100644
--- a/ggml/src/ggml-hexagon/htp/cumsum-ops.c
+++ b/ggml/src/ggml-hexagon/htp/cumsum-ops.c
@@ -14,7 +14,7 @@
 #include "htp-tensor.h"
 #include "hvx-types.h"
 #include "hvx-utils.h"
-#include "hex-dma.h"
+#include "dma-queue.h"

 #define htp_cumsum_tensors_preamble                         \
     const struct htp_tensor * restrict src0 = octx->src[0]; \
@@ -55,7 +55,7 @@ struct htp_cumsum_context {
     struct htp_cumsum_context * cctx = (struct htp_cumsum_context *) data; \
     struct htp_ops_context *    octx = cctx->octx;                         \
     htp_cumsum_tensors_preamble;                                           \
-    dma_queue * dma_queue = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

 // ---------------------------------------------------------------------------
 // HVX prefix scan helpers
@@ -131,47 +131,47 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
     const size_t src_row_size_aligned = cctx->src_row_size_aligned;
     const size_t dst_row_size_aligned = cctx->dst_row_size_aligned;

-    const uint8_t * src_data = (const uint8_t *) src0->data;
-    uint8_t *       dst_data = (uint8_t *) dst->data;
+    const dma_addr_t src_data = src0->data;
+    const dma_addr_t dst_data = dst->data;

     uint8_t * src_spad = octx->src0_spad.data + (ith * src_row_size_aligned * 2);
     uint8_t * dst_spad = octx->dst_spad.data  + (ith * dst_row_size_aligned * 2);

     for (uint32_t ir = ir0, spad_idx = 0; ir < ir1 && spad_idx < 2; ir++, spad_idx++) {
         // Dummy dst writeback to establish queue ordering
-        dma_queue_push_vtcm_to_ddr(dma_queue,
-                                   dma_make_ptr(dst_data, dst_spad + (spad_idx * dst_row_size_aligned)),
-                                   dst_row_size, dst_row_size_aligned, 0);
-
-        dma_queue_push_ddr_to_vtcm(dma_queue,
-                                   dma_make_ptr(src_spad + (spad_idx * src_row_size_aligned),
-                                                src_data + (ir * src_row_size)),
-                                   src_row_size_aligned, src_row_size, 1);
+        dma_queue_push(dma_q,
+                       dma_make_data(dst_data, dst_spad + (spad_idx * dst_row_size_aligned)),
+                       dst_row_size, dst_row_size_aligned, dst_row_size, 0);
+
+        dma_queue_push(dma_q,
+                       dma_make_data(src_spad + (spad_idx * src_row_size_aligned),
+                                     src_data + (ir * src_row_size)),
+                       src_row_size_aligned, src_row_size, src_row_size, 1);
     }

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

     for (uint32_t ir = ir0; ir < ir1; ir++) {
-        float * dst_spad_row = (float *) dma_queue_pop(dma_queue).src;
-        float * src_spad_row = (float *) dma_queue_pop(dma_queue).dst;
+        float * dst_spad_row = (float *) dma_queue_pop(dma_q).src;
+        float * src_spad_row = (float *) dma_queue_pop(dma_q).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         hvx_cumsum_row_f32(src_spad_row, dst_spad_row, ne00);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        dma_queue_push_vtcm_to_ddr(dma_queue,
-                                   dma_make_ptr(dst_data + (ir * dst_row_size), (uint8_t *) dst_spad_row),
-                                   dst_row_size, dst_row_size_aligned, 1);
+        dma_queue_push(dma_q,
+                       dma_make_data(dst_data + (ir * dst_row_size), dst_spad_row),
+                       dst_row_size, dst_row_size_aligned, dst_row_size, 1);

         const uint32_t next_row = ir + 2;
         if (next_row < ir1) {
-            dma_queue_push_ddr_to_vtcm(dma_queue,
-                                       dma_make_ptr((uint8_t *) src_spad_row, src_data + (next_row * src_row_size)),
-                                       src_row_size_aligned, src_row_size, 1);
+            dma_queue_push(dma_q,
+                           dma_make_data(src_spad_row, src_data + (next_row * src_row_size)),
+                           src_row_size_aligned, src_row_size, src_row_size, 1);
         }
     }

-    dma_queue_flush(dma_queue);
+    dma_queue_flush(dma_q);

     FARF(HIGH, "cumsum-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
@@ -211,10 +211,6 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
     const struct htp_tensor * src0 = octx->src[0];
     const struct htp_tensor * dst  = octx->dst;

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const uint32_t total_rows      = src0->ne[1] * src0->ne[2] * src0->ne[3];
     const size_t dst_data_row_size = dst->ne[0] * sizeof(float);

@@ -264,6 +260,9 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
     };

     if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
+        if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
         work_queue_run(octx->ctx->work_queue, cumsum_thread_f32, &cctx, n_threads);
     } else {
         work_queue_run(octx->ctx->work_queue, cumsum_thread_f32_dma, &cctx, n_threads);
diff --git a/ggml/src/ggml-hexagon/htp/diag-ops.c b/ggml/src/ggml-hexagon/htp/diag-ops.c
index a69fd89d3..162214d3e 100644
--- a/ggml/src/ggml-hexagon/htp/diag-ops.c
+++ b/ggml/src/ggml-hexagon/htp/diag-ops.c
@@ -13,7 +13,7 @@
 #include "hvx-types.h"
 #include "hex-utils.h"
 #include "hvx-copy.h"
-#include "hex-dma.h"
+#include "dma-queue.h"

 #define htp_diag_tensors_preamble                           \
     const struct htp_tensor * restrict src0 = octx->src[0]; \
@@ -59,7 +59,7 @@ static inline void hvx_diag_row_f32(const float * restrict src, float * restrict

 static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data) {
     htp_diag_preamble;
-    dma_queue * dma_queue = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

     const uint32_t ib0 = dctx->batch_start + dctx->batches_per_thread * ith;
     const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->batch_start + dctx->total_batches);
@@ -73,8 +73,8 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
     const size_t src_batch_size_aligned = dctx->src_batch_size_aligned;
     const size_t dst_row_size_aligned   = dctx->dst_row_size_aligned;

-    const uint8_t * src_data = (const uint8_t *) src0->data;
-    uint8_t *       dst_data = (uint8_t *) dst->data;
+    const dma_addr_t src_data = src0->data;
+    const dma_addr_t dst_data = dst->data;

     // 1 src buffer + 1 dst row buffer per thread in VTCM
     uint8_t * src_spad = octx->src0_spad.data + (ith * src_batch_size_aligned);
@@ -86,13 +86,13 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
         const uint32_t i3 = ib / ne02;
         const uint32_t i2 = ib % ne02;

-        const uint8_t * src_batch = src_data + i3 * nb03 + i2 * nb02;
+        const dma_addr_t src_batch = src_data + i3 * nb03 + i2 * nb02;

         // Fetch source vector into VTCM
-        dma_queue_push_ddr_to_vtcm(dma_queue,
-                                   dma_make_ptr(src_spad, src_batch),
-                                   src_batch_size_aligned, src_batch_size, 1);
-        dma_queue_flush(dma_queue);
+        dma_queue_push(dma_q,
+                       dma_make_data(src_spad, src_batch),
+                       src_batch_size_aligned, src_batch_size, src_batch_size, 1);
+        dma_queue_flush(dma_q);

         const float * src_spad_f32 = (const float *) src_spad;
         float       * dst_spad_f32 = (float *) dst_spad;
@@ -104,11 +104,11 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (ib * ne1 + i1));

             // Write completed row back to DDR
-            uint8_t * dst_row = dst_data + i3 * nb3 + i2 * nb2 + i1 * nb1;
-            dma_queue_push_vtcm_to_ddr(dma_queue,
-                                       dma_make_ptr(dst_row, dst_spad),
-                                       dst_row_size, dst_row_size_aligned, 1);
-            dma_queue_flush(dma_queue);
+            const dma_addr_t dst_row = dst_data + i3 * nb3 + i2 * nb2 + i1 * nb1;
+            dma_queue_push(dma_q,
+                           dma_make_data(dst_row, dst_spad),
+                           dst_row_size, dst_row_size_aligned, dst_row_size, 1);
+            dma_queue_flush(dma_q);
         }
     }

@@ -156,10 +156,6 @@ int op_diag_f32(struct htp_ops_context * octx) {
     const struct htp_tensor * src0 = octx->src[0];
     const struct htp_tensor * dst  = octx->dst;

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const uint32_t total_batches = src0->ne[2] * src0->ne[3];
     const size_t dst_batch_size  = dst->ne[1] * dst->nb[1];

@@ -221,6 +217,9 @@ int op_diag_f32(struct htp_ops_context * octx) {
     };

     if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
+        if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
         work_queue_run(octx->ctx->work_queue, diag_thread_f32, &dctx, n_threads);
     } else {
         work_queue_run(octx->ctx->work_queue, diag_thread_f32_dma, &dctx, n_threads);
diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.c b/ggml/src/ggml-hexagon/htp/dma-queue.c
index 4beded1de..464e4b849 100644
--- a/ggml/src/ggml-hexagon/htp/dma-queue.c
+++ b/ggml/src/ggml-hexagon/htp/dma-queue.c
@@ -22,58 +22,69 @@ static inline uintptr_t align_up(uintptr_t addr, size_t align) {
     return (addr + align - 1) & ~(align - 1);
 }

-size_t dma_queue_sizeof(size_t capacity) {
+static inline size_t dma_ring_sizeof(size_t capacity) {
     capacity = pow2_ceil(capacity);

-    size_t size_q      = sizeof(dma_queue);
-    size_t offset_r    = align_up(size_q, HEX_L2_LINE_SIZE);
     size_t size_r      = sizeof(dma_ring);
-    size_t offset_desc = align_up(offset_r + size_r, HEX_L2_LINE_SIZE);
+    size_t offset_desc = align_up(size_r, HEX_L2_LINE_SIZE);
     size_t size_desc   = capacity * sizeof(dma_descriptor_2d);
-    size_t offset_dptr = align_up(offset_desc + size_desc, HEX_L2_LINE_SIZE);
-    size_t size_dptr   = capacity * sizeof(dma_ptr);
+    size_t offset_data = align_up(offset_desc + size_desc, HEX_L2_LINE_SIZE);
+    size_t size_data   = capacity * sizeof(dma_data);

-    return offset_dptr + size_dptr;
+    return offset_data + size_data;
 }

-size_t dma_queue_alignof(void) {
-    return HEX_L2_LINE_SIZE;
-}
-
-dma_queue_t dma_queue_init(void * ptr, size_t capacity, uintptr_t vtcm_base, size_t vtcm_size, struct htp_thread_trace * trace) {
+static inline dma_ring * dma_ring_init(void * ptr, size_t capacity, struct htp_thread_trace * trace) {
     capacity = pow2_ceil(capacity);

-    size_t size_q      = sizeof(dma_queue);
-    size_t offset_r    = align_up(size_q, HEX_L2_LINE_SIZE);
     size_t size_r      = sizeof(dma_ring);
-    size_t offset_desc = align_up(offset_r + size_r, HEX_L2_LINE_SIZE);
+    size_t offset_desc = align_up(size_r, HEX_L2_LINE_SIZE);
     size_t size_desc   = capacity * sizeof(dma_descriptor_2d);
-    size_t offset_dptr = align_up(offset_desc + size_desc, HEX_L2_LINE_SIZE);
-    size_t size_dptr   = capacity * sizeof(dma_ptr);
-
-    size_t total_size = offset_dptr + size_dptr;
-    memset(ptr, 0, total_size);
-
-    dma_queue * q = (dma_queue *) ptr;
-    dma_ring * r  = (dma_ring *) ((uintptr_t) ptr + offset_r);
-
-    q->ring    = r;
-    q->nocache = 0;
-    q->alias   = false;
+    size_t offset_data = align_up(offset_desc + size_desc, HEX_L2_LINE_SIZE);

+    dma_ring * r = (dma_ring *) ptr;
     r->trace     = trace;
-    r->vtcm_base = vtcm_base;
-    r->vtcm_end  = vtcm_base + vtcm_size;
     r->capacity  = capacity;
     r->idx_mask  = capacity - 1;
     r->push_idx  = 0;
     r->pop_idx   = 0;
+    r->desc      = (dma_descriptor_2d *) ((uintptr_t) ptr + offset_desc);
+    r->data      = (dma_data *) ((uintptr_t) ptr + offset_data);
+    r->tail      = &r->desc[capacity - 1];
+
+    return r;
+}
+
+size_t dma_queue_sizeof(size_t capacity) {
+    size_t size_q    = sizeof(dma_queue);
+    size_t offset_r0 = align_up(size_q, HEX_L2_LINE_SIZE);
+    size_t size_r0   = dma_ring_sizeof(capacity);
+    size_t offset_r1 = align_up(offset_r0 + size_r0, HEX_L2_LINE_SIZE);
+    size_t size_r1   = dma_ring_sizeof(DMA_FALLBACK_CAPACITY);
+
+    return offset_r1 + size_r1;
+}
+
+size_t dma_queue_alignof(void) {
+    return HEX_L2_LINE_SIZE;
+}
+
+dma_queue_t dma_queue_init(void * ptr, size_t capacity, struct htp_thread_trace * trace) {
+    size_t total_size = dma_queue_sizeof(capacity);
+    memset(ptr, 0, total_size);
+
+    dma_queue * q = (dma_queue *) ptr;
+
+    size_t size_q    = sizeof(dma_queue);
+    size_t offset_r0 = align_up(size_q, HEX_L2_LINE_SIZE);
+    size_t size_r0   = dma_ring_sizeof(capacity);
+    size_t offset_r1 = align_up(offset_r0 + size_r0, HEX_L2_LINE_SIZE);

-    r->desc = (dma_descriptor_2d *) ((uintptr_t) ptr + offset_desc);
-    r->dptr = (dma_ptr *) ((uintptr_t) ptr + offset_dptr);
-    r->tail = &r->desc[capacity - 1];
+    q->ring0 = dma_ring_init((void *) ((uintptr_t) ptr + offset_r0), capacity, trace);
+    q->ring1 = dma_ring_init((void *) ((uintptr_t) ptr + offset_r1), DMA_FALLBACK_CAPACITY, trace);
+    q->alias = false;

-    FARF(HIGH, "dma-queue: capacity %u, unified memory size %zu\n", capacity, total_size);
+    FARF(HIGH, "dma-queue: capacity %u, unified memory size %zu\n", (unsigned) capacity, total_size);

     return q;
 }
@@ -86,13 +97,13 @@ size_t dma_queue_alias_sizeof(void) {
     return sizeof(dma_queue);
 }

-dma_queue_t dma_queue_alias_init(void * ptr, dma_queue_t main_q, uint8_t nocache) {
+dma_queue_t dma_queue_alias_init(void * ptr, dma_queue_t main_q) {
     dma_queue * q = (dma_queue *) ptr;
     memset(q, 0, sizeof(dma_queue));

-    q->ring    = main_q->ring;
-    q->nocache = nocache;
-    q->alias   = true;
+    q->ring0 = main_q->ring0;
+    q->ring1 = main_q->ring1;
+    q->alias = true;

     return q;
 }
@@ -101,4 +112,101 @@ void dma_queue_alias_free(dma_queue_t q) {
     (void) q;
 }

+bool dma_queue_push_fallback_2d(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
+    dma_ring * r0 = q->ring0;
+    dma_ring * r1 = q->ring1;
+
+    if (((r0->push_idx + 1) & r0->idx_mask) == r0->pop_idx) {
+        return false;
+    }
+
+    r1->tail = r0->tail;
+
+    size_t rem_rows    = nrows;
+    dma_addr_t cur_dst = ddata.dst;
+    dma_addr_t cur_src = ddata.src;
+
+    while (rem_rows > 0) {
+        const uint32_t cur_rows = MIN(rem_rows, DMA_MAX_NROWS);
+        dma_data cur_data = dma_make_data(cur_dst, cur_src);
+        if (!dma_ring_push_single_2d(r1, cur_data, dst_stride, src_stride, row_size, cur_rows)) {
+            dma_ring_flush(r1);
+            dma_ring_push_single_2d(r1, cur_data, dst_stride, src_stride, row_size, cur_rows);
+        }
+        cur_dst  += cur_rows * dst_stride;
+        cur_src  += cur_rows * src_stride;
+        rem_rows -= cur_rows;
+    }
+
+    dma_ring_flush(r1);
+    r0->tail = r1->tail;
+
+    return dma_ring_push_single_2d(r0, ddata, 0, 0, 0, /*nrows=*/ 0);
+}
+
+bool dma_queue_push_fallback_contig(dma_queue * q, dma_data ddata, size_t total) {
+    dma_ring * r0 = q->ring0;
+    dma_ring * r1 = q->ring1;
+
+    if (((r0->push_idx + 1) & r0->idx_mask) == r0->pop_idx) {
+        return false;
+    }
+
+    r1->tail = r0->tail;
+
+    size_t rem_bytes   = total;
+    dma_addr_t cur_dst = ddata.dst;
+    dma_addr_t cur_src = ddata.src;
+
+    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)) {
+            dma_ring_flush(r1);
+            dma_ring_push_single_1d(r1, cur_data, cur_bytes);
+        }
+        cur_dst   += cur_bytes;
+        cur_src   += cur_bytes;
+        rem_bytes -= cur_bytes;
+    }
+
+    dma_ring_flush(r1);
+    r0->tail = r1->tail;
+
+    return dma_ring_push_single_1d(r0, ddata, /*size=*/ 0);
+}
+
+#if __HVX_ARCH__ < 75
+
+bool dma_queue_push_fallback_1d(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
+    dma_ring * r0 = q->ring0;
+    dma_ring * r1 = q->ring1;
+
+    if (((r0->push_idx + 1) & r0->idx_mask) == r0->pop_idx) {
+        return false;
+    }
+
+    r1->tail = r0->tail;
+
+    size_t rem_rows    = nrows;
+    dma_addr_t cur_dst = ddata.dst;
+    dma_addr_t cur_src = ddata.src;
+
+    while (rem_rows > 0) {
+        dma_data cur_data = dma_make_data(cur_dst, cur_src);
+        if (!dma_ring_push_single_1d(r1, cur_data, row_size)) {
+            dma_ring_flush(r1);
+            dma_ring_push_single_1d(r1, cur_data, row_size);
+        }
+        cur_dst  += dst_stride;
+        cur_src  += src_stride;
+        rem_rows -= 1;
+    }
+
+    dma_ring_flush(r1);
+    r0->tail = r1->tail;
+
+    return dma_ring_push_single_1d(r0, ddata, /*size=*/ 0);
+}

+#endif
diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.h b/ggml/src/ggml-hexagon/htp/dma-queue.h
index 190ca3a9b..d256e6bef 100644
--- a/ggml/src/ggml-hexagon/htp/dma-queue.h
+++ b/ggml/src/ggml-hexagon/htp/dma-queue.h
@@ -3,8 +3,10 @@

 #include <HAP_farf.h>
 #include <hexagon_types.h>
+#include <assert.h>
 #include <stdbool.h>
 #include <stdint.h>
+#include <string.h>
 #include "hex-utils.h"

 #include "hex-profile.h"
@@ -24,8 +26,8 @@ typedef struct dma_descriptor_1d_s {
     uint32_t src_bypass:1;
     uint32_t order:1;
     uint32_t done:1;
-    void *   src;
-    void *   dst;
+    uint32_t src;
+    uint32_t dst;
 } dma_descriptor_1d;

 #if __HVX_ARCH__ < 75
@@ -40,8 +42,8 @@ typedef struct dma_descriptor_2d_s {
     uint32_t src_bypass:1;
     uint32_t order:1;
     uint32_t done:1;
-    void *   src;
-    void *   dst;
+    uint32_t src;
+    uint32_t dst;
     uint32_t desc_type:8;
     uint32_t reserved1:24;
     uint32_t row_size:16;
@@ -64,10 +66,18 @@ typedef struct dma_descriptor_2d_s {
     uint32_t src_bypass:1;
     uint32_t order:1;
     uint32_t done:1;
-    void *   src;
-    void *   dst;
+    uint32_t src;
+    uint32_t dst;
     uint32_t desc_type:8;
+#if __HVX_ARCH__ > 79
+    uint32_t src_upper:8;
+    uint32_t dst_upper:8;
+    uint32_t allocation:2;
+    uint32_t reserved0:2;
+    uint32_t transform:4;
+#else
     uint32_t reserved0:24;
+#endif
     uint32_t row_size:24;
     uint32_t nrows_lo:8;
     uint32_t nrows_hi:8;
@@ -78,45 +88,63 @@ typedef struct dma_descriptor_2d_s {

 #endif

+#if __HVX_ARCH__ > 79
+typedef uint64_t dma_addr_t;
+#else
+typedef uint32_t dma_addr_t;
+#endif
+
 typedef struct {
-    void       *dst;
-    const void *src;
-} dma_ptr;
+    dma_addr_t dst;
+    dma_addr_t src;
+} dma_data;
+
+// Hardware descriptor field limits
+#define DMA_MAX_NROWS          0xFFFFu        // 16-bit HW descriptor limit (65535)
+#define DMA_MAX_SIZE_16B       0xFFFFu        // 16-bit HW descriptor limit for row_size (65535)
+#define DMA_MAX_STRIDE_16B     0xFFFFu        // 16-bit HW descriptor limit for strides (65535)
+#define DMA_MAX_SIZE_24B       0x00FFFFFFu    // 24-bit HW descriptor limit for row_size / 1D size (16MB - 1)
+#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
+
+#define DMA_FALLBACK_CAPACITY  16u            // descriptors in secondary fallback ring

 typedef struct dma_ring_s dma_ring;
 struct dma_ring_s {
     dma_descriptor_2d * desc;      // descriptor pointers
     dma_descriptor_2d * tail;      // tail pointer
-    dma_ptr           * dptr;      // dst/src pointers
+    dma_data          * data;      // dst/src data
     uint32_t            push_idx;
     uint32_t            pop_idx;
     uint32_t            capacity;
     uint32_t            idx_mask;
     struct htp_thread_trace * trace;
-    uintptr_t           vtcm_base;
-    uintptr_t           vtcm_end;
 };

 typedef struct dma_queue_s dma_queue;
 typedef dma_queue * dma_queue_t;

 struct dma_queue_s {
-    dma_ring *          ring;      // Points to the descriptor ring state
-    uint8_t             nocache;   // Queue-specific bypass flag
+    dma_ring *          ring0;     // Main descriptor ring state
+    dma_ring *          ring1;     // Secondary fallback descriptor ring state
     bool                alias;     // When set, dma_queue_delete will not free the ring
 };

-
-
 size_t      dma_queue_sizeof(size_t capacity);
 size_t      dma_queue_alignof(void);
-dma_queue_t dma_queue_init(void * ptr, size_t capacity, uintptr_t vtcm_base, size_t vtcm_size, struct htp_thread_trace * trace);
+dma_queue_t dma_queue_init(void * ptr, size_t capacity, struct htp_thread_trace * trace);
 void        dma_queue_free(dma_queue_t q);

 size_t      dma_queue_alias_sizeof(void);
-dma_queue_t dma_queue_alias_init(void * ptr, dma_queue_t main_q, uint8_t nocache);
+dma_queue_t dma_queue_alias_init(void * ptr, dma_queue_t main_q);
 void        dma_queue_alias_free(dma_queue_t q);

+bool        dma_queue_push_fallback_2d(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows);
+bool        dma_queue_push_fallback_contig(dma_queue * q, dma_data ddata, size_t total);
+#if __HVX_ARCH__ < 75
+bool        dma_queue_push_fallback_1d(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows);
+#endif
+
 // TODO: technically we don't need these and could use Q6_dmstart/wait/etc instead
 // but those do not seem to always compiler properly.
 static inline void dmstart(void * next) {
@@ -141,36 +169,37 @@ static inline unsigned int dmwait(void) {
     return ret;
 }

-static inline dma_ptr dma_make_ptr(void *dst, const void *src)
+static inline dma_data dma_make_data_impl(dma_addr_t dst, dma_addr_t src)
 {
-    dma_ptr p = { dst, src };
-    return p;
+    dma_data d = { dst, src };
+    return d;
 }

-static inline bool dma_is_vtcm(const dma_queue * q, const void * ptr) {
-    return (uintptr_t) ptr >= q->ring->vtcm_base && (uintptr_t) ptr < q->ring->vtcm_end;
-}
+#define dma_make_data(dst, src) dma_make_data_impl((dma_addr_t) (dst), (dma_addr_t) (src))
+
+static inline bool dma_ring_push_single_1d(dma_ring * r, dma_data ddata, size_t size) {
+#if __HVX_ARCH__ > 79
+    assert(!((ddata.src | ddata.dst) >> 32) || size == 0);
+#endif

-static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t size) {
-    dma_ring * r = q->ring;
     if (((r->push_idx + 1) & r->idx_mask) == r->pop_idx) {
         return false;
     }

     dma_descriptor_1d * desc = (dma_descriptor_1d *) &r->desc[r->push_idx];
-    desc->src  = (void *) dptr.src;
-    desc->dst  = (void *) dptr.dst;
+    desc->src  = (uint32_t) ddata.src;
+    desc->dst  = (uint32_t) ddata.dst;
     desc->size = size;

-    r->dptr[r->push_idx] = dptr;
+    r->data[r->push_idx] = ddata;

     htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx);

     if (size) {
         desc->next       = NULL;
         desc->desc_size  = 0; // 1D mode
-        desc->src_bypass = dma_is_vtcm(q, dptr.src) ? 1 : q->nocache;
-        desc->dst_bypass = dma_is_vtcm(q, dptr.dst) ? 1 : q->nocache;
+        desc->src_bypass = 1;
+        desc->dst_bypass = 1;
         desc->order      = 0;
         desc->done       = 0;

@@ -185,8 +214,17 @@ static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t
     return true;
 }

-static inline bool dma_queue_push_single_2d(dma_queue * q, dma_ptr dptr, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
-    dma_ring * r = q->ring;
+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
+    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) {
         return false;
     }
@@ -194,34 +232,44 @@ static inline bool dma_queue_push_single_2d(dma_queue * q, dma_ptr dptr, size_t
     dma_descriptor_2d * desc = &r->desc[r->push_idx];

     desc->next           = NULL;
-    desc->reserved0      = 0;
     desc->reserved1      = 0;
     desc->desc_size      = 1; // 2d mode
-    desc->src_bypass     = dma_is_vtcm(q, dptr.src) ? 1 : q->nocache;
-    desc->dst_bypass     = dma_is_vtcm(q, dptr.dst) ? 1 : q->nocache;
+    desc->src_bypass     = 1;
+    desc->dst_bypass     = 1;
     desc->src_comp       = 0;
     desc->dst_comp       = 0;
     desc->order          = 0;
     desc->done           = 0;
     desc->src_stride     = src_stride;
     desc->dst_stride     = dst_stride;
-    desc->src            = (void *) dptr.src;
-    desc->dst            = (void *) dptr.dst;
+    desc->src            = (uint32_t) ddata.src;
+    desc->dst            = (uint32_t) ddata.dst;
     desc->row_size       = row_size;

 #if __HVX_ARCH__ < 75
+    desc->reserved0      = 0;
     desc->desc_type      = 0; // 2d (16-bit) mode
     desc->nrows          = nrows;
     desc->src_offset     = 0;
     desc->dst_offset     = 0;
 #else
+#if __HVX_ARCH__ > 79
+    desc->src_upper      = src_hi;
+    desc->dst_upper      = dst_hi;
+    desc->allocation     = 0;
+    desc->reserved0      = 0;
+    desc->transform      = 0;
+    desc->desc_type      = is_ext ? 10 : 9; // 2d 40-bit or 24-bit mode
+#else
+    desc->reserved0      = 0;
     desc->desc_type      = 9; // 2d (24-bit) mode
+#endif
     desc->nrows_lo       = (nrows & 0xff);
     desc->nrows_hi       = (nrows >> 8);
     desc->offset         = 0;
 #endif

-    r->dptr[r->push_idx] = dptr;
+    r->data[r->push_idx] = ddata;

     htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx);

@@ -236,21 +284,20 @@ static inline bool dma_queue_push_single_2d(dma_queue * q, dma_ptr dptr, size_t
     return true;
 }

-static inline dma_ptr dma_queue_pop(dma_queue * q) {
-    dma_ring * r = q->ring;
-    dma_ptr dptr  = { NULL };
+static inline dma_data dma_ring_pop(dma_ring * r) {
+    dma_data ddata = { 0 };

     if (r->push_idx == r->pop_idx) {
-        return dptr;
+        return ddata;
     }

-    dptr = r->dptr[r->pop_idx];
+    ddata = r->data[r->pop_idx];

     volatile dma_descriptor_2d * desc = &r->desc[r->pop_idx];

     // Wait for desc to complete
     if (!desc->done) {
-        // FARF(ALWAYS, "dma-poll: idx %u dst %p src %p", r->pop_idx, dptr.dst, dptr.src);
+        // FARF(ALWAYS, "dma-poll: idx %u dst %p src %p", r->pop_idx, ddata.dst, ddata.src);
         while (!desc->done) {
             dmpoll();
         }
@@ -259,100 +306,125 @@ static inline dma_ptr dma_queue_pop(dma_queue * q) {
     htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx);

     r->pop_idx = (r->pop_idx + 1) & r->idx_mask;
-    return dptr;
+    return ddata;
 }

-static inline dma_ptr dma_queue_pop_nowait(dma_queue * q) {
-    dma_ring * r = q->ring;
-    dma_ptr dptr  = { NULL };
+static inline dma_data dma_ring_pop_nowait(dma_ring * r) {
+    dma_data ddata = { 0 };

     if (r->push_idx == r->pop_idx) {
-        return dptr;
+        return ddata;
     }

-    dptr = r->dptr[r->pop_idx];
+    ddata = r->data[r->pop_idx];

     htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx);

     r->pop_idx = (r->pop_idx + 1) & r->idx_mask;
-    return dptr;
+    return ddata;
+}
+
+static inline bool dma_ring_empty(dma_ring * r) {
+    return r->push_idx == r->pop_idx;
+}
+
+static inline void dma_ring_flush(dma_ring * r) {
+    while (!dma_ring_empty(r)) {
+        dma_ring_pop(r);
+    }
+}
+
+static inline uint32_t dma_ring_depth(dma_ring * r) {
+    return (r->push_idx - r->pop_idx) & r->idx_mask;
+}
+
+static inline uint32_t dma_ring_capacity(dma_ring * r) {
+    return r->capacity;
+}
+
+static inline bool dma_queue_push_single_1d(dma_queue * q, dma_data ddata, size_t size) {
+    return dma_ring_push_single_1d(q->ring0, ddata, size);
+}
+
+static inline bool dma_queue_push_single_2d(dma_queue * q, dma_data ddata, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
+    return dma_ring_push_single_2d(q->ring0, ddata, dst_stride, src_stride, row_size, nrows);
+}
+
+static inline dma_data dma_queue_pop(dma_queue * q) {
+    return dma_ring_pop(q->ring0);
+}
+
+static inline dma_data dma_queue_pop_nowait(dma_queue * q) {
+    return dma_ring_pop_nowait(q->ring0);
 }

 static inline bool dma_queue_empty(dma_queue * q) {
-    return q->ring->push_idx == q->ring->pop_idx;
+    return dma_ring_empty(q->ring0);
 }

 static inline void dma_queue_flush(dma_queue * q) {
-    while (dma_queue_pop(q).dst != NULL) ;
+    dma_ring_flush(q->ring0);
 }

 static inline uint32_t dma_queue_depth(dma_queue * q) {
-    return (q->ring->push_idx - q->ring->pop_idx) & q->ring->idx_mask;
+    return dma_ring_depth(q->ring0);
 }

 static inline uint32_t dma_queue_capacity(dma_queue * q) {
-    return q->ring->capacity;
+    return dma_ring_capacity(q->ring0);
 }

 #if __HVX_ARCH__ < 75

-// Overflow-safe DMA push: all 2d descriptor fields (row_size, nrows, src_stride, dst_stride) are 16-bit, max 65535.
-// This version transparently handles values that exceed the 16-bit limit and submits chained DMA transtions.
-
-#define DMA_MAX_FIELD_VAL 65535u
-
-static inline bool dma_queue_push(dma_queue *q, dma_ptr dptr, size_t dst_stride, size_t src_stride, size_t row_size, size_t 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) {
     // Fast path: everything fits in 16 bits
     if (nrows == 0 || __builtin_expect(
-            row_size   <= DMA_MAX_FIELD_VAL &&
-            nrows      <= DMA_MAX_FIELD_VAL &&
-            src_stride <= DMA_MAX_FIELD_VAL &&
-            dst_stride <= DMA_MAX_FIELD_VAL, 1)) {
-        return dma_queue_push_single_2d(q, dptr, dst_stride, src_stride, row_size, nrows);
+            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);
     }

-    // Contiguous block
-    // Use 1d DMA mode which supports sizes up to 24-bits (16MB)
+    // 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;
-        return dma_queue_push_single_1d(q, dptr, total);
+        if (total <= DMA_MAX_SIZE_24B) {
+            return dma_ring_push_single_1d(q->ring0, ddata, total);
+        }
+        return dma_queue_push_fallback_contig(q, ddata, total);
     }

-    // Stride overflow - fall back to row-by-row.
-    {
-        const uint8_t *src = (const uint8_t *) dptr.src;
-        uint8_t       *dst = (uint8_t *)       dptr.dst;
-        size_t r = 0;
-        while (r + 1 < nrows) {
-            dma_ptr p = dma_make_ptr(dst + r * dst_stride, src + r * src_stride);
-            if (!dma_queue_push_single_1d(q, p, row_size)) {
-                dma_queue_flush(q);
-            } else {
-                r++;
-            }
-        }
-        dma_queue_flush(q);
-        dma_ptr p = dma_make_ptr(dst + r * dst_stride, src + r * src_stride);
-        return dma_queue_push_single_1d(q, p, row_size);
+    // 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_ptr dptr, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) {
-    // On v75 and up we always use 2d 24-bit mode
-    return dma_queue_push_single_2d(q, dptr, 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 (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)) {
+        return dma_ring_push_single_2d(q->ring0, ddata, dst_stride, src_stride, row_size, nrows);
+    }

-#endif
+    // 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);
+    }

-static inline bool dma_queue_push_ddr_to_vtcm(dma_queue * q, dma_ptr dptr, size_t dst_row_size, size_t src_row_size, size_t nrows) {
-    return dma_queue_push(q, dptr, dst_row_size, src_row_size, src_row_size, nrows);
+    return dma_queue_push_fallback_2d(q, ddata, dst_stride, src_stride, row_size, nrows);
 }

-static inline bool dma_queue_push_vtcm_to_ddr(dma_queue * q, dma_ptr dptr, size_t dst_row_size, size_t src_row_size, size_t nrows) {
-    return dma_queue_push(q, dptr, dst_row_size, src_row_size, dst_row_size, nrows);
-}
+#endif

 #define DMA_CACHE_MAX_SIZE 256U

@@ -360,7 +432,7 @@ typedef struct {
     uint8_t *base;
     uint32_t line_size;
     uint32_t capacity;
-    uint32_t src[DMA_CACHE_MAX_SIZE];
+    dma_addr_t src[DMA_CACHE_MAX_SIZE];
     uint16_t age[DMA_CACHE_MAX_SIZE];
 } dma_cache;

@@ -376,14 +448,14 @@ static inline void dma_cache_init(dma_cache *c, uint8_t *base, uint32_t line_siz
     }
 }

-static inline bool dma_cache_push(dma_queue *q, dma_cache *c, const uint8_t * src, uint32_t dst_stride, uint32_t src_stride, uint32_t row_size, uint32_t nrows)
+static inline bool dma_cache_push(dma_queue *q, dma_cache *c, dma_addr_t src_addr, uint32_t dst_stride, uint32_t src_stride, uint32_t row_size, uint32_t nrows)
 {
     uint32_t o_idx = 0;
     uint16_t o_age = 0;
     uint8_t *  dst = 0;

     for (unsigned i=0; i < c->capacity; i++) {
-        if (c->src[i] == (uint32_t) src) {
+        if (c->src[i] == src_addr) {
             c->age[i] = 0;
             dst = c->base + (i * c->line_size); nrows = 0; // dummy dma
         } else {
@@ -393,12 +465,12 @@ static inline bool dma_cache_push(dma_queue *q, dma_cache *c, const uint8_t * sr
     }
     if (!dst) {
         c->age[o_idx] = 0;
-        c->src[o_idx] = (uint32_t) src;
+        c->src[o_idx] = src_addr;
         dst = c->base + o_idx * c->line_size; // normal nrows dma
-        return dma_queue_push(q, dma_make_ptr(dst, src), dst_stride, src_stride, row_size, nrows);
+        return dma_queue_push(q, dma_make_data(dst, src_addr), dst_stride, src_stride, row_size, nrows);
     }

-    return dma_queue_push_single_1d(q, dma_make_ptr(dst, src), 0);
+    return dma_queue_push_single_1d(q, dma_make_data(dst, src_addr), 0);
 }

 #ifdef __cplusplus
diff --git a/ggml/src/ggml-hexagon/htp/fill-ops.c b/ggml/src/ggml-hexagon/htp/fill-ops.c
index 1f6eaafad..212104a23 100644
--- a/ggml/src/ggml-hexagon/htp/fill-ops.c
+++ b/ggml/src/ggml-hexagon/htp/fill-ops.c
@@ -88,8 +88,8 @@ int op_fill(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     uint32_t row_start = 0;
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
index 75422f420..988886082 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
@@ -12,7 +12,7 @@
 #include <stdint.h>
 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hex-fastdiv.h"
 #include "hex-profile.h"
 #include "hmx-queue.h"
@@ -86,6 +86,7 @@ struct htp_fa_context {
     uint8_t * spad_v;
     uint8_t * spad_m;
     uint8_t * spad_a;
+    float *   spad_sinks;

     const struct htp_tensor * k;
     const struct htp_tensor * v;
@@ -149,6 +150,7 @@ struct hmx_fa_context {
     uint8_t *    vtcm_hmx_scales_qk;   // HMX output scales (qk_scale)
     __fp16 *     vtcm_mask_buf;        // VTCM mask buffer [Br * m_line], DMA'd per KV block
     __fp16 *     vtcm_slopes;          // ALiBi slopes [g_br]
+    float *      vtcm_sinks;           // Attention sinks
     size_t       row_buf_stride;       // HVX vectors per row buffer (Bc/64)
     size_t       mask_buf_row_stride;  // elements (__fp16) per row in mask buffer
     size_t       q_tile_bytes;
@@ -213,7 +215,7 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *

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

-    dma_queue * dma = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

     const uint32_t DK = nek0;
     const uint32_t DV = nev0;
@@ -243,27 +245,27 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
         const uint32_t iv3 = fastdiv(iq3, &factx->broadcast_rv3);
         const uint32_t iv2 = fastdiv(iq2, &factx->broadcast_rv2);

-        const __fp16 * mp_base = NULL;
+        dma_addr_t mp_base = 0;
         if (mask) {
             const uint32_t im2 = fastmodulo(iq2, mask->ne[2], &factx->src3_div2);
             const uint32_t im3 = fastmodulo(iq3, mask->ne[3], &factx->src3_div3);
-            mp_base = (const __fp16 *) ((const uint8_t *) mask->data + iq1*mask->nb[1] + im2*mask->nb[2] + im3*mask->nb[3]);
+            mp_base = mask->data + iq1*mask->nb[1] + im2*mask->nb[2] + im3*mask->nb[3];
         }

         // Precalculate next row variables if there is a next row
         bool has_next_ir = (ir + 1 < ir1);
         uint32_t next_ik2 = 0, next_ik3 = 0, next_iv2 = 0, next_iv3 = 0;
-        const uint8_t * next_q_row_ptr = NULL;
-        const __fp16 * next_mp_base = NULL;
+        dma_addr_t next_q_row_ptr = 0;
+        dma_addr_t next_mp_base = 0;

-        const uint8_t * next_k_src0 = NULL;
-        const uint8_t * next_v_src0 = NULL;
-        const uint8_t * next_m_src0 = NULL;
+        dma_addr_t next_k_src0 = 0;
+        dma_addr_t next_v_src0 = 0;
+        dma_addr_t next_m_src0 = 0;
         uint32_t next_block_size0 = 0;

-        const uint8_t * next_k_src1 = NULL;
-        const uint8_t * next_v_src1 = NULL;
-        const uint8_t * next_m_src1 = NULL;
+        dma_addr_t next_k_src1 = 0;
+        dma_addr_t next_v_src1 = 0;
+        dma_addr_t next_m_src1 = 0;
         uint32_t next_block_size1 = 0;

         if (has_next_ir) {
@@ -278,22 +280,22 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
             next_iv3 = fastdiv(next_iq3, &factx->broadcast_rv3);
             next_iv2 = fastdiv(next_iq2, &factx->broadcast_rv2);

-            next_q_row_ptr = (const uint8_t *) q->data + (next_iq1*nbq1 + next_iq2*nbq2 + next_iq3*nbq3);
+            next_q_row_ptr = q->data + next_iq1*nbq1 + next_iq2*nbq2 + next_iq3*nbq3;

             if (mask) {
                 const uint32_t next_im2 = fastmodulo(next_iq2, mask->ne[2], &factx->src3_div2);
                 const uint32_t next_im3 = fastmodulo(next_iq3, mask->ne[3], &factx->src3_div3);
-                next_mp_base = (const __fp16 *) ((const uint8_t *) mask->data + next_iq1*mask->nb[1] + next_im2*mask->nb[2] + next_im3*mask->nb[3]);
+                next_mp_base = mask->data + next_iq1*mask->nb[1] + next_im2*mask->nb[2] + next_im3*mask->nb[3];
             }

             // Precalculate next K/V block 0 source pointers
             {
                 const uint32_t ic_start = 0;
                 next_block_size0 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
-                next_k_src0 = (const uint8_t *) k->data + (ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3);
-                next_v_src0 = (const uint8_t *) v->data + (ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3);
+                next_k_src0 = k->data + ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3;
+                next_v_src0 = v->data + ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3;
                 if (mask) {
-                    next_m_src0 = (const uint8_t *) (next_mp_base + ic_start);
+                    next_m_src0 = next_mp_base + ic_start * sizeof(__fp16);
                 }
             }

@@ -301,18 +303,18 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
             if (factx->n_blocks > 1) {
                 const uint32_t ic_start = 1 * FLASH_ATTN_BLOCK_SIZE;
                 next_block_size1 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
-                next_k_src1 = (const uint8_t *) k->data + (ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3);
-                next_v_src1 = (const uint8_t *) v->data + (ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3);
+                next_k_src1 = k->data + ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3;
+                next_v_src1 = v->data + ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3;
                 if (mask) {
-                    next_m_src1 = (const uint8_t *) (next_mp_base + ic_start);
+                    next_m_src1 = next_mp_base + ic_start * sizeof(__fp16);
                 }
             }
         }

         if (ir == ir0) {
             // Fetch Q row
-            const uint8_t * q_row_ptr = (const uint8_t *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3);
-            dma_queue_push(dma, dma_make_ptr(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
+            const dma_addr_t q_row_ptr = q->data + iq1*nbq1 + iq2*nbq2 + iq3*nbq3;
+            dma_queue_push(dma_q, dma_make_data(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);

             // Prefetch first two blocks
             for (uint32_t ib = 0; ib < MIN(factx->n_blocks, 2); ++ib) {
@@ -320,20 +322,20 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                 const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);

                 // K
-                const uint8_t * k_src = (const uint8_t *) k->data + (ic_start*nbk1 + ik2*nbk2 + ik3*nbk3);
+                const dma_addr_t k_src = k->data + ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
                 uint8_t * k_dst = spad_k + (ib % 2) * factx->size_k_block;
-                dma_queue_push(dma, dma_make_ptr(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size);
+                dma_queue_push(dma_q, dma_make_data(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size);

                 // V
-                const uint8_t * v_src = (const uint8_t *) v->data + (ic_start*nbv1 + iv2*nbv2 + iv3*nbv3);
+                const dma_addr_t v_src = v->data + ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
                 uint8_t * v_dst = spad_v + (ib % 2) * factx->size_v_block;
-                dma_queue_push(dma, dma_make_ptr(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size);
+                dma_queue_push(dma_q, dma_make_data(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size);

                 // Mask
                 if (mask) {
-                    const uint8_t * m_src = (const uint8_t *) (mp_base + ic_start);
+                    const dma_addr_t m_src = mp_base + ic_start * sizeof(__fp16);
                     // Mask is 1D contiguous for this row
-                    dma_cache_push(dma, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1);
+                    dma_cache_push(dma_q, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1);
                 }
             }
         }
@@ -348,7 +350,7 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
         hvx_splat_f32_a(spad_a, 0, DV);
         float * VKQ32 = (float *) (spad_a + 0);

-        uint8_t * q_ptr_vtcm = dma_queue_pop(dma).dst;
+        uint8_t * q_ptr_vtcm = (void *) dma_queue_pop(dma_q).dst;
         if (factx->is_q_fp32) {
             hvx_copy_f16_f32_aa(q_ptr_vtcm, q_ptr_vtcm, DK);  // inplace convert f32 to f16
         }
@@ -365,9 +367,9 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
             const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);

             // Wait for DMA
-            uint8_t * k_base = dma_queue_pop(dma).dst; // K
-            uint8_t * v_base = dma_queue_pop(dma).dst; // V
-            __fp16  * m_base = mask ? dma_queue_pop(dma).dst : NULL; // M
+            uint8_t * k_base = (void *) dma_queue_pop(dma_q).dst; // K
+            uint8_t * v_base = (void *) dma_queue_pop(dma_q).dst; // V
+            __fp16  * m_base = mask ? (__fp16 *) dma_queue_pop(dma_q).dst : NULL; // M

             if (factx->k->type == HTP_TYPE_Q8_0) {
                 htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, ir);
@@ -424,7 +426,7 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *

             if (ib + 1 == factx->n_blocks && has_next_ir) {
                 // Queue next row's Q row!
-                dma_queue_push(dma, dma_make_ptr(spad_q, next_q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
+                dma_queue_push(dma_q, dma_make_data(spad_q, next_q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);

                 if (factx->n_blocks % 2 == 0) {
                     // Queue next row's block 0 (into buffer slot 0)
@@ -432,14 +434,14 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                     uint8_t * v_dst = spad_v + 0 * factx->size_v_block;

                     // K (block 0 of next row)
-                    dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);
+                    dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);

                     // V (block 0 of next row)
-                    dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);
+                    dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);

                     // Mask (block 0 of next row)
                     if (mask) {
-                        dma_cache_push(dma, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
+                        dma_cache_push(dma_q, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
                     }
                 }
             }
@@ -502,17 +504,17 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                 const uint32_t next_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - next_ic_start);

                 // K
-                const uint8_t * k_src = (const uint8_t *) k->data + (next_ic_start*nbk1 + ik2*nbk2 + ik3*nbk3);
-                dma_queue_push(dma, dma_make_ptr(k_base, k_src), factx->size_k_row_padded, nbk1, size_k_row, next_block_size);
+                const dma_addr_t k_src = k->data + next_ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
+                dma_queue_push(dma_q, dma_make_data(k_base, k_src), factx->size_k_row_padded, nbk1, size_k_row, next_block_size);

                 // V
-                const uint8_t * v_src = (const uint8_t *) v->data + (next_ic_start*nbv1 + iv2*nbv2 + iv3*nbv3);
-                dma_queue_push(dma, dma_make_ptr(v_base, v_src), factx->size_v_row_padded, nbv1, size_v_row, next_block_size);
+                const dma_addr_t v_src = v->data + next_ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
+                dma_queue_push(dma_q, dma_make_data(v_base, v_src), factx->size_v_row_padded, nbv1, size_v_row, next_block_size);

                 // Mask
                 if (mask) {
-                    const uint8_t * m_src = (const uint8_t *) (mp_base + next_ic_start);
-                    dma_cache_push(dma, &m_cache, m_src, next_block_size * 2, next_block_size * 2, next_block_size * 2, 1);
+                    const dma_addr_t m_src = mp_base + next_ic_start * sizeof(__fp16);
+                    dma_cache_push(dma_q, &m_cache, m_src, next_block_size * 2, next_block_size * 2, next_block_size * 2, 1);
                 }
             }
         }
@@ -525,14 +527,14 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                     uint8_t * v_dst = spad_v + 1 * factx->size_v_block;

                     // K (block 1 of next row)
-                    dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);
+                    dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);

                     // V (block 1 of next row)
-                    dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);
+                    dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);

                     // Mask (block 1 of next row)
                     if (mask) {
-                        dma_cache_push(dma, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
+                        dma_cache_push(dma_q, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
                     }
                 }
             } else {
@@ -542,14 +544,14 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                     uint8_t * v_dst = spad_v + 0 * factx->size_v_block;

                     // K (block 0 of next row)
-                    dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);
+                    dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);

                     // V (block 0 of next row)
-                    dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);
+                    dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);

                     // Mask (block 0 of next row)
                     if (mask) {
-                        dma_cache_push(dma, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
+                        dma_cache_push(dma_q, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
                     }
                 }

@@ -559,14 +561,14 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
                     uint8_t * v_dst = spad_v + 1 * factx->size_v_block;

                     // K (block 1 of next row)
-                    dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);
+                    dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);

                     // V (block 1 of next row)
-                    dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);
+                    dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);

                     // Mask (block 1 of next row)
                     if (mask) {
-                        dma_cache_push(dma, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
+                        dma_cache_push(dma_q, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
                     }
                 }
             }
@@ -578,7 +580,7 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
         float S = hvx_vec_get_f32(S_vec);

         if (sinks) {
-            const float s = ((float *)((char *) sinks->data))[h];
+            const float s = factx->spad_sinks[h];

             float vs = 1.0f;

@@ -781,7 +783,7 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) {
         const size_t m_end         = hex_smin(m_start + m_bytes_per_t, col_vec_bytes);

         if (factx->sinks) {
-            const float * sinks_data = (const float *) (uintptr_t) factx->sinks->data;
+            const float * sinks_data = factx->vtcm_sinks;
             float *       m_vec      = (float *) factx->vtcm_m_vec;
             const size_t  r_start    = l_start / sizeof(float);
             const size_t  r_end      = l_end / sizeof(float);
@@ -1779,7 +1781,7 @@ static __attribute__((noinline)) void fa_compute_slopes(
 }

 static void fa_push_mask_dma_gqa(
-    dma_queue *               dma,
+    dma_queue *               dma_q,
     const struct htp_tensor * mask,
     uint32_t                  q_start,
     uint32_t                  im3,
@@ -1794,36 +1796,36 @@ static void fa_push_mask_dma_gqa(
     for (uint32_t g = 0; g < G; ++g) {
         const uint32_t h_idx = kv_head * G + g;
         const uint32_t im2 = fastmodulo(h_idx, mask->ne[2], &factx->src3_div2);
-        const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] +
-                                 im2 * mask->nb[2] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
+        const dma_addr_t ms_src = mask->data + q_start * mask->nb[1] +
+                                  im2 * mask->nb[2] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
         uint8_t * ms_dst = (uint8_t *) factx->vtcm_mask_buf + g * m_line_bytes;
-        dma_queue_push(dma, dma_make_ptr(ms_dst, ms_src), G * m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
+        dma_queue_push(dma_q, dma_make_data(ms_dst, ms_src), G * m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
     }
 }

-static void fa_pop_mask_dma_gqa(dma_queue * dma, uint32_t G) {
+static void fa_pop_mask_dma_gqa(dma_queue * dma_q, uint32_t G) {
     for (uint32_t g = 0; g < G; ++g) {
-        dma_queue_pop(dma);
+        dma_queue_pop(dma_q);
     }
 }

-static inline void fa_prefetch_block(dma_queue * dma, const struct htp_tensor * k, const struct htp_tensor * v, const struct htp_tensor * mask,
+static inline void fa_prefetch_block(dma_queue * dma_q, const struct htp_tensor * k, const struct htp_tensor * v, const struct htp_tensor * mask,
                                      uint32_t b, size_t Bc, size_t size_k_row_padded, size_t size_k_row, size_t size_v_row_padded, size_t size_v_row,
                                      uint32_t ik2, uint32_t ik3, uint32_t iv2, uint32_t iv3, uint32_t q_start, uint32_t im3, uint32_t kv_head, uint32_t G,
                                      size_t m_line_bytes, size_t n_rows_q, size_t nek1, size_t prefetch_buf, struct hmx_fa_context * factx) {
     const uint32_t  prefetch_start = b * Bc;
     const uint32_t  prefetch_rows  = hex_smin(Bc, nek1 - prefetch_start);
-    const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
-    dma_queue_push(dma, dma_make_ptr(factx->vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
-    const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
-    dma_queue_push(dma, dma_make_ptr(factx->vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);
+    const dma_addr_t k_prefetch_src = k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
+    dma_queue_push(dma_q, dma_make_data(factx->vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
+    const dma_addr_t v_prefetch_src = v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
+    dma_queue_push(dma_q, dma_make_data(factx->vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);

     if (mask) {
         if (__builtin_expect(factx->mask_broadcast, true)) {
-            const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + prefetch_start * sizeof(__fp16);
-            dma_cache_push(dma, &factx->m_cache, ms_src, m_line_bytes, mask->nb[1], prefetch_rows * sizeof(__fp16), n_rows_q);
+            const dma_addr_t ms_src = mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + prefetch_start * sizeof(__fp16);
+            dma_cache_push(dma_q, &factx->m_cache, ms_src, m_line_bytes, mask->nb[1], prefetch_rows * sizeof(__fp16), n_rows_q);
         } else {
-            fa_push_mask_dma_gqa(dma, mask, q_start, im3, prefetch_start, kv_head, G, m_line_bytes, prefetch_rows, n_rows_q, factx);
+            fa_push_mask_dma_gqa(dma_q, mask, q_start, im3, prefetch_start, kv_head, G, m_line_bytes, prefetch_rows, n_rows_q, factx);
         }
     }
 }
@@ -1953,7 +1955,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     // Build the VTCM layout once (shared with the host estimator) and place every
     // scratch buffer at its computed offset. Padded head dims size the HMX tiles.
     struct hmx_fa_vtcm_layout L;
-    hmx_fa_vtcm_layout_build(&L, G, DK_pad, DV_pad, Br, Bc, n_threads, pipeline, factx.is_q_fp32);
+    hmx_fa_vtcm_layout_build(&L, G, DK_pad, DV_pad, Br, Bc, n_threads, pipeline, factx.is_q_fp32, factx.sinks != NULL, factx.n_heads);

     if (L.total_bytes > ctx->vtcm_size) {
         return HTP_STATUS_VTCM_TOO_SMALL;
@@ -1995,6 +1997,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     factx.col_vec_bytes       = L.col_vec_bytes;
     factx.d_tile_bytes        = L.d_tile_bytes;
     factx.vtcm_slopes         = VTCM_LAYOUT_PTR(__fp16, base, L.off_slopes);
+    factx.vtcm_sinks          = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_sinks, factx.sinks != NULL);

     const size_t m_line_bytes = L.m_line_bytes;  // used by the mask DMAs in the KV loop

@@ -2022,13 +2025,13 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     hmx_init_column_scales(factx.vtcm_hmx_scales_id, Q6_V_vsplat_R(0x3c00)); // 1.0
     hmx_init_column_scales(factx.vtcm_hmx_scales_qk, hvx_vec_splat_f16(factx.scale));

-    // ======== Skip compute if profiling ========
-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     // ======== DMA setup ========
-    dma_queue * const dma = ctx->dma[0];
+    dma_queue * const dma_q = ctx->dma[0];
+
+    if (factx.sinks) {
+        dma_queue_push(dma_q, dma_make_data(factx.vtcm_sinks, factx.sinks->data), L.sinks_bytes, 0, factx.sinks->size, 1);
+        dma_queue_pop(dma_q);
+    }

     const size_t n_row_tiles_g_br = g_br / HMX_FP16_TILE_N_ROWS;
     const size_t n_tiles_per_bc   = Bc / HMX_FP16_TILE_N_COLS;
@@ -2064,32 +2067,32 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                 // 1. Push Q and KV DMAs for the very first iteration.
                 // Subsequent iterations are enqueued early at the end of the previous iteration.
                 if (ib3 == 0 && q_start == q_start_min && kv_head == 0) {
-                    const uint8_t * q_ptr = (const uint8_t *) q->data + q_start * q->nb[1] +
+                    const dma_addr_t q_ptr = q->data + q_start * q->nb[1] +
                                             (kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3];
                     const size_t q_row_bytes = q_transposed ? n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed;
                     const size_t n_rows      = q_transposed ? factx.G : n_rows_q;
-                    dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, q_ptr), q_row_bytes, hex_smax(q_src_stride, q_row_bytes), q_row_bytes, n_rows);
+                    dma_queue_push(dma_q, dma_make_data(factx.vtcm_q_dma, q_ptr), q_row_bytes, hex_smax(q_src_stride, q_row_bytes), q_row_bytes, n_rows);

                     if (factx.n_kv_blocks > 0) {
-                        const uint8_t * k_src = (const uint8_t *) k->data + ik2 * k->nb[2] + ik3 * k->nb[3];
-                        dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0);
+                        const dma_addr_t k_src = k->data + ik2 * k->nb[2] + ik3 * k->nb[3];
+                        dma_queue_push(dma_q, dma_make_data(factx.vtcm_k_fp16[0], k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0);

-                        const uint8_t * v_src = (const uint8_t *) v->data + iv2 * v->nb[2] + iv3 * v->nb[3];
-                        dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0);
+                        const dma_addr_t v_src = v->data + iv2 * v->nb[2] + iv3 * v->nb[3];
+                        dma_queue_push(dma_q, dma_make_data(factx.vtcm_v_fp16[0], v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0);

                         if (factx.pipeline && mask) {
                             if (__builtin_expect(factx.mask_broadcast, true)) {
-                                const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + 0;
-                                dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), n_rows_q);
+                                const dma_addr_t ms_src = mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + 0;
+                                dma_cache_push(dma_q, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), n_rows_q);
                             } else {
-                                fa_push_mask_dma_gqa(dma, mask, q_start, im3, 0, kv_head, G, m_line_bytes, kv_rows0, n_rows_q, &factx);
+                                fa_push_mask_dma_gqa(dma_q, mask, q_start, im3, 0, kv_head, G, m_line_bytes, kv_rows0, n_rows_q, &factx);
                             }
                         }
                     }
                 }

                 // 2. Pop Q DMA (blocks until Q is loaded)
-                dma_queue_pop(dma);
+                dma_queue_pop(dma_q);

                 // ---- Load Q block & Initialize per-block state ----
                 fa_phase_q_load(&factx, q, q_start, kv_head, ib3, n_rows_g);
@@ -2116,12 +2119,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {

                     // Prefetch block 1 early if there are multiple blocks
                     if (factx.n_kv_blocks > 1) {
-                        fa_prefetch_block(dma, k, v, mask, 1, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row,
+                        fa_prefetch_block(dma_q, k, v, mask, 1, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row,
                                           ik2, ik3, iv2, iv3, q_start, im3, kv_head, G, m_line_bytes, n_rows_q, nek1, 1, &factx);
                     }

                     // Prep and start QK-dot(0)
-                    void * curr_k0 = dma_queue_pop(dma).dst;
+                    void * curr_k0 = (void *) dma_queue_pop(dma_q).dst;
                     fa_phase_k_interleave(&factx, kv_rows0, k_src_stride, curr_k0, 0, 0);

                     qk_job[0].q_tiles        = factx.vtcm_q_tiles;
@@ -2140,16 +2143,16 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                         const size_t   n_col_tiles = hmx_ceil_div(kv_rows, HMX_FP16_TILE_N_COLS);

                         // ---- 1. Pop and run V-prep for current block ----
-                        void * curr_v = dma_queue_pop(dma).dst;
+                        void * curr_v = (void *) dma_queue_pop(dma_q).dst;
                         fa_phase_v_interleave(&factx, kv_rows, v_src_stride, curr_v, factx.vtcm_v_tiles[buf_idx], n_tiles_per_bc, kv_start);

                         // ---- 2. Pop and run mask-prep for current block ----
                         __fp16 * current_mask_vtcm = NULL;
                         if (mask) {
                             if (__builtin_expect(factx.mask_broadcast, true)) {
-                                current_mask_vtcm = (__fp16 *) dma_queue_pop(dma).dst;
+                                current_mask_vtcm = (__fp16 *) dma_queue_pop(dma_q).dst;
                             } else {
-                                fa_pop_mask_dma_gqa(dma, G);
+                                fa_pop_mask_dma_gqa(dma_q, G);
                                 current_mask_vtcm = factx.vtcm_mask_buf;
                             }
                         }
@@ -2183,7 +2186,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                             const uint32_t next_rows  = hex_smin(Bc, nek1 - next_start);
                             const size_t   next_buf   = 1 - buf_idx;

-                            void * next_k = dma_queue_pop(dma).dst;
+                            void * next_k = (void *) dma_queue_pop(dma_q).dst;
                             fa_phase_k_interleave(&factx, next_rows, k_src_stride, next_k, next_start, next_buf);

                             qk_job[next_buf].q_tiles        = factx.vtcm_q_tiles;
@@ -2234,7 +2237,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {

                         // Prefetch block kv_blk + 2
                         if (kv_blk + 2 < factx.n_kv_blocks) {
-                            fa_prefetch_block(dma, k, v, mask, kv_blk + 2, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row,
+                            fa_prefetch_block(dma_q, k, v, mask, kv_blk + 2, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row,
                                               ik2, ik3, iv2, iv3, q_start, im3, kv_head, G, m_line_bytes, n_rows_q, nek1, buf_idx, &factx);
                         }

@@ -2276,10 +2279,10 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {

                         if (mask) {
                             if (__builtin_expect(factx.mask_broadcast, true)) {
-                                const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
-                                dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
+                                const dma_addr_t ms_src = mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
+                                dma_cache_push(dma_q, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
                             } else {
-                                fa_push_mask_dma_gqa(dma, mask, q_start, im3, kv_start, kv_head, G, m_line_bytes, kv_rows, n_rows_q, &factx);
+                                fa_push_mask_dma_gqa(dma_q, mask, q_start, im3, kv_start, kv_head, G, m_line_bytes, kv_rows, n_rows_q, &factx);
                             }
                         }

@@ -2287,14 +2290,14 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                             const uint32_t  prefetch_start = (kv_blk + 1) * Bc;
                             const uint32_t  prefetch_rows  = hex_smin(Bc, nek1 - prefetch_start);
                             const size_t    prefetch_buf   = 1 - buf_idx;
-                            const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
-                            dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
-                            const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
-                            dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);
+                            const dma_addr_t k_prefetch_src = k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
+                            dma_queue_push(dma_q, dma_make_data(factx.vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
+                            const dma_addr_t v_prefetch_src = v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
+                            dma_queue_push(dma_q, dma_make_data(factx.vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);
                         }

                         // Wait for current K DMA and interleave
-                        void * curr_k = dma_queue_pop(dma).dst;
+                        void * curr_k = (void *) dma_queue_pop(dma_q).dst;
                         fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start, 0);

                         {
@@ -2312,16 +2315,16 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                         }

                         // Wait for current V DMA and interleave
-                        void * curr_v = dma_queue_pop(dma).dst;
+                        void * curr_v = (void *) dma_queue_pop(dma_q).dst;
                         fa_phase_v_interleave(&factx, kv_rows, v_src_stride, curr_v, factx.vtcm_v_tiles[0], n_tiles_per_bc, kv_start);

                         // ---- Phase 3: softmax + build_D ----
                         __fp16 * current_mask_vtcm = NULL;
                         if (mask) {
                             if (__builtin_expect(factx.mask_broadcast, true)) {
-                                current_mask_vtcm = (__fp16 *) dma_queue_pop(dma).dst;
+                                current_mask_vtcm = (__fp16 *) dma_queue_pop(dma_q).dst;
                             } else {
-                                fa_pop_mask_dma_gqa(dma, G);
+                                fa_pop_mask_dma_gqa(dma_q, G);
                                 current_mask_vtcm = factx.vtcm_mask_buf;
                             }
                         }
@@ -2393,10 +2396,10 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {

                 if (has_next) {
                     const uint32_t next_n_rows_q = hex_smin(Br, neq1 - next_q_start);
-                    const uint8_t * next_q_ptr = (const uint8_t *) q->data + next_q_start * q->nb[1] + (next_kv_head * factx.G) * q->nb[2] + next_ib3 * q->nb[3];
+                    const dma_addr_t next_q_ptr = q->data + next_q_start * q->nb[1] + (next_kv_head * factx.G) * q->nb[2] + next_ib3 * q->nb[3];
                     const size_t next_q_row_bytes = q_transposed ? next_n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed;
                     const size_t next_n_rows      = q_transposed ? factx.G : next_n_rows_q;
-                    dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, next_q_ptr), next_q_row_bytes, hex_smax(q_src_stride, next_q_row_bytes), next_q_row_bytes, next_n_rows);
+                    dma_queue_push(dma_q, dma_make_data(factx.vtcm_q_dma, next_q_ptr), next_q_row_bytes, hex_smax(q_src_stride, next_q_row_bytes), next_q_row_bytes, next_n_rows);

                     if (factx.n_kv_blocks > 0) {
                         const uint32_t next_ik2 = next_kv_head;
@@ -2408,11 +2411,11 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                             next_iv3 = fastdiv(next_ib3, &kparams->broadcast_rv3);
                         }

-                        const uint8_t * next_k_src = (const uint8_t *) k->data + next_ik2 * k->nb[2] + next_ik3 * k->nb[3];
-                        dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], next_k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0);
+                        const dma_addr_t next_k_src = k->data + next_ik2 * k->nb[2] + next_ik3 * k->nb[3];
+                        dma_queue_push(dma_q, dma_make_data(factx.vtcm_k_fp16[0], next_k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0);

-                        const uint8_t * next_v_src = (const uint8_t *) v->data + next_iv2 * v->nb[2] + next_iv3 * v->nb[3];
-                        dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], next_v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0);
+                        const dma_addr_t next_v_src = v->data + next_iv2 * v->nb[2] + next_iv3 * v->nb[3];
+                        dma_queue_push(dma_q, dma_make_data(factx.vtcm_v_fp16[0], next_v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0);

                         if (factx.pipeline && mask) {
                             uint32_t next_im3 = im3;
@@ -2420,10 +2423,10 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                                 next_im3 = fastmodulo(next_ib3, mask->ne[3], &factx.src3_div3);
                             }
                             if (__builtin_expect(factx.mask_broadcast, true)) {
-                                const uint8_t * ms_src = (const uint8_t *) mask->data + next_q_start * mask->nb[1] + next_im3 * mask->nb[3] + 0;
-                                dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), next_n_rows_q);
+                                const dma_addr_t ms_src = mask->data + next_q_start * mask->nb[1] + next_im3 * mask->nb[3] + 0;
+                                dma_cache_push(dma_q, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), next_n_rows_q);
                             } else {
-                                fa_push_mask_dma_gqa(dma, mask, next_q_start, next_im3, 0, next_kv_head, G, m_line_bytes, kv_rows0, next_n_rows_q, &factx);
+                                fa_push_mask_dma_gqa(dma_q, mask, next_q_start, next_im3, 0, next_kv_head, G, m_line_bytes, kv_rows0, next_n_rows_q, &factx);
                             }
                         }
                     }
@@ -2465,6 +2468,10 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

+    if (htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     const struct htp_fa_kernel_params * kparams = (const struct htp_fa_kernel_params *) octx->kernel_params;

     if (kparams->kernel_type == HTP_FA_KERNEL_UNSUPPORTED) {
@@ -2502,11 +2509,6 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     factx.size_k_row_padded = kparams->u.hvx.size_k_row_padded;
     factx.size_v_row_padded = kparams->u.hvx.size_v_row_padded;

-    size_t size_q_block = factx.size_q_row_padded * 1; // single row for now
-    factx.size_k_block = factx.size_k_row_padded * FLASH_ATTN_BLOCK_SIZE;
-    factx.size_v_block = factx.size_v_row_padded * FLASH_ATTN_BLOCK_SIZE;
-    factx.size_m_block = hex_round_up(FLASH_ATTN_BLOCK_SIZE * sizeof(__fp16), 128);
-
     factx.n_blocks = kparams->n_kv_blocks;

     factx.scale = kparams->scale;
@@ -2551,26 +2553,38 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     factx.qrow_start       = qrow_start;
     factx.qrows_per_thread = fastdiv(qrows + n_threads - 1, &octx->n_threads_div);

-    size_t size_vkq_acc = hex_round_up(v->ne[0] * sizeof(float), 128); // VKQ32
-
-    factx.size_q_block = size_q_block;
-    factx.size_vkq_acc = size_vkq_acc;
+    const bool has_mask  = (mask != NULL);
+    const bool has_sinks = (octx->src[4] != NULL);
+    struct hvx_fa_vtcm_layout L;
+    hvx_fa_vtcm_layout_build(&L, k->ne[0], v->ne[0], factx.is_q_fp32, has_mask, has_sinks, n_head, n_threads);

-    uint8_t * vtcm_cur = octx->ctx->vtcm_base;
-
-    factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * n_threads);
-    factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * n_threads);
-    factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * n_threads);
-    factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * n_threads);
-    factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * n_threads);
-
-    if ((size_t) (vtcm_cur - octx->ctx->vtcm_base) > octx->ctx->vtcm_size) {
+    if (L.total_bytes > octx->ctx->vtcm_size) {
         return HTP_STATUS_VTCM_TOO_SMALL;
     }

-    if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, n_threads);
+    factx.size_q_block = L.size_q_block;
+    factx.size_k_block = L.size_k_block;
+    factx.size_v_block = L.size_v_block;
+    factx.size_m_block = L.size_m_block;
+    factx.size_vkq_acc = L.size_vkq_acc;
+
+    uint8_t * const base = octx->ctx->vtcm_base;
+
+    factx.spad_q     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_q);
+    factx.spad_k     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_k);
+    factx.spad_v     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_v);
+    factx.spad_m     = VTCM_LAYOUT_PTR_OPTIONAL(uint8_t, base, L.off_m, has_mask);
+    factx.spad_a     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_a);
+    factx.spad_sinks = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_sinks, has_sinks);
+
+    if (has_sinks) {
+        const struct htp_tensor * sinks = octx->src[4];
+        dma_queue * dma_q = octx->ctx->dma[0];
+        dma_queue_push(dma_q, dma_make_data(factx.spad_sinks, sinks->data), L.size_sinks, 0, sinks->size, 1);
+        dma_queue_pop(dma_q);
     }

+    work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, n_threads);
+
     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
index 027845411..2bd232190 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
@@ -121,6 +121,7 @@ struct hmx_fa_vtcm_layout {
     size_t off_hmx_scales_qk;
     size_t off_mask_buf;
     size_t off_slopes;
+    size_t off_sinks;

     // Region byte sizes reused by the device at runtime (not just for allocation).
     size_t q_tile_bytes;
@@ -130,6 +131,7 @@ struct hmx_fa_vtcm_layout {
     size_t m_line_bytes;       // one mask row
     size_t m_buf_slot_bytes;   // one dma_cache slot = align_up(Br * m_line_bytes, 4096)
     size_t col_vec_bytes;
+    size_t sinks_bytes;

     // Derived strides.
     size_t row_buf_stride;       // HVX vectors (128B) per row buffer
@@ -142,8 +144,10 @@ struct hmx_fa_vtcm_layout {
 // Build the VTCM layout.

 static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L,
-                                       size_t gqa_factor, size_t DK, size_t DV,
-                                       size_t Br, size_t Bc, size_t n_threads, bool pipeline, bool is_q_fp32) {
+                                            size_t gqa_factor, size_t DK, size_t DV,
+                                            size_t Br, size_t Bc, size_t n_threads,
+                                            bool pipeline, bool is_q_fp32,
+                                            bool has_sinks, size_t n_heads) {
     const size_t g_br         = hex_align_up(gqa_factor * Br, HMX_FP16_TILE_N_ROWS);
     const size_t q_tile_size  = hex_align_up(g_br * DK   * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE);
     const size_t o_tile_size  = hex_align_up(g_br * DV   * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE);
@@ -166,6 +170,7 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L,
     const size_t m_buf_slot   = hex_align_up(Br * m_line_size, 256);
     const size_t m_buf_size   = m_buf_slot * HMX_FA_DMA_CACHE_SIZE;
     const size_t slopes_size  = hex_align_up(g_br * sizeof(__fp16), 128);
+    const size_t sinks_size   = hex_round_up(n_heads * sizeof(float), 128);

     size_t off = 0;

@@ -215,6 +220,7 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L,
     VTCM_LAYOUT_ALLOC(off, off_hmx_scales_qk, 256);
     VTCM_LAYOUT_ALLOC(off, off_mask_buf,      m_buf_size);
     VTCM_LAYOUT_ALLOC(off, off_slopes,        slopes_size);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_sinks, sinks_size, has_sinks);

     L->q_tile_bytes        = q_tile_size;
     L->o_tile_bytes        = o_tile_size;
@@ -228,20 +234,43 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L,
     L->m_buf_slot_bytes    = m_buf_slot;
     L->row_buf_stride      = row_vec_size / 128;
     L->mask_buf_row_stride = m_line_size / sizeof(__fp16);
+    L->sinks_bytes         = has_sinks ? sinks_size : 0;
     L->pipeline            = pipeline;
     L->total_bytes         = off;
 }

 // Exact VTCM usage for a given (gqa_factor, DK, DV, Br, Bc) configuration.
-static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, size_t DV, size_t Br, size_t Bc, size_t n_threads, bool pipeline, bool is_q_fp32) {
+static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, size_t DV, size_t Br, size_t Bc, size_t n_threads, bool pipeline, bool is_q_fp32, bool has_sinks, size_t n_heads) {
     struct hmx_fa_vtcm_layout L;
-    hmx_fa_vtcm_layout_build(&L, gqa_factor, DK, DV, Br, Bc, n_threads, pipeline, is_q_fp32);
+    hmx_fa_vtcm_layout_build(&L, gqa_factor, DK, DV, Br, Bc, n_threads, pipeline, is_q_fp32, has_sinks, n_heads);
     return L.total_bytes;
 }

 #define FA_HVX_BLOCK_SIZE 64

-static inline size_t hvx_fa_compute_vtcm_usage(size_t DK, size_t DV, bool is_q_fp32, bool has_mask, size_t n_threads) {
+struct hvx_fa_vtcm_layout {
+    size_t off_q;
+    size_t off_k;
+    size_t off_v;
+    size_t off_m;
+    size_t off_a;
+    size_t off_sinks;
+
+    size_t size_q_block;
+    size_t size_k_block;
+    size_t size_v_block;
+    size_t size_m_block;
+    size_t size_vkq_acc;
+    size_t size_sinks;
+
+    size_t total_bytes;
+};
+
+static inline void hvx_fa_vtcm_layout_build(struct hvx_fa_vtcm_layout * L,
+                                            size_t DK, size_t DV,
+                                            bool is_q_fp32, bool has_mask,
+                                            bool has_sinks, size_t n_heads,
+                                            size_t n_threads) {
     const size_t size_q_row_padded = hex_round_up(DK * (is_q_fp32 ? 4 : 2), 128);
     const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128);
     const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128);
@@ -251,29 +280,47 @@ static inline size_t hvx_fa_compute_vtcm_usage(size_t DK, size_t DV, bool is_q_f
     const size_t size_v_block = size_v_row_padded * FA_HVX_BLOCK_SIZE;
     const size_t size_m_block = hex_round_up(FA_HVX_BLOCK_SIZE * sizeof(__fp16), 128);
     const size_t size_vkq_acc = hex_round_up(DV * sizeof(float), 128);
+    const size_t size_sinks   = hex_round_up(n_heads * sizeof(float), 128);

-    const size_t size_per_thread = size_q_block * 1
-                                 + size_k_block * 2
-                                 + size_v_block * 2
-                                 + (has_mask ? size_m_block * HVX_FA_DMA_CACHE_SIZE : 0)
-                                 + size_vkq_acc;
+    size_t off = 0;
+
+    VTCM_LAYOUT_ALLOC(off, off_q, size_q_block * n_threads);
+    VTCM_LAYOUT_ALLOC(off, off_k, size_k_block * 2 * n_threads);
+    VTCM_LAYOUT_ALLOC(off, off_v, size_v_block * 2 * n_threads);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_m, size_m_block * HVX_FA_DMA_CACHE_SIZE * n_threads, has_mask);
+    VTCM_LAYOUT_ALLOC(off, off_a, size_vkq_acc * n_threads);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_sinks, size_sinks, has_sinks);
+
+    L->size_q_block = size_q_block;
+    L->size_k_block = size_k_block;
+    L->size_v_block = size_v_block;
+    L->size_m_block = size_m_block;
+    L->size_vkq_acc = size_vkq_acc;
+    L->size_sinks   = has_sinks ? size_sinks : 0;
+    L->total_bytes  = off;
+}

-    return size_per_thread * n_threads;
+static inline size_t hvx_fa_compute_vtcm_usage(size_t DK, size_t DV, bool is_q_fp32, bool has_mask, bool has_sinks, size_t n_heads, size_t n_threads) {
+    struct hvx_fa_vtcm_layout L;
+    hvx_fa_vtcm_layout_build(&L, DK, DV, is_q_fp32, has_mask, has_sinks, n_heads, n_threads);
+    return L.total_bytes;
 }

 #define FA_MIN_KV_BLOCKS 3

 // Cost-based (Br, Bc) search for flash attention with pipeline constraint.
 static inline int hmx_fa_find_chunk_size(size_t * Br_out,
-                                  size_t * Bc_out,
-                                  size_t   gqa_factor,
-                                  size_t   DK,
-                                  size_t   DV,
-                                  size_t   qo_len,
-                                  size_t   kv_len,
-                                  size_t   vtcm_budget,
-                                  size_t   n_threads,
-                                  bool     is_q_fp32) {
+                                         size_t * Bc_out,
+                                         size_t   gqa_factor,
+                                         size_t   DK,
+                                         size_t   DV,
+                                         size_t   qo_len,
+                                         size_t   kv_len,
+                                         size_t   vtcm_budget,
+                                         size_t   n_threads,
+                                         bool     is_q_fp32,
+                                         bool     has_sinks,
+                                         size_t   n_heads) {
     const size_t T       = HMX_FP16_TILE_N_ROWS;  // 32
     const size_t br_unit = hmx_ceil_div(T, gqa_factor);
     const size_t bc_unit = HMX_FP16_TILE_N_COLS * 2;  // 64
@@ -297,7 +344,7 @@ static inline int hmx_fa_find_chunk_size(size_t * Br_out,
     for (size_t Br = Br_max; Br >= br_unit; Br -= br_unit) {
         // Try all Bc candidates from Bc_limit down to bc_unit
         for (size_t Bc = Bc_limit; Bc >= bc_unit; Bc -= bc_unit) {
-            size_t vtcm_needed = hmx_fa_compute_vtcm_usage(gqa_factor, DK, DV, Br, Bc, n_threads, can_pipeline, is_q_fp32);
+            size_t vtcm_needed = hmx_fa_compute_vtcm_usage(gqa_factor, DK, DV, Br, Bc, n_threads, can_pipeline, is_q_fp32, has_sinks, n_heads);
             if (vtcm_needed <= vtcm_budget) {
                 // This Bc fits for this Br!
                 const size_t q_blocks       = (qo_len + Br - 1) / Br;
diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
index 0b6529571..b37313370 100644
--- a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
+++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
@@ -1,44 +1,40 @@
-#include <math.h>
 #include <stdint.h>
+#include <stdbool.h>
 #include <string.h>
+#include <math.h>

-#include "hvx-utils.h"
-#include "hex-fastdiv.h"
-#include "hex-common.h"
-#include "hex-profile.h"
-
-#define GGML_COMMON_DECL_C
+#include "hvx-base.h"
+#include "hvx-copy.h"
+#include "hvx-reduce.h"
+#include "hvx-exp.h"
+#include "dma-queue.h"
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "htp-tensor.h"
+#include "gated-delta-net-ops.h"

 #ifndef MIN
 #define MIN(a, b) ((a) < (b) ? (a) : (b))
 #endif

-#define HTP_GDN_MAX_SV 128
-
-
 struct htp_gdn_context {
     struct htp_ops_context * octx;
-    uint32_t rows_per_thread;
-    size_t   state_bytes;
+    const struct htp_gdn_kernel_params * kparams;
+    struct htp_gdn_vtcm_layout layout;
     uint8_t * vtcm_base;
-    size_t   vtcm_per_thread;
     uint32_t row_start;
     uint32_t nrows;
 };

-static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const float * restrict mul, const float * restrict dot, uint32_t n) {
+static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const HVX_Vector * restrict mul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc = Q6_V_vzero();
-
-    const uint32_t epv  = 128 / sizeof(float);
+    const uint32_t epv = 128 / sizeof(float);
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
         HVX_Vector vd   = hvx_vmemu(dst + i * epv);
-        HVX_Vector vm   = hvx_vmem(mul + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vm   = mul[i];
+        HVX_Vector vdot = dot[i];
         HVX_Vector out  = hvx_vec_mul_f32_f32(vd, vm);
         hvx_vmemu(dst + i * epv) = out;
         acc = hvx_vec_add_f32_f32(acc, hvx_vec_mul_f32_f32(out, vdot));
@@ -46,29 +42,28 @@ static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const float * res

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vd   = hvx_vmemu(dst + off);
-        HVX_Vector vm   = hvx_vmem(mul + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
-        HVX_Vector out  = hvx_vec_mul_f32_f32(vd, vm);
-        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        HVX_Vector vm = mul[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
-        HVX_Vector prod = hvx_vec_mul_f32_f32(out, vdot);
-        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, prod, Q6_V_vzero()));
+        HVX_Vector zero = Q6_V_vzero();
+
+        HVX_Vector out = hvx_vec_mul_f32_f32(hvx_vmemu(dst + off), vm);
+        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, hvx_vec_mul_f32_f32(out, vdot), zero));
     }

     return hvx_vec_reduce_sum_f32(acc);
 }

-static inline HVX_Vector gdn_mul_scalar_dot_f32(float * restrict dst, float mul, const float * restrict dot, uint32_t n) {
+static inline HVX_Vector gdn_mul_scalar_dot_f32(float * restrict dst, float mul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc = Q6_V_vzero();
     const HVX_Vector vmul = hvx_vec_splat_f32(mul);
-
-    const uint32_t epv  = 128 / sizeof(float);
+    const uint32_t epv = 128 / sizeof(float);
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
         HVX_Vector vd   = hvx_vmemu(dst + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vdot = dot[i];
         HVX_Vector out  = hvx_vec_mul_f32_f32(vd, vmul);
         hvx_vmemu(dst + i * epv) = out;
         acc = hvx_vec_add_f32_f32(acc, hvx_vec_mul_f32_f32(out, vdot));
@@ -76,29 +71,28 @@ static inline HVX_Vector gdn_mul_scalar_dot_f32(float * restrict dst, float mul,

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vd   = hvx_vmemu(dst + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
-        HVX_Vector out  = hvx_vec_mul_f32_f32(vd, vmul);
-        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
-        HVX_Vector prod = hvx_vec_mul_f32_f32(out, vdot);
-        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, prod, Q6_V_vzero()));
+        HVX_Vector zero = Q6_V_vzero();
+
+        HVX_Vector out = hvx_vec_mul_f32_f32(hvx_vmemu(dst + off), vmul);
+        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, hvx_vec_mul_f32_f32(out, vdot), zero));
     }

     return hvx_vec_reduce_sum_f32(acc);
 }

-static inline HVX_Vector gdn_add_scaled_dot_f32(float * restrict dst, const float * restrict src,
-        HVX_Vector vscale, const float * restrict dot, uint32_t n) {
+static inline HVX_Vector gdn_add_scaled_dot_f32(float * restrict dst, const HVX_Vector * restrict src,
+        HVX_Vector vscale, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc = Q6_V_vzero();
-
-    const uint32_t epv  = 128 / sizeof(float);
+    const uint32_t epv = 128 / sizeof(float);
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
         HVX_Vector vd   = hvx_vmemu(dst + i * epv);
-        HVX_Vector vs   = hvx_vmem(src + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vs   = src[i];
+        HVX_Vector vdot = dot[i];
         HVX_Vector out  = hvx_vec_add_f32_f32(vd, hvx_vec_mul_f32_f32(vs, vscale));
         hvx_vmemu(dst + i * epv) = out;
         acc = hvx_vec_add_f32_f32(acc, hvx_vec_mul_f32_f32(out, vdot));
@@ -106,22 +100,22 @@ static inline HVX_Vector gdn_add_scaled_dot_f32(float * restrict dst, const floa

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vd   = hvx_vmemu(dst + off);
-        HVX_Vector vs   = hvx_vmem(src + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
-        HVX_Vector out  = hvx_vec_add_f32_f32(vd, hvx_vec_mul_f32_f32(vs, vscale));
-        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        HVX_Vector vs   = src[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
-        HVX_Vector prod = hvx_vec_mul_f32_f32(out, vdot);
-        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, prod, Q6_V_vzero()));
+        HVX_Vector zero = Q6_V_vzero();
+
+        HVX_Vector out = hvx_vec_add_f32_f32(hvx_vmemu(dst + off), hvx_vec_mul_f32_f32(vs, vscale));
+        hvx_vec_store_u(dst + off, nloe * sizeof(float), out);
+        acc = hvx_vec_add_f32_f32(acc, Q6_V_vmux_QVV(mask, hvx_vec_mul_f32_f32(out, vdot), zero));
     }

     return hvx_vec_reduce_sum_f32(acc);
 }

-static inline void gdn_mul_dot4_f32(float * restrict dst0, float * restrict dst1,
-        float * restrict dst2, float * restrict dst3, const float * restrict mul,
-        const float * restrict dot, uint32_t n, float * restrict sums) {
+static inline HVX_Vector gdn_mul_dot4_f32(float * restrict dst0, float * restrict dst1,
+        float * restrict dst2, float * restrict dst3,
+        const HVX_Vector * restrict mul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
@@ -131,8 +125,8 @@ static inline void gdn_mul_dot4_f32(float * restrict dst0, float * restrict dst1
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vm = hvx_vmem(mul + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vm   = mul[i];
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_mul_f32_f32(hvx_vmemu(dst0 + i * epv), vm);
         HVX_Vector out1 = hvx_vec_mul_f32_f32(hvx_vmemu(dst1 + i * epv), vm);
@@ -152,8 +146,8 @@ static inline void gdn_mul_dot4_f32(float * restrict dst0, float * restrict dst1

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vm   = hvx_vmem(mul + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vm = mul[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -174,23 +168,22 @@ static inline void gdn_mul_dot4_f32(float * restrict dst0, float * restrict dst1
     }

     HVX_Vector_x4 acc = { .v = { acc0, acc1, acc2, acc3 } };
-    hvx_vec_store_u(sums, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(acc));
+    return hvx_vec_reduce_sum_f32x4(acc);
 }

-static inline void gdn_mul_scalar_dot4_f32(float * restrict dst0, float * restrict dst1,
-        float * restrict dst2, float * restrict dst3, float mul,
-        const float * restrict dot, uint32_t n, float * restrict sums) {
+static inline HVX_Vector gdn_mul_scalar_dot4_f32(float * restrict dst0, float * restrict dst1,
+        float * restrict dst2, float * restrict dst3,
+        HVX_Vector vmul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
     HVX_Vector acc3 = Q6_V_vzero();
-    const HVX_Vector vmul = hvx_vec_splat_f32(mul);

     const uint32_t epv = 128 / sizeof(float);
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_mul_f32_f32(hvx_vmemu(dst0 + i * epv), vmul);
         HVX_Vector out1 = hvx_vec_mul_f32_f32(hvx_vmemu(dst1 + i * epv), vmul);
@@ -210,7 +203,7 @@ static inline void gdn_mul_scalar_dot4_f32(float * restrict dst0, float * restri

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -231,13 +224,13 @@ static inline void gdn_mul_scalar_dot4_f32(float * restrict dst0, float * restri
     }

     HVX_Vector_x4 acc = { .v = { acc0, acc1, acc2, acc3 } };
-    hvx_vec_store_u(sums, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(acc));
+    return hvx_vec_reduce_sum_f32x4(acc);
 }

-static inline void gdn_add_scaled_dot4_f32(float * restrict dst0, float * restrict dst1,
-        float * restrict dst2, float * restrict dst3, const float * restrict src,
-        const float * restrict scale, const float * restrict dot, uint32_t n,
-        float * restrict sums) {
+static inline HVX_Vector gdn_add_scaled_dot4_f32(float * restrict dst0, float * restrict dst1,
+        float * restrict dst2, float * restrict dst3,
+        const HVX_Vector * restrict src, const float * restrict scale,
+        const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
@@ -251,8 +244,8 @@ static inline void gdn_add_scaled_dot4_f32(float * restrict dst0, float * restri
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vs = hvx_vmem(src + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vs   = src[i];
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_add_f32_f32(hvx_vmemu(dst0 + i * epv), hvx_vec_mul_f32_f32(vs, scale0));
         HVX_Vector out1 = hvx_vec_add_f32_f32(hvx_vmemu(dst1 + i * epv), hvx_vec_mul_f32_f32(vs, scale1));
@@ -272,8 +265,8 @@ static inline void gdn_add_scaled_dot4_f32(float * restrict dst0, float * restri

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vs = hvx_vmem(src + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vs   = src[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -294,14 +287,13 @@ static inline void gdn_add_scaled_dot4_f32(float * restrict dst0, float * restri
     }

     HVX_Vector_x4 acc = { .v = { acc0, acc1, acc2, acc3 } };
-    hvx_vec_store_u(sums, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(acc));
+    return hvx_vec_reduce_sum_f32x4(acc);
 }

-static inline void gdn_mul_dot8_f32(float * restrict dst0, float * restrict dst1,
+static inline HVX_Vector gdn_mul_dot8_f32(float * restrict dst0, float * restrict dst1,
         float * restrict dst2, float * restrict dst3, float * restrict dst4,
         float * restrict dst5, float * restrict dst6, float * restrict dst7,
-        const float * restrict mul, const float * restrict dot, uint32_t n,
-        float * restrict sums) {
+        const HVX_Vector * restrict mul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
@@ -315,8 +307,8 @@ static inline void gdn_mul_dot8_f32(float * restrict dst0, float * restrict dst1
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vm = hvx_vmem(mul + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vm   = mul[i];
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_mul_f32_f32(hvx_vmemu(dst0 + i * epv), vm);
         HVX_Vector out1 = hvx_vec_mul_f32_f32(hvx_vmemu(dst1 + i * epv), vm);
@@ -348,8 +340,8 @@ static inline void gdn_mul_dot8_f32(float * restrict dst0, float * restrict dst1

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vm = hvx_vmem(mul + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vm = mul[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -383,14 +375,16 @@ static inline void gdn_mul_dot8_f32(float * restrict dst0, float * restrict dst1

     HVX_Vector_x4 accA = { .v = { acc0, acc1, acc2, acc3 } };
     HVX_Vector_x4 accB = { .v = { acc4, acc5, acc6, acc7 } };
-    hvx_vec_store_u(sums + 0, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accA));
-    hvx_vec_store_u(sums + 4, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accB));
+    HVX_Vector rA = hvx_vec_reduce_sum_f32x4(accA);
+    HVX_Vector rB = hvx_vec_reduce_sum_f32x4(accB);
+    HVX_VectorPred q16 = Q6_Q_vsetq2_R(16);
+    return Q6_V_vmux_QVV(q16, rA, Q6_V_vror_VR(rB, 128 - 16));
 }

-static inline void gdn_mul_scalar_dot8_f32(float * restrict dst0, float * restrict dst1,
+static inline HVX_Vector gdn_mul_scalar_dot8_f32(float * restrict dst0, float * restrict dst1,
         float * restrict dst2, float * restrict dst3, float * restrict dst4,
         float * restrict dst5, float * restrict dst6, float * restrict dst7,
-        float mul, const float * restrict dot, uint32_t n, float * restrict sums) {
+        HVX_Vector vmul, const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
@@ -399,13 +393,12 @@ static inline void gdn_mul_scalar_dot8_f32(float * restrict dst0, float * restri
     HVX_Vector acc5 = Q6_V_vzero();
     HVX_Vector acc6 = Q6_V_vzero();
     HVX_Vector acc7 = Q6_V_vzero();
-    const HVX_Vector vmul = hvx_vec_splat_f32(mul);

     const uint32_t epv = 128 / sizeof(float);
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_mul_f32_f32(hvx_vmemu(dst0 + i * epv), vmul);
         HVX_Vector out1 = hvx_vec_mul_f32_f32(hvx_vmemu(dst1 + i * epv), vmul);
@@ -437,7 +430,7 @@ static inline void gdn_mul_scalar_dot8_f32(float * restrict dst0, float * restri

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -471,15 +464,17 @@ static inline void gdn_mul_scalar_dot8_f32(float * restrict dst0, float * restri

     HVX_Vector_x4 accA = { .v = { acc0, acc1, acc2, acc3 } };
     HVX_Vector_x4 accB = { .v = { acc4, acc5, acc6, acc7 } };
-    hvx_vec_store_u(sums + 0, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accA));
-    hvx_vec_store_u(sums + 4, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accB));
+    HVX_Vector rA = hvx_vec_reduce_sum_f32x4(accA);
+    HVX_Vector rB = hvx_vec_reduce_sum_f32x4(accB);
+    HVX_VectorPred q16 = Q6_Q_vsetq2_R(16);
+    return Q6_V_vmux_QVV(q16, rA, Q6_V_vror_VR(rB, 128 - 16));
 }

-static inline void gdn_add_scaled_dot8_f32(float * restrict dst0, float * restrict dst1,
+static inline HVX_Vector gdn_add_scaled_dot8_f32(float * restrict dst0, float * restrict dst1,
         float * restrict dst2, float * restrict dst3, float * restrict dst4,
         float * restrict dst5, float * restrict dst6, float * restrict dst7,
-        const float * restrict src, const float * restrict scale,
-        const float * restrict dot, uint32_t n, float * restrict sums) {
+        const HVX_Vector * restrict src, const float * restrict scale,
+        const HVX_Vector * restrict dot, uint32_t n) {
     HVX_Vector acc0 = Q6_V_vzero();
     HVX_Vector acc1 = Q6_V_vzero();
     HVX_Vector acc2 = Q6_V_vzero();
@@ -501,8 +496,8 @@ static inline void gdn_add_scaled_dot8_f32(float * restrict dst0, float * restri
     const uint32_t nvec = n / epv;
     const uint32_t nloe = n % epv;
     for (uint32_t i = 0; i < nvec; ++i) {
-        HVX_Vector vs = hvx_vmem(src + i * epv);
-        HVX_Vector vdot = hvx_vmem(dot + i * epv);
+        HVX_Vector vs   = src[i];
+        HVX_Vector vdot = dot[i];

         HVX_Vector out0 = hvx_vec_add_f32_f32(hvx_vmemu(dst0 + i * epv), hvx_vec_mul_f32_f32(vs, scale0));
         HVX_Vector out1 = hvx_vec_add_f32_f32(hvx_vmemu(dst1 + i * epv), hvx_vec_mul_f32_f32(vs, scale1));
@@ -534,8 +529,8 @@ static inline void gdn_add_scaled_dot8_f32(float * restrict dst0, float * restri

     if (nloe) {
         const uint32_t off = nvec * epv;
-        HVX_Vector vs = hvx_vmem(src + off);
-        HVX_Vector vdot = hvx_vmem(dot + off);
+        HVX_Vector vs   = src[nvec];
+        HVX_Vector vdot = dot[nvec];
         HVX_VectorPred mask = Q6_Q_vsetq2_R(nloe * sizeof(float));
         HVX_Vector zero = Q6_V_vzero();

@@ -569,13 +564,196 @@ static inline void gdn_add_scaled_dot8_f32(float * restrict dst0, float * restri

     HVX_Vector_x4 accA = { .v = { acc0, acc1, acc2, acc3 } };
     HVX_Vector_x4 accB = { .v = { acc4, acc5, acc6, acc7 } };
-    hvx_vec_store_u(sums + 0, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accA));
-    hvx_vec_store_u(sums + 4, 4 * sizeof(float), hvx_vec_reduce_sum_f32x4(accB));
+    HVX_Vector rA = hvx_vec_reduce_sum_f32x4(accA);
+    HVX_Vector rB = hvx_vec_reduce_sum_f32x4(accB);
+    HVX_VectorPred q16 = Q6_Q_vsetq2_R(16);
+    return Q6_V_vmux_QVV(q16, rA, Q6_V_vror_VR(rB, 128 - 16));
+}
+
+static inline void gdn_step_kda_f32(
+    float * restrict s_work,
+    float * restrict attn_out,
+    const float * restrict q_t,
+    const float * restrict k_t,
+    const float * restrict v_t,
+    const float * restrict g_t,
+    float beta_val,
+    float scale,
+    uint32_t S_v
+) {
+    const uint32_t epv  = 128 / sizeof(float);
+    const uint32_t nvec = S_v / epv;
+    const uint32_t nloe = S_v % epv;
+
+    HVX_Vector vq[4];
+    HVX_Vector vk[4];
+    HVX_Vector vg[4];
+
+    static const float kInf    = INFINITY;
+    static const float kMaxExp = 88.7228f;
+    const HVX_Vector max_exp = hvx_vec_splat_f32(kMaxExp);
+    const HVX_Vector inf     = hvx_vec_splat_f32(kInf);
+
+    for (uint32_t i = 0; i < nvec; ++i) {
+        vq[i] = hvx_vmemu(q_t + i * epv);
+        vk[i] = hvx_vmemu(k_t + i * epv);
+        vg[i] = hvx_vec_exp_f32_guard(hvx_vmemu(g_t + i * epv), max_exp, inf);
+    }
+    if (nloe) {
+        vq[nvec] = hvx_vmemu(q_t + nvec * epv);
+        vk[nvec] = hvx_vmemu(k_t + nvec * epv);
+        vg[nvec] = hvx_vec_exp_f32_guard(hvx_vmemu(g_t + nvec * epv), max_exp, inf);
+    }
+
+    const HVX_Vector vbeta  = hvx_vec_splat_f32(beta_val);
+    const HVX_Vector vscale = hvx_vec_splat_f32(scale);
+
+    float delta[8] __attribute__((aligned(128)));
+
+    uint32_t j = 0;
+    for (; j + 8 <= S_v; j += 8) {
+        float * row0 = s_work + (uint64_t) (j + 0) * S_v;
+        float * row1 = s_work + (uint64_t) (j + 1) * S_v;
+        float * row2 = s_work + (uint64_t) (j + 2) * S_v;
+        float * row3 = s_work + (uint64_t) (j + 3) * S_v;
+        float * row4 = s_work + (uint64_t) (j + 4) * S_v;
+        float * row5 = s_work + (uint64_t) (j + 5) * S_v;
+        float * row6 = s_work + (uint64_t) (j + 6) * S_v;
+        float * row7 = s_work + (uint64_t) (j + 7) * S_v;
+
+        HVX_Vector vsums = gdn_mul_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
+                                            vg, vk, S_v);
+
+        HVX_Vector vv_t   = hvx_vmemu(v_t + j);
+        HVX_Vector diff   = hvx_vec_sub_f32_f32(vv_t, vsums);
+        HVX_Vector vdelta = hvx_vec_mul_f32_f32(diff, vbeta);
+        hvx_vec_store_u(delta, 8 * sizeof(float), vdelta);
+
+        HVX_Vector vattn = gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
+                                                   vk, delta, vq, S_v);
+
+        HVX_Vector res_attn = hvx_vec_mul_f32_f32(vattn, vscale);
+        hvx_vec_store_u(attn_out + j, 8 * sizeof(float), res_attn);
+    }
+    for (; j + 4 <= S_v; j += 4) {
+        float * row0 = s_work + (uint64_t) (j + 0) * S_v;
+        float * row1 = s_work + (uint64_t) (j + 1) * S_v;
+        float * row2 = s_work + (uint64_t) (j + 2) * S_v;
+        float * row3 = s_work + (uint64_t) (j + 3) * S_v;
+
+        HVX_Vector vsums = gdn_mul_dot4_f32(row0, row1, row2, row3, vg, vk, S_v);
+
+        HVX_Vector vv_t   = hvx_vmemu(v_t + j);
+        HVX_Vector diff   = hvx_vec_sub_f32_f32(vv_t, vsums);
+        HVX_Vector vdelta = hvx_vec_mul_f32_f32(diff, vbeta);
+        hvx_vec_store_u(delta, 4 * sizeof(float), vdelta);
+
+        HVX_Vector vattn = gdn_add_scaled_dot4_f32(row0, row1, row2, row3, vk, delta, vq, S_v);
+
+        HVX_Vector res_attn = hvx_vec_mul_f32_f32(vattn, vscale);
+        hvx_vec_store_u(attn_out + j, 4 * sizeof(float), res_attn);
+    }
+    for (; j < S_v; ++j) {
+        float * row = s_work + (uint64_t) j * S_v;
+        HVX_Vector vsum = gdn_mul_dot_f32(row, vg, vk, S_v);
+        HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
+        HVX_Vector vdj  = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), vbeta);
+        HVX_Vector vres = gdn_add_scaled_dot_f32(row, vk, vdj, vq, S_v);
+        attn_out[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale));
+    }
+}
+
+static inline void gdn_step_scalar_f32(
+    float * restrict s_work,
+    float * restrict attn_out,
+    const float * restrict q_t,
+    const float * restrict k_t,
+    const float * restrict v_t,
+    const float * restrict g_t,
+    float beta_val,
+    float scale,
+    uint32_t S_v
+) {
+    const uint32_t epv  = 128 / sizeof(float);
+    const uint32_t nvec = S_v / epv;
+    const uint32_t nloe = S_v % epv;
+
+    HVX_Vector vq[4];
+    HVX_Vector vk[4];
+
+    for (uint32_t i = 0; i < nvec; ++i) {
+        vq[i] = hvx_vmemu(q_t + i * epv);
+        vk[i] = hvx_vmemu(k_t + i * epv);
+    }
+    if (nloe) {
+        vq[nvec] = hvx_vmemu(q_t + nvec * epv);
+        vk[nvec] = hvx_vmemu(k_t + nvec * epv);
+    }
+
+    const float gate       = expf(g_t[0]);
+    const HVX_Vector vgate = hvx_vec_splat_f32(gate);
+    const HVX_Vector vbeta = hvx_vec_splat_f32(beta_val);
+    const HVX_Vector vscale = hvx_vec_splat_f32(scale);
+
+    float delta[8] __attribute__((aligned(128)));
+
+    uint32_t j = 0;
+    for (; j + 8 <= S_v; j += 8) {
+        float * row0 = s_work + (uint64_t) (j + 0) * S_v;
+        float * row1 = s_work + (uint64_t) (j + 1) * S_v;
+        float * row2 = s_work + (uint64_t) (j + 2) * S_v;
+        float * row3 = s_work + (uint64_t) (j + 3) * S_v;
+        float * row4 = s_work + (uint64_t) (j + 4) * S_v;
+        float * row5 = s_work + (uint64_t) (j + 5) * S_v;
+        float * row6 = s_work + (uint64_t) (j + 6) * S_v;
+        float * row7 = s_work + (uint64_t) (j + 7) * S_v;
+
+        HVX_Vector vsums = gdn_mul_scalar_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
+                                                   vgate, vk, S_v);
+
+        HVX_Vector vv_t   = hvx_vmemu(v_t + j);
+        HVX_Vector diff   = hvx_vec_sub_f32_f32(vv_t, vsums);
+        HVX_Vector vdelta = hvx_vec_mul_f32_f32(diff, vbeta);
+        hvx_vec_store_u(delta, 8 * sizeof(float), vdelta);
+
+        HVX_Vector vattn = gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
+                                                   vk, delta, vq, S_v);
+
+        HVX_Vector res_attn = hvx_vec_mul_f32_f32(vattn, vscale);
+        hvx_vec_store_u(attn_out + j, 8 * sizeof(float), res_attn);
+    }
+    for (; j + 4 <= S_v; j += 4) {
+        float * row0 = s_work + (uint64_t) (j + 0) * S_v;
+        float * row1 = s_work + (uint64_t) (j + 1) * S_v;
+        float * row2 = s_work + (uint64_t) (j + 2) * S_v;
+        float * row3 = s_work + (uint64_t) (j + 3) * S_v;
+
+        HVX_Vector vsums = gdn_mul_scalar_dot4_f32(row0, row1, row2, row3, vgate, vk, S_v);
+
+        HVX_Vector vv_t   = hvx_vmemu(v_t + j);
+        HVX_Vector diff   = hvx_vec_sub_f32_f32(vv_t, vsums);
+        HVX_Vector vdelta = hvx_vec_mul_f32_f32(diff, vbeta);
+        hvx_vec_store_u(delta, 4 * sizeof(float), vdelta);
+
+        HVX_Vector vattn = gdn_add_scaled_dot4_f32(row0, row1, row2, row3, vk, delta, vq, S_v);
+
+        HVX_Vector res_attn = hvx_vec_mul_f32_f32(vattn, vscale);
+        hvx_vec_store_u(attn_out + j, 4 * sizeof(float), res_attn);
+    }
+    for (; j < S_v; ++j) {
+        float * row = s_work + (uint64_t) j * S_v;
+        HVX_Vector vsum = gdn_mul_scalar_dot_f32(row, gate, vk, S_v);
+        HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
+        HVX_Vector vdj  = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), vbeta);
+        HVX_Vector vres = gdn_add_scaled_dot_f32(row, vk, vdj, vq, S_v);
+        attn_out[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale));
+    }
 }

 static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, void * data) {
     struct htp_gdn_context * gctx = (struct htp_gdn_context *) data;
     struct htp_ops_context * octx = gctx->octx;
+    const struct htp_gdn_kernel_params * kparams = gctx->kparams;

     const struct htp_tensor * q     = octx->src[0];
     const struct htp_tensor * k     = octx->src[1];
@@ -585,66 +763,55 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
     const struct htp_tensor * state = octx->src[5];
     const struct htp_tensor * dst   = octx->dst;

-    const uint32_t S_v      = v->ne[0];
-    const uint32_t H        = v->ne[1];
-    const uint32_t n_tokens = v->ne[2];
-    const uint32_t n_seqs   = v->ne[3];
-    const uint32_t K        = octx->op_params[0];
-
-    const uint32_t row_end = gctx->row_start + gctx->nrows;
+    const uint32_t S_v      = kparams->S_v;
+    const uint32_t H        = kparams->H;
+    const uint32_t n_tokens = kparams->n_tokens;
+    const uint32_t n_seqs   = kparams->n_seqs;
+    const uint32_t K        = kparams->K;
+    const uint32_t row_end  = gctx->row_start + gctx->nrows;

     if (ith >= gctx->nrows) {
         return;
     }

-    const uint32_t rq3 = n_seqs / q->ne[3];
-    const uint32_t rk3 = n_seqs / k->ne[3];
-    const float scale = 1.0f / sqrtf((float) S_v);
-
+    const struct htp_tensor * dst_cache = octx->dsts[1];
+    const float scale = kparams->scale;
     float * dst_base       = (float *) (uintptr_t) dst->data;
-    float * state_out_base = dst_base + (uint64_t) S_v * H * n_tokens * n_seqs;
-    const float * state_in_base = (const float *) (uintptr_t) state->data;
-
-    const bool kda = (g->ne[0] == S_v);
-    float local_gate[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_q[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_k[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_sums[32] __attribute__((aligned(128)));
-
-    dma_queue * dma = octx->ctx->dma[ith];
-    size_t state_aligned = (size_t) S_v * S_v * sizeof(float);
-    state_aligned = (state_aligned + 127) & ~(size_t)127;
+    float * state_out_base = dst_cache ? (float *) (uintptr_t) dst_cache->data : (dst_base + S_v * H * n_tokens * n_seqs);
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+    const struct htp_gdn_vtcm_layout * layout = &gctx->layout;
     float * s_work[2];
-    s_work[0] = (float *) (gctx->vtcm_base + gctx->vtcm_per_thread * ith);
-    s_work[1] = s_work[0] + state_aligned / sizeof(float);
+    s_work[0] = (float *) (gctx->vtcm_base + layout->bytes_per_thread * ith);
+    s_work[1] = s_work[0] + layout->state_aligned / sizeof(float);

-    struct fastdiv_values fd_H = init_fastdiv_values(H);
-    struct fastdiv_values fd_q1 = init_fastdiv_values(q->ne[1]);
-    struct fastdiv_values fd_k1 = init_fastdiv_values(k->ne[1]);
-    struct fastdiv_values fd_rq3 = init_fastdiv_values(rq3);
-    struct fastdiv_values fd_rk3 = init_fastdiv_values(rk3);
+    const struct fastdiv_values * fd_H   = &kparams->div_H;
+    const struct fastdiv_values * fd_q1  = &kparams->div_q1;
+    const struct fastdiv_values * fd_k1  = &kparams->div_k1;
+    const struct fastdiv_values * fd_rq3 = &kparams->div_rq3;
+    const struct fastdiv_values * fd_rk3 = &kparams->div_rk3;

-    const uint64_t state_seq_stride = state->nb[3] / sizeof(float);
-    const uint64_t state_size_per_snap = (uint64_t) S_v * S_v * H * n_seqs;
+    const uint32_t state_seq_stride = kparams->state_seq_stride;
+    const uint64_t state_size_per_snap = (uint64_t) kparams->state_size_per_snap;
+    const dma_addr_t state_out_dma_base = dst_cache ? dst_cache->data : (dst->data + S_v * H * n_tokens * n_seqs * sizeof(float));

     uint32_t ir_prefetch = gctx->row_start + ith;
     int spad_idx = 0;

     // Prefetch preamble (up to 2 steps)
-    for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
-        const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
-        const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
-        const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
-        // final state lands in snapshot slot 0 (most-recent-first ordering)
-        float * ps_out = state_out_base + ((uint64_t) piv3 * H + piv1) * S_v * S_v;
+    for (int step = 0; step < 2 && ir_prefetch < row_end; step++) {
+        const uint32_t piv1 = fastmodulo(ir_prefetch, H, fd_H);
+        const uint32_t piv3 = fastdiv(ir_prefetch, fd_H);
+        dma_addr_t ps_in  = state->data + ((uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v) * sizeof(float);
+        dma_addr_t ps_out = state_out_dma_base + ((uint64_t) piv3 * H + piv1) * S_v * S_v * sizeof(float);

         // Push dummy write-back
-        dma_queue_push(dma, dma_make_ptr(ps_out, s_work[spad_idx]),
+        dma_queue_push(dma_q, dma_make_data(ps_out, s_work[spad_idx]),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), 0);

         // Push fetch
-        dma_queue_push(dma, dma_make_ptr(s_work[spad_idx], ps_in),
+        dma_queue_push(dma_q, dma_make_data(s_work[spad_idx], ps_in),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), S_v);

@@ -653,28 +820,26 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));

     int curr_spad_idx = 0;
     for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
-        dma_queue_pop(dma);
-        dma_queue_pop(dma);
+        dma_queue_pop(dma_q);
+        dma_queue_pop(dma_q);

         float * s_work_curr = s_work[curr_spad_idx];

-        const uint32_t iv1 = fastmodulo(ir, H, &fd_H);
-        const uint32_t iv3 = fastdiv(ir, &fd_H);
-
-        const uint32_t iq1 = fastmodulo(iv1, q->ne[1], &fd_q1);
-        const uint32_t ik1 = fastmodulo(iv1, k->ne[1], &fd_k1);
-        const uint32_t iq3 = fastdiv(iv3, &fd_rq3);
-        const uint32_t ik3 = fastdiv(iv3, &fd_rk3);
+        const uint32_t iv1 = fastmodulo(ir, H, fd_H);
+        const uint32_t iv3 = fastdiv(ir, fd_H);

-        // final state lands in snapshot slot 0 (most-recent-first ordering)
-        float * s_out = state_out_base + ((uint64_t) iv3 * H + iv1) * S_v * S_v;
+        const uint32_t iq1 = fastmodulo(iv1, q->ne[1], fd_q1);
+        const uint32_t ik1 = fastmodulo(iv1, k->ne[1], fd_k1);
+        const uint32_t iq3 = fastdiv(iv3, fd_rq3);
+        const uint32_t ik3 = fastdiv(iv3, fd_rk3);

+        dma_addr_t s_out  = state_out_dma_base + ((uint64_t) iv3 * H + iv1) * S_v * S_v * sizeof(float);
         float * attn_data = dst_base + ((uint64_t) iv3 * n_tokens * H + iv1) * S_v;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t t = 0; t < n_tokens; ++t) {
             const float * q_t = (const float *) ((const uint8_t *) (uintptr_t) q->data +
                     (uint64_t) iq3 * q->nb[3] + (uint64_t) t * q->nb[2] + (uint64_t) iq1 * q->nb[1]);
@@ -687,146 +852,36 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
             const float beta_val = *(const float *) ((const uint8_t *) (uintptr_t) beta->data +
                     (uint64_t) iv3 * beta->nb[3] + (uint64_t) t * beta->nb[2] + (uint64_t) iv1 * beta->nb[1]);

-            hvx_copy_f32_au((uint8_t *) local_q, (const uint8_t *) q_t, S_v);
-            hvx_copy_f32_au((uint8_t *) local_k, (const uint8_t *) k_t, S_v);
-
-            if (kda) {
-                hvx_exp_f32((uint8_t *) local_gate, (const uint8_t *) g_t, S_v, false);
-
-                uint32_t j = 0;
-                for (; j + 8 <= S_v; j += 8) {
-                    float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                    float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                    float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                    float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                    float * row4 = s_work_curr + (uint64_t) (j + 4) * S_v;
-                    float * row5 = s_work_curr + (uint64_t) (j + 5) * S_v;
-                    float * row6 = s_work_curr + (uint64_t) (j + 6) * S_v;
-                    float * row7 = s_work_curr + (uint64_t) (j + 7) * S_v;
-                    gdn_mul_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                     local_gate, local_k, S_v, local_sums);
-
-                    float local_delta_b[32] __attribute__((aligned(128)));
-                    HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                    HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                    HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                    hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                    gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                            local_k, local_delta_b, local_q, S_v, local_sums);
-
-                    HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                    hvx_vec_store_u(attn_data + j, 8 * sizeof(float), res_attn);
-                }
-                for (; j + 4 <= S_v; j += 4) {
-                    float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                    float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                    float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                    float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                    gdn_mul_dot4_f32(row0, row1, row2, row3, local_gate, local_k, S_v, local_sums);
-
-                    float local_delta_b[32] __attribute__((aligned(128)));
-                    HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                    HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                    HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                    hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                    gdn_add_scaled_dot4_f32(row0, row1, row2, row3, local_k, local_delta_b, local_q, S_v, local_sums);
-
-                    HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                    hvx_vec_store_u(attn_data + j, 4 * sizeof(float), res_attn);
-                }
-                HVX_Vector vscale_splat = hvx_vec_splat_f32(scale);
-                for (; j < S_v; ++j) {
-                    float * row = s_work_curr + (uint64_t) j * S_v;
-                    HVX_Vector vsum = gdn_mul_dot_f32(row, local_gate, local_k, S_v);
-                    HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
-                    HVX_Vector vdj = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), hvx_vec_splat_f32(beta_val));
-                    HVX_Vector vres = gdn_add_scaled_dot_f32(row, local_k, vdj, local_q, S_v);
-                    attn_data[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale_splat));
-                }
+            if (kparams->kda) {
+                gdn_step_kda_f32(s_work_curr, attn_data, q_t, k_t, v_t, g_t, beta_val, scale, S_v);
             } else {
-                const float gate = expf(g_t[0]);
-                uint32_t j = 0;
-                for (; j + 8 <= S_v; j += 8) {
-                    float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                    float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                    float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                    float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                    float * row4 = s_work_curr + (uint64_t) (j + 4) * S_v;
-                    float * row5 = s_work_curr + (uint64_t) (j + 5) * S_v;
-                    float * row6 = s_work_curr + (uint64_t) (j + 6) * S_v;
-                    float * row7 = s_work_curr + (uint64_t) (j + 7) * S_v;
-                    gdn_mul_scalar_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                            gate, local_k, S_v, local_sums);
-
-                    float local_delta_b[32] __attribute__((aligned(128)));
-                    HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                    HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                    HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                    hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                    gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                            local_k, local_delta_b, local_q, S_v, local_sums);
-
-                    HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                    hvx_vec_store_u(attn_data + j, 8 * sizeof(float), res_attn);
-                }
-                for (; j + 4 <= S_v; j += 4) {
-                    float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                    float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                    float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                    float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                    gdn_mul_scalar_dot4_f32(row0, row1, row2, row3, gate, local_k, S_v, local_sums);
-
-                    float local_delta_b[32] __attribute__((aligned(128)));
-                    HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                    HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                    HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                    hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                    gdn_add_scaled_dot4_f32(row0, row1, row2, row3, local_k, local_delta_b, local_q, S_v, local_sums);
-
-                    HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                    hvx_vec_store_u(attn_data + j, 4 * sizeof(float), res_attn);
-                }
-                HVX_Vector vscale_splat = hvx_vec_splat_f32(scale);
-                for (; j < S_v; ++j) {
-                    float * row = s_work_curr + (uint64_t) j * S_v;
-                    HVX_Vector vsum = gdn_mul_scalar_dot_f32(row, gate, local_k, S_v);
-                    HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
-                    HVX_Vector vdj = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), hvx_vec_splat_f32(beta_val));
-                    HVX_Vector vres = gdn_add_scaled_dot_f32(row, local_k, vdj, local_q, S_v);
-                    attn_data[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale_splat));
-                }
+                gdn_step_scalar_f32(s_work_curr, attn_data, q_t, k_t, v_t, g_t, beta_val, scale, S_v);
             }

             if (K > 1) {
-                // snapshot slot mapping: slot 0 = most recent state, slot s = s tokens back.
                 const int64_t target_slot = (int64_t) n_tokens - 1 - (int64_t) t;
-                if (target_slot >= 0 && target_slot < (int64_t) K) {
+                if (target_slot > 0 && target_slot < (int64_t) K) {
                     float * curr_state_o = state_out_base + (uint64_t) target_slot * state_size_per_snap + ((uint64_t) iv3 * H + iv1) * S_v * S_v;
-                    if (curr_state_o != s_out) {
-                        hvx_copy_f32_uu((uint8_t *) curr_state_o, (const uint8_t *) s_work_curr, S_v * S_v);
-                    }
+                    hvx_copy_f32_uu((uint8_t *) curr_state_o, (const uint8_t *) s_work_curr, S_v * S_v);
                 }
             }

             attn_data += (uint64_t) S_v * H;
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         // Push real write-back
-        dma_queue_push(dma, dma_make_ptr(s_out, s_work_curr),
+        dma_queue_push(dma_q, dma_make_data(s_out, s_work_curr),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), S_v);

         // Prefetch next block (if any)
         if (ir_prefetch < row_end) {
-            const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
-            const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
-            const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
+            const uint32_t piv1 = fastmodulo(ir_prefetch, H, fd_H);
+            const uint32_t piv3 = fastdiv(ir_prefetch, fd_H);
+            dma_addr_t ps_in = state->data + ((uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v) * sizeof(float);

-            dma_queue_push(dma, dma_make_ptr(s_work[spad_idx], ps_in),
+            dma_queue_push(dma_q, dma_make_data(s_work[spad_idx], ps_in),
                            S_v * sizeof(float), S_v * sizeof(float),
                            S_v * sizeof(float), S_v);

@@ -836,14 +891,13 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo

         curr_spad_idx ^= 1;
     }
-    dma_queue_flush(dma);
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
+    dma_queue_flush(dma_q);
 }

-
 static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, void * data) {
     struct htp_gdn_context * gctx = (struct htp_gdn_context *) data;
     struct htp_ops_context * octx = gctx->octx;
+    const struct htp_gdn_kernel_params * kparams = gctx->kparams;

     const struct htp_tensor * q     = octx->src[0];
     const struct htp_tensor * k     = octx->src[1];
@@ -853,63 +907,51 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
     const struct htp_tensor * state = octx->src[5];
     const struct htp_tensor * dst   = octx->dst;

-    const uint32_t S_v      = v->ne[0];
-    const uint32_t H        = v->ne[1];
-    const uint32_t n_seqs   = v->ne[3];
-
-    const uint32_t row_end = gctx->row_start + gctx->nrows;
+    const uint32_t S_v      = kparams->S_v;
+    const uint32_t H        = kparams->H;
+    const uint32_t n_seqs   = kparams->n_seqs;
+    const uint32_t row_end  = gctx->row_start + gctx->nrows;

     if (ith >= gctx->nrows) {
         return;
     }

-    const uint32_t rq3 = n_seqs / q->ne[3];
-    const uint32_t rk3 = n_seqs / k->ne[3];
-    const float scale = 1.0f / sqrtf((float) S_v);
+    const struct htp_tensor * dst_cache = octx->dsts[1];
+    const float scale = kparams->scale;
+    float * dst_base  = (float *) (uintptr_t) dst->data;

-    float * dst_base       = (float *) (uintptr_t) dst->data;
-    float * state_out_base = dst_base + (uint64_t) S_v * H * n_seqs;
-    const float * state_in_base = (const float *) (uintptr_t) state->data;
-
-    const bool kda = (g->ne[0] == S_v);
-    float local_gate[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_q[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_k[HTP_GDN_MAX_SV] __attribute__((aligned(128)));
-    float local_sums[32] __attribute__((aligned(128)));
-
-    dma_queue * dma = octx->ctx->dma[ith];
-    size_t state_aligned = (size_t) S_v * S_v * sizeof(float);
-    state_aligned = (state_aligned + 127) & ~(size_t)127;
+    dma_queue * dma_q = octx->ctx->dma[ith];
+    const struct htp_gdn_vtcm_layout * layout = &gctx->layout;
     float * s_work[2];
-    s_work[0] = (float *) (gctx->vtcm_base + gctx->vtcm_per_thread * ith);
-    s_work[1] = s_work[0] + state_aligned / sizeof(float);
+    s_work[0] = (float *) (gctx->vtcm_base + layout->bytes_per_thread * ith);
+    s_work[1] = s_work[0] + layout->state_aligned / sizeof(float);

-    struct fastdiv_values fd_H = init_fastdiv_values(H);
-    struct fastdiv_values fd_q1 = init_fastdiv_values(q->ne[1]);
-    struct fastdiv_values fd_k1 = init_fastdiv_values(k->ne[1]);
-    struct fastdiv_values fd_rq3 = init_fastdiv_values(rq3);
-    struct fastdiv_values fd_rk3 = init_fastdiv_values(rk3);
+    const struct fastdiv_values * fd_H   = &kparams->div_H;
+    const struct fastdiv_values * fd_q1  = &kparams->div_q1;
+    const struct fastdiv_values * fd_k1  = &kparams->div_k1;
+    const struct fastdiv_values * fd_rq3 = &kparams->div_rq3;
+    const struct fastdiv_values * fd_rk3 = &kparams->div_rk3;

-    const uint64_t state_seq_stride = state->nb[3] / sizeof(float);
+    const uint32_t state_seq_stride = kparams->state_seq_stride;
+    const dma_addr_t state_out_dma_base = dst_cache ? dst_cache->data : (dst->data + S_v * H * n_seqs * sizeof(float));

     uint32_t ir_prefetch = gctx->row_start + ith;
     int spad_idx = 0;

     // Prefetch preamble (up to 2 steps)
-    for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
-        const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
-        const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
-        const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
-        // final state lands in snapshot slot 0 (most-recent-first ordering)
-        float * ps_out = state_out_base + ((uint64_t) piv3 * H + piv1) * S_v * S_v;
+    for (int step = 0; step < 2 && ir_prefetch < row_end; step++) {
+        const uint32_t piv1 = fastmodulo(ir_prefetch, H, fd_H);
+        const uint32_t piv3 = fastdiv(ir_prefetch, fd_H);
+        dma_addr_t ps_in  = state->data + ((uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v) * sizeof(float);
+        dma_addr_t ps_out = state_out_dma_base + ((uint64_t) piv3 * H + piv1) * S_v * S_v * sizeof(float);

         // Push dummy write-back
-        dma_queue_push(dma, dma_make_ptr(ps_out, s_work[spad_idx]),
+        dma_queue_push(dma_q, dma_make_data(ps_out, s_work[spad_idx]),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), 0);

         // Push fetch
-        dma_queue_push(dma, dma_make_ptr(s_work[spad_idx], ps_in),
+        dma_queue_push(dma_q, dma_make_data(s_work[spad_idx], ps_in),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), S_v);

@@ -918,26 +960,23 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
     }

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));

     int curr_spad_idx = 0;
     for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
-        dma_queue_pop(dma);
-        dma_queue_pop(dma);
+        dma_queue_pop(dma_q);
+        dma_queue_pop(dma_q);

         float * s_work_curr = s_work[curr_spad_idx];

-        const uint32_t iv1 = fastmodulo(ir, H, &fd_H);
-        const uint32_t iv3 = fastdiv(ir, &fd_H);
+        const uint32_t iv1 = fastmodulo(ir, H, fd_H);
+        const uint32_t iv3 = fastdiv(ir, fd_H);

-        const uint32_t iq1 = fastmodulo(iv1, q->ne[1], &fd_q1);
-        const uint32_t ik1 = fastmodulo(iv1, k->ne[1], &fd_k1);
-        const uint32_t iq3 = fastdiv(iv3, &fd_rq3);
-        const uint32_t ik3 = fastdiv(iv3, &fd_rk3);
-
-        // final state lands in snapshot slot 0 (most-recent-first ordering)
-        float * s_out = state_out_base + ((uint64_t) iv3 * H + iv1) * S_v * S_v;
+        const uint32_t iq1 = fastmodulo(iv1, q->ne[1], fd_q1);
+        const uint32_t ik1 = fastmodulo(iv1, k->ne[1], fd_k1);
+        const uint32_t iq3 = fastdiv(iv3, fd_rq3);
+        const uint32_t ik3 = fastdiv(iv3, fd_rk3);

+        dma_addr_t s_out  = state_out_dma_base + ((uint64_t) iv3 * H + iv1) * S_v * S_v * sizeof(float);
         float * attn_data = dst_base + ((uint64_t) iv3 * H + iv1) * S_v;

         const float * q_t = (const float *) ((const uint8_t *) (uintptr_t) q->data +
@@ -951,132 +990,26 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
         const float beta_val = *(const float *) ((const uint8_t *) (uintptr_t) beta->data +
                 (uint64_t) iv3 * beta->nb[3] + (uint64_t) iv1 * beta->nb[1]);

-        hvx_copy_f32_au((uint8_t *) local_q, (const uint8_t *) q_t, S_v);
-        hvx_copy_f32_au((uint8_t *) local_k, (const uint8_t *) k_t, S_v);
-
-        if (kda) {
-            hvx_exp_f32((uint8_t *) local_gate, (const uint8_t *) g_t, S_v, false);
-
-            uint32_t j = 0;
-            for (; j + 8 <= S_v; j += 8) {
-                float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                float * row4 = s_work_curr + (uint64_t) (j + 4) * S_v;
-                float * row5 = s_work_curr + (uint64_t) (j + 5) * S_v;
-                float * row6 = s_work_curr + (uint64_t) (j + 6) * S_v;
-                float * row7 = s_work_curr + (uint64_t) (j + 7) * S_v;
-                gdn_mul_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                 local_gate, local_k, S_v, local_sums);
-
-                float local_delta_b[32] __attribute__((aligned(128)));
-                HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                        local_k, local_delta_b, local_q, S_v, local_sums);
-
-                HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                hvx_vec_store_u(attn_data + j, 8 * sizeof(float), res_attn);
-            }
-            for (; j + 4 <= S_v; j += 4) {
-                float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                gdn_mul_dot4_f32(row0, row1, row2, row3, local_gate, local_k, S_v, local_sums);
-
-                float local_delta_b[32] __attribute__((aligned(128)));
-                HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                gdn_add_scaled_dot4_f32(row0, row1, row2, row3, local_k, local_delta_b, local_q, S_v, local_sums);
-
-                HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                hvx_vec_store_u(attn_data + j, 4 * sizeof(float), res_attn);
-            }
-            HVX_Vector vscale_splat = hvx_vec_splat_f32(scale);
-            for (; j < S_v; ++j) {
-                float * row = s_work_curr + (uint64_t) j * S_v;
-                HVX_Vector vsum = gdn_mul_dot_f32(row, local_gate, local_k, S_v);
-                HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
-                HVX_Vector vdj = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), hvx_vec_splat_f32(beta_val));
-                HVX_Vector vres = gdn_add_scaled_dot_f32(row, local_k, vdj, local_q, S_v);
-                attn_data[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale_splat));
-            }
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
+        if (kparams->kda) {
+            gdn_step_kda_f32(s_work_curr, attn_data, q_t, k_t, v_t, g_t, beta_val, scale, S_v);
         } else {
-            const float gate = expf(g_t[0]);
-            uint32_t j = 0;
-            for (; j + 8 <= S_v; j += 8) {
-                float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                float * row4 = s_work_curr + (uint64_t) (j + 4) * S_v;
-                float * row5 = s_work_curr + (uint64_t) (j + 5) * S_v;
-                float * row6 = s_work_curr + (uint64_t) (j + 6) * S_v;
-                float * row7 = s_work_curr + (uint64_t) (j + 7) * S_v;
-                gdn_mul_scalar_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                        gate, local_k, S_v, local_sums);
-
-                float local_delta_b[32] __attribute__((aligned(128)));
-                HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                gdn_add_scaled_dot8_f32(row0, row1, row2, row3, row4, row5, row6, row7,
-                                        local_k, local_delta_b, local_q, S_v, local_sums);
-
-                HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                hvx_vec_store_u(attn_data + j, 8 * sizeof(float), res_attn);
-            }
-            for (; j + 4 <= S_v; j += 4) {
-                float * row0 = s_work_curr + (uint64_t) (j + 0) * S_v;
-                float * row1 = s_work_curr + (uint64_t) (j + 1) * S_v;
-                float * row2 = s_work_curr + (uint64_t) (j + 2) * S_v;
-                float * row3 = s_work_curr + (uint64_t) (j + 3) * S_v;
-                gdn_mul_scalar_dot4_f32(row0, row1, row2, row3, gate, local_k, S_v, local_sums);
-
-                float local_delta_b[32] __attribute__((aligned(128)));
-                HVX_Vector vv_t = hvx_vmemu(v_t + j);
-                HVX_Vector v_local_sums = hvx_vmem(local_sums);
-                HVX_Vector diff = hvx_vec_sub_f32_f32(vv_t, v_local_sums);
-                hvx_vmem(local_delta_b) = hvx_vec_mul_f32_f32(diff, hvx_vec_splat_f32(beta_val));
-
-                gdn_add_scaled_dot4_f32(row0, row1, row2, row3, local_k, local_delta_b, local_q, S_v, local_sums);
-
-                HVX_Vector res_attn = hvx_vec_mul_f32_f32(hvx_vmem(local_sums), hvx_vec_splat_f32(scale));
-                hvx_vec_store_u(attn_data + j, 4 * sizeof(float), res_attn);
-            }
-            HVX_Vector vscale_splat = hvx_vec_splat_f32(scale);
-            for (; j < S_v; ++j) {
-                float * row = s_work_curr + (uint64_t) j * S_v;
-                HVX_Vector vsum = gdn_mul_scalar_dot_f32(row, gate, local_k, S_v);
-                HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]);
-                HVX_Vector vdj = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), hvx_vec_splat_f32(beta_val));
-                HVX_Vector vres = gdn_add_scaled_dot_f32(row, local_k, vdj, local_q, S_v);
-                attn_data[j] = hvx_vec_get_f32(hvx_vec_mul_f32_f32(vres, vscale_splat));
-            }
+            gdn_step_scalar_f32(s_work_curr, attn_data, q_t, k_t, v_t, g_t, beta_val, scale, S_v);
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         // Push real write-back
-        dma_queue_push(dma, dma_make_ptr(s_out, s_work_curr),
+        dma_queue_push(dma_q, dma_make_data(s_out, s_work_curr),
                        S_v * sizeof(float), S_v * sizeof(float),
                        S_v * sizeof(float), S_v);

         // Prefetch next block (if any)
         if (ir_prefetch < row_end) {
-            const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
-            const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
-            const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
+            const uint32_t piv1 = fastmodulo(ir_prefetch, H, fd_H);
+            const uint32_t piv3 = fastdiv(ir_prefetch, fd_H);
+            dma_addr_t ps_in = state->data + ((uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v) * sizeof(float);

-            dma_queue_push(dma, dma_make_ptr(s_work[spad_idx], ps_in),
+            dma_queue_push(dma_q, dma_make_data(s_work[spad_idx], ps_in),
                            S_v * sizeof(float), S_v * sizeof(float),
                            S_v * sizeof(float), S_v);

@@ -1086,11 +1019,9 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo

         curr_spad_idx ^= 1;
     }
-    dma_queue_flush(dma);
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
+    dma_queue_flush(dma_q);
 }

-
 int op_gated_delta_net(struct htp_ops_context * octx) {
     const struct htp_tensor * q     = octx->src[0];
     const struct htp_tensor * k     = octx->src[1];
@@ -1131,44 +1062,88 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    for (int i = 0; i < 5; i++) {
+        if (htp_tensor_is_extended(octx->src[i])) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
+    }
+    if (htp_tensor_is_extended(octx->dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+    if (octx->dsts[1]) {
+        const struct htp_tensor * dst_cache = octx->dsts[1];
+        if (dst_cache->type != HTP_TYPE_F32 || htp_tensor_is_extended(dst_cache)) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
     }

-    const uint32_t total_rows = H * n_seqs;
+    const struct htp_gdn_kernel_params * kparams = (const struct htp_gdn_kernel_params *) octx->kernel_params;
+    struct htp_gdn_kernel_params kparams_local;
+    if (!kparams || kparams->S_v == 0) {
+        const uint32_t rq3 = n_seqs / q->ne[3];
+        const uint32_t rk3 = n_seqs / k->ne[3];
+        const uint32_t total_rows = H * n_seqs;
+        uint32_t n_threads = (total_rows < octx->n_threads) ? total_rows : octx->n_threads;
+        if (n_threads == 0) {
+            n_threads = 1;
+        }
+
+        memset(&kparams_local, 0, sizeof(kparams_local));
+        kparams_local.n_threads           = n_threads;
+        kparams_local.S_v                 = S_v;
+        kparams_local.H                   = H;
+        kparams_local.n_tokens            = n_tokens;
+        kparams_local.n_seqs              = n_seqs;
+        kparams_local.K                   = K;
+        kparams_local.total_rows          = total_rows;
+        kparams_local.rows_per_thread     = (total_rows + n_threads - 1) / n_threads;
+        struct htp_gdn_vtcm_layout layout_local;
+        htp_gdn_vtcm_layout_build(&layout_local, S_v, n_threads);
+        kparams_local.state_aligned       = (uint32_t) layout_local.state_aligned;
+        kparams_local.vtcm_per_thread     = (uint32_t) layout_local.bytes_per_thread;
+        kparams_local.vtcm_size           = (uint32_t) layout_local.total_bytes;
+        kparams_local.kda                 = (g->ne[0] == S_v) ? 1 : 0;
+        kparams_local.scale               = 1.0f / sqrtf((float) S_v);
+        kparams_local.state_seq_stride    = (uint32_t) (state->nb[3] / sizeof(float));
+        kparams_local.state_size_per_snap = S_v * S_v * H * n_seqs;
+
+        kparams_local.div_H         = init_fastdiv_values(H);
+        kparams_local.div_q1        = init_fastdiv_values(q->ne[1]);
+        kparams_local.div_k1        = init_fastdiv_values(k->ne[1]);
+        kparams_local.div_rq3       = init_fastdiv_values(rq3);
+        kparams_local.div_rk3       = init_fastdiv_values(rk3);
+        kparams_local.div_n_threads = init_fastdiv_values(n_threads);
+
+        kparams = &kparams_local;
+    }

+    const uint32_t total_rows = kparams->total_rows;
     uint32_t row_start = 0;
     uint32_t nrows     = total_rows;

-    if (octx->ctx->mdev.count > 1) {
-        const uint32_t head_bytes = S_v * sizeof(float);
-        const uint32_t rows_per_chunk = (head_bytes > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(head_bytes, HEX_L2_LINE_SIZE)) : 1;
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0,
-                                                                             octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
-        row_start = range.start;
-        nrows     = range.count;
+    if (octx->op_params[1] != 0) {
+        row_start = octx->op_params[1];
+        nrows     = octx->op_params[2];
     }

     if (nrows == 0) {
         return HTP_STATUS_OK;
     }

-    const uint32_t n_threads = octx->n_threads;
+    const uint32_t n_threads = (nrows < kparams->n_threads) ? nrows : kparams->n_threads;

     struct htp_gdn_context gctx;
-    gctx.octx = octx;
-    gctx.row_start       = row_start;
-    gctx.nrows           = nrows;
-    gctx.rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
-    gctx.state_bytes = (size_t) S_v * S_v * sizeof(float);
-
-    size_t state_aligned = (size_t) S_v * S_v * sizeof(float);
-    state_aligned = (state_aligned + 127) & ~(size_t)127;
+    gctx.octx      = octx;
+    gctx.kparams   = kparams;
+    gctx.row_start = row_start;
+    gctx.nrows     = nrows;
+    gctx.vtcm_base = octx->ctx->vtcm_base;

-    assert(octx->ctx->vtcm_size >= 2 * state_aligned * n_threads);
+    htp_gdn_vtcm_layout_build(&gctx.layout, S_v, n_threads);

-    gctx.vtcm_base = octx->ctx->vtcm_base;
-    gctx.vtcm_per_thread = 2 * state_aligned;
+    if (gctx.layout.total_bytes > octx->ctx->vtcm_size) {
+        return HTP_STATUS_VTCM_TOO_SMALL;
+    }

     FARF(HIGH, "gated-delta-net-f32: q(%ux%ux%ux%u) k(%ux%ux%ux%u) v(%ux%ux%ux%u) state(%ux%ux%ux%u) -> (%ux%ux%ux%u) : "
          "vtcm-size %zu n_threads %u\n",
@@ -1177,7 +1152,7 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
          v->ne[0], v->ne[1], v->ne[2], v->ne[3],
          state->ne[0], state->ne[1], state->ne[2], state->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         gctx.vtcm_per_thread * octx->n_threads, octx->n_threads);
+         gctx.layout.total_bytes, n_threads);

     if (n_tokens == 1) {
         work_queue_run(octx->ctx->work_queue, gated_delta_net_f32_tg_thread, &gctx, n_threads);
diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h
new file mode 100644
index 000000000..fd703142e
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h
@@ -0,0 +1,297 @@
+#ifndef HTP_GATED_DELTA_NET_OPS_H
+#define HTP_GATED_DELTA_NET_OPS_H
+
+#include <stdint.h>
+#include <stddef.h>
+#include <stdbool.h>
+
+#include "hex-fastdiv.h"
+#include "hex-common.h"
+#include "htp-vtcm.h"
+
+#define HTP_GDN_MAX_SV     128
+#define HTP_GDN_CHUNK_SIZE 64
+
+#ifndef HMX_FP16_TILE_SIZE
+#define HMX_FP16_TILE_SIZE 2048
+#endif
+
+enum htp_gdn_kernel_type {
+    HTP_GDN_KERNEL_HVX_RECURRENT = 0,
+    HTP_GDN_KERNEL_HMX_CHUNKED   = 1,
+};
+
+struct htp_gdn_kernel_params {
+    uint8_t  kernel_type;
+    uint8_t  pipeline;
+    uint16_t chunk_size;
+    uint16_t n_chunks;
+    uint16_t n_heads_batch;
+
+    uint32_t n_threads;
+    uint32_t S_v;
+    uint32_t H;
+    uint32_t n_tokens;
+    uint32_t n_seqs;
+    uint32_t K;
+
+    uint32_t total_rows;
+    uint32_t row_start;
+    uint32_t nrows;
+    uint32_t rows_per_thread;
+
+    uint32_t kda;
+    uint32_t state_aligned;
+    uint32_t vtcm_per_thread;
+    uint32_t vtcm_size;
+    uint32_t state_seq_stride;
+    uint32_t state_size_per_snap;
+
+    float    scale;
+
+    struct fastdiv_values div_H;
+    struct fastdiv_values div_q1;
+    struct fastdiv_values div_k1;
+    struct fastdiv_values div_rq3;
+    struct fastdiv_values div_rk3;
+    struct fastdiv_values div_n_threads;
+};
+
+#if defined(__cplusplus)
+static_assert(sizeof(struct htp_gdn_kernel_params) <= 128, "htp_gdn_kernel_params is too large for kernel_params blob");
+#else
+_Static_assert(sizeof(struct htp_gdn_kernel_params) <= 128, "htp_gdn_kernel_params is too large for kernel_params blob");
+#endif
+
+struct htp_gdn_vtcm_layout {
+    size_t state_aligned;
+    size_t bytes_per_thread;
+    size_t total_bytes;
+};
+
+static inline void htp_gdn_vtcm_layout_build(
+    struct htp_gdn_vtcm_layout * layout,
+    uint32_t S_v,
+    uint32_t n_threads
+) {
+    size_t state_bytes = (size_t) S_v * S_v * sizeof(float);
+    layout->state_aligned = hex_round_up(state_bytes, 128);
+    layout->bytes_per_thread = 2 * layout->state_aligned;
+    layout->total_bytes = layout->bytes_per_thread * n_threads;
+}
+
+struct htp_gdn_hmx_vtcm_layout {
+    size_t off_s_state;
+    size_t off_s_f16;
+    size_t off_s_col_tiles;
+    size_t off_s_update_f32;
+    size_t off_s_update_tiles;
+
+    size_t off_q_f32[2];
+    size_t off_k_f32[2];
+    size_t off_v_f32[2];
+    size_t off_g_f32[2];
+    size_t off_b_f32[2];
+    size_t off_g_raw[2];
+    size_t off_b_raw[2];
+    size_t off_o_f32[2];
+
+    size_t off_v_inter_f32;
+    size_t off_o_inter_f32;
+    size_t off_o_intra_f32;
+    size_t off_k_f16;
+    size_t off_v_prime_f16;
+    size_t off_delta_f16;
+    size_t off_d_f16;
+
+    size_t off_q_row_tiles;
+    size_t off_q_prime_row_tiles;
+    size_t off_k_row_tiles;
+    size_t off_k_col_tiles;
+    size_t off_k_prime_row_tiles;
+    size_t off_k_col_tiles_64x128;
+    size_t off_kk_tiles;
+    size_t off_qk_tiles;
+    size_t off_v_inter_tiles;
+    size_t off_o_inter_tiles;
+    size_t off_inv_row_tiles;
+    size_t off_a_row_tiles;
+    size_t off_v_prime_col_tiles;
+    size_t off_delta_tiles;
+    size_t off_delta_col_tiles;
+    size_t off_o_intra_tiles;
+    size_t off_d_row_tiles;
+
+    size_t off_gamma;
+    size_t off_lambda_init;
+    size_t off_decay_m;
+    size_t off_decay_a;
+    size_t off_rows_kk;
+    size_t off_rows_qk;
+    size_t off_rows_inv;
+    size_t off_rows_a;
+
+    size_t off_thread_scratch;
+    size_t off_attn_rem;
+    size_t off_scales_1;
+
+    size_t state_f32_bytes;
+    size_t state_f16_bytes;
+    size_t state_tiles_bytes;
+    size_t dma_chunk_bytes;
+    size_t act_f16_bytes;
+    size_t tile_64xSv_bytes;
+    size_t tile_64x64_bytes;
+
+    uint32_t n_heads_batch;
+    uint32_t n_threads;
+    bool     pipeline;
+    size_t   total_bytes;
+};
+
+static inline void htp_gdn_hmx_vtcm_layout_build(
+    struct htp_gdn_hmx_vtcm_layout * L,
+    uint32_t S_v,
+    uint32_t chunk_size,
+    uint32_t n_heads_batch,
+    uint32_t n_threads,
+    bool     pipeline
+) {
+    memset(L, 0, sizeof(*L));
+    L->n_heads_batch = n_heads_batch;
+    L->n_threads     = n_threads;
+    L->pipeline      = pipeline;
+
+    const size_t bh = (size_t) n_heads_batch;
+    const size_t nth = (size_t) (n_threads > 0 ? n_threads : 1);
+
+    const size_t state_f32_sz   = hex_round_up(S_v * S_v * sizeof(float), 2048);
+    const size_t state_f16_sz   = hex_round_up(S_v * S_v * sizeof(__fp16), 2048);
+    const size_t n_sv_tiles     = S_v / 32;
+    const size_t state_tiles_sz = n_sv_tiles * n_sv_tiles * 2048;
+
+    const size_t dma_chunk_sz   = hex_round_up(chunk_size * S_v * sizeof(float), 2048);
+    const size_t dma_scalar_sz  = hex_round_up(chunk_size * sizeof(float), 128);
+
+    const size_t act_f16_sz     = hex_round_up(chunk_size * S_v * sizeof(__fp16), 2048);
+    const size_t tile_64xSv_sz  = 2 * n_sv_tiles * 2048;
+    const size_t tile_64x64_sz  = 4 * 2048;
+
+    const size_t decay_sz       = 64 * 64 * sizeof(__fp16);
+    const size_t row_vecs_sz    = 64 * 128;
+
+    L->state_f32_bytes   = state_f32_sz;
+    L->state_f16_bytes   = state_f16_sz;
+    L->state_tiles_bytes = state_tiles_sz;
+    L->dma_chunk_bytes   = dma_chunk_sz;
+    L->act_f16_bytes     = act_f16_sz;
+    L->tile_64xSv_bytes  = tile_64xSv_sz;
+    L->tile_64x64_bytes  = tile_64x64_sz;
+
+    size_t off = 0;
+
+    VTCM_LAYOUT_ALLOC(off, off_s_state,        bh * state_f32_sz);
+    VTCM_LAYOUT_ALLOC(off, off_s_f16,          bh * state_f16_sz);
+    VTCM_LAYOUT_ALLOC(off, off_s_col_tiles,    bh * state_tiles_sz);
+    VTCM_LAYOUT_ALLOC(off, off_s_update_f32,   bh * state_f32_sz);
+    VTCM_LAYOUT_ALLOC(off, off_s_update_tiles, bh * state_tiles_sz);
+
+    VTCM_LAYOUT_ALLOC(off, off_q_f32[0], bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_q_f32[1], bh * dma_chunk_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_k_f32[0], bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_k_f32[1], bh * dma_chunk_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_v_f32[0], bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_v_f32[1], bh * dma_chunk_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_g_f32[0], bh * dma_scalar_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_g_f32[1], bh * dma_scalar_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_b_f32[0], bh * dma_scalar_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_b_f32[1], bh * dma_scalar_sz, pipeline);
+    const size_t raw_gb_sz = hex_round_up(bh * chunk_size * sizeof(float), 128);
+    VTCM_LAYOUT_ALLOC(off, off_g_raw[0], raw_gb_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_g_raw[1], raw_gb_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_b_raw[0], raw_gb_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_b_raw[1], raw_gb_sz, pipeline);
+    VTCM_LAYOUT_ALLOC(off, off_o_f32[0], bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_o_f32[1], bh * dma_chunk_sz, pipeline);
+
+    VTCM_LAYOUT_ALLOC(off, off_v_inter_f32, bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC(off, off_o_inter_f32, bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC(off, off_o_intra_f32, bh * dma_chunk_sz);
+    VTCM_LAYOUT_ALLOC(off, off_k_f16,       bh * act_f16_sz);
+    VTCM_LAYOUT_ALLOC(off, off_v_prime_f16, bh * act_f16_sz);
+    VTCM_LAYOUT_ALLOC(off, off_delta_f16,   bh * act_f16_sz);
+    VTCM_LAYOUT_ALLOC(off, off_d_f16,       bh * act_f16_sz);
+
+    VTCM_LAYOUT_ALLOC(off, off_q_row_tiles,        bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_q_prime_row_tiles,  bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_k_row_tiles,        bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_k_col_tiles,        bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_k_prime_row_tiles,  bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_k_col_tiles_64x128, bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_kk_tiles,           bh * tile_64x64_sz);
+    VTCM_LAYOUT_ALLOC(off, off_qk_tiles,           bh * tile_64x64_sz);
+    VTCM_LAYOUT_ALLOC(off, off_v_inter_tiles,      bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_o_inter_tiles,      bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_inv_row_tiles,      bh * tile_64x64_sz);
+    VTCM_LAYOUT_ALLOC(off, off_a_row_tiles,        bh * tile_64x64_sz);
+    VTCM_LAYOUT_ALLOC(off, off_v_prime_col_tiles,  bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_delta_tiles,        bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_delta_col_tiles,    bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_o_intra_tiles,      bh * tile_64xSv_sz);
+    VTCM_LAYOUT_ALLOC(off, off_d_row_tiles,        bh * tile_64xSv_sz);
+
+    VTCM_LAYOUT_ALLOC(off, off_gamma,       bh * hex_round_up(chunk_size * sizeof(float), 128));
+    VTCM_LAYOUT_ALLOC(off, off_lambda_init, bh * hex_round_up(chunk_size * sizeof(float), 128));
+    VTCM_LAYOUT_ALLOC(off, off_decay_m,     bh * decay_sz);
+    VTCM_LAYOUT_ALLOC(off, off_decay_a,     bh * decay_sz);
+    VTCM_LAYOUT_ALLOC(off, off_rows_kk,     bh * row_vecs_sz);
+    VTCM_LAYOUT_ALLOC(off, off_rows_qk,     bh * row_vecs_sz);
+    VTCM_LAYOUT_ALLOC(off, off_rows_inv,    bh * row_vecs_sz);
+    VTCM_LAYOUT_ALLOC(off, off_rows_a,      bh * row_vecs_sz);
+
+    const size_t thread_scratch_sz = 64 * 128;
+    VTCM_LAYOUT_ALLOC(off, off_thread_scratch, nth * thread_scratch_sz);
+    VTCM_LAYOUT_ALLOC(off, off_attn_rem,       nth * (128 * sizeof(float)));
+    VTCM_LAYOUT_ALLOC(off, off_scales_1,       256);
+
+    L->total_bytes = off;
+}
+
+static inline bool htp_gdn_hmx_solve_layout(
+    struct htp_gdn_hmx_vtcm_layout * layout_out,
+    uint32_t S_v,
+    uint32_t chunk_size,
+    uint32_t total_rows,
+    size_t   vtcm_budget,
+    uint32_t n_threads,
+    bool     pipeline,
+    uint32_t * n_heads_batch_out
+) {
+    uint32_t max_batch = 8;
+    if (max_batch > total_rows) {
+        max_batch = total_rows;
+    }
+    if (max_batch > n_threads) {
+        max_batch = n_threads;
+    }
+    static const uint32_t candidates[] = { 8, 6, 4, 2, 1 };
+    for (size_t i = 0; i < sizeof(candidates) / sizeof(candidates[0]); ++i) {
+        uint32_t bh = candidates[i];
+        if (bh > max_batch) {
+            continue;
+        }
+        struct htp_gdn_hmx_vtcm_layout L;
+        htp_gdn_hmx_vtcm_layout_build(&L, S_v, chunk_size, bh, n_threads, pipeline);
+        if (L.total_bytes <= vtcm_budget) {
+            *layout_out = L;
+            *n_heads_batch_out = bh;
+            return true;
+        }
+    }
+    if (pipeline) {
+        return htp_gdn_hmx_solve_layout(layout_out, S_v, chunk_size, total_rows, vtcm_budget, n_threads, false, n_heads_batch_out);
+    }
+    return false;
+}
+
+#endif // HTP_GATED_DELTA_NET_OPS_H
diff --git a/ggml/src/ggml-hexagon/htp/get-rows-ops.c b/ggml/src/ggml-hexagon/htp/get-rows-ops.c
index 958ecac3f..f51e00c15 100644
--- a/ggml/src/ggml-hexagon/htp/get-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/get-rows-ops.c
@@ -11,6 +11,7 @@
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "hex-common.h"
+#include "dma-queue.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
@@ -59,140 +60,139 @@ struct get_rows_context {
                                                \
     const uint32_t nr = ne10 * ne11 * ne12;

-#define GET_ROWS_THREAD_ST_FN(IDX_TYPE)                                                                                \
-static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                            \
-    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                 \
-    struct htp_ops_context * octx = grctx->octx;                                                                       \
-    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                \
-    get_rows_preamble;                                                                                                 \
-    const uint32_t dr  = grctx->tasks_per_thread;                                                                      \
-    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                 \
-    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                     \
-        return;                                                                                                        \
-    }                                                                                                                  \
-    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                              \
-    const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00);                                 \
-    dma_queue * dma_queue = octx->ctx->dma[ith];                                                                       \
-    for (uint32_t i = ir0; i < ir1; ++i) {                                                                             \
-        const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);                                                      \
-        const uint32_t rem = i - i12 * ne11 * ne10;                                                                    \
-        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                         \
-        const uint32_t i10 = rem - i11 * ne10;                                                                         \
-        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12);           \
-        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                      \
-        assert(i01 < ne01);                                                                                            \
-        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                         \
-        const uint32_t i02 = i11 - q02 * ne02;                                                                         \
-        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                         \
-        const uint32_t i03 = i12 - q03 * ne03;                                                                         \
-        const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03;                                \
-        const uintptr_t dst_ptr  = octx->dst->data    + i10*nb1  + i11*nb2  + i12*nb3;                                 \
-        while (!dma_queue_push(dma_queue, dma_make_ptr((void *)dst_ptr, (const void *)src0_ptr), nb1, nb01,            \
-                               row_size_bytes, 1)) {                                                                   \
-            dma_queue_pop(dma_queue);                                                                                  \
-        }                                                                                                              \
-    }                                                                                                                  \
-    dma_queue_flush(dma_queue);                                                                                        \
+#define GET_ROWS_THREAD_ST_FN(IDX_TYPE)                                                                                 \
+static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                             \
+    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                  \
+    struct htp_ops_context * octx = grctx->octx;                                                                        \
+    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                 \
+    get_rows_preamble;                                                                                                  \
+    const uint32_t dr  = grctx->tasks_per_thread;                                                                       \
+    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                  \
+    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                      \
+        return;                                                                                                         \
+    }                                                                                                                   \
+    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                               \
+    const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00);                                  \
+    dma_queue * dma_q = octx->ctx->dma[ith];                                                                            \
+    for (uint32_t i = ir0; i < ir1; ++i) {                                                                              \
+        const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);                                                       \
+        const uint32_t rem = i - i12 * ne11 * ne10;                                                                     \
+        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                          \
+        const uint32_t i10 = rem - i11 * ne10;                                                                          \
+        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12); \
+        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                       \
+        assert(i01 < ne01);                                                                                             \
+        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                          \
+        const uint32_t i02 = i11 - q02 * ne02;                                                                          \
+        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                          \
+        const uint32_t i03 = i12 - q03 * ne03;                                                                          \
+        const dma_addr_t src0_data = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03;                               \
+        const dma_addr_t dst_data  = octx->dst->data    + i10*nb1  + i11*nb2  + i12*nb3;                                \
+        while (!dma_queue_push(dma_q, dma_make_data(dst_data, src0_data), nb1, nb01,                                    \
+                               row_size_bytes, 1)) {                                                                    \
+            dma_queue_pop(dma_q);                                                                                       \
+        }                                                                                                               \
+    }                                                                                                                   \
+    dma_queue_flush(dma_q);                                                                                             \
 }

 GET_ROWS_THREAD_ST_FN(int32_t)
 GET_ROWS_THREAD_ST_FN(int64_t)

-#define GET_ROWS_THREAD_DT_FN(TYPE_NAME, SRC0_SIZE_EXPR, IDX_TYPE, COMPUTE_EXPR)                                       \
-static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                 \
-    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                 \
-    struct htp_ops_context * octx = grctx->octx;                                                                       \
-    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                \
-    get_rows_preamble;                                                                                                 \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                             \
-    const uint32_t dr  = grctx->tasks_per_thread;                                                                      \
-    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                 \
-    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                     \
-        return;                                                                                                        \
-    }                                                                                                                  \
-    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                              \
-    const uint32_t chunks_per_row = kparams->chunks_per_row;                                                           \
-    const uint32_t chunk_size     = kparams->chunk_size;                                                               \
-    dma_queue * dma_queue = octx->ctx->dma[ith];                                                                       \
-    const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;                                         \
-    uint8_t * vtcm_src0 = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;         \
-    uint8_t * vtcm_dst  = grctx->vtcm_base + vtcm_layout->off_dst  + ith * vtcm_layout->dst_bytes_per_thread;          \
-    for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, spad_idx++) {                      \
-        const uint32_t i = ir0 + step;                                                                                 \
-        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                           \
-        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                       \
-        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                \
-        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                              \
-        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                         \
-        const uint32_t i10 = rem - i11 * ne10;                                                                         \
-        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12);           \
-        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                      \
-        assert(i01 < ne01);                                                                                            \
-        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                         \
-        const uint32_t i02 = i11 - q02 * ne02;                                                                         \
-        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                         \
-        const uint32_t i03 = i12 - q03 * ne03;                                                                         \
-        const uint32_t offset = chunk_idx * chunk_size;                                                                \
-        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                               \
-        const uint32_t cur_src0_bytes = SRC0_SIZE_EXPR(cur_elems);                                                     \
-        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                     \
-        const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03 + SRC0_SIZE_EXPR(offset);       \
-        dma_queue_push(dma_queue,                                                                                      \
-                       dma_make_ptr((void *)(uintptr_t)octx->dst->data,                                                \
-                                    vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size),                            \
-                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 0);                              \
-        dma_queue_push(dma_queue,                                                                                      \
-                       dma_make_ptr((void *)(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size),                 \
-                                    (const void *)src0_ptr),                                                           \
-                       vtcm_layout->src0_spad_half_size, cur_src0_bytes, cur_src0_bytes, 1);                           \
-    }                                                                                                                  \
-    for (uint32_t step = 0; step < ir1 - ir0; ++step) {                                                                \
-        const uint32_t i = ir0 + step;                                                                                 \
-        void * dst_spad = (void *) dma_queue_pop(dma_queue).src;                                                       \
-        void * src_spad = (void *) dma_queue_pop(dma_queue).dst;                                                       \
-        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                           \
-        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                       \
-        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                \
-        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                              \
-        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                         \
-        const uint32_t i10 = rem - i11 * ne10;                                                                         \
-        const uint32_t offset = chunk_idx * chunk_size;                                                                \
-        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                               \
-        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                     \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                          \
-        COMPUTE_EXPR;                                                                                                  \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                           \
-        const uintptr_t dst_ptr  = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3 + offset * sizeof(float);             \
-        dma_queue_push(dma_queue,                                                                                      \
-                       dma_make_ptr((void *)dst_ptr, (const void *)dst_spad),                                          \
-                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 1);                              \
-        const uint32_t next_step = step + 2;                                                                           \
-        if (next_step < ir1 - ir0) {                                                                                   \
-            const uint32_t pi = ir0 + next_step;                                                                       \
-            const uint32_t prow_idx   = fastdiv(pi, &kparams->div_chunks_per_row);                                     \
-            const uint32_t pchunk_idx = pi - prow_idx * chunks_per_row;                                                \
-            const uint32_t pi12 = fastdiv(prow_idx, &kparams->div_ne10_ne11);                                          \
-            const uint32_t prem = prow_idx - pi12 * ne11 * ne10;                                                       \
-            const uint32_t pi11 = fastdiv(prem, &kparams->div_ne10);                                                   \
-            const uint32_t pi10 = prem - pi11 * ne10;                                                                  \
-            const IDX_TYPE * psrc1_ptr = (const IDX_TYPE *)(octx->src[1]->data + pi10*nb10 + pi11*nb11 + pi12*nb12);   \
-            const uint32_t pi01 = (uint32_t)*psrc1_ptr;                                                                \
-            assert(pi01 < ne01);                                                                                       \
-            const uint32_t pq02 = fastdiv(pi11, &kparams->div_ne02);                                                   \
-            const uint32_t pi02 = pi11 - pq02 * ne02;                                                                  \
-            const uint32_t pq03 = fastdiv(pi12, &kparams->div_ne03);                                                   \
-            const uint32_t pi03 = pi12 - pq03 * ne03;                                                                  \
-            const uint32_t poffset = pchunk_idx * chunk_size;                                                          \
-            const uint32_t pcur_elems = (poffset < ne00) ? MIN(chunk_size, ne00 - poffset) : 0;                        \
-            const uint32_t pcur_src0_bytes = SRC0_SIZE_EXPR(pcur_elems);                                               \
-            const uintptr_t psrc0_ptr =                                                                                \
-                octx->src[0]->data + pi01*nb01 + pi02*nb02 + pi03*nb03 + SRC0_SIZE_EXPR(poffset);                      \
-            dma_queue_push(dma_queue,                                                                                  \
-                           dma_make_ptr((void *)src_spad, (const void *)psrc0_ptr),                                    \
-                           vtcm_layout->src0_spad_half_size, pcur_src0_bytes, pcur_src0_bytes, 1);                     \
-        }                                                                                                              \
-    }                                                                                                                  \
-    dma_queue_flush(dma_queue);                                                                                        \
+#define GET_ROWS_THREAD_DT_FN(TYPE_NAME, SRC0_SIZE_EXPR, IDX_TYPE, COMPUTE_EXPR)                                                \
+static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) {                          \
+    struct get_rows_context * grctx = (struct get_rows_context *)data;                                                          \
+    struct htp_ops_context * octx = grctx->octx;                                                                                \
+    const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                         \
+    get_rows_preamble;                                                                                                          \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                      \
+    const uint32_t dr  = grctx->tasks_per_thread;                                                                               \
+    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                          \
+    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                              \
+        return;                                                                                                                 \
+    }                                                                                                                           \
+    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                                       \
+    const uint32_t chunks_per_row = kparams->chunks_per_row;                                                                    \
+    const uint32_t chunk_size     = kparams->chunk_size;                                                                        \
+    dma_queue * dma_q = octx->ctx->dma[ith];                                                                                    \
+    const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;                                                  \
+    uint8_t * vtcm_src0 = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;                  \
+    uint8_t * vtcm_dst  = grctx->vtcm_base + vtcm_layout->off_dst  + ith * vtcm_layout->dst_bytes_per_thread;                   \
+    for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, spad_idx++) {                               \
+        const uint32_t i = ir0 + step;                                                                                          \
+        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                                    \
+        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                                \
+        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                         \
+        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                                       \
+        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                                  \
+        const uint32_t i10 = rem - i11 * ne10;                                                                                  \
+        const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12);         \
+        const uint32_t i01 = (uint32_t)*src1_ptr;                                                                               \
+        assert(i01 < ne01);                                                                                                     \
+        const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);                                                                  \
+        const uint32_t i02 = i11 - q02 * ne02;                                                                                  \
+        const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);                                                                  \
+        const uint32_t i03 = i12 - q03 * ne03;                                                                                  \
+        const uint32_t offset = chunk_idx * chunk_size;                                                                         \
+        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                                        \
+        const uint32_t cur_src0_bytes = SRC0_SIZE_EXPR(cur_elems);                                                              \
+        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                              \
+        const dma_addr_t src0_data = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03 + SRC0_SIZE_EXPR(offset);              \
+        dma_queue_push(dma_q,                                                                                                   \
+                       dma_make_data(octx->dst->data,                                                                           \
+                                     vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size),                                    \
+                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 0);                                       \
+        dma_queue_push(dma_q,                                                                                                   \
+                       dma_make_data(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size, src0_data),                       \
+                       vtcm_layout->src0_spad_half_size, cur_src0_bytes, cur_src0_bytes, 1);                                    \
+    }                                                                                                                           \
+    for (uint32_t step = 0; step < ir1 - ir0; ++step) {                                                                         \
+        const uint32_t i = ir0 + step;                                                                                          \
+        void * dst_spad = (void *) dma_queue_pop(dma_q).src;                                                                    \
+        void * src_spad = (void *) dma_queue_pop(dma_q).dst;                                                                    \
+        const uint32_t row_idx   = fastdiv(i, &kparams->div_chunks_per_row);                                                    \
+        const uint32_t chunk_idx = i - row_idx * chunks_per_row;                                                                \
+        const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11);                                                         \
+        const uint32_t rem = row_idx - i12 * ne11 * ne10;                                                                       \
+        const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);                                                                  \
+        const uint32_t i10 = rem - i11 * ne10;                                                                                  \
+        const uint32_t offset = chunk_idx * chunk_size;                                                                         \
+        const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0;                                        \
+        const uint32_t cur_dst_bytes  = cur_elems * sizeof(float);                                                              \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                                   \
+        COMPUTE_EXPR;                                                                                                           \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, i);                                                                    \
+        const dma_addr_t dst_data  = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3 + offset * sizeof(float);                    \
+        dma_queue_push(dma_q,                                                                                                   \
+                       dma_make_data(dst_data, dst_spad),                                                                       \
+                       cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 1);                                       \
+        const uint32_t next_step = step + 2;                                                                                    \
+        if (next_step < ir1 - ir0) {                                                                                            \
+            const uint32_t pi = ir0 + next_step;                                                                                \
+            const uint32_t prow_idx   = fastdiv(pi, &kparams->div_chunks_per_row);                                              \
+            const uint32_t pchunk_idx = pi - prow_idx * chunks_per_row;                                                         \
+            const uint32_t pi12 = fastdiv(prow_idx, &kparams->div_ne10_ne11);                                                   \
+            const uint32_t prem = prow_idx - pi12 * ne11 * ne10;                                                                \
+            const uint32_t pi11 = fastdiv(prem, &kparams->div_ne10);                                                            \
+            const uint32_t pi10 = prem - pi11 * ne10;                                                                           \
+            const IDX_TYPE * psrc1_ptr = (const IDX_TYPE *)(uintptr_t)(octx->src[1]->data + pi10*nb10 + pi11*nb11 + pi12*nb12); \
+            const uint32_t pi01 = (uint32_t)*psrc1_ptr;                                                                         \
+            assert(pi01 < ne01);                                                                                                \
+            const uint32_t pq02 = fastdiv(pi11, &kparams->div_ne02);                                                            \
+            const uint32_t pi02 = pi11 - pq02 * ne02;                                                                           \
+            const uint32_t pq03 = fastdiv(pi12, &kparams->div_ne03);                                                            \
+            const uint32_t pi03 = pi12 - pq03 * ne03;                                                                           \
+            const uint32_t poffset = pchunk_idx * chunk_size;                                                                   \
+            const uint32_t pcur_elems = (poffset < ne00) ? MIN(chunk_size, ne00 - poffset) : 0;                                 \
+            const uint32_t pcur_src0_bytes = SRC0_SIZE_EXPR(pcur_elems);                                                        \
+            const dma_addr_t psrc0_data =                                                                                       \
+                octx->src[0]->data + pi01*nb01 + pi02*nb02 + pi03*nb03 + SRC0_SIZE_EXPR(poffset);                               \
+            dma_queue_push(dma_q,                                                                                               \
+                           dma_make_data(src_spad, psrc0_data),                                                                 \
+                           vtcm_layout->src0_spad_half_size, pcur_src0_bytes, pcur_src0_bytes, 1);                              \
+        }                                                                                                                       \
+    }                                                                                                                           \
+    dma_queue_flush(dma_q);                                                                                                     \
 }

 #define F32_BYTES(n)  ((n) * sizeof(float))
@@ -227,8 +227,8 @@ int op_get_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(octx->src[1])) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     const struct htp_tensor * dst = octx->dst;
diff --git a/ggml/src/ggml-hexagon/htp/hex-dma.h b/ggml/src/ggml-hexagon/htp/hex-dma.h
deleted file mode 100644
index 9e9a5f950..000000000
--- a/ggml/src/ggml-hexagon/htp/hex-dma.h
+++ /dev/null
@@ -1,2 +0,0 @@
-#pragma once
-#include "dma-queue.h"
diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h
index a3d5e8cef..814feef7c 100644
--- a/ggml/src/ggml-hexagon/htp/htp-ctx.h
+++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h
@@ -1,7 +1,7 @@
 #ifndef HTP_CTX_H
 #define HTP_CTX_H

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hmx-queue.h"
 #include "htp-ops.h"
 #include "hex-profile.h"
@@ -17,16 +17,16 @@
 #ifndef HTP_MAX_NTHREADS
 #define HTP_MAX_NTHREADS 10
 #endif
-#define HTP_MAX_MMAPS    16

-#define HTP_MAX_DIRTY_RANGES 32
+#define HTP_MAX_MMAPS        64
+#define HTP_MAX_DIRTY_RANGES 64

 // Memory mapping
 struct htp_mmap {
     uint64_t size;
     uint64_t base;
     uint32_t fd;
-    uint32_t reserved;
+    uint32_t flags;
 };

 struct htp_dirty_range {
@@ -68,9 +68,6 @@ struct htp_ops_context {
         const struct htp_tensor * dsts[HTP_OP_MAX_OUTPUTS];
     };

-    dma_queue **    src_dma[HTP_OP_MAX_INPUTS];
-    dma_queue **    dst_dma[HTP_OP_MAX_OUTPUTS];
-
     // TODO convert these to an array
     struct htp_spad src0_spad;
     struct htp_spad src1_spad;
@@ -90,7 +87,6 @@ struct htp_context {

     struct htp_mmap        mmap[HTP_MAX_MMAPS];
     dma_queue_t            dma[HTP_MAX_NTHREADS];
-    dma_queue_t            dma_cached[HTP_MAX_NTHREADS];
     struct htp_thread_trace trace[HTP_MAX_NTHREADS + 1];
     work_queue_t           work_queue;
     hmx_queue_t            hmx_queue;
diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h
index faf3118c4..0e63febdd 100644
--- a/ggml/src/ggml-hexagon/htp/htp-ops.h
+++ b/ggml/src/ggml-hexagon/htp/htp-ops.h
@@ -133,10 +133,13 @@ enum htp_tensor_flags {
     HTP_TENSOR_FENCE   = (1U << 2)  // Tensor is synchronization fence (explicitly managed)
 };

+enum htp_buf_flags {
+    HTP_BUF_EXTENDED = (1U << 0),
+};
+
 // Tensor descriptor
 struct htp_tensor {
-    uint32_t data;                 // Buffer offset in the messages, and data pointer on the NPU
-    uint32_t reserved;             // Reserved for alignment padding (must be multiple of 8)
+    uint64_t data;                 // Buffer offset in the messages, and data pointer on the NPU
     uint32_t size;                 // Data size in bytes
     uint32_t flags;                // Buffer / tensor flags
     uint32_t type;                 // Data type
@@ -150,12 +153,12 @@ struct htp_tensor {
 struct htp_buf_desc {
     uint64_t base;     // base address
     uint64_t size;     // total size
-    uint32_t flags;    // buffer flags (unused)
+    uint32_t flags;    // HTP_BUF_*
     uint32_t fd;       // file descriptor
 };

 enum htp_op_flags {
-    HTP_OPFLAGS_SKIP_COMPUTE  = (1U << 0), // Skip actual computation (used for profiling)
+    HTP_OPFLAGS_STUB  = (1U << 0),
 };

 // Op descriptor
diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.c b/ggml/src/ggml-hexagon/htp/htp-tensor.c
index 760ccd831..03b0070d7 100644
--- a/ggml/src/ggml-hexagon/htp/htp-tensor.c
+++ b/ggml/src/ggml-hexagon/htp/htp-tensor.c
@@ -226,6 +226,10 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
 }

 static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor * t) {
+    if (!t || (t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE))) {
+        return;
+    }
+
     uint32_t t_start = t->data;
     uint32_t t_end   = t_start + t->size;

@@ -236,6 +240,7 @@ static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor
         if (r->start < t_end && t_start < r->end) {
             if (t_start <= r->start && r->end <= t_end) {
                 r->start = 0;
+                r->end   = 0;
             } else if (t_start <= r->start) {
                 r->start = t_end;
             } else if (r->end <= t_end) {
@@ -246,6 +251,10 @@ static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor
 }

 static inline bool is_tensor_dirty(struct htp_context * ctx, const struct htp_tensor * t) {
+    if (!t || (t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE))) {
+        return false;
+    }
+
     uint32_t t_start = t->data;
     uint32_t t_end   = t_start + t->size;

@@ -327,7 +336,7 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co

     for (uint32_t i = 0; i < n; i++) {
         const struct htp_tensor * t = tensors[i];
-        if (t && is_tensor_dirty(ctx, t)) {
+        if (is_tensor_dirty(ctx, t)) {
             dirty_tensors[n_dirty++] = t;
             ranges[n_dirty - 1].start = t->data;
             ranges[n_dirty - 1].end   = t->data + t->size;
diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.h b/ggml/src/ggml-hexagon/htp/htp-tensor.h
index 1e32bf09f..f7a966836 100644
--- a/ggml/src/ggml-hexagon/htp/htp-tensor.h
+++ b/ggml/src/ggml-hexagon/htp/htp-tensor.h
@@ -21,6 +21,10 @@ static inline void * htp_tensor_data(const struct htp_tensor * t) {
     return (void *) (uintptr_t) t->data;
 }

+static inline bool htp_tensor_is_extended(const struct htp_tensor * t) {
+    return t && (t->data >> 32) != 0;
+}
+
 static inline uint32_t * htp_tensor_flags(const struct htp_tensor * t) {
     return (uint32_t *) &t->flags;
 }
diff --git a/ggml/src/ggml-hexagon/htp/htp_iface.idl b/ggml/src/ggml-hexagon/htp/htp_iface.idl
index 47693d8b8..b46e25296 100644
--- a/ggml/src/ggml-hexagon/htp/htp_iface.idl
+++ b/ggml/src/ggml-hexagon/htp/htp_iface.idl
@@ -13,7 +13,7 @@ struct htp_iface_pmu_conf {
 interface htp_iface : remote_handle64 {
     AEEResult start(in uint32 sess_id, in uint64 dsp_queue_id, in uint32 n_hvx, in uint32 n_hmx, in uint64 max_vmem);
     AEEResult stop();
-    AEEResult mmap(in uint32 fd, in uint32 size);
+    AEEResult mmap(in uint32 fd, in uint64 size);
     AEEResult munmap(in uint32 fd);
     AEEResult profiler(in uint32 mode, in htp_iface_pmu_conf pmu);
     AEEResult etm(in uint32 enable);
diff --git a/ggml/src/ggml-hexagon/htp/hvx-exp.h b/ggml/src/ggml-hexagon/htp/hvx-exp.h
index bcd3d2d32..93ca8cf51 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-exp.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-exp.h
@@ -173,7 +173,7 @@ static inline void hvx_exp_f32(uint8_t * restrict dst, const uint8_t * restrict
         HVX_Vector * p_vec_in1 = (HVX_Vector *) src;
         HVX_Vector * p_vec_out = (HVX_Vector *) dst;

-        #pragma unroll(4)
+        #pragma unroll(2)
         for (int i = 0; i < num_elems_whole; i += VLEN_FP32) {
             if (true == negate) {
                 HVX_Vector neg_vec_in = hvx_vec_neg_f32(*p_vec_in1++);
@@ -183,7 +183,7 @@ static inline void hvx_exp_f32(uint8_t * restrict dst, const uint8_t * restrict
             }
         }
     } else {
-        #pragma unroll(4)
+        #pragma unroll(2)
         for (int i = 0; i < num_elems_whole; i += VLEN_FP32) {
             HVX_Vector in = *(HVX_UVector *) (src + i * SIZEOF_FP32);

diff --git a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h
deleted file mode 100644
index 5c1372cf1..000000000
--- a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h
+++ /dev/null
@@ -1,1648 +0,0 @@
-// Dynamic quantizers that produce flat (non-tiled) activations
-
-static inline void quantize_block_f32_q8_0_flat(
-    float * restrict x,
-    uint8_t * restrict y_quants,
-    __fp16 * restrict y_scales,
-    uint32_t block_idx
-) {
-    HVX_Vector * vx = (HVX_Vector *) x;
-    HVX_Vector zero = Q6_V_vzero();
-
-    HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
-    HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
-    HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
-    HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
-
-    HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
-    HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
-    HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
-    HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
-
-    HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
-    HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
-    HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
-    HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
-
-    HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
-    HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
-
-    HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
-    HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
-
-    HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008));  // 1.0 / 127.0
-    HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008));  // 1.0 / 127.0
-    HVX_Vector vd01_hf   = Q6_Vhf_equals_Vqf16(vd01_qf16);
-    HVX_Vector vd23_hf   = Q6_Vhf_equals_Vqf16(vd23_qf16);
-
-    HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
-    HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
-    vx01_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
-    vx23_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));
-
-    HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
-    HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
-    HVX_Vector vx_i8    = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
-
-    * (HVX_Vector *) (y_quants + block_idx * 128) = vx_i8;
-
-    HVX_VectorPair vp1 = Q6_W_vshuff_VVR(vd23_hf, vd01_hf, -2);
-    HVX_VectorPair vp2 = Q6_W_vshuff_VVR(Q6_V_hi_W(vp1), Q6_V_lo_W(vp1), -2);
-    HVX_Vector v_scales = Q6_V_lo_W(vp2);
-    hvx_vec_store_u(y_scales + block_idx * 4, 8, v_scales);
-}
-
-static inline void quantize_block_f32_q8_1_flat(
-    float * restrict x,
-    uint8_t * restrict y_quants,
-    __fp16 * restrict y_scales,
-    uint32_t block_idx
-) {
-    HVX_Vector * vx = (HVX_Vector *) x;
-    HVX_Vector zero = Q6_V_vzero();
-
-    HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
-    HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
-    HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
-    HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
-
-    HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
-    HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
-    HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
-    HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
-
-    HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
-    HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
-    HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
-    HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
-
-    HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
-    HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
-
-    HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
-    HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
-
-    HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008));  // 1.0 / 127.0
-    HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008));  // 1.0 / 127.0
-    HVX_Vector vd01_hf   = Q6_Vhf_equals_Vqf16(vd01_qf16);
-    HVX_Vector vd23_hf   = Q6_Vhf_equals_Vqf16(vd23_qf16);
-
-    HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
-    HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
-    vx01_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
-    vx23_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));
-
-    HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
-    HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
-    HVX_Vector vx_i8    = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
-
-    const HVX_Vector ones = Q6_Vb_vsplat_R(1);
-    HVX_Vector v_sums = Q6_Vw_vrmpy_VbVb(vx_i8, ones);
-    v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 4));
-    v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 8));
-    v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 16));
-
-    * (HVX_Vector *) (y_quants + block_idx * 128) = vx_i8;
-
-    HVX_VectorPair vp1 = Q6_W_vshuff_VVR(vd23_hf, vd01_hf, -2);
-    HVX_VectorPair vp2 = Q6_W_vshuff_VVR(Q6_V_hi_W(vp1), Q6_V_lo_W(vp1), -2);
-    HVX_Vector v_scales = Q6_V_lo_W(vp2);
-
-    HVX_VectorPair v_deal1 = Q6_W_vdeal_VVR(v_sums, v_sums, -4);
-    HVX_Vector v_even1 = Q6_V_lo_W(v_deal1);
-    HVX_VectorPair v_deal2 = Q6_W_vdeal_VVR(v_even1, v_even1, -4);
-    HVX_Vector v_even2 = Q6_V_lo_W(v_deal2);
-    HVX_VectorPair v_deal3 = Q6_W_vdeal_VVR(v_even2, v_even2, -4);
-    HVX_Vector v_sums_shuffled = Q6_V_lo_W(v_deal3);
-
-    HVX_Vector v_sums_sf = Q6_Vsf_equals_Vw(v_sums_shuffled);
-    HVX_Vector v_sums_hf = hvx_vec_f32_to_f16(v_sums_sf, Q6_V_vzero());
-
-    HVX_Vector v_prod = hvx_vec_mul_f16_f16(v_scales, v_sums_hf);
-
-    HVX_VectorPair vp_scales = Q6_W_vshuff_VVR(v_prod, v_scales, -2);
-    HVX_Vector v_final = Q6_V_lo_W(vp_scales);
-
-    hvx_vec_store_u(y_scales + block_idx * 8, 16, v_final);
-}
-
-static inline void quantize_row_f32_q8_0_flat(float * restrict x, uint8_t * restrict y, uint32_t k) {
-    assert(k % 32 == 0);
-    const uint32_t quants_size = hex_round_up(k, 128);
-    uint8_t * restrict y_quants = y;
-    __fp16 * restrict y_scales = (__fp16 *) (y + quants_size);
-
-    const uint32_t nb = (k + 127) / 128;
-    for (uint32_t i = 0; i < nb; i++) {
-        quantize_block_f32_q8_0_flat(x + i * 128, y_quants, y_scales, i);
-    }
-}
-
-static inline void quantize_row_f32_q8_1_flat(float * restrict x, uint8_t * restrict y, uint32_t k) {
-    assert(k % 32 == 0);
-    const uint32_t quants_size = hex_round_up(k, 128);
-    uint8_t * restrict y_quants = y;
-    __fp16 * restrict y_scales = (__fp16 *) (y + quants_size);
-
-    const uint32_t nb = (k + 127) / 128;
-    for (uint32_t i = 0; i < nb; i++) {
-        quantize_block_f32_q8_1_flat(x + i * 128, y_quants, y_scales, i);
-    }
-}
-
-static inline void quantize_f32_q8_0_flat_kernel(
-    const uint8_t * restrict src_data,
-    uint8_t * restrict dst_data,
-    uint8_t * restrict tmp_data,
-    uint32_t ne0,
-    uint32_t nrows,
-    size_t src_row_size,
-    size_t dst_row_size
-) {
-    const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
-    hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
-
-    for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_row_size, 2);
-        hvx_copy_f32_aa(tmp_data, src_data, ne0);
-
-        quantize_row_f32_q8_0_flat((float *) tmp_data, dst_data, ne0);
-        dst_data += dst_row_size;
-        src_data += src_row_size;
-    }
-}
-
-static inline void quantize_f32_q8_1_flat_kernel(
-    const uint8_t * restrict src_data,
-    uint8_t * restrict dst_data,
-    uint8_t * restrict tmp_data,
-    uint32_t ne0,
-    uint32_t nrows,
-    size_t src_row_size,
-    size_t dst_row_size
-) {
-    const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
-    hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
-
-    for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_row_size, 2);
-        hvx_copy_f32_aa(tmp_data, src_data, ne0);
-
-        quantize_row_f32_q8_1_flat((float *) tmp_data, dst_data, ne0);
-        dst_data += dst_row_size;
-        src_data += src_row_size;
-    }
-}
-
-static inline void quantize_f32_f32_flat_kernel(
-    const uint8_t * restrict src_data,
-    uint8_t * restrict dst_data,
-    uint8_t * restrict tmp_data,
-    uint32_t ne0,
-    uint32_t nrows,
-    size_t src_stride,
-    size_t dst_stride
-) {
-    (void) tmp_data;
-    const size_t src_row_size = ne0 * sizeof(float);
-    for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_stride, 2);
-        hvx_copy_f32_au(dst_data, src_data, ne0);
-
-        dst_data += dst_stride;
-        src_data += src_stride;
-    }
-}
-
-static inline void quantize_f32_f16_flat_kernel(
-    const uint8_t * restrict src_data,
-    uint8_t * restrict dst_data,
-    uint8_t * restrict tmp_data,
-    uint32_t ne0,
-    uint32_t nrows,
-    size_t src_stride,
-    size_t dst_stride
-) {
-    (void) tmp_data;
-    const size_t src_row_size = ne0 * sizeof(float);
-    for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_stride, 2);
-        hvx_copy_f16_f32_au(dst_data, src_data, ne0);
-
-        dst_data += dst_stride;
-        src_data += src_stride;
-    }
-}
-
-static inline void quantize_f16_f16_flat_kernel(
-    const uint8_t * restrict src_data,
-    uint8_t * restrict dst_data,
-    uint8_t * restrict tmp_data,
-    uint32_t ne0,
-    uint32_t nrows,
-    size_t src_stride,
-    size_t dst_stride
-) {
-    (void) tmp_data;
-    const size_t src_row_size = ne0 * sizeof(float);
-    for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_stride, 2);
-        hvx_copy_f16_au(dst_data, src_data, ne0);
-
-        dst_data += dst_stride;
-        src_data += src_stride;
-    }
-}
-
-// Dot kernels that consume flat (non-tiled) activations
-
-static void flat_vec_dot_q4_0_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-    HVX_Vector i8 = Q6_Vb_vsplat_R(8);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx_i8 = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx_i8, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_sum = accum_4bit_32x1(vptr, v_act_rep, i8);
-        HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
-
-        HVX_Vector v_scale_w = vptr[4];
-
-        __fp16 scale_a_val = y_scales[kt];
-        HVX_Vector v_scale_a = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-
-        HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
-        HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
-    }
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_q4_0_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-    HVX_Vector i8 = Q6_Vb_vsplat_R(8);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0_i8 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1_i8 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0_i8, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1_i8, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_act1_rep[8];
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_4bit_32x2(vptr, v_act0_rep, v_act1_rep, i8);
-        HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
-        HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
-
-        HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
-        HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
-
-        HVX_Vector v_scale_w = vptr[4];
-
-        __fp16 scale_a0_val = y0_scales[kt];
-        __fp16 scale_a1_val = y1_scales[kt];
-        HVX_Vector v_scale_a0 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_scale_a1 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-
-        HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a0);
-        HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a1);
-
-        HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
-        HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
-    }
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-static void flat_vec_dot_q4_1_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx_i8 = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx_i8, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_sum = accum_4bit_32x1(vptr, v_act_rep, Q6_V_vzero());
-        HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
-
-        HVX_Vector v_scale_offset = vptr[4];
-        HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
-        HVX_Vector v_scale = Q6_V_lo_W(p_deal);
-        HVX_Vector v_offset = Q6_V_hi_W(p_deal);
-
-        __fp16 scale_a_val = y_scales[kt * 2 + 0];
-        __fp16 sum_a_val   = y_scales[kt * 2 + 1];
-        HVX_Vector v_scale_a = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-        HVX_Vector v_sum_a   = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&sum_a_val));
-
-        HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a);
-        HVX_Vector v_offset_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a);
-
-        HVX_Vector v_scaled_dot = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
-        HVX_Vector v_sum_scaled = hvx_vec_add_f32_f32(v_scaled_dot, v_offset_comb);
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
-    }
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_q4_1_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0_i8 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1_i8 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0_i8, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1_i8, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_act1_rep[8];
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_4bit_32x2(vptr, v_act0_rep, v_act1_rep, Q6_V_vzero());
-        HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
-        HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
-
-        HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
-        HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
-
-        HVX_Vector v_scale_offset = vptr[4];
-        HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
-        HVX_Vector v_scale = Q6_V_lo_W(p_deal);
-        HVX_Vector v_offset = Q6_V_hi_W(p_deal);
-
-        __fp16 scale_a0_val = y0_scales[kt * 2 + 0];
-        __fp16 sum_a0_val   = y0_scales[kt * 2 + 1];
-        __fp16 scale_a1_val = y1_scales[kt * 2 + 0];
-        __fp16 sum_a1_val   = y1_scales[kt * 2 + 1];
-
-        HVX_Vector v_scale_a0 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_sum_a0   = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&sum_a0_val));
-        HVX_Vector v_scale_a1 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-        HVX_Vector v_sum_a1   = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&sum_a1_val));
-
-        HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a0);
-        HVX_Vector v_offset_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a0);
-        HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a1);
-        HVX_Vector v_offset_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a1);
-
-        HVX_Vector v_scaled_dot_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
-        HVX_Vector v_sum_scaled_c0 = hvx_vec_add_f32_f32(v_scaled_dot_c0, v_offset_comb_c0);
-
-        HVX_Vector v_scaled_dot_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
-        HVX_Vector v_sum_scaled_c1 = hvx_vec_add_f32_f32(v_scaled_dot_c1, v_offset_comb_c1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
-    }
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-static void flat_vec_dot_q8_0_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 1152);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx_i8 = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx_i8, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_sum = accum_q8_0_32x1(vptr, v_act_rep);
-        HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
-
-        HVX_Vector v_scale_w = vptr[8];
-
-        __fp16 scale_a_val = y_scales[kt];
-        HVX_Vector v_scale_a = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-
-        HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
-        HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
-    }
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_q8_0_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 1152);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0_i8 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1_i8 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0_i8, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1_i8, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_act1_rep[8];
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_q8_0_32x2(vptr, v_act0_rep, v_act1_rep);
-        HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
-        HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
-
-        HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
-        HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
-
-        HVX_Vector v_scale_w = vptr[8];
-
-        __fp16 scale_a0_val = y0_scales[kt];
-        __fp16 scale_a1_val = y1_scales[kt];
-        HVX_Vector v_scale_a0 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_scale_a1 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-
-        HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a0);
-        HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a1);
-
-        HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
-        HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
-    }
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-static void flat_vec_dot_q6_k_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-    HVX_Vector i32 = Q6_Vb_vsplat_R(32);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 896);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx_i8 = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx_i8, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_q6_k_32x1(vptr, v_act_rep, i32);
-
-        __fp16 scale_a_val = y_scales[kt];
-        HVX_Vector v_scale_a = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, scale_q6_k_32x1(v_sums, vptr[6], v_scale_a));
-    }
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_q6_k_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-    HVX_Vector i32 = Q6_Vb_vsplat_R(32);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 896);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0_i8 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1_i8 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0_i8, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1_i8, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        HVX_Vector v_act1_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums0, v_sums1;
-        accum_q6_k_32x2(vptr, v_act0_rep, v_act1_rep, i32, &v_sums0, &v_sums1);
-
-        __fp16 scale_a0_val = y0_scales[kt];
-        __fp16 scale_a1_val = y1_scales[kt];
-        HVX_Vector v_scale_a0 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_scale_a1 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, scale_q6_k_32x1(v_sums0, vptr[6], v_scale_a0));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, scale_q6_k_32x1(v_sums1, vptr[6], v_scale_a1));
-    }
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-static void flat_vec_dot_iq4nl_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-    HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
-    HVX_Vector lut = *(const HVX_Vector *) kvalues_iq4nl_lut;
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_sum = accum_4bit_32x1_lut(vptr, v_act_rep, mask_h4, lut);
-        HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
-
-        HVX_Vector v_scale_w = vptr[4];
-
-        __fp16 scale_a_val = y_scales[kt];
-        HVX_Vector v_scale_a = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-
-        HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
-        HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
-    }
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_iq4nl_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-    HVX_Vector mask_h4        = Q6_Vb_vsplat_R(0x0F);
-    HVX_Vector lut            = *(const HVX_Vector *) kvalues_iq4nl_lut;
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_act1_rep[8];
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_4bit_32x2_lut(vptr, v_act0_rep, v_act1_rep, mask_h4, lut);
-        HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
-        HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
-
-        HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
-        HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
-
-        HVX_Vector v_scale_w = vptr[4];
-
-        __fp16 scale_a0_val = y0_scales[kt];
-        __fp16 scale_a1_val = y1_scales[kt];
-        HVX_Vector v_scale_a0 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_scale_a1 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-
-        HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a0);
-        HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a1);
-
-        HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
-        HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
-    }
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-static void flat_vec_dot_mxfp4_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y_q = vy;
-
-    HVX_Vector v_sum_float = Q6_V_vzero();
-    HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
-    HVX_Vector lut = *(const HVX_Vector *) kvalues_mxfp4_lut;
-    HVX_Vector expand = *(const HVX_Vector *) expand_x32_e8m0;
-    HVX_Vector e8m0_mask = Q6_V_vsplat_R(0x000000ff);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y_scales = (const __fp16 *) (y_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx = * (const HVX_Vector *) (y_q + block_idx * 128);
-        HVX_Vector v_act_raw = Q6_V_vror_VR(vx, sub_idx * 32);
-
-        HVX_Vector v_act_rep[8];
-        v_act_rep[0] = Q6_V_vdelta_VV(v_act_raw, v_repl_ctrl);
-        v_act_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 4), v_repl_ctrl);
-        v_act_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 8), v_repl_ctrl);
-        v_act_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 12), v_repl_ctrl);
-        v_act_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 16), v_repl_ctrl);
-        v_act_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 20), v_repl_ctrl);
-        v_act_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 24), v_repl_ctrl);
-        v_act_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_sum = accum_4bit_32x1_lut(vptr, v_act_rep, mask_h4, lut);
-        HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
-
-        HVX_Vector v_scale_w = hvx_vmem(tile_ptr + kt * 640 + 512);
-        HVX_Vector r0_d = Q6_V_vdelta_VV(v_scale_w, expand);
-        r0_d = Q6_V_vand_VV(r0_d, e8m0_mask);
-        HVX_Vector v_scale_w_f32 = Q6_Vw_vasl_VwR(r0_d, 23);
-
-        __fp16 scale_a_val = y_scales[kt];
-        HVX_Vector v_scale_a_f16 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a_val));
-        HVX_VectorPair p_scale_a_f32 = hvx_vec_f16_to_f32(v_scale_a_f16);
-        HVX_Vector v_scale_a = Q6_V_lo_W(p_scale_a_f32);
-
-        HVX_Vector v_scale_comb = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a);
-        HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
-
-        v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
-    }
-
-    v_sum_float = hvx_vec_mul_f32_f32(v_sum_float, hvx_vec_splat_f32(0.5f));
-
-    if (sz) {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
-    } else {
-        hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
-    }
-}
-
-static void flat_vec_dot_mxfp4_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
-    const uint8_t * restrict tile_ptr = vx;
-    const uint8_t * restrict y0_q = vy0;
-    const uint8_t * restrict y1_q = vy1;
-
-    HVX_Vector v_sum_float_c0 = Q6_V_vzero();
-    HVX_Vector v_sum_float_c1 = Q6_V_vzero();
-    HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
-    HVX_Vector lut = *(const HVX_Vector *) kvalues_mxfp4_lut;
-    HVX_Vector expand = *(const HVX_Vector *) expand_x32_e8m0;
-    HVX_Vector e8m0_mask = Q6_V_vsplat_R(0x000000ff);
-
-    static const uint8_t __attribute__((aligned(128))) repl[128] = {
-        0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-        0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
-    };
-    HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
-
-    const uint32_t quants_size = hex_round_up(n, 128);
-    const __fp16 * restrict y0_scales = (const __fp16 *) (y0_q + quants_size);
-    const __fp16 * restrict y1_scales = (const __fp16 *) (y1_q + quants_size);
-
-    uint32_t n_k_tiles = n / 32;
-    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
-        const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
-
-        uint32_t block_idx = kt / 4;
-        uint32_t sub_idx = kt % 4;
-
-        HVX_Vector vx0 = * (const HVX_Vector *) (y0_q + block_idx * 128);
-        HVX_Vector vx1 = * (const HVX_Vector *) (y1_q + block_idx * 128);
-
-        HVX_Vector v_act0_raw = Q6_V_vror_VR(vx0, sub_idx * 32);
-        HVX_Vector v_act1_raw = Q6_V_vror_VR(vx1, sub_idx * 32);
-
-        HVX_Vector v_act0_rep[8];
-        v_act0_rep[0] = Q6_V_vdelta_VV(v_act0_raw, v_repl_ctrl);
-        v_act0_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 4), v_repl_ctrl);
-        v_act0_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 8), v_repl_ctrl);
-        v_act0_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 12), v_repl_ctrl);
-        v_act0_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 16), v_repl_ctrl);
-        v_act0_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 20), v_repl_ctrl);
-        v_act0_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 24), v_repl_ctrl);
-        v_act0_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act0_raw, 28), v_repl_ctrl);
-
-        HVX_Vector v_act1_rep[8];
-        v_act1_rep[0] = Q6_V_vdelta_VV(v_act1_raw, v_repl_ctrl);
-        v_act1_rep[1] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 4), v_repl_ctrl);
-        v_act1_rep[2] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 8), v_repl_ctrl);
-        v_act1_rep[3] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 12), v_repl_ctrl);
-        v_act1_rep[4] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 16), v_repl_ctrl);
-        v_act1_rep[5] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 20), v_repl_ctrl);
-        v_act1_rep[6] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 24), v_repl_ctrl);
-        v_act1_rep[7] = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act1_raw, 28), v_repl_ctrl);
-
-        HVX_VectorPair v_sums = accum_4bit_32x2_lut(vptr, v_act0_rep, v_act1_rep, mask_h4, lut);
-        HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
-        HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
-
-        HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
-        HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
-
-        HVX_Vector v_scale_w = hvx_vmem(tile_ptr + kt * 640 + 512);
-        HVX_Vector r0_d = Q6_V_vdelta_VV(v_scale_w, expand);
-        r0_d = Q6_V_vand_VV(r0_d, e8m0_mask);
-        HVX_Vector v_scale_w_f32 = Q6_Vw_vasl_VwR(r0_d, 23);
-
-        __fp16 scale_a0_val = y0_scales[kt];
-        __fp16 scale_a1_val = y1_scales[kt];
-        HVX_Vector v_scale_a0_f16 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a0_val));
-        HVX_Vector v_scale_a1_f16 = hvx_vec_repl_f16(Q6_Vh_vsplat_R(*(const int16_t *)&scale_a1_val));
-        HVX_VectorPair p_scale_a0_f32 = hvx_vec_f16_to_f32(v_scale_a0_f16);
-        HVX_VectorPair p_scale_a1_f32 = hvx_vec_f16_to_f32(v_scale_a1_f16);
-        HVX_Vector v_scale_a0 = Q6_V_lo_W(p_scale_a0_f32);
-        HVX_Vector v_scale_a1 = Q6_V_lo_W(p_scale_a1_f32);
-
-        HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a0);
-        HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a1);
-
-        HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
-        HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
-    }
-
-    v_sum_float_c0 = hvx_vec_mul_f32_f32(v_sum_float_c0, hvx_vec_splat_f32(0.5f));
-    v_sum_float_c1 = hvx_vec_mul_f32_f32(v_sum_float_c1, hvx_vec_splat_f32(0.5f));
-
-    if (sz0) {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
-    } else {
-        hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
-    }
-    if (sz1) {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
-    } else {
-        hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
-    }
-}
-
-#if __HVX_ARCH__ < 79
-#define HVX_OP_ADD_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(a, b))
-#define HVX_OP_MUL_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(a, b))
-#else
-#define HVX_OP_ADD_F32(a, b) Q6_Vsf_vadd_VsfVsf(a, b)
-#define HVX_OP_MUL_F32(a, b) Q6_Vsf_vmpy_VsfVsf(a, b)
-#endif
-
-static inline void vec_dot_f32_f32_aa_1x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy) {
-    const HVX_Vector * restrict x = (const HVX_Vector *) vx;
-    const HVX_Vector * restrict y = (const HVX_Vector *) vy;
-
-    uint32_t nvec = n / VLEN_FP32; // num full fp32 hvx vectors
-    uint32_t nloe = n % VLEN_FP32; // leftover elements
-
-    HVX_Vector rsum = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(4)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector prod = HVX_OP_MUL_F32(x[i], y[i]);
-        rsum = HVX_OP_ADD_F32(rsum, prod);
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
-        HVX_Vector x_sf = Q6_V_vand_QV(bmask, x[i]);
-        HVX_Vector y_sf = Q6_V_vand_QV(bmask, y[i]);
-        HVX_Vector prod = HVX_OP_MUL_F32(x_sf, y_sf);
-        rsum = HVX_OP_ADD_F32(rsum, prod);
-    }
-
-    *s = hvx_vec_get_f32(hvx_vec_reduce_sum_f32(rsum));
-}
-
-static inline void vec_dot_f32_f32_aa_2x1(const uint32_t n, float * restrict s0,
-                                const void * restrict vx0, const void * restrict vx1,
-                                const void * restrict vy0) {
-    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
-    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
-    const HVX_Vector * restrict y  = (const HVX_Vector *) vy0;
-
-    uint32_t nvec = n / VLEN_FP32;
-    uint32_t nloe = n % VLEN_FP32;
-
-    HVX_Vector rsum0 = Q6_V_vzero();
-    HVX_Vector rsum1 = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector y_sf = y[i];
-        HVX_Vector prod0 = HVX_OP_MUL_F32(x0[i], y_sf);
-        HVX_Vector prod1 = HVX_OP_MUL_F32(x1[i], y_sf);
-        rsum0 = HVX_OP_ADD_F32(rsum0, prod0);
-        rsum1 = HVX_OP_ADD_F32(rsum1, prod1);
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
-        HVX_Vector y_sf  = Q6_V_vand_QV(bmask, y[i]);
-        HVX_Vector x0_sf = Q6_V_vand_QV(bmask, x0[i]);
-        HVX_Vector x1_sf = Q6_V_vand_QV(bmask, x1[i]);
-        HVX_Vector prod0 = HVX_OP_MUL_F32(x0_sf, y_sf);
-        HVX_Vector prod1 = HVX_OP_MUL_F32(x1_sf, y_sf);
-        rsum0 = HVX_OP_ADD_F32(rsum0, prod0);
-        rsum1 = HVX_OP_ADD_F32(rsum1, prod1);
-    }
-
-    HVX_Vector rsum = hvx_vec_reduce_sum_f32x2(rsum0, rsum1);
-    hvx_vec_store_u(s0, 8, rsum);
-}
-
-static inline void vec_dot_f32_f32_aa_2x2(const uint32_t n, float * restrict s0, float * restrict s1,
-                                const void * restrict vx0, const void * restrict vx1,
-                                const void * restrict vy0, const void * restrict vy1) {
-    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
-    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
-    const HVX_Vector * restrict y0 = (const HVX_Vector *) vy0;
-    const HVX_Vector * restrict y1 = (const HVX_Vector *) vy1;
-
-    uint32_t nvec = n / VLEN_FP32;
-    uint32_t nloe = n % VLEN_FP32;
-
-    HVX_Vector r0_c0_sum = Q6_V_vzero();
-    HVX_Vector r0_c1_sum = Q6_V_vzero();
-    HVX_Vector r1_c0_sum = Q6_V_vzero();
-    HVX_Vector r1_c1_sum = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector r0_sf = x0[i];
-        HVX_Vector r1_sf = x1[i];
-        HVX_Vector c0_sf = y0[i];
-        HVX_Vector c1_sf = y1[i];
-
-        r0_c0_sum = HVX_OP_ADD_F32(r0_c0_sum, HVX_OP_MUL_F32(r0_sf, c0_sf));
-        r0_c1_sum = HVX_OP_ADD_F32(r0_c1_sum, HVX_OP_MUL_F32(r0_sf, c1_sf));
-        r1_c0_sum = HVX_OP_ADD_F32(r1_c0_sum, HVX_OP_MUL_F32(r1_sf, c0_sf));
-        r1_c1_sum = HVX_OP_ADD_F32(r1_c1_sum, HVX_OP_MUL_F32(r1_sf, c1_sf));
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
-
-        HVX_Vector r0_sf = Q6_V_vand_QV(bmask, x0[i]);
-        HVX_Vector r1_sf = Q6_V_vand_QV(bmask, x1[i]);
-        HVX_Vector c0_sf = Q6_V_vand_QV(bmask, y0[i]);
-        HVX_Vector c1_sf = Q6_V_vand_QV(bmask, y1[i]);
-
-        r0_c0_sum = HVX_OP_ADD_F32(r0_c0_sum, HVX_OP_MUL_F32(r0_sf, c0_sf));
-        r0_c1_sum = HVX_OP_ADD_F32(r0_c1_sum, HVX_OP_MUL_F32(r0_sf, c1_sf));
-        r1_c0_sum = HVX_OP_ADD_F32(r1_c0_sum, HVX_OP_MUL_F32(r1_sf, c0_sf));
-        r1_c1_sum = HVX_OP_ADD_F32(r1_c1_sum, HVX_OP_MUL_F32(r1_sf, c1_sf));
-    }
-
-    // Reduce and store results
-    HVX_Vector r0_r1_c0_sum = hvx_vec_reduce_sum_f32x2(r0_c0_sum, r1_c0_sum);
-    HVX_Vector r0_r1_c1_sum = hvx_vec_reduce_sum_f32x2(r0_c1_sum, r1_c1_sum);
-
-    hvx_vec_store_u(s0, 8, r0_r1_c0_sum);
-    hvx_vec_store_u(s1, 8, r0_r1_c1_sum);
-}
-
-static inline void vec_dot_f32_f32_uu_1x1(const uint32_t n, float * restrict s, const void * restrict x, const void * restrict y) {
-    const HVX_UVector * restrict vx = (const HVX_UVector * restrict) x;
-    const HVX_UVector * restrict vy = (const HVX_UVector * restrict) y;
-
-    uint32_t nvec = n / VLEN_FP32; // num full fp32 hvx vectors
-    uint32_t nloe = n % VLEN_FP32; // leftover elements
-
-    HVX_Vector       rsum = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector x_sf = vx[i];
-        HVX_Vector y_sf = vy[i];
-
-        rsum = HVX_OP_ADD_F32(rsum, HVX_OP_MUL_F32(x_sf, y_sf));
-    }
-
-    if (nloe) {
-        HVX_Vector x_sf = vx[i];
-        HVX_Vector y_sf = vy[i];
-
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
-        x_sf = Q6_V_vand_QV(bmask, x_sf);
-        y_sf = Q6_V_vand_QV(bmask, y_sf);
-
-        rsum = HVX_OP_ADD_F32(rsum, HVX_OP_MUL_F32(x_sf, y_sf));
-    }
-
-    rsum = hvx_vec_reduce_sum_f32(rsum);
-    hvx_vec_store_u(&s[0], 4, rsum);
-}
-
-#undef HVX_OP_ADD_F32
-#undef HVX_OP_MUL_F32
-
-static inline void vec_dot_f16_f16_aa_1x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy) {
-    const HVX_Vector * restrict x = (const HVX_Vector *) vx;
-    const HVX_Vector * restrict y = (const HVX_Vector *) vy;
-
-    uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
-    uint32_t nloe = n % VLEN_FP16; // leftover elements
-
-    HVX_VectorPair rsum_p = Q6_W_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(4)
-    for (i = 0; i < nvec; i++) {
-        rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, x[i], y[i]);
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
-        HVX_Vector x_hf = Q6_V_vand_QV(bmask, x[i]);
-        HVX_Vector y_hf = Q6_V_vand_QV(bmask, y[i]);
-        rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, x_hf, y_hf);
-    }
-
-    HVX_Vector rsum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum_p), Q6_V_hi_W(rsum_p)));
-    hvx_vec_store_u(s, 4, hvx_vec_reduce_sum_f32(rsum));
-}
-
-static inline void vec_dot_f16_f16_aa_2x1(const uint32_t n, float * restrict s0,
-                                const void * restrict vx0, const void * restrict vx1,
-                                const void * restrict vy0) {
-    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
-    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
-    const HVX_Vector * restrict y  = (const HVX_Vector *) vy0;
-
-    uint32_t nvec = n / VLEN_FP16;
-    uint32_t nloe = n % VLEN_FP16;
-
-    HVX_VectorPair rsum0_p = Q6_W_vzero();
-    HVX_VectorPair rsum1_p = Q6_W_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector y_hf = y[i];
-        rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0[i], y_hf);
-        rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1[i], y_hf);
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
-        HVX_Vector y_hf  = Q6_V_vand_QV(bmask, y[i]);
-        HVX_Vector x0_hf = Q6_V_vand_QV(bmask, x0[i]);
-        HVX_Vector x1_hf = Q6_V_vand_QV(bmask, x1[i]);
-        rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0_hf, y_hf);
-        rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1_hf, y_hf);
-    }
-
-    HVX_Vector rsum0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum0_p), Q6_V_hi_W(rsum0_p)));
-    HVX_Vector rsum1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum1_p), Q6_V_hi_W(rsum1_p)));
-    HVX_Vector rsum  = hvx_vec_reduce_sum_f32x2(rsum0, rsum1);
-    hvx_vec_store_u(s0, 8, rsum);
-}
-
-static inline void vec_dot_f16_f16_aa_2x2(const uint32_t n, float * restrict s0, float * restrict s1,
-                                const void * restrict vx0, const void * restrict vx1,
-                                const void * restrict vy0, const void * restrict vy1) {
-    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
-    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
-    const HVX_Vector * restrict y0 = (const HVX_Vector *) vy0;
-    const HVX_Vector * restrict y1 = (const HVX_Vector *) vy1;
-
-    uint32_t nvec = n / VLEN_FP16;
-    uint32_t nloe = n % VLEN_FP16;
-
-    // Row sums (sf) - 4 accumulators for 2x2 tile
-    HVX_VectorPair r0_c0_sum_p = Q6_W_vzero();
-    HVX_VectorPair r0_c1_sum_p = Q6_W_vzero();
-    HVX_VectorPair r1_c0_sum_p = Q6_W_vzero();
-    HVX_VectorPair r1_c1_sum_p = Q6_W_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        HVX_Vector r0_hf = x0[i];
-        HVX_Vector r1_hf = x1[i];
-        HVX_Vector c0_hf = y0[i];
-        HVX_Vector c1_hf = y1[i];
-
-        // Compute 4 dot products: r0xc0, r0xc1, r1xc0, r1xc1
-        r0_c0_sum_p = hvx_vec_mpyacc_f32_f16(r0_c0_sum_p, r0_hf, c0_hf);
-        r0_c1_sum_p = hvx_vec_mpyacc_f32_f16(r0_c1_sum_p, r0_hf, c1_hf);
-        r1_c0_sum_p = hvx_vec_mpyacc_f32_f16(r1_c0_sum_p, r1_hf, c0_hf);
-        r1_c1_sum_p = hvx_vec_mpyacc_f32_f16(r1_c1_sum_p, r1_hf, c1_hf);
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
-
-        HVX_Vector r0_hf = Q6_V_vand_QV(bmask, x0[i]);
-        HVX_Vector r1_hf = Q6_V_vand_QV(bmask, x1[i]);
-        HVX_Vector c0_hf = Q6_V_vand_QV(bmask, y0[i]);
-        HVX_Vector c1_hf = Q6_V_vand_QV(bmask, y1[i]);
-
-        r0_c0_sum_p = hvx_vec_mpyacc_f32_f16(r0_c0_sum_p, r0_hf, c0_hf);
-        r0_c1_sum_p = hvx_vec_mpyacc_f32_f16(r0_c1_sum_p, r0_hf, c1_hf);
-        r1_c0_sum_p = hvx_vec_mpyacc_f32_f16(r1_c0_sum_p, r1_hf, c0_hf);
-        r1_c1_sum_p = hvx_vec_mpyacc_f32_f16(r1_c1_sum_p, r1_hf, c1_hf);
-    }
-
-    HVX_Vector r0_c0_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r0_c0_sum_p), Q6_V_hi_W(r0_c0_sum_p)));
-    HVX_Vector r0_c1_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r0_c1_sum_p), Q6_V_hi_W(r0_c1_sum_p)));
-    HVX_Vector r1_c0_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r1_c0_sum_p), Q6_V_hi_W(r1_c0_sum_p)));
-    HVX_Vector r1_c1_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r1_c1_sum_p), Q6_V_hi_W(r1_c1_sum_p)));
-
-    // Reduce and store results
-    HVX_Vector r0_r1_c0_sum = hvx_vec_reduce_sum_f32x2(r0_c0_sum, r1_c0_sum);
-    HVX_Vector r0_r1_c1_sum = hvx_vec_reduce_sum_f32x2(r0_c1_sum, r1_c1_sum);
-
-    hvx_vec_store_u(&s0[0], 8, r0_r1_c0_sum);  // row0,col0 row1,col0
-    hvx_vec_store_u(&s1[0], 8, r0_r1_c1_sum);  // row0,col1 row1,col1
-}
-
-static inline void vec_dot_f16_f16_uu_1x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy) {
-    const HVX_UVector * restrict x = (const HVX_UVector *) vx;
-    const HVX_UVector * restrict y = (const HVX_UVector *) vy;
-
-    uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
-    uint32_t nloe = n % VLEN_FP16; // leftover elements
-
-    HVX_Vector rsum = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(4)
-    for (i = 0; i < nvec; i++) {
-        HVX_VectorPair xy_qf = Q6_Wqf32_vmpy_VhfVhf(x[i], y[i]);
-        rsum = Q6_Vqf32_vadd_Vqf32Vqf32(rsum, Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_lo_W(xy_qf),  Q6_V_hi_W(xy_qf)));
-    }
-
-    if (nloe) {
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
-        HVX_Vector x_hf = Q6_V_vand_QV(bmask, x[i]);
-        HVX_Vector y_hf = Q6_V_vand_QV(bmask, y[i]);
-
-        HVX_VectorPair xy_qf = Q6_Wqf32_vmpy_VhfVhf(x_hf, y_hf);
-        rsum = Q6_Vqf32_vadd_Vqf32Vqf32(rsum, Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_lo_W(xy_qf),  Q6_V_hi_W(xy_qf)));
-    }
-
-    rsum = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(rsum));
-    hvx_vec_store_u(&s[0], 4, rsum);
-}
-
-static inline void vec_dot_f16_f32_uu_1x1(const uint32_t n, float * restrict s, const void * restrict x, const void * restrict y) {
-    const HVX_UVector * restrict vx = (const HVX_UVector * restrict) x;
-    const HVX_UVector * restrict vy = (const HVX_UVector * restrict) y;
-
-    uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
-    uint32_t nloe = n % VLEN_FP16; // leftover elements
-
-    const HVX_Vector zero = Q6_V_vzero();
-
-    HVX_Vector       rsum = Q6_V_vzero();
-
-    uint32_t i = 0;
-
-    #pragma unroll(2)
-    for (i = 0; i < nvec; i++) {
-        // Load y (fp32) and convert into fp16
-        HVX_Vector y0_qf = Q6_Vqf32_vsub_VsfVsf(vy[i*2+0], zero);  // 32 elements
-        HVX_Vector y1_qf = Q6_Vqf32_vsub_VsfVsf(vy[i*2+1], zero);  // 32 elements
-        HVX_Vector y_hf  = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(y1_qf, y0_qf)));
-
-        // Load x (fp16)
-        HVX_Vector x_hf  = vx[i];
-
-        HVX_VectorPair xy_qf = Q6_Wqf32_vmpy_VhfVhf(x_hf, y_hf);
-
-        rsum = Q6_Vqf32_vadd_Vqf32Vqf32(rsum, Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_lo_W(xy_qf),  Q6_V_hi_W(xy_qf)));
-    }
-
-    if (nloe) {
-        // Load y (fp32) and convert into fp16
-        HVX_Vector y0_qf = Q6_Vqf32_vsub_VsfVsf(vy[i*2+0], zero);  // 32 elements
-        HVX_Vector y1_qf = Q6_Vqf32_vsub_VsfVsf(vy[i*2+1], zero);  // 32 elements
-        HVX_Vector y_hf  = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(y1_qf, y0_qf)));
-
-        // Load x (fp16)
-        HVX_Vector x_hf  = vx[i];
-
-        // Zero-out unused elements
-        // Note that we need to clear both x and y because they may contain NANs
-        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
-        x_hf = Q6_V_vand_QV(bmask, x_hf);
-        y_hf = Q6_V_vand_QV(bmask, y_hf);
-
-        HVX_VectorPair xy_qf = Q6_Wqf32_vmpy_VhfVhf(x_hf, y_hf);
-
-        rsum = Q6_Vqf32_vadd_Vqf32Vqf32(rsum, Q6_Vqf32_vadd_Vqf32Vqf32(Q6_V_lo_W(xy_qf),  Q6_V_hi_W(xy_qf)));
-    }
-
-    // Convert into fp32 and reduce
-    rsum = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(rsum));
-    hvx_vec_store_u(&s[0], 4, rsum);
-}
-
-static inline void hvx_tensor_add_f32_grid(
-    const struct htp_tensor * restrict dst,
-    const struct htp_tensor * restrict src2,
-    uint32_t start_row,
-    uint32_t end_row,
-    uint32_t start_col,
-    uint32_t end_col,
-    const struct fastdiv_values * div_ne11_12,
-    const struct fastdiv_values * div_ne11
-) {
-    if (start_row >= end_row || start_col >= end_col) return;
-    const uint32_t nb1 = dst->nb[1]; // row stride in bytes
-
-    const uint32_t ne11 = dst->ne[1];
-    const uint32_t ne12 = dst->ne[2];
-    const uint32_t ne11_12 = ne11 * ne12;
-
-    const bool is_broadcast1 = (src2->ne[1] == 1);
-    const bool is_broadcast2 = (src2->ne[2] == 1);
-    const bool is_broadcast3 = (src2->ne[3] == 1);
-
-    for (uint32_t r = start_row; r < end_row; r++) {
-        float * dst_row = (float *) ((uint8_t *) dst->data + r * nb1);
-
-        uint32_t i13 = fastdiv(r, div_ne11_12);
-        uint32_t i12 = fastdiv(r - i13 * ne11_12, div_ne11);
-        uint32_t i11 = r - i13 * ne11_12 - i12 * ne11;
-
-        uint32_t i23 = is_broadcast3 ? 0 : i13;
-        uint32_t i22 = is_broadcast2 ? 0 : i12;
-        uint32_t i21 = is_broadcast1 ? 0 : i11;
-
-        const float * src2_row = (const float *) ((const uint8_t *) src2->data +
-                                  i21 * src2->nb[1] + i22 * src2->nb[2] + i23 * src2->nb[3]);
-
-        float * dst_ptr = &dst_row[start_col];
-        const float * src2_ptr = &src2_row[start_col];
-        int remaining = end_col - start_col;
-        while (remaining >= 32) {
-            HVX_Vector v_out = hvx_vmemu(dst_ptr);
-            HVX_Vector v_z   = hvx_vmemu(src2_ptr);
-            hvx_vmemu(dst_ptr) = hvx_vec_add_f32_f32(v_out, v_z);
-            dst_ptr += 32;
-            src2_ptr += 32;
-            remaining -= 32;
-        }
-        if (remaining > 0) {
-            HVX_Vector v_out = hvx_vmemu(dst_ptr);
-            HVX_Vector v_z   = hvx_vmemu(src2_ptr);
-            hvx_vec_store_u(dst_ptr, remaining * sizeof(float), hvx_vec_add_f32_f32(v_out, v_z));
-        }
-    }
-}
-
diff --git a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-float.h b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-float.h
new file mode 100644
index 000000000..605892aa7
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-float.h
@@ -0,0 +1,382 @@
+#ifndef HVX_MM_KERNELS_FLOAT_H
+#define HVX_MM_KERNELS_FLOAT_H
+
+#include "hvx-utils.h"
+#include "htp-tensor.h"
+
+// Float activation copy/quantization kernels (DDR -> VTCM)
+
+static inline void quantize_f32_f32_kernel(
+    const uint8_t * restrict src_data,
+    uint8_t * restrict dst_data,
+    uint8_t * restrict tmp_data,
+    uint32_t ne0,
+    uint32_t nrows,
+    size_t src_stride,
+    size_t dst_stride
+) {
+    (void) tmp_data;
+    const size_t src_row_size = ne0 * sizeof(float);
+    for (uint32_t i = 0; i < nrows; ++i) {
+        hex_l2fetch(src_data, src_row_size, src_stride, 2);
+        hvx_copy_f32_au(dst_data, src_data, ne0);
+
+        dst_data += dst_stride;
+        src_data += src_stride;
+    }
+}
+
+static inline void quantize_f32_f16_kernel(
+    const uint8_t * restrict src_data,
+    uint8_t * restrict dst_data,
+    uint8_t * restrict tmp_data,
+    uint32_t ne0,
+    uint32_t nrows,
+    size_t src_stride,
+    size_t dst_stride
+) {
+    (void) tmp_data;
+    const size_t src_row_size = ne0 * sizeof(float);
+    for (uint32_t i = 0; i < nrows; ++i) {
+        hex_l2fetch(src_data, src_row_size, src_stride, 2);
+        hvx_copy_f16_f32_au(dst_data, src_data, ne0);
+
+        dst_data += dst_stride;
+        src_data += src_stride;
+    }
+}
+
+static inline void quantize_f16_f16_kernel(
+    const uint8_t * restrict src_data,
+    uint8_t * restrict dst_data,
+    uint8_t * restrict tmp_data,
+    uint32_t ne0,
+    uint32_t nrows,
+    size_t src_stride,
+    size_t dst_stride
+) {
+    (void) tmp_data;
+    const size_t src_row_size = ne0 * sizeof(float);
+    for (uint32_t i = 0; i < nrows; ++i) {
+        hex_l2fetch(src_data, src_row_size, src_stride, 2);
+        hvx_copy_f16_au(dst_data, src_data, ne0);
+
+        dst_data += dst_stride;
+        src_data += src_stride;
+    }
+}
+
+// Float dot product kernels (HVX)
+
+#if __HVX_ARCH__ < 79
+#define HVX_OP_ADD_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(a, b))
+#define HVX_OP_MUL_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(a, b))
+#else
+#define HVX_OP_ADD_F32(a, b) Q6_Vsf_vadd_VsfVsf(a, b)
+#define HVX_OP_MUL_F32(a, b) Q6_Vsf_vmpy_VsfVsf(a, b)
+#endif
+
+static inline void vec_dot_f32_f32_aa_1x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy) {
+    const HVX_Vector * restrict x = (const HVX_Vector *) vx;
+    const HVX_Vector * restrict y = (const HVX_Vector *) vy;
+
+    uint32_t nvec = n / VLEN_FP32; // num full fp32 hvx vectors
+    uint32_t nloe = n % VLEN_FP32; // leftover elements
+
+    HVX_Vector rsum = Q6_V_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(4)
+    for (i = 0; i < nvec; i++) {
+        HVX_Vector prod = HVX_OP_MUL_F32(x[i], y[i]);
+        rsum = HVX_OP_ADD_F32(rsum, prod);
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
+        HVX_Vector x_sf = Q6_V_vand_QV(bmask, x[i]);
+        HVX_Vector y_sf = Q6_V_vand_QV(bmask, y[i]);
+        HVX_Vector prod = HVX_OP_MUL_F32(x_sf, y_sf);
+        rsum = HVX_OP_ADD_F32(rsum, prod);
+    }
+
+    *s = hvx_vec_get_f32(hvx_vec_reduce_sum_f32(rsum));
+}
+
+static inline void vec_dot_f32_f32_aa_2x1(const uint32_t n, float * restrict s0,
+                                const void * restrict vx0, const void * restrict vx1,
+                                const void * restrict vy0) {
+    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
+    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
+    const HVX_Vector * restrict y  = (const HVX_Vector *) vy0;
+
+    uint32_t nvec = n / VLEN_FP32;
+    uint32_t nloe = n % VLEN_FP32;
+
+    HVX_Vector rsum0 = Q6_V_vzero();
+    HVX_Vector rsum1 = Q6_V_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(2)
+    for (i = 0; i < nvec; i++) {
+        HVX_Vector y_sf = y[i];
+        HVX_Vector prod0 = HVX_OP_MUL_F32(x0[i], y_sf);
+        HVX_Vector prod1 = HVX_OP_MUL_F32(x1[i], y_sf);
+        rsum0 = HVX_OP_ADD_F32(rsum0, prod0);
+        rsum1 = HVX_OP_ADD_F32(rsum1, prod1);
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
+        HVX_Vector y_sf  = Q6_V_vand_QV(bmask, y[i]);
+        HVX_Vector x0_sf = Q6_V_vand_QV(bmask, x0[i]);
+        HVX_Vector x1_sf = Q6_V_vand_QV(bmask, x1[i]);
+        HVX_Vector prod0 = HVX_OP_MUL_F32(x0_sf, y_sf);
+        HVX_Vector prod1 = HVX_OP_MUL_F32(x1_sf, y_sf);
+        rsum0 = HVX_OP_ADD_F32(rsum0, prod0);
+        rsum1 = HVX_OP_ADD_F32(rsum1, prod1);
+    }
+
+    HVX_Vector rsum = hvx_vec_reduce_sum_f32x2(rsum0, rsum1);
+    hvx_vec_store_u(s0, 8, rsum);
+}
+
+static inline void vec_dot_f32_f32_aa_2x2(const uint32_t n, float * restrict s0, float * restrict s1,
+                                const void * restrict vx0, const void * restrict vx1,
+                                const void * restrict vy0, const void * restrict vy1) {
+    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
+    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
+    const HVX_Vector * restrict y0 = (const HVX_Vector *) vy0;
+    const HVX_Vector * restrict y1 = (const HVX_Vector *) vy1;
+
+    uint32_t nvec = n / VLEN_FP32;
+    uint32_t nloe = n % VLEN_FP32;
+
+    HVX_Vector r0_c0_sum = Q6_V_vzero();
+    HVX_Vector r0_c1_sum = Q6_V_vzero();
+    HVX_Vector r1_c0_sum = Q6_V_vzero();
+    HVX_Vector r1_c1_sum = Q6_V_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(2)
+    for (i = 0; i < nvec; i++) {
+        HVX_Vector r0_sf = x0[i];
+        HVX_Vector r1_sf = x1[i];
+        HVX_Vector c0_sf = y0[i];
+        HVX_Vector c1_sf = y1[i];
+
+        r0_c0_sum = HVX_OP_ADD_F32(r0_c0_sum, HVX_OP_MUL_F32(r0_sf, c0_sf));
+        r0_c1_sum = HVX_OP_ADD_F32(r0_c1_sum, HVX_OP_MUL_F32(r0_sf, c1_sf));
+        r1_c0_sum = HVX_OP_ADD_F32(r1_c0_sum, HVX_OP_MUL_F32(r1_sf, c0_sf));
+        r1_c1_sum = HVX_OP_ADD_F32(r1_c1_sum, HVX_OP_MUL_F32(r1_sf, c1_sf));
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 4);
+
+        HVX_Vector r0_sf = Q6_V_vand_QV(bmask, x0[i]);
+        HVX_Vector r1_sf = Q6_V_vand_QV(bmask, x1[i]);
+        HVX_Vector c0_sf = Q6_V_vand_QV(bmask, y0[i]);
+        HVX_Vector c1_sf = Q6_V_vand_QV(bmask, y1[i]);
+
+        r0_c0_sum = HVX_OP_ADD_F32(r0_c0_sum, HVX_OP_MUL_F32(r0_sf, c0_sf));
+        r0_c1_sum = HVX_OP_ADD_F32(r0_c1_sum, HVX_OP_MUL_F32(r0_sf, c1_sf));
+        r1_c0_sum = HVX_OP_ADD_F32(r1_c0_sum, HVX_OP_MUL_F32(r1_sf, c0_sf));
+        r1_c1_sum = HVX_OP_ADD_F32(r1_c1_sum, HVX_OP_MUL_F32(r1_sf, c1_sf));
+    }
+
+    // Reduce and store results
+    HVX_Vector r0_r1_c0_sum = hvx_vec_reduce_sum_f32x2(r0_c0_sum, r1_c0_sum);
+    HVX_Vector r0_r1_c1_sum = hvx_vec_reduce_sum_f32x2(r0_c1_sum, r1_c1_sum);
+
+    hvx_vec_store_u(s0, 8, r0_r1_c0_sum);
+    hvx_vec_store_u(s1, 8, r0_r1_c1_sum);
+}
+
+#undef HVX_OP_ADD_F32
+#undef HVX_OP_MUL_F32
+
+static inline void vec_dot_f16_f16_aa_1x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy) {
+    const HVX_Vector * restrict x = (const HVX_Vector *) vx;
+    const HVX_Vector * restrict y = (const HVX_Vector *) vy;
+
+    uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
+    uint32_t nloe = n % VLEN_FP16; // leftover elements
+
+    HVX_VectorPair rsum_p = Q6_W_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(4)
+    for (i = 0; i < nvec; i++) {
+        rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, x[i], y[i]);
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
+        HVX_Vector x_hf = Q6_V_vand_QV(bmask, x[i]);
+        HVX_Vector y_hf = Q6_V_vand_QV(bmask, y[i]);
+        rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, x_hf, y_hf);
+    }
+
+    HVX_Vector rsum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum_p), Q6_V_hi_W(rsum_p)));
+    hvx_vec_store_u(s, 4, hvx_vec_reduce_sum_f32(rsum));
+}
+
+static inline void vec_dot_f16_f16_aa_2x1(const uint32_t n, float * restrict s0,
+                                const void * restrict vx0, const void * restrict vx1,
+                                const void * restrict vy0) {
+    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
+    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
+    const HVX_Vector * restrict y  = (const HVX_Vector *) vy0;
+
+    uint32_t nvec = n / VLEN_FP16;
+    uint32_t nloe = n % VLEN_FP16;
+
+    HVX_VectorPair rsum0_p = Q6_W_vzero();
+    HVX_VectorPair rsum1_p = Q6_W_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(2)
+    for (i = 0; i < nvec; i++) {
+        HVX_Vector y_hf = y[i];
+        rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0[i], y_hf);
+        rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1[i], y_hf);
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
+        HVX_Vector y_hf  = Q6_V_vand_QV(bmask, y[i]);
+        HVX_Vector x0_hf = Q6_V_vand_QV(bmask, x0[i]);
+        HVX_Vector x1_hf = Q6_V_vand_QV(bmask, x1[i]);
+        rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0_hf, y_hf);
+        rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1_hf, y_hf);
+    }
+
+    HVX_Vector rsum0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum0_p), Q6_V_hi_W(rsum0_p)));
+    HVX_Vector rsum1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(rsum1_p), Q6_V_hi_W(rsum1_p)));
+    HVX_Vector rsum  = hvx_vec_reduce_sum_f32x2(rsum0, rsum1);
+    hvx_vec_store_u(s0, 8, rsum);
+}
+
+static inline void vec_dot_f16_f16_aa_2x2(const uint32_t n, float * restrict s0, float * restrict s1,
+                                const void * restrict vx0, const void * restrict vx1,
+                                const void * restrict vy0, const void * restrict vy1) {
+    const HVX_Vector * restrict x0 = (const HVX_Vector *) vx0;
+    const HVX_Vector * restrict x1 = (const HVX_Vector *) vx1;
+    const HVX_Vector * restrict y0 = (const HVX_Vector *) vy0;
+    const HVX_Vector * restrict y1 = (const HVX_Vector *) vy1;
+
+    uint32_t nvec = n / VLEN_FP16;
+    uint32_t nloe = n % VLEN_FP16;
+
+    // Row sums (sf) - 4 accumulators for 2x2 tile
+    HVX_VectorPair r0_c0_sum_p = Q6_W_vzero();
+    HVX_VectorPair r0_c1_sum_p = Q6_W_vzero();
+    HVX_VectorPair r1_c0_sum_p = Q6_W_vzero();
+    HVX_VectorPair r1_c1_sum_p = Q6_W_vzero();
+
+    uint32_t i = 0;
+
+    #pragma unroll(2)
+    for (i = 0; i < nvec; i++) {
+        HVX_Vector r0_hf = x0[i];
+        HVX_Vector r1_hf = x1[i];
+        HVX_Vector c0_hf = y0[i];
+        HVX_Vector c1_hf = y1[i];
+
+        // Compute 4 dot products: r0xc0, r0xc1, r1xc0, r1xc1
+        r0_c0_sum_p = hvx_vec_mpyacc_f32_f16(r0_c0_sum_p, r0_hf, c0_hf);
+        r0_c1_sum_p = hvx_vec_mpyacc_f32_f16(r0_c1_sum_p, r0_hf, c1_hf);
+        r1_c0_sum_p = hvx_vec_mpyacc_f32_f16(r1_c0_sum_p, r1_hf, c0_hf);
+        r1_c1_sum_p = hvx_vec_mpyacc_f32_f16(r1_c1_sum_p, r1_hf, c1_hf);
+    }
+
+    if (nloe) {
+        HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
+
+        HVX_Vector r0_hf = Q6_V_vand_QV(bmask, x0[i]);
+        HVX_Vector r1_hf = Q6_V_vand_QV(bmask, x1[i]);
+        HVX_Vector c0_hf = Q6_V_vand_QV(bmask, y0[i]);
+        HVX_Vector c1_hf = Q6_V_vand_QV(bmask, y1[i]);
+
+        r0_c0_sum_p = hvx_vec_mpyacc_f32_f16(r0_c0_sum_p, r0_hf, c0_hf);
+        r0_c1_sum_p = hvx_vec_mpyacc_f32_f16(r0_c1_sum_p, r0_hf, c1_hf);
+        r1_c0_sum_p = hvx_vec_mpyacc_f32_f16(r1_c0_sum_p, r1_hf, c0_hf);
+        r1_c1_sum_p = hvx_vec_mpyacc_f32_f16(r1_c1_sum_p, r1_hf, c1_hf);
+    }
+
+    HVX_Vector r0_c0_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r0_c0_sum_p), Q6_V_hi_W(r0_c0_sum_p)));
+    HVX_Vector r0_c1_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r0_c1_sum_p), Q6_V_hi_W(r0_c1_sum_p)));
+    HVX_Vector r1_c0_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r1_c0_sum_p), Q6_V_hi_W(r1_c0_sum_p)));
+    HVX_Vector r1_c1_sum = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(r1_c1_sum_p), Q6_V_hi_W(r1_c1_sum_p)));
+
+    // Reduce and store results
+    HVX_Vector r0_r1_c0_sum = hvx_vec_reduce_sum_f32x2(r0_c0_sum, r1_c0_sum);
+    HVX_Vector r0_r1_c1_sum = hvx_vec_reduce_sum_f32x2(r0_c1_sum, r1_c1_sum);
+
+    hvx_vec_store_u(&s0[0], 8, r0_r1_c0_sum);  // row0,col0 row1,col0
+    hvx_vec_store_u(&s1[0], 8, r0_r1_c1_sum);  // row0,col1 row1,col1
+}
+
+
+
+static inline void hvx_tensor_add_f32_grid(
+    const struct htp_tensor * restrict dst,
+    const struct htp_tensor * restrict src2,
+    uint32_t start_row,
+    uint32_t end_row,
+    uint32_t start_col,
+    uint32_t end_col,
+    const struct fastdiv_values * div_ne11_12,
+    const struct fastdiv_values * div_ne11
+) {
+    if (start_row >= end_row || start_col >= end_col) return;
+    const uint32_t nb1 = dst->nb[1]; // row stride in bytes
+
+    const uint32_t ne11 = dst->ne[1];
+    const uint32_t ne12 = dst->ne[2];
+    const uint32_t ne11_12 = ne11 * ne12;
+
+    const bool is_broadcast1 = (src2->ne[1] == 1);
+    const bool is_broadcast2 = (src2->ne[2] == 1);
+    const bool is_broadcast3 = (src2->ne[3] == 1);
+
+    for (uint32_t r = start_row; r < end_row; r++) {
+        float * dst_row = (float *) ((uint8_t *) dst->data + (size_t) r * nb1);
+
+        uint32_t i13 = fastdiv(r, div_ne11_12);
+        uint32_t i12 = fastdiv(r - i13 * ne11_12, div_ne11);
+        uint32_t i11 = r - i13 * ne11_12 - i12 * ne11;
+
+        uint32_t i23 = is_broadcast3 ? 0 : i13;
+        uint32_t i22 = is_broadcast2 ? 0 : i12;
+        uint32_t i21 = is_broadcast1 ? 0 : i11;
+
+        const float * src2_row = (const float *) ((const uint8_t *) src2->data +
+                                  (size_t) i21 * src2->nb[1] + (size_t) i22 * src2->nb[2] + (size_t) i23 * src2->nb[3]);
+
+        float * dst_ptr = &dst_row[start_col];
+        const float * src2_ptr = &src2_row[start_col];
+        int remaining = end_col - start_col;
+        while (remaining >= 32) {
+            HVX_Vector v_out = hvx_vmemu(dst_ptr);
+            HVX_Vector v_z   = hvx_vmemu(src2_ptr);
+            hvx_vmemu(dst_ptr) = hvx_vec_add_f32_f32(v_out, v_z);
+            dst_ptr += 32;
+            src2_ptr += 32;
+            remaining -= 32;
+        }
+        if (remaining > 0) {
+            HVX_Vector v_out = hvx_vmemu(dst_ptr);
+            HVX_Vector v_z   = hvx_vmemu(src2_ptr);
+            hvx_vec_store_u(dst_ptr, remaining * sizeof(float), hvx_vec_add_f32_f32(v_out, v_z));
+        }
+    }
+}
+
+#endif // HVX_MM_KERNELS_FLOAT_H
diff --git a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
index c889538ac..4d6110ffa 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
@@ -48,22 +48,33 @@ static inline void quantize_block_f32_q8_1_tiled(float * restrict x, uint8_t * r
     v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 8));
     v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 16));

-    float vmax0[32]  __attribute__((aligned(128)));
-    float vmax1[32]  __attribute__((aligned(128)));
-    float vmax2[32]  __attribute__((aligned(128)));
-    float vmax3[32]  __attribute__((aligned(128)));
-    int32_t sums[32] __attribute__((aligned(128)));
-
-    hvx_vec_store_u(vmax0, 128, vmax0_sf);
-    hvx_vec_store_u(vmax1, 128, vmax1_sf);
-    hvx_vec_store_u(vmax2, 128, vmax2_sf);
-    hvx_vec_store_u(vmax3, 128, vmax3_sf);
-    hvx_vec_store_u(sums,  128, v_sums);
-
-    float d0 = vmax0[0] / 127.0f;
-    float d1 = vmax1[0] / 127.0f;
-    float d2 = vmax2[0] / 127.0f;
-    float d3 = vmax3[0] / 127.0f;
+    const HVX_Vector v_inv127 = hvx_vec_splat_f32(1.0f / 127.0f);
+    HVX_Vector vd0_sf = hvx_vec_mul_f32_f32(vmax0_sf, v_inv127);
+    HVX_Vector vd1_sf = hvx_vec_mul_f32_f32(vmax1_sf, v_inv127);
+    HVX_Vector vd2_sf = hvx_vec_mul_f32_f32(vmax2_sf, v_inv127);
+    HVX_Vector vd3_sf = hvx_vec_mul_f32_f32(vmax3_sf, v_inv127);
+
+    HVX_Vector v_sums_sf = Q6_Vsf_equals_Vw(v_sums);
+    HVX_Vector voff0_sf = hvx_vec_mul_f32_f32(vd0_sf, v_sums_sf);
+    HVX_Vector voff1_sf = hvx_vec_mul_f32_f32(vd1_sf, Q6_V_vror_VR(v_sums_sf, 32));
+    HVX_Vector voff2_sf = hvx_vec_mul_f32_f32(vd2_sf, Q6_V_vror_VR(v_sums_sf, 64));
+    HVX_Vector voff3_sf = hvx_vec_mul_f32_f32(vd3_sf, Q6_V_vror_VR(v_sums_sf, 96));
+
+    HVX_Vector voff01_hf = hvx_vec_f32_to_f16(voff0_sf, voff1_sf);
+    HVX_Vector voff23_hf = hvx_vec_f32_to_f16(voff2_sf, voff3_sf);
+
+    HVX_Vector r_scale[4] = {
+        hvx_vec_repl_f16(vd01_hf),
+        hvx_vec_repl_f16(Q6_V_vror_VR(vd01_hf, 64)),
+        hvx_vec_repl_f16(vd23_hf),
+        hvx_vec_repl_f16(Q6_V_vror_VR(vd23_hf, 64)),
+    };
+    HVX_Vector r_offset[4] = {
+        hvx_vec_repl_f16(voff01_hf),
+        hvx_vec_repl_f16(Q6_V_vror_VR(voff01_hf, 64)),
+        hvx_vec_repl_f16(voff23_hf),
+        hvx_vec_repl_f16(Q6_V_vror_VR(voff23_hf, 64)),
+    };

     static const uint8_t __attribute__((aligned(128))) repl[128] = {
         0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
@@ -89,24 +100,6 @@ static inline void quantize_block_f32_q8_1_tiled(float * restrict x, uint8_t * r
         HVX_Vector r6 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 24), v_repl_ctrl);
         HVX_Vector r7 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 28), v_repl_ctrl);

-        __fp16 scale_h, offset_h;
-        if (b == 0) {
-            scale_h  = (__fp16) d0;
-            offset_h = (__fp16) (sums[0] * d0);
-        } else if (b == 1) {
-            scale_h  = (__fp16) d1;
-            offset_h = (__fp16) (sums[8] * d1);
-        } else if (b == 2) {
-            scale_h  = (__fp16) d2;
-            offset_h = (__fp16) (sums[16] * d2);
-        } else {
-            scale_h  = (__fp16) d3;
-            offset_h = (__fp16) (sums[24] * d3);
-        }
-
-        HVX_Vector r_scale  = Q6_Vh_vsplat_R(*(int16_t *)&scale_h);
-        HVX_Vector r_offset = Q6_Vh_vsplat_R(*(int16_t *)&offset_h);
-
         HVX_Vector * restrict dst = (HVX_Vector *) (y_block + b * 1280);
         dst[0] = r0;
         dst[1] = r1;
@@ -116,8 +109,8 @@ static inline void quantize_block_f32_q8_1_tiled(float * restrict x, uint8_t * r
         dst[5] = r5;
         dst[6] = r6;
         dst[7] = r7;
-        dst[8] = r_scale;
-        dst[9] = r_offset;
+        dst[8] = r_scale[b];
+        dst[9] = r_offset[b];
     }
 }

@@ -486,51 +479,7 @@ static void tiled_vec_dot_q4_0_32x2(const uint32_t n, float * restrict s0, float
     HVX_Vector i8 = Q6_Vb_vsplat_R(8);

     uint32_t n_k_tiles = n / 32;
-    uint32_t kt = 0;
-    for (; kt + 1 < n_k_tiles; kt += 2) {
-        const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
-        const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
-        const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
-
-        const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
-        const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
-        const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
-
-        HVX_VectorPair v_sums0 = accum_4bit_32x2(vptr0, v_act0_0, v_act1_0, i8);
-        HVX_VectorPair v_sums1 = accum_4bit_32x2(vptr1, v_act0_1, v_act1_1, i8);
-
-        HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
-        HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
-        HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
-        HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
-
-        HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
-        HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
-        HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
-        HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
-
-        HVX_Vector v_scale_w0 = vptr0[4];
-        HVX_Vector v_scale_w1 = vptr1[4];
-        HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
-        HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
-        HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
-        HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
-
-        HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
-        HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
-        HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
-        HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
-
-        HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
-        HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
-        HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
-        HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
-    }
-
-    for (; kt < n_k_tiles; kt++) {
+    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
         const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
         const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
         const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
@@ -615,76 +564,7 @@ static void tiled_vec_dot_q4_1_32x2(const uint32_t n, float * restrict s0, float
     HVX_Vector v_sum_float_c1 = Q6_V_vzero();

     uint32_t n_k_tiles = n / 32;
-    uint32_t kt = 0;
-    for (; kt + 1 < n_k_tiles; kt += 2) {
-        const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
-        const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1280);
-        const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1280);
-
-        const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
-        const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1280);
-        const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1280);
-
-        HVX_VectorPair v_sums0 = accum_4bit_32x2(vptr0, v_act0_0, v_act1_0, Q6_V_vzero());
-        HVX_VectorPair v_sums1 = accum_4bit_32x2(vptr1, v_act0_1, v_act1_1, Q6_V_vzero());
-
-        HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
-        HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
-        HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
-        HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
-
-        HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
-        HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
-        HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
-        HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
-
-        HVX_Vector v_scale_offset0 = vptr0[4];
-        HVX_VectorPair p_deal0 = Q6_W_vdeal_VVR(v_scale_offset0, v_scale_offset0, -2);
-        HVX_Vector v_scale0 = Q6_V_lo_W(p_deal0);
-        HVX_Vector v_offset0 = Q6_V_hi_W(p_deal0);
-
-        HVX_Vector v_scale_offset1 = vptr1[4];
-        HVX_VectorPair p_deal1 = Q6_W_vdeal_VVR(v_scale_offset1, v_scale_offset1, -2);
-        HVX_Vector v_scale1 = Q6_V_lo_W(p_deal1);
-        HVX_Vector v_offset1 = Q6_V_hi_W(p_deal1);
-
-        HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
-        HVX_Vector v_sum_a_c0_0   = v_act0_0[9];
-        HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
-        HVX_Vector v_sum_a_c1_0   = v_act1_0[9];
-
-        HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
-        HVX_Vector v_sum_a_c0_1   = v_act0_1[9];
-        HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
-        HVX_Vector v_sum_a_c1_1   = v_act1_1[9];
-
-        HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale0, v_scale_a_c0_0);
-        HVX_Vector v_offset_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset0, v_sum_a_c0_0);
-        HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale0, v_scale_a_c1_0);
-        HVX_Vector v_offset_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset0, v_sum_a_c1_0);
-
-        HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale1, v_scale_a_c0_1);
-        HVX_Vector v_offset_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset1, v_sum_a_c0_1);
-        HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale1, v_scale_a_c1_1);
-        HVX_Vector v_offset_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset1, v_sum_a_c1_1);
-
-        HVX_Vector v_scaled_dot_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
-        HVX_Vector v_sum_scaled_c0_0 = hvx_vec_add_f32_f32(v_scaled_dot_c0_0, v_offset_comb_c0_0);
-
-        HVX_Vector v_scaled_dot_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
-        HVX_Vector v_sum_scaled_c1_0 = hvx_vec_add_f32_f32(v_scaled_dot_c1_0, v_offset_comb_c1_0);
-
-        HVX_Vector v_scaled_dot_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
-        HVX_Vector v_sum_scaled_c0_1 = hvx_vec_add_f32_f32(v_scaled_dot_c0_1, v_offset_comb_c0_1);
-
-        HVX_Vector v_scaled_dot_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
-        HVX_Vector v_sum_scaled_c1_1 = hvx_vec_add_f32_f32(v_scaled_dot_c1_1, v_offset_comb_c1_1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
-    }
-
-    for (; kt < n_k_tiles; kt++) {
+    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
         const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
         const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1280);
         const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1280);
@@ -771,51 +651,7 @@ static void tiled_vec_dot_q8_0_32x2(const uint32_t n, float * restrict s0, float
     HVX_Vector v_sum_float_c1 = Q6_V_vzero();

     uint32_t n_k_tiles = n / 32;
-    uint32_t kt = 0;
-    for (; kt + 1 < n_k_tiles; kt += 2) {
-        const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 1152);
-        const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
-        const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
-
-        const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 1152);
-        const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
-        const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
-
-        HVX_VectorPair v_sums0 = accum_q8_0_32x2(vptr0, v_act0_0, v_act1_0);
-        HVX_VectorPair v_sums1 = accum_q8_0_32x2(vptr1, v_act0_1, v_act1_1);
-
-        HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
-        HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
-        HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
-        HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
-
-        HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
-        HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
-        HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
-        HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
-
-        HVX_Vector v_scale_w0 = vptr0[8];
-        HVX_Vector v_scale_w1 = vptr1[8];
-        HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
-        HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
-        HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
-        HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
-
-        HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
-        HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
-        HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
-        HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
-
-        HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
-        HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
-        HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
-        HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
-    }
-
-    for (; kt < n_k_tiles; kt++) {
+    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
         const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 1152);
         const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
         const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
@@ -952,51 +788,7 @@ static void tiled_vec_dot_iq4nl_32x2(const uint32_t n, float * restrict s0, floa
     HVX_Vector lut = *(const HVX_Vector *) kvalues_iq4nl_lut;

     uint32_t n_k_tiles = n / 32;
-    uint32_t kt = 0;
-    for (; kt + 1 < n_k_tiles; kt += 2) {
-        const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
-        const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
-        const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
-
-        const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
-        const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
-        const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
-
-        HVX_VectorPair v_sums0 = accum_4bit_32x2_lut(vptr0, v_act0_0, v_act1_0, mask_h4, lut);
-        HVX_VectorPair v_sums1 = accum_4bit_32x2_lut(vptr1, v_act0_1, v_act1_1, mask_h4, lut);
-
-        HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
-        HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
-        HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
-        HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
-
-        HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
-        HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
-        HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
-        HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
-
-        HVX_Vector v_scale_w0 = vptr0[4];
-        HVX_Vector v_scale_w1 = vptr1[4];
-        HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
-        HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
-        HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
-        HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
-
-        HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
-        HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
-        HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
-        HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
-
-        HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
-        HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
-        HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
-        HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
-    }
-
-    for (; kt < n_k_tiles; kt++) {
+    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
         const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
         const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
         const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
@@ -1089,69 +881,7 @@ static void tiled_vec_dot_mxfp4_32x2(const uint32_t n, float * restrict s0, floa
     HVX_Vector e8m0_mask = Q6_V_vsplat_R(0x000000ff);

     uint32_t n_k_tiles = n / 32;
-    uint32_t kt = 0;
-    for (; kt + 1 < n_k_tiles; kt += 2) {
-        const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
-        const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
-        const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
-
-        const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
-        const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
-        const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
-
-        HVX_VectorPair v_sums0 = accum_4bit_32x2_lut(vptr0, v_act0_0, v_act1_0, mask_h4, lut);
-        HVX_VectorPair v_sums1 = accum_4bit_32x2_lut(vptr1, v_act0_1, v_act1_1, mask_h4, lut);
-
-        HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
-        HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
-        HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
-        HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
-
-        HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
-        HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
-        HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
-        HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
-
-        HVX_Vector v_scale_w0 = hvx_vmem(tile_ptr + (kt + 0) * 640 + 512);
-        HVX_Vector r0_d0 = Q6_V_vdelta_VV(v_scale_w0, expand);
-        r0_d0 = Q6_V_vand_VV(r0_d0, e8m0_mask);
-        HVX_Vector v_scale_w_f32_0 = Q6_Vw_vasl_VwR(r0_d0, 23);
-
-        HVX_Vector v_scale_w1 = hvx_vmem(tile_ptr + (kt + 1) * 640 + 512);
-        HVX_Vector r0_d1 = Q6_V_vdelta_VV(v_scale_w1, expand);
-        r0_d1 = Q6_V_vand_VV(r0_d1, e8m0_mask);
-        HVX_Vector v_scale_w_f32_1 = Q6_Vw_vasl_VwR(r0_d1, 23);
-
-        HVX_Vector v_scale_a_c0_f16_0 = v_act0_0[8];
-        HVX_Vector v_scale_a_c1_f16_0 = v_act1_0[8];
-        HVX_Vector v_scale_a_c0_f16_1 = v_act0_1[8];
-        HVX_Vector v_scale_a_c1_f16_1 = v_act1_1[8];
-
-        HVX_VectorPair p_scale_a_c0_f32_0 = hvx_vec_f16_to_f32_shuff(v_scale_a_c0_f16_0);
-        HVX_VectorPair p_scale_a_c1_f32_0 = hvx_vec_f16_to_f32_shuff(v_scale_a_c1_f16_0);
-        HVX_VectorPair p_scale_a_c0_f32_1 = hvx_vec_f16_to_f32_shuff(v_scale_a_c0_f16_1);
-        HVX_VectorPair p_scale_a_c1_f32_1 = hvx_vec_f16_to_f32_shuff(v_scale_a_c1_f16_1);
-
-        HVX_Vector v_scale_a_c0_0 = Q6_V_lo_W(p_scale_a_c0_f32_0);
-        HVX_Vector v_scale_a_c1_0 = Q6_V_lo_W(p_scale_a_c1_f32_0);
-        HVX_Vector v_scale_a_c0_1 = Q6_V_lo_W(p_scale_a_c0_f32_1);
-        HVX_Vector v_scale_a_c1_1 = Q6_V_lo_W(p_scale_a_c1_f32_1);
-
-        HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f32_f32(v_scale_w_f32_0, v_scale_a_c0_0);
-        HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f32_f32(v_scale_w_f32_0, v_scale_a_c1_0);
-        HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f32_f32(v_scale_w_f32_1, v_scale_a_c0_1);
-        HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f32_f32(v_scale_w_f32_1, v_scale_a_c1_1);
-
-        HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
-        HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
-        HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
-        HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
-
-        v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
-        v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
-    }
-
-    for (; kt < n_k_tiles; kt++) {
+    for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
         const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
         const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
         const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
diff --git a/ggml/src/ggml-hexagon/htp/im2col-ops.c b/ggml/src/ggml-hexagon/htp/im2col-ops.c
index 26af14ed5..2e05cf3e1 100644
--- a/ggml/src/ggml-hexagon/htp/im2col-ops.c
+++ b/ggml/src/ggml-hexagon/htp/im2col-ops.c
@@ -14,7 +14,7 @@
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "hvx-utils.h"
-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hex-profile.h"
 #include "htp-vtcm.h"
 #include "htp-tensor.h"
@@ -86,8 +86,8 @@ static inline void htp_im2col_vtcm_layout_build(struct htp_im2col_vtcm_layout *
         const uint32_t OH                       = is_2D ? dst->ne[2] : 1;                                  \
         const uint32_t OW                       = dst->ne[1];                                              \
         const uint32_t patch_stride             = IC * KH * KW;                                            \
-        const float * restrict src_data         = (const float *) src1->data;                              \
-        DST_CTYPE * restrict dst_data           = (DST_CTYPE *) dst->data;                                 \
+        const float * restrict src_data         = (const float *) (uintptr_t) src1->data;                  \
+        DST_CTYPE * restrict dst_data           = (DST_CTYPE *) (uintptr_t) dst->data;                     \
         const uint32_t patch_end                = ictx->patch_base + ictx->npatches;                       \
         const uint32_t patch_start              = ictx->patch_base + ictx->npatches_per_thread * ith;      \
         const uint32_t patch_stop               = MIN(patch_start + ictx->npatches_per_thread, patch_end); \
@@ -153,155 +153,156 @@ IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx
 // this block's store-out) waits for both - safe because the ring is strict
 // FIFO and each buffer slot is only reused after its prior consumer (compute
 // or store-out) already finished in program order.
-#define IM2COL_BLOCKED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG)                                  \
-    static void FNAME(unsigned int nth, unsigned int ith, void * data) {                                             \
-        struct htp_im2col_context * ictx        = (struct htp_im2col_context *) data;                                \
-        struct htp_ops_context *    octx        = ictx->octx;                                                        \
-        struct htp_thread_trace * restrict tr   = &octx->ctx->trace[ith];                                            \
-        const struct htp_tensor * restrict src1 = octx->src[1];                                                      \
-        const struct htp_tensor * restrict dst  = octx->dst;                                                         \
-        const int32_t  s0 = octx->op_params[0], s1 = octx->op_params[1];                                             \
-        const int32_t  p0 = octx->op_params[2], p1 = octx->op_params[3];                                             \
-        const int32_t  d0 = octx->op_params[4], d1 = octx->op_params[5];                                             \
-        const int32_t  is_2D = octx->op_params[6] == 1;                                                              \
-        const uint32_t N     = is_2D ? src1->ne[3] : src1->ne[2];                                                    \
-        const uint32_t IC    = is_2D ? src1->ne[2] : src1->ne[1];                                                    \
-        const uint32_t IH    = is_2D ? src1->ne[1] : 1;                                                              \
-        const uint32_t IW    = src1->ne[0];                                                                          \
-        const uint32_t KH    = is_2D ? octx->src[0]->ne[1] : 1;                                                      \
-        const uint32_t KW    = octx->src[0]->ne[0];                                                                  \
-        const uint32_t OH    = is_2D ? dst->ne[2] : 1;                                                               \
-        const uint32_t OW    = dst->ne[1];                                                                           \
-        const uint32_t owb = ictx->pe_owb, Wb = ictx->pe_wb;                                                         \
-        const uint32_t patch_stride     = IC * KH * KW;                                                              \
-        const float * restrict src_data = (const float *) src1->data;                                                \
-        DST_CTYPE * restrict dst_data   = (DST_CTYPE *) dst->data;                                                   \
-        dma_queue *    dmaq             = octx->ctx->dma[ith];                                                       \
-        uint8_t *      srcb_base        = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread;                    \
-        uint8_t *      dstb_base        = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread;                    \
-        float *        srcb2[2]         = { (float *) srcb_base, (float *) (srcb_base + ictx->pe_src_row_bytes) };   \
-        DST_CTYPE *    dstb2[2]   = { (DST_CTYPE *) dstb_base, (DST_CTYPE *) (dstb_base + ictx->pe_dst_row_bytes) }; \
-        const uint32_t nrows      = N * OH;                                                                          \
-        const uint32_t per_thread = ictx->pe_rows_per_thread;                                                        \
-        const uint32_t row_start  = per_thread * ith;                                                                \
-        const uint32_t row_end    = MIN(row_start + per_thread, nrows);                                              \
-        if (row_start >= row_end)                                                                                    \
-            return;                                                                                                  \
-        const uint32_t nbpr         = (OW + owb - 1) / owb;                                                          \
-        const uint32_t nrows_local  = row_end - row_start;                                                           \
-        const uint32_t total_blocks = nrows_local * nbpr;                                                            \
-        for (uint32_t bi = 0; bi < total_blocks; bi++) {                                                             \
-            const uint32_t buf  = bi & 1u;                                                                           \
-            float *        srcb = srcb2[buf];                                                                        \
-            DST_CTYPE *    dstb = dstb2[buf];                                                                        \
-            const uint32_t r    = row_start + bi / nbpr;                                                             \
-            const uint32_t in   = r / OH;                                                                            \
-            const uint32_t ioh  = r % OH;                                                                            \
-            const uint32_t c0   = (bi % nbpr) * owb;                                                                 \
-            const uint32_t nb   = MIN(owb, OW - c0);                                                                 \
-            const int32_t  win0 = (int32_t) c0 * s0 - p0;                                                            \
-            if (bi == 0) {                                                                                           \
-                /* prologue: stage block 0 and wait - nothing to overlap with yet */                                 \
-                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                            \
-                    const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1;                                \
-                    if (iih < 0 || iih >= (int32_t) IH)                                                              \
-                        continue;                                                                                    \
-                    const int32_t lo = win0 < 0 ? -win0 : 0;                                                         \
-                    int32_t       hi = (int32_t) IW - win0;                                                          \
-                    if (hi > (int32_t) Wb)                                                                           \
-                        hi = (int32_t) Wb;                                                                           \
-                    if (hi <= lo)                                                                                    \
-                        continue;                                                                                    \
-                    const uint32_t cpw  = (uint32_t) (hi - lo);                                                      \
-                    float *        vdst = srcb + (uint64_t) ikh * Wb + (uint32_t) lo;                                \
-                    const float *  vsrc = src_data + ((uint64_t) (in * IC) * IH + iih) * IW + (win0 + lo);           \
-                    while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc),             \
-                                           (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float),       \
-                                           cpw * sizeof(float), IC)) {                                               \
-                        dma_queue_pop(dmaq);                                                                         \
-                    }                                                                                                \
-                }                                                                                                    \
-                dma_queue_flush(dmaq);                                                                               \
-            }                                                                                                        \
-            if (bi + 1 < total_blocks) {                                                                             \
-                /* prefetch: stage block bi+1 into the other slot; overlaps with this block's compute below */       \
-                const uint32_t nbuf  = 1u - buf;                                                                     \
-                float *        nsrcb = srcb2[nbuf];                                                                  \
-                const uint32_t nr    = row_start + (bi + 1) / nbpr;                                                  \
-                const uint32_t nin   = nr / OH;                                                                      \
-                const uint32_t nioh  = nr % OH;                                                                      \
-                const uint32_t nc0   = ((bi + 1) % nbpr) * owb;                                                      \
-                const int32_t  nwin0 = (int32_t) nc0 * s0 - p0;                                                      \
-                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                            \
-                    const int32_t iih = (int32_t) nioh * s1 + (int32_t) ikh * d1 - p1;                               \
-                    if (iih < 0 || iih >= (int32_t) IH)                                                              \
-                        continue;                                                                                    \
-                    const int32_t lo = nwin0 < 0 ? -nwin0 : 0;                                                       \
-                    int32_t       hi = (int32_t) IW - nwin0;                                                         \
-                    if (hi > (int32_t) Wb)                                                                           \
-                        hi = (int32_t) Wb;                                                                           \
-                    if (hi <= lo)                                                                                    \
-                        continue;                                                                                    \
-                    const uint32_t cpw  = (uint32_t) (hi - lo);                                                      \
-                    float *        vdst = nsrcb + (uint64_t) ikh * Wb + (uint32_t) lo;                               \
-                    const float *  vsrc = src_data + ((uint64_t) (nin * IC) * IH + iih) * IW + (nwin0 + lo);         \
-                    while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc),             \
-                                           (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float),       \
-                                           cpw * sizeof(float), IC)) {                                               \
-                        dma_queue_pop(dmaq);                                                                         \
-                    }                                                                                                \
-                }                                                                                                    \
-            }                                                                                                        \
-            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                    \
-            for (uint32_t j = 0; j < nb; j++) {                                                                      \
-                const uint32_t iow       = c0 + j;                                                                   \
-                DST_CTYPE *    dst_patch = dstb + (uint64_t) j * patch_stride;                                       \
-                const int32_t  iiw0      = (int32_t) iow * s0 - p0;                                                  \
-                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                            \
-                    const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1;                                \
-                    const int     okh = (iih >= 0 && iih < (int32_t) IH);                                            \
-                    for (uint32_t iic = 0; iic < IC; iic++) {                                                        \
-                        DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW;                                \
-                        if (!okh) {                                                                                  \
-                            SPLAT_FN(out_run, 0.0f, KW);                                                             \
-                            continue;                                                                                \
-                        }                                                                                            \
-                        const float * vrow = srcb + ((uint64_t) (iic * KH + ikh)) * Wb; /* col win0 at idx 0*/       \
-                        if (d0 == 1) {                                                                               \
-                            /* contiguous run within the staged window: [lo,hi) in-bounds, tails zero pad */         \
-                            const int32_t lo = iiw0 < 0 ? -iiw0 : 0;                                                 \
-                            int32_t       hi = (int32_t) IW - iiw0;                                                  \
-                            if (hi > (int32_t) KW) {                                                                 \
-                                hi = (int32_t) KW;                                                                   \
-                            }                                                                                        \
-                            if (hi <= lo) {                                                                          \
-                                SPLAT_FN(out_run, 0.0f, KW);                                                         \
-                            } else {                                                                                 \
-                                if (lo > 0) {                                                                        \
-                                    SPLAT_FN(out_run, 0.0f, (uint32_t) lo);                                          \
-                                }                                                                                    \
-                                COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (vrow + (iiw0 + lo - win0)),   \
-                                        (uint32_t) (hi - lo));                                                       \
-                                if (hi < (int32_t) KW) {                                                             \
-                                    SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi));                              \
-                                }                                                                                    \
-                            }                                                                                        \
-                            continue;                                                                                \
-                        }                                                                                            \
-                        for (uint32_t ikw = 0; ikw < KW; ikw++) {                                                    \
-                            const int32_t iiw = iiw0 + (int32_t) ikw * d0;                                           \
-                            out_run[ikw] =                                                                           \
-                                (iiw < 0 || iiw >= (int32_t) IW) ? (DST_CTYPE) 0.0f : (DST_CTYPE) vrow[iiw - win0];  \
-                        }                                                                                            \
-                    }                                                                                                \
-                }                                                                                                    \
-            }                                                                                                        \
-            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                     \
-            DST_CTYPE * ddr = dst_data + ((uint64_t) (in * OH + ioh) * OW + c0) * patch_stride;                      \
-            dma_queue_push_vtcm_to_ddr(dmaq, dma_make_ptr((uint8_t *) ddr, (uint8_t *) dstb),                        \
-                                       nb * patch_stride * (DST_ELEM), nb * patch_stride * (DST_ELEM), 1);           \
-            dma_queue_flush(dmaq);                                                                                   \
-        }                                                                                                            \
+#define IM2COL_BLOCKED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG)                                            \
+    static void FNAME(unsigned int nth, unsigned int ith, void * data) {                                                       \
+        struct htp_im2col_context * ictx        = (struct htp_im2col_context *) data;                                          \
+        struct htp_ops_context *    octx        = ictx->octx;                                                                  \
+        struct htp_thread_trace * restrict tr   = &octx->ctx->trace[ith];                                                      \
+        const struct htp_tensor * restrict src1 = octx->src[1];                                                                \
+        const struct htp_tensor * restrict dst  = octx->dst;                                                                   \
+        const int32_t  s0 = octx->op_params[0], s1 = octx->op_params[1];                                                       \
+        const int32_t  p0 = octx->op_params[2], p1 = octx->op_params[3];                                                       \
+        const int32_t  d0 = octx->op_params[4], d1 = octx->op_params[5];                                                       \
+        const int32_t  is_2D = octx->op_params[6] == 1;                                                                        \
+        const uint32_t N     = is_2D ? src1->ne[3] : src1->ne[2];                                                              \
+        const uint32_t IC    = is_2D ? src1->ne[2] : src1->ne[1];                                                              \
+        const uint32_t IH    = is_2D ? src1->ne[1] : 1;                                                                        \
+        const uint32_t IW    = src1->ne[0];                                                                                    \
+        const uint32_t KH    = is_2D ? octx->src[0]->ne[1] : 1;                                                                \
+        const uint32_t KW    = octx->src[0]->ne[0];                                                                            \
+        const uint32_t OH    = is_2D ? dst->ne[2] : 1;                                                                         \
+        const uint32_t OW    = dst->ne[1];                                                                                     \
+        const uint32_t owb = ictx->pe_owb, Wb = ictx->pe_wb;                                                                   \
+        const uint32_t patch_stride     = IC * KH * KW;                                                                        \
+        const dma_addr_t src_data       = src1->data;                                                                          \
+        const dma_addr_t dst_data       = dst->data;                                                                           \
+        dma_queue *    dmaq             = octx->ctx->dma[ith];                                                                 \
+        uint8_t *      srcb_base        = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread;                              \
+        uint8_t *      dstb_base        = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread;                              \
+        float *        srcb2[2]         = { (float *) srcb_base, (float *) (srcb_base + ictx->pe_src_row_bytes) };             \
+        DST_CTYPE *    dstb2[2]   = { (DST_CTYPE *) dstb_base, (DST_CTYPE *) (dstb_base + ictx->pe_dst_row_bytes) };           \
+        const uint32_t nrows      = N * OH;                                                                                    \
+        const uint32_t per_thread = ictx->pe_rows_per_thread;                                                                  \
+        const uint32_t row_start  = per_thread * ith;                                                                          \
+        const uint32_t row_end    = MIN(row_start + per_thread, nrows);                                                        \
+        if (row_start >= row_end)                                                                                              \
+            return;                                                                                                            \
+        const uint32_t nbpr         = (OW + owb - 1) / owb;                                                                    \
+        const uint32_t nrows_local  = row_end - row_start;                                                                     \
+        const uint32_t total_blocks = nrows_local * nbpr;                                                                      \
+        for (uint32_t bi = 0; bi < total_blocks; bi++) {                                                                       \
+            const uint32_t buf  = bi & 1u;                                                                                     \
+            float *        srcb = srcb2[buf];                                                                                  \
+            DST_CTYPE *    dstb = dstb2[buf];                                                                                  \
+            const uint32_t r    = row_start + bi / nbpr;                                                                       \
+            const uint32_t in   = r / OH;                                                                                      \
+            const uint32_t ioh  = r % OH;                                                                                      \
+            const uint32_t c0   = (bi % nbpr) * owb;                                                                           \
+            const uint32_t nb   = MIN(owb, OW - c0);                                                                           \
+            const int32_t  win0 = (int32_t) c0 * s0 - p0;                                                                      \
+            if (bi == 0) {                                                                                                     \
+                /* prologue: stage block 0 and wait - nothing to overlap with yet */                                           \
+                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                                      \
+                    const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1;                                          \
+                    if (iih < 0 || iih >= (int32_t) IH)                                                                        \
+                        continue;                                                                                              \
+                    const int32_t lo = win0 < 0 ? -win0 : 0;                                                                   \
+                    int32_t       hi = (int32_t) IW - win0;                                                                    \
+                    if (hi > (int32_t) Wb)                                                                                     \
+                        hi = (int32_t) Wb;                                                                                     \
+                    if (hi <= lo)                                                                                              \
+                        continue;                                                                                              \
+                    const uint32_t cpw  = (uint32_t) (hi - lo);                                                                \
+                    float *        vdst = srcb + (size_t) ikh * Wb + lo;                                                       \
+                    const dma_addr_t vsrc = src_data + (size_t) (((in * IC) * IH + iih) * IW + (win0 + lo)) * sizeof(float);   \
+                    while (!dma_queue_push(dmaq, dma_make_data(vdst, vsrc),                                                    \
+                                           (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float),                 \
+                                           cpw * sizeof(float), IC)) {                                                         \
+                        dma_queue_pop(dmaq);                                                                                   \
+                    }                                                                                                          \
+                }                                                                                                              \
+                dma_queue_flush(dmaq);                                                                                         \
+            }                                                                                                                  \
+            if (bi + 1 < total_blocks) {                                                                                       \
+                /* prefetch: stage block bi+1 into the other slot; overlaps with this block's compute below */                 \
+                const uint32_t nbuf  = 1u - buf;                                                                               \
+                float *        nsrcb = srcb2[nbuf];                                                                            \
+                const uint32_t nr    = row_start + (bi + 1) / nbpr;                                                            \
+                const uint32_t nin   = nr / OH;                                                                                \
+                const uint32_t nioh  = nr % OH;                                                                                \
+                const uint32_t nc0   = ((bi + 1) % nbpr) * owb;                                                                \
+                const int32_t  nwin0 = (int32_t) nc0 * s0 - p0;                                                                \
+                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                                      \
+                    const int32_t iih = (int32_t) nioh * s1 + (int32_t) ikh * d1 - p1;                                         \
+                    if (iih < 0 || iih >= (int32_t) IH)                                                                        \
+                        continue;                                                                                              \
+                    const int32_t lo = nwin0 < 0 ? -nwin0 : 0;                                                                 \
+                    int32_t       hi = (int32_t) IW - nwin0;                                                                   \
+                    if (hi > (int32_t) Wb)                                                                                     \
+                        hi = (int32_t) Wb;                                                                                     \
+                    if (hi <= lo)                                                                                              \
+                        continue;                                                                                              \
+                    const uint32_t cpw  = (uint32_t) (hi - lo);                                                                \
+                    float *        vdst = nsrcb + (size_t) ikh * Wb + lo;                                                      \
+                    const dma_addr_t vsrc = src_data + (size_t) (((nin * IC) * IH + iih) * IW + (nwin0 + lo)) * sizeof(float); \
+                    while (!dma_queue_push(dmaq, dma_make_data(vdst, vsrc),                                                    \
+                                           (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float),                 \
+                                           cpw * sizeof(float), IC)) {                                                         \
+                        dma_queue_pop(dmaq);                                                                                   \
+                    }                                                                                                          \
+                }                                                                                                              \
+            }                                                                                                                  \
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                              \
+            for (uint32_t j = 0; j < nb; j++) {                                                                                \
+                const uint32_t iow       = c0 + j;                                                                             \
+                DST_CTYPE *    dst_patch = dstb + (uint64_t) j * patch_stride;                                                 \
+                const int32_t  iiw0      = (int32_t) iow * s0 - p0;                                                            \
+                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                                      \
+                    const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1;                                          \
+                    const int     okh = (iih >= 0 && iih < (int32_t) IH);                                                      \
+                    for (uint32_t iic = 0; iic < IC; iic++) {                                                                  \
+                        DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW;                                          \
+                        if (!okh) {                                                                                            \
+                            SPLAT_FN(out_run, 0.0f, KW);                                                                       \
+                            continue;                                                                                          \
+                        }                                                                                                      \
+                        const float * vrow = srcb + ((uint64_t) (iic * KH + ikh)) * Wb; /* col win0 at idx 0*/                 \
+                        if (d0 == 1) {                                                                                         \
+                            /* contiguous run within the staged window: [lo,hi) in-bounds, tails zero pad */                   \
+                            const int32_t lo = iiw0 < 0 ? -iiw0 : 0;                                                           \
+                            int32_t       hi = (int32_t) IW - iiw0;                                                            \
+                            if (hi > (int32_t) KW) {                                                                           \
+                                hi = (int32_t) KW;                                                                             \
+                            }                                                                                                  \
+                            if (hi <= lo) {                                                                                    \
+                                SPLAT_FN(out_run, 0.0f, KW);                                                                   \
+                            } else {                                                                                           \
+                                if (lo > 0) {                                                                                  \
+                                    SPLAT_FN(out_run, 0.0f, (uint32_t) lo);                                                    \
+                                }                                                                                              \
+                                COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (vrow + (iiw0 + lo - win0)),             \
+                                        (uint32_t) (hi - lo));                                                                 \
+                                if (hi < (int32_t) KW) {                                                                       \
+                                    SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi));                                        \
+                                }                                                                                              \
+                            }                                                                                                  \
+                            continue;                                                                                          \
+                        }                                                                                                      \
+                        for (uint32_t ikw = 0; ikw < KW; ikw++) {                                                              \
+                            const int32_t iiw = iiw0 + (int32_t) ikw * d0;                                                     \
+                            out_run[ikw] =                                                                                     \
+                                (iiw < 0 || iiw >= (int32_t) IW) ? (DST_CTYPE) 0.0f : (DST_CTYPE) vrow[iiw - win0];            \
+                        }                                                                                                      \
+                    }                                                                                                          \
+                }                                                                                                              \
+            }                                                                                                                  \
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                               \
+            const dma_addr_t ddr = dst_data + (size_t) ((in * OH + ioh) * OW + c0) * patch_stride * (DST_ELEM);                \
+            dma_queue_push(dmaq, dma_make_data(ddr, dstb),                                                                     \
+                           nb * patch_stride * (DST_ELEM), nb * patch_stride * (DST_ELEM),                                     \
+                           nb * patch_stride * (DST_ELEM), 1);                                                                 \
+            dma_queue_flush(dmaq);                                                                                             \
+        }                                                                                                                      \
     }
 IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_thread,     __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "blk-dma-f16")
 IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float,  hvx_copy_f32_uu,     hvx_splat_f32_u, sizeof(float),  "blk-dma-f32")
@@ -309,75 +310,76 @@ IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float,  hvx_copy_f32_uu,
 // Exact-tiling patch-embed DMA fast path (s0==KW, p0=0, d0=1; and 2D s1==KH,
 // p1=0, d1=1). Intentionally reads no stride/pad/dilation params so the inner
 // copy stays tight and fully hoisted - do NOT graft the general gather in here.
-#define IM2COL_PATCHEMBED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG)                               \
-    static void FNAME(unsigned int nth, unsigned int ith, void * data) {                                             \
-        struct htp_im2col_context * ictx        = (struct htp_im2col_context *) data;                                \
-        struct htp_ops_context *    octx        = ictx->octx;                                                        \
-        struct htp_thread_trace * restrict tr   = &octx->ctx->trace[ith];                                            \
-        const struct htp_tensor * restrict src1 = octx->src[1];                                                      \
-        const struct htp_tensor * restrict dst  = octx->dst;                                                         \
-        const int32_t  is_2D                    = octx->op_params[6] == 1;                                           \
-        const uint32_t N                        = is_2D ? src1->ne[3] : src1->ne[2];                                 \
-        const uint32_t IC                       = is_2D ? src1->ne[2] : src1->ne[1];                                 \
-        const uint32_t IH                       = is_2D ? src1->ne[1] : 1;                                           \
-        const uint32_t IW                       = src1->ne[0];                                                       \
-        const uint32_t KH                       = is_2D ? octx->src[0]->ne[1] : 1;                                   \
-        const uint32_t KW                       = octx->src[0]->ne[0];                                               \
-        const uint32_t OH                       = is_2D ? dst->ne[2] : 1;                                            \
-        const uint32_t OW                       = dst->ne[1];                                                        \
-        const uint32_t patch_stride             = IC * KH * KW;                                                      \
-        const float * restrict src_data         = (const float *) src1->data;                                        \
-        DST_CTYPE * restrict dst_data           = (DST_CTYPE *) dst->data;                                           \
-        dma_queue *    dmaq                     = octx->ctx->dma[ith];                                               \
-        uint8_t *      src_base                 = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread;            \
-        uint8_t *      dst_base                 = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread;            \
-        float *        srcb                     = (float *) src_base;                                                \
-        DST_CTYPE *    dstb                     = (DST_CTYPE *) dst_base;                                            \
-        const uint32_t row_end_max              = ictx->pe_row_base + ictx->pe_nrows;                                \
-        const uint32_t per_thread               = ictx->pe_rows_per_thread;                                          \
-        const uint32_t row_start                = ictx->pe_row_base + per_thread * ith;                              \
-        const uint32_t row_end                  = MIN(row_start + per_thread, row_end_max);                          \
-        if (row_start >= row_end)                                                                                    \
-            return;                                                                                                  \
-        for (uint32_t r = row_start; r < row_end; r++) {                                                             \
-            const uint32_t in  = r / OH;                                                                             \
-            const uint32_t ioh = r % OH;                                                                             \
-            for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                                \
-                int32_t iih = (int32_t) ioh * (int32_t) KH + (int32_t) ikh;                                          \
-                int     ok  = (iih >= 0 && iih < (int32_t) IH);                                                      \
-                for (uint32_t iic = 0; iic < IC; iic++) {                                                            \
-                    float *       vdst = srcb + ((uint64_t) (iic * KH + ikh)) * IW;                                  \
-                    const float * _vsrc =                                                                            \
-                        ok ? (src_data + ((uint64_t) (in * IC + iic) * IH + iih) * IW) : (const float *) vdst;       \
-                    dma_queue_push_ddr_to_vtcm(                                                                      \
-                        dmaq, dma_make_ptr((uint8_t *) vdst, ok ? (const uint8_t *) _vsrc : (const uint8_t *) vdst), \
-                        IW * sizeof(float), IW * sizeof(float), ok ? 1 : 0);                                         \
-                }                                                                                                    \
-            }                                                                                                        \
-            for (uint32_t i = 0; i < IC * KH; i++)                                                                   \
-                dma_queue_pop(dmaq);                                                                                 \
-            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                    \
-            for (uint32_t iow = 0; iow < OW; iow++) {                                                                \
-                DST_CTYPE * dst_patch = dstb + (uint64_t) iow * patch_stride;                                        \
-                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                            \
-                    int32_t iih = (int32_t) ioh * (int32_t) KH + (int32_t) ikh;                                      \
-                    for (uint32_t iic = 0; iic < IC; iic++) {                                                        \
-                        DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW;                                \
-                        if (iih < 0 || iih >= (int32_t) IH) {                                                        \
-                            SPLAT_FN(out_run, 0.0f, KW);                                                             \
-                            continue;                                                                                \
-                        }                                                                                            \
-                        const float * src_run = srcb + ((uint64_t) (iic * KH + ikh)) * IW + (uint64_t) iow * KW;     \
-                        COPY_FN((uint8_t *) out_run, (const uint8_t *) src_run, KW);                                 \
-                    }                                                                                                \
-                }                                                                                                    \
-            }                                                                                                        \
-            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                     \
-            DST_CTYPE * ddr_row = dst_data + ((uint64_t) (in * OH + ioh) * OW) * patch_stride;                       \
-            dma_queue_push_vtcm_to_ddr(dmaq, dma_make_ptr((uint8_t *) ddr_row, (uint8_t *) dstb),                    \
-                                       OW * patch_stride * (DST_ELEM), OW * patch_stride * (DST_ELEM), 1);           \
-            dma_queue_flush(dmaq);                                                                                   \
-        }                                                                                                            \
+#define IM2COL_PATCHEMBED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG)                           \
+    static void FNAME(unsigned int nth, unsigned int ith, void * data) {                                         \
+        struct htp_im2col_context * ictx        = (struct htp_im2col_context *) data;                            \
+        struct htp_ops_context *    octx        = ictx->octx;                                                    \
+        struct htp_thread_trace * restrict tr   = &octx->ctx->trace[ith];                                        \
+        const struct htp_tensor * restrict src1 = octx->src[1];                                                  \
+        const struct htp_tensor * restrict dst  = octx->dst;                                                     \
+        const int32_t  is_2D                    = octx->op_params[6] == 1;                                       \
+        const uint32_t N                        = is_2D ? src1->ne[3] : src1->ne[2];                             \
+        const uint32_t IC                       = is_2D ? src1->ne[2] : src1->ne[1];                             \
+        const uint32_t IH                       = is_2D ? src1->ne[1] : 1;                                       \
+        const uint32_t IW                       = src1->ne[0];                                                   \
+        const uint32_t KH                       = is_2D ? octx->src[0]->ne[1] : 1;                               \
+        const uint32_t KW                       = octx->src[0]->ne[0];                                           \
+        const uint32_t OH                       = is_2D ? dst->ne[2] : 1;                                        \
+        const uint32_t OW                       = dst->ne[1];                                                    \
+        const uint32_t patch_stride             = IC * KH * KW;                                                  \
+        const dma_addr_t src_data       = src1->data;                                                            \
+        const dma_addr_t dst_data       = dst->data;                                                             \
+        dma_queue *    dma_q            = octx->ctx->dma[ith];                                                   \
+        uint8_t *      src_base         = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread;                \
+        uint8_t *      dst_base         = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread;                \
+        float *        srcb             = (float *) src_base;                                                    \
+        DST_CTYPE *    dstb             = (DST_CTYPE *) dst_base;                                                \
+        const uint32_t row_end_max              = ictx->pe_row_base + ictx->pe_nrows;                            \
+        const uint32_t per_thread               = ictx->pe_rows_per_thread;                                      \
+        const uint32_t row_start                = ictx->pe_row_base + per_thread * ith;                          \
+        const uint32_t row_end                  = MIN(row_start + per_thread, row_end_max);                      \
+        if (row_start >= row_end)                                                                                \
+            return;                                                                                              \
+        for (uint32_t r = row_start; r < row_end; r++) {                                                         \
+            const uint32_t in  = r / OH;                                                                         \
+            const uint32_t ioh = r % OH;                                                                         \
+            for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                            \
+                int32_t iih = (int32_t) ioh * (int32_t) KH + (int32_t) ikh;                                      \
+                int     ok  = (iih >= 0 && iih < (int32_t) IH);                                                  \
+                for (uint32_t iic = 0; iic < IC; iic++) {                                                        \
+                    float *       vdst = srcb + (size_t) (iic * KH + ikh) * IW;                                  \
+                    const dma_addr_t vsrc = ok                                                                   \
+                        ? (src_data + (size_t) ((in * IC + iic) * IH + iih) * IW * sizeof(float))                \
+                        : src_data;                                                                              \
+                    dma_queue_push(dma_q, dma_make_data(vdst, vsrc),                                             \
+                                   IW * sizeof(float), IW * sizeof(float), IW * sizeof(float), ok ? 1 : 0);      \
+                }                                                                                                \
+            }                                                                                                    \
+            for (uint32_t i = 0; i < IC * KH; i++)                                                               \
+                dma_queue_pop(dma_q);                                                                            \
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                \
+            for (uint32_t iow = 0; iow < OW; iow++) {                                                            \
+                DST_CTYPE * dst_patch = dstb + (uint64_t) iow * patch_stride;                                    \
+                for (uint32_t ikh = 0; ikh < KH; ikh++) {                                                        \
+                    int32_t iih = (int32_t) ioh * (int32_t) KH + (int32_t) ikh;                                  \
+                    for (uint32_t iic = 0; iic < IC; iic++) {                                                    \
+                        DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW;                            \
+                        if (iih < 0 || iih >= (int32_t) IH) {                                                    \
+                            SPLAT_FN(out_run, 0.0f, KW);                                                         \
+                            continue;                                                                            \
+                        }                                                                                        \
+                        const float * src_run = srcb + ((uint64_t) (iic * KH + ikh)) * IW + (uint64_t) iow * KW; \
+                        COPY_FN((uint8_t *) out_run, (const uint8_t *) src_run, KW);                             \
+                    }                                                                                            \
+                }                                                                                                \
+            }                                                                                                    \
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);                                                 \
+            const dma_addr_t ddr_row = dst_data + (size_t) (in * OH + ioh) * OW * patch_stride * (DST_ELEM);     \
+            dma_queue_push(dma_q, dma_make_data(ddr_row, dstb),                                                  \
+                           OW * patch_stride * (DST_ELEM), OW * patch_stride * (DST_ELEM),                       \
+                           OW * patch_stride * (DST_ELEM), 1);                                                   \
+            dma_queue_flush(dma_q);                                                                              \
+        }                                                                                                        \
     }

 IM2COL_PATCHEMBED_DMA_BODY(im2col_patchembed_dma_thread,     __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "pe-dma-f16")
@@ -496,10 +498,6 @@ int op_im2col(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const int32_t  is_2D         = octx->op_params[6] == 1;
     const uint32_t N             = is_2D ? src1->ne[3] : src1->ne[2];
     const uint32_t OH            = is_2D ? dst->ne[2] : 1;
@@ -571,12 +569,14 @@ int op_im2col(struct htp_ops_context * octx) {
         }
     }
     // Fall through to pure-DDR.
-
-
     if (npatches == 0) {
         return HTP_STATUS_OK;
     }

+    if (htp_tensor_is_extended(src1) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     if (dst->type == HTP_TYPE_F16) {
         work_queue_run(octx->ctx->work_queue, im2col_patchembed_thread, &ictx, n_threads);
     } else {
diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c
index b324cfd3a..653c9a250 100644
--- a/ggml/src/ggml-hexagon/htp/main.c
+++ b/ggml/src/ggml-hexagon/htp/main.c
@@ -21,7 +21,7 @@
 #include <stdatomic.h>

 #include "hex-utils.h"
-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hmx-queue.h"

 #define GGML_COMMON_DECL_C
@@ -49,33 +49,74 @@ struct htp_handle {
     struct htp_context * ctx;
 };

-static inline void * htp_mmap(uint32_t fd, uint32_t size) {
+static inline uint64_t htp_mmap(uint32_t fd, uint64_t size, uint32_t flags) {
+#if __HVX_ARCH__ > 79
+    if (flags & HTP_BUF_EXTENDED) {
+        HAP_mem_req_payload_t payload;
+        memset(&payload, 0, sizeof(payload));
+        payload.request_id = HAP_MEM_MAP;
+        payload.mmap.len   = size;
+        payload.mmap.prot  = HAP_MEM_CACHE_NON_SHARED | HAP_PROT_READ;
+        payload.mmap.flags = HAP_MEM_FLAGS_EXTENDED_MAP;
+        payload.mmap.fd    = fd;
+
+        if (HAP_mem_request(&payload) != 0) {
+            FARF(ERROR, "extended mmap failed : fd %u size %llu", fd, (unsigned long long) size);
+            return 0;
+        }
+
+        return payload.mmap.dsp_va;
+    }
+#else
+    if (flags & HTP_BUF_EXTENDED) {
+        FARF(ERROR, "extended mmap is unsupported on v%d", __HVX_ARCH__);
+        return 0;
+    }
+#endif
+
+    if (size > UINT32_MAX) {
+        FARF(ERROR, "mmap failed : size %llu exceeds 32-bit limit", (unsigned long long) size);
+        return 0;
+    }
+
     void * va = (void *)-1;
     for (int retry = 0; retry < 2; retry++) {
 #if __HVX_ARCH__ > 73
-        va = HAP_mmap2(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0);
+        va = HAP_mmap2(NULL, (size_t) size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0);
 #else
         if (size > HTP_MMAP_MAX_VMEM) {
-            FARF(ERROR, "mmap failed : size %u exceeds 2GB limit for HAP_mmap", (uint32_t) size);
+            FARF(ERROR, "mmap failed : size %llu exceeds 2GB limit for HAP_mmap", (unsigned long long) size);
             abort();
         }
-        va = HAP_mmap(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0);
+        va = HAP_mmap(NULL, (int) size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0);
 #endif
         if (va != (void *)-1 && va != NULL) {
-            return va;
+            return (uint64_t) (uintptr_t) va;
         }
         if (retry == 0) {
-            FARF(HIGH, "mmap failed first try (va %p fd %u size %u), retrying...", va, fd, size);
+            FARF(HIGH, "mmap failed first try (va %p fd %u size %llu), retrying...", va, fd, (unsigned long long) size);
         }
     }
-    return NULL;
+    return 0;
 }

-static inline void htp_munmap(void * va, uint32_t size) {
+static inline void htp_munmap(uint64_t va, uint64_t size, uint32_t flags) {
+#if __HVX_ARCH__ > 79
+    if (flags & HTP_BUF_EXTENDED) {
+        HAP_mem_req_payload_t payload;
+        memset(&payload, 0, sizeof(payload));
+        payload.request_id   = HAP_MEM_UNMAP;
+        payload.munmap.dsp_va = va;
+        payload.munmap.len    = size;
+        HAP_mem_request(&payload);
+        return;
+    }
+#endif
+
 #if __HVX_ARCH__ > 73
-    HAP_munmap2(va, size);
+    HAP_munmap2((void *) (uintptr_t) va, (size_t) size);
 #else
-    HAP_munmap(va, size);
+    HAP_munmap((void *) (uintptr_t) va, (int) size);
 #endif
 }

@@ -160,10 +201,11 @@ AEEResult htp_iface_close(remote_handle64 handle) {
         // release the mmaps (if any)
         for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) {
             if (ctx->mmap[i].size) {
-                htp_munmap((void *) ctx->mmap[i].base, ctx->mmap[i].size);
+                htp_munmap(ctx->mmap[i].base, ctx->mmap[i].size, ctx->mmap[i].flags);
                 ctx->mmap[i].size = 0;
-                ctx->mmap[i].base = NULL;
+                ctx->mmap[i].base = 0;
                 ctx->mmap[i].fd   = -1;
+                ctx->mmap[i].flags = 0;
             }
         }

@@ -184,7 +226,7 @@ AEEResult htp_iface_close(remote_handle64 handle) {
     return AEE_SUCCESS;
 }

-AEEResult htp_iface_mmap(remote_handle64 handle, uint32_t fd, uint32_t size) {
+AEEResult htp_iface_mmap(remote_handle64 handle, uint32_t fd, uint64_t size) {
     struct htp_handle * h = (struct htp_handle *) handle;
     if (!h || !h->ctx) {
         return AEE_EBADPARM;
@@ -203,16 +245,17 @@ AEEResult htp_iface_mmap(remote_handle64 handle, uint32_t fd, uint32_t size) {
     for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) {
         struct htp_mmap *m = &ctx->mmap[i];
         if (!m->size) {
-            FARF(HIGH, "mmap : fd %u size %u", fd, size);
-            void *va = htp_mmap(fd, size);
-            if (va == NULL) {
-                FARF(ERROR, "mmap failed : fd %u size %u", fd, (uint32_t) size);
+            FARF(HIGH, "mmap : fd %u size %llu", fd, (unsigned long long) size);
+            uint64_t va = htp_mmap(fd, size, 0);
+            if (va == 0) {
+                FARF(ERROR, "mmap failed : fd %u size %llu", fd, (unsigned long long) size);
                 return AEE_EFAILED;
             }

-            m->base   = (uint64_t) va;
+            m->base   = va;
             m->fd     = fd;
             m->size   = size;
+            m->flags  = 0;

             return AEE_SUCCESS;
         }
@@ -231,11 +274,12 @@ AEEResult htp_iface_munmap(remote_handle64 handle, uint32 fd) {
     for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) {
         struct htp_mmap *m = &ctx->mmap[i];
         if (fd < 0 || m->fd == fd) {
-            FARF(HIGH, "unmmap : base %p fd %u size %u", (void*) m->base, m->fd, (uint32_t) m->size);
-            htp_munmap((void *) m->base, m->size);
+            FARF(HIGH, "unmmap : base 0x%llx fd %u size %llu", (unsigned long long) m->base, m->fd, (unsigned long long) m->size);
+            htp_munmap(m->base, m->size, m->flags);
             m->size   = 0;
             m->base   = NULL;
             m->fd     = -1;
+            m->flags  = 0;
         }
     }

@@ -394,8 +438,6 @@ AEEResult htp_iface_start(remote_handle64 handle, uint32_t sess_id, uint64_t dsp
     for (uint32_t i = 0; i < n_hvx; i++) {
         size_dma  = hex_align_up(size_dma, dma_queue_alignof());
         size_dma += dma_queue_sizeof(256);
-        size_dma  = hex_align_up(size_dma, dma_queue_alignof());
-        size_dma += dma_queue_alias_sizeof();
     }
     offset = offset_dma + size_dma;

@@ -538,16 +580,11 @@ AEEResult htp_iface_start(remote_handle64 handle, uint32_t sess_id, uint64_t dsp
     // Initialize DMA queues
     uint8_t * dma_ptr_curr = (uint8_t *) ((uintptr_t) block + offset_dma);
     size_t size_dma_q = dma_queue_sizeof(256);
-    size_t size_dma_alias = dma_queue_alias_sizeof();

     for (int i = 0; i < ctx->n_threads; i++) {
         dma_ptr_curr = (uint8_t *) hex_align_up((uintptr_t) dma_ptr_curr, dma_queue_alignof());
-        ctx->dma_cached[i] = dma_queue_init(dma_ptr_curr, 256, (uintptr_t) ctx->vtcm_base, ctx->vtcm_size, &ctx->trace[i]);
+        ctx->dma[i] = dma_queue_init(dma_ptr_curr, 256, &ctx->trace[i]);
         dma_ptr_curr += size_dma_q;
-
-        dma_ptr_curr = (uint8_t *) hex_align_up((uintptr_t) dma_ptr_curr, dma_queue_alignof());
-        ctx->dma[i] = dma_queue_alias_init(dma_ptr_curr, ctx->dma_cached[i], 1);
-        dma_ptr_curr += size_dma_alias;
     }

     ctx->ddr_spad_size = 512 * 1024; // 512 KB
@@ -608,8 +645,7 @@ AEEResult htp_iface_stop(remote_handle64 handle) {
     work_queue_free(ctx->work_queue);

     for (int i = 0; i < ctx->n_threads; i++) {
-        dma_queue_alias_free(ctx->dma[i]);
-        dma_queue_free(ctx->dma_cached[i]);
+        dma_queue_free(ctx->dma[i]);
     }

     if (ctx->hmx_queue) {
@@ -908,7 +944,7 @@ static inline bool reuse_buf(struct htp_context *ctx, uint32_t *m_reuse, struct

     for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) {
         struct htp_mmap *m = ctx->mmap + i;
-        if (m->size && m->fd == b->fd) {
+        if (m->size && m->fd == b->fd && m->flags == b->flags) {
             b->base   = m->base;
             *m_reuse |= (1 << i);
             return true;
@@ -920,11 +956,12 @@ static inline bool reuse_buf(struct htp_context *ctx, uint32_t *m_reuse, struct

 static inline void drop_mmap(struct htp_context *ctx, struct htp_mmap *m) {
     if (m->size) {
-        FARF(ALWAYS, "unmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size);
-        htp_munmap((void *) m->base, m->size);
+        FARF(ALWAYS, "unmap : fd %u base 0x%llx size %llu", m->fd, (unsigned long long) m->base, (unsigned long long) m->size);
+        htp_munmap(m->base, m->size, m->flags);
         m->size = 0;
         m->base = 0;
         m->fd   = -1;
+        m->flags = 0;
     }
 }

@@ -935,17 +972,18 @@ static inline bool mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) {
     for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) {
         struct htp_mmap *m = &ctx->mmap[i];
         if (!m->size) {
-            void *va = htp_mmap(b->fd, b->size);
-            if (va == NULL) {
-                FARF(HIGH, "mmap failed (will attempt defrag) : fd %u size %u", b->fd, (uint32_t) b->size);
+            uint64_t va = htp_mmap(b->fd, b->size, b->flags);
+            if (va == 0) {
+                FARF(HIGH, "mmap failed (will attempt defrag) : fd %u size %llu", b->fd, (unsigned long long) b->size);
                 return false;
             }

-            m->base   = b->base = (uint64_t) va;
+            m->base   = b->base = va;
             m->fd     = b->fd;
             m->size   = b->size;
+            m->flags  = b->flags;

-            FARF(ALWAYS, "mmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size);
+            FARF(ALWAYS, "mmap : fd %u base 0x%llx size %llu flags 0x%x", m->fd, (unsigned long long) m->base, (unsigned long long) m->size, m->flags);
             return true;
         }
     }
@@ -964,14 +1002,22 @@ static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uin
     // See what we can reuse
     for (uint32_t i=0; i < n_bufs; i++) {
         struct htp_buf_desc *b = bufs + i;
-        if (reuse_buf(ctx, &m_reuse, b)) { b_reuse++; } else { e_vmem += b->size; }
-        FARF(HIGH, "prep-buf #%u : pass0 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags);
+        if (reuse_buf(ctx, &m_reuse, b)) {
+            b_reuse++;
+        } else if (!(b->flags & HTP_BUF_EXTENDED)) {
+            e_vmem += b->size;
+        }
+        FARF(HIGH, "prep-buf #%u : pass0 fd %u base 0x%llx size %llu flags 0x%x", i, b->fd, (unsigned long long) b->base, (unsigned long long) b->size, b->flags);
     }

     if (b_reuse == n_bufs) return; // all bufs reuse existing mappings

     // See how much vmem we have mmaped right now
-    for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) { m_vmem += ctx->mmap[i].size; }
+    for (uint32_t i=0; i<HTP_MAX_MMAPS; i++) {
+        if (!(ctx->mmap[i].flags & HTP_BUF_EXTENDED)) {
+            m_vmem += ctx->mmap[i].size;
+        }
+    }

     FARF(HIGH, "prep-bufs : pass1 mmap-vmem %zu extra-vmem %zu max-vmem %zu : n-bufs %u b-reuse %u",
             (size_t) m_vmem, (size_t) e_vmem, (size_t) ctx->max_vmem, n_bufs, b_reuse);
@@ -980,7 +1026,9 @@ static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uin
         // Drop unused mappings
         for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) {
             bool used = m_reuse & (1<<i);
-            if (!used) { drop_mmap(ctx, ctx->mmap + i); }
+            if (!used && !(ctx->mmap[i].flags & HTP_BUF_EXTENDED)) {
+                drop_mmap(ctx, ctx->mmap + i);
+            }
         }
     }

@@ -992,35 +1040,40 @@ static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uin
             mmap_ok = false;
             break;
         }
-        FARF(HIGH, "prep-buf #%u : pass1 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags);
+        FARF(HIGH, "prep-buf #%u : pass1 fd %u base 0x%llx size %llu flags 0x%x", i, b->fd, (unsigned long long) b->base, (unsigned long long) b->size, b->flags);
     }

     if (!mmap_ok) {
-        // Attempt clean defragmentation: drop all mappings and remap (pass 2)
-        FARF(HIGH, "prep-bufs : dropping all mappings to defragment address space");
-        for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) { drop_mmap(ctx, ctx->mmap + i); }
+        // Attempt defragmentation: drop 32-bit mappings and remap (pass 2)
+        FARF(HIGH, "prep-bufs : dropping 32-bit mappings to defragment address space");
+        for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) {
+            if (!(ctx->mmap[i].flags & HTP_BUF_EXTENDED)) {
+                drop_mmap(ctx, ctx->mmap + i);
+            }
+        }

         for (uint32_t i=0; i < n_bufs; i++) {
             struct htp_buf_desc *b = bufs + i;
-            b->base = 0;
+            if (!(b->flags & HTP_BUF_EXTENDED)) {
+                b->base = 0;
+            }
             if (!mmap_buf(ctx, b)) {
-                FARF(ERROR, "prep-bufs : mmap failed after defragmentation (fd %u size %u)", b->fd, (uint32_t) b->size);
+                FARF(ERROR, "prep-bufs : mmap failed after defragmentation (fd %u size %llu)", b->fd, (unsigned long long) b->size);
                 abort();
             }
-            FARF(HIGH, "prep-buf #%u : pass2 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags);
+            FARF(HIGH, "prep-buf #%u : pass2 fd %u base 0x%llx size %llu flags 0x%x", i, b->fd, (unsigned long long) b->base, (unsigned long long) b->size, b->flags);
         }
     }
 }

 static void prep_tensor(struct htp_context *ctx, struct htp_buf_desc *bufs, struct htp_tensor *tens, uint32_t idx, struct htp_tensor *t) {
-    uint32_t offset = t->data;
-    uint32_t size   = t->size;
+    uint64_t offset = t->data;
     uint32_t bi     = t->bi;

-    t->data  = (uint32_t) (bufs[bi].base + offset);  // update data to the actual pointer
+    t->data  = bufs[bi].base + offset;  // update data to the actual pointer

-    FARF(HIGH, "prep-tensor #%u: bi %u offset %u size %u data %p : %u:%u:%u:%u", idx, t->bi, offset, t->size, (void*) t->data,
-        t->ne[0], t->ne[1], t->ne[3], t->ne[3]);
+    FARF(HIGH, "prep-tensor #%u: bi %u offset %llu size %u data 0x%llx : %u:%u:%u:%u", idx, t->bi, (unsigned long long) offset, t->size, (unsigned long long) t->data,
+        t->ne[0], t->ne[1], t->ne[2], t->ne[3]);
 }

 static void prep_tensors(struct htp_context *ctx, struct htp_buf_desc *bufs, struct htp_tensor *tens, uint32_t n_tens) {
@@ -1050,15 +1103,13 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_buf_desc * bufs
         uint16_t src_idx = op->src[i];
         if (src_idx == 0xffff) {
             octx->src[i]     = NULL;
-            octx->src_dma[i] = NULL;
             continue;
         }

         struct htp_tensor *src = tens + src_idx;
         octx->src[i]     = src;
-        octx->src_dma[i] = octx->ctx->dma; // FIXME: ? octx->ctx->dma_cached : octx->ctx->dma;

-        FARF(HIGH, "prep-src #%u: data %p size %u : %u:%u:%u:%u", op->src[i], (void*) src->data, src->size,
+        FARF(HIGH, "prep-src #%u: data 0x%llx size %u : %u:%u:%u:%u", op->src[i], (unsigned long long) src->data, src->size,
             src->ne[0], src->ne[1], src->ne[2], src->ne[3]);
     }

@@ -1069,14 +1120,12 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_buf_desc * bufs
         uint16_t dst_idx = op->dst[i];
         if (dst_idx == 0xffff) {
             octx->dsts[i]    = NULL;
-            octx->dst_dma[i] = NULL;
             continue;
         }
         struct htp_tensor *dst = tens + dst_idx;
         octx->dsts[i]    = dst;
-        octx->dst_dma[i] = octx->ctx->dma; // FIXME: ? octx->ctx->dma_cached : octx->ctx->dma;

-        FARF(HIGH, "prep-dst[%u] #%u: data %p size %u : %u:%u:%u:%u", i, dst_idx, (void*) dst->data, dst->size,
+        FARF(HIGH, "prep-dst[%u] #%u: data 0x%llx size %u : %u:%u:%u:%u", i, dst_idx, (unsigned long long) dst->data, dst->size,
             dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
     }

diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c
index e16cfdcbe..a09bc7a28 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.c
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c
@@ -11,7 +11,7 @@
 #include <string.h>
 #include <stdatomic.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"
 #include "hvx-dump.h"
 #include "hvx-arith.h"
@@ -25,22 +25,13 @@
 #include "matmul-ops.h"
 #include "htp-vtcm.h"

-static void hvx_tensor_add_f32_grid(
-    const struct htp_tensor * restrict dst,
-    const struct htp_tensor * restrict src2,
-    uint32_t start_row,
-    uint32_t end_row,
-    uint32_t start_col,
-    uint32_t end_col,
-    const struct fastdiv_values * div_ne11_12,
-    const struct fastdiv_values * div_ne11
-);
-
 typedef struct {
     float        *dst;
-    const float  *src2;
+    dma_addr_t    src2_addr;
+    size_t        src2_bytes;
     const float  *activation;
-    const __fp16 *weight;
+    dma_addr_t    weight;
+    dma_queue *   weight_dma;
     int           m;
     int           k;
     int           n;
@@ -66,6 +57,15 @@ typedef struct {
     struct fastdiv_values div_r3;
 } hmx_mm_f16_f32_batched_params_t;

+static bool htp_matmul_has_extended_weight(const struct htp_ops_context * octx, uint32_t n_weights) {
+    for (uint32_t i = 0; i < n_weights; ++i) {
+        if (htp_tensor_is_extended(octx->src[i])) {
+            return true;
+        }
+    }
+    return false;
+}
+
 struct htp_mm_context {
     const char * type;
     struct htp_ops_context * octx;
@@ -94,6 +94,8 @@ struct htp_mm_context {
     uint32_t src0_row_end;
     uint32_t src0_row_size_padded;
     uint32_t src1_nrows;
+    uint32_t cur_m_start;
+    uint32_t cur_m_rows;

     struct fastdiv_values mm_div_ne12_ne1;
     struct fastdiv_values mm_div_ne1;
@@ -228,7 +230,7 @@ static const uint8_t __attribute__((aligned(VLEN))) kvalues_mxfp4_lut[] = {
 #define htp_matmul_preamble                                         \
     struct htp_mm_context * mmctx  = data;                          \
     struct htp_ops_context * octx  = mmctx->octx;                   \
-    dma_queue *dma_queue           = octx->ctx->dma[ith];           \
+    dma_queue * dma_q              = octx->ctx->dma[ith];           \
     uint32_t src0_nrows_per_thread = mmctx->src0_nrows_per_thread;  \
     htp_matmul_tensors_preamble;

@@ -244,291 +246,213 @@ static inline void hvx_mm_run_quant_task(struct htp_mm_context * mmctx, unsigned
     }
 }

-// *** matmul with support for 4d tensors and full broadcasting
-
-static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
-    htp_matmul_preamble;
-
-    assert(ne12 % ne02 == 0);
-    assert(ne13 % ne03 == 0);
-
-    // This is the size of the first dimension of the result, so we can iterate that way. (see the ASSERT above, these are the same numbers)
-    const uint32_t nr0 = ne0;
-
-    // This is the size of the rest of the dimensions of the result
-    const uint32_t nr1 = ne1 * ne2 * ne3;
-
-    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;
-
-    // distribute the thread work across the inner or outer loop based on which one is larger
-    uint32_t dr0, dr1, ith0, ith1;
-    if (nr0 > nr1) {
-        dr0  = fastdiv(src0_nrows + nth - 1, &octx->n_threads_div);
-        dr1  = nr1;
-        ith0 = ith;
-        ith1 = 0;
-    } else {
-        dr0  = src0_nrows;
-        dr1  = fastdiv(nr1 + nth - 1, &octx->n_threads_div);
-        ith0 = 0;
-        ith1 = ith;
-    }
-
-    const uint32_t ir0_start = mmctx->src0_row_start + dr0 * ith0;
-    const uint32_t ir0_end   = MIN(ir0_start + dr0, mmctx->src0_row_end);
-
-    const uint32_t ir1_start = dr1 * ith1;
-    const uint32_t ir1_end   = MIN(ir1_start + dr1, nr1);
-
-    // no work for this thread
-    if (ir0_start >= ir0_end || ir1_start >= ir1_end) {
-        return;
-    }
-
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0_start);
-
-    const uint32_t blck_0 = 64;
-    const uint32_t blck_1 = 64;
-
-    for (uint32_t iir1 = ir1_start; iir1 < ir1_end; iir1 += blck_1) {
-        for (uint32_t iir0 = ir0_start; iir0 < ir0_end; iir0 += blck_0) {
-            for (uint32_t ir1 = iir1; ir1 < MIN(iir1 + blck_1, ir1_end); ir1++) {
-                const uint32_t i13 = fastdiv(ir1, &mmctx->mm_div_ne12_ne1);
-                const uint32_t i12 = fastdiv(ir1 - i13 * ne12 * ne1, &mmctx->mm_div_ne1);
-                const uint32_t i11 = (ir1 - i13 * ne12 * ne1 - i12 * ne1);
-
-                // broadcast src0 into src1
-                const uint32_t i03 = fastdiv(i13, &mmctx->mm_div_r3);
-                const uint32_t i02 = fastdiv(i12, &mmctx->mm_div_r2);
-
-                const uint32_t i1 = i11;
-                const uint32_t i2 = i12;
-                const uint32_t i3 = i13;
-
-                const uint8_t * restrict src0_base = (const uint8_t *) src0->data + (0 + i02 * nb02 + i03 * nb03);
-                const uint8_t * restrict src1_col  = (const uint8_t *) src1->data + (i11 * nb11 + i12 * nb12 + i13 * nb13);
-                float * dst_col = (float *) ((uint8_t * restrict) dst->data + (i1 * nb1 + i2 * nb2 + i3 * nb3));
-
-                const uint32_t ir0_block_end = MIN(iir0 + blck_0, ir0_end);
-                for (uint32_t ir0 = iir0; ir0 < ir0_block_end; ir0++) {
-                    const uint8_t * restrict src0_row = src0_base + ir0 * nb01;
-                    mmctx->vec_dot_1x1(ne00, &dst_col[ir0], src0_row, src1_col);
-                }
-            }
-        }
-    }

-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0_start);
-    if (src2) {
-        hvx_tensor_add_f32_grid(dst, src2, ir1_start, ir1_end, ir0_start, ir0_end, &mmctx->mm_div_ne12_ne1, &mmctx->mm_div_ne1);
-    }
-}

 // hvx kernels first: the HMX Q6_K dequantizer reuses unpack_q6_k_group from there
 #include "hvx-mm-kernels-tiled.h"
+#include "hvx-mm-kernels-float.h"
 #include "hmx-mm-kernels-tiled.h"
-#include "hvx-mm-kernels-flat.h"

 // Specialized repacked matmul macros
-#define MATMUL_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1)                                                              \
-static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                        \
-    htp_matmul_preamble;                                                                                                          \
-                                                                                                                                  \
-    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                      \
-    const uint32_t src1_nrows = ne11 * ne12 * ne13;                                                                               \
-                                                                                                                                  \
-    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                         \
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                            \
-                                                                                                                                  \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                        \
-                                                                                                                                  \
-    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                      \
-    const uint32_t n_prefetch = kparams->n_prefetch;                                                                              \
-    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);                         \
-                                                                                                                                  \
-    const size_t dst_row_size  = nb1;                                                                                             \
-    const size_t src1_row_size = nb11;                                                                                            \
-    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                           \
-    const size_t src2_stride = src2 ? ((src2->ne[1] == 1) ? 0 : src2->nb[1]) : 0;                                                 \
-                                                                                                                                  \
-    uint8_t * restrict vtcm_dst_ptr  = mmctx->vtcm_dst  + mmctx->vtcm_dst_size_per_thread  * ith;                                 \
-    uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                                 \
-    uint8_t * restrict src1_data = mmctx->vtcm_src1;                                                                              \
-                                                                                                                                  \
-    const uint8_t * restrict src0_row = (const uint8_t *) src0->data;                                                             \
-                                                                                                                                  \
-    const uint32_t tile_size = TILE_SIZE;                                                                                         \
-    const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                              \
-                                                                                                                                  \
-    uint32_t n_k_tiles_w = ne00 / 32;                                                                                             \
-    uint32_t n_k_tiles_a = ne10 / 32;                                                                                             \
-    uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                           \
-    uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                    \
-                                                                                                                                  \
-    uint32_t ct_start = src0_start_row / 32;                                                                                      \
-    uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                                 \
-                                                                                                                                  \
-    uint32_t push_ct = ct_start;                                                                                                  \
-    if (src0_start_row < src0_end_row) {                                                                                          \
-        for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                                \
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned,                            \
-                           src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a);          \
-        }                                                                                                                         \
-    }                                                                                                                             \
-                                                                                                                                  \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                            \
-                                                                                                                                  \
-    if (src0_start_row >= src0_end_row) {                                                                                         \
-        return;                                                                                                                   \
-    }                                                                                                                             \
-                                                                                                                                  \
-    for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                             \
-        const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;                                                                    \
-                                                                                                                                  \
-        int valid_rows = (int)ne0 - (int)(ct * 32);                                                                               \
-        valid_rows = MIN(32, MAX(0, valid_rows));                                                                                 \
-                                                                                                                                  \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                    \
-        uint32_t ir1 = 0;                                                                                                         \
-        for (; ir1 + 1 < src1_nrows; ir1 += 2) {                                                                                  \
-            const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);                           \
-            const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride);                           \
-            float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size));                                         \
-            float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size));                                         \
-                                                                                                                                  \
-            float * dst_ptr0 = &dst_row0[ct * 32];                                                                                \
-            float * dst_ptr1 = &dst_row1[ct * 32];                                                                                \
-                                                                                                                                  \
-            const float * src2_ptr0 = NULL;                                                                                       \
-            const float * src2_ptr1 = NULL;                                                                                       \
-            if (src2) {                                                                                                           \
-                const float * restrict src2_row0 = (const float *) ((const uint8_t *) src2->data + ((ir1+0) * src2_stride));      \
-                const float * restrict src2_row1 = (const float *) ((const uint8_t *) src2->data + ((ir1+1) * src2_stride));      \
-                src2_ptr0 = &src2_row0[ct * 32];                                                                                  \
-                src2_ptr1 = &src2_row1[ct * 32];                                                                                  \
-            }                                                                                                                     \
-            DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, src2_ptr0, src2_ptr1);                    \
-        }                                                                                                                         \
-                                                                                                                                  \
-        for (; ir1 < src1_nrows; ++ir1) {                                                                                         \
-            const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);                                \
-            float * restrict dst_row          = (float *) (dst->data + (ir1 * dst_row_size));                                     \
-            float * dst_ptr = &dst_row[ct * 32];                                                                                  \
-                                                                                                                                  \
-            const float * src2_ptr = NULL;                                                                                        \
-            if (src2) {                                                                                                           \
-                const float * restrict src2_row = (const float *) ((const uint8_t *) src2->data + (ir1 * src2_stride));           \
-                src2_ptr = &src2_row[ct * 32];                                                                                    \
-            }                                                                                                                     \
-            DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, src2_ptr);                                                       \
-        }                                                                                                                         \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                     \
-                                                                                                                                  \
-        if (push_ct < ct_end) {                                                                                                   \
-            dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),                      \
-                           aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                                 \
-            push_ct++;                                                                                                            \
-        }                                                                                                                         \
-    }                                                                                                                             \
+#define MATMUL_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1)                                                                       \
+static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                                 \
+    htp_matmul_preamble;                                                                                                                   \
+                                                                                                                                           \
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                               \
+    const uint32_t src1_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : (ne11 * ne12 * ne13);                                              \
+    const uint32_t cur_m_start = mmctx->cur_m_start;                                                                                       \
+                                                                                                                                           \
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                                  \
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                                     \
+                                                                                                                                           \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                                 \
+                                                                                                                                           \
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                               \
+    const uint32_t n_prefetch = kparams->n_prefetch;                                                                                       \
+    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);                                  \
+                                                                                                                                           \
+    const size_t dst_row_size  = nb1;                                                                                                      \
+    const size_t src1_row_size = nb11;                                                                                                     \
+    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                                    \
+    const size_t src2_stride = src2 ? ((src2->ne[1] == 1) ? 0 : src2->nb[1]) : 0;                                                          \
+                                                                                                                                           \
+    uint8_t * restrict vtcm_dst_ptr  = mmctx->vtcm_dst  + mmctx->vtcm_dst_size_per_thread  * ith;                                          \
+    uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                                          \
+    uint8_t * restrict src1_data = mmctx->vtcm_src1;                                                                                       \
+                                                                                                                                           \
+    const dma_addr_t src0_row = src0->data;                                                                                                \
+                                                                                                                                           \
+    const uint32_t tile_size = TILE_SIZE;                                                                                                  \
+    const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                                       \
+                                                                                                                                           \
+    uint32_t n_k_tiles_w = ne00 / 32;                                                                                                      \
+    uint32_t n_k_tiles_a = ne10 / 32;                                                                                                      \
+    uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                                    \
+    uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                             \
+                                                                                                                                           \
+    uint32_t ct_start = src0_start_row / 32;                                                                                               \
+    uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                                          \
+                                                                                                                                           \
+    uint32_t push_ct = ct_start;                                                                                                           \
+    if (src0_start_row < src0_end_row) {                                                                                                   \
+        for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                                         \
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned,                                        \
+                           src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                   \
+        }                                                                                                                                  \
+    }                                                                                                                                      \
+                                                                                                                                           \
+    hvx_mm_run_quant_task(mmctx, ith);                                                                                                     \
+                                                                                                                                           \
+    if (src0_start_row >= src0_end_row) {                                                                                                  \
+        return;                                                                                                                            \
+    }                                                                                                                                      \
+                                                                                                                                           \
+    for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                                      \
+        const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;                                                                        \
+                                                                                                                                           \
+        int valid_rows = (int)ne0 - (int)(ct * 32);                                                                                        \
+        valid_rows = MIN(32, MAX(0, valid_rows));                                                                                          \
+                                                                                                                                           \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                             \
+        uint32_t ir1 = 0;                                                                                                                  \
+        for (; ir1 + 1 < src1_nrows; ir1 += 2) {                                                                                           \
+            const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);                                    \
+            const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride);                                    \
+            float * restrict dst_row0 = (float *) (dst->data + ((cur_m_start + ir1+0) * dst_row_size));                                    \
+            float * restrict dst_row1 = (float *) (dst->data + ((cur_m_start + ir1+1) * dst_row_size));                                    \
+                                                                                                                                           \
+            float * dst_ptr0 = &dst_row0[ct * 32];                                                                                         \
+            float * dst_ptr1 = &dst_row1[ct * 32];                                                                                         \
+                                                                                                                                           \
+            const float * src2_ptr0 = NULL;                                                                                                \
+            const float * src2_ptr1 = NULL;                                                                                                \
+            if (src2) {                                                                                                                    \
+                const float * restrict src2_row0 = (const float *) ((const uint8_t *) src2->data + ((cur_m_start + ir1+0) * src2_stride)); \
+                const float * restrict src2_row1 = (const float *) ((const uint8_t *) src2->data + ((cur_m_start + ir1+1) * src2_stride)); \
+                src2_ptr0 = &src2_row0[ct * 32];                                                                                           \
+                src2_ptr1 = &src2_row1[ct * 32];                                                                                           \
+            }                                                                                                                              \
+            DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, src2_ptr0, src2_ptr1);                             \
+        }                                                                                                                                  \
+                                                                                                                                           \
+        for (; ir1 < src1_nrows; ++ir1) {                                                                                                  \
+            const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);                                         \
+            float * restrict dst_row          = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));                              \
+            float * dst_ptr = &dst_row[ct * 32];                                                                                           \
+                                                                                                                                           \
+            const float * src2_ptr = NULL;                                                                                                 \
+            if (src2) {                                                                                                                    \
+                const float * restrict src2_row = (const float *) ((const uint8_t *) src2->data + ((cur_m_start + ir1) * src2_stride));    \
+                src2_ptr = &src2_row[ct * 32];                                                                                             \
+            }                                                                                                                              \
+            DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, src2_ptr);                                                                \
+        }                                                                                                                                  \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                              \
+                                                                                                                                           \
+        if (push_ct < ct_end) {                                                                                                            \
+            dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),                                             \
+                           aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                                          \
+            push_ct++;                                                                                                                     \
+        }                                                                                                                                  \
+    }                                                                                                                                      \
 }

-#define MATVEC_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X1)                                                                       \
-static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                        \
-    htp_matmul_preamble;                                                                                                          \
-                                                                                                                                  \
-    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                      \
-                                                                                                                                  \
-    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                         \
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                            \
-                                                                                                                                  \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                        \
-                                                                                                                                  \
-    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                      \
-    const uint32_t n_prefetch = kparams->n_prefetch;                                                                              \
-    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);                         \
-                                                                                                                                  \
-    const size_t dst_row_size  = nb1;                                                                                             \
-    const size_t src1_row_size = nb11;                                                                                            \
-    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                           \
-                                                                                                                                  \
-    uint8_t * vtcm_dst_ptr  = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith;                                            \
-    uint8_t * vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                                          \
-    uint8_t * src1_data = mmctx->vtcm_src1;                                                                                       \
-                                                                                                                                  \
-    float * tmp = (float *) vtcm_dst_ptr;                                                                                         \
-                                                                                                                                  \
-    const uint8_t * restrict src0_row = (const uint8_t *) src0->data;                                                             \
-                                                                                                                                  \
-    const uint8_t * restrict src1_col = (const uint8_t *) src1_data;                                                              \
-    float * restrict dst_col          = (float *) dst->data;                                                                      \
-                                                                                                                                  \
-    const uint32_t tile_size = TILE_SIZE;                                                                                         \
-    const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                              \
-                                                                                                                                  \
-    uint32_t n_k_tiles_w = ne00 / 32;                                                                                             \
-    uint32_t n_k_tiles_a = ne10 / 32;                                                                                             \
-    uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                           \
-    uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                    \
-                                                                                                                                  \
-    uint32_t ct_start = src0_start_row / 32;                                                                                      \
-    uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                                 \
-                                                                                                                                  \
-    uint32_t push_ct = ct_start;                                                                                                  \
-    if (src0_start_row < src0_end_row) {                                                                                          \
-        if (src2) {                                                                                                               \
-            float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row;                                                  \
-            const float * src2_ptr = (const float *) src2->data + src0_start_row;                                                 \
-            int slice_size = (int)MIN(src0_end_row, ne0) - (int)src0_start_row;                                                   \
-            if (slice_size > 0) {                                                                                                 \
-                dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr, src2_ptr),                                                  \
-                               slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);            \
-                dma_queue_pop_nowait(dma_queue);                                                                                  \
-            }                                                                                                                     \
-        }                                                                                                                         \
-        for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                                \
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned,                            \
-                           src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a);          \
-        }                                                                                                                         \
-    }                                                                                                                             \
-                                                                                                                                  \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                            \
-                                                                                                                                  \
-    if (src0_start_row >= src0_end_row) {                                                                                         \
-        return;                                                                                                                   \
-    }                                                                                                                             \
-                                                                                                                                  \
-    for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                             \
-        const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;                                                                    \
-                                                                                                                                  \
-        float * dst_ptr = &tmp[ct * 32 - src0_start_row];                                                                         \
-        int valid_rows = (int)ne0 - (int)(ct * 32);                                                                               \
-        valid_rows = MIN(32, MAX(0, valid_rows));                                                                                 \
-                                                                                                                                  \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                    \
-        DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL);                                                               \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                     \
-                                                                                                                                  \
-        if (push_ct < ct_end) {                                                                                                   \
-            dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),                      \
-                           aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                                 \
-            push_ct++;                                                                                                            \
-        }                                                                                                                         \
-    }                                                                                                                             \
-                                                                                                                                  \
-    int copy_cnt = (int)MIN(src0_end_row, ne0) - (int)src0_start_row;                                                             \
-    if (copy_cnt > 0) {                                                                                                           \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct_end);                                                                \
-        if (src2) {                                                                                                               \
-            hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row],                                                                 \
-                            (const uint8_t *) tmp,                                                                                \
-                            (const uint8_t *) ((const float *) mmctx->vtcm_src2 + src0_start_row),                                \
-                            copy_cnt);                                                                                            \
-        } else {                                                                                                                  \
-            hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt);                                     \
-        }                                                                                                                         \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct_end);                                                                 \
-    }                                                                                                                             \
+#define MATVEC_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X1)                                                              \
+static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                               \
+    htp_matmul_preamble;                                                                                                 \
+                                                                                                                         \
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                             \
+                                                                                                                         \
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                \
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                   \
+                                                                                                                         \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                               \
+                                                                                                                         \
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;             \
+    const uint32_t n_prefetch = kparams->n_prefetch;                                                                     \
+    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);                \
+                                                                                                                         \
+    const size_t dst_row_size  = nb1;                                                                                    \
+    const size_t src1_row_size = nb11;                                                                                   \
+    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                  \
+                                                                                                                         \
+    uint8_t * vtcm_dst_ptr  = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith;                                   \
+    uint8_t * vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                                 \
+    uint8_t * src1_data = mmctx->vtcm_src1;                                                                              \
+                                                                                                                         \
+    float * tmp = (float *) vtcm_dst_ptr;                                                                                \
+                                                                                                                         \
+    const dma_addr_t src0_row = src0->data;                                                                              \
+                                                                                                                         \
+    const uint8_t * restrict src1_col = (const uint8_t *) src1_data;                                                     \
+    float * restrict dst_col          = (float *) dst->data;                                                             \
+                                                                                                                         \
+    const uint32_t tile_size = TILE_SIZE;                                                                                \
+    const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                     \
+                                                                                                                         \
+    uint32_t n_k_tiles_w = ne00 / 32;                                                                                    \
+    uint32_t n_k_tiles_a = ne10 / 32;                                                                                    \
+    uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                  \
+    uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                           \
+                                                                                                                         \
+    uint32_t ct_start = src0_start_row / 32;                                                                             \
+    uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                        \
+                                                                                                                         \
+    uint32_t push_ct = ct_start;                                                                                         \
+    if (src0_start_row < src0_end_row) {                                                                                 \
+        if (src2) {                                                                                                      \
+            float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row;                                         \
+            const dma_addr_t src2_addr = src2->data + src0_start_row * sizeof(float);                                    \
+            int slice_size = (int)MIN(src0_end_row, ne0) - (int)src0_start_row;                                          \
+            if (slice_size > 0) {                                                                                        \
+                dma_queue_push(dma_q, dma_make_data(vtcm_src2_ptr, src2_addr),                                           \
+                               slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);   \
+                dma_queue_pop_nowait(dma_q);                                                                             \
+            }                                                                                                            \
+        }                                                                                                                \
+        for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                       \
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned,                      \
+                           src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \
+        }                                                                                                                \
+    }                                                                                                                    \
+                                                                                                                         \
+    hvx_mm_run_quant_task(mmctx, ith);                                                                                   \
+                                                                                                                         \
+    if (src0_start_row >= src0_end_row) {                                                                                \
+        return;                                                                                                          \
+    }                                                                                                                    \
+                                                                                                                         \
+    for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                    \
+        const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;                                                      \
+                                                                                                                         \
+        float * dst_ptr = &tmp[ct * 32 - src0_start_row];                                                                \
+        int valid_rows = (int)ne0 - (int)(ct * 32);                                                                      \
+        valid_rows = MIN(32, MAX(0, valid_rows));                                                                        \
+                                                                                                                         \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                           \
+        DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL);                                                      \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                            \
+                                                                                                                         \
+        if (push_ct < ct_end) {                                                                                          \
+            dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),                           \
+                           aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                        \
+            push_ct++;                                                                                                   \
+        }                                                                                                                \
+    }                                                                                                                    \
+                                                                                                                         \
+    int copy_cnt = (int)MIN(src0_end_row, ne0) - (int)src0_start_row;                                                    \
+    if (copy_cnt > 0) {                                                                                                  \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct_end);                                                       \
+        if (src2) {                                                                                                      \
+            hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row],                                                        \
+                            (const uint8_t *) tmp,                                                                       \
+                            (const uint8_t *) ((const float *) mmctx->vtcm_src2 + src0_start_row),                       \
+                            copy_cnt);                                                                                   \
+        } else {                                                                                                         \
+            hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt);                            \
+        }                                                                                                                \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct_end);                                                        \
+    }                                                                                                                    \
 }

 #define MATMUL_NX_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1)                                                           \
@@ -555,19 +479,18 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
     uint32_t n_k_tiles_a = ne10 / 32;                                                                                             \
     uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                    \
                                                                                                                                   \
-    dma_queue * dma_queue = octx->ctx->dma[ith];                                                                                  \
-                                                                                                                                  \
     hvx_mm_run_quant_task(mmctx, ith);                                                                                            \
                                                                                                                                   \
     for (uint32_t widx = 0; widx < n_weights; widx++) {                                                                           \
         const struct htp_tensor * restrict src_w = octx->src[widx];                                                               \
         const struct htp_tensor * restrict dst   = octx->dsts[widx];                                                              \
         if (!src_w || !dst) continue;                                                                                             \
+        dma_queue * dma_q = octx->ctx->dma[ith];                                                                                  \
                                                                                                                                   \
         const uint32_t ne00 = src_w->ne[0];                                                                                       \
         const uint32_t ne01 = src_w->ne[1];                                                                                       \
         const size_t dst_row_size = dst->nb[1];                                                                                   \
-        const uint8_t * restrict src_w_row = (const uint8_t *) src_w->data;                                                       \
+        const dma_addr_t src_w_row = src_w->data;                                                                                 \
                                                                                                                                   \
         uint32_t n_k_tiles_w = ne00 / 32;                                                                                         \
         uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                       \
@@ -595,12 +518,12 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
                                                                                                                                   \
         uint32_t push_ct = ct_start;                                                                                              \
         for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                                \
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_weight_ptr + d * tile_row_transfer_size_aligned,                          \
+            dma_queue_push(dma_q, dma_make_data(vtcm_weight_ptr + d * tile_row_transfer_size_aligned,                             \
                            src_w_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a);         \
         }                                                                                                                         \
                                                                                                                                   \
         for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                         \
-            const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;                                                                \
+            const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;                                                           \
             int valid_rows = (int)ne01 - (int)(ct * 32);                                                                          \
             valid_rows = MIN(32, MAX(0, valid_rows));                                                                             \
                                                                                                                                   \
@@ -627,7 +550,7 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                 \
                                                                                                                                   \
             if (push_ct < ct_end) {                                                                                               \
-                dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src_w_row + push_ct * tile_row_stride),                 \
+                dma_queue_push(dma_q, dma_make_data(w_tile, src_w_row + push_ct * tile_row_stride),                               \
                                aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                             \
                 push_ct++;                                                                                                        \
             }                                                                                                                     \
@@ -642,52 +565,57 @@ MATMUL_2D_REPACKED_IMPL(q6_k,       896,  tiled_vec_dot_q6_k_32x2,  tiled_vec_do
 MATMUL_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
 MATMUL_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)

-MATMUL_2D_REPACKED_IMPL(q4_0_flat,  576,  flat_vec_dot_q4_0_32x2,   flat_vec_dot_q4_0_32x1)
-MATMUL_2D_REPACKED_IMPL(q4_1_flat,  640,  flat_vec_dot_q4_1_32x2,   flat_vec_dot_q4_1_32x1)
-MATMUL_2D_REPACKED_IMPL(q8_0_flat,  1088, flat_vec_dot_q8_0_32x2,   flat_vec_dot_q8_0_32x1)
-MATMUL_2D_REPACKED_IMPL(q6_k_flat,  896,  flat_vec_dot_q6_k_32x2,   flat_vec_dot_q6_k_32x1)
-MATMUL_2D_REPACKED_IMPL(iq4nl_flat, 576,  flat_vec_dot_iq4nl_32x2,  flat_vec_dot_iq4nl_32x1)
-MATMUL_2D_REPACKED_IMPL(mxfp4_flat, 544,  flat_vec_dot_mxfp4_32x2,  flat_vec_dot_mxfp4_32x1)
-
-#define QUANTIZE_IMPL(name, log_name, kernel_fn, dst_row_size_expr)                                        \
-static void name(unsigned int nth, unsigned int ith, void * data) {                                        \
-    struct htp_mm_context * mmctx = data;                                                                  \
-    struct htp_ops_context * octx = mmctx->octx;                                                           \
-    const struct htp_tensor * src = mmctx->act;                                                            \
-    const uint32_t ne0 = src->ne[0];                                                                       \
-    const uint32_t ne1 = src->ne[1];                                                                       \
-    const uint32_t ne2 = src->ne[2];                                                                       \
-    const uint32_t ne3 = src->ne[3];                                                                       \
-    const uint32_t nrows = ne1 * ne2 * ne3;                                                                \
-    const uint32_t nrows_per_thread = mmctx->n_quant_rows_per_thread;                                      \
-                                                                                                           \
-    const uint32_t ir_first = nrows_per_thread * ith;                                                      \
-    if (ir_first >= nrows) {                                                                               \
-        return;                                                                                            \
-    }                                                                                                      \
-                                                                                                           \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                 \
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                        \
-                                                                                                           \
-    uint8_t * restrict dst = mmctx->vtcm_src1;                                                             \
-    const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows);                                      \
-    const size_t src_row_size = src->nb[1];                                                                \
-    const size_t dst_row_size = (dst_row_size_expr);                                                       \
-    const uint8_t * restrict src_data = (const uint8_t *) src->data + (src_row_size * ir_first);           \
-    uint8_t * restrict dst_data = (uint8_t *) dst + (dst_row_size * ir_first);                             \
-    uint8_t * restrict tmp_data = (uint8_t *) mmctx->vtcm_dst + (mmctx->vtcm_dst_size_per_thread * ith);   \
-    kernel_fn(src_data, dst_data, tmp_data, ne0, ir_last - ir_first, src_row_size, dst_row_size);          \
-                                                                                                           \
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                         \
+#define QUANTIZE_IMPL(name, log_name, kernel_fn, dst_row_size_expr)                                                                                               \
+static void name(unsigned int nth, unsigned int ith, void * data) {                                                                                               \
+    struct htp_mm_context * mmctx = data;                                                                                                                         \
+    struct htp_ops_context * octx = mmctx->octx;                                                                                                                  \
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                                                      \
+    const struct htp_tensor * src = mmctx->act;                                                                                                                   \
+    const uint32_t ne0 = src->ne[0];                                                                                                                              \
+    const uint32_t nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->src1_nrows;                                                                             \
+    const uint32_t nrows_per_thread = mmctx->n_quant_rows_per_thread;                                                                                             \
+                                                                                                                                                                  \
+    const uint32_t ir_first = nrows_per_thread * ith;                                                                                                             \
+    if (ir_first >= nrows) {                                                                                                                                      \
+        return;                                                                                                                                                   \
+    }                                                                                                                                                             \
+                                                                                                                                                                  \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                                                        \
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                                                                               \
+                                                                                                                                                                  \
+    uint8_t * restrict dst = mmctx->vtcm_src1;                                                                                                                    \
+    const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows);                                                                                             \
+    const size_t src_row_size = src->nb[1];                                                                                                                       \
+    const size_t dst_row_size = (dst_row_size_expr);                                                                                                              \
+    uint8_t * restrict tmp_data = (uint8_t *) mmctx->vtcm_dst + (mmctx->vtcm_dst_size_per_thread * ith);                                                          \
+                                                                                                                                                                  \
+    const bool is_contiguous = (src->nb[2] == src->ne[1] * src->nb[1]) && (src->nb[3] == src->ne[2] * src->nb[2]);                                                \
+    if (is_contiguous) {                                                                                                                                          \
+        const uint8_t * restrict src_data = (const uint8_t *) src->data + (src_row_size * (mmctx->cur_m_start + ir_first));                                       \
+        uint8_t * restrict dst_data = (uint8_t *) dst + (dst_row_size * ir_first);                                                                                \
+        kernel_fn(src_data, dst_data, tmp_data, ne0, ir_last - ir_first, src_row_size, dst_row_size);                                                             \
+    } else {                                                                                                                                                      \
+        const uint32_t ne12_ne1 = src->ne[2] * src->ne[1];                                                                                                        \
+        for (uint32_t ir = ir_first; ir < ir_last; ++ir) {                                                                                                        \
+            const uint32_t ir1 = mmctx->cur_m_start + ir;                                                                                                         \
+            const uint32_t i13 = fastdiv(ir1, &kparams->div_ne12_ne1);                                                                                            \
+            const uint32_t rem = ir1 - i13 * ne12_ne1;                                                                                                            \
+            const uint32_t i12 = fastdiv(rem, &kparams->div_ne1);                                                                                                 \
+            const uint32_t i11 = rem - i12 * src->ne[1];                                                                                                          \
+            const uint8_t * restrict row_src = (const uint8_t *) src->data + ((size_t) i11 * src->nb[1] + (size_t) i12 * src->nb[2] + (size_t) i13 * src->nb[3]); \
+            uint8_t * restrict row_dst = dst + (dst_row_size * ir);                                                                                               \
+            kernel_fn(row_src, row_dst, tmp_data, ne0, 1, src_row_size, dst_row_size);                                                                            \
+        }                                                                                                                                                         \
+    }                                                                                                                                                             \
+                                                                                                                                                                  \
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                                                                                \
 }

 QUANTIZE_IMPL(quantize_f32_q8_0_tiled, "quantize-f32-q8_0_tiled", quantize_f32_q8_0_tiled_kernel, htp_mm_q8_0_tiled_row_size(ne0))
 QUANTIZE_IMPL(quantize_f32_q8_1_tiled, "quantize-f32-q8_1_tiled", quantize_f32_q8_1_tiled_kernel, htp_mm_q8_1_tiled_row_size(ne0))
-QUANTIZE_IMPL(quantize_f32_q8_0_flat,  "quantize-f32-q8_0_flat",  quantize_f32_q8_0_flat_kernel,  htp_mm_q8_0_flat_row_size(ne0))
-QUANTIZE_IMPL(quantize_f32_q8_1_flat,  "quantize-f32-q8_1_flat",  quantize_f32_q8_1_flat_kernel,  htp_mm_q8_1_flat_row_size(ne0))
-QUANTIZE_IMPL(quantize_f32_f32_flat,   "quantize-f32-f32",        quantize_f32_f32_flat_kernel,   mmctx->vtcm_src1_stride)
-QUANTIZE_IMPL(quantize_f32_f16_flat,   "quantize-f32-f16",        quantize_f32_f16_flat_kernel,   mmctx->vtcm_src1_stride)
-QUANTIZE_IMPL(quantize_f16_f16_flat,   "quantize-f16-f16",        quantize_f16_f16_flat_kernel,   mmctx->vtcm_src1_stride)
+QUANTIZE_IMPL(quantize_f32_f32,       "quantize-f32-f32",       quantize_f32_f32_kernel,       mmctx->vtcm_src1_stride)
+QUANTIZE_IMPL(quantize_f32_f16,       "quantize-f32-f16",       quantize_f32_f16_kernel,       mmctx->vtcm_src1_stride)
+QUANTIZE_IMPL(quantize_f16_f16,       "quantize-f16-f16",       quantize_f16_f16_kernel,       mmctx->vtcm_src1_stride)

 static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, void * data) {
     struct htp_mm_context * mmctx = data;
@@ -744,25 +672,146 @@ MATVEC_2D_REPACKED_IMPL(q6_k,       896,  tiled_vec_dot_q6_k_32x1)
 MATVEC_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x1)
 MATVEC_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x1)

-MATVEC_2D_REPACKED_IMPL(q4_0_flat,  576,  flat_vec_dot_q4_0_32x1)
-MATVEC_2D_REPACKED_IMPL(q4_1_flat,  640,  flat_vec_dot_q4_1_32x1)
-MATVEC_2D_REPACKED_IMPL(q8_0_flat,  1088, flat_vec_dot_q8_0_32x1)
-MATVEC_2D_REPACKED_IMPL(q6_k_flat,  896,  flat_vec_dot_q6_k_32x1)
-MATVEC_2D_REPACKED_IMPL(iq4nl_flat, 576,  flat_vec_dot_iq4nl_32x1)
-MATVEC_2D_REPACKED_IMPL(mxfp4_flat, 544,  flat_vec_dot_mxfp4_32x1)
-
-
 MATMUL_NX_2D_REPACKED_IMPL(q4_0,       576,  tiled_vec_dot_q4_0_32x2,  tiled_vec_dot_q4_0_32x1)
 MATMUL_NX_2D_REPACKED_IMPL(q4_1,       640,  tiled_vec_dot_q4_1_32x2,  tiled_vec_dot_q4_1_32x1)
 MATMUL_NX_2D_REPACKED_IMPL(q8_0,       1088, tiled_vec_dot_q8_0_32x2,  tiled_vec_dot_q8_0_32x1)
 MATMUL_NX_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
 MATMUL_NX_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)

-MATMUL_NX_2D_REPACKED_IMPL(q4_0_flat,  576,  flat_vec_dot_q4_0_32x2,   flat_vec_dot_q4_0_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(q4_1_flat,  640,  flat_vec_dot_q4_1_32x2,   flat_vec_dot_q4_1_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(q8_0_flat,  1088, flat_vec_dot_q8_0_32x2,   flat_vec_dot_q8_0_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(iq4nl_flat, 576,  flat_vec_dot_iq4nl_32x2,  flat_vec_dot_iq4nl_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(mxfp4_flat, 544,  flat_vec_dot_mxfp4_32x2,  flat_vec_dot_mxfp4_32x1)
+#define MATMUL_4D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1)                                                                                        \
+static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                                                  \
+    htp_matmul_preamble;                                                                                                                                    \
+                                                                                                                                                            \
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                                                \
+    const uint32_t cur_m_rows = mmctx->cur_m_rows ? mmctx->cur_m_rows : (ne11 * ne12 * ne13);                                                               \
+    const uint32_t cur_m_start = mmctx->cur_m_start;                                                                                                        \
+                                                                                                                                                            \
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                                                   \
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                                                      \
+                                                                                                                                                            \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                                                  \
+                                                                                                                                                            \
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                                                \
+    const uint32_t n_prefetch = kparams->n_prefetch;                                                                                                        \
+    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);                                                   \
+                                                                                                                                                            \
+    const size_t dst_row_size = nb1;                                                                                                                        \
+    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                                                     \
+                                                                                                                                                            \
+    uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                                                           \
+    uint8_t * restrict src1_data = mmctx->vtcm_src1;                                                                                                        \
+                                                                                                                                                            \
+    const uint32_t tile_size = TILE_SIZE;                                                                                                                   \
+    const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                                                        \
+                                                                                                                                                            \
+    const uint32_t n_k_tiles_w = ne00 / 32;                                                                                                                 \
+    const uint32_t n_k_tiles_a = ne10 / 32;                                                                                                                 \
+    const uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                                               \
+    const uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                                        \
+    const uint32_t src0_slice_stride = ((ne01 + 31) / 32) * tile_row_stride;                                                                                \
+                                                                                                                                                            \
+    const uint32_t ct_start = src0_start_row / 32;                                                                                                          \
+    const uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                                                     \
+                                                                                                                                                            \
+    hvx_mm_run_quant_task(mmctx, ith);                                                                                                                      \
+                                                                                                                                                            \
+    if (src0_start_row >= src0_end_row || cur_m_rows == 0) {                                                                                                \
+        return;                                                                                                                                             \
+    }                                                                                                                                                       \
+                                                                                                                                                            \
+    const uint32_t total_batches = ne12 * ne13;                                                                                                             \
+    const uint32_t b_start = fastdiv(cur_m_start, &kparams->div_ne1);                                                                                       \
+    uint32_t b_end = fastdiv(cur_m_start + cur_m_rows + ne11 - 1, &kparams->div_ne1);                                                                       \
+    b_end = MIN(b_end, total_batches);                                                                                                                      \
+                                                                                                                                                            \
+    uint32_t b_grp_start = b_start;                                                                                                                         \
+    while (b_grp_start < b_end) {                                                                                                                           \
+        const uint32_t b3 = fastdiv(b_grp_start, &kparams->div_ne12);                                                                                       \
+        const uint32_t b2 = b_grp_start - b3 * ne12;                                                                                                        \
+        const uint32_t i02 = fastdiv(b2, &kparams->div_r2);                                                                                                 \
+        const uint32_t i03 = fastdiv(b3, &kparams->div_r3);                                                                                                 \
+                                                                                                                                                            \
+        uint32_t b_grp_end = b_grp_start + 1;                                                                                                               \
+        while (b_grp_end < b_end) {                                                                                                                         \
+            const uint32_t cur_b3 = fastdiv(b_grp_end, &kparams->div_ne12);                                                                                 \
+            const uint32_t cur_b2 = b_grp_end - cur_b3 * ne12;                                                                                              \
+            const uint32_t cur_i02 = fastdiv(cur_b2, &kparams->div_r2);                                                                                     \
+            const uint32_t cur_i03 = fastdiv(cur_b3, &kparams->div_r3);                                                                                     \
+            if (cur_i02 != i02 || cur_i03 != i03) {                                                                                                         \
+                break;                                                                                                                                      \
+            }                                                                                                                                               \
+            b_grp_end++;                                                                                                                                    \
+        }                                                                                                                                                   \
+                                                                                                                                                            \
+        const uint32_t slice_idx = i03 * ne02 + i02;                                                                                                        \
+        const dma_addr_t src0_slice = src0->data + (size_t) slice_idx * src0_slice_stride;                                                                  \
+                                                                                                                                                            \
+        uint32_t push_ct = ct_start;                                                                                                                        \
+        for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {                                                                          \
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned,                                                         \
+                           src0_slice + (size_t) push_ct * tile_row_stride),                                                                                \
+                           aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                                                           \
+        }                                                                                                                                                   \
+                                                                                                                                                            \
+        for (uint32_t ct = ct_start; ct < ct_end; ct++) {                                                                                                   \
+            const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;                                                                                     \
+                                                                                                                                                            \
+            int valid_rows = (int)ne0 - (int)(ct * 32);                                                                                                     \
+            valid_rows = MIN(32, MAX(0, valid_rows));                                                                                                       \
+                                                                                                                                                            \
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                                          \
+            for (uint32_t b = b_grp_start; b < b_grp_end; b++) {                                                                                            \
+                const uint32_t b_m_start = b * ne11;                                                                                                        \
+                const uint32_t m_first = MAX(cur_m_start, b_m_start);                                                                                       \
+                const uint32_t m_last  = MIN(cur_m_start + cur_m_rows, b_m_start + ne11);                                                                   \
+                if (m_first >= m_last) continue;                                                                                                            \
+                                                                                                                                                            \
+                const uint32_t cur_b3 = fastdiv(b, &kparams->div_ne12);                                                                                     \
+                const uint32_t cur_b2 = b - cur_b3 * ne12;                                                                                                  \
+                uint8_t * dst_batch_base = (uint8_t *) dst->data + (size_t) cur_b2 * nb2 + (size_t) cur_b3 * nb3;                                           \
+                                                                                                                                                            \
+                const uint32_t chunk_m_offset = m_first - cur_m_start;                                                                                      \
+                const uint32_t dst_m_offset   = m_first - b_m_start;                                                                                        \
+                const uint32_t batch_nrows    = m_last - m_first;                                                                                           \
+                                                                                                                                                            \
+                uint32_t ir1 = 0;                                                                                                                           \
+                for (; ir1 + 1 < batch_nrows; ir1 += 2) {                                                                                                   \
+                    const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 0) * src1_stride);                          \
+                    const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 1) * src1_stride);                          \
+                    float * restrict dst_row0 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 0) * dst_row_size);                                       \
+                    float * restrict dst_row1 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 1) * dst_row_size);                                       \
+                    float * dst_ptr0 = &dst_row0[ct * 32];                                                                                                  \
+                    float * dst_ptr1 = &dst_row1[ct * 32];                                                                                                  \
+                    DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, NULL, NULL);                                                \
+                }                                                                                                                                           \
+                for (; ir1 < batch_nrows; ++ir1) {                                                                                                          \
+                    const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);                               \
+                    float * restrict dst_row = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size);                                            \
+                    float * dst_ptr = &dst_row[ct * 32];                                                                                                    \
+                    DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL);                                                                             \
+                }                                                                                                                                           \
+            }                                                                                                                                               \
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                                           \
+                                                                                                                                                            \
+            if (push_ct < ct_end) {                                                                                                                         \
+                dma_queue_push(dma_q, dma_make_data(w_tile, src0_slice + (size_t) push_ct * tile_row_stride),                                               \
+                               aligned_tile_size, tile_size, tile_size, n_k_tiles_a);                                                                       \
+                push_ct++;                                                                                                                                  \
+            }                                                                                                                                               \
+        }                                                                                                                                                   \
+        b_grp_start = b_grp_end;                                                                                                                            \
+    }                                                                                                                                                       \
+    if (src2) {                                                                                                                                             \
+        hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + cur_m_rows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1); \
+    }                                                                                                                                                       \
+}
+
+MATMUL_4D_REPACKED_IMPL(q4_0,       576,  tiled_vec_dot_q4_0_32x2,  tiled_vec_dot_q4_0_32x1)
+MATMUL_4D_REPACKED_IMPL(q4_1,       640,  tiled_vec_dot_q4_1_32x2,  tiled_vec_dot_q4_1_32x1)
+MATMUL_4D_REPACKED_IMPL(q8_0,       1088, tiled_vec_dot_q8_0_32x2,  tiled_vec_dot_q8_0_32x1)
+MATMUL_4D_REPACKED_IMPL(q6_k,       896,  tiled_vec_dot_q6_k_32x2,  tiled_vec_dot_q6_k_32x1)
+MATMUL_4D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
+MATMUL_4D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)

 static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
     htp_matmul_preamble;
@@ -773,7 +822,8 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t prefetch_mask = n_prefetch - 1;

     const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;  // src0 rows
-    const uint32_t src1_nrows = ne11 * ne12 * ne13;                          // src1 rows
+    const uint32_t src1_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->src1_nrows;                          // src1 rows
+    const uint32_t cur_m_start = mmctx->cur_m_start;

     const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
     const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
@@ -793,7 +843,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
     uint8_t * restrict src1_data     = mmctx->vtcm_src1;

-    const uint8_t * restrict src0_row = (const uint8_t *) src0->data;
+    const dma_addr_t src0_row = src0->data;

     // Prefill vtcm with src0 rows
     if (src0_start_row < src0_end_row) {
@@ -802,7 +852,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
             if (is0 >= (int)n_prefetch) {
                 break;
             }
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 2);
         }
     }
@@ -815,7 +865,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {

     // Process src0 rows
     for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
-        const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+        const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
         // Process src1 columns in pairs (2x2 tiling)
@@ -823,15 +873,15 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
         for (; ir1 + 1 < src1_nrows; ir1 += 2) {
             const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);
             const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride);
-            float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size));
-            float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size));
+            float * restrict dst_row0 = (float *) (dst->data + ((cur_m_start + ir1+0) * dst_row_size));
+            float * restrict dst_row1 = (float *) (dst->data + ((cur_m_start + ir1+1) * dst_row_size));
             mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
         }

         // Handle remaining src1 rows (fallback to 2x1)
         for (; ir1 < src1_nrows; ++ir1) {
             const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
-            float * restrict dst_row          = (float *) (dst->data + (ir1 * dst_row_size));
+            float * restrict dst_row          = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));
             mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -840,7 +890,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
         const int pr0 = (ir0 + n_prefetch);
         const int is0 = (pr0 - src0_start_row) & prefetch_mask;
         if (pr0 < src0_end_row_x2) {
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 2);
         }
     }
@@ -849,21 +899,21 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
     if (src0_end_row != src0_end_row_x2) {
         uint32_t  ir0 = src0_end_row_x2;
         const int is0 = (ir0 - src0_start_row) & prefetch_mask;
-        dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+        dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                        src0_stride, src0_row_size, src0_row_size, 1);
-        const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+        const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
         #pragma unroll(2)
         for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) {
             const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
-            float * restrict dst_row          = (float *) (dst->data + (ir1 * dst_row_size));
+            float * restrict dst_row          = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));
             mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
     }
     if (src2) {
-        hvx_tensor_add_f32_grid(dst, src2, 0, src1_nrows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
+        hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + src1_nrows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
     }
 }

@@ -891,7 +941,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {

     float * tmp = (float *) vtcm_dst_ptr;

-    const uint8_t * restrict src0_row = (const uint8_t *) src0->data;
+    const dma_addr_t src0_row = src0->data;
     const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
     float * restrict dst_col          = (float *) dst->data;

@@ -906,12 +956,12 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     if (src0_start_row < src0_end_row) {
         if (src2) {
             float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row;
-            const float * src2_ptr = (const float *) src2->data + src0_start_row;
+            const dma_addr_t src2_addr = src2->data + src0_start_row * sizeof(float);
             int slice_size = (int)src0_end_row - (int)src0_start_row;
             if (slice_size > 0) {
-                dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr, src2_ptr),
+                dma_queue_push(dma_q, dma_make_data(vtcm_src2_ptr, src2_addr),
                                slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);
-                dma_queue_pop_nowait(dma_queue);
+                dma_queue_pop_nowait(dma_q);
             }
         }
         for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
@@ -919,7 +969,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
             if (is0 >= n_prefetch) {
                 break;
             }
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 2);
         }
     }
@@ -932,7 +982,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {

     // Process src0 rows
     for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
-        const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+        const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
         mmctx->vec_dot_2x1(ne00, &tmp[ir0 - src0_start_row], ss0, ss0 + src0_stride, src1_col);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -941,7 +991,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
         const uint32_t pr0 = (ir0 + n_prefetch);
         const uint32_t is0 = (pr0 - src0_start_row) & prefetch_mask;
         if (pr0 < src0_end_row_x2) {
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 2);
         }
     }
@@ -950,9 +1000,9 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     if (src0_end_row != src0_end_row_x2) {
         const uint32_t ir0 = src0_end_row_x2;
         const uint32_t is0 = (ir0 - src0_start_row) & prefetch_mask;
-        dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+        dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                        src0_stride, src0_row_size, src0_row_size, 1);
-        const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+        const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
         mmctx->vec_dot_1x1(ne00, &tmp[ir0 - src0_start_row], ss0, src1_col);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -973,6 +1023,154 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     }
 }

+static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
+    htp_matmul_preamble;
+
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+    const uint32_t n_prefetch = kparams->n_prefetch;
+    assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
+    const uint32_t prefetch_mask = n_prefetch - 1;
+
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;
+    const uint32_t cur_m_rows = mmctx->cur_m_rows ? mmctx->cur_m_rows : (ne11 * ne12 * ne13);
+    const uint32_t cur_m_start = mmctx->cur_m_start;
+
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
+    const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);
+
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    const size_t dst_row_size  = nb1;
+    const size_t src0_row_size = nb01;
+    const size_t src0_stride = mmctx->vtcm_src0_stride;
+    const size_t src1_stride = mmctx->vtcm_src1_stride;
+
+    uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
+    uint8_t * restrict src1_data     = mmctx->vtcm_src1;
+
+    hvx_mm_run_quant_task(mmctx, ith);
+
+    if (src0_start_row >= src0_end_row || cur_m_rows == 0) {
+        return;
+    }
+
+    const uint32_t total_batches = ne12 * ne13;
+    const uint32_t b_start = fastdiv(cur_m_start, &kparams->div_ne1);
+    uint32_t b_end = fastdiv(cur_m_start + cur_m_rows + ne11 - 1, &kparams->div_ne1);
+    b_end = MIN(b_end, total_batches);
+
+    uint32_t b_grp_start = b_start;
+    while (b_grp_start < b_end) {
+        const uint32_t b3 = fastdiv(b_grp_start, &kparams->div_ne12);
+        const uint32_t b2 = b_grp_start - b3 * ne12;
+        const uint32_t i02 = fastdiv(b2, &kparams->div_r2);
+        const uint32_t i03 = fastdiv(b3, &kparams->div_r3);
+
+        uint32_t b_grp_end = b_grp_start + 1;
+        while (b_grp_end < b_end) {
+            const uint32_t cur_b3 = fastdiv(b_grp_end, &kparams->div_ne12);
+            const uint32_t cur_b2 = b_grp_end - cur_b3 * ne12;
+            const uint32_t cur_i02 = fastdiv(cur_b2, &kparams->div_r2);
+            const uint32_t cur_i03 = fastdiv(cur_b3, &kparams->div_r3);
+            if (cur_i02 != i02 || cur_i03 != i03) {
+                break;
+            }
+            b_grp_end++;
+        }
+
+        const dma_addr_t src0_row = src0->data + ((size_t) i02 * nb02 + (size_t) i03 * nb03);
+
+        for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
+            const int is0 = (ir0 - src0_start_row);
+            if (is0 >= (int)n_prefetch) {
+                break;
+            }
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + (size_t) ir0 * src0_row_size),
+                           src0_stride, src0_row_size, src0_row_size, 2);
+        }
+
+        for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
+            const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
+
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
+            for (uint32_t b = b_grp_start; b < b_grp_end; b++) {
+                const uint32_t b_m_start = b * ne11;
+                const uint32_t m_first = MAX(cur_m_start, b_m_start);
+                const uint32_t m_last  = MIN(cur_m_start + cur_m_rows, b_m_start + ne11);
+                if (m_first >= m_last) continue;
+
+                const uint32_t cur_b3 = fastdiv(b, &kparams->div_ne12);
+                const uint32_t cur_b2 = b - cur_b3 * ne12;
+                uint8_t * dst_batch_base = (uint8_t *) dst->data + (size_t) cur_b2 * nb2 + (size_t) cur_b3 * nb3;
+
+                const uint32_t chunk_m_offset = m_first - cur_m_start;
+                const uint32_t dst_m_offset   = m_first - b_m_start;
+                const uint32_t batch_nrows    = m_last - m_first;
+
+                uint32_t ir1 = 0;
+                for (; ir1 + 1 < batch_nrows; ir1 += 2) {
+                    const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 0) * src1_stride);
+                    const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 1) * src1_stride);
+                    float * restrict dst_row0 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 0) * dst_row_size);
+                    float * restrict dst_row1 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 1) * dst_row_size);
+                    mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+                }
+                for (; ir1 < batch_nrows; ++ir1) {
+                    const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);
+                    float * restrict dst_row          = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size);
+                    mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
+                }
+            }
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
+
+            const int pr0 = (ir0 + n_prefetch);
+            const int is0 = (pr0 - src0_start_row) & prefetch_mask;
+            if (pr0 < src0_end_row_x2) {
+                dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + (size_t) pr0 * src0_row_size),
+                               src0_stride, src0_row_size, src0_row_size, 2);
+            }
+        }
+
+        if (src0_end_row != src0_end_row_x2) {
+            uint32_t ir0 = src0_end_row_x2;
+            const int is0 = (ir0 - src0_start_row) & prefetch_mask;
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + (size_t) ir0 * src0_row_size),
+                           src0_stride, src0_row_size, src0_row_size, 1);
+            const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
+
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
+            for (uint32_t b = b_grp_start; b < b_grp_end; b++) {
+                const uint32_t b_m_start = b * ne11;
+                const uint32_t m_first = MAX(cur_m_start, b_m_start);
+                const uint32_t m_last  = MIN(cur_m_start + cur_m_rows, b_m_start + ne11);
+                if (m_first >= m_last) continue;
+
+                const uint32_t cur_b3 = fastdiv(b, &kparams->div_ne12);
+                const uint32_t cur_b2 = b - cur_b3 * ne12;
+                uint8_t * dst_batch_base = (uint8_t *) dst->data + (size_t) cur_b2 * nb2 + (size_t) cur_b3 * nb3;
+
+                const uint32_t chunk_m_offset = m_first - cur_m_start;
+                const uint32_t dst_m_offset   = m_first - b_m_start;
+                const uint32_t batch_nrows    = m_last - m_first;
+
+                for (uint32_t ir1 = 0; ir1 < batch_nrows; ++ir1) {
+                    const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);
+                    float * restrict dst_row          = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size);
+                    mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
+                }
+            }
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
+        }
+
+        b_grp_start = b_grp_end;
+    }
+
+    if (src2) {
+        hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + cur_m_rows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
+    }
+}
+
 #define MMID_MATRIX_ROW(row_id, i1) matrix_rows[(row_id) * mmctx->mapping_stride + (i1)]

 static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
@@ -1018,7 +1216,7 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
             continue;
         }

-        const uint8_t * src0_row = (const uint8_t *) src0->data + cur_a * nb02;
+        const dma_addr_t src0_row = src0->data + cur_a * nb02;

         const uint32_t tile_size = htp_mm_get_weight_tile_size(src0->type);
         const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(src0->type);
@@ -1032,12 +1230,12 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {

         uint32_t push_ct = ct_start;
         for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
                            aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
         }

         for (uint32_t ct = ct_start; ct < ct_end; ct++) {
-            const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;
+            const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;

             int valid_rows = (int)ne01 - (int)(ct * 32);
             valid_rows = MIN(32, MAX(0, valid_rows));
@@ -1057,7 +1255,7 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

             if (push_ct < ct_end) {
-                dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),
+                dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),
                                aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
                 push_ct++;
             }
@@ -1105,7 +1303,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
         }
         assert(eid < (int32_t) n_ids);

-        const uint8_t * restrict src0_row = (const uint8_t *) src0->data + eid * nb02;
+        const dma_addr_t src0_row = src0->data + eid * nb02;
         const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
         float * restrict dst_row          = (float *) (dst->data + ie1 * nb1);

@@ -1121,12 +1319,12 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {

         uint32_t push_ct = ct_start;
         for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
                            aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
         }

         for (uint32_t ct = ct_start; ct < ct_end; ct++) {
-            const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;
+            const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;

             int valid_rows = (int)ne01 - (int)(ct * 32);
             valid_rows = MIN(32, MAX(0, valid_rows));
@@ -1136,7 +1334,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

             if (push_ct < ct_end) {
-                dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),
+                dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),
                                aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
                 push_ct++;
             }
@@ -1147,7 +1345,6 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
 static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
     struct htp_mm_context * mmctx = (struct htp_mm_context *) data;
     struct htp_ops_context * octx = mmctx->octx;
-    dma_queue * dma_queue         = octx->ctx->dma[ith];
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
     const uint32_t n_weights      = kparams->n_weights;
     const struct htp_tensor * restrict src0 = octx->src[0];
@@ -1176,6 +1373,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const struct htp_tensor * restrict src_w = octx->src[p];
             const struct htp_tensor * restrict dst   = octx->dsts[p];
             if (!src_w || !dst) continue;
+            dma_queue * dma_q = octx->ctx->dma[ith];

             const uint32_t ne01 = src_w->ne[1];
             uint32_t start_row = 0;
@@ -1195,7 +1393,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, end_row);
             if (src0_start_row >= src0_end_row) continue;

-            const uint8_t * restrict src0_row = (const uint8_t *) src_w->data + eid * src_w->nb[2];
+            const dma_addr_t src0_row = src_w->data + eid * src_w->nb[2];
             const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
             float * restrict dst_row = (float *) (dst->data + ie1 * dst->nb[1]);

@@ -1211,12 +1409,12 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {

             uint32_t push_ct = ct_start;
             for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {
-                dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
+                dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
                                aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
             }

             for (uint32_t ct = ct_start; ct < ct_end; ct++) {
-                const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;
+                const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;

                 int valid_rows = (int)src_w->ne[1] - (int)(ct * 32);
                 valid_rows = MIN(32, MAX(0, valid_rows));
@@ -1226,7 +1424,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
                 htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

                 if (push_ct < ct_end) {
-                    dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),
+                    dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),
                                    aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
                     push_ct++;
                 }
@@ -1238,7 +1436,6 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
 static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
     struct htp_mm_context * mmctx = (struct htp_mm_context *) data;
     struct htp_ops_context * octx = mmctx->octx;
-    dma_queue * dma_queue         = octx->ctx->dma[ith];
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
     const uint32_t n_weights      = kparams->n_weights;
     const struct htp_tensor * restrict src0 = octx->src[0];
@@ -1270,6 +1467,7 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const struct htp_tensor * restrict src_w = octx->src[p];
             const struct htp_tensor * restrict dst   = octx->dsts[p];
             if (!src_w || !dst) continue;
+            dma_queue * dma_q = octx->ctx->dma[ith];

             const uint32_t ne01 = src_w->ne[1];
             uint32_t start_row = 0;
@@ -1289,7 +1487,7 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, end_row);
             if (src0_start_row >= src0_end_row) continue;

-            const uint8_t * src0_row = (const uint8_t *) src_w->data + cur_a * src_w->nb[2];
+            const dma_addr_t src0_row = src_w->data + cur_a * src_w->nb[2];

             const uint32_t tile_size = htp_mm_get_weight_tile_size(src_w->type);
             const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(src_w->type);
@@ -1303,12 +1501,12 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {

             uint32_t push_ct = ct_start;
             for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) {
-                dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
+                dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride),
                                aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
             }

             for (uint32_t ct = ct_start; ct < ct_end; ct++) {
-                const uint8_t * w_tile = dma_queue_pop(dma_queue).dst;
+                const uint8_t * w_tile = (void *) dma_queue_pop(dma_q).dst;

                 int valid_rows = (int)src_w->ne[1] - (int)(ct * 32);
                 valid_rows = MIN(32, MAX(0, valid_rows));
@@ -1328,7 +1526,7 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
                 htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

                 if (push_ct < ct_end) {
-                    dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride),
+                    dma_queue_push(dma_q, dma_make_data(w_tile, src0_row + push_ct * tile_row_stride),
                                    aligned_tile_size, tile_size, tile_size, n_k_tiles_a);
                     push_ct++;
                 }
@@ -1384,6 +1582,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     const uint32_t src0_nrows = ne01;
     const uint32_t src1_nrows = ne11 * ne12 * ne13;
+    mmctx->src1_nrows = src1_nrows;

     uint32_t src0_row_start = 0;
     uint32_t src0_row_end   = src0_nrows;
@@ -1426,7 +1625,23 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
     worker_callback_t quant_task_func;
     worker_callback_t matmul_job_func;
     uint32_t n_quant_tasks = 1;
-    if (src1_nrows > 1) {
+    const bool is_batched = (ne12 > 1 || ne13 > 1 || ne02 > 1 || ne03 > 1);
+    if (is_batched) {
+        if (is_repacked) {
+            switch (src0->type) {
+                case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_4d_repacked_q4_0;   break;
+                case HTP_TYPE_Q4_1:
+                case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mm_4d_repacked_q4_1;   break;
+                case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_4d_repacked_q8_0;   break;
+                case HTP_TYPE_Q6_K:   matmul_job_func = hvx_mm_4d_repacked_q6_k;   break;
+                case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_4d_repacked_iq4nl;  break;
+                case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_4d_repacked_mxfp4;  break;
+                default:              return HTP_STATUS_NO_SUPPORT;
+            }
+        } else {
+            matmul_job_func = hvx_mm_4d;
+        }
+    } else if (src1_nrows > 1) {
         if (is_repacked) {
             switch (src0->type) {
                 case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_2d_repacked_q4_0;   break;
@@ -1462,7 +1677,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     switch (kparams->kernel_type) {
         case HTP_MM_KERNEL_HVX_F16_F16_VTCM:
-            quant_task_func        = (src1->type == HTP_TYPE_F32) ? quantize_f32_f16_flat : quantize_f16_f16_flat;
+            quant_task_func        = (src1->type == HTP_TYPE_F32) ? quantize_f32_f16 : quantize_f16_f16;
             mmctx->type            = "f16-f16";
             mmctx->vec_dot_1x1     = vec_dot_f16_f16_aa_1x1;
             mmctx->vec_dot_2x1     = vec_dot_f16_f16_aa_2x1;
@@ -1470,34 +1685,8 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             src1_row_size          = hex_round_up(ne10 * 2, 128);
             break;

-        case HTP_MM_KERNEL_HVX_F16_F32_DDR:
-            mmctx->type            = "f16-f32";
-            mmctx->vec_dot_1x1     = vec_dot_f16_f32_uu_1x1;
-            matmul_job_func        = hvx_mm_4d;
-            mmctx->mm_div_ne12_ne1 = kparams->div_ne12_ne1;
-            mmctx->mm_div_ne1      = kparams->div_ne1;
-            mmctx->mm_div_r2       = kparams->div_r2;
-            mmctx->mm_div_r3       = kparams->div_r3;
-            need_quant             = false;
-            quant_task_func        = NULL;
-            src1_row_size          = nb11;
-            break;
-
-        case HTP_MM_KERNEL_HVX_F16_F16_DDR:
-            mmctx->type            = "f16-f16";
-            mmctx->vec_dot_1x1     = vec_dot_f16_f16_uu_1x1;
-            matmul_job_func        = hvx_mm_4d;
-            mmctx->mm_div_ne12_ne1 = kparams->div_ne12_ne1;
-            mmctx->mm_div_ne1      = kparams->div_ne1;
-            mmctx->mm_div_r2       = kparams->div_r2;
-            mmctx->mm_div_r3       = kparams->div_r3;
-            src1_row_size          = nb11;
-            need_quant             = false;
-            quant_task_func        = NULL;
-            break;
-
         case HTP_MM_KERNEL_HVX_F32_F32_VTCM:
-            quant_task_func        = quantize_f32_f32_flat;
+            quant_task_func        = quantize_f32_f32;
             mmctx->type            = "f32-f32";
             mmctx->vec_dot_1x1     = vec_dot_f32_f32_aa_1x1;
             mmctx->vec_dot_2x1     = vec_dot_f32_f32_aa_2x1;
@@ -1505,50 +1694,6 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             src1_row_size          = hex_round_up(ne10 * 4, 128);
             break;

-        case HTP_MM_KERNEL_HVX_F32_F32_DDR:
-            quant_task_func        = NULL;
-            mmctx->type            = "f32-f32";
-            mmctx->vec_dot_1x1     = vec_dot_f32_f32_uu_1x1;
-            mmctx->mm_div_ne12_ne1 = kparams->div_ne12_ne1;
-            mmctx->mm_div_ne1      = kparams->div_ne1;
-            mmctx->mm_div_r2       = kparams->div_r2;
-            mmctx->mm_div_r3       = kparams->div_r3;
-            src1_row_size          = nb11;
-            need_quant             = false;
-            matmul_job_func        = hvx_mm_4d;
-            break;
-
-        case HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT: {
-            n_quant_tasks = MIN(src1_nrows, octx->n_threads);
-            quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_flat : quantize_f32_q8_0_flat;
-            src1_row_size = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
-
-            if (src1_nrows > 1) {
-                switch (src0->type) {
-                    case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_2d_repacked_q4_0_flat;   break;
-                    case HTP_TYPE_Q4_1:
-                    case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mm_2d_repacked_q4_1_flat;   break;
-                    case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_2d_repacked_q8_0_flat;   break;
-                    case HTP_TYPE_Q6_K:   matmul_job_func = hvx_mm_2d_repacked_q6_k_flat;   break;
-                    case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_2d_repacked_iq4nl_flat;  break;
-                    case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_2d_repacked_mxfp4_flat;  break;
-                    default:              return HTP_STATUS_NO_SUPPORT;
-                }
-            } else {
-                switch (src0->type) {
-                    case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mv_2d_repacked_q4_0_flat;   break;
-                    case HTP_TYPE_Q4_1:
-                    case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mv_2d_repacked_q4_1_flat;   break;
-                    case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mv_2d_repacked_q8_0_flat;   break;
-                    case HTP_TYPE_Q6_K:   matmul_job_func = hvx_mv_2d_repacked_q6_k_flat;   break;
-                    case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mv_2d_repacked_iq4nl_flat;  break;
-                    case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mv_2d_repacked_mxfp4_flat;  break;
-                    default:              return HTP_STATUS_NO_SUPPORT;
-                }
-            }
-            break;
-        }
-
         case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
         case HTP_MM_KERNEL_HVX_QUANT_ROW:
         default:
@@ -1560,7 +1705,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             const uint32_t nb = (ne10 + qk - 1) / qk;
             const uint32_t total_nb = src1_nrows * nb;

-            if (src1_nrows < octx->n_threads) {
+            if (src1_nrows < octx->n_threads && !is_batched) {
                 n_quant_tasks = MIN(total_nb, octx->n_threads);
                 quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
                 for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -1579,8 +1724,12 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             break;
     }

+    const uint32_t m_chunk = (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows)
+                           ? (uint32_t) kparams->m_chunk : src1_nrows;
+    const uint32_t m_layout_rows = m_chunk;
+
     struct htp_mm_hvx_vtcm_layout L;
-    htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads,
+    htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, m_layout_rows, octx->n_threads,
                                  dst_row_size, src0_row_size, src1_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false);

     if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM ||
@@ -1623,19 +1772,46 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
     mmctx->vtcm_src0_stride = src0_row_size_padded;
     mmctx->vtcm_src1_stride = src1_row_size;

-    if (need_quant) {
-        mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
-        mmctx->quant_task_func = quant_task_func;
-        mmctx->n_quant_tasks = n_quant_tasks;
-        atomic_init(&mmctx->quant_barrier, n_quant_tasks);
+    if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows) {
+        atomic_init(&mmctx->quant_barrier, 0);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
+
+        for (uint32_t m_start = 0; m_start < src1_nrows; m_start += m_chunk) {
+            const uint32_t cur_m_rows = MIN(src1_nrows - m_start, m_chunk);
+            mmctx->cur_m_start = m_start;
+            mmctx->cur_m_rows  = cur_m_rows;
+
+            if (need_quant) {
+                const uint32_t quant_tasks = MIN(cur_m_rows, octx->n_threads);
+                mmctx->n_quant_rows_per_thread = (cur_m_rows + quant_tasks - 1) / quant_tasks;
+                mmctx->n_quant_tasks = quant_tasks;
+                atomic_store(&mmctx->quant_barrier, quant_tasks);
+                mmctx->quant_task_func = quant_task_func;
+            } else {
+                mmctx->quant_task_func = NULL;
+                mmctx->n_quant_tasks = 0;
+            }
+
+            worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, octx->n_threads);
+        }
     } else {
-        mmctx->quant_task_func = NULL;
-        mmctx->n_quant_tasks = 0;
-    }
+        mmctx->cur_m_start = 0;
+        mmctx->cur_m_rows  = src1_nrows;
+
+        if (need_quant) {
+            mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
+            mmctx->quant_task_func = quant_task_func;
+            mmctx->n_quant_tasks = n_quant_tasks;
+            atomic_init(&mmctx->quant_barrier, n_quant_tasks);
+        } else {
+            mmctx->quant_task_func = NULL;
+            mmctx->n_quant_tasks = 0;
+        }

-    htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, octx->n_threads);
+        worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, octx->n_threads);
+    }

     return HTP_STATUS_OK;
 }
@@ -1653,7 +1829,6 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
     uint8_t * restrict src1_data     = mmctx->vtcm_src1;

-    dma_queue * dma_queue = octx->ctx->dma[ith];
     const uint32_t n_prefetch = kparams->n_prefetch;
     assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
     const uint32_t prefetch_mask = n_prefetch - 1;
@@ -1666,6 +1841,7 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
         const struct htp_tensor * restrict src_w = octx->src[widx];
         const struct htp_tensor * restrict dst   = octx->dsts[widx];
         if (!src_w || !dst) continue;
+        dma_queue * dma_q = octx->ctx->dma[ith];

         const uint32_t ne00 = src_w->ne[0];
         const uint32_t ne01 = src_w->ne[1];
@@ -1691,17 +1867,17 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
         const size_t src0_row_size = src_w->nb[1];
         const size_t src0_stride   = hex_round_up(src0_row_size, 128);

-        const uint8_t * restrict src0_row = (const uint8_t *) src_w->data;
+        const dma_addr_t src0_row = src_w->data;

         for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
             const int is0 = (ir0 - src0_start_row);
             if (is0 >= (int)n_prefetch) break;
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 2);
         }

         for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
-            const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+            const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
             htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
             uint32_t ir1 = 0;
             for (; ir1 + 1 < src1_nrows; ir1 += 2) {
@@ -1721,7 +1897,7 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
             const int pr0 = (ir0 + n_prefetch);
             const int is0 = (pr0 - src0_start_row) & prefetch_mask;
             if (pr0 < src0_end_row_x2) {
-                dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
+                dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size),
                                src0_stride, src0_row_size, src0_row_size, 2);
             }
         }
@@ -1729,9 +1905,9 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
         if (src0_end_row != src0_end_row_x2) {
             uint32_t ir0 = src0_end_row_x2;
             const int is0 = (ir0 - src0_start_row) & prefetch_mask;
-            dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
+            dma_queue_push(dma_q, dma_make_data(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size),
                            src0_stride, src0_row_size, src0_row_size, 1);
-            const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
+            const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
             htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
             for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) {
                 const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
@@ -1868,7 +2044,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
     // Push step 0
     if (n_steps > 0 && n_rows > 0) {
         uint32_t nrows_to_fetch = hex_smin(n_rows, R);
-        dma_queue_push(dma_q, dma_make_ptr(thread_f32_act, src + c_first),
+        dma_queue_push(dma_q, dma_make_data(thread_f32_act, src + c_first),
                        c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
     }
     // Push step 1
@@ -1878,7 +2054,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
             const float *next_src = src + next_r * k_stride + c_first;
             float *next_buf = thread_f32_act + 1 * R * c_len;
-            dma_queue_push(dma_q, dma_make_ptr(next_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(next_buf, next_src),
                            c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
         }
     }
@@ -1909,7 +2085,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
             const float *next_src = src + next_r * k_stride + c_first;
-            dma_queue_push(dma_q, dma_make_ptr(curr_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(curr_buf, next_src),
                            c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
         }
     }
@@ -2013,7 +2189,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
     // Push step 0
     if (n_steps > 0 && n_rows > 0) {
         uint32_t nrows_to_fetch = hex_smin(n_rows, R);
-        dma_queue_push(dma_q, dma_make_ptr(thread_f32_act, src),
+        dma_queue_push(dma_q, dma_make_data(thread_f32_act, src),
                        k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
     }
     // Push step 1 (if valid)
@@ -2023,7 +2199,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
             const float *next_src = src + next_r * k_stride;
             float *next_buf = thread_f32_act + 1 * R * k_block;
-            dma_queue_push(dma_q, dma_make_ptr(next_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(next_buf, next_src),
                            k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
         }
     }
@@ -2052,7 +2228,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
             const float *next_src = src + next_r * k_stride;
-            dma_queue_push(dma_q, dma_make_ptr(curr_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(curr_buf, next_src),
                            k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
         }
     }
@@ -2459,10 +2635,12 @@ static inline void hmx_matmul_job_init(hmx_matmul_job_t * job,
 }

 static int hmx_mm_2d_f32(struct htp_context *ctx,
+                                  dma_queue *weight_dma,
                                   float *restrict dst,
-                                  const float *restrict src2,
+                                  dma_addr_t src2_addr,
+                                  size_t src2_bytes,
                                   const float *activation,
-                                  const uint8_t *weight,
+                                  dma_addr_t weight,
                                   int m, int k, int n,
                                   int act_stride,
                                   int weight_stride,
@@ -2525,7 +2703,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
     const size_t qweight_row_stride = is_quant ? (size_t)(n_k_tiles * aligned_tile_size) / 32 : 0;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size, src2_bytes);

     vtcm_used = L.total_bytes;
     if (vtcm_used > vtcm_budget) {
@@ -2550,6 +2728,13 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,

     hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00));  // scale: 1.0, bias: 0.0 in FP16

+    const bool has_src2      = (src2_bytes > 0 && src2_addr != 0);
+    float   *vtcm_src2       = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_src2, has_src2);
+    if (has_src2) {
+        dma_queue_push(weight_dma, dma_make_data(vtcm_src2, src2_addr), hex_align_up(src2_bytes, 128), 0, src2_bytes, 1);
+        dma_queue_pop(weight_dma);
+    }
+
     FARF(HIGH, "hmx-mm-2d: m %d k %d n %d wtype %d mc %zu nc %zu vtcm %zu/%zu",
          m, k, n, weight_type, m_chunk_n_rows, n_chunk_n_cols, vtcm_used, vtcm_budget);

@@ -2586,13 +2771,13 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             // Prologue: push A0 and optionally A1 (if n_chunk_cnt > 1)
             const size_t   n_cols_A0 = hex_smin(n - 0 * n_chunk_n_cols, n_chunk_n_cols);
             const uint32_t height_A0 = is_quant ? (n_cols_A0 / 32) * n_k_tiles : n_cols_A0;
-            dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight),
+            dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight),
                            dma_dst_stride, dma_src_stride, dma_width_bytes, height_A0);

             if (1 < n_chunk_cnt) {
                 const size_t   n_cols_A1 = hex_smin(n - 1 * n_chunk_n_cols, n_chunk_n_cols);
                 const uint32_t height_A1 = is_quant ? (n_cols_A1 / 32) * n_k_tiles : n_cols_A1;
-                dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
+                dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
                                dma_dst_stride, dma_src_stride, dma_width_bytes, height_A1);
             }

@@ -2605,7 +2790,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 const size_t n_cols_p2 = hex_smin(n - nc_p2, n_chunk_n_cols);

                 // 1. pop A_i
-                void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+                void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

                 // 2. dequantize A_i
                 dequantize_tiled_weight_chunk_to_fp16_tiles(
@@ -2616,7 +2801,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 // 3. push A_{i+2} (if i+2 < n_chunk_cnt)
                 if (i + 2 < n_chunk_cnt) {
                     const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2;
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_p2 * weight_stride),
+                    dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_p2 * weight_stride),
                                    dma_dst_stride, dma_src_stride, dma_width_bytes, height_p2);
                 }

@@ -2633,7 +2818,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                     const size_t nc_prev = (i - 1) * n_chunk_n_cols;
                     const size_t n_cols_prev = hex_smin(n - nc_prev, n_chunk_n_cols);
                     float *output_chunk = dst + (mr * dst_stride + nc_prev);
-                    const float *src2_chunk = src2 ? (src2 + mr * src2_stride + nc_prev) : NULL;
+                    const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_prev) : NULL;
                     int chunk_dst_cols = dst_cols - (int)nc_prev;
                     if (chunk_dst_cols > 0) {
                         transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output_bufs[(i - 1) % 2], n_rows, n_cols_prev, dst_stride, src2_stride, chunk_dst_cols, n_threads);
@@ -2646,7 +2831,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             const size_t nc_last = (n_chunk_cnt - 1) * n_chunk_n_cols;
             const size_t n_cols_last = hex_smin(n - nc_last, n_chunk_n_cols);
             float *output_chunk = dst + (mr * dst_stride + nc_last);
-            const float *src2_chunk = src2 ? (src2 + mr * src2_stride + nc_last) : NULL;
+            const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_last) : NULL;
             int chunk_dst_cols = dst_cols - (int)nc_last;
             if (chunk_dst_cols > 0) {
                 transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output_bufs[(n_chunk_cnt - 1) % 2], n_rows, n_cols_last, dst_stride, src2_stride, chunk_dst_cols, n_threads);
@@ -2678,7 +2863,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             if (n > 0) {
                 const size_t n_cols = hex_smin(n, n_chunk_n_cols);
                 const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
-                dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
+                dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
             }

             for (size_t nc = 0; nc < n; nc += n_chunk_n_cols) {
@@ -2687,7 +2872,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 const size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS);

                 // A: Wait for weight DMA
-                void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+                void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

                 // B: Weight Dequantize (Threaded)
                 dequantize_tiled_weight_chunk_to_fp16_tiles(
@@ -2700,7 +2885,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 if (nc_next < n) {
                     const size_t n_cols_next = hex_smin(n - nc_next, n_chunk_n_cols);
                     const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
+                    dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
                 }

                 // C: HMX Compute (Queue-based)
@@ -2710,7 +2895,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,

                 // D: Output Store
                 float *output_chunk = dst + (mr * dst_stride + nc);
-                const float *src2_chunk = src2 ? (src2 + mr * src2_stride + nc) : NULL;
+                const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc) : NULL;
                 int chunk_dst_cols = dst_cols - (int)nc;
                 if (chunk_dst_cols > 0) {
                     transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output, n_rows, n_cols, dst_stride, src2_stride, chunk_dst_cols, n_threads);
@@ -2785,7 +2970,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     const uint32_t dma_width_bytes = is_quant ? tile_size : row_stride;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size, 0);

     if (L.total_bytes > vtcm_budget) {
         FARF(ERROR, "hmx-mm-nx-2d: VTCM overflow: used %zu budget %zu, m %d k %d mc %d nc %d",
@@ -2859,7 +3044,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 const struct htp_tensor * restrict dst   = octx->dsts[p];
                 if (!src_w || !dst) continue;

-                const uint8_t * weight     = (const uint8_t *) src_w->data;
+                const dma_addr_t weight       = src_w->data;
+                dma_queue * weight_dma      = octx->ctx->dma[0];
                 float * dst_ptr            = (float *) dst->data;
                 const size_t n             = src_w->ne[1];
                 if (n == 0) continue;
@@ -2872,13 +3058,13 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k

                 const size_t   n_cols_A0 = hex_smin(n - 0 * n_chunk_n_cols, n_chunk_n_cols);
                 const uint32_t height_A0 = is_quant ? (n_cols_A0 / 32) * n_k_tiles : n_cols_A0;
-                dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight),
+                dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight),
                                dma_dst_stride, dma_src_stride, dma_width_bytes, height_A0);

                 if (1 < n_chunk_cnt) {
                     const size_t   n_cols_A1 = hex_smin(n - 1 * n_chunk_n_cols, n_chunk_n_cols);
                     const uint32_t height_A1 = is_quant ? (n_cols_A1 / 32) * n_k_tiles : n_cols_A1;
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
+                    dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
                                    dma_dst_stride, dma_src_stride, dma_width_bytes, height_A1);
                 }

@@ -2889,7 +3075,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                     const size_t n_cols    = hex_smin(n - nc, n_chunk_n_cols);
                     const size_t n_cols_p2 = hex_smin(n - nc_p2, n_chunk_n_cols);

-                    void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+                    void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

                     dequantize_tiled_weight_chunk_to_fp16_tiles(
                         ctx, vtcm_weight_bufs[i % 2], curr_raw,
@@ -2898,7 +3084,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k

                     if (i + 2 < n_chunk_cnt) {
                         const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2;
-                        dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_p2 * weight_stride),
+                        dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_p2 * weight_stride),
                                        dma_dst_stride, dma_src_stride, dma_width_bytes, height_p2);
                     }

@@ -2956,7 +3142,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 const struct htp_tensor * restrict dst   = octx->dsts[p];
                 if (!src_w || !dst) continue;

-                const uint8_t * weight     = (const uint8_t *) src_w->data;
+                const dma_addr_t weight       = src_w->data;
+                dma_queue * weight_dma      = octx->ctx->dma[0];
                 float * dst_ptr            = (float *) dst->data;
                 const size_t n             = src_w->ne[1];
                 if (n == 0) continue;
@@ -2969,7 +3156,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 if (n > 0) {
                     const size_t n_cols = hex_smin(n, n_chunk_n_cols);
                     const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
+                    dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
                 }

                 for (size_t nc = 0; nc < n; nc += n_chunk_n_cols) {
@@ -2977,7 +3164,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                     const size_t n_row_tiles = hmx_ceil_div(n_rows, HTP_MM_HMX_TILE_N_ROWS);
                     const size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS);

-                    void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+                    void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

                     dequantize_tiled_weight_chunk_to_fp16_tiles(
                         ctx, vtcm_scratch0, curr_raw,
@@ -2988,7 +3175,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                     if (nc_next < n) {
                         const size_t n_cols_next = hex_smin(n - nc_next, n_chunk_n_cols);
                         const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
-                        dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
+                        dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
                     }

                     hmx_matmul_job_init(&job, vtcm_output, vtcm_f16_act, vtcm_scratch0, vtcm_scales, n_row_tiles, n_col_tiles, k / HTP_MM_HMX_TILE_N_ROWS);
@@ -3008,13 +3195,11 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     return HTP_STATUS_OK;
 }

-static inline const __fp16 *hmx_mm_weight_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params,
-                                                        int dst_b2, int dst_b3) {
+static inline dma_addr_t hmx_mm_weight_batch_data(const hmx_mm_f16_f32_batched_params_t *params,
+                                                         int dst_b2, int dst_b3) {
     const size_t b2_idx = (params->r2 <= 1) ? (size_t) dst_b2 : (size_t) fastdiv((uint32_t) dst_b2, &params->div_r2);
     const size_t b3_idx = (params->r3 <= 1) ? (size_t) dst_b3 : (size_t) fastdiv((uint32_t) dst_b3, &params->div_r3);
-    return (const __fp16 *) ((const uint8_t *) params->weight +
-                             b2_idx * params->src0_nb2 +
-                             b3_idx * params->src0_nb3);
+    return params->weight + b2_idx * params->src0_nb2 + b3_idx * params->src0_nb3;
 }

 static inline const float *hmx_mm_activation_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params,
@@ -3031,13 +3216,6 @@ static inline float *hmx_mm_dst_batch_ptr(const hmx_mm_f16_f32_batched_params_t
                       (size_t) dst_b3 * params->dst_nb3);
 }

-static inline const float *hmx_mm_src2_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params,
-                                               int src2_b2, int src2_b3) {
-    return params->src2 ? (const float *) ((const uint8_t *) params->src2 +
-                      (size_t) src2_b2 * params->src2_nb2 +
-                      (size_t) src2_b3 * params->src2_nb3) : NULL;
-}
-
 static int hmx_mm_f16_f32_batched_simple(struct htp_context *ctx,
                                                         const hmx_mm_f16_f32_batched_params_t *params,
                                                         int m_chunk, int n_chunk, int pipeline, int n_threads, int act_threads, int vtcm_size,
@@ -3045,15 +3223,18 @@ static int hmx_mm_f16_f32_batched_simple(struct htp_context *ctx,
     int ret = 0;
     for (int b3 = 0; b3 < params->ne13 && ret == 0; ++b3) {
         for (int b2 = 0; b2 < params->ne12 && ret == 0; ++b2) {
-            ret = hmx_mm_2d_f32(ctx, hmx_mm_dst_batch_ptr(params, b2, b3),
-                                           hmx_mm_src2_batch_ptr(params, b2, b3),
-                                           hmx_mm_activation_batch_ptr(params, b2, b3),
-                                           (const uint8_t *)hmx_mm_weight_batch_ptr(params, b2, b3),
-                                           params->m, params->k, params->n,
-                                           params->act_stride, params->weight_stride * (int)sizeof(__fp16),
-                                           HTP_TYPE_F16, params->k, params->dst_stride, params->src2_stride, params->n,
-                                           m_chunk, n_chunk, pipeline, n_threads, act_threads,
-                                           act_threads_div, k_div, 0, 0, vtcm_size);
+            dma_addr_t cur_src2_addr = params->src2_addr ? (params->src2_addr +
+                                       (dma_addr_t) b2 * params->src2_nb2 +
+                                       (dma_addr_t) b3 * params->src2_nb3) : 0;
+            ret = hmx_mm_2d_f32(ctx, params->weight_dma, hmx_mm_dst_batch_ptr(params, b2, b3),
+                                cur_src2_addr, params->src2_bytes,
+                                hmx_mm_activation_batch_ptr(params, b2, b3),
+                                hmx_mm_weight_batch_data(params, b2, b3),
+                                params->m, params->k, params->n,
+                                params->act_stride, params->weight_stride * (int)sizeof(__fp16),
+                                HTP_TYPE_F16, params->k, params->dst_stride, params->src2_stride, params->n,
+                                m_chunk, n_chunk, pipeline, n_threads, act_threads,
+                                act_threads_div, k_div, 0, 0, vtcm_size);
         }
     }
     return ret;
@@ -3095,7 +3276,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
     size_t vtcm_used = vtcm_size;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, HTP_TYPE_F16, params->k, m_chunk_n_rows, n_chunk_n_cols, group_size, use_dma_activation, false, act_threads, 0);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, HTP_TYPE_F16, params->k, m_chunk_n_rows, n_chunk_n_cols, group_size, use_dma_activation, false, act_threads, 0, params->src2_bytes);

     if (L.total_bytes > vtcm_budget) {
         FARF(HIGH, "%s: grouped layout overflowed VTCM, falling back to simple batched loop", __func__);
@@ -3112,6 +3293,13 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
     __fp16  *vtcm_scales     = VTCM_LAYOUT_PTR(__fp16, base, L.off_scales);
     float   *vtcm_f32_act    = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_act_f32, use_dma_activation);

+    const bool has_src2      = (params->src2_bytes > 0 && params->src2_addr != 0);
+    float   *vtcm_src2       = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_src2, has_src2);
+    if (has_src2) {
+        dma_queue_push(params->weight_dma, dma_make_data(vtcm_src2, params->src2_addr), hex_align_up(params->src2_bytes, 128), 0, params->src2_bytes, 1);
+        dma_queue_pop(params->weight_dma);
+    }
+
     hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00));  // scale: 1.0, bias: 0.0 in FP16

     FARF(HIGH, "%s: grouped path m=%d k=%d n=%d group=%d streams=%d mc=%zu nc=%zu vtcm=%zu/%zu",
@@ -3128,7 +3316,8 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_

     for (int b3 = 0; b3 < params->ne13; ++b3) {
         for (int b2_base = 0; b2_base < params->ne12; b2_base += group_size) {
-            const __fp16 *weight_group = hmx_mm_weight_batch_ptr(params, b2_base, b3);
+            const dma_addr_t weight_group = hmx_mm_weight_batch_data(params, b2_base, b3);
+            dma_queue * weight_dma = params->weight_dma;

             for (size_t mr = 0; mr < (size_t) params->m; mr += m_chunk_n_rows) {
                 const size_t n_rows = hex_smin((size_t) params->m - mr, m_chunk_n_rows);
@@ -3162,12 +3351,12 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                 // Prologue: Push A0 and A1 (if exists)
                 {
                     const size_t n_cols_first = hex_smin((size_t) params->n, n_chunk_n_cols);
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_scratch0, weight_group),
+                    dma_queue_push(weight_dma, dma_make_data(vtcm_scratch0, weight_group),
                                       fp16_row_bytes, weight_row_bytes, fp16_row_bytes, n_cols_first);
                 }
                 if (n_chunk_n_cols < (size_t) params->n) {
                     const size_t n_cols_second = hex_smin((size_t) params->n - n_chunk_n_cols, n_chunk_n_cols);
-                    dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_scratch1, weight_group + params->weight_stride),
+                    dma_queue_push(weight_dma, dma_make_data(vtcm_scratch1, weight_group + params->weight_stride * sizeof(__fp16)),
                                       fp16_row_bytes, weight_row_bytes, fp16_row_bytes, n_cols_second);
                 }

@@ -3176,16 +3365,16 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                     const size_t n_col_tiles = hmx_ceil_div((int) n_cols, HTP_MM_HMX_TILE_N_COLS);

                     {
-                        void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+                        void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

                         hmx_interleave_rows_to_tiles(vtcm_weight, (const __fp16 *) curr_raw, n_cols, params->k, params->k, 0, n_cols);

                         const size_t nc_next = nc + n_chunk_n_cols * 2;
                         if (nc_next < (size_t) params->n) {
                             const size_t n_cols_next = hex_smin((size_t) params->n - nc_next, n_chunk_n_cols);
-                            const __fp16 *next_weight_chunk = weight_group + nc_next * params->weight_stride;
+                            const dma_addr_t next_weight_chunk = weight_group + nc_next * params->weight_stride * sizeof(__fp16);

-                            dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, next_weight_chunk),
+                            dma_queue_push(weight_dma, dma_make_data(curr_raw, next_weight_chunk),
                                               fp16_row_bytes, weight_row_bytes, fp16_row_bytes, n_cols_next);
                         }
                     }
@@ -3201,7 +3390,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_

                         {
                             float *output = hmx_mm_dst_batch_ptr(params, b2_base + g, b3) + mr * params->dst_stride + nc;
-                            const float *src2_chunk = params->src2 ? (hmx_mm_src2_batch_ptr(params, b2_base + g, b3) + mr * params->src2_stride + nc) : NULL;
+                            const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * params->src2_stride + nc) : NULL;
                             int chunk_dst_cols = params->n - (int)nc;
                             if (chunk_dst_cols > 0) {
                                 transfer_output_chunk_threaded(ctx, output, src2_chunk, vtcm_output, (int) n_rows, (int) n_cols,
@@ -3314,9 +3503,10 @@ static void transfer_output_chunk_scattered_threaded(
 }

 static int hmx_mm_id_2d_f32(struct htp_context *ctx,
+                                         dma_queue *weight_dma,
                                          float *restrict dst,
                                          const float *activation,
-                                         const uint8_t *weight,
+                                         dma_addr_t weight,
                                          int m, int k, int n,
                                          int k_valid,
                                          int ne11,
@@ -3429,7 +3619,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
         if (n > 0) {
             const size_t n_cols = hex_smin((size_t) n, n_chunk_n_cols);
             const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
-            dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight, weight),
+            dma_queue_push(weight_dma, dma_make_data(vtcm_weight, weight),
                            dma_dst_stride, dma_src_stride, dma_width_bytes, height);
         }

@@ -3438,7 +3628,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
             const size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS);

             // A: Wait for weight DMA
-            void * curr_raw = dma_queue_pop(ctx->dma[0]).dst;
+            void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;

             // B: Weight Dequantize (Threaded)
             dequantize_tiled_weight_chunk_to_fp16_tiles(
@@ -3452,7 +3642,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
             if (nc_next < (size_t) n) {
                 const size_t n_cols_next = hex_smin((size_t) n - nc_next, n_chunk_n_cols);
                 const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
-                dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_next * weight_stride),
+                dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride),
                                dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
             }

@@ -3495,13 +3685,15 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
         return HTP_STATUS_OK;
     }

-    const float * src2_ptr = NULL;
+    dma_addr_t src2_addr = 0;
+    size_t src2_bytes = 0;
     uint32_t src2_stride = 0;
     size_t src2_nb2 = 0;
     size_t src2_nb3 = 0;
     if (src2) {
         src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
-        src2_ptr = (const float *) src2->data + m_start * src2_stride;
+        src2_addr = src2->data + (dma_addr_t) m_start * src2_stride * sizeof(float);
+        src2_bytes = (size_t) kparams->vtcm_src2_size;
         src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
         src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
     }
@@ -3515,9 +3707,11 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
     if (kparams->kernel_type == HTP_MM_KERNEL_HMX_F16_BATCHED) {
         hmx_mm_f16_f32_batched_params_t batch_params = {
             .dst             = dst_ptr,
-            .src2            = src2_ptr,
+            .src2_addr       = src2_addr,
+            .src2_bytes      = src2_bytes,
             .activation      = act_ptr,
-            .weight          = (const __fp16 *) src0->data,
+            .weight          = src0->data,
+            .weight_dma      = octx->ctx->dma[0],
             .m               = m_rows,
             .k               = k,
             .n               = n,
@@ -3551,7 +3745,7 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
                                      kparams->vtcm_size);
     } else {
         ret = hmx_mm_2d_f32(
-            octx->ctx, dst_ptr, src2_ptr, act_ptr, (const uint8_t *) src0->data,
+            octx->ctx, octx->ctx->dma[0], dst_ptr, src2_addr, src2_bytes, act_ptr, src0->data,
             m_rows, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
             dst_stride, src2_stride, (int)dst->ne[0],
             kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
@@ -3609,8 +3803,8 @@ static int hmx_mm_op_matmul_id(
         }
         if (m_start >= m_end) continue;

-        int ret = hmx_mm_id_2d_f32(octx->ctx, (float*) dst->data, (float*) src1->data,
-                                   (const uint8_t *) src0->data + cur_a * nb02,
+        int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) src1->data,
+                                   src0->data + cur_a * nb02,
                                    cne1, ne00, ne01,
                                    ne10,
                                    ne11,
@@ -3706,6 +3900,9 @@ static int hvx_mm_matmul_id(
     mmctx->vtcm_src2_size_per_thread = 0;
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

+    mmctx->cur_m_start = 0;
+    mmctx->cur_m_rows  = src1_nrows;
+
     mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->quant_task_func = quant_task_func;
     mmctx->n_quant_tasks = n_quant_tasks;
@@ -3753,8 +3950,8 @@ static int hmx_mm_op_matmul_id_nx(
             const struct htp_tensor * restrict dst   = octx->dsts[p];
             if (!src_w || !dst) continue;

-            int ret = hmx_mm_id_2d_f32(octx->ctx, (float*) dst->data, (float*) act->data,
-                                       (const uint8_t *) src_w->data + cur_a * src_w->nb[2],
+            int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) act->data,
+                                       src_w->data + cur_a * src_w->nb[2],
                                        cne1, src_w->ne[0], src_w->ne[1],
                                        act->ne[0],
                                        act->ne[1],
@@ -3844,6 +4041,9 @@ static int hvx_mm_matmul_id_nx(
     mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

+    mmctx->cur_m_start = 0;
+    mmctx->cur_m_rows  = src1_nrows;
+
     mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->quant_task_func         = quant_task_func;
     mmctx->n_quant_tasks           = n_quant_tasks;
@@ -3947,6 +4147,9 @@ int op_matmul_id(struct htp_ops_context * octx) {
     mmctx->act = src1;

     const struct htp_tensor * restrict ids = octx->src[2];
+    if (htp_tensor_is_extended(ids) || htp_tensor_is_extended(src1) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }

     const size_t src0_row_size = nb01;
     const size_t dst_row_size  = nb1;
@@ -3997,9 +4200,11 @@ int op_matmul_id(struct htp_ops_context * octx) {
     mmctx->matrix_row_counts    = matrix_row_counts;
     mmctx->matrix_rows          = matrix_rows;
     mmctx->mapping_stride       = mapping_stride;
-    mmctx->mm_div_ne11          = kparams->div_ne11;
+    mmctx->mm_div_ne11          = kparams->div_ne1;
     mmctx->src0_row_size_padded = src0_row_size_padded;
     mmctx->src1_nrows           = src1_nrows;
+    mmctx->cur_m_start          = 0;
+    mmctx->cur_m_rows           = src1_nrows;

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

@@ -4062,6 +4267,14 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
     const struct htp_tensor * restrict src0 = octx->src[0];
     const struct htp_tensor * restrict act  = octx->src[n_weights];
     const struct htp_tensor * restrict ids  = octx->src[n_weights + 1];
+    if (htp_tensor_is_extended(ids) || htp_tensor_is_extended(act)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+    for (uint32_t p = 0; p < n_weights; p++) {
+        if (octx->dsts[p] && htp_tensor_is_extended(octx->dsts[p])) {
+            return HTP_STATUS_NO_SUPPORT;
+        }
+    }

     mmctx->act = act;

@@ -4110,9 +4323,11 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
     mmctx->matrix_row_counts    = matrix_row_counts;
     mmctx->matrix_rows          = matrix_rows;
     mmctx->mapping_stride       = mapping_stride;
-    mmctx->mm_div_ne11          = kparams->div_ne11;
+    mmctx->mm_div_ne11          = kparams->div_ne1;
     mmctx->src0_row_size_padded = src0_row_size_padded;
     mmctx->src1_nrows           = src1_nrows;
+    mmctx->cur_m_start          = 0;
+    mmctx->cur_m_rows           = src1_nrows;

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

@@ -4163,6 +4378,9 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     mmctx->act  = act;

     const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
+    mmctx->src1_nrows  = src1_nrows;
+    mmctx->cur_m_start = 0;
+    mmctx->cur_m_rows  = src1_nrows;

     const size_t src0_row_size = src0->nb[1];
     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
@@ -4177,10 +4395,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     worker_callback_t quant_task_func;
     uint32_t n_quant_tasks = 1;
-    if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) {
-        n_quant_tasks = MIN(src1_nrows, octx->n_threads);
-        quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_flat : quantize_f32_q8_0_flat;
-    } else if (src1_nrows < octx->n_threads) {
+    if (src1_nrows < octx->n_threads) {
         n_quant_tasks = MIN(total_nb, octx->n_threads);
         quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
         for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -4196,12 +4411,9 @@ int op_matmul_nx(struct htp_ops_context * octx) {
         quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled;
     }

-    size_t src1_row_size;
-    if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) {
-        src1_row_size = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(act->ne[0]) : htp_mm_q8_0_flat_row_size(act->ne[0]);
-    } else {
-        src1_row_size = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]);
-    }
+    const size_t src1_row_size = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K)
+                               ? htp_mm_q8_1_tiled_row_size(act->ne[0])
+                               : htp_mm_q8_0_tiled_row_size(act->ne[0]);

     struct htp_mm_hvx_vtcm_layout L;
     htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
@@ -4242,26 +4454,14 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     const uint32_t n_matmul_jobs = octx->n_threads;
     worker_callback_t matmul_job_func;
     if (is_repacked) {
-        if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) {
-            switch (src0->type) {
-                case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_0_flat;   break;
-                case HTP_TYPE_Q4_1:
-                case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_1_flat;   break;
-                case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q8_0_flat;   break;
-                case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl_flat;  break;
-                case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4_flat;  break;
-                default:              return HTP_STATUS_NO_SUPPORT;
-            }
-        } else {
-            switch (src0->type) {
-                case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_0;   break;
-                case HTP_TYPE_Q4_1:
-                case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_1;   break;
-                case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q8_0;   break;
-                case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl;  break;
-                case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4;  break;
-                default:              return HTP_STATUS_NO_SUPPORT;
-            }
+        switch (src0->type) {
+            case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_0;   break;
+            case HTP_TYPE_Q4_1:
+            case HTP_TYPE_Q4_K:   matmul_job_func = hvx_mm_nx_2d_repacked_q4_1;   break;
+            case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q8_0;   break;
+            case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl;  break;
+            case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4;  break;
+            default:              return HTP_STATUS_NO_SUPPORT;
         }
     } else {
         matmul_job_func = hvx_mm_nx_2d;
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h
index 1df8c2933..fe9dbb61c 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.h
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h
@@ -62,17 +62,11 @@ enum htp_mm_kernel_type {

     // HVX floating-point paths
     HTP_MM_KERNEL_HVX_F16_F16_VTCM,
-    HTP_MM_KERNEL_HVX_F16_F16_DDR,
-    HTP_MM_KERNEL_HVX_F16_F32_DDR,
-
     HTP_MM_KERNEL_HVX_F32_F32_VTCM,
-    HTP_MM_KERNEL_HVX_F32_F32_DDR,
-    HTP_MM_KERNEL_HVX_F32_F16_DDR,

     // HVX quantized paths
     HTP_MM_KERNEL_HVX_QUANT_ROW,      // standard row-wise parallel quantization
     HTP_MM_KERNEL_HVX_QUANT_BLOCK,    // parallel block-wise quantization
-    HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, // row-wise fallback flat quantization
 };

 // Op-specific struct for precomputed matmul params
@@ -101,7 +95,7 @@ struct htp_mm_kernel_params {
     struct fastdiv_values div_ne1;
     struct fastdiv_values div_r2;
     struct fastdiv_values div_r3;
-    struct fastdiv_values div_ne11;
+    struct fastdiv_values div_ne12;
     struct fastdiv_values div_n_act_threads;
     struct fastdiv_values div_ne00_padded;
 };
@@ -246,20 +240,6 @@ static inline size_t htp_mm_q8_1_tiled_row_size(uint32_t ne) {
     return nb_32 * HTP_MM_ACT_TILE_SIZE_Q8_1;
 }

-static inline size_t htp_mm_q8_0_flat_row_size(uint32_t ne) {
-    const uint32_t quants_size = hex_align_up(ne, 128);
-    const uint32_t num_scales = (ne + 31) / 32;
-    const uint32_t scales_size = hex_align_up(num_scales * 2, 128);
-    return quants_size + scales_size;
-}
-
-static inline size_t htp_mm_q8_1_flat_row_size(uint32_t ne) {
-    const uint32_t quants_size = hex_align_up(ne, 128);
-    const uint32_t num_scales = (ne + 31) / 32;
-    const uint32_t scales_size = hex_align_up(num_scales * 4, 128);
-    return quants_size + scales_size;
-}
-
 static inline size_t htp_mm_get_tiled_row_stride(int weight_type, uint32_t k) {
     uint32_t nb = (k + QK_Q4_0_TILED - 1) / QK_Q4_0_TILED;
     switch (weight_type) {
@@ -331,6 +311,7 @@ struct htp_mm_hmx_vtcm_layout {
     size_t off_dst[2];        // [1] is only used when pipelined
     size_t off_scratch[2];    // dequantization scratch pads
     size_t off_scales;        // HMX scales (256 bytes)
+    size_t off_src2;          // src2 bias in VTCM

     // Cached sizes of regions for HMX kernel use
     size_t weight_area_bytes;
@@ -339,6 +320,7 @@ struct htp_mm_hmx_vtcm_layout {
     size_t output_area_bytes;
     size_t scratch_bytes[2];
     size_t act_head_stride;
+    size_t src2_bytes;

     size_t total_bytes;
 };
@@ -372,7 +354,8 @@ static inline void htp_mm_hmx_vtcm_layout_build(
     bool use_dma_activation,
     bool pipeline,
     uint32_t act_threads,
-    uint32_t aligned_tile_size
+    uint32_t aligned_tile_size,
+    size_t src2_size
 ) {
     size_t off = 0;

@@ -390,6 +373,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
         size_t off_group_a = 0;
         VTCM_LAYOUT_ALLOC(off_group_a, off_act, activation_area_size);
         VTCM_LAYOUT_ALLOC(off_group_a, off_scales, HTP_MM_HMX_TILE_SIZE); // Padded to 2K for alignment and future persistent data
+        VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE), src2_size > 0);

         // Group B: Compute-only buffers (starts at off_group_a)
         size_t off_group_b = off_group_a;
@@ -418,6 +402,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
         L->scratch_bytes[0]  = scratch_area_size;
         L->scratch_bytes[1]  = scratch_area_size;
         L->act_head_stride   = act_head_stride;
+        L->src2_bytes        = src2_size;

         off = off_group_a + hex_smax(group_b_size, group_c_size);
     } else {
@@ -441,6 +426,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
         size_t off_group_a = 0;
         VTCM_LAYOUT_ALLOC(off_group_a, off_scales, HTP_MM_HMX_TILE_SIZE); // Padded to 2K for alignment and future persistent data
         VTCM_LAYOUT_ALLOC(off_group_a, off_act, act_area_size);
+        VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE), src2_size > 0);

         // Group B: Compute-only buffers (starts at off_group_a)
         size_t off_group_b = off_group_a;
@@ -468,6 +454,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
         L->scratch_bytes[0]  = scratch0_size;
         L->scratch_bytes[1]  = scratch1_size;
         L->act_head_stride   = 0;
+        L->src2_bytes        = src2_size;

         off = off_group_a + hex_smax(group_b_size, group_c_size);
     }
@@ -490,6 +477,7 @@ static inline void htp_mm_hvx_vtcm_layout_build(
     bool is_matmul_id,
     bool is_fused_nx
 ) {
+    (void)src1_row_size;
     size_t src0_sz = 0;
     size_t src1_sz = 0;
     size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
@@ -517,12 +505,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
             weight_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128);
         }

-        size_t flat_act_row_size  = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(ne10)  : htp_mm_q8_0_flat_row_size(ne10);
         size_t tiled_act_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
-
-        size_t act_sz = (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT)
-            ? hex_round_up(flat_act_row_size  * src1_nrows, 128)
-            : hex_round_up(tiled_act_row_size * src1_nrows, 128);
+        size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);

         src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
         src1_sz = act_sz;                           // quantized activation buffer
@@ -547,6 +531,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(

         src0_sz = src0_sz_per_thread * n_threads;
         dst_sz  = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
+        src2_sz = 0;
+        src3_sz = 0;
     } else {
         const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
         const size_t dst_nrows = (src1_nrows > 1) ? 0 : 1;
@@ -559,15 +545,6 @@ static inline void htp_mm_hvx_vtcm_layout_build(
                 dst_sz  = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
                 break;
             }
-            case HTP_MM_KERNEL_HVX_F16_F32_DDR:
-            case HTP_MM_KERNEL_HVX_F16_F16_DDR:
-            case HTP_MM_KERNEL_HVX_F32_F32_DDR:
-            case HTP_MM_KERNEL_HVX_F32_F16_DDR: {
-                src0_sz = htp_mm_round_up(n_prefetch * src0_row_size, 256) * n_threads;
-                src1_sz = htp_mm_round_up(n_prefetch * src1_row_size, 256) * n_threads;
-                dst_sz  = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
-                break;
-            }
             case HTP_MM_KERNEL_HVX_F32_F32_VTCM: {
                 size_t f32_src1_row_size = htp_mm_round_up(ne10 * 4, 128);
                 src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
@@ -598,28 +575,6 @@ static inline void htp_mm_hvx_vtcm_layout_build(
                 dst_sz = dst_size_per_thread * n_threads;
                 break;
             }
-            case HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT: {
-                size_t q_src1_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
-
-                src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
-                src1_sz = htp_mm_round_up(q_src1_row_size * src1_nrows, 256);
-
-                src0_sz = src0_sz * n_threads;
-
-                if (is_repack) {
-                    uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype);
-                    uint32_t n_k_tiles = ne10 / 32;
-                    uint32_t tile_row_size = n_k_tiles * aligned_tile_size;
-                    size_t repacked_vtcm_size = htp_mm_round_up(n_prefetch * tile_row_size, 256);
-                    src0_sz = repacked_vtcm_size * n_threads;
-                }
-
-                size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
-                size_t dst_slice_per_thread = dst_nrows > 0 ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
-                size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
-                dst_sz = dst_size_per_thread * n_threads;
-                break;
-            }
             default:
                 break;
         }
@@ -640,19 +595,99 @@ static inline void htp_mm_hvx_vtcm_layout_build(
     L->total_bytes = off;
 }

+static inline bool htp_mm_hvx_solve_vtcm_params(
+    int kernel_type,
+    int wtype,
+    uint32_t ne10,
+    uint32_t src1_nrows,
+    uint32_t n_threads,
+    size_t dst_row_size,
+    size_t src0_row_size,
+    size_t src1_row_size,
+    size_t src2_row_size,
+    uint32_t n_prefetch,
+    size_t vtcm_budget,
+    struct htp_mm_hvx_vtcm_layout * L_out,
+    uint32_t * m_chunk_out
+) {
+    struct htp_mm_hvx_vtcm_layout L;
+    htp_mm_hvx_vtcm_layout_build(
+        &L, kernel_type, wtype, ne10, src1_nrows, n_threads,
+        dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+    );
+
+    if (L.total_bytes <= vtcm_budget) {
+        *L_out = L;
+        *m_chunk_out = src1_nrows;
+        return true;
+    }
+
+    const size_t fixed_bytes = L.src0_bytes + L.src2_bytes + L.dst_bytes;
+    if (vtcm_budget <= fixed_bytes) {
+        return false;
+    }
+
+    const size_t avail_act = vtcm_budget - fixed_bytes;
+    size_t row_size = 0;
+    if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
+        row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K)
+                 ? htp_mm_q8_1_tiled_row_size(ne10)
+                 : htp_mm_q8_0_tiled_row_size(ne10);
+    } else if (kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM) {
+        row_size = hex_round_up(ne10 * 2, 128);
+    } else {
+        row_size = hex_round_up(ne10 * 4, 128);
+    }
+    if (row_size == 0) {
+        return false;
+    }
+
+    uint32_t m_chunk = (uint32_t) (avail_act / row_size);
+    if (m_chunk > 1) {
+        m_chunk &= ~1U;
+    }
+    if (m_chunk > src1_nrows) {
+        m_chunk = src1_nrows;
+    }
+    if (m_chunk < 1) {
+        return false;
+    }
+
+    htp_mm_hvx_vtcm_layout_build(
+        &L, kernel_type, wtype, ne10, m_chunk, n_threads,
+        dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+    );
+
+    while (m_chunk > 2 && L.total_bytes > vtcm_budget) {
+        m_chunk -= 2;
+        htp_mm_hvx_vtcm_layout_build(
+            &L, kernel_type, wtype, ne10, m_chunk, n_threads,
+            dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+        );
+    }
+
+    if (L.total_bytes <= vtcm_budget) {
+        *L_out = L;
+        *m_chunk_out = m_chunk;
+        return true;
+    }
+
+    return false;
+}
+
 static inline size_t htp_mm_hmx_get_2d_vtcm_size(
-    int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size
+    int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size, size_t src2_size
 ) {
     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, false, pipeline, act_threads, aligned_tile_size);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, false, pipeline, act_threads, aligned_tile_size, src2_size);
     return L.total_bytes;
 }

 static inline size_t htp_mm_hmx_get_batched_vtcm_size(
-    int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool use_dma_activation, bool pipeline, uint32_t act_threads) {
+    int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool use_dma_activation, bool pipeline, uint32_t act_threads, size_t src2_size) {
     (void)pipeline;
     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, use_dma_activation, false, act_threads, 0);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, use_dma_activation, false, act_threads, 0, src2_size);
     return L.total_bytes;
 }

@@ -665,6 +700,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
     bool use_dma_activation,
     int n_threads,
     bool pipeline,
+    size_t src2_size,
     size_t vtcm_budget,
     size_t * m_chunk_out,
     size_t * n_chunk_out,
@@ -679,7 +715,7 @@ static inline bool htp_mm_hmx_solve_batched_params(

     int act_threads = n_threads;
     while (act_threads >= 1) {
-        size_t group_overhead = htp_mm_hmx_get_batched_overhead();
+        size_t group_overhead = htp_mm_hmx_get_batched_overhead() + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
         size_t group_size_per_n, group_size_per_m, group_size_per_mn;
         htp_mm_hmx_get_batched_chunk_costs(k, group_size, &group_size_per_n, &group_size_per_m, &group_size_per_mn);

@@ -690,7 +726,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
         if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
                                (size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
                                &m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
-            size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads);
+            size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads, src2_size);
             if (exact_size <= vtcm_budget) {
                 size_t mblocks = ((size_t) ne11 + m_chunk_candidate - 1) / m_chunk_candidate;
                 if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
@@ -730,6 +766,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
     bool pipeline,
     bool is_matmul_id,
     uint32_t aligned_tile_size,
+    size_t src2_size,
     size_t vtcm_budget,
     size_t * m_chunk_out,
     size_t * n_chunk_out,
@@ -746,7 +783,7 @@ static inline bool htp_mm_hmx_solve_2d_params(

     int act_threads = n_threads;
     while (act_threads >= 1) {
-        size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id);
+        size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id) + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
         size_t simple_2d_size_per_n, simple_2d_size_per_m, simple_2d_size_per_mn;
         htp_mm_hmx_get_2d_chunk_costs(wtype, k, pipeline, aligned_tile_size, &simple_2d_size_per_n, &simple_2d_size_per_m, &simple_2d_size_per_mn);

@@ -757,7 +794,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
         if (htp_mm_hmx_compute_chunks(vtcm_budget, simple_2d_overhead, simple_2d_size_per_n, simple_2d_size_per_m, simple_2d_size_per_mn, m_for_chunks, ne01_padded,
                                (size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) m_for_cost * HTP_MM_HMX_COST_A_CONVERT,
                                &m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
-            size_t exact_size = htp_mm_hmx_get_2d_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, pipeline, is_matmul_id ? 0 : act_threads, aligned_tile_size);
+            size_t exact_size = htp_mm_hmx_get_2d_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, pipeline, is_matmul_id ? 0 : act_threads, aligned_tile_size, src2_size);
             if (exact_size <= vtcm_budget) {
                 size_t mblocks = ((size_t) m_for_cost + m_chunk_candidate - 1) / m_chunk_candidate;
                 if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
diff --git a/ggml/src/ggml-hexagon/htp/pad-ops.c b/ggml/src/ggml-hexagon/htp/pad-ops.c
index 0222f24dc..85f25a8eb 100644
--- a/ggml/src/ggml-hexagon/htp/pad-ops.c
+++ b/ggml/src/ggml-hexagon/htp/pad-ops.c
@@ -7,7 +7,7 @@

 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"

 #define GGML_COMMON_DECL_C
@@ -51,6 +51,15 @@ static inline const uint8_t * pad_src_row_ptr(const struct htp_tensor * src,
         + (i3 - (uint32_t)lp3) * src->nb[3];
 }

+static inline dma_addr_t pad_src_row_data(const struct htp_tensor * src,
+                                          uint32_t i1, uint32_t i2, uint32_t i3,
+                                          int32_t lp1, int32_t lp2, int32_t lp3) {
+    return src->data
+        + (i1 - (uint32_t)lp1) * src->nb[1]
+        + (i2 - (uint32_t)lp2) * src->nb[2]
+        + (i3 - (uint32_t)lp3) * src->nb[3];
+}
+
 /* Compute the DDR src row pointer for a circular row (wrap-around indexing) */
 static inline const uint8_t * pad_circ_src_row_ptr(const struct htp_tensor * src,
                                                     uint32_t i1, uint32_t i2, uint32_t i3,
@@ -61,6 +70,15 @@ static inline const uint8_t * pad_circ_src_row_ptr(const struct htp_tensor * src
         + wrap_around((int32_t)i3 - lp3, src->ne[3]) * src->nb[3];
 }

+static inline dma_addr_t pad_circ_src_row_data(const struct htp_tensor * src,
+                                               uint32_t i1, uint32_t i2, uint32_t i3,
+                                               int32_t lp1, int32_t lp2, int32_t lp3) {
+    return src->data
+        + wrap_around((int32_t)i1 - lp1, src->ne[1]) * src->nb[1]
+        + wrap_around((int32_t)i2 - lp2, src->ne[2]) * src->nb[2]
+        + wrap_around((int32_t)i3 - lp3, src->ne[3]) * src->nb[3];
+}
+
 struct htp_pad_context {
     struct htp_ops_context * octx;

@@ -118,7 +136,7 @@ struct htp_pad_context {
     uint8_t * src_spad_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread; \
     uint8_t * dst_spad_base = octx->dst_spad.data  + ith * octx->dst_spad.size_per_thread;  \
                                                                                             \
-    dma_queue * dma = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];

 // ---------------------------------------------------------------------------
 // HVX vectorized PAD kernel
@@ -196,9 +214,9 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
         uint8_t * src_spad_cur = src_spad_base + spad_idx * src_row_size_aligned;
         uint8_t * dst_spad_cur = dst_spad_base + spad_idx * dst_row_size_aligned;

-        dma_queue_push_vtcm_to_ddr(dma,
-            dma_make_ptr((uint8_t *)dst->data, dst_spad_cur),
-            dst_row_size, dst_row_size_aligned, 0);
+        dma_queue_push(dma_q,
+            dma_make_data(dst->data, dst_spad_cur),
+            dst_row_size, dst_row_size_aligned, dst_row_size, 0);

         uint32_t i1, i2, i3;
         pad_decompose_row(ir, ne1, ne2, &i1, &i2, &i3);
@@ -207,15 +225,14 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
                                              lp2, rp2, ne2,
                                              lp3, rp3, ne3);

-        const uint8_t * src_ptr = interior
-            ? pad_src_row_ptr(src, i1, i2, i3, lp1, lp2, lp3) : NULL;
+        const dma_addr_t src_data = interior
+            ? pad_src_row_data(src, i1, i2, i3, lp1, lp2, lp3) : src->data;

         // Interior row: real DMA (1 row) from DDR to VTCM.
         // Border row: null DMA (nrows=0)
-        dma_queue_push_ddr_to_vtcm(dma,
-            dma_make_ptr(src_spad_cur,
-                         src_ptr ? src_ptr : (const uint8_t *)src_spad_cur),
-            src_row_size_aligned, src_row_size, src_ptr ? 1 : 0);
+        dma_queue_push(dma_q,
+            dma_make_data(src_spad_cur, src_data),
+            src_row_size_aligned, src_row_size, src_row_size, interior ? 1 : 0);
     }

     // -----------------------------------------------------------------------
@@ -225,13 +242,13 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];

     for (uint32_t ir = row_start; ir < row_end; ir++) {
-        uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
-        uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
+        uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i1, i2, i3;
         pad_decompose_row(ir, ne1, ne2, &i1, &i2, &i3);

-        uint8_t * dst_ptr = (uint8_t *) dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;
+        const dma_addr_t dst_data = dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;

         const int interior = pad_is_interior(i1, i2, i3,
                                              lp1, rp1, ne1,
@@ -254,9 +271,9 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        dma_queue_push_vtcm_to_ddr(dma,
-            dma_make_ptr(dst_ptr, dst_spad_cur),
-            dst_row_size, dst_row_size_aligned, 1);
+        dma_queue_push(dma_q,
+            dma_make_data(dst_data, dst_spad_cur),
+            dst_row_size, dst_row_size_aligned, dst_row_size, 1);

         const uint32_t next_row = ir + 2;
         if (next_row < row_end) {
@@ -266,17 +283,16 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
                                                       lp1, rp1, ne1,
                                                       lp2, rp2, ne2,
                                                       lp3, rp3, ne3);
-            const uint8_t * next_src_ptr = next_interior
-                ? pad_src_row_ptr(src, ni1, ni2, ni3, lp1, lp2, lp3) : NULL;
+            const dma_addr_t next_src_data = next_interior
+                ? pad_src_row_data(src, ni1, ni2, ni3, lp1, lp2, lp3) : src->data;

-            dma_queue_push_ddr_to_vtcm(dma,
-                dma_make_ptr(src_spad_cur,
-                             next_src_ptr ? next_src_ptr : (const uint8_t *)src_spad_cur),
-                src_row_size_aligned, src_row_size, next_src_ptr ? 1 : 0);
+            dma_queue_push(dma_q,
+                dma_make_data(src_spad_cur, next_src_data),
+                src_row_size_aligned, src_row_size, src_row_size, next_interior ? 1 : 0);
         }
     }

-    dma_queue_flush(dma);
+    dma_queue_flush(dma_q);

     FARF(HIGH, "pad-hvx-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
@@ -372,15 +388,16 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
         uint8_t * src_spad_cur = src_spad_base + spad_idx * src_row_size_aligned;
         uint8_t * dst_spad_cur = dst_spad_base + spad_idx * dst_row_size_aligned;

-        dma_queue_push_vtcm_to_ddr(dma,
-            dma_make_ptr((uint8_t *)dst->data, dst_spad_cur),
-            dst_row_size, dst_row_size_aligned, 0);
+        dma_queue_push(dma_q,
+            dma_make_data(dst->data, dst_spad_cur),
+            dst_row_size, dst_row_size_aligned, dst_row_size, 0);

         uint32_t pi1, pi2, pi3;
         pad_decompose_row(ir, ne1, ne2, &pi1, &pi2, &pi3);
-        dma_queue_push_ddr_to_vtcm(dma,
-            dma_make_ptr(src_spad_cur, pad_circ_src_row_ptr(src, pi1, pi2, pi3, lp1, lp2, lp3)),
-            src_row_size_aligned, src_row_size, 1);
+        const dma_addr_t src_data = pad_circ_src_row_data(src, pi1, pi2, pi3, lp1, lp2, lp3);
+        dma_queue_push(dma_q,
+            dma_make_data(src_spad_cur, src_data),
+            src_row_size_aligned, src_row_size, src_row_size, 1);
     }

     // -----------------------------------------------------------------------
@@ -390,12 +407,12 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];

     for (uint32_t ir = row_start; ir < row_end; ir++) {
-        uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
-        uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
+        uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma_q).src;
+        uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma_q).dst;

         uint32_t i1, i2, i3;
         pad_decompose_row(ir, ne1, ne2, &i1, &i2, &i3);
-        uint8_t * dst_ptr = (uint8_t *) dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;
+        const dma_addr_t dst_data = dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         if (lp0 > 0) {
@@ -431,22 +448,22 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

-        dma_queue_push_vtcm_to_ddr(dma,
-            dma_make_ptr(dst_ptr, dst_spad_cur),
-            dst_row_size, dst_row_size_aligned, 1);
+        dma_queue_push(dma_q,
+            dma_make_data(dst_data, dst_spad_cur),
+            dst_row_size, dst_row_size_aligned, dst_row_size, 1);

         const uint32_t next_row = ir + 2;
         if (next_row < row_end) {
             uint32_t nri1, nri2, nri3;
             pad_decompose_row(next_row, ne1, ne2, &nri1, &nri2, &nri3);
-            dma_queue_push_ddr_to_vtcm(dma,
-                dma_make_ptr(src_spad_cur,
-                             pad_circ_src_row_ptr(src, nri1, nri2, nri3, lp1, lp2, lp3)),
-                src_row_size_aligned, src_row_size, 1);
+            const dma_addr_t next_src_data = pad_circ_src_row_data(src, nri1, nri2, nri3, lp1, lp2, lp3);
+            dma_queue_push(dma_q,
+                dma_make_data(src_spad_cur, next_src_data),
+                src_row_size_aligned, src_row_size, src_row_size, 1);
         }
     }

-    dma_queue_flush(dma);
+    dma_queue_flush(dma_q);

     FARF(HIGH, "pad-hvx-circ-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
@@ -468,10 +485,6 @@ int op_pad(struct htp_ops_context * octx) {
             return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const int32_t lp0 = octx->op_params[0];
     const int32_t rp0 = octx->op_params[1];
     const int32_t lp1 = octx->op_params[2];
@@ -515,6 +528,10 @@ int op_pad(struct htp_ops_context * octx) {
     const int use_dma = (src0->nb[0] == (uint32_t)type_size) && (ne00 >= 512) &&
                         (octx->ctx->vtcm_size >= vtcm_needed);

+    if (!use_dma && (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst))) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     if (use_dma) {
         octx->src0_spad.size_per_thread = 2 * src_row_size_aligned;
         octx->dst_spad.size_per_thread  = 2 * dst_row_size_aligned;
diff --git a/ggml/src/ggml-hexagon/htp/repeat-ops.c b/ggml/src/ggml-hexagon/htp/repeat-ops.c
index 530279d65..2551be225 100644
--- a/ggml/src/ggml-hexagon/htp/repeat-ops.c
+++ b/ggml/src/ggml-hexagon/htp/repeat-ops.c
@@ -122,8 +122,8 @@ int op_repeat(struct htp_ops_context * octx) {
             return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     const uint32_t total_dst_rows  = dst->ne[1] * dst->ne[2] * dst->ne[3];
diff --git a/ggml/src/ggml-hexagon/htp/roll-ops.c b/ggml/src/ggml-hexagon/htp/roll-ops.c
index 6faf2ac47..9c373f56d 100644
--- a/ggml/src/ggml-hexagon/htp/roll-ops.c
+++ b/ggml/src/ggml-hexagon/htp/roll-ops.c
@@ -67,8 +67,8 @@ static inline uint32_t htp_roll_wrap(int32_t i, uint32_t ne) {
 #define htp_roll_dma_preamble dma_queue * q = octx->ctx->dma[0];

 static inline void roll_dma_push(dma_queue * q,
-                                 uintptr_t   dst,
-                                 uintptr_t   src,
+                                 dma_addr_t  dst,
+                                 dma_addr_t  src,
                                  uint32_t    dst_stride,
                                  uint32_t    src_stride,
                                  uint32_t    bytes,
@@ -77,10 +77,10 @@ static inline void roll_dma_push(dma_queue * q,
         return;
     }

-    if (!dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src), dst_stride, src_stride, bytes, nrows)) {
+    if (!dma_queue_push(q, dma_make_data(dst, src), dst_stride, src_stride, bytes, nrows)) {
         dma_queue_flush(q);
-        dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src),
-                                       dst_stride, src_stride, bytes, nrows);
+        dma_queue_push(q, dma_make_data(dst, src),
+                       dst_stride, src_stride, bytes, nrows);
     }
 }

@@ -92,29 +92,29 @@ static inline void roll_dma_push_rows(dma_queue *               q,
                                       uint32_t                  nrows,
                                       uint32_t                  row_size,
                                       uint32_t                  i0_src0) {
-    const uintptr_t dst_base = dst->data + (uintptr_t) dst_row * row_size;
-    const uintptr_t src_base = src0->data + (uintptr_t) src_row * row_size;
-    const uint32_t  n0       = src0->ne[0] - i0_src0;
+    const dma_addr_t dst_base = dst->data + (size_t) dst_row * row_size;
+    const dma_addr_t src_base = src0->data + (size_t) src_row * row_size;
+    const uint32_t   n0       = src0->ne[0] - i0_src0;

-    roll_dma_push(q, dst_base, src_base + (uintptr_t) i0_src0 * sizeof(float),
+    roll_dma_push(q, dst_base, src_base + (size_t) i0_src0 * sizeof(float),
                   row_size, row_size, n0 * sizeof(float), nrows);
-    roll_dma_push(q, dst_base + (uintptr_t) n0 * sizeof(float), src_base,
+    roll_dma_push(q, dst_base + (size_t) n0 * sizeof(float), src_base,
                   row_size, row_size, i0_src0 * sizeof(float), nrows);
 }

 // Same row-wrap split as roll_dma_push_rows, but addressed with explicit byte strides so it
 // also works for a src0 that is row-contiguous only (e.g. a permuted view) rather than fully packed.
 static inline void roll_dma_push_range(dma_queue * q,
-                                       uintptr_t   dst_row,
-                                       uintptr_t   src_row,
+                                       dma_addr_t  dst_row,
+                                       dma_addr_t  src_row,
                                        uint32_t    dst_stride,
                                        uint32_t    src_stride,
                                        uint32_t    nrows,
                                        uint32_t    i0_src0,
                                        uint32_t    n0) {
-    roll_dma_push(q, dst_row, src_row + (uintptr_t) i0_src0 * sizeof(float),
+    roll_dma_push(q, dst_row, src_row + (size_t) i0_src0 * sizeof(float),
                   dst_stride, src_stride, n0 * sizeof(float), nrows);
-    roll_dma_push(q, dst_row + (uintptr_t) n0 * sizeof(float), src_row,
+    roll_dma_push(q, dst_row + (size_t) n0 * sizeof(float), src_row,
                   dst_stride, src_stride, i0_src0 * sizeof(float), nrows);
 }

@@ -184,12 +184,12 @@ static int roll_dma_f32_strided(struct htp_ops_context * octx) {
         for (uint32_t i2 = 0; i2 < ne2; i2++) {
             const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2);

-            const uintptr_t dst_row0 = dst->data  + (uintptr_t) i2  * nb2  + (uintptr_t) i3  * nb3;
-            const uintptr_t src_row0 = src0->data + (uintptr_t) i02 * nb02 + (uintptr_t) i03 * nb03;
+            const dma_addr_t dst_row0 = dst->data  + (size_t) i2  * nb2  + (size_t) i3  * nb3;
+            const dma_addr_t src_row0 = src0->data + (size_t) i02 * nb02 + (size_t) i03 * nb03;

-            roll_dma_push_range(q, dst_row0, src_row0 + (uintptr_t) i1_src0 * nb01,
+            roll_dma_push_range(q, dst_row0, src_row0 + (size_t) i1_src0 * nb01,
                                 nb1, nb01, n1_first, i0_src0, n0);
-            roll_dma_push_range(q, dst_row0 + (uintptr_t) n1_first * nb1, src_row0,
+            roll_dma_push_range(q, dst_row0 + (size_t) n1_first * nb1, src_row0,
                                 nb1, nb01, i1_src0, i0_src0, n0);
         }
     }
@@ -223,8 +223,8 @@ static void roll_thread_f32(unsigned int nth, unsigned int ith, void * data) {
         const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2);
         const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3);

-        const uint8_t * src_row = (const uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
-        uint8_t * dst_row = (uint8_t *) dst->data + i1*nb1 + i2*nb2 + i3*nb3;
+        const uint8_t * src_row = (const uint8_t *) (uintptr_t) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
+        uint8_t * dst_row = (uint8_t *) (uintptr_t) dst->data + i1*nb1 + i2*nb2 + i3*nb3;

         hex_l2fetch(src_row + i0_src0 * sizeof(float), n0 * sizeof(float), ne0 * sizeof(float), 1);
         hvx_copy_uu(dst_row, src_row + i0_src0 * sizeof(float), n0, sizeof(float));
@@ -261,10 +261,6 @@ int execute_op_roll_f32(struct htp_ops_context * octx) {
         return HTP_STATUS_INVAL_PARAMS;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
-    }
-
     const uint32_t total_rows = ne1 * ne2 * ne3;
     const size_t dst_row_size = ne0 * sizeof(float);

@@ -290,6 +286,10 @@ int execute_op_roll_f32(struct htp_ops_context * octx) {
         return roll_dma_f32_strided(octx);
     }

+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     const uint32_t n_threads = octx->n_threads;
     struct htp_roll_context rctx = {
         .octx             = octx,
diff --git a/ggml/src/ggml-hexagon/htp/rope-ops.c b/ggml/src/ggml-hexagon/htp/rope-ops.c
index c36976ed0..f6b4d383c 100644
--- a/ggml/src/ggml-hexagon/htp/rope-ops.c
+++ b/ggml/src/ggml-hexagon/htp/rope-ops.c
@@ -9,7 +9,7 @@
 #include <string.h>
 #include <stdlib.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"
 #include "hex-fastdiv.h"

@@ -85,6 +85,8 @@ struct htp_rope_context {

     struct fastdiv_values div_ne2_ne1;
     struct fastdiv_values div_ne1;
+
+    const float * freq_factors;
 };

 static float rope_yarn_ramp(const float low, const float high, const int i0) {
@@ -562,10 +564,10 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
     float *   theta_cache    = (float *) (src0_spad_base);
               src0_spad_base = src0_spad_base + rctx->theta_cache_offset;

-    dma_queue * dma_queue = octx->ctx->dma[ith];
+    dma_queue * dma_q = octx->ctx->dma[ith];
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    const int32_t * pos = (const int32_t *) src1->data;
-    const float * freq_factors = src2 ? (const float *) src2->data : NULL;
+    const int32_t * pos = (const int32_t *) (uintptr_t) src1->data;
+    const float * freq_factors = rctx->freq_factors;

     const uint32_t i3_start = fastdiv(src0_start_row, &rctx->div_ne2_ne1);
     const uint32_t rem      = fastmodulo(src0_start_row, ne2 * ne1, &rctx->div_ne2_ne1);
@@ -587,7 +589,7 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                 const uint32_t nrows = MIN(src0_end_row - ir, ne1 - i1);

                 // Depth before prefetch
-                const uint32_t dma_depth = dma_queue_depth(dma_queue);
+                const uint32_t dma_depth = dma_queue_depth(dma_q);

                 // Prefetch up to 2 blocks
                 const uint32_t p_nrows = MIN(nrows, 2 * HTP_ROPE_SPAD_BLOCK);
@@ -595,12 +597,12 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                     const uint32_t pnr = MIN(nrows - pr, HTP_ROPE_SPAD_BLOCK);
                     const uint32_t slot = (cur_slot + pr / HTP_ROPE_SPAD_BLOCK) % HTP_ROPE_SPAD_NSLOTS;
                     uint8_t * spad_slot = rope_spad_slot(src0_spad_base, slot, rctx->src0_row_size_aligned);
-                    const uint8_t * src_addr = (const uint8_t *) src0->data + i3 * nb03 + i2 * nb02 + (i1 + pr) * nb01;
+                    const dma_addr_t src0_data = src0->data + i3 * nb03 + i2 * nb02 + (i1 + pr) * nb01;

                     // Dummy DMA transaction for sequencing (interleaving wr, rd, wr, rd, ...)
-                    dma_queue_push(dma_queue, dma_make_ptr((void *) dst->data, spad_slot), 0, 0, 0, 0);
+                    dma_queue_push(dma_q, dma_make_data(dst->data, spad_slot), 0, 0, 0, 0);

-                    dma_queue_push(dma_queue, dma_make_ptr(spad_slot, src_addr),
+                    dma_queue_push(dma_q, dma_make_data(spad_slot, src0_data),
                         rctx->src0_row_size_aligned, rctx->src0_row_stride, rctx->src0_row_size, pnr);
                 }

@@ -634,7 +636,7 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                 }

                 // Skip output DMA transactions from prev block (if any)
-                for (uint32_t d = 0; d < dma_depth; d++) { dma_queue_pop_nowait(dma_queue); }
+                for (uint32_t d = 0; d < dma_depth; d++) { dma_queue_pop_nowait(dma_q); }

                 // Compute loop
                 const uint32_t ne = is_vision ? ne0 : rctx->n_dims;
@@ -647,8 +649,8 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                     const uint32_t cur_ir   = base_ir + cr;
                     const uint32_t cur_i1   = base_i1 + cr;

-                    dma_queue_pop(dma_queue);
-                    uint8_t * cur_spad = (uint8_t *) dma_queue_pop(dma_queue).dst;
+                    dma_queue_pop(dma_q);
+                    uint8_t * cur_spad = (uint8_t *) dma_queue_pop(dma_q).dst;

                     htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, cur_ir);
                     if (is_neox || is_vision) {
@@ -658,8 +660,8 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                     }
                     htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, cur_ir);

-                    uint8_t * dst_addr = (uint8_t *) dst->data + i3 * nb3 + i2 * nb2 + cur_i1 * nb1;
-                    dma_queue_push(dma_queue, dma_make_ptr(dst_addr, cur_spad),
+                    const dma_addr_t dst_data = dst->data + i3 * nb3 + i2 * nb2 + cur_i1 * nb1;
+                    dma_queue_push(dma_q, dma_make_data(dst_data, cur_spad),
                         rctx->dst_row_stride, rctx->src0_row_size_aligned, rctx->dst_row_size, cnr);

                     // Prefetch 2 blocks ahead into the slot just freed
@@ -668,9 +670,9 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
                         const uint32_t pnr    = MIN(nrows - p_cr, HTP_ROPE_SPAD_BLOCK);
                         const uint32_t p_slot = (cur_slot + p_cr / HTP_ROPE_SPAD_BLOCK) % HTP_ROPE_SPAD_NSLOTS;
                         uint8_t * p_spad      = rope_spad_slot(src0_spad_base, p_slot, rctx->src0_row_size_aligned);
-                        const uint8_t * src_addr = (const uint8_t *) src0->data + i3 * nb03 + i2 * nb02 + (base_i1 + p_cr) * nb01;
+                        const dma_addr_t p_src0_data = src0->data + i3 * nb03 + i2 * nb02 + (base_i1 + p_cr) * nb01;

-                        dma_queue_push(dma_queue, dma_make_ptr(p_spad, src_addr),
+                        dma_queue_push(dma_q, dma_make_data(p_spad, p_src0_data),
                             rctx->src0_row_size_aligned, rctx->src0_row_stride, rctx->src0_row_size, pnr);
                     }
                 }
@@ -685,7 +687,7 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
     }

 done:
-    dma_queue_flush(dma_queue);
+    dma_queue_flush(dma_q);

     FARF(HIGH, "rope-f32: %d/%d: (%u:%u)\n", ith, nth, src0_start_row, src0_end_row);
 }
@@ -713,6 +715,10 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
     }
     assert(octx->ctx->vtcm_size >= kparams->vtcm_size);

+    if (htp_tensor_is_extended(src1)) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
     const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
     const size_t dst_data_row_size = dst->ne[0] * sizeof(float);

@@ -748,6 +754,16 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
     rctx.spad_per_thread       = kparams->spad_per_thread;
     rctx.theta_cache_offset    = kparams->theta_cache_offset;

+    if (src2) {
+        dma_queue * dma_q = octx->ctx->dma[0];
+        const size_t ff_size = src2->ne[0] * sizeof(float);
+        float * vtcm_freq_factors = (float *) (rctx.vtcm_base + kparams->freq_factors_offset);
+        dma_queue_push(dma_q, dma_make_data(vtcm_freq_factors, src2->data),
+                       kparams->freq_factors_size, 0, ff_size, 1);
+        dma_queue_pop(dma_q);
+        rctx.freq_factors = vtcm_freq_factors;
+    }
+
     const int32_t * op_params = &octx->op_params[0];
     rctx.n_dims     = ((const int32_t *) op_params)[1];
     rctx.mode       = ((const int32_t *) op_params)[2];
diff --git a/ggml/src/ggml-hexagon/htp/rope-ops.h b/ggml/src/ggml-hexagon/htp/rope-ops.h
index 476653d05..ee055ccbc 100644
--- a/ggml/src/ggml-hexagon/htp/rope-ops.h
+++ b/ggml/src/ggml-hexagon/htp/rope-ops.h
@@ -16,6 +16,8 @@ struct htp_rope_kernel_params {
     uint32_t spad_per_thread;
     uint32_t theta_cache_offset;
     uint32_t src0_row_size_aligned;
+    uint32_t freq_factors_offset;
+    uint32_t freq_factors_size;

     struct fastdiv_values div_ne2_ne1;
     struct fastdiv_values div_ne1;
@@ -32,21 +34,25 @@ struct htp_rope_vtcm_layout {
     size_t bytes_per_thread;
     size_t theta_cache_size_aligned;
     size_t src0_row_size_aligned;
+    size_t freq_factors_size_aligned;
 };

 static inline void htp_rope_vtcm_layout_build(
     struct htp_rope_vtcm_layout * layout,
     uint32_t ne00,
-    uint32_t n_threads
+    uint32_t n_threads,
+    uint32_t n_freq_factors
 ) {
-    const size_t src0_row_size            = ne00 * sizeof(float);
-    const size_t src0_row_size_aligned    = hex_round_up((uint32_t) src0_row_size, 128);
-    const size_t theta_cache_size_aligned = hex_round_up((uint32_t) src0_row_size, 256);
-
-    layout->src0_row_size_aligned    = src0_row_size_aligned;
-    layout->theta_cache_size_aligned = theta_cache_size_aligned;
-    layout->bytes_per_thread         = theta_cache_size_aligned + HTP_ROPE_SPAD_NROWS * src0_row_size_aligned;
-    layout->total_bytes              = layout->bytes_per_thread * n_threads;
+    const size_t src0_row_size             = ne00 * sizeof(float);
+    const size_t src0_row_size_aligned     = hex_round_up((uint32_t) src0_row_size, 128);
+    const size_t theta_cache_size_aligned  = hex_round_up((uint32_t) src0_row_size, 256);
+    const size_t freq_factors_size_aligned = hex_round_up(n_freq_factors * sizeof(float), 256);
+
+    layout->src0_row_size_aligned     = src0_row_size_aligned;
+    layout->theta_cache_size_aligned  = theta_cache_size_aligned;
+    layout->freq_factors_size_aligned = freq_factors_size_aligned;
+    layout->bytes_per_thread          = theta_cache_size_aligned + HTP_ROPE_SPAD_NROWS * src0_row_size_aligned;
+    layout->total_bytes               = layout->bytes_per_thread * n_threads + freq_factors_size_aligned;
 }

 static inline uint8_t * rope_spad_slot(uint8_t * base, uint32_t slot, size_t row_size_aligned) {
diff --git a/ggml/src/ggml-hexagon/htp/set-rows-ops.c b/ggml/src/ggml-hexagon/htp/set-rows-ops.c
index fbd5162a7..1d72538f1 100644
--- a/ggml/src/ggml-hexagon/htp/set-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/set-rows-ops.c
@@ -77,7 +77,7 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
         return;                                                                                                  \
     }                                                                                                            \
     const uint32_t ir1 = MIN(ir0 + dr, srctx->task_start + srctx->tasks);                                        \
-    dma_queue * dma_queue = octx->ctx->dma[ith];                                                                 \
+    dma_queue * dma_q = octx->ctx->dma[ith];                                                                     \
     const struct htp_set_rows_vtcm_layout * vtcm_layout = &srctx->vtcm_layout;                                   \
     uint8_t * vtcm_src0 = srctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;   \
     uint8_t * vtcm_dst  = srctx->vtcm_base + vtcm_layout->off_dst  + ith * vtcm_layout->dst_bytes_per_thread;    \
@@ -90,14 +90,14 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
     uint32_t pi03 = 0;                                                                                           \
     for (uint32_t step = 0, spad_idx = 0; step < total_steps && spad_idx < 2; ++step, spad_idx++) {              \
         uint32_t i = ir0 + pi_step;                                                                              \
-        const uintptr_t src0_ptr = octx->src[0]->data + i*nb01 + pi02*nb02 + pi03*nb03;                          \
-        dma_queue_push(dma_queue,                                                                                \
-                       dma_make_ptr((void *)octx->dst->data,                                                     \
-                                    vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size),                      \
+        const dma_addr_t src0_data = octx->src[0]->data + i*nb01 + pi02*nb02 + pi03*nb03;                        \
+        dma_queue_push(dma_q,                                                                                    \
+                       dma_make_data(octx->dst->data,                                                            \
+                                     vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size),                     \
                        dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 0);                          \
-        dma_queue_push(dma_queue,                                                                                \
-                       dma_make_ptr((void *)(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size),           \
-                                    (const void *)src0_ptr),                                                     \
+        dma_queue_push(dma_q,                                                                                    \
+                       dma_make_data(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size,                    \
+                                     src0_data),                                                                 \
                        vtcm_layout->src0_spad_half_size, src0_row_size, src0_row_size, 1);                       \
         pi_step++;                                                                                               \
         if (pi_step == nrows_per_thread) {                                                                       \
@@ -115,8 +115,8 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
     uint32_t ci11_base = 0;                                                                                      \
     uint32_t ci12_base = 0;                                                                                      \
     for (uint32_t step = 0; step < total_steps; ++step) {                                                        \
-        void * dst_spad = (void *) dma_queue_pop(dma_queue).src;                                                 \
-        void * src_spad = (void *) dma_queue_pop(dma_queue).dst;                                                 \
+        void * dst_spad = (void *) dma_queue_pop(dma_q).src;                                                     \
+        void * src_spad = (void *) dma_queue_pop(dma_q).dst;                                                     \
         uint32_t i = ir0 + ci_step;                                                                              \
         const uintptr_t src1_addr = octx->src[1]->data + i*nb10 + ci11_base*nb11 + ci12_base*nb12;               \
         const IDX_TYPE i1 = *(const IDX_TYPE *)src1_addr;                                                        \
@@ -128,21 +128,21 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
         }                                                                                                        \
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, step);                                                  \
         if (valid_i1) {                                                                                          \
-            const uintptr_t dst_ptr = octx->dst->data + target_i1*nb1 + ci02*nb2 + ci03*nb3;                     \
-            dma_queue_push(dma_queue,                                                                            \
-                           dma_make_ptr((void *)dst_ptr, (const void *)dst_spad),                                \
+            const dma_addr_t dst_data = octx->dst->data + target_i1*nb1 + ci02*nb2 + ci03*nb3;                   \
+            dma_queue_push(dma_q,                                                                                \
+                           dma_make_data(dst_data, dst_spad),                                                    \
                            dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 1);                      \
         } else {                                                                                                 \
-            dma_queue_push(dma_queue,                                                                            \
-                           dma_make_ptr((void *)octx->dst->data, (const void *)dst_spad),                        \
+            dma_queue_push(dma_q,                                                                                \
+                           dma_make_data(octx->dst->data, dst_spad),                                             \
                            dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 0);                      \
         }                                                                                                        \
         const uint32_t next_step = step + 2;                                                                     \
         if (next_step < total_steps) {                                                                           \
             uint32_t ni = ir0 + pi_step;                                                                         \
-            const uintptr_t psrc0_ptr = octx->src[0]->data + ni*nb01 + pi02*nb02 + pi03*nb03;                    \
-            dma_queue_push(dma_queue,                                                                            \
-                           dma_make_ptr((void *)src_spad, (const void *)psrc0_ptr),                              \
+            const dma_addr_t psrc0_data = octx->src[0]->data + ni*nb01 + pi02*nb02 + pi03*nb03;                  \
+            dma_queue_push(dma_q,                                                                                \
+                           dma_make_data(src_spad, psrc0_data),                                                  \
                            vtcm_layout->src0_spad_half_size, src0_row_size, src0_row_size, 1);                   \
             pi_step++;                                                                                           \
             if (pi_step == nrows_per_thread) {                                                                   \
@@ -172,7 +172,7 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
             }                                                                                                    \
         }                                                                                                        \
     }                                                                                                            \
-    dma_queue_flush(dma_queue);                                                                                  \
+    dma_queue_flush(dma_q);                                                                                      \
 }

 SET_ROWS_THREAD_DMA_FN(f32,  int32_t, { hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, ne00); })
@@ -196,8 +196,8 @@ int op_set_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(octx->src[1])) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     const struct htp_tensor * dst = octx->dst;
diff --git a/ggml/src/ggml-hexagon/htp/softmax-ops.c b/ggml/src/ggml-hexagon/htp/softmax-ops.c
index 2497ec763..48be5d725 100644
--- a/ggml/src/ggml-hexagon/htp/softmax-ops.c
+++ b/ggml/src/ggml-hexagon/htp/softmax-ops.c
@@ -8,54 +8,49 @@
 #include <math.h>
 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
+#include "work-queue.h"
 #include "hvx-utils.h"
 #include "hex-fastdiv.h"
+#include "hex-common.h"
+#include "hex-profile.h"

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
-#include "hex-common.h"
-#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
-
-#define htp_softmax_preamble3                     \
-    const uint32_t ne00 = src0->ne[0];            \
-    const uint32_t ne01 = src0->ne[1];            \
-    const uint32_t ne02 = src0->ne[2];            \
-    const uint32_t ne03 = src0->ne[3];            \
-                                                  \
-    const uint32_t nb00 = src0->nb[0];            \
-    const uint32_t nb01 = src0->nb[1];            \
-    const uint32_t nb02 = src0->nb[2];            \
-    const uint32_t nb03 = src0->nb[3];            \
-                                                  \
-    const uint32_t ne10 = src1 ? src1->ne[0] : 1; \
-    const uint32_t ne11 = src1 ? src1->ne[1] : 1; \
-    const uint32_t ne12 = src1 ? src1->ne[2] : 1; \
-    const uint32_t ne13 = src1 ? src1->ne[3] : 1; \
-                                                  \
-    const uint32_t nb10 = src1 ? src1->nb[0] : 1; \
-    const uint32_t nb11 = src1 ? src1->nb[1] : 1; \
-    const uint32_t nb12 = src1 ? src1->nb[2] : 1; \
-    const uint32_t nb13 = src1 ? src1->nb[3] : 1; \
-                                                  \
-    const uint32_t ne0 = dst->ne[0];              \
-    const uint32_t ne1 = dst->ne[1];              \
-    const uint32_t ne2 = dst->ne[2];              \
-    const uint32_t ne3 = dst->ne[3];              \
-                                                  \
-    const uint32_t nb0 = dst->nb[0];              \
-    const uint32_t nb1 = dst->nb[1];              \
-    const uint32_t nb2 = dst->nb[2];              \
-    const uint32_t nb3 = dst->nb[3];
+#include "htp-vtcm.h"
+#include "htp/softmax-ops.h"
+#include "hvx-flash-attn.h"

 struct htp_softmax_context {
     struct htp_ops_context * octx;
+    const struct htp_softmax_kernel_params * kparams;
+
+    void *   compute;
+
+    dma_addr_t data_src0;
+    dma_addr_t data_src1;
+    dma_addr_t data_dst;
+
+    uint8_t * vtcm_src0;
+    uint8_t * vtcm_src1;
+    uint8_t * vtcm_dst;
+
+    uint32_t  vtcm_src0_size_per_thread;
+    uint32_t  vtcm_src1_size_per_thread;
+    uint32_t  vtcm_dst_size_per_thread;
+
+    uint32_t  src0_spad_half_size;
+    uint32_t  src1_spad_half_size;
+    uint32_t  dst_spad_half_size;
+
+    uint32_t  src0_row_size_aligned;
+    uint32_t  src1_row_size_aligned;
+    uint32_t  dst_row_size_aligned;

     bool     use_f16;
-    bool     use_src1;

     uint32_t n_head;
     uint32_t n_head_log2;
@@ -65,68 +60,78 @@ struct htp_softmax_context {
     float    m0;
     float    m1;

-    struct fastdiv_values fastdiv_ne01;
-    struct fastdiv_values fastdiv_ne02;
-    struct fastdiv_values fastdiv_ne12; // For mask broadcasting
-    struct fastdiv_values fastdiv_ne13; // For mask broadcasting
+    struct fastdiv_values div_ne01;
+    struct fastdiv_values div_ne02;
+    struct fastdiv_values div_ne12;
+    struct fastdiv_values div_ne13;

     uint32_t src0_nrows_per_thread;
     uint32_t row_start;
     uint32_t nrows;
+
+    float    slopes[512] __attribute__((aligned(128)));
 };

-static void apply_mask(float * restrict wp0,
-                       const float * restrict mp_f32,
-                       const __fp16 * restrict mp_f16,
-                       uint32_t ne00,
-                       float slope,
-                       bool use_f16) {
-    if (!mp_f32) {
-        return;
-    }
-    if (use_f16) {
-        for (uint32_t i = 0; i < ne00; ++i) {
-            wp0[i] += slope * (float) mp_f16[i];
-        }
-    } else {
-        for (uint32_t i = 0; i < ne00; ++i) {
-            wp0[i] += slope * mp_f32[i];
-        }
-    }
-}
+typedef void (*softmax_compute_fn_t)(
+    void * restrict dst,
+    const void * restrict src0,
+    const void * restrict mask,
+    uint32_t ne00,
+    float scale,
+    float slope
+);

-static void init_softmax_ctx(struct htp_softmax_context * smctx, struct htp_ops_context * octx) {
-    const struct htp_tensor * src0 = octx->src[0];
-    const struct htp_tensor * src1 = octx->src[1];
+static void hvx_fast_softmax_prep_f16(const uint8_t * restrict src,
+                                      uint8_t * restrict dst,
+                                      const int num_elems,
+                                      float     scale,
+                                      const uint8_t * restrict mask,
+                                      float slope) {
+    const HVX_Vector * restrict v_src  = (const HVX_Vector *) src;
+    HVX_Vector * restrict v_dst        = (HVX_Vector *) dst;
+    const HVX_Vector * restrict v_mask = (const HVX_Vector *) mask;

-    memset(smctx, 0, sizeof(struct htp_softmax_context));
+    HVX_Vector scale_vec = hvx_vec_splat_f32(scale);
+    HVX_Vector slope_vec = hvx_vec_splat_f32(slope);

-    memcpy(&smctx->scale,    (float *) octx->op_params,     sizeof(float));
-    memcpy(&smctx->max_bias, (float *) octx->op_params + 1, sizeof(float));
+    const int nvec_64 = num_elems / VLEN_FP16;
+    const int nloe_64 = num_elems % VLEN_FP16;

-    smctx->n_head      = src0->ne[2];
-    smctx->n_head_log2 = 1u << (uint32_t) floor(log2(smctx->n_head));
+    #pragma unroll(2)
+    for (int i = 0; i < nvec_64; i++) {
+        HVX_VectorPair p = hvx_vec_f16_to_f32(v_mask[i]);
+        HVX_Vector m0 = Q6_V_lo_W(p);
+        HVX_Vector m1 = Q6_V_hi_W(p);

-    smctx->m0 = powf(2.0f, -(smctx->max_bias) / smctx->n_head_log2);
-    smctx->m1 = powf(2.0f, -(smctx->max_bias / 2.0f) / smctx->n_head_log2);
+        HVX_Vector s0 = v_src[2 * i];
+        HVX_Vector s1 = v_src[2 * i + 1];

-    smctx->use_src1 = (src1 != 0);
-    smctx->use_f16  = (src1 != 0) && (src1->type == HTP_TYPE_F16);
+        HVX_Vector v0 = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(s0, scale_vec), Q6_Vqf32_vmpy_VsfVsf(m0, slope_vec));
+        HVX_Vector v1 = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(s1, scale_vec), Q6_Vqf32_vmpy_VsfVsf(m1, slope_vec));

-    smctx->octx = octx;
+        v_dst[2 * i]     = Q6_Vsf_equals_Vqf32(v0);
+        v_dst[2 * i + 1] = Q6_Vsf_equals_Vqf32(v1);
+    }

-    // Initialize fastdiv values
-    const uint32_t ne01 = src0->ne[1];
-    const uint32_t ne02 = src0->ne[2];
+    if (nloe_64 > 0) {
+        HVX_VectorPair p = hvx_vec_f16_to_f32(v_mask[nvec_64]);
+        HVX_Vector m0 = Q6_V_lo_W(p);
+
+        HVX_Vector s0 = v_src[2 * nvec_64];
+        HVX_Vector v0 = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(s0, scale_vec), Q6_Vqf32_vmpy_VsfVsf(m0, slope_vec));

-    if (ne01 > 0) smctx->fastdiv_ne01 = init_fastdiv_values(ne01);
-    if (ne02 > 0) smctx->fastdiv_ne02 = init_fastdiv_values(ne02);
+        if (nloe_64 <= VLEN_FP32) {
+            hvx_vec_store_a(&v_dst[2 * nvec_64], nloe_64 * sizeof(float), Q6_Vsf_equals_Vqf32(v0));
+        } else {
+            v_dst[2 * nvec_64] = Q6_Vsf_equals_Vqf32(v0);

-    const uint32_t ne12 = src1 ? src1->ne[2] : 1;
-    const uint32_t ne13 = src1 ? src1->ne[3] : 1;
+            HVX_Vector m1 = Q6_V_hi_W(p);
+            HVX_Vector s1 = v_src[2 * nvec_64 + 1];
+            HVX_Vector v1 = Q6_Vqf32_vadd_Vqf32Vqf32(Q6_Vqf32_vmpy_VsfVsf(s1, scale_vec), Q6_Vqf32_vmpy_VsfVsf(m1, slope_vec));

-    if (ne12 > 0) smctx->fastdiv_ne12 = init_fastdiv_values(ne12);
-    if (ne13 > 0) smctx->fastdiv_ne13 = init_fastdiv_values(ne13);
+            hvx_vec_store_a(&v_dst[2 * nvec_64 + 1], (nloe_64 - VLEN_FP32) * sizeof(float), Q6_Vsf_equals_Vqf32(v1));
+        }
+    }
 }

 static void hvx_fast_softmax_prep_f32(const uint8_t * restrict src,
@@ -135,57 +140,68 @@ static void hvx_fast_softmax_prep_f32(const uint8_t * restrict src,
                                       float     scale,
                                       const uint8_t * restrict mask,
                                       float slope) {
-    const uint8_t * restrict src_curr  = src;
-    uint8_t * restrict dst_curr        = dst;
-    const uint8_t * restrict mask_curr = mask;
+    const HVX_Vector * restrict v_src  = (const HVX_Vector *) src;
+    HVX_Vector * restrict v_dst        = (HVX_Vector *) dst;
+    const HVX_Vector * restrict v_mask = (const HVX_Vector *) mask;

     HVX_Vector scale_vec = hvx_vec_splat_f32(scale);
     HVX_Vector slope_vec = hvx_vec_splat_f32(slope);

-    int step_of_1 = num_elems >> 5;
+    const int nvec = num_elems / VLEN_FP32;
+    const int nloe = num_elems % VLEN_FP32;

     #pragma unroll(4)
-    for (int i = 0; i < step_of_1; i++) {
-        HVX_Vector v1 = *(HVX_Vector *) src_curr;
-
-        HVX_Vector v3 = *(HVX_Vector *) mask_curr;
+    for (int i = 0; i < nvec; i++) {
+        HVX_Vector v1 = v_src[i];
+        HVX_Vector v3 = v_mask[i];

         HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(v1, scale_vec);
-
         HVX_Vector v4 = Q6_Vqf32_vmpy_VsfVsf(v3, slope_vec);
-
         HVX_Vector v5 = Q6_Vqf32_vadd_Vqf32Vqf32(v2, v4);

-        *(HVX_Vector *) dst_curr = Q6_Vsf_equals_Vqf32(v5);
+        v_dst[i] = Q6_Vsf_equals_Vqf32(v5);
+    }
+
+    if (nloe > 0) {
+        HVX_Vector v1 = v_src[nvec];
+        HVX_Vector v3 = v_mask[nvec];
+
+        HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(v1, scale_vec);
+        HVX_Vector v4 = Q6_Vqf32_vmpy_VsfVsf(v3, slope_vec);
+        HVX_Vector v5 = Q6_Vqf32_vadd_Vqf32Vqf32(v2, v4);

-        src_curr += VLEN;
-        dst_curr += VLEN;
-        mask_curr += VLEN;
+        hvx_vec_store_a(&v_dst[nvec], nloe * sizeof(float), Q6_Vsf_equals_Vqf32(v5));
     }
 }

-static void hvx_fast_softmax_f32(const uint8_t * restrict src, uint8_t * restrict dst, uint8_t * restrict pad, const int num_elems) {
-    const HVX_Vector * restrict v_src = (HVX_Vector *) src;
-    HVX_Vector * restrict v_pad       = (HVX_Vector *) pad;
+static void hvx_fast_softmax_f32(const uint8_t * restrict src, uint8_t * restrict dst, const int num_elems) {
+    const HVX_Vector * restrict v_src = (const HVX_Vector *) src;
     HVX_Vector * restrict v_dst       = (HVX_Vector *) dst;

-    HVX_Vector sum_vec = Q6_V_vsplat_R(0x00000000);
-    HVX_Vector max_vec = hvx_vec_splat_f32(((const float *) src)[0]);
-    HVX_Vector zero_v  = Q6_V_vzero();
-    HVX_Vector one_v   = hvx_vec_splat_f32(1.0);
+    const int nvec = num_elems / VLEN_FP32;
+    const int nloe = num_elems % VLEN_FP32;

-    int step_of_1 = num_elems >> 5;
+    HVX_Vector max_vec = hvx_vec_splat_f32(((const float *) src)[0]);

-    #pragma unroll(4)
-    for (int i = 0; i < step_of_1; i++) {
+    #pragma unroll(2)
+    for (int i = 0; i < nvec; i++) {
         HVX_Vector v1 = v_src[i];
         max_vec       = Q6_Vsf_vmax_VsfVsf(max_vec, v1);
     }

-    max_vec = hvx_vec_reduce_max_f32(max_vec); // replicated over all lanes
+    if (nloe > 0) {
+        HVX_VectorPred q_mask = Q6_Q_vsetq_R(nloe * sizeof(float));
+        HVX_Vector neg_inf    = hvx_vec_splat_f32(-INFINITY);
+        HVX_Vector v_tail     = Q6_V_vmux_QVV(q_mask, v_src[nvec], neg_inf);
+        max_vec               = Q6_Vsf_vmax_VsfVsf(max_vec, v_tail);
+    }

-    #pragma unroll(4)
-    for (int i = 0; i < step_of_1; i++) {
+    max_vec = hvx_vec_reduce_max_f32(max_vec);
+
+    HVX_Vector sum_vec = Q6_V_vsplat_R(0x00000000);
+
+    #pragma unroll(2)
+    for (int i = 0; i < nvec; i++) {
         HVX_Vector v1 = v_src[i];
         HVX_Vector v2 = Q6_Vqf32_vsub_VsfVsf(v1, max_vec);

@@ -193,39 +209,91 @@ static void hvx_fast_softmax_f32(const uint8_t * restrict src, uint8_t * restric

         sum_vec = Q6_Vqf32_vadd_VsfVsf(Q6_Vsf_equals_Vqf32(sum_vec), v3);

-        v_pad[i] = v3;
+        v_dst[i] = v3;
     }

-    sum_vec = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_vec)); // replicated over all lanes
+    if (nloe > 0) {
+        HVX_VectorPred q_mask = Q6_Q_vsetq_R(nloe * sizeof(float));
+        HVX_Vector v1     = v_src[nvec];
+        HVX_Vector v2     = Q6_Vqf32_vsub_VsfVsf(v1, max_vec);
+        HVX_Vector v3     = hvx_vec_exp_f32(Q6_Vsf_equals_Vqf32(v2));
+        HVX_Vector v3_pad = Q6_V_vmux_QVV(q_mask, v3, Q6_V_vzero());

-    HVX_VectorPred pos_sum   = Q6_Q_vcmp_gt_VwVw(sum_vec, zero_v);
+        sum_vec     = Q6_Vqf32_vadd_VsfVsf(Q6_Vsf_equals_Vqf32(sum_vec), v3_pad);
+        v_dst[nvec] = v3_pad;
+    }
+
+    sum_vec = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_vec));
+
+    HVX_VectorPred pos_sum   = Q6_Q_vcmp_gt_VwVw(sum_vec, Q6_V_vzero());
     HVX_Vector     v4        = hvx_vec_inverse_f32(sum_vec);
-    HVX_Vector     scale_vec = Q6_V_vmux_QVV(pos_sum, v4, one_v);
+    HVX_Vector     scale_vec = Q6_V_vmux_QVV(pos_sum, v4, hvx_vec_splat_f32(1.0f));

-    #pragma unroll(4)
-    for (int i = 0; i < step_of_1; i++) {
-        HVX_Vector v1 = v_pad[i];
+    #pragma unroll(2)
+    for (int i = 0; i < nvec; i++) {
+        HVX_Vector v1 = v_dst[i];
         HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(v1, scale_vec);
         v_dst[i]      = Q6_Vsf_equals_Vqf32(v2);
     }
+
+    if (nloe > 0) {
+        HVX_Vector v1 = v_dst[nvec];
+        HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(v1, scale_vec);
+        hvx_vec_store_a(&v_dst[nvec], nloe * sizeof(float), Q6_Vsf_equals_Vqf32(v2));
+    }
 }

-static float hvx_softmax_f32(const uint8_t * restrict src, uint8_t * restrict dst, uint8_t * restrict spad, const int  num_elems, const float max) {
-    hvx_sub_scalar_f32(spad, src, max, num_elems);
+static void compute_fast_softmax_f32_nomask(
+    void * restrict dst,
+    const void * restrict src0,
+    const void * restrict mask,
+    uint32_t ne00,
+    float scale,
+    float slope
+) {
+    (void) mask;
+    (void) slope;
+    hvx_scale_f32((uint8_t *) dst, (const uint8_t *) src0, ne00, scale);
+    hvx_fast_softmax_f32((const uint8_t *) dst, (uint8_t *) dst, ne00);
+}

-    hvx_exp_f32(dst, spad, num_elems, false);
-    return hvx_reduce_sum_f32(dst, num_elems);
+static void compute_fast_softmax_f32_mask_f32(
+    void * restrict dst,
+    const void * restrict src0,
+    const void * restrict mask,
+    uint32_t ne00,
+    float scale,
+    float slope
+) {
+    hvx_fast_softmax_prep_f32((const uint8_t *) src0, (uint8_t *) dst, ne00, scale, (const uint8_t *) mask, slope);
+    hvx_fast_softmax_f32((const uint8_t *) dst, (uint8_t *) dst, ne00);
 }

-static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
-    struct htp_softmax_context * smctx = (struct htp_softmax_context *) data;
-    struct htp_ops_context * octx = smctx->octx;
+static void compute_fast_softmax_f32_mask_f16(
+    void * restrict dst,
+    const void * restrict src0,
+    const void * restrict mask,
+    uint32_t ne00,
+    float scale,
+    float slope
+) {
+    hvx_fast_softmax_prep_f16((const uint8_t *) src0, (uint8_t *) dst, ne00, scale, (const uint8_t *) mask, slope);
+    hvx_fast_softmax_f32((const uint8_t *) dst, (uint8_t *) dst, ne00);
+}
+
+static const softmax_compute_fn_t softmax_kernels[HTP_SOFTMAX_KERNEL_COUNT] = {
+    [HTP_SOFTMAX_KERNEL_NOMASK]   = compute_fast_softmax_f32_nomask,
+    [HTP_SOFTMAX_KERNEL_MASK_F32] = compute_fast_softmax_f32_mask_f32,
+    [HTP_SOFTMAX_KERNEL_MASK_F16] = compute_fast_softmax_f32_mask_f16,
+};

+static void softmax_thread_dma(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_softmax_context * smctx = (const struct htp_softmax_context *) data;
+    struct htp_ops_context * octx = smctx->octx;
     const struct htp_tensor * src0 = octx->src[0];
-    const struct htp_tensor * src1 = octx->src[1];
     const struct htp_tensor * dst  = octx->dst;
-
-    htp_softmax_preamble3;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];

     const uint32_t src0_nrows            = smctx->nrows;
     const uint32_t src0_nrows_per_thread = smctx->src0_nrows_per_thread;
@@ -233,122 +301,213 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t src0_start_row = smctx->row_start + src0_nrows_per_thread * ith;
     const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, smctx->row_start + src0_nrows);

-    // no work for this thread
     if (src0_start_row >= src0_end_row) {
         return;
     }

-    int is_aligned = 1;
-    int opt_path   = 0;
+    const dma_addr_t data_src0 = smctx->data_src0;
+    const dma_addr_t data_dst  = smctx->data_dst;

-    if (!hex_is_aligned((void *) src0->data, VLEN) || !hex_is_aligned((void *) dst->data, VLEN)) {
-        is_aligned = 0;
-        FARF(HIGH, "softmax-f32: unaligned addresses in elementwise op, possibly slower execution\n");
-    }
+    const size_t src0_row_size = src0->ne[0] * sizeof(float);
+    const size_t dst_row_size  = src0->ne[0] * sizeof(float);

-    // Only use the fast path when aligned AND row size is multiple of VLEN (128 bytes)
-    // The fast path (hvx_fast_softmax_f32) doesn't handle tail elements
-    // The non-opt path uses hvx_softmax_f32 which properly handles all sizes via its helper functions
-    if ((1 == is_aligned) && !(nb01 & (VLEN - 1))) {
-        opt_path = 1;
-    }
+    uint8_t * src0_vtcm_base = smctx->vtcm_src0 + (ith * smctx->vtcm_src0_size_per_thread);
+    uint8_t * dst_vtcm_base  = smctx->vtcm_dst  + (ith * smctx->vtcm_dst_size_per_thread);

-    uint8_t * src0_spad_data = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
-    uint8_t * src1_spad_data = octx->src1_spad.data + (ith * octx->src1_spad.size_per_thread);
-    uint8_t * dst_spad_data  = octx->dst_spad.data  + (ith * octx->dst_spad.size_per_thread);
+    const size_t src0_vtcm_half = smctx->src0_spad_half_size;
+    const size_t dst_vtcm_half  = smctx->dst_spad_half_size;

-    float * wp0 = (float *) src0_spad_data;
-    float * wp1 = (float *) src1_spad_data;
-    float * wp2 = (float *) dst_spad_data;
+    dma_queue * dma_q = octx->ctx->dma[ith];

-    uint32_t prev_i2 = (uint32_t)-1;
-    float slope = 1.0f;
+    for (uint32_t r = src0_start_row, idx = 0; r < src0_end_row && idx < 2; r++, idx++) {
+        dma_addr_t cur_dst  = data_dst  + r * dst_row_size;
+        dma_addr_t cur_src0 = data_src0 + r * src0_row_size;
+        void * d_spad = dst_vtcm_base  + idx * dst_vtcm_half;
+        void * s_spad = src0_vtcm_base + idx * src0_vtcm_half;

-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
+        dma_queue_push(dma_q, dma_make_data(cur_dst, d_spad),
+                       dst_row_size, smctx->dst_row_size_aligned, dst_row_size, 0);
+        dma_queue_push(dma_q, dma_make_data(s_spad, cur_src0),
+                       smctx->src0_row_size_aligned, src0_row_size, src0_row_size, 1);
+    }
+
+    softmax_compute_fn_t compute = (softmax_compute_fn_t) smctx->compute;
+    const uint32_t ne00 = src0->ne[0];

     for (uint32_t r = src0_start_row; r < src0_end_row; ++r) {
-        uint32_t i1 = fastmodulo(r, ne01, &smctx->fastdiv_ne01);
-        uint32_t r_div_ne01 = fastdiv(r, &smctx->fastdiv_ne01);
-        uint32_t i2 = fastmodulo(r_div_ne01, ne02, &smctx->fastdiv_ne02);
-        uint32_t i3 = fastdiv(r_div_ne01, &smctx->fastdiv_ne02);
-
-        // Map to original logic indices
-        // i01 = i1
-        // i02 = i2
-        // i03 = i3
-
-        const uint32_t i11 = i1;
-        // const uint32_t i12 = i2 % ne12;
-        // const uint32_t i13 = i3 % ne13;
-
-        uint32_t i12, i13;
-        if (ne12 == ne02) {
-             i12 = i2;
-        } else {
-             i12 = fastmodulo(i2, ne12, &smctx->fastdiv_ne12);
+        void * d_spad = (void *) dma_queue_pop(dma_q).src;
+        void * s_spad = (void *) dma_queue_pop(dma_q).dst;
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);
+        compute(d_spad, s_spad, NULL, ne00, smctx->scale, 1.0f);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);
+
+        dma_addr_t cur_dst = data_dst + r * dst_row_size;
+        dma_queue_push(dma_q, dma_make_data(cur_dst, d_spad),
+                       dst_row_size, smctx->dst_row_size_aligned, dst_row_size, 1);
+
+        const uint32_t next_r = r + 2;
+        if (next_r < src0_end_row) {
+            dma_addr_t next_src0 = data_src0 + next_r * src0_row_size;
+            dma_queue_push(dma_q, dma_make_data(s_spad, next_src0),
+                           smctx->src0_row_size_aligned, src0_row_size, src0_row_size, 1);
         }
+    }

-        if (ne13 == ne03) {
-             i13 = i3;
-        } else {
-             i13 = fastmodulo(i3, ne13, &smctx->fastdiv_ne13);
-        }
+    dma_queue_flush(dma_q);
+}

-        // ALiBi
-        if (i2 != prev_i2) {
-            const uint32_t h = i2;  // head
-            slope = (smctx->max_bias > 0.0f) ? h < smctx->n_head_log2 ? powf(smctx->m0, h + 1) : powf(smctx->m1, 2 * (h - smctx->n_head_log2) + 1) : 1.0f;
-            prev_i2 = i2;
-        }
+static void softmax_thread_mask_dma(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_softmax_context * smctx = (const struct htp_softmax_context *) data;
+    struct htp_ops_context * octx = smctx->octx;
+    const struct htp_tensor * src0 = octx->src[0];
+    const struct htp_tensor * src1 = octx->src[1];
+    const struct htp_tensor * dst  = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];

-        float * sp = (float *) ((char *) src0->data + i1 * nb01 + i2 * nb02 + i3 * nb03);
-        float * dp = (float *) ((char *) dst->data  + i1 * nb1  + i2 * nb2  + i3 * nb3);
+    const uint32_t src0_nrows            = smctx->nrows;
+    const uint32_t src0_nrows_per_thread = smctx->src0_nrows_per_thread;

-        // broadcast the mask across rows
-        __fp16 * mp_f16 = (smctx->use_src1) ? (__fp16 *) ((char *) src1->data + i11 * nb11 + i12 * nb12 + i13 * nb13) : NULL;
-        float *  mp_f32 = (smctx->use_src1) ? (float *)  ((char *) src1->data + i11 * nb11 + i12 * nb12 + i13 * nb13) : NULL;
+    const uint32_t src0_start_row = smctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, smctx->row_start + src0_nrows);

-        if ((1 == opt_path) && (mp_f32) && !(smctx->use_f16)) {
-            hvx_fast_softmax_prep_f32((const uint8_t *) sp, (uint8_t *) wp0, ne00, smctx->scale, (const uint8_t *) mp_f32, slope);
-            hvx_fast_softmax_f32((const uint8_t *) wp0, (uint8_t *) dp, (uint8_t *) wp1, ne00);
-        } else if (1 == opt_path) {
-            hvx_scale_f32((uint8_t *) wp0, (const uint8_t *) sp, ne00, smctx->scale);
-            apply_mask(wp0, mp_f32, mp_f16, ne00, slope, smctx->use_f16);
-            hvx_fast_softmax_f32((const uint8_t *) wp0, (uint8_t *) dp, (uint8_t *) wp1, ne00);
-        } else {
-            // Non-optimized path: uses HVX helper functions that properly handle all tensor sizes
-            // including non-multiples of 32 (the HVX vector lane count for f32)
-            hvx_scale_f32((uint8_t *) wp0, (const uint8_t *) sp, ne00, smctx->scale);
-            apply_mask(wp0, mp_f32, mp_f16, ne00, slope, smctx->use_f16);
-            float max = hvx_reduce_max_f32((const uint8_t *) wp0, ne00);
-            float sum = hvx_softmax_f32((const uint8_t *) wp0, (uint8_t *) wp2, (uint8_t *) wp1, ne00, max);
-            sum       = sum > 0.0 ? (1.0 / sum) : 1;
-            hvx_scale_f32((uint8_t *) dp, (const uint8_t *) wp2, ne00, sum);
-        }
+    if (src0_start_row >= src0_end_row) {
+        return;
     }

-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
+    const dma_addr_t data_src0 = smctx->data_src0;
+    const dma_addr_t data_src1 = smctx->data_src1;
+    const dma_addr_t data_dst  = smctx->data_dst;
+
+    const size_t src0_row_size = src0->ne[0] * sizeof(float);
+    const size_t dst_row_size  = src0->ne[0] * sizeof(float);
+    const size_t mask_row_size = smctx->use_f16 ? (src1->ne[0] * sizeof(__fp16)) : (src1->ne[0] * sizeof(float));
+
+    uint8_t * src0_vtcm_base = smctx->vtcm_src0 + (ith * smctx->vtcm_src0_size_per_thread);
+    uint8_t * src1_vtcm_base = smctx->vtcm_src1 + (ith * smctx->vtcm_src1_size_per_thread);
+    uint8_t * dst_vtcm_base  = smctx->vtcm_dst  + (ith * smctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_vtcm_half = smctx->src0_spad_half_size;
+    const size_t src1_vtcm_half = smctx->src1_spad_half_size;
+    const size_t dst_vtcm_half  = smctx->dst_spad_half_size;
+
+    const uint32_t nb11 = src1->nb[1];
+    const uint32_t nb12 = src1->nb[2];
+    const uint32_t nb13 = src1->nb[3];
+
+    const uint32_t ne00 = src0->ne[0];
+    const uint32_t ne01 = src0->ne[1];
+    const uint32_t ne02 = src0->ne[2];
+    const uint32_t ne03 = src0->ne[3];
+    const uint32_t ne12 = src1->ne[2];
+    const uint32_t ne13 = src1->ne[3];
+
+    const struct fastdiv_values * div_ne01 = &smctx->div_ne01;
+    const struct fastdiv_values * div_ne02 = &smctx->div_ne02;
+    const struct fastdiv_values * div_ne12 = &smctx->div_ne12;
+    const struct fastdiv_values * div_ne13 = &smctx->div_ne13;
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    for (uint32_t r = src0_start_row, idx = 0; r < src0_end_row && idx < 2; r++, idx++) {
+        dma_addr_t cur_dst  = data_dst  + r * dst_row_size;
+        dma_addr_t cur_src0 = data_src0 + r * src0_row_size;
+
+        uint32_t i1 = fastmodulo(r, ne01, div_ne01);
+        uint32_t r_div_ne01 = fastdiv(r, div_ne01);
+        uint32_t i2 = fastmodulo(r_div_ne01, ne02, div_ne02);
+        uint32_t i3 = fastdiv(r_div_ne01, div_ne02);
+        uint32_t i12 = (ne12 == ne02) ? i2 : fastmodulo(i2, ne12, div_ne12);
+        uint32_t i13 = (ne13 == ne03) ? i3 : fastmodulo(i3, ne13, div_ne13);
+        dma_addr_t cur_src1 = data_src1 + i1 * nb11 + i12 * nb12 + i13 * nb13;
+
+        void * d_spad = dst_vtcm_base  + idx * dst_vtcm_half;
+        void * s_spad = src0_vtcm_base + idx * src0_vtcm_half;
+        void * m_spad = src1_vtcm_base + idx * src1_vtcm_half;
+
+        dma_queue_push(dma_q, dma_make_data(cur_dst, d_spad),
+                       dst_row_size, smctx->dst_row_size_aligned, dst_row_size, 0);
+        dma_queue_push(dma_q, dma_make_data(s_spad, cur_src0),
+                       smctx->src0_row_size_aligned, src0_row_size, src0_row_size, 1);
+        dma_queue_push(dma_q, dma_make_data(m_spad, cur_src1),
+                       smctx->src1_row_size_aligned, mask_row_size, mask_row_size, 1);
+    }
+
+    softmax_compute_fn_t compute = (softmax_compute_fn_t) smctx->compute;
+    const bool has_bias = smctx->max_bias > 0.0f;
+    uint32_t prev_i2 = (uint32_t)-1;
+    float slope = 1.0f;
+
+    for (uint32_t r = src0_start_row; r < src0_end_row; ++r) {
+        void * d_spad = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * s_spad = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+        void * m_spad = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+
+        if (has_bias) {
+            uint32_t r_div_ne01 = fastdiv(r, div_ne01);
+            uint32_t i2 = fastmodulo(r_div_ne01, ne02, div_ne02);
+            if (i2 != prev_i2) {
+                slope = smctx->slopes[i2];
+                prev_i2 = i2;
+            }
+        }

-    FARF(HIGH, "softmax-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u : opt %u f16 %u\n", ith, nth,
-         ne00, ne01, ne02, ne03, src0_start_row, src0_end_row, ne10, ne11, ne12, ne13,
-         ne0, ne1, ne2, ne3, opt_path, smctx->use_f16);
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r);
+        compute(d_spad, s_spad, m_spad, ne00, smctx->scale, slope);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r);
+
+        dma_addr_t cur_dst = data_dst + r * dst_row_size;
+        dma_queue_push(dma_q, dma_make_data(cur_dst, d_spad),
+                       dst_row_size, smctx->dst_row_size_aligned, dst_row_size, 1);
+
+        const uint32_t next_r = r + 2;
+        if (next_r < src0_end_row) {
+            dma_addr_t next_src0 = data_src0 + next_r * src0_row_size;
+
+            uint32_t ni1 = fastmodulo(next_r, ne01, div_ne01);
+            uint32_t nr_div_ne01 = fastdiv(next_r, div_ne01);
+            uint32_t ni2 = fastmodulo(nr_div_ne01, ne02, div_ne02);
+            uint32_t ni3 = fastdiv(nr_div_ne01, div_ne02);
+            uint32_t ni12 = (ne12 == ne02) ? ni2 : fastmodulo(ni2, ne12, div_ne12);
+            uint32_t ni13 = (ne13 == ne03) ? ni3 : fastmodulo(ni3, ne13, div_ne13);
+            dma_addr_t next_src1 = data_src1 + ni1 * nb11 + ni12 * nb12 + ni13 * nb13;
+
+            dma_queue_push(dma_q, dma_make_data(s_spad, next_src0),
+                           smctx->src0_row_size_aligned, src0_row_size, src0_row_size, 1);
+            dma_queue_push(dma_q, dma_make_data(m_spad, next_src1),
+                           smctx->src1_row_size_aligned, mask_row_size, mask_row_size, 1);
+        }
+    }
+
+    dma_queue_flush(dma_q);
 }

 static int execute_op_softmax_f32(struct htp_ops_context * octx) {
-    int err = HTP_STATUS_OK;
-
     const struct htp_tensor * src0 = octx->src[0];
-    const struct htp_tensor * src1 = octx->src[1];
     const struct htp_tensor * dst  = octx->dst;

-    struct htp_softmax_context smctx;
     const char * op_type = "softmax-f32";

-    init_softmax_ctx(&smctx, octx);
+    const struct htp_softmax_kernel_params * kparams =
+        (const struct htp_softmax_kernel_params *) octx->kernel_params;
+
+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    if (kparams->kernel_id >= HTP_SOFTMAX_KERNEL_COUNT) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    if (octx->ctx->vtcm_size < (size_t) kparams->vtcm_size) {
+        FARF(ERROR, "%s : current VTCM reservation %zu is too small, needed %u\n",
+             op_type, octx->ctx->vtcm_size, kparams->vtcm_size);
+        return HTP_STATUS_VTCM_TOO_SMALL;
+    }

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const size_t elem_size = sizeof(float);
+    const size_t elem_size    = sizeof(float);
     const size_t dst_row_size = dst->nb[1];

     uint32_t row_start = 0;
@@ -357,9 +516,13 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
     if (octx->ctx->mdev.count > 1) {
         uint32_t rows_per_chunk = 0;
         htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
+            src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
         row_start = range.start;
         nrows     = range.count;
+        if (nrows < octx->n_threads) {
+            htp_ops_context_set_n_threads(octx, nrows ? nrows : 1);
+        }
     }

     if (nrows == 0) {
@@ -367,50 +530,71 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
     }

     const uint32_t n_threads = octx->n_threads;
+    uint8_t * const vtcm_base = (uint8_t *) octx->ctx->vtcm_base;

-    smctx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
-    smctx.row_start             = row_start;
-    smctx.nrows                 = nrows;
-
-    const size_t src0_row_size = src0->nb[1];
-    const size_t src1_row_size = src0_row_size;
-
-    // VTCM scratchpads for all tensors
-    // 4 rows per thread, padded to HVX vector size
-    octx->src0_spad.size_per_thread = hex_round_up(4 * src0_row_size, 128);
-    octx->src1_spad.size_per_thread = hex_round_up(4 * src1_row_size, 128);
-    octx->dst_spad.size_per_thread  = hex_round_up(4 * dst_row_size, 128);
-
-    octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
-    octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
-    octx->dst_spad.size  = octx->dst_spad.size_per_thread  * n_threads;
-
-    size_t spad_size = octx->src0_spad.size + octx->src1_spad.size + octx->dst_spad.size;
-
-    if (src1) {
-        FARF(HIGH, "%s: %ux%ux%ux%u x %ux%ux%ux%u -> %ux%ux%ux%u : src0-spad-size %u src1-spad-size %u dst-spad-size %u\n",
-             op_type, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2],
-             src1->ne[3], dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], octx->src0_spad.size, octx->src1_spad.size,
-             octx->dst_spad.size);
-    } else {
-        FARF(HIGH, "%s: %ux%ux%ux%u -> %ux%ux%ux%u : src0-spad-size %u src1-spad-size %u dst-spad-size %u\n", op_type,
-             src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-             octx->src0_spad.size, octx->src1_spad.size, octx->dst_spad.size);
-    }
+    const uint32_t off_src0 = 0;
+    const uint32_t off_dst  = off_src0 + kparams->vtcm_src0_size_per_thread * kparams->n_threads;
+    const uint32_t off_src1 = off_dst  + kparams->vtcm_dst_size_per_thread  * kparams->n_threads;

-    // Make sure the reserved vtcm size is sufficient
-    if (octx->ctx->vtcm_size < spad_size) {
-        FARF(ERROR, "%s : current VTCM reservation %zu is too small, needed %zu\n", op_type, octx->ctx->vtcm_size, spad_size);
-        return HTP_STATUS_VTCM_TOO_SMALL;
-    }
+    struct htp_softmax_context smctx = {
+        .octx                  = octx,
+        .kparams               = kparams,
+        .compute               = (void *) softmax_kernels[kparams->kernel_id],

-    octx->src0_spad.data = octx->ctx->vtcm_base;                        octx->src0_spad.src = NULL;
-    octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size; octx->src1_spad.src = NULL;
-    octx->dst_spad.data  = octx->src1_spad.data + octx->src1_spad.size; octx->dst_spad.src  = NULL;
+        .data_src0             = src0->data,
+        .data_src1             = kparams->use_src1 ? octx->src[1]->data : 0,
+        .data_dst              = dst->data,

-    work_queue_run(octx->ctx->work_queue, softmax_job_f32, &smctx, n_threads);
+        .vtcm_src0             = VTCM_LAYOUT_PTR(uint8_t, vtcm_base, off_src0),
+        .vtcm_dst              = VTCM_LAYOUT_PTR(uint8_t, vtcm_base, off_dst),
+        .vtcm_src1             = VTCM_LAYOUT_PTR_OPTIONAL(uint8_t, vtcm_base, off_src1, kparams->use_src1),

-    return err;
+        .vtcm_src0_size_per_thread = kparams->vtcm_src0_size_per_thread,
+        .vtcm_src1_size_per_thread = kparams->vtcm_src1_size_per_thread,
+        .vtcm_dst_size_per_thread  = kparams->vtcm_dst_size_per_thread,
+
+        .src0_spad_half_size   = kparams->src0_spad_half_size,
+        .src1_spad_half_size   = kparams->src1_spad_half_size,
+        .dst_spad_half_size    = kparams->dst_spad_half_size,
+
+        .src0_row_size_aligned = kparams->src0_row_size_aligned,
+        .src1_row_size_aligned = kparams->src1_row_size_aligned,
+        .dst_row_size_aligned  = kparams->dst_row_size_aligned,
+
+        .use_f16               = kparams->use_f16 != 0,
+
+        .n_head                = kparams->n_head,
+        .n_head_log2           = kparams->n_head_log2,
+
+        .scale                 = kparams->scale,
+        .max_bias              = kparams->max_bias,
+        .m0                    = kparams->m0,
+        .m1                    = kparams->m1,
+
+        .div_ne01              = kparams->div_ne01,
+        .div_ne02              = kparams->div_ne02,
+        .div_ne12              = kparams->div_ne12,
+        .div_ne13              = kparams->div_ne13,
+
+        .src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .row_start             = row_start,
+        .nrows                 = nrows,
+    };
+
+    if (kparams->max_bias > 0.0f && kparams->use_src1) {
+        if (kparams->n_head > 512) {
+            return HTP_STATUS_INVAL_PARAMS;
+        }
+        for (uint32_t h = 0; h < kparams->n_head; h += 32) {
+            HVX_Vector v_slopes = hvx_alibi_slopes(h, 1, kparams->n_head_log2, kparams->m0, kparams->m1);
+            hvx_vmem(&smctx.slopes[h]) = v_slopes;
+        }
+    }
+
+    work_queue_func_t task_func = kparams->use_src1 ? softmax_thread_mask_dma : softmax_thread_dma;
+    work_queue_run(octx->ctx->work_queue, task_func, &smctx, n_threads);
+
+    return HTP_STATUS_OK;
 }

 int op_softmax(struct htp_ops_context * octx) {
diff --git a/ggml/src/ggml-hexagon/htp/softmax-ops.h b/ggml/src/ggml-hexagon/htp/softmax-ops.h
new file mode 100644
index 000000000..8d976adb8
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/softmax-ops.h
@@ -0,0 +1,106 @@
+#ifndef HTP_SOFTMAX_OPS_H
+#define HTP_SOFTMAX_OPS_H
+
+#include <stdint.h>
+#include <stddef.h>
+#include <stdbool.h>
+#include <math.h>
+#include "hex-fastdiv.h"
+#include "hex-common.h"
+
+enum htp_softmax_kernel_id {
+    HTP_SOFTMAX_KERNEL_NOMASK = 0,
+    HTP_SOFTMAX_KERNEL_MASK_F32,
+    HTP_SOFTMAX_KERNEL_MASK_F16,
+    HTP_SOFTMAX_KERNEL_COUNT,
+};
+
+struct htp_softmax_kernel_params {
+    uint32_t n_threads;
+    uint32_t src0_nrows;
+    uint32_t src0_nrows_per_thread;
+    uint32_t vtcm_size;
+
+    uint32_t vtcm_src0_size_per_thread;
+    uint32_t vtcm_src1_size_per_thread;
+    uint32_t vtcm_dst_size_per_thread;
+
+    uint32_t src0_row_size_aligned;
+    uint32_t src1_row_size_aligned;
+    uint32_t dst_row_size_aligned;
+
+    uint32_t src0_spad_half_size;
+    uint32_t src1_spad_half_size;
+    uint32_t dst_spad_half_size;
+
+    uint32_t n_head;
+    uint32_t n_head_log2;
+    uint32_t use_src1;
+    uint32_t use_f16;
+    uint32_t kernel_id;
+
+    float    scale;
+    float    max_bias;
+    float    m0;
+    float    m1;
+
+    struct fastdiv_values div_ne01;
+    struct fastdiv_values div_ne02;
+    struct fastdiv_values div_ne12;
+    struct fastdiv_values div_ne13;
+};
+
+#if defined(__cplusplus)
+static_assert(sizeof(struct htp_softmax_kernel_params) <= 128, "htp_softmax_kernel_params is too large for kernel_params blob");
+#else
+_Static_assert(sizeof(struct htp_softmax_kernel_params) <= 128, "htp_softmax_kernel_params is too large for kernel_params blob");
+#endif
+
+struct htp_softmax_vtcm_layout {
+    size_t total_bytes;
+    size_t off_src0;
+    size_t off_dst;
+    size_t off_src1;
+
+    size_t src0_bytes_per_thread;
+    size_t dst_bytes_per_thread;
+    size_t src1_bytes_per_thread;
+
+    size_t src0_spad_half_size;
+    size_t dst_spad_half_size;
+    size_t src1_spad_half_size;
+};
+
+static inline void htp_softmax_vtcm_layout_build(
+    struct htp_softmax_vtcm_layout * layout,
+    uint32_t ne00,
+    uint32_t ne10,
+    bool use_src1,
+    bool use_f16,
+    uint32_t n_threads
+) {
+    size_t src0_row_size = ne00 * sizeof(float);
+    size_t dst_row_size  = ne00 * sizeof(float);
+    size_t src1_row_size = use_src1 ? (ne10 * (use_f16 ? 2 : 4)) : 0;
+
+    size_t src0_row_size_aligned = hex_round_up(src0_row_size, 128);
+    size_t dst_row_size_aligned  = hex_round_up(dst_row_size,  128);
+    size_t src1_row_size_aligned = use_src1 ? hex_round_up(src1_row_size, 128) : 0;
+
+    layout->src0_spad_half_size = src0_row_size_aligned;
+    layout->dst_spad_half_size  = dst_row_size_aligned;
+    layout->src1_spad_half_size = src1_row_size_aligned;
+
+    // Double buffering: 2 half-buffers per thread
+    layout->src0_bytes_per_thread = src0_row_size_aligned * 2;
+    layout->dst_bytes_per_thread  = dst_row_size_aligned  * 2;
+    layout->src1_bytes_per_thread = src1_row_size_aligned * 2;
+
+    layout->off_src0 = 0;
+    layout->off_dst  = layout->off_src0 + layout->src0_bytes_per_thread * n_threads;
+    layout->off_src1 = layout->off_dst  + layout->dst_bytes_per_thread  * n_threads;
+
+    layout->total_bytes = layout->off_src1 + layout->src1_bytes_per_thread * n_threads;
+}
+
+#endif // HTP_SOFTMAX_OPS_H
diff --git a/ggml/src/ggml-hexagon/htp/solve-tri-ops.c b/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
index 847a78712..182982fcd 100644
--- a/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
+++ b/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
@@ -218,8 +218,8 @@ int op_solve_tri(struct htp_ops_context * octx) {
         return HTP_STATUS_INVAL_PARAMS;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(src1) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     const uint32_t k = src1->ne[0];
diff --git a/ggml/src/ggml-hexagon/htp/ssm-conv.c b/ggml/src/ggml-hexagon/htp/ssm-conv.c
index bef142536..931aa406e 100644
--- a/ggml/src/ggml-hexagon/htp/ssm-conv.c
+++ b/ggml/src/ggml-hexagon/htp/ssm-conv.c
@@ -14,124 +14,22 @@
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "htp-ctx.h"
-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hex-profile.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
 #include "hvx-utils.h"
-
-#define htp_ssm_conv_tensors_preamble                           \
-    const struct htp_tensor * restrict src0 = octx->src[0];     \
-    const struct htp_tensor * restrict src1 = octx->src[1];     \
-    const struct htp_tensor * restrict dst  = octx->dst;        \
-    struct htp_spad * restrict src0_spad    = &octx->src0_spad; \
-    struct htp_spad * restrict src1_spad    = &octx->src1_spad; \
-    struct htp_spad * restrict dst_spad     = &octx->dst_spad;  \
-                                                                \
-    const uint32_t ne00 = src0->ne[0];                          \
-    const uint32_t ne01 = src0->ne[1];                          \
-    const uint32_t ne02 = src0->ne[2];                          \
-    const uint32_t ne03 = src0->ne[3];                          \
-                                                                \
-    const uint32_t ne10 = src1->ne[0];                          \
-    const uint32_t ne11 = src1->ne[1];                          \
-    const uint32_t ne12 = src1->ne[2];                          \
-    const uint32_t ne13 = src1->ne[3];                          \
-                                                                \
-    const uint32_t ne0 = dst->ne[0];                            \
-    const uint32_t ne1 = dst->ne[1];                            \
-    const uint32_t ne2 = dst->ne[2];                            \
-    const uint32_t ne3 = dst->ne[3];                            \
-                                                                \
-    const uint32_t nb00 = src0->nb[0];                          \
-    const uint32_t nb01 = src0->nb[1];                          \
-    const uint32_t nb02 = src0->nb[2];                          \
-    const uint32_t nb03 = src0->nb[3];                          \
-                                                                \
-    const uint32_t nb10 = src1->nb[0];                          \
-    const uint32_t nb11 = src1->nb[1];                          \
-    const uint32_t nb12 = src1->nb[2];                          \
-    const uint32_t nb13 = src1->nb[3];                          \
-                                                                \
-    const uint32_t nb0 = dst->nb[0];                            \
-    const uint32_t nb1 = dst->nb[1];                            \
-    const uint32_t nb2 = dst->nb[2];                            \
-    const uint32_t nb3 = dst->nb[3];
+#include "ssm-conv.h"

 struct htp_ssm_conv_context {
-    struct htp_ops_context * octx;
-    uint32_t nrows_per_thread;
-    uint32_t d_inner_tile;
-    uint64_t t_start;
-    uint32_t row_start;
-    uint32_t nrows;
+    struct htp_ops_context *                  octx;
+    const struct htp_ssm_conv_kernel_params * kparams;
+    uint32_t                                  nrows_per_thread;
+    uint32_t                                  d_inner_tile;
+    uint32_t                                  row_start;
+    uint32_t                                  nrows;
 };

-#define htp_ssm_conv_preamble                                                   \
-    struct htp_ssm_conv_context * scctx = (struct htp_ssm_conv_context *) data; \
-    struct htp_ops_context *      octx  = scctx->octx;                          \
-    htp_ssm_conv_tensors_preamble;                                              \
-    dma_queue * dma_queue = octx->ctx->dma[ith];
-
-// Scalar FP32 SSM_CONV implementation
-static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *data) {
-    htp_ssm_conv_preamble;
-
-    const uint32_t d_conv  = src1->ne[0];
-    const uint32_t d_inner = src0->ne[1];
-    const uint32_t n_t     = dst->ne[1];
-    const uint32_t n_s     = dst->ne[2];
-
-    const uint32_t src0_stride_inner = src0->nb[1] / sizeof(float); // stride for inner dimension
-    const uint32_t src0_stride_seq   = src0->nb[2] / sizeof(float); // stride for sequence dimension
-    const uint32_t src1_stride_inner = src1->nb[1] / sizeof(float); // stride for inner dimension
-    const uint32_t dst_stride_token  = dst->nb[1]  / sizeof(float); // stride for token dimension
-    const uint32_t dst_stride_seq    = dst->nb[2]  / sizeof(float); // stride for sequence dimension
-
-    const float * src0_data = (const float *) src0->data;
-    const float * src1_data = (const float *) src1->data;
-    float *       dst_data  = (float *) dst->data;
-
-    // Calculate row range for this thread
-    const uint32_t d_inner_per_thread = scctx->nrows_per_thread;
-    const uint32_t d_inner_start = scctx->row_start + d_inner_per_thread * ith;
-    const uint32_t d_inner_end   = MIN(d_inner_start + d_inner_per_thread, scctx->row_start + scctx->nrows);
-
-    // No work for this thread
-    if (d_inner_start >= d_inner_end) {
-        return;
-    }
-
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_start);
-
-    for (uint32_t i3 = 0; i3 < n_s; ++i3) {
-        for (uint32_t i2 = 0; i2 < n_t; ++i2) {
-            for (uint32_t i1 = d_inner_start; i1 < d_inner_end; ++i1) {
-                float sumf = 0.0f;
-
-                for (uint32_t i0 = 0; i0 < d_conv; ++i0) {
-                    const uint32_t src0_idx = (i2 + i0) + i1 * src0_stride_inner + i3 * src0_stride_seq;
-                    const uint32_t src1_idx = i0 + i1 * src1_stride_inner;
-
-                    sumf += src0_data[src0_idx] * src1_data[src1_idx];
-                }
-
-                const uint32_t dst_idx = i1 + i2 * dst_stride_token + i3 * dst_stride_seq;
-                dst_data[dst_idx] = sumf;
-            }
-        }
-    }
-
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_end);
-
-    FARF(HIGH, "ssm-conv-f32 %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], d_inner_start, d_inner_end,
-         src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
-         dst->ne[2], dst->ne[3]);
-}
-
-
 // In-register 32x32 fp32 transpose using std 5-stage HVX vshuff butterfly.
 static inline void hvx_transpose_32x32_f32(HVX_Vector m[32]) {
     HVX_Vector tmp[32];
@@ -181,40 +79,69 @@ static inline void hvx_transpose_32x32_f32(HVX_Vector m[32]) {
     }
 }

-// HVX FP32 SSM_CONV implementation - channel-vectorized HVX kernel with src0/src1
-// transposed into VTCM.
-//
-// VTCM layouts (per thread):
-//   src1_T : {d_inner_stride, d_conv}       - staged once per launch (small).
-//   src0_T : {d_inner_tile,     ncs}        - staged per d_inner-tile.
-//
-// d_inner_tile is chosen so that per-thread VTCM stays under the budget.
-// Each thread iterates ceil(d_inner_per_thread d_inner_tile) tiles serially.
-#define HTP_SSM_CONV_VTCM_BUDGET (1u << 20) // 1 MiB per thread
-
-// Scalar transpose: src1 {d_conv, d_inner} (DDR) -> {d_inner_stride, d_conv} (VTCM)
-static inline void transpose_src1(const float * src1_data,
-                                  uint32_t      src1_stride_inner,
-                                  uint32_t      i1_off,
-                                  uint32_t      d_inner_per_thread,
-                                  uint32_t      d_inner_stride,
-                                  uint32_t      d_conv,
-                                  float *       src1_T) {
-    for (uint32_t i = 0; i < d_inner_per_thread; ++i) {
-        const float * src_row = src1_data + (i1_off + i) * src1_stride_inner;
+// HVX deinterleave for d_conv == 4: channel-major raw VTCM -> tap-major T VTCM
+static inline void hvx_ssm_conv_unpack_to_T_4(const float * raw, float * T, uint32_t d_inner_per_thread, uint32_t d_inner_stride) {
+    for (uint32_t cb = 0; cb < d_inner_per_thread; cb += VLEN_FP32) {
+        HVX_Vector v0 = *(const HVX_Vector *)(raw + (cb + 0) * 4);
+        HVX_Vector v1 = *(const HVX_Vector *)(raw + (cb + 8) * 4);
+        HVX_Vector v2 = *(const HVX_Vector *)(raw + (cb + 16) * 4);
+        HVX_Vector v3 = *(const HVX_Vector *)(raw + (cb + 24) * 4);
+
+        HVX_VectorPair p01 = Q6_W_vdeal_VVR(v1, v0, -4);
+        HVX_VectorPair p23 = Q6_W_vdeal_VVR(v3, v2, -4);
+
+        HVX_VectorPair p_w02 = Q6_W_vdeal_VVR(Q6_V_lo_W(p23), Q6_V_lo_W(p01), -4);
+        HVX_VectorPair p_w13 = Q6_W_vdeal_VVR(Q6_V_hi_W(p23), Q6_V_hi_W(p01), -4);
+
+        *(HVX_Vector *)(T + 0 * d_inner_stride + cb) = Q6_V_lo_W(p_w02);
+        *(HVX_Vector *)(T + 1 * d_inner_stride + cb) = Q6_V_lo_W(p_w13);
+        *(HVX_Vector *)(T + 2 * d_inner_stride + cb) = Q6_V_hi_W(p_w02);
+        *(HVX_Vector *)(T + 3 * d_inner_stride + cb) = Q6_V_hi_W(p_w13);
+    }
+}
+
+// HVX transpose for general d_conv <= 32: channel-major raw VTCM -> tap-major T VTCM
+static inline void hvx_ssm_conv_unpack_to_T_gen(const float * raw, float * T, uint32_t d_inner_per_thread, uint32_t d_inner_stride, uint32_t d_conv) {
+    uint32_t __attribute__((aligned(VLEN))) mask_buf[VLEN_FP32] = { 0 };
+    for (uint32_t j = 0; j < d_conv; ++j) {
+        mask_buf[j] = 0xFFFFFFFF;
+    }
+    const HVX_Vector mask = *(const HVX_Vector *) mask_buf;
+
+    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 sub[32];
+        for (uint32_t r = 0; r < cb_n; ++r) {
+            const float * ch_ptr = raw + (cb + r) * d_conv;
+            sub[r] = Q6_V_vand_VV(*(const HVX_UVector *) ch_ptr, mask);
+        }
+        for (uint32_t r = cb_n; r < 32; ++r) {
+            sub[r] = hvx_vec_splat_f32(0.0f);
+        }
+
+        hvx_transpose_32x32_f32(sub);
+
         for (uint32_t j = 0; j < d_conv; ++j) {
-            src1_T[j * d_inner_stride + i] = src_row[j];
+            *(HVX_Vector *)(T + j * d_inner_stride + cb) = sub[j];
         }
     }
 }

-// HVX 32x32 src0 transpose: src0 {ncs, d_inner} (DDR) -> src0_T {d_inner_tile, ncs} (VTCM)
+static inline void hvx_ssm_conv_unpack_to_T(const float * raw, float * T, uint32_t d_inner_per_thread, uint32_t d_inner_stride, uint32_t d_conv) {
+    if (d_conv == 4 && (d_inner_per_thread % VLEN_FP32 == 0)) {
+        hvx_ssm_conv_unpack_to_T_4(raw, T, d_inner_per_thread, d_inner_stride);
+    } else {
+        hvx_ssm_conv_unpack_to_T_gen(raw, T, d_inner_per_thread, d_inner_stride, d_conv);
+    }
+}
+
+// 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 /* dst column offset */) {
+                                        uint32_t      cb) {
     const uint32_t T_TILE = VLEN_FP32;

     HVX_Vector __attribute__((aligned(VLEN))) sub[32];
@@ -222,20 +149,15 @@ static inline void transpose_src0_block(const float * src0_block,
     for (uint32_t t0 = 0; t0 < ncs; t0 += T_TILE) {
         const uint32_t t_n = MIN(T_TILE, ncs - t0);

-        // Load 32 rows (channels) of T_TILE samples; pad missing channels with zeros.
+        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;
+
         for (uint32_t r = 0; r < cb_n; ++r) {
             const float * src_row = src0_block + r * ncs + t0;
-            if (t_n == T_TILE) {
-                sub[r] = *(const HVX_UVector *) src_row;
-            } else {
-                HVX_Vector v = hvx_vec_splat_f32(0.0f);
-                hvx_vec_store_u(&v, t_n * sizeof(float), hvx_vec_splat_f32(0.0f));
-
-                float __attribute__((aligned(VLEN))) tmp[VLEN_FP32] = { 0 };
-                for (uint32_t k = 0; k < t_n; ++k) tmp[k] = src_row[k];
-                v = *(const HVX_Vector *) tmp;
-                sub[r] = v;
-            }
+            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);
@@ -243,8 +165,6 @@ static inline void transpose_src0_block(const float * src0_block,

         hvx_transpose_32x32_f32(sub);

-        // Store transposed sub-tile to src0_T at offsets (t0 + j) * d_inner_tile + cb.
-        // Only write the valid t_n rows of the transposed result.
         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) {
@@ -256,20 +176,21 @@ static inline void transpose_src0_block(const float * src0_block,
     }
 }

-static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void *data) {
-    htp_ssm_conv_preamble;
+// Single-row 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;
+    const struct htp_ssm_conv_kernel_params * kparams = scctx->kparams;

-    const uint32_t d_conv  = src1->ne[0];
-    const uint32_t d_inner = src0->ne[1];
-    const uint32_t n_t     = dst->ne[1];
-    const uint32_t n_s     = dst->ne[2];
-    const uint32_t ncs     = src0->ne[0];
+    const struct htp_tensor * restrict src0 = octx->src[0];
+    const struct htp_tensor * restrict src1 = octx->src[1];
+    const struct htp_tensor * restrict dst  = octx->dst;

-    const uint32_t src0_stride_inner = src0->nb[1] / sizeof(float);
-    const uint32_t src0_stride_seq   = src0->nb[2] / sizeof(float);
-    const uint32_t src1_stride_inner = src1->nb[1] / sizeof(float);
-    const uint32_t dst_stride_token  = dst->nb[1]  / sizeof(float);
-    const uint32_t dst_stride_seq    = dst->nb[2]  / sizeof(float);
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    const uint32_t d_conv  = kparams->d_conv;
+    const uint32_t d_inner = kparams->d_inner;
+    const uint32_t n_s     = kparams->n_s;

     const uint32_t dr  = scctx->nrows_per_thread;
     const uint32_t ir0 = scctx->row_start + dr * ith;
@@ -279,23 +200,141 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
         return;
     }

+    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 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;
+    uint8_t * dst_spad_base  = octx->dst_spad.data  + ith * octx->dst_spad.size_per_thread;
+
+    const size_t weight_bytes    = (size_t) d_inner_per_thread * d_conv * sizeof(float);
+    const size_t weight_raw_size = hex_round_up(weight_bytes, 128);
+
+    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);
+
+    float * dst_spad = (float *) dst_spad_base;
+
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    // 1. Fetch weights src1 from DDR into VTCM via DMA (DMA64-safe)
+    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);
+            }
+        }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) s);
+
+        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);
+    }
+
+    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],
+         dst->ne[2], dst->ne[3]);
+}
+
+// Multi-token prefill worker (n_t > 1)
+static void ssm_conv_thread_f32_prefill(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;
+    const struct htp_ssm_conv_kernel_params * kparams = scctx->kparams;
+
+    const struct htp_tensor * restrict src0 = octx->src[0];
+    const struct htp_tensor * restrict src1 = octx->src[1];
+    const struct htp_tensor * restrict dst  = octx->dst;
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    const uint32_t d_conv  = kparams->d_conv;
+    const uint32_t d_inner = kparams->d_inner;
+    const uint32_t n_t     = kparams->n_t;
+    const uint32_t n_s     = kparams->n_s;
+    const uint32_t ncs     = src0->ne[0];
+
+    const uint32_t dr  = scctx->nrows_per_thread;
+    const uint32_t ir0 = scctx->row_start + dr * ith;
+    const uint32_t ir1 = MIN(ir0 + dr, scctx->row_start + scctx->nrows);
+
+    if (ir0 >= ir1) {
+        return;
+    }

     const uint32_t d_inner_per_thread = ir1 - ir0;
-    const uint32_t d_inner_stride     = scctx->nrows_per_thread;
+    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 float * src0_data = (const float *) src0->data;
-    const float * src1_data = (const float *) src1->data;
-    float       * dst_data  = (float       *) dst->data;
+    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_token_bytes  = dst->nb[1];
+    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;
+    uint8_t * dst_spad_base  = octx->dst_spad.data  + ith * octx->dst_spad.size_per_thread;

-    // Per-thread VTCM regions.
-    float * src0_T = (float *)(octx->src0_spad.data + ith * octx->src0_spad.size_per_thread);
-    float * src1_T = (float *)(octx->src1_spad.data + ith * octx->src1_spad.size_per_thread);
+    const size_t weight_bytes    = (size_t) d_inner_per_thread * d_conv * sizeof(float);
+    const size_t weight_raw_size = hex_round_up(weight_bytes, 128);

-    // Stage src1 weights once into VTCM in {d_inner_stride, d_conv} layout.
-    transpose_src1(src1_data, src1_stride_inner, ir0, d_inner_per_thread, d_inner_stride, d_conv, src1_T);
+    float * src1_raw = (float *) src1_spad_base;
+    float * src1_T   = (float *) (src1_spad_base + weight_raw_size);
+
+    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);
+
+    float * dst_tile = (float *) dst_spad_base;
+
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    // 1. Fetch weights src1 from DDR into VTCM via DMA (DMA64-safe)
+    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 uint32_t C_TILE = VLEN_FP32;

@@ -303,14 +342,24 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
         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);

-            // Place src0 chunk into VTCM in {d_inner_tile, ncs} layout.
-            const float * src0_block = src0_data + i3 * src0_stride_seq + (ir0 + tile_off) * src0_stride_inner;
+            // 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);

+            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);
+
+            // 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_block + cb * src0_stride_inner, ncs, cb_n, d_inner_tile, src0_T, cb);
+                transpose_src0_block(src0_tile_raw + cb * ncs, ncs, cb_n, d_inner_tile, src0_T, cb);
             }

+            // 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);
@@ -323,21 +372,29 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
                     }

                     HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
-
-                    float * dst_ptr = dst_data + (ir0 + tile_off + cb) + t * dst_stride_token + i3 * dst_stride_seq;
+                    float * dst_tile_ptr = dst_tile + t * tile_n + cb;
                     if (cb_n == C_TILE) {
-                        *(HVX_UVector *) dst_ptr = y;
+                        *(HVX_Vector *) dst_tile_ptr = y;
                     } else {
-                        hvx_vec_store_u(dst_ptr, cb_n * sizeof(float), y);
+                        hvx_vec_store_u(dst_tile_ptr, cb_n * sizeof(float), y);
                     }
                 }
             }
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) tile_off);
+
+            // 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);
+
+            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);
         }
     }

-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir1);
-
-    FARF(HIGH, "ssm-conv-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
+    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],
          dst->ne[2], dst->ne[3]);
@@ -352,21 +409,25 @@ int op_ssm_conv_f32(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    const uint32_t d_conv  = src1->ne[0];
-    const uint32_t d_inner = src0->ne[1];
-    const uint32_t n_t     = dst->ne[1];  // tokens per sequence
-    const uint32_t n_s     = dst->ne[2];  // number of sequences in the batch
+    const struct htp_ssm_conv_kernel_params * kparams = (const struct htp_ssm_conv_kernel_params *) octx->kernel_params;

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+
+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
     }

     uint32_t row_start = 0;
-    uint32_t nrows     = d_inner;
+    uint32_t nrows     = kparams->d_inner;

     if (octx->ctx->mdev.count > 1) {
         const uint32_t elems_per_chunk = VLEN_FP32;
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(d_inner, htp_tensor_mdev_data_aligned(dst) ? elems_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
+            kparams->d_inner,
+            htp_tensor_mdev_data_aligned(dst) ? elems_per_chunk : 0,
+            octx->ctx->mdev.idx,
+            octx->ctx->mdev.count,
+            &octx->ctx->mdev.count_div
+        );
         row_start = range.start;
         nrows     = range.count;
     }
@@ -375,64 +436,49 @@ int op_ssm_conv_f32(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

-    const uint32_t n_threads = octx->n_threads;
-
-    struct htp_ssm_conv_context scctx = { 0 };
-    scctx.octx      = octx;
-    scctx.row_start = row_start;
-    scctx.nrows     = nrows;
-
-    uint32_t use_hvx = 0;
-    if (nrows >= VLEN_FP32 && n_t >= VLEN_FP32) {
-        use_hvx = 1;
+    if (kparams->vtcm_size > octx->ctx->vtcm_size) {
+        return HTP_STATUS_VTCM_TOO_SMALL;
     }

-    const uint32_t raw_rpt = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
-    scctx.nrows_per_thread = hex_round_up(raw_rpt, VLEN_FP32);
-
-    const uint32_t d_inner_per_thread = scctx.nrows_per_thread;
-    const uint32_t ncs                = src0->ne[0];
-
-    const uint32_t src1_T_size = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 256);
-    const uint32_t src0_T_max = HTP_SSM_CONV_VTCM_BUDGET > src1_T_size ? HTP_SSM_CONV_VTCM_BUDGET - src1_T_size : 0;
-
-    uint32_t d_inner_tile = (src0_T_max / sizeof(float)) / ncs;
-    d_inner_tile -= (d_inner_tile % VLEN_FP32);
-    if (d_inner_tile == 0) {
-        FARF(HIGH, "ssm_conv-f32: inner tile rounds to 0 (ncs=%u), falling back to scalar\n", ncs);
-        use_hvx = 0;
-    } else {
-        scctx.d_inner_tile = d_inner_tile;
-
-        octx->src0_spad.size_per_thread = hex_round_up(d_inner_tile * ncs * sizeof(float), 256);
-        octx->src1_spad.size_per_thread = src1_T_size;
-        octx->dst_spad.size_per_thread  = 0;
-
-        octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
-        octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
-        octx->dst_spad.size  = 0;
-
-        octx->src0_spad.data = octx->ctx->vtcm_base;
-        octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
-        octx->src0_spad.src  = NULL;
-        octx->src1_spad.src  = NULL;
-
-        const size_t total_spad = octx->src0_spad.size + octx->src1_spad.size;
-        if (total_spad > octx->ctx->vtcm_size) {
-            FARF(HIGH, "ssm_conv-f32: scratchpad %zu exceeds VTCM %zu, falling back to scalar\n",
-                 total_spad, octx->ctx->vtcm_size);
-            use_hvx = 0;
-        }
-    }
-
-    FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_hvx %d\n", src0->ne[0],
-         src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
-         dst->ne[1], dst->ne[2], dst->ne[3], use_hvx);
+    const uint32_t n_threads = octx->n_threads;

-    if (use_hvx) {
-        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32_hvx, &scctx, n_threads);
+    octx->src0_spad.size_per_thread = kparams->vtcm_src0_size_per_thread;
+    octx->src1_spad.size_per_thread = kparams->vtcm_src1_size_per_thread;
+    octx->dst_spad.size_per_thread  = kparams->vtcm_dst_size_per_thread;
+
+    octx->src0_spad.size = kparams->vtcm_src0_size;
+    octx->src1_spad.size = kparams->vtcm_src1_size;
+    octx->dst_spad.size  = kparams->vtcm_dst_size;
+
+    octx->src0_spad.data = octx->ctx->vtcm_base;
+    octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
+    octx->dst_spad.data  = octx->src1_spad.data + octx->src1_spad.size;
+    octx->src0_spad.src  = NULL;
+    octx->src1_spad.src  = NULL;
+    octx->dst_spad.src   = NULL;
+
+    const uint32_t raw_rpt            = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+    const uint32_t d_inner_per_thread = hex_round_up(raw_rpt, VLEN_FP32);
+
+    struct htp_ssm_conv_context scctx = {
+        .octx             = octx,
+        .kparams          = kparams,
+        .nrows_per_thread = d_inner_per_thread,
+        .d_inner_tile     = kparams->d_inner_tile,
+        .row_start        = row_start,
+        .nrows            = nrows,
+    };
+
+    FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : mode %s\n",
+         src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3],
+         src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3],
+         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
+         kparams->n_t == 1 ? "decode" : "prefill");
+
+    if (kparams->n_t == 1) {
+        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_decode, &scctx, n_threads);
     } else {
-        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32, &scctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_prefill, &scctx, n_threads);
     }

     return HTP_STATUS_OK;
@@ -441,16 +487,10 @@ int op_ssm_conv_f32(struct htp_ops_context * octx) {
 int op_ssm_conv(struct htp_ops_context * octx) {
     const struct htp_tensor * dst = octx->dst;

-    int err = HTP_STATUS_OK;
-
     switch (dst->type) {
         case HTP_TYPE_F32:
-            err = op_ssm_conv_f32(octx);
-            break;
+            return op_ssm_conv_f32(octx);
         default:
-            err = HTP_STATUS_NO_SUPPORT;
-            break;
+            return HTP_STATUS_NO_SUPPORT;
     }
-
-    return err;
 }
diff --git a/ggml/src/ggml-hexagon/htp/ssm-conv.h b/ggml/src/ggml-hexagon/htp/ssm-conv.h
new file mode 100644
index 000000000..be62d7bf5
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/ssm-conv.h
@@ -0,0 +1,40 @@
+#ifndef HTP_SSM_CONV_H
+#define HTP_SSM_CONV_H
+
+#include <stdint.h>
+
+#include "hex-fastdiv.h"
+#include "htp-ops.h"
+
+struct htp_ssm_conv_kernel_params {
+    uint32_t n_threads;
+    uint32_t d_conv;
+    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;
+
+    uint32_t vtcm_src0_size;
+    uint32_t vtcm_src1_size;
+    uint32_t vtcm_dst_size;
+    uint32_t vtcm_size;
+
+    struct fastdiv_values div_n_threads;
+};
+
+#if defined(__cplusplus)
+static_assert(sizeof(struct htp_ssm_conv_kernel_params) <= 128, "htp_ssm_conv_kernel_params is too large for kernel_params blob");
+#else
+_Static_assert(sizeof(struct htp_ssm_conv_kernel_params) <= 128, "htp_ssm_conv_kernel_params is too large for kernel_params blob");
+#endif
+
+#endif // HTP_SSM_CONV_H
diff --git a/ggml/src/ggml-hexagon/htp/sum-rows-ops.c b/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
index faf716b4b..9b8e04a0f 100644
--- a/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
@@ -8,7 +8,7 @@
 #include <string.h>
 #include <math.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hvx-utils.h"

 #define GGML_COMMON_DECL_C
@@ -106,8 +106,8 @@ int op_sum_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
-        return HTP_STATUS_OK;
+    if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
+        return HTP_STATUS_NO_SUPPORT;
     }

     const uint32_t src0_nrows      = ne01 * ne02 * ne03;
diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c
index cb82bfa3c..9a1479e29 100644
--- a/ggml/src/ggml-hexagon/htp/unary-ops.c
+++ b/ggml/src/ggml-hexagon/htp/unary-ops.c
@@ -8,7 +8,7 @@
 #include <math.h>
 #include <string.h>

-#include "hex-dma.h"
+#include "dma-queue.h"
 #include "hex-fastdiv.h"
 #include "hvx-exp.h"
 #include "hvx-sigmoid.h"
@@ -23,13 +23,47 @@
 #include "htp-vtcm.h"
 #include "hex-profile.h"

+struct htp_unary_context;
+
+typedef void (*unary_compute_fn_t)(const void * restrict src,
+                                   void * restrict dst,
+                                   uint32_t num_rows,
+                                   const struct htp_unary_context * uctx);
+
+typedef void (*unary_rms_norm_mul_compute_fn_t)(const void * restrict src,
+                                                const void * restrict weight,
+                                                void * restrict dst,
+                                                uint32_t num_rows,
+                                                const struct htp_unary_context * uctx);
+
+typedef void (*unary_tri_compute_fn_t)(const void * restrict src,
+                                       void * restrict dst,
+                                       uint32_t num_rows,
+                                       uint32_t ir,
+                                       const struct htp_unary_context * uctx);
+
+typedef void (*unary_tile_compute_fn_t)(void * restrict dst,
+                                        const void * restrict src,
+                                        uint32_t tw,
+                                        const struct htp_unary_context * uctx);
+
+typedef void (*unary_tiled_tri_compute_fn_t)(const void * restrict src,
+                                             void * restrict dst,
+                                             uint32_t tile_elems,
+                                             uint32_t col_start,
+                                             uint32_t i01,
+                                             uint32_t ne0,
+                                             int32_t ttype);
+
 struct htp_unary_context {
     struct htp_ops_context * octx;
     const struct htp_unary_kernel_params * kparams;

-    const uint8_t *           data_src0;
-    const uint8_t *           data_src1;            // weight/scale tensor for RMS_NORM_MUL
-    uint8_t *                 data_dst;
+    void *                    compute;
+
+    dma_addr_t                data_src0;
+    dma_addr_t                data_src1;            // weight/scale tensor for RMS_NORM_MUL
+    dma_addr_t                data_dst;

     size_t                    src0_data_row_size;   // actual data bytes per row
     size_t                    src1_data_row_size;
@@ -121,8 +155,8 @@ static inline uint32_t unary_block_size(uint32_t ir,
     const size_t src0_row_size_aligned = uctx->src0_row_size_aligned; \
     const size_t dst_row_size_aligned = uctx->dst_row_size_aligned;

-static void scale_f32(const float * restrict src,
-                      float * restrict dst,
+static void scale_f32(const void * restrict src,
+                      void * restrict dst,
                       const uint32_t num_rows,
                       const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -139,8 +173,8 @@ static void scale_f32(const float * restrict src,
     }
 }

-static void clamp_f32(const float * restrict src,
-                      float * restrict dst,
+static void clamp_f32(const void * restrict src,
+                      void * restrict dst,
                       const uint32_t num_rows,
                       const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -157,8 +191,8 @@ static void clamp_f32(const float * restrict src,
     }
 }

-static void leaky_relu_f32(const float * restrict src,
-                           float * restrict dst,
+static void leaky_relu_f32(const void * restrict src,
+                           void * restrict dst,
                            const uint32_t num_rows,
                            const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -173,8 +207,8 @@ static void leaky_relu_f32(const float * restrict src,
     }
 }

-static void rms_norm_f32(const float * restrict src,
-                         float * restrict dst,
+static void rms_norm_f32(const void * restrict src,
+                         void * restrict dst,
                          const uint32_t num_rows,
                          const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -189,9 +223,9 @@ static void rms_norm_f32(const float * restrict src,
     }
 }

-static void rms_norm_mul_f32(const float * restrict src,
-                             const float * restrict weight,
-                             float * restrict dst,
+static void rms_norm_mul_f32(const void * restrict src,
+                             const void * restrict weight,
+                             void * restrict dst,
                              const uint32_t num_rows,
                              const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -207,8 +241,8 @@ static void rms_norm_mul_f32(const float * restrict src,
     }
 }

-static void norm_f32(const float * restrict src,
-                     float * restrict dst,
+static void norm_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -223,8 +257,8 @@ static void norm_f32(const float * restrict src,
     }
 }

-static void sqr_f32(const float * restrict src,
-                    float * restrict dst,
+static void sqr_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -237,8 +271,8 @@ static void sqr_f32(const float * restrict src,
     }
 }

-static void sqrt_f32(const float * restrict src,
-                     float * restrict dst,
+static void sqrt_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -251,8 +285,8 @@ static void sqrt_f32(const float * restrict src,
     }
 }

-static void scale_f16(const _Float16 * restrict src,
-                      _Float16 * restrict dst,
+static void scale_f16(const void * restrict src,
+                      void * restrict dst,
                       const uint32_t num_rows,
                       const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -269,8 +303,8 @@ static void scale_f16(const _Float16 * restrict src,
     }
 }

-static void clamp_f16(const _Float16 * restrict src,
-                      _Float16 * restrict dst,
+static void clamp_f16(const void * restrict src,
+                      void * restrict dst,
                       const uint32_t num_rows,
                       const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -287,8 +321,8 @@ static void clamp_f16(const _Float16 * restrict src,
     }
 }

-static void rms_norm_f16(const _Float16 * restrict src,
-                         _Float16 * restrict dst,
+static void rms_norm_f16(const void * restrict src,
+                         void * restrict dst,
                          const uint32_t num_rows,
                          const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -303,8 +337,8 @@ static void rms_norm_f16(const _Float16 * restrict src,
     }
 }

-static void norm_f16(const _Float16 * restrict src,
-                     _Float16 * restrict dst,
+static void norm_f16(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -319,8 +353,8 @@ static void norm_f16(const _Float16 * restrict src,
     }
 }

-static void sqr_f16(const _Float16 * restrict src,
-                    _Float16 * restrict dst,
+static void sqr_f16(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -333,8 +367,8 @@ static void sqr_f16(const _Float16 * restrict src,
     }
 }

-static void sqrt_f16(const _Float16 * restrict src,
-                     _Float16 * restrict dst,
+static void sqrt_f16(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -347,8 +381,8 @@ static void sqrt_f16(const _Float16 * restrict src,
     }
 }

-static void abs_f16(const _Float16 * restrict src,
-                    _Float16 * restrict dst,
+static void abs_f16(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -361,8 +395,8 @@ static void abs_f16(const _Float16 * restrict src,
     }
 }

-static void log_f16(const _Float16 * restrict src,
-                    _Float16 * restrict dst,
+static void log_f16(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -375,8 +409,8 @@ static void log_f16(const _Float16 * restrict src,
     }
 }

-static void l2_norm_f16(const _Float16 * restrict src,
-                        _Float16 * restrict dst,
+static void l2_norm_f16(const void * restrict src,
+                        void * restrict dst,
                         const uint32_t num_rows,
                         const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -391,8 +425,8 @@ static void l2_norm_f16(const _Float16 * restrict src,
     }
 }

-static void neg_f32(const float * restrict src,
-                    float * restrict dst,
+static void neg_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -405,8 +439,8 @@ static void neg_f32(const float * restrict src,
     }
 }

-static void exp_f32(const float * restrict src,
-                    float * restrict dst,
+static void exp_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -419,8 +453,8 @@ static void exp_f32(const float * restrict src,
     }
 }

-static void sigmoid_f32(const float * restrict src,
-                        float * restrict dst,
+static void sigmoid_f32(const void * restrict src,
+                        void * restrict dst,
                         const uint32_t num_rows,
                         const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -434,8 +468,8 @@ static void sigmoid_f32(const float * restrict src,
 }

 // silu(x) = x * sigmoid(x)
-static void silu_f32(const float * restrict src,
-                     float * restrict dst,
+static void silu_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -450,8 +484,8 @@ static void silu_f32(const float * restrict src,
 }

 // gelu(x) = x * sigmoid(1.702 * x)  (quick/sigmoid approximation, matches CPU GELU_QUICK reference)
-static void gelu_f32(const float * restrict src,
-                     float * restrict dst,
+static void gelu_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -466,8 +500,8 @@ static void gelu_f32(const float * restrict src,
     }
 }

-static void tri_f32(const float * restrict src,
-                    float * restrict dst,
+static void tri_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const uint32_t ir,
                     const struct htp_unary_context * uctx) {
@@ -551,8 +585,8 @@ static void tri_f32(const float * restrict src,
     }
 }

-static void softplus_f32(const float * restrict src,
-                         float * restrict dst,
+static void softplus_f32(const void * restrict src,
+                         void * restrict dst,
                          const uint32_t num_rows,
                          const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -570,8 +604,8 @@ static void softplus_f32(const float * restrict src,
     }
 }

-static void l2_norm_f32(const float * restrict src,
-                        float * restrict dst,
+static void l2_norm_f32(const void * restrict src,
+                        void * restrict dst,
                         const uint32_t num_rows,
                         const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -579,15 +613,15 @@ static void l2_norm_f32(const float * restrict src,
     memcpy(&epsilon, op_params, sizeof(float));

     for (uint32_t ir = 0; ir < num_rows; ir++) {
-        const float * restrict src_f = (const float *)((const uint8_t *)src + (ir * src0_row_size_aligned));
-        float * restrict dst_f       = (float *)((uint8_t *)dst + (ir * dst_row_size_aligned));
+        const uint8_t * restrict src_f = (const uint8_t *)src + (ir * src0_row_size_aligned);
+        uint8_t * restrict dst_f       = (uint8_t *)dst + (ir * dst_row_size_aligned);

         hvx_fast_l2_norm_f32((const uint8_t *)src_f, (uint8_t *)dst_f, ne0, epsilon);
     }
 }

-static void tanh_f32(const float * restrict src,
-                     float * restrict dst,
+static void tanh_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -600,8 +634,8 @@ static void tanh_f32(const float * restrict src,
     }
 }

-static void abs_f32(const float * restrict src,
-                    float * restrict dst,
+static void abs_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -614,8 +648,8 @@ static void abs_f32(const float * restrict src,
     }
 }

-static void relu_f32(const float * restrict src,
-                     float * restrict dst,
+static void relu_f32(const void * restrict src,
+                     void * restrict dst,
                      const uint32_t num_rows,
                      const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -628,8 +662,8 @@ static void relu_f32(const float * restrict src,
     }
 }

-static void log_f32(const float * restrict src,
-                    float * restrict dst,
+static void log_f32(const void * restrict src,
+                    void * restrict dst,
                     const uint32_t num_rows,
                     const struct htp_unary_context * uctx) {
     htp_unary_op_preamble;
@@ -642,369 +676,100 @@ static void log_f32(const float * restrict src,
     }
 }

-#define DEFINE_UNARY_TASK_IMPL(NAME, TYPE, SUFFIX, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR)                              \
-static void unary_task_##SUFFIX##_##NAME(unsigned int nth, unsigned int ith, void * data) {                         \
-    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;                                \
-    struct htp_ops_context * octx = uctx->octx;                                                                     \
-    const struct htp_tensor * src = octx->src[0];                                                                   \
-    const struct htp_tensor * dst = octx->dst;                                                                      \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                          \
-                                                                                                                    \
-    htp_unary_preamble;                                                                                             \
-                                                                                                                    \
-    int32_t *    op_params = octx->op_params;                                                                       \
-    uint32_t     src0_nrows_per_thread = uctx->src0_nrows_per_thread;                                               \
-                                                                                                                    \
-    const size_t src0_data_row_size = uctx->src0_data_row_size;                                                     \
-    const size_t dst_data_row_size  = uctx->dst_data_row_size;                                                      \
-                                                                                                                    \
-    const size_t src0_row_size_aligned = uctx->src0_row_size_aligned;                                               \
-    const size_t dst_row_size_aligned  = uctx->dst_row_size_aligned;                                                \
-                                                                                                                    \
-    const uint32_t src0_nrows = uctx->src0_nrows;                                                                   \
-    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;                                  \
-    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);      \
-                                                                                                                    \
-    if (src0_start_row >= src0_end_row) {                                                                           \
-        return;                                                                                                     \
-    }                                                                                                               \
-                                                                                                                    \
-    const uint8_t * restrict data_src = uctx->data_src0;                                                            \
-    const uint8_t * restrict data_src1 = uctx->data_src1;                                                           \
-    uint8_t * restrict       data_dst = uctx->data_dst;                                                             \
-                                                                                                                    \
-    const struct htp_tensor * src1 = (IS_RMS_NORM_MUL) ? octx->src[1] : NULL;                                       \
-    const uint32_t nb11 = src1 ? src1->nb[1] : 0;                                                                   \
-    const uint32_t nb12 = src1 ? src1->nb[2] : 0;                                                                   \
-    const uint32_t nb13 = src1 ? src1->nb[3] : 0;                                                                   \
-    const uint32_t nb11_bc = (src1 && src1->ne[1] > 1) ? nb11 : 0;                                                  \
-    const uint32_t nb12_bc = (src1 && src1->ne[2] > 1) ? nb12 : 0;                                                  \
-    const uint32_t nb13_bc = (src1 && src1->ne[3] > 1) ? nb13 : 0;                                                  \
-    const bool src1_contig = src1 ? ((nb12 == (size_t)ne01 * nb11) && (nb13 == (size_t)ne02 * nb12)) : false;       \
-                                                                                                                    \
-    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);                           \
-    uint8_t * src1_vtcm_data = uctx->vtcm_src1 ? (uctx->vtcm_src1 + (ith * uctx->vtcm_src1_size_per_thread)) : NULL;\
-    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);                             \
-                                                                                                                    \
-    size_t src0_vtcm_half_size = uctx->src0_vtcm_half_size;                                                         \
-    size_t src1_vtcm_half_size = uctx->src1_vtcm_half_size;                                                         \
-    size_t dst_vtcm_half_size  = uctx->dst_vtcm_half_size;                                                          \
-                                                                                                                    \
-    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&                                                       \
-                             (nb03 == (size_t)ne02 * nb02);                                                         \
-    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&                                                       \
-                             (nb3  == (size_t)ne2  * nb2);                                                          \
-                                                                                                                    \
-    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;                                             \
-    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;                                             \
-    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;                                            \
-                                                                                                                    \
-    const bool src1_needs_row_clip = (IS_RMS_NORM_MUL) && !uctx->broadcast_weight && !src1_contig;                  \
-    const bool block_src0_contig = src0_contig && !src1_needs_row_clip;                                             \
-    const bool block_dst_contig  = dst_contig  && !src1_needs_row_clip;                                             \
-                                                                                                                    \
-    const uint32_t src0_max_block = block_src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01);             \
-    const uint32_t dst_max_block  = block_dst_contig  ? uctx->block : MIN((uint32_t)uctx->block, ne1);              \
-    const uint32_t BLOCK = MIN(src0_max_block, dst_max_block);                                                      \
-    if (BLOCK == 0) {                                                                                               \
-        FARF(ERROR, "unary-" #SUFFIX " : current VTCM reservation %zu is too small, needed at least %zu\n",         \
-             uctx->vtcm_src0_size_per_thread, src0_row_size_aligned);                                               \
-        return;                                                                                                     \
-    }                                                                                                               \
-                                                                                                                    \
-    dma_queue * dma_queue = octx->ctx->dma[ith];                                                                    \
-                                                                                                                    \
-    if ((IS_RMS_NORM_MUL) && uctx->broadcast_weight) {                                                              \
-        dma_queue_push(dma_queue, dma_make_ptr(src1_vtcm_data, data_src1),                                          \
-                       uctx->src1_row_size_aligned, 0, uctx->src1_data_row_size, 1);                                \
-        dma_queue_flush(dma_queue);                                                                                 \
-    }                                                                                                               \
-                                                                                                                    \
-    for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) {               \
-        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig,  \
-                                                     ne01, div_ne01);                                               \
-                                                                                                                    \
-        dma_queue_push(dma_queue,                                                                                   \
-            dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)),                                \
-            nb1, dst_row_size_aligned, dst_data_row_size, 0);                                                       \
-                                                                                                                    \
-        const size_t src0_off = src0_contig ? (ir * nb01) :                                                         \
-            unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);                      \
-        dma_queue_push(dma_queue,                                                                                   \
-            dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_vtcm_half_size), data_src + src0_off),                   \
-            src0_row_size_aligned, nb01, src0_data_row_size, block_size);                                           \
-                                                                                                                    \
-        if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) {                                                         \
-            const size_t src1_off = src1_contig ? (ir * nb11) :                                                     \
-                unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, nb13_bc);         \
-            dma_queue_push(dma_queue,                                                                               \
-                dma_make_ptr(src1_vtcm_data + (vtcm_idx * src1_vtcm_half_size), data_src1 + src1_off),              \
-                uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, block_size);                           \
-        }                                                                                                           \
-                                                                                                                    \
-        ir += block_size;                                                                                           \
-    }                                                                                                               \
-                                                                                                                    \
-    for (uint32_t ir = src0_start_row; ir < src0_end_row; ) {                                                       \
-        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig,  \
-                                                     ne01, div_ne01);                                               \
-                                                                                                                    \
-        TYPE * dst_vtcm  = (TYPE *) dma_queue_pop(dma_queue).src;                                                   \
-        TYPE * src0_vtcm = (TYPE *) dma_queue_pop(dma_queue).dst;                                                   \
-        TYPE * src1_vtcm = NULL;                                                                                    \
-        if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) {                                                         \
-            src1_vtcm = (TYPE *) dma_queue_pop(dma_queue).dst;                                                      \
-        }                                                                                                           \
-                                                                                                                    \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                      \
-        CORE_EXPR;                                                                                                  \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                       \
-                                                                                                                    \
-        const size_t dst_off = dst_contig ? (ir * nb1) :                                                            \
-            unary_row_offset(ir, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3);                           \
-        dma_queue_push(dma_queue,                                                                                   \
-            dma_make_ptr(data_dst + dst_off, dst_vtcm),                                                             \
-            nb1, dst_row_size_aligned, dst_data_row_size, block_size);                                              \
-                                                                                                                    \
-        const uint32_t next_ir = ir + block_size;                                                                   \
-        if (next_ir < src0_end_row) {                                                                               \
-            const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, block_src0_contig,      \
-                                                              block_dst_contig, ne01, div_ne01);                    \
-            const uint32_t pref_ir = next_ir + next_block_size;                                                     \
-            if (pref_ir < src0_end_row) {                                                                           \
-                const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, block_src0_contig,  \
-                                                                  block_dst_contig, ne01, div_ne01);                \
-                const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) :                                       \
-                    unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);         \
-                dma_queue_push(dma_queue,                                                                           \
-                    dma_make_ptr(src0_vtcm, data_src + src0_pref_off),                                              \
-                    src0_row_size_aligned, nb01, src0_data_row_size, pref_block_size);                              \
-                                                                                                                    \
-                if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) {                                                 \
-                    const size_t src1_pref_off = src1_contig ? (pref_ir * nb11) :                                   \
-                        unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc,      \
-                                          nb13_bc);                                                                 \
-                    dma_queue_push(dma_queue,                                                                       \
-                        dma_make_ptr(src1_vtcm, data_src1 + src1_pref_off),                                         \
-                        uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, pref_block_size);              \
-                }                                                                                                   \
-            }                                                                                                       \
-        }                                                                                                           \
-        ir += block_size;                                                                                           \
-    }                                                                                                               \
-                                                                                                                    \
-    dma_queue_flush(dma_queue);                                                                                     \
-}
-
-// F32 unary task: row-block DMA/VTCM plumbing, float-typed VTCM buffers.
-#define DEFINE_UNARY_TASK(NAME, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) \
-    DEFINE_UNARY_TASK_IMPL(NAME, float, f32, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR)
-
-DEFINE_UNARY_TASK(norm,           false, false, norm_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(rms_norm,       false, false, rms_norm_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(rms_norm_mul,   true,  false, rms_norm_mul_f32(src0_vtcm, uctx->broadcast_weight ? (const float *) src1_vtcm_data : src1_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(scale,          false, false, scale_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(clamp,          false, false, clamp_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(leaky_relu,     false, false, leaky_relu_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(sqr,            false, false, sqr_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(sqrt,           false, false, sqrt_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_neg,      false, false, neg_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_exp,      false, false, exp_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_sigmoid,  false, false, sigmoid_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_silu,     false, false, silu_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_gelu,     false, false, gelu_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_softplus, false, false, softplus_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_tanh,     false, false, tanh_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_abs,      false, false, abs_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_log,      false, false, log_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(unary_relu,     false, false, relu_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(l2_norm,        false, false, l2_norm_f32(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK(tri,            false, true,  tri_f32(src0_vtcm, dst_vtcm, block_size, ir, uctx))
-
-// F16 unary tasks: same DMA/VTCM plumbing as DEFINE_UNARY_TASK, but VTCM buffers are
-// _Float16-typed. None of the current F16 ops need RMS_NORM_MUL or TRI support.
-DEFINE_UNARY_TASK_IMPL(norm,      _Float16, f16, false, false, norm_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(rms_norm,  _Float16, f16, false, false, rms_norm_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(scale,     _Float16, f16, false, false, scale_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(clamp,     _Float16, f16, false, false, clamp_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(sqr,       _Float16, f16, false, false, sqr_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(sqrt,      _Float16, f16, false, false, sqrt_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(l2_norm,   _Float16, f16, false, false, l2_norm_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(unary_abs, _Float16, f16, false, false, abs_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-DEFINE_UNARY_TASK_IMPL(unary_log, _Float16, f16, false, false, log_f16(src0_vtcm, dst_vtcm, block_size, uctx))
-
-// Apply a pointwise unary op to one column tile that is already in VTCM.
-#define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR)                                                         \
-static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) {                            \
-    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;                                  \
-    struct htp_ops_context * octx = uctx->octx;                                                                       \
-    const struct htp_tensor * src = octx->src[0];                                                                     \
-    const struct htp_tensor * dst = octx->dst;                                                                        \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                            \
-                                                                                                                      \
-    htp_unary_preamble;                                                                                               \
-                                                                                                                      \
-    uint32_t     src0_nrows_per_thread = uctx->src0_nrows_per_thread;                                                 \
-                                                                                                                      \
-    int32_t *      op_params = octx->op_params;                                                                       \
-    const uint32_t col_tile  = uctx->col_tile;                                                                        \
-                                                                                                                      \
-    const uint32_t src0_nrows = uctx->src0_nrows;                                                                     \
-    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;                                    \
-    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);        \
-                                                                                                                      \
-    if (src0_start_row >= src0_end_row) {                                                                             \
-        return;                                                                                                       \
-    }                                                                                                                 \
-                                                                                                                      \
-    const uint8_t * restrict data_src = uctx->data_src0;                                                              \
-    uint8_t * restrict       data_dst = uctx->data_dst;                                                               \
-                                                                                                                      \
-    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);                             \
-    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);                               \
-                                                                                                                      \
-    const size_t src0_half = uctx->src0_vtcm_half_size;                                                               \
-    const size_t dst_half  = uctx->dst_vtcm_half_size;                                                                \
-                                                                                                                      \
-    dma_queue * dmaq = octx->ctx->dma[ith];                                                                           \
-                                                                                                                      \
-    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;                                               \
-    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;                                               \
-    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;                                              \
-    const struct fastdiv_values * div_tpr   = &uctx->kparams->div_tpr;                                                \
-                                                                                                                      \
-    const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile;                                                   \
-    const int32_t  tri_ttype     = (IS_TRI) ? op_params[0] : 0;                                                       \
-                                                                                                                      \
-    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&                                                         \
-                             (nb03 == (size_t)ne02 * nb02);                                                           \
-    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&                                                         \
-                             (nb3  == (size_t)ne2  * nb2);                                                            \
-                                                                                                                      \
-    const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row;                                     \
-                                                                                                                      \
-    for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) {                            \
-        const uint32_t row  = src0_start_row + t / tiles_per_row;                                                     \
-        const uint32_t col  = (t % tiles_per_row) * col_tile;                                                         \
-        const uint32_t tw   = MIN(col_tile, ne0 - col);                                                               \
-        const size_t   tb   = (size_t) tw * sizeof(float);                                                            \
-        const size_t   soff = (src0_contig ? (row * nb01) :                                                           \
-                               unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +  \
-                               (size_t) col * sizeof(float);                                                          \
-                                                                                                                      \
-        dma_queue_push(dmaq, dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0);              \
-        dma_queue_push(dmaq, dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);  \
-    }                                                                                                                 \
-                                                                                                                      \
-    uint32_t row = src0_start_row;                                                                                    \
-    uint32_t col = 0;                                                                                                 \
-    uint32_t tile_in_row = 0;                                                                                         \
-    uint32_t i01 = fastmodulo(row, ne01, div_ne01);                                                                   \
-                                                                                                                      \
-    uint32_t prow = src0_start_row + fastdiv(2, div_tpr);                                                             \
-    uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile;                                                 \
-    uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr);                                                    \
-                                                                                                                      \
-    for (uint32_t t = 0; t < total_tiles; t++) {                                                                      \
-        uint8_t * dst_vtcm = (uint8_t *) dma_queue_pop(dmaq).src;                                                     \
-        uint8_t * src_vtcm = (uint8_t *) dma_queue_pop(dmaq).dst;                                                     \
-                                                                                                                      \
-        const uint32_t tw  = MIN(col_tile, ne0 - col);                                                                \
-                                                                                                                      \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                         \
-        CORE_TILE_EXPR;                                                                                               \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                          \
-                                                                                                                      \
-        const size_t doff = (dst_contig ? (row * nb1) :                                                               \
-                             unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) +         \
-                             (size_t) col * sizeof(float);                                                            \
-        const size_t tb   = (size_t) tw * sizeof(float);                                                              \
-        dma_queue_push(dmaq, dma_make_ptr(data_dst + doff, dst_vtcm), tb, tb, tb, 1);                                 \
-                                                                                                                      \
-        const uint32_t pt = t + 2;                                                                                    \
-        if (pt < total_tiles) {                                                                                       \
-            const uint32_t ptw  = MIN(col_tile, ne0 - pcol);                                                          \
-            const size_t   ptb  = (size_t) ptw * sizeof(float);                                                       \
-            const size_t   psoff = (src0_contig ? (prow * nb01) :                                                     \
-                                    unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02,     \
-                                                     nb03)) +                                                         \
-                                   (size_t) pcol * sizeof(float);                                                     \
-            dma_queue_push(dmaq, dma_make_ptr(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1);                         \
-        }                                                                                                             \
-                                                                                                                      \
-        tile_in_row++;                                                                                                \
-        col += col_tile;                                                                                              \
-        if (tile_in_row == tiles_per_row) {                                                                           \
-            tile_in_row = 0;                                                                                          \
-            col = 0;                                                                                                  \
-            row++;                                                                                                    \
-            i01++;                                                                                                    \
-            if (i01 == ne01) {                                                                                        \
-                i01 = 0;                                                                                              \
-            }                                                                                                         \
-        }                                                                                                             \
-                                                                                                                      \
-        ptile_in_row++;                                                                                               \
-        pcol += col_tile;                                                                                             \
-        if (ptile_in_row == tiles_per_row) {                                                                          \
-            ptile_in_row = 0;                                                                                         \
-            pcol = 0;                                                                                                 \
-            prow++;                                                                                                   \
-        }                                                                                                             \
-    }                                                                                                                 \
-                                                                                                                      \
-    dma_queue_flush(dmaq);                                                                                            \
-}
-
-static inline void tile_scale_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw, const int32_t * op_params) {
+#// Pointwise unary ops on one column tile in VTCM.
+static void tile_scale_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
     float scale = 0.f;
-    float bias = 0.f;
-    memcpy(&scale, &op_params[0], sizeof(float));
-    memcpy(&bias,  &op_params[1], sizeof(float));
-    hvx_scale_offset_f32_aa(dst_vtcm, src_vtcm, tw, scale, bias);
+    float bias  = 0.f;
+    memcpy(&scale, &uctx->octx->op_params[0], sizeof(float));
+    memcpy(&bias,  &uctx->octx->op_params[1], sizeof(float));
+    hvx_scale_offset_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw, scale, bias);
 }

-static inline void tile_clamp_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw, const int32_t * op_params) {
+static void tile_clamp_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
     float min = 0.f;
     float max = 0.f;
-    memcpy(&min, &op_params[0], sizeof(float));
-    memcpy(&max, &op_params[1], sizeof(float));
-    hvx_clamp_scalar_f32(dst_vtcm, src_vtcm, min, max, tw);
+    memcpy(&min, &uctx->octx->op_params[0], sizeof(float));
+    memcpy(&max, &uctx->octx->op_params[1], sizeof(float));
+    hvx_clamp_scalar_f32((uint8_t *) dst, (const uint8_t *) src, min, max, tw);
 }

-static inline void tile_leaky_relu_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw, const int32_t * op_params) {
+static void tile_leaky_relu_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
     float negative_slope = 0.f;
-    memcpy(&negative_slope, &op_params[0], sizeof(float));
-    hvx_leaky_relu_scalar_f32(dst_vtcm, src_vtcm, negative_slope, tw);
+    memcpy(&negative_slope, &uctx->octx->op_params[0], sizeof(float));
+    hvx_leaky_relu_scalar_f32((uint8_t *) dst, (const uint8_t *) src, negative_slope, tw);
+}
+
+static void tile_sqr_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_sqr_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
 }

-static inline void tile_unary_softplus_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw) {
-    const float * restrict sf = (const float *) src_vtcm;
-    float * restrict df       = (float *) dst_vtcm;
+static void tile_sqrt_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_sqrt_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
+}
+
+static void tile_neg_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_scale_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw, -1.0f);
+}
+
+static void tile_exp_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_exp_f32((uint8_t *) dst, (const uint8_t *) src, tw, false);
+}
+
+static void tile_sigmoid_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_sigmoid_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
+}
+
+static void tile_silu_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_sigmoid_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
+    hvx_mul_f32_aaa((uint8_t *) dst, (const uint8_t *) src, (uint8_t *) dst, tw);
+}
+
+static void tile_gelu_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_mul_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 1.702f, tw);
+    hvx_sigmoid_f32_aa((uint8_t *) dst, (uint8_t *) dst, tw);
+    hvx_mul_f32_aaa((uint8_t *) dst, (const uint8_t *) src, (uint8_t *) dst, tw);
+}
+
+static void tile_softplus_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    const float * restrict sf = (const float *) src;
+    float * restrict df       = (float *) dst;
     for (uint32_t i = 0; i < tw; i++) {
         float x = sf[i];
         df[i] = (x > 20.0f) ? x : logf(1.0f + expf(x));
     }
 }

-// silu(x) = x * sigmoid(x)
-static inline void tile_silu_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw) {
-    hvx_sigmoid_f32_aa(dst_vtcm, src_vtcm, tw);
-    hvx_mul_f32_aaa(dst_vtcm, src_vtcm, dst_vtcm, tw);
+static void tile_tanh_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_tanh_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
 }

-// gelu(x) = x * sigmoid(1.702 * x)  (quick/sigmoid approximation, matches CPU GELU_QUICK reference)
-static inline void tile_gelu_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw) {
-    hvx_mul_scalar_f32(dst_vtcm, src_vtcm, 1.702f, tw);
-    hvx_sigmoid_f32_aa(dst_vtcm, dst_vtcm, tw);
-    hvx_mul_f32_aaa(dst_vtcm, src_vtcm, dst_vtcm, tw);
+static void tile_abs_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_abs_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
+}
+
+static void tile_log_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_log_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
 }

-// Triangular mask applied to one column tile. Boundary is an absolute column index, so
-// each vector compares against its absolute column position (col_start + i*VLEN_FP32).
-static inline void tri_apply_tile_f32(const uint8_t * restrict src, uint8_t * restrict dst,
-                                      uint32_t tile_elems, uint32_t col_start, uint32_t i01,
-                                      uint32_t ne0, int32_t ttype) {
+static void tile_relu_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
+    (void) uctx;
+    hvx_max_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 0.0f, tw);
+}
+
+static void tri_apply_tile_f32(const void * restrict src, void * restrict dst,
+                               uint32_t tile_elems, uint32_t col_start, uint32_t i01,
+                               uint32_t ne0, int32_t ttype) {
     const HVX_Vector * restrict v_src = (const HVX_Vector *) src;
     HVX_Vector * restrict v_dst       = (HVX_Vector *) dst;
     const HVX_Vector zero = hvx_vec_splat_f32(0.0f);
@@ -1074,22 +839,623 @@ static inline void tri_apply_tile_f32(const uint8_t * restrict src, uint8_t * re
     }
 }

-DEFINE_UNARY_TILED_TASK(scale,          false, tile_scale_f32(dst_vtcm, src_vtcm, tw, op_params))
-DEFINE_UNARY_TILED_TASK(clamp,          false, tile_clamp_f32(dst_vtcm, src_vtcm, tw, op_params))
-DEFINE_UNARY_TILED_TASK(leaky_relu,     false, tile_leaky_relu_f32(dst_vtcm, src_vtcm, tw, op_params))
-DEFINE_UNARY_TILED_TASK(sqr,            false, hvx_sqr_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(sqrt,           false, hvx_sqrt_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_neg,      false, hvx_scale_f32_aa(dst_vtcm, src_vtcm, tw, -1.0f))
-DEFINE_UNARY_TILED_TASK(unary_exp,      false, hvx_exp_f32(dst_vtcm, src_vtcm, tw, false))
-DEFINE_UNARY_TILED_TASK(unary_sigmoid,  false, hvx_sigmoid_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_silu,     false, tile_silu_f32(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_gelu,     false, tile_gelu_f32(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_softplus, false, tile_unary_softplus_f32(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_tanh,     false, hvx_tanh_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_abs,      false, hvx_abs_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_log,      false, hvx_log_f32_aa(dst_vtcm, src_vtcm, tw))
-DEFINE_UNARY_TILED_TASK(unary_relu,     false, hvx_max_scalar_f32(dst_vtcm, src_vtcm, 0.0f, tw))
-DEFINE_UNARY_TILED_TASK(tri,            true,  tri_apply_tile_f32(src_vtcm, dst_vtcm, tw, col, i01, ne0, tri_ttype))
+// 1. Standard row-block unary task (F32 and F16).
+static void unary_thread_row_block(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;
+    struct htp_ops_context * octx = uctx->octx;
+    const struct htp_tensor * src = octx->src[0];
+    const struct htp_tensor * dst = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    htp_unary_preamble;
+
+    const uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread;
+    const size_t src0_data_row_size = uctx->src0_data_row_size;
+    const size_t dst_data_row_size  = uctx->dst_data_row_size;
+    const size_t src0_row_size_aligned = uctx->src0_row_size_aligned;
+    const size_t dst_row_size_aligned  = uctx->dst_row_size_aligned;
+
+    const uint32_t src0_nrows = uctx->src0_nrows;
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);
+
+    if (src0_start_row >= src0_end_row) {
+        return;
+    }
+
+    const dma_addr_t data_src = uctx->data_src0;
+    const dma_addr_t data_dst = uctx->data_dst;
+
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_vtcm_half_size = uctx->src0_vtcm_half_size;
+    const size_t dst_vtcm_half_size  = uctx->dst_vtcm_half_size;
+
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&
+                             (nb03 == (size_t)ne02 * nb02);
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&
+                             (nb3  == (size_t)ne2  * nb2);
+
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;
+
+    const uint32_t src0_max_block = src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01);
+    const uint32_t dst_max_block  = dst_contig  ? uctx->block : MIN((uint32_t)uctx->block, ne1);
+    const uint32_t BLOCK = MIN(src0_max_block, dst_max_block);
+    if (BLOCK == 0) {
+        FARF(ERROR, "unary-row-block : current VTCM reservation %zu is too small, needed at least %zu\n",
+             uctx->vtcm_src0_size_per_thread, src0_row_size_aligned);
+        return;
+    }
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig,
+                                                     ne01, div_ne01);
+
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)),
+            nb1, dst_row_size_aligned, dst_data_row_size, 0);
+
+        const size_t src0_off = src0_contig ? (ir * nb01) :
+            unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+        dma_queue_push(dma_q,
+            dma_make_data(src0_vtcm_data + (vtcm_idx * src0_vtcm_half_size), data_src + src0_off),
+            src0_row_size_aligned, nb01, src0_data_row_size, block_size);
+
+        ir += block_size;
+    }
+
+    unary_compute_fn_t compute = (unary_compute_fn_t) uctx->compute;
+
+    for (uint32_t ir = src0_start_row; ir < src0_end_row; ) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig,
+                                                     ne01, div_ne01);
+
+        void * dst_vtcm  = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * src0_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+        compute(src0_vtcm, dst_vtcm, block_size, uctx);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+
+        const size_t dst_off = dst_contig ? (ir * nb1) :
+            unary_row_offset(ir, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3);
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst + dst_off, dst_vtcm),
+            nb1, dst_row_size_aligned, dst_data_row_size, block_size);
+
+        const uint32_t next_ir = ir + block_size;
+        if (next_ir < src0_end_row) {
+            const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, src0_contig,
+                                                              dst_contig, ne01, div_ne01);
+            const uint32_t pref_ir = next_ir + next_block_size;
+            if (pref_ir < src0_end_row) {
+                const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, src0_contig,
+                                                                  dst_contig, ne01, div_ne01);
+                const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) :
+                    unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+                dma_queue_push(dma_q,
+                    dma_make_data(src0_vtcm, data_src + src0_pref_off),
+                    src0_row_size_aligned, nb01, src0_data_row_size, pref_block_size);
+            }
+        }
+        ir += block_size;
+    }
+
+    dma_queue_flush(dma_q);
+}
+
+// 2. RMS_NORM_MUL row-block task with weight buffer.
+static void unary_thread_rms_norm_mul_f32(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;
+    struct htp_ops_context * octx = uctx->octx;
+    const struct htp_tensor * src = octx->src[0];
+    const struct htp_tensor * dst = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    htp_unary_preamble;
+
+    const uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread;
+    const size_t src0_data_row_size = uctx->src0_data_row_size;
+    const size_t dst_data_row_size  = uctx->dst_data_row_size;
+    const size_t src0_row_size_aligned = uctx->src0_row_size_aligned;
+    const size_t dst_row_size_aligned  = uctx->dst_row_size_aligned;
+
+    const uint32_t src0_nrows = uctx->src0_nrows;
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);
+
+    if (src0_start_row >= src0_end_row) {
+        return;
+    }
+
+    const dma_addr_t data_src  = uctx->data_src0;
+    const dma_addr_t data_src1 = uctx->data_src1;
+    const dma_addr_t data_dst  = uctx->data_dst;
+
+    const struct htp_tensor * src1 = octx->src[1];
+    const uint32_t nb11 = src1->nb[1];
+    const uint32_t nb12 = src1->nb[2];
+    const uint32_t nb13 = src1->nb[3];
+    const uint32_t nb11_bc = (src1->ne[1] > 1) ? nb11 : 0;
+    const uint32_t nb12_bc = (src1->ne[2] > 1) ? nb12 : 0;
+    const uint32_t nb13_bc = (src1->ne[3] > 1) ? nb13 : 0;
+    const bool src1_contig = ((nb12 == (size_t)ne01 * nb11) && (nb13 == (size_t)ne02 * nb12));
+
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);
+    uint8_t * src1_vtcm_data = uctx->vtcm_src1 ? (uctx->vtcm_src1 + (ith * uctx->vtcm_src1_size_per_thread)) : NULL;
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_vtcm_half_size = uctx->src0_vtcm_half_size;
+    const size_t src1_vtcm_half_size = uctx->src1_vtcm_half_size;
+    const size_t dst_vtcm_half_size  = uctx->dst_vtcm_half_size;
+
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&
+                             (nb03 == (size_t)ne02 * nb02);
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&
+                             (nb3  == (size_t)ne2  * nb2);
+
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;
+
+    const bool src1_needs_row_clip = !uctx->broadcast_weight && !src1_contig;
+    const bool block_src0_contig = src0_contig && !src1_needs_row_clip;
+    const bool block_dst_contig  = dst_contig  && !src1_needs_row_clip;
+
+    const uint32_t src0_max_block = block_src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01);
+    const uint32_t dst_max_block  = block_dst_contig  ? uctx->block : MIN((uint32_t)uctx->block, ne1);
+    const uint32_t BLOCK = MIN(src0_max_block, dst_max_block);
+    if (BLOCK == 0) {
+        FARF(ERROR, "unary-rms-norm-mul : current VTCM reservation %zu is too small, needed at least %zu\n",
+             uctx->vtcm_src0_size_per_thread, src0_row_size_aligned);
+        return;
+    }
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    if (uctx->broadcast_weight) {
+        dma_queue_push(dma_q, dma_make_data(src1_vtcm_data, data_src1),
+                       uctx->src1_row_size_aligned, 0, uctx->src1_data_row_size, 1);
+        dma_queue_flush(dma_q);
+    }
+
+    for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig,
+                                                     ne01, div_ne01);
+
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)),
+            nb1, dst_row_size_aligned, dst_data_row_size, 0);
+
+        const size_t src0_off = src0_contig ? (ir * nb01) :
+            unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+        dma_queue_push(dma_q,
+            dma_make_data(src0_vtcm_data + (vtcm_idx * src0_vtcm_half_size), data_src + src0_off),
+            src0_row_size_aligned, nb01, src0_data_row_size, block_size);
+
+        if (!uctx->broadcast_weight) {
+            const size_t src1_off = src1_contig ? (ir * nb11) :
+                unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, nb13_bc);
+            dma_queue_push(dma_q,
+                dma_make_data(src1_vtcm_data + (vtcm_idx * src1_vtcm_half_size), data_src1 + src1_off),
+                uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, block_size);
+        }
+
+        ir += block_size;
+    }
+
+    unary_rms_norm_mul_compute_fn_t compute = (unary_rms_norm_mul_compute_fn_t) uctx->compute;
+
+    for (uint32_t ir = src0_start_row; ir < src0_end_row; ) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig,
+                                                     ne01, div_ne01);
+
+        void * dst_vtcm  = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * src0_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+        void * src1_vtcm = NULL;
+        if (!uctx->broadcast_weight) {
+            src1_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+        }
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+        const void * w = uctx->broadcast_weight ? (const void *) src1_vtcm_data : src1_vtcm;
+        compute(src0_vtcm, w, dst_vtcm, block_size, uctx);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+
+        const size_t dst_off = dst_contig ? (ir * nb1) :
+            unary_row_offset(ir, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3);
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst + dst_off, dst_vtcm),
+            nb1, dst_row_size_aligned, dst_data_row_size, block_size);
+
+        const uint32_t next_ir = ir + block_size;
+        if (next_ir < src0_end_row) {
+            const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, block_src0_contig,
+                                                              block_dst_contig, ne01, div_ne01);
+            const uint32_t pref_ir = next_ir + next_block_size;
+            if (pref_ir < src0_end_row) {
+                const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, block_src0_contig,
+                                                                  block_dst_contig, ne01, div_ne01);
+                const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) :
+                    unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+                dma_queue_push(dma_q,
+                    dma_make_data(src0_vtcm, data_src + src0_pref_off),
+                    src0_row_size_aligned, nb01, src0_data_row_size, pref_block_size);
+
+                if (!uctx->broadcast_weight) {
+                    const size_t src1_pref_off = src1_contig ? (pref_ir * nb11) :
+                        unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc,
+                                          nb13_bc);
+                    dma_queue_push(dma_q,
+                        dma_make_data(src1_vtcm, data_src1 + src1_pref_off),
+                        uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, pref_block_size);
+                }
+            }
+        }
+        ir += block_size;
+    }
+
+    dma_queue_flush(dma_q);
+}
+
+// 3. TRI row-block task with row index ir.
+static void unary_thread_tri_f32(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;
+    struct htp_ops_context * octx = uctx->octx;
+    const struct htp_tensor * src = octx->src[0];
+    const struct htp_tensor * dst = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    htp_unary_preamble;
+
+    const uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread;
+    const size_t src0_data_row_size = uctx->src0_data_row_size;
+    const size_t dst_data_row_size  = uctx->dst_data_row_size;
+    const size_t src0_row_size_aligned = uctx->src0_row_size_aligned;
+    const size_t dst_row_size_aligned  = uctx->dst_row_size_aligned;
+
+    const uint32_t src0_nrows = uctx->src0_nrows;
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);
+
+    if (src0_start_row >= src0_end_row) {
+        return;
+    }
+
+    const dma_addr_t data_src = uctx->data_src0;
+    const dma_addr_t data_dst = uctx->data_dst;
+
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_vtcm_half_size = uctx->src0_vtcm_half_size;
+    const size_t dst_vtcm_half_size  = uctx->dst_vtcm_half_size;
+
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&
+                             (nb03 == (size_t)ne02 * nb02);
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&
+                             (nb3  == (size_t)ne2  * nb2);
+
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;
+
+    const uint32_t src0_max_block = src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01);
+    const uint32_t dst_max_block  = dst_contig  ? uctx->block : MIN((uint32_t)uctx->block, ne1);
+    const uint32_t BLOCK = MIN(src0_max_block, dst_max_block);
+    if (BLOCK == 0) {
+        FARF(ERROR, "unary-tri : current VTCM reservation %zu is too small, needed at least %zu\n",
+             uctx->vtcm_src0_size_per_thread, src0_row_size_aligned);
+        return;
+    }
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig,
+                                                     ne01, div_ne01);
+
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)),
+            nb1, dst_row_size_aligned, dst_data_row_size, 0);
+
+        const size_t src0_off = src0_contig ? (ir * nb01) :
+            unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+        dma_queue_push(dma_q,
+            dma_make_data(src0_vtcm_data + (vtcm_idx * src0_vtcm_half_size), data_src + src0_off),
+            src0_row_size_aligned, nb01, src0_data_row_size, block_size);
+
+        ir += block_size;
+    }
+
+    unary_tri_compute_fn_t compute = (unary_tri_compute_fn_t) uctx->compute;
+
+    for (uint32_t ir = src0_start_row; ir < src0_end_row; ) {
+        const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig,
+                                                     ne01, div_ne01);
+
+        void * dst_vtcm  = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * src0_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+        compute(src0_vtcm, dst_vtcm, block_size, ir, uctx);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);
+
+        const size_t dst_off = dst_contig ? (ir * nb1) :
+            unary_row_offset(ir, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3);
+        dma_queue_push(dma_q,
+            dma_make_data(data_dst + dst_off, dst_vtcm),
+            nb1, dst_row_size_aligned, dst_data_row_size, block_size);
+
+        const uint32_t next_ir = ir + block_size;
+        if (next_ir < src0_end_row) {
+            const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, src0_contig,
+                                                              dst_contig, ne01, div_ne01);
+            const uint32_t pref_ir = next_ir + next_block_size;
+            if (pref_ir < src0_end_row) {
+                const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, src0_contig,
+                                                                  dst_contig, ne01, div_ne01);
+                const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) :
+                    unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03);
+                dma_queue_push(dma_q,
+                    dma_make_data(src0_vtcm, data_src + src0_pref_off),
+                    src0_row_size_aligned, nb01, src0_data_row_size, pref_block_size);
+            }
+        }
+        ir += block_size;
+    }
+
+    dma_queue_flush(dma_q);
+}
+
+// 4. Pointwise tiled unary task.
+static void unary_thread_tiled(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;
+    struct htp_ops_context * octx = uctx->octx;
+    const struct htp_tensor * src = octx->src[0];
+    const struct htp_tensor * dst = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    htp_unary_preamble;
+
+    const uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread;
+    const uint32_t col_tile              = uctx->col_tile;
+
+    const uint32_t src0_nrows     = uctx->src0_nrows;
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);
+
+    if (src0_start_row >= src0_end_row) {
+        return;
+    }
+
+    const dma_addr_t data_src = uctx->data_src0;
+    const dma_addr_t data_dst = uctx->data_dst;
+
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_half = uctx->src0_vtcm_half_size;
+    const size_t dst_half  = uctx->dst_vtcm_half_size;
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;
+    const struct fastdiv_values * div_tpr   = &uctx->kparams->div_tpr;
+
+    const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile;
+
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&
+                             (nb03 == (size_t)ne02 * nb02);
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&
+                             (nb3  == (size_t)ne2  * nb2);
+
+    const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row;
+
+    for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) {
+        const uint32_t row  = src0_start_row + t / tiles_per_row;
+        const uint32_t col  = (t % tiles_per_row) * col_tile;
+        const uint32_t tw   = MIN(col_tile, ne0 - col);
+        const size_t   tb   = (size_t) tw * sizeof(float);
+        const size_t   soff = (src0_contig ? (row * nb01) :
+                               unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +
+                               (size_t) col * sizeof(float);
+
+        dma_queue_push(dma_q, dma_make_data(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0);
+        dma_queue_push(dma_q, dma_make_data(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);
+    }
+
+    unary_tile_compute_fn_t compute = (unary_tile_compute_fn_t) uctx->compute;
+
+    uint32_t row = src0_start_row;
+    uint32_t col = 0;
+    uint32_t tile_in_row = 0;
+
+    uint32_t prow = src0_start_row + fastdiv(2, div_tpr);
+    uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile;
+    uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr);
+
+    for (uint32_t t = 0; t < total_tiles; t++) {
+        void * dst_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * src_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+
+        const uint32_t tw = MIN(col_tile, ne0 - col);
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t);
+        compute(dst_vtcm, src_vtcm, tw, uctx);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t);
+
+        const size_t doff = (dst_contig ? (row * nb1) :
+                             unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) +
+                             (size_t) col * sizeof(float);
+        const size_t tb   = (size_t) tw * sizeof(float);
+        dma_queue_push(dma_q, dma_make_data(data_dst + doff, dst_vtcm), tb, tb, tb, 1);
+
+        const uint32_t pt = t + 2;
+        if (pt < total_tiles) {
+            const uint32_t ptw  = MIN(col_tile, ne0 - pcol);
+            const size_t   ptb  = (size_t) ptw * sizeof(float);
+            const size_t   psoff = (src0_contig ? (prow * nb01) :
+                                    unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02,
+                                                     nb03)) +
+                                   (size_t) pcol * sizeof(float);
+            dma_queue_push(dma_q, dma_make_data(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1);
+        }
+
+        tile_in_row++;
+        col += col_tile;
+        if (tile_in_row == tiles_per_row) {
+            tile_in_row = 0;
+            col = 0;
+            row++;
+        }
+
+        ptile_in_row++;
+        pcol += col_tile;
+        if (ptile_in_row == tiles_per_row) {
+            ptile_in_row = 0;
+            pcol = 0;
+            prow++;
+        }
+    }
+
+    dma_queue_flush(dma_q);
+}
+
+// 5. TRI tiled task.
+static void unary_thread_tiled_tri_f32(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;
+    struct htp_ops_context * octx = uctx->octx;
+    const struct htp_tensor * src = octx->src[0];
+    const struct htp_tensor * dst = octx->dst;
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
+    htp_unary_preamble;
+
+    const uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread;
+    const int32_t * op_params            = octx->op_params;
+    const uint32_t col_tile              = uctx->col_tile;
+
+    const uint32_t src0_nrows     = uctx->src0_nrows;
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);
+
+    if (src0_start_row >= src0_end_row) {
+        return;
+    }
+
+    const dma_addr_t data_src = uctx->data_src0;
+    const dma_addr_t data_dst = uctx->data_dst;
+
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);
+
+    const size_t src0_half = uctx->src0_vtcm_half_size;
+    const size_t dst_half  = uctx->dst_vtcm_half_size;
+
+    dma_queue * dma_q = octx->ctx->dma[ith];
+
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;
+    const struct fastdiv_values * div_tpr   = &uctx->kparams->div_tpr;
+
+    const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile;
+    const int32_t  tri_ttype     = op_params[0];
+
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&
+                             (nb03 == (size_t)ne02 * nb02);
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&
+                             (nb3  == (size_t)ne2  * nb2);
+
+    const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row;
+
+    for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) {
+        const uint32_t row  = src0_start_row + t / tiles_per_row;
+        const uint32_t col  = (t % tiles_per_row) * col_tile;
+        const uint32_t tw   = MIN(col_tile, ne0 - col);
+        const size_t   tb   = (size_t) tw * sizeof(float);
+        const size_t   soff = (src0_contig ? (row * nb01) :
+                               unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +
+                               (size_t) col * sizeof(float);
+
+        dma_queue_push(dma_q, dma_make_data(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0);
+        dma_queue_push(dma_q, dma_make_data(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);
+    }
+
+    unary_tiled_tri_compute_fn_t compute = (unary_tiled_tri_compute_fn_t) uctx->compute;
+
+    uint32_t row = src0_start_row;
+    uint32_t col = 0;
+    uint32_t tile_in_row = 0;
+    uint32_t i01 = fastmodulo(row, ne01, div_ne01);
+
+    uint32_t prow = src0_start_row + fastdiv(2, div_tpr);
+    uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile;
+    uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr);
+
+    for (uint32_t t = 0; t < total_tiles; t++) {
+        void * dst_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).src;
+        void * src_vtcm = (void *) (uintptr_t) dma_queue_pop(dma_q).dst;
+
+        const uint32_t tw = MIN(col_tile, ne0 - col);
+
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t);
+        compute(src_vtcm, dst_vtcm, tw, col, i01, ne0, tri_ttype);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t);
+
+        const size_t doff = (dst_contig ? (row * nb1) :
+                             unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) +
+                             (size_t) col * sizeof(float);
+        const size_t tb   = (size_t) tw * sizeof(float);
+        dma_queue_push(dma_q, dma_make_data(data_dst + doff, dst_vtcm), tb, tb, tb, 1);
+
+        const uint32_t pt = t + 2;
+        if (pt < total_tiles) {
+            const uint32_t ptw  = MIN(col_tile, ne0 - pcol);
+            const size_t   ptb  = (size_t) ptw * sizeof(float);
+            const size_t   psoff = (src0_contig ? (prow * nb01) :
+                                    unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02,
+                                                     nb03)) +
+                                   (size_t) pcol * sizeof(float);
+            dma_queue_push(dma_q, dma_make_data(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1);
+        }
+
+        tile_in_row++;
+        col += col_tile;
+        if (tile_in_row == tiles_per_row) {
+            tile_in_row = 0;
+            col = 0;
+            row++;
+            i01++;
+            if (i01 == ne01) {
+                i01 = 0;
+            }
+        }
+
+        ptile_in_row++;
+        pcol += col_tile;
+        if (ptile_in_row == tiles_per_row) {
+            ptile_in_row = 0;
+            pcol = 0;
+            prow++;
+        }
+    }
+
+    dma_queue_flush(dma_q);
+}

 static int execute_op_unary(struct htp_ops_context * octx) {
     int err = HTP_STATUS_OK;
@@ -1207,115 +1573,128 @@ static int execute_op_unary(struct htp_ops_context * octx) {
          src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
          kparams->vtcm_src0_size, kparams->vtcm_src1_size, kparams->vtcm_dst_size);

-    if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
-        struct htp_unary_context uctx = {
-            .octx                  = octx,
-            .kparams               = kparams,
-            .src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
-            .src0_nrows            = nrows,
-            .row_start             = row_start,
-
-            .data_src0             = (const uint8_t *)src0->data,
-            .data_src1             = (octx->op == HTP_OP_RMS_NORM_MUL) ? (const uint8_t *)src1->data : NULL,
-            .data_dst              = (uint8_t *)dst->data,
-
-            .src0_data_row_size    = src0_data_row_size,
-            .src1_data_row_size    = src1_data_row_size,
-            .dst_data_row_size     = dst_data_row_size,
-
-            .src0_row_size_aligned = src0_row_size_aligned,
-            .src1_row_size_aligned = src1_row_size_aligned,
-            .dst_row_size_aligned  = dst_row_size_aligned,
-
-            .src0_vtcm_half_size   = kparams->vtcm_src0_size_per_thread / 2,
-            .src1_vtcm_half_size   = (octx->op == HTP_OP_RMS_NORM_MUL) ? (kparams->vtcm_src1_size_per_thread / (broadcast_weight ? 1 : 2)) : 0,
-            .dst_vtcm_half_size    = kparams->vtcm_dst_size_per_thread / 2,
-
-            .block                 = kparams->block,
-            .nc                    = src0->ne[0],
-            .col_tile              = col_tile,
-            .broadcast_weight      = broadcast_weight,
-
-            .vtcm_src0             = VTCM_LAYOUT_PTR(uint8_t, base, 0),
-            .vtcm_src1             = VTCM_LAYOUT_PTR_OPTIONAL(uint8_t, base, kparams->vtcm_src0_size, kparams->vtcm_src1_size > 0),
-            .vtcm_dst              = VTCM_LAYOUT_PTR(uint8_t, base, kparams->vtcm_src0_size + kparams->vtcm_src1_size),
-
-            .vtcm_src0_size_per_thread = kparams->vtcm_src0_size_per_thread,
-            .vtcm_src1_size_per_thread = kparams->vtcm_src1_size_per_thread,
-            .vtcm_dst_size_per_thread  = kparams->vtcm_dst_size_per_thread,
-        };
-
-        FARF(HIGH, "%s: %s mode (col_tile %u)\n", op_type, col_tile ? "tiled" : "row-block", col_tile);
-
-        worker_callback_t task_func = NULL;
-        if (col_tile) {
-            switch (octx->op) {
-                case HTP_OP_SCALE:           task_func = unary_task_f32_tiled_scale;          break;
-                case HTP_OP_CLAMP:           task_func = unary_task_f32_tiled_clamp;          break;
-                case HTP_OP_LEAKY_RELU:      task_func = unary_task_f32_tiled_leaky_relu;     break;
-                case HTP_OP_SQR:             task_func = unary_task_f32_tiled_sqr;            break;
-                case HTP_OP_SQRT:            task_func = unary_task_f32_tiled_sqrt;           break;
-                case HTP_OP_UNARY_NEG:       task_func = unary_task_f32_tiled_unary_neg;      break;
-                case HTP_OP_UNARY_EXP:       task_func = unary_task_f32_tiled_unary_exp;      break;
-                case HTP_OP_UNARY_SIGMOID:   task_func = unary_task_f32_tiled_unary_sigmoid;  break;
-                case HTP_OP_UNARY_SILU:      task_func = unary_task_f32_tiled_unary_silu;     break;
-                case HTP_OP_UNARY_GELU:      task_func = unary_task_f32_tiled_unary_gelu;     break;
-                case HTP_OP_UNARY_SOFTPLUS:  task_func = unary_task_f32_tiled_unary_softplus; break;
-                case HTP_OP_UNARY_TANH:      task_func = unary_task_f32_tiled_unary_tanh;     break;
-                case HTP_OP_UNARY_ABS:       task_func = unary_task_f32_tiled_unary_abs;      break;
-                case HTP_OP_UNARY_LOG:       task_func = unary_task_f32_tiled_unary_log;      break;
-                case HTP_OP_UNARY_RELU:      task_func = unary_task_f32_tiled_unary_relu;     break;
-                case HTP_OP_TRI:             task_func = unary_task_f32_tiled_tri;            break;
-                default:                     break;
-            }
-        } else if (is_f16) {
-            switch (octx->op) {
-                case HTP_OP_NORM:            task_func = unary_task_f16_norm;                 break;
-                case HTP_OP_RMS_NORM:        task_func = unary_task_f16_rms_norm;             break;
-                case HTP_OP_SCALE:           task_func = unary_task_f16_scale;                break;
-                case HTP_OP_CLAMP:           task_func = unary_task_f16_clamp;                break;
-                case HTP_OP_SQR:             task_func = unary_task_f16_sqr;                  break;
-                case HTP_OP_SQRT:            task_func = unary_task_f16_sqrt;                 break;
-                case HTP_OP_L2_NORM:         task_func = unary_task_f16_l2_norm;              break;
-                case HTP_OP_UNARY_ABS:       task_func = unary_task_f16_unary_abs;            break;
-                case HTP_OP_UNARY_LOG:       task_func = unary_task_f16_unary_log;            break;
-                default:                     break;
-            }
-        } else {
-            switch (octx->op) {
-                case HTP_OP_NORM:            task_func = unary_task_f32_norm;                 break;
-                case HTP_OP_RMS_NORM:        task_func = unary_task_f32_rms_norm;             break;
-                case HTP_OP_RMS_NORM_MUL:    task_func = unary_task_f32_rms_norm_mul;         break;
-                case HTP_OP_SCALE:           task_func = unary_task_f32_scale;                break;
-                case HTP_OP_CLAMP:           task_func = unary_task_f32_clamp;                break;
-                case HTP_OP_LEAKY_RELU:      task_func = unary_task_f32_leaky_relu;           break;
-                case HTP_OP_SQR:             task_func = unary_task_f32_sqr;                  break;
-                case HTP_OP_SQRT:            task_func = unary_task_f32_sqrt;                 break;
-                case HTP_OP_UNARY_NEG:       task_func = unary_task_f32_unary_neg;            break;
-                case HTP_OP_UNARY_EXP:       task_func = unary_task_f32_unary_exp;            break;
-                case HTP_OP_UNARY_SIGMOID:   task_func = unary_task_f32_unary_sigmoid;        break;
-                case HTP_OP_UNARY_SILU:      task_func = unary_task_f32_unary_silu;           break;
-                case HTP_OP_UNARY_GELU:      task_func = unary_task_f32_unary_gelu;           break;
-                case HTP_OP_UNARY_SOFTPLUS:  task_func = unary_task_f32_unary_softplus;       break;
-                case HTP_OP_UNARY_TANH:      task_func = unary_task_f32_unary_tanh;           break;
-                case HTP_OP_UNARY_ABS:       task_func = unary_task_f32_unary_abs;            break;
-                case HTP_OP_UNARY_LOG:       task_func = unary_task_f32_unary_log;            break;
-                case HTP_OP_UNARY_RELU:      task_func = unary_task_f32_unary_relu;           break;
-                case HTP_OP_L2_NORM:         task_func = unary_task_f32_l2_norm;              break;
-                case HTP_OP_TRI:             task_func = unary_task_f32_tri;                  break;
-                default:                     break;
-            }
-        }
+    uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
+    struct htp_unary_context uctx = {
+        .octx                  = octx,
+        .kparams               = kparams,
+        .src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .src0_nrows            = nrows,
+        .row_start             = row_start,

-        if (task_func) {
-            work_queue_run(octx->ctx->work_queue, task_func, &uctx, n_threads);
-        } else {
-            FARF(ERROR, "execute_op_unary: task function is NULL for op %d\n", octx->op);
-            err = HTP_STATUS_NO_SUPPORT;
+        .data_src0             = src0->data,
+        .data_src1             = (octx->op == HTP_OP_RMS_NORM_MUL) ? src1->data : 0,
+        .data_dst              = dst->data,
+
+        .src0_data_row_size    = src0_data_row_size,
+        .src1_data_row_size    = src1_data_row_size,
+        .dst_data_row_size     = dst_data_row_size,
+
+        .src0_row_size_aligned = src0_row_size_aligned,
+        .src1_row_size_aligned = src1_row_size_aligned,
+        .dst_row_size_aligned  = dst_row_size_aligned,
+
+        .src0_vtcm_half_size   = kparams->vtcm_src0_size_per_thread / 2,
+        .src1_vtcm_half_size   = (octx->op == HTP_OP_RMS_NORM_MUL) ? (kparams->vtcm_src1_size_per_thread / (broadcast_weight ? 1 : 2)) : 0,
+        .dst_vtcm_half_size    = kparams->vtcm_dst_size_per_thread / 2,
+
+        .block                 = kparams->block,
+        .nc                    = src0->ne[0],
+        .col_tile              = col_tile,
+        .broadcast_weight      = broadcast_weight,
+
+        .vtcm_src0             = VTCM_LAYOUT_PTR(uint8_t, base, 0),
+        .vtcm_src1             = VTCM_LAYOUT_PTR_OPTIONAL(uint8_t, base, kparams->vtcm_src0_size, kparams->vtcm_src1_size > 0),
+        .vtcm_dst              = VTCM_LAYOUT_PTR(uint8_t, base, kparams->vtcm_src0_size + kparams->vtcm_src1_size),
+
+        .vtcm_src0_size_per_thread = kparams->vtcm_src0_size_per_thread,
+        .vtcm_src1_size_per_thread = kparams->vtcm_src1_size_per_thread,
+        .vtcm_dst_size_per_thread  = kparams->vtcm_dst_size_per_thread,
+    };
+
+    FARF(HIGH, "%s: %s mode (col_tile %u)\n", op_type, col_tile ? "tiled" : "row-block", col_tile);
+
+    worker_callback_t task_func = NULL;
+    void * compute_func = NULL;
+
+    if (col_tile) {
+        task_func = unary_thread_tiled;
+        switch (octx->op) {
+            case HTP_OP_SCALE:           compute_func = (void *) tile_scale_f32;          break;
+            case HTP_OP_CLAMP:           compute_func = (void *) tile_clamp_f32;          break;
+            case HTP_OP_LEAKY_RELU:      compute_func = (void *) tile_leaky_relu_f32;     break;
+            case HTP_OP_SQR:             compute_func = (void *) tile_sqr_f32;            break;
+            case HTP_OP_SQRT:            compute_func = (void *) tile_sqrt_f32;           break;
+            case HTP_OP_UNARY_NEG:       compute_func = (void *) tile_neg_f32;            break;
+            case HTP_OP_UNARY_EXP:       compute_func = (void *) tile_exp_f32;            break;
+            case HTP_OP_UNARY_SIGMOID:   compute_func = (void *) tile_sigmoid_f32;        break;
+            case HTP_OP_UNARY_SILU:      compute_func = (void *) tile_silu_f32;           break;
+            case HTP_OP_UNARY_GELU:      compute_func = (void *) tile_gelu_f32;           break;
+            case HTP_OP_UNARY_SOFTPLUS:  compute_func = (void *) tile_softplus_f32;       break;
+            case HTP_OP_UNARY_TANH:      compute_func = (void *) tile_tanh_f32;           break;
+            case HTP_OP_UNARY_ABS:       compute_func = (void *) tile_abs_f32;            break;
+            case HTP_OP_UNARY_LOG:       compute_func = (void *) tile_log_f32;            break;
+            case HTP_OP_UNARY_RELU:      compute_func = (void *) tile_relu_f32;           break;
+            case HTP_OP_TRI:
+                task_func    = unary_thread_tiled_tri_f32;
+                compute_func = (void *) tri_apply_tile_f32;
+                break;
+            default:                     break;
+        }
+    } else if (is_f16) {
+        task_func = unary_thread_row_block;
+        switch (octx->op) {
+            case HTP_OP_NORM:            compute_func = (void *) norm_f16;                break;
+            case HTP_OP_RMS_NORM:        compute_func = (void *) rms_norm_f16;            break;
+            case HTP_OP_SCALE:           compute_func = (void *) scale_f16;               break;
+            case HTP_OP_CLAMP:           compute_func = (void *) clamp_f16;               break;
+            case HTP_OP_SQR:             compute_func = (void *) sqr_f16;                 break;
+            case HTP_OP_SQRT:            compute_func = (void *) sqrt_f16;                break;
+            case HTP_OP_L2_NORM:         compute_func = (void *) l2_norm_f16;             break;
+            case HTP_OP_UNARY_ABS:       compute_func = (void *) abs_f16;                 break;
+            case HTP_OP_UNARY_LOG:       compute_func = (void *) log_f16;                 break;
+            default:                     break;
+        }
+    } else {
+        task_func = unary_thread_row_block;
+        switch (octx->op) {
+            case HTP_OP_NORM:            compute_func = (void *) norm_f32;                break;
+            case HTP_OP_RMS_NORM:        compute_func = (void *) rms_norm_f32;            break;
+            case HTP_OP_RMS_NORM_MUL:
+                task_func    = unary_thread_rms_norm_mul_f32;
+                compute_func = (void *) rms_norm_mul_f32;
+                break;
+            case HTP_OP_SCALE:           compute_func = (void *) scale_f32;               break;
+            case HTP_OP_CLAMP:           compute_func = (void *) clamp_f32;               break;
+            case HTP_OP_LEAKY_RELU:      compute_func = (void *) leaky_relu_f32;          break;
+            case HTP_OP_SQR:             compute_func = (void *) sqr_f32;                 break;
+            case HTP_OP_SQRT:            compute_func = (void *) sqrt_f32;                break;
+            case HTP_OP_UNARY_NEG:       compute_func = (void *) neg_f32;                 break;
+            case HTP_OP_UNARY_EXP:       compute_func = (void *) exp_f32;                 break;
+            case HTP_OP_UNARY_SIGMOID:   compute_func = (void *) sigmoid_f32;             break;
+            case HTP_OP_UNARY_SILU:      compute_func = (void *) silu_f32;                break;
+            case HTP_OP_UNARY_GELU:      compute_func = (void *) gelu_f32;                break;
+            case HTP_OP_UNARY_SOFTPLUS:  compute_func = (void *) softplus_f32;            break;
+            case HTP_OP_UNARY_TANH:      compute_func = (void *) tanh_f32;                break;
+            case HTP_OP_UNARY_ABS:       compute_func = (void *) abs_f32;                 break;
+            case HTP_OP_UNARY_LOG:       compute_func = (void *) log_f32;                 break;
+            case HTP_OP_UNARY_RELU:      compute_func = (void *) relu_f32;                break;
+            case HTP_OP_L2_NORM:         compute_func = (void *) l2_norm_f32;             break;
+            case HTP_OP_TRI:
+                task_func    = unary_thread_tri_f32;
+                compute_func = (void *) tri_f32;
+                break;
+            default:                     break;
         }
     }

+    if (!task_func || !compute_func) {
+        FARF(ERROR, "execute_op_unary: task function is NULL for op %d\n", octx->op);
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
+    uctx.compute = compute_func;
+    work_queue_run(octx->ctx->work_queue, task_func, &uctx, n_threads);
+
     return err;
 }

diff --git a/scripts/snapdragon/ggml-hexagon-inspect.py b/scripts/snapdragon/ggml-hexagon-inspect.py
new file mode 100755
index 000000000..3afda8a09
--- /dev/null
+++ b/scripts/snapdragon/ggml-hexagon-inspect.py
@@ -0,0 +1,1106 @@
+#!/usr/bin/env python3
+"""
+ggml-hexagon-inspect.py - Hexagon DSP binary inspection and diagnostic tool.
+
+Inspects Hexagon ELF binaries (libggml-htp-v*.so) for:
+  - Register spills (--spills): counts scalar and HVX vector stack spills,
+    separating in-loop spills from frame setup/teardown.
+  - Function disassembly (--disasm <func>): annotated disassembly showing
+    hardware loop bounds, packet boundaries, and spill instructions.
+  - Crash address resolution (--addr2line <addr...>): maps hex crash offsets
+    to function symbols, offsets, and source lines.
+  - CI verification (--strict): fails with non-zero exit if in-loop vector
+    spills or DMA worker vector instructions are detected.
+
+Usage:
+  # Check spills across all functions or specific operations
+  ./scripts/snapdragon/ggml-hexagon-inspect.py --spills
+  ./scripts/snapdragon/ggml-hexagon-inspect.py --spills --func "^compute_"
+  ./scripts/snapdragon/ggml-hexagon-inspect.py --spills --func "^compute_" --strict
+
+  # Disassemble a function with annotated loop and spill markers
+  ./scripts/snapdragon/ggml-hexagon-inspect.py --disasm compute_same_shape_div_f32
+
+  # Resolve crash addresses (CLI arguments or piped logcat/FARF logs)
+  ./scripts/snapdragon/ggml-hexagon-inspect.py --addr2line 0x51a30 0x5ba54
+  adb logcat | ./scripts/snapdragon/ggml-hexagon-inspect.py --addr2line
+"""
+
+import argparse
+import logging
+import os
+import platform
+import re
+import shutil
+import signal
+import subprocess
+import sys
+from pathlib import Path
+from typing import Dict, List, NamedTuple, Optional, Tuple
+
+# Ignore SIGPIPE to handle pipes (e.g. head, grep) gracefully
+if hasattr(signal, "SIGPIPE"):
+    signal.signal(signal.SIGPIPE, signal.SIG_DFL)
+
+logger = logging.getLogger("ggml-hexagon-inspect")
+
+
+class InsnInfo(NamedTuple):
+    address: int
+    asm_text: str
+    is_vec: bool
+    is_vspill: bool
+    is_sspill: bool
+    is_store: bool
+    is_load: bool
+    in_loop: bool
+
+
+class FuncStats:
+    def __init__(self, name: str, address: int, size: int):
+        self.name = name
+        self.address = address
+        self.size = size
+        self.packet_count = 0
+        self.insn_count = 0
+        self.vec_insn_count = 0
+        self.loop_count = 0
+        self.vspills_in_loop = 0
+        self.vspills_total = 0
+        self.sspills_in_loop = 0
+        self.sspills_total = 0
+        self.promotions_in_loop = 0
+        self.promotions_total = 0
+        self.promotion_targets: Dict[str, int] = {}
+        self.calls_in_loop = 0
+        self.calls_total = 0
+        self.insns: List[InsnInfo] = []
+
+
+class SymbolEntry(NamedTuple):
+    address: int
+    size: int
+    name: str
+
+
+# Regular expression patterns for Hexagon disassembly parsing
+RE_SYMBOL_HEADER = re.compile(r"^([0-9a-fA-F]+)\s+<([^>]+)>:", re.MULTILINE)
+RE_INSN_LINE = re.compile(
+    r"^\s*([0-9a-fA-F]+):\s+([0-9a-fA-F]{2}(?:\s+[0-9a-fA-F]{2}){3})\s+([0-9a-fA-F]{8})\s*(.*)$"
+)
+RE_LOOP0_START = re.compile(r"\bloop0\((0x[0-9a-fA-F]+)")
+RE_LOOP1_START = re.compile(r"\bloop1\((0x[0-9a-fA-F]+)")
+RE_VSPILL = re.compile(r"\bvmemu?\s*\(\s*r(?:29|30)\b")
+RE_SSPILL = re.compile(r"\bmem[bwhd]\s*\(\s*r(?:29|30)\b")
+RE_VEC_OP = re.compile(r"\b(v[0-9]+|w[0-9]+|q[0-3]|vmemu?)\b")
+RE_STORE = re.compile(r"=\s*(?:v[0-9]|r[0-9]|w[0-9]|#)")
+RE_PROMOTION_CALL = re.compile(
+    r"\b(?:call|jump)\s+(?:0x[0-9a-fA-F]+\s+)?<(__(?:trunc|extend)[a-zA-Z0-9_]+)(?:@plt)?>"
+)
+RE_ANY_CALL = re.compile(r"\bcallr?\b")
+
+
+def get_repo_root() -> Path:
+    # Resolve repository root from script location
+    return Path(__file__).resolve().parent.parent.parent
+
+
+def extract_arch_num(p: Path) -> int:
+    # Extract integer architecture version (e.g. v81 -> 81)
+    m = re.search(r"-v([0-9]+)\.so$", p.name)
+    return int(m.group(1)) if m else 0
+
+
+def find_default_lib(repo_root: Path, arch_filter: Optional[str] = None) -> Optional[Path]:
+    # Search for built Hexagon shared libraries in build and pkg directories
+    candidates = []
+    search_dirs = [
+        repo_root / "build-adb" / "ggml" / "src" / "ggml-hexagon",
+        repo_root / "build-android" / "ggml" / "src" / "ggml-hexagon",
+        repo_root / "build-ubuntu" / "ggml" / "src" / "ggml-hexagon",
+        repo_root / "build-linux" / "ggml" / "src" / "ggml-hexagon",
+        repo_root / "pkg-adb" / "llama.cpp" / "lib",
+        repo_root / "pkg-android" / "llama.cpp" / "lib",
+        repo_root / "pkg-ubuntu" / "llama.cpp" / "lib",
+    ]
+
+    arch_needle = None
+    if arch_filter:
+        arch_needle = arch_filter if arch_filter.startswith("v") else f"v{arch_filter}"
+
+    for d in search_dirs:
+        if not d.is_dir():
+            continue
+        for p in d.glob("libggml-htp-*.so"):
+            if arch_needle and arch_needle not in p.name:
+                continue
+            candidates.append(p)
+
+    if not candidates:
+        for p in repo_root.glob("build-*/ggml/src/ggml-hexagon/libggml-htp-*.so"):
+            if arch_needle and arch_needle not in p.name:
+                continue
+            candidates.append(p)
+
+    if not candidates:
+        return None
+
+    # Group latest build candidates (within 60s of max mtime) and pick highest arch
+    max_mtime = max(p.stat().st_mtime for p in candidates)
+    recent = [p for p in candidates if max_mtime - p.stat().st_mtime <= 60]
+    recent.sort(key=lambda p: extract_arch_num(p), reverse=True)
+    return recent[0]
+
+
+def translate_container_arg(arg: str, repo_root: Path) -> str:
+    # Do not translate non-path command flags
+    if arg.startswith("-") and "=" not in arg:
+        return arg
+    if arg.startswith("--") and "=" in arg:
+        flag, val = arg.split("=", 1)
+        return f"{flag}={translate_container_arg(val, repo_root)}"
+    try:
+        p = Path(arg)
+        if (p.is_absolute() and p.exists()) or (p.exists() and ("/" in arg or "\\" in arg)):
+            resolved = p.resolve()
+            if resolved.is_relative_to(repo_root):
+                rel = resolved.relative_to(repo_root)
+                return f"/workspace/{rel.as_posix()}"
+    except Exception:
+        pass
+    return arg
+
+
+class HexagonToolchain:
+    def __init__(
+        self,
+        repo_root: Path,
+        use_docker: bool = False,
+        image_url: str = "ghcr.io/snapdragon-toolchain",
+        image_name: str = "arm64-android",
+        image_ver: str = "v0.7",
+    ):
+        self.repo_root = repo_root
+        self.image = f"{image_url}/{image_name}:{image_ver}"
+        self.docker_bin = shutil.which("docker")
+        self.use_docker = use_docker
+
+        if not use_docker:
+            self.native_objdump, self.native_addr2line = self._discover_native_tools()
+        else:
+            self.native_objdump = None
+            self.native_addr2line = None
+
+        if not self.native_objdump and not self.native_addr2line:
+            self.use_docker = True
+
+    def _discover_native_tools(self) -> Tuple[Optional[str], Optional[str]]:
+        # Check system PATH
+        objdump = shutil.which("hexagon-llvm-objdump")
+        addr2line = shutil.which("hexagon-addr2line") or shutil.which("hexagon-llvm-addr2line")
+
+        # Check HEXAGON_TOOLS_ROOT environment variable
+        tools_root = os.environ.get("HEXAGON_TOOLS_ROOT")
+        if tools_root:
+            bin_dir = Path(tools_root) / "Tools" / "bin"
+            objdump_path = bin_dir / "hexagon-llvm-objdump"
+            addr2line_path = bin_dir / "hexagon-addr2line"
+            if objdump_path.is_file() and not objdump:
+                objdump = str(objdump_path)
+            if addr2line_path.is_file() and not addr2line:
+                addr2line = str(addr2line_path)
+
+        # Check HEXAGON_SDK_ROOT environment variable
+        sdk_root = os.environ.get("HEXAGON_SDK_ROOT")
+        if sdk_root:
+            tools_parent = Path(sdk_root) / "tools" / "HEXAGON_Tools"
+            if tools_parent.is_dir():
+                for t_dir in tools_parent.iterdir():
+                    bin_dir = t_dir / "Tools" / "bin"
+                    objdump_path = bin_dir / "hexagon-llvm-objdump"
+                    addr2line_path = bin_dir / "hexagon-addr2line"
+                    if objdump_path.is_file() and not objdump:
+                        objdump = str(objdump_path)
+                    if addr2line_path.is_file() and not addr2line:
+                        addr2line = str(addr2line_path)
+
+        return objdump, addr2line
+
+    def run_tool(self, tool_name: str, args: List[str], stdin_data: Optional[str] = None) -> str:
+        # Execute tool either natively or inside Docker container
+        if not self.use_docker:
+            tool_path = self.native_objdump if "objdump" in tool_name else self.native_addr2line
+            if not tool_path:
+                tool_path = shutil.which(tool_name)
+            if not tool_path:
+                raise RuntimeError(f"Tool {tool_name} not found natively. Use Docker instead.")
+
+            cmd = [tool_path] + args
+            res = subprocess.run(cmd, capture_output=True, text=True, input=stdin_data)
+            if res.returncode != 0:
+                raise RuntimeError(f"Tool {tool_name} failed: {res.stderr.strip()}")
+            return res.stdout
+
+        # Running via Docker container
+        if not self.docker_bin:
+            raise RuntimeError("Docker is required but not installed or found on PATH.")
+
+        container_tools_dir = "/opt/hexagon/6.6.0.0/tools/HEXAGON_Tools/19.0.07/Tools/bin"
+        if "objdump" in tool_name:
+            container_tool = f"{container_tools_dir}/hexagon-llvm-objdump"
+        elif "addr2line" in tool_name:
+            container_tool = f"{container_tools_dir}/hexagon-addr2line"
+        elif "nm" in tool_name:
+            container_tool = f"{container_tools_dir}/llvm-nm"
+        else:
+            container_tool = f"{container_tools_dir}/{tool_name}"
+
+        # Translate file paths from host to /workspace
+        translated_args = [translate_container_arg(arg, self.repo_root) for arg in args]
+
+        docker_cmd = [
+            "docker",
+            "run",
+            "--rm",
+            "--platform",
+            "linux/amd64",
+            "-v",
+            f"{self.repo_root}:/workspace",
+            "-w",
+            "/workspace",
+        ]
+
+        if platform.system() != "Windows":
+            docker_cmd += ["-u", f"{os.getuid()}:{os.getgid()}"]
+
+        docker_cmd += [self.image, container_tool] + translated_args
+
+        res = subprocess.run(docker_cmd, capture_output=True, text=True, input=stdin_data)
+        if res.returncode != 0:
+            raise RuntimeError(f"Docker command failed: {res.stderr.strip()}")
+        return res.stdout
+
+
+def parse_symbols(toolchain: HexagonToolchain, lib_path: Path) -> List[SymbolEntry]:
+    # Parse function symbols from library using objdump -t
+    output = toolchain.run_tool("hexagon-llvm-objdump", ["-t", str(lib_path)])
+    sym_re = re.compile(r"^([0-9a-fA-F]+)\s+[lgw! ]+\s+F\s+\.text\s+([0-9a-fA-F]+)\s+(.+)$")
+
+    symbols = []
+    for line in output.splitlines():
+        m = sym_re.match(line.strip())
+        if m:
+            addr = int(m.group(1), 16)
+            size = int(m.group(2), 16)
+            name = m.group(3).strip()
+            symbols.append(SymbolEntry(addr, size, name))
+
+    symbols.sort(key=lambda s: s.address)
+    return symbols
+
+
+def find_enclosing_symbol(symbols: List[SymbolEntry], address: int) -> Optional[Tuple[str, int]]:
+    # Binary search enclosing function symbol and compute offset
+    low = 0
+    high = len(symbols) - 1
+    best = None
+
+    while low <= high:
+        mid = (low + high) // 2
+        s = symbols[mid]
+        if s.address <= address:
+            if address < s.address + s.size:
+                return (s.name, address - s.address)
+            best = s
+            low = mid + 1
+        else:
+            high = mid - 1
+
+    if best and address < best.address + best.size:
+        return (best.name, address - best.address)
+    return None
+
+
+def parse_disassembly(
+    disasm_text: str, func_filter: Optional[re.Pattern] = None
+) -> List[FuncStats]:
+    # Parse disassembly text into structured function statistics
+    matches = list(RE_SYMBOL_HEADER.finditer(disasm_text))
+    funcs: List[FuncStats] = []
+
+    for i, m in enumerate(matches):
+        name = m.group(2)
+        if func_filter and not func_filter.search(name):
+            continue
+
+        addr = int(m.group(1), 16)
+        start_idx = m.end()
+        end_idx = matches[i + 1].start() if i + 1 < len(matches) else len(disasm_text)
+        chunk = disasm_text[start_idx:end_idx]
+
+        # Calculate rough byte size from line addresses
+        stats = FuncStats(name=name, address=addr, size=0)
+
+        loop0_target: Optional[int] = None
+        loop1_target: Optional[int] = None
+        loop0_active = False
+        loop1_active = False
+
+        first_addr = None
+        last_addr = None
+
+        for raw_line in chunk.splitlines():
+            lm = RE_INSN_LINE.match(raw_line)
+            if not lm:
+                continue
+
+            cur_addr = int(lm.group(1), 16)
+            asm_chunk = lm.group(4)
+
+            if first_addr is None:
+                first_addr = cur_addr
+            last_addr = cur_addr
+
+            # Track packet count
+            if "{" in asm_chunk:
+                stats.packet_count += 1
+
+            # Check loop starts
+            m0 = RE_LOOP0_START.search(asm_chunk)
+            if m0:
+                loop0_target = int(m0.group(1), 16)
+                stats.loop_count += 1
+
+            m1 = RE_LOOP1_START.search(asm_chunk)
+            if m1:
+                loop1_target = int(m1.group(1), 16)
+                stats.loop_count += 1
+
+            if loop0_target is not None and cur_addr >= loop0_target:
+                loop0_active = True
+            if loop1_target is not None and cur_addr >= loop1_target:
+                loop1_active = True
+
+            in_loop = loop0_active or loop1_active
+
+            # Parse instructions within packet line
+            cleaned = re.sub(r"[{}\s]|:endloop[01]", " ", asm_chunk)
+            sub_insns = [p.strip() for p in cleaned.split(";") if p.strip()]
+
+            for insn in sub_insns:
+                stats.insn_count += 1
+                is_vec = bool(RE_VEC_OP.search(insn))
+                if is_vec:
+                    stats.vec_insn_count += 1
+
+                is_vspill = bool(RE_VSPILL.search(insn))
+                is_sspill = bool(RE_SSPILL.search(insn))
+
+                # Identify store vs load
+                is_store = False
+                is_load = False
+                if is_vspill or is_sspill:
+                    if RE_STORE.search(insn):
+                        is_store = True
+                    else:
+                        is_load = True
+
+                if is_vspill:
+                    stats.vspills_total += 1
+                    if in_loop:
+                        stats.vspills_in_loop += 1
+                elif is_sspill:
+                    stats.sspills_total += 1
+                    if in_loop:
+                        stats.sspills_in_loop += 1
+
+                is_call = bool(RE_ANY_CALL.search(insn))
+                prom_m = RE_PROMOTION_CALL.search(insn)
+                if is_call:
+                    stats.calls_total += 1
+                    if in_loop:
+                        stats.calls_in_loop += 1
+                if prom_m:
+                    stats.promotions_total += 1
+                    ptarget = prom_m.group(1)
+                    stats.promotion_targets[ptarget] = stats.promotion_targets.get(ptarget, 0) + 1
+                    if in_loop:
+                        stats.promotions_in_loop += 1
+
+                stats.insns.append(
+                    InsnInfo(
+                        address=cur_addr,
+                        asm_text=insn,
+                        is_vec=is_vec,
+                        is_vspill=is_vspill,
+                        is_sspill=is_sspill,
+                        is_store=is_store,
+                        is_load=is_load,
+                        in_loop=in_loop,
+                    )
+                )
+
+            # Check loop ends
+            if ":endloop0" in asm_chunk:
+                loop0_active = False
+                loop0_target = None
+            if ":endloop1" in asm_chunk:
+                loop1_active = False
+                loop1_target = None
+
+        if first_addr is not None and last_addr is not None:
+            stats.size = (last_addr - first_addr) + 4
+
+        funcs.append(stats)
+
+    return funcs
+
+
+def annotate_disasm_line(
+    raw_line: str,
+    loop0_target: Optional[int],
+    loop1_target: Optional[int],
+    loop0_active: bool,
+    loop1_active: bool,
+    use_color: bool = True,
+) -> Tuple[str, Optional[int], Optional[int], bool, bool]:
+    # Annotate disassembly line with spill and loop tags
+    lm = RE_INSN_LINE.match(raw_line)
+    if not lm:
+        return raw_line, loop0_target, loop1_target, loop0_active, loop1_active
+
+    cur_addr = int(lm.group(1), 16)
+    asm_chunk = lm.group(4)
+
+    # Check loop starts
+    m0 = RE_LOOP0_START.search(asm_chunk)
+    if m0:
+        loop0_target = int(m0.group(1), 16)
+    m1 = RE_LOOP1_START.search(asm_chunk)
+    if m1:
+        loop1_target = int(m1.group(1), 16)
+
+    if loop0_target is not None and cur_addr >= loop0_target:
+        loop0_active = True
+    if loop1_target is not None and cur_addr >= loop1_target:
+        loop1_active = True
+
+    in_loop = loop0_active or loop1_active
+
+    tags = []
+    if m0:
+        tags.append("[LOOP0-START]")
+    if m1:
+        tags.append("[LOOP1-START]")
+
+    if RE_VSPILL.search(asm_chunk):
+        if in_loop:
+            tags.append("[V-SPILL:IN-LOOP]" if not use_color else "\033[1;31m[V-SPILL:IN-LOOP]\033[0m")
+        else:
+            tags.append("[V-SPILL]" if not use_color else "\033[1;33m[V-SPILL]\033[0m")
+    elif RE_SSPILL.search(asm_chunk):
+        if in_loop:
+            tags.append("[S-SPILL:IN-LOOP]" if not use_color else "\033[1;35m[S-SPILL:IN-LOOP]\033[0m")
+
+    prom_m = RE_PROMOTION_CALL.search(asm_chunk)
+    if prom_m:
+        ptarget = prom_m.group(1)
+        if in_loop:
+            tags.append(f"[PROMOTION:{ptarget}:IN-LOOP]" if not use_color else f"\033[1;31m[PROMOTION:{ptarget}:IN-LOOP]\033[0m")
+        else:
+            tags.append(f"[PROMOTION:{ptarget}]" if not use_color else f"\033[1;35m[PROMOTION:{ptarget}]\033[0m")
+    elif RE_ANY_CALL.search(asm_chunk):
+        if in_loop:
+            tags.append("[CALL:IN-LOOP]" if not use_color else "\033[1;31m[CALL:IN-LOOP]\033[0m")
+        else:
+            tags.append("[CALL]" if not use_color else "\033[1;36m[CALL]\033[0m")
+
+    if ":endloop0" in asm_chunk:
+        tags.append("[LOOP0-END]")
+        loop0_active = False
+        loop0_target = None
+    if ":endloop1" in asm_chunk:
+        tags.append("[LOOP1-END]")
+        loop1_active = False
+        loop1_target = None
+
+    tag_str = " ".join(tags)
+    if tag_str:
+        annotated = f"{raw_line:<80}  {tag_str}"
+    else:
+        annotated = raw_line
+
+    return annotated, loop0_target, loop1_target, loop0_active, loop1_active
+
+
+def run_spills(
+    toolchain: HexagonToolchain,
+    lib_path: Path,
+    args: argparse.Namespace,
+) -> int:
+    # Scan and report register spills across binary functions
+    logger.info(f"Inspecting library: {lib_path}")
+    disasm_text = toolchain.run_tool("hexagon-llvm-objdump", ["-d", str(lib_path)])
+
+    func_re = re.compile(args.func) if args.func else None
+    funcs = parse_disassembly(disasm_text, func_re)
+
+    # Filter functions
+    reported = []
+    for f in funcs:
+        has_spills = f.vspills_total > 0 or f.sspills_in_loop > 0 or f.sspills_total > 0
+        if args.all or args.func or has_spills:
+            reported.append(f)
+
+    # Sort: in-loop vector spills desc, then total vector spills desc, then in-loop scalar spills desc
+    reported.sort(
+        key=lambda x: (x.vspills_in_loop, x.vspills_total, x.sspills_in_loop, x.sspills_total),
+        reverse=True,
+    )
+
+    use_color = not args.no_color and sys.stdout.isatty()
+
+    # Print summary table
+    col_addr = "Address"
+    col_name = "Function"
+    col_pkts = "Packets"
+    col_insn = "Insns"
+    col_vec = "HVX Ops"
+    col_vloop = "V-Loop"
+    col_vtot = "V-Tot"
+    col_sloop = "S-Loop"
+    col_stot = "S-Tot"
+
+    hdr = (
+        f"{col_addr:<10} | {col_name:<44} | {col_pkts:>7} | {col_insn:>6} | "
+        f"{col_vec:>7} | {col_vloop:>6} | {col_vtot:>5} | {col_sloop:>6} | {col_stot:>5}"
+    )
+    sep = "-" * len(hdr)
+
+    logger.info("\n" + sep)
+    logger.info(hdr)
+    logger.info(sep)
+
+    tot_vloop = 0
+    tot_sloop = 0
+    tot_funcs_with_vloop = 0
+    strict_violations = []
+
+    dma_re: Optional[re.Pattern[str]] = re.compile(args.dma_pattern) if args.dma_pattern else None
+
+    for f in reported:
+        tot_vloop += f.vspills_in_loop
+        tot_sloop += f.sspills_in_loop
+        if f.vspills_in_loop > 0:
+            tot_funcs_with_vloop += 1
+
+        # Check strict criteria
+        if args.strict:
+            if f.vspills_in_loop > args.max_inloop_vspills:
+                strict_violations.append(
+                    f"{f.name}: {f.vspills_in_loop} in-loop vector spills (max allowed: {args.max_inloop_vspills})"
+                )
+            if dma_re and dma_re.search(f.name):
+                if f.vec_insn_count > args.max_dma_vec_ops:
+                    strict_violations.append(
+                        f"{f.name}: DMA worker contains {f.vec_insn_count} HVX vector ops (max allowed: {args.max_dma_vec_ops})"
+                    )
+
+        # Highlight in-loop vector spills
+        vloop_str = f"{f.vspills_in_loop:>6}"
+        if f.vspills_in_loop > 0 and use_color:
+            vloop_str = f"\033[1;31m{vloop_str}\033[0m"
+
+        logger.info(
+            f"0x{f.address:08x} | {f.name:<44} | {f.packet_count:>7} | {f.insn_count:>6} | "
+            f"{f.vec_insn_count:>7} | {vloop_str} | {f.vspills_total:>5} | {f.sspills_in_loop:>6} | {f.sspills_total:>5}"
+        )
+
+    logger.info(sep)
+    logger.info(
+        f"Total functions analyzed: {len(funcs)} | Reported: {len(reported)} | "
+        f"Functions with in-loop vector spills: {tot_funcs_with_vloop} | "
+        f"Total in-loop vector spills: {tot_vloop} | Total in-loop scalar spills: {tot_sloop}"
+    )
+
+    if args.strict:
+        logger.info("\n" + "=" * 50)
+        if strict_violations:
+            if use_color:
+                logger.error("\033[1;31mSTRICT CHECK FAILED\033[0m")
+            else:
+                logger.error("STRICT CHECK FAILED")
+            for v in strict_violations:
+                logger.error(f"  - {v}")
+            logger.info("=" * 50)
+            return 1
+        else:
+            if use_color:
+                logger.info("\033[1;32mSTRICT CHECK PASSED: 0 violations\033[0m")
+            else:
+                logger.info("STRICT CHECK PASSED: 0 violations")
+            logger.info("=" * 50)
+
+    return 0
+
+
+def run_promotions(
+    toolchain: HexagonToolchain,
+    lib_path: Path,
+    args: argparse.Namespace,
+) -> int:
+    # Scan and report soft-float promotion calls across binary functions
+    logger.info(f"Inspecting library: {lib_path}")
+    disasm_text = toolchain.run_tool("hexagon-llvm-objdump", ["-d", str(lib_path)])
+
+    func_re = re.compile(args.func) if args.func else None
+    funcs = parse_disassembly(disasm_text, func_re)
+
+    reported = []
+    for f in funcs:
+        if args.all or f.promotions_total > 0:
+            reported.append(f)
+
+    # Sort: in-loop promotions desc, then total promotions desc
+    reported.sort(
+        key=lambda x: (x.promotions_in_loop, x.promotions_total),
+        reverse=True,
+    )
+
+    use_color = not args.no_color and sys.stdout.isatty()
+
+    col_addr = "Address"
+    col_name = "Function"
+    col_loop = "Loops"
+    col_inloop = "In-Loop"
+    col_tot = "Total"
+    col_targets = "Promotion Targets"
+
+    hdr = f"{col_addr:<10} | {col_name:<44} | {col_loop:>5} | {col_inloop:>7} | {col_tot:>5} | {col_targets}"
+    sep = "-" * max(len(hdr), 110)
+
+    logger.info("\n" + sep)
+    logger.info(hdr)
+    logger.info(sep)
+
+    tot_inloop = 0
+    tot_prom = 0
+    tot_funcs_with_prom = 0
+    strict_violations = []
+
+    for f in reported:
+        tot_inloop += f.promotions_in_loop
+        tot_prom += f.promotions_total
+        if f.promotions_total > 0:
+            tot_funcs_with_prom += 1
+
+        if args.strict:
+            max_p = args.max_promotions if args.max_promotions is not None else 0
+            if f.promotions_total > max_p:
+                strict_violations.append(
+                    f"{f.name}: {f.promotions_total} float promotion calls (max allowed: {max_p})"
+                )
+
+        inloop_str = f"{f.promotions_in_loop:>7}"
+        if f.promotions_in_loop > 0 and use_color:
+            inloop_str = f"\033[1;31m{inloop_str}\033[0m"
+
+        targets_str = ", ".join(f"{t}: {c}" for t, c in sorted(f.promotion_targets.items()))
+        logger.info(
+            f"0x{f.address:08x} | {f.name:<44} | {f.loop_count:>5} | {inloop_str} | {f.promotions_total:>5} | {targets_str}"
+        )
+
+    logger.info(sep)
+    logger.info(
+        f"Total functions analyzed: {len(funcs)} | Reported: {len(reported)} | "
+        f"Functions with float promotions: {tot_funcs_with_prom} | "
+        f"Total promotion calls: {tot_prom} | In-loop: {tot_inloop}"
+    )
+
+    if args.strict:
+        logger.info("\n" + "=" * 50)
+        if strict_violations:
+            if use_color:
+                logger.error("\033[1;31mSTRICT CHECK FAILED\033[0m")
+            else:
+                logger.error("STRICT CHECK FAILED")
+            for v in strict_violations:
+                logger.error(f"  - {v}")
+            logger.info("=" * 50)
+            return 1
+        else:
+            if use_color:
+                logger.info("\033[1;32mSTRICT CHECK PASSED: 0 violations\033[0m")
+            else:
+                logger.info("STRICT CHECK PASSED: 0 violations")
+            logger.info("=" * 50)
+
+    return 0
+
+
+def run_disasm(
+    toolchain: HexagonToolchain,
+    lib_path: Path,
+    args: argparse.Namespace,
+) -> int:
+    # Disassemble matching function(s) with annotated loop and spill markers
+    func_pattern = args.disasm
+    logger.info(f"Inspecting library: {lib_path}")
+    logger.info(f"Disassembling functions matching: '{func_pattern}'\n")
+
+    # Disassemble symbol
+    disasm_text = toolchain.run_tool(
+        "hexagon-llvm-objdump",
+        ["-d", f"--disassemble-symbols={func_pattern}", str(lib_path)],
+    )
+
+    # If --disassemble-symbols yielded nothing (e.g. pattern was a regex), dump whole binary and filter
+    matches = list(RE_SYMBOL_HEADER.finditer(disasm_text))
+    if not matches:
+        all_disasm = toolchain.run_tool("hexagon-llvm-objdump", ["-d", str(lib_path)])
+        pat = re.compile(func_pattern)
+        all_matches = list(RE_SYMBOL_HEADER.finditer(all_disasm))
+        matched_symbols = [m.group(2) for m in all_matches if pat.search(m.group(2))]
+        if not matched_symbols:
+            logger.error(f"Error: No symbols found matching '{func_pattern}'.")
+            return 1
+        # Re-run with symbol list bounded by limit
+        sym_limit = args.limit if hasattr(args, "limit") and args.limit and args.limit > 0 else len(matched_symbols)
+        sym_arg = ",".join(matched_symbols[:sym_limit])
+        disasm_text = toolchain.run_tool(
+            "hexagon-llvm-objdump",
+            ["-d", f"--disassemble-symbols={sym_arg}", str(lib_path)],
+        )
+        matches = list(RE_SYMBOL_HEADER.finditer(disasm_text))
+
+    use_color = not args.no_color and sys.stdout.isatty()
+
+    # Parse and log annotated functions
+    for i, m in enumerate(matches):
+        name = m.group(2)
+        addr = int(m.group(1), 16)
+        start_idx = m.end()
+        end_idx = matches[i + 1].start() if i + 1 < len(matches) else len(disasm_text)
+        chunk = disasm_text[start_idx:end_idx]
+
+        # Parse statistics for this function
+        func_stats = parse_disassembly(disasm_text[m.start():end_idx])[0]
+
+        # Log header
+        hdr_border = "=" * 80
+        logger.info(hdr_border)
+        logger.info(f"Function: {name}")
+        logger.info(f"Address:  0x{addr:08x} - 0x{addr + func_stats.size:08x} ({func_stats.size} bytes)")
+        logger.info(f"Packets:  {func_stats.packet_count} | Instructions: {func_stats.insn_count} | Loops: {func_stats.loop_count}")
+        vec_pct = (func_stats.vec_insn_count / func_stats.insn_count * 100.0) if func_stats.insn_count else 0.0
+        logger.info(f"HVX Ops:  {func_stats.vec_insn_count} ({vec_pct:.1f}% of instructions)")
+        logger.info(
+            f"Spills:   Vector in-loop: {func_stats.vspills_in_loop} | Vector total: {func_stats.vspills_total} | "
+            f"Scalar in-loop: {func_stats.sspills_in_loop} | Scalar total: {func_stats.sspills_total}"
+        )
+        logger.info(
+            f"Calls:    Total: {func_stats.calls_total} (in-loop: {func_stats.calls_in_loop}) | "
+            f"Float promotions: {func_stats.promotions_total} (in-loop: {func_stats.promotions_in_loop})"
+        )
+        logger.info(hdr_border)
+
+        # Log annotated disassembly
+        loop0_target: Optional[int] = None
+        loop1_target: Optional[int] = None
+        loop0_active = False
+        loop1_active = False
+
+        for line in chunk.splitlines():
+            ann_line, loop0_target, loop1_target, loop0_active, loop1_active = annotate_disasm_line(
+                line, loop0_target, loop1_target, loop0_active, loop1_active, use_color
+            )
+            logger.info(ann_line)
+        logger.info("")
+
+    return 0
+
+
+def extract_addresses_from_input(lines: List[str]) -> List[int]:
+    # Extract hex program counter addresses from input lines
+    re_pc = re.compile(r"\b(?:pc|PC|ip|IP)\s*(?:=|:|\s)\s*0*(?:0x)?([0-9a-fA-F]{3,8})\b")
+    re_plus_hex = re.compile(r"\+0x([0-9a-fA-F]{3,8})\b")
+    re_hex = re.compile(r"\b0x([0-9a-fA-F]{3,8})\b")
+    re_bare_hex = re.compile(r"^\s*0*([0-9a-fA-F]{3,8})\s*$")
+
+    addrs = []
+    seen = set()
+
+    for line in lines:
+        matched = False
+        for m in re_pc.finditer(line):
+            val = int(m.group(1), 16)
+            if val not in seen:
+                seen.add(val)
+                addrs.append(val)
+            matched = True
+
+        if not matched:
+            for m in re_plus_hex.finditer(line):
+                val = int(m.group(1), 16)
+                if val not in seen:
+                    seen.add(val)
+                    addrs.append(val)
+                matched = True
+
+        if not matched:
+            for m in re_hex.finditer(line):
+                val = int(m.group(1), 16)
+                if val not in seen:
+                    seen.add(val)
+                    addrs.append(val)
+                matched = True
+
+        if not matched:
+            m = re_bare_hex.match(line)
+            if m:
+                val = int(m.group(1), 16)
+                if val not in seen:
+                    seen.add(val)
+                    addrs.append(val)
+
+    return addrs
+
+
+def run_addr2line(
+    toolchain: HexagonToolchain,
+    lib_path: Path,
+    args: argparse.Namespace,
+) -> int:
+    # Resolve addresses or crash logs to source locations and symbols
+    input_addrs: List[int] = []
+
+    if args.addr2line:
+        for arg in args.addr2line:
+            if arg == "-":
+                continue
+            try:
+                val = int(arg, 16)
+                input_addrs.append(val)
+            except ValueError:
+                # Treat as text line and search for hex addresses
+                input_addrs.extend(extract_addresses_from_input([arg]))
+
+    # Read from stdin if piped or requested via '-'
+    if not sys.stdin.isatty() or "-" in (args.addr2line or []):
+        stdin_lines = sys.stdin.readlines()
+        input_addrs.extend(extract_addresses_from_input(stdin_lines))
+
+    if not input_addrs:
+        logger.error("Error: No addresses found to resolve. Provide hex addresses or pipe crash logs to stdin.")
+        logger.error("Example: ./scripts/snapdragon/ggml-hexagon-inspect.py --addr2line 0x51a30 0x5ba54")
+        return 1
+
+    logger.info(f"Resolving {len(input_addrs)} address(es) against: {lib_path}\n")
+
+    # Load symbol table for symbol + offset fallback
+    symbols = parse_symbols(toolchain, lib_path)
+
+    # Format addresses for addr2line tool (prefixed with 0x)
+    addr_strs = [f"0x{a:x}" for a in input_addrs]
+    tool_args = ["-e", str(lib_path), "-f", "-C", "-p", "-a"] + addr_strs
+
+    raw_output = toolchain.run_tool("hexagon-addr2line", tool_args)
+
+    # Parse output lines
+    # Format: 0x51a30: binary_thread_add_id_f32 at /path/file.c:123
+    re_out = re.compile(r"^(0x[0-9a-fA-F]+):\s+(.*?)\s+at\s+(.*)$")
+
+    for line in raw_output.splitlines():
+        line = line.strip()
+        if not line:
+            continue
+        m = re_out.match(line)
+        if m:
+            addr_hex = m.group(1)
+            addr_val = int(addr_hex, 16)
+            func_name = m.group(2)
+            src_loc = m.group(3)
+
+            # Check if function name is unknown or generic, look up symbol table
+            sym_info = find_enclosing_symbol(symbols, addr_val)
+            if sym_info:
+                sym_name, sym_offset = sym_info
+                sym_display = f"{sym_name}+0x{sym_offset:x}"
+            else:
+                sym_display = func_name
+
+            logger.info(f"{addr_hex:<12} -> {sym_display:<40} ({src_loc})")
+        else:
+            logger.info(line)
+
+    return 0
+
+
+def main():
+    parser = argparse.ArgumentParser(
+        description="Inspect Hexagon DSP binaries for register spills, function disassembly, and crash analysis."
+    )
+
+    # Target library
+    parser.add_argument(
+        "--lib",
+        help="Path to Hexagon shared library (e.g. libggml-htp-v81.so). Auto-detected if omitted.",
+    )
+    parser.add_argument(
+        "--arch",
+        help="Architecture version filter for auto-detection (e.g. v75, v79, v81).",
+    )
+
+    # Modes
+    parser.add_argument(
+        "--spills",
+        action="store_true",
+        help="Scan binary and report scalar/vector stack spills table.",
+    )
+    parser.add_argument(
+        "--promotions",
+        action="store_true",
+        help="Scan binary and report functions with soft-float promotion calls (__trunc*, __extend*).",
+    )
+    parser.add_argument(
+        "--disasm",
+        metavar="FUNC",
+        help="Disassemble function symbol or regex pattern with annotated loop and spill markers.",
+    )
+    parser.add_argument(
+        "--limit",
+        type=int,
+        default=20,
+        help="Maximum symbols to disassemble when using pattern in --disasm (default: 20, 0 for unlimited).",
+    )
+    parser.add_argument(
+        "--addr2line",
+        nargs="*",
+        metavar="ADDR",
+        help="Resolve hex addresses or piped crash traces to symbols and source lines.",
+    )
+
+    # Filtering & Display
+    parser.add_argument(
+        "--func",
+        "-f",
+        help="Regex filter for function names in --spills or --promotions.",
+    )
+    parser.add_argument(
+        "--all",
+        "-a",
+        action="store_true",
+        help="Show all functions in table, even those with 0 spills/promotions.",
+    )
+    parser.add_argument(
+        "--no-color",
+        action="store_true",
+        help="Disable ANSI color output.",
+    )
+
+    # Strict check options
+    parser.add_argument(
+        "--strict",
+        action="store_true",
+        help="CI mode: exit with non-zero status if violations (in-loop vector spills, DMA worker vector ops) occur.",
+    )
+    parser.add_argument(
+        "--max-inloop-vspills",
+        type=int,
+        default=0,
+        help="Maximum allowed in-loop vector spills in --strict mode (default: 0).",
+    )
+    parser.add_argument(
+        "--max-dma-vec-ops",
+        type=int,
+        default=0,
+        help="Maximum allowed vector instructions in DMA workers in --strict mode (default: 0).",
+    )
+    parser.add_argument(
+        "--max-promotions",
+        type=int,
+        default=None,
+        help="Maximum allowed float promotion calls in --strict mode (default: 0).",
+    )
+    parser.add_argument(
+        "--dma-pattern",
+        default=r"^.*_thread(?:_.*)?$",
+        help="Regex pattern identifying DMA worker functions (default: '^.*_thread(?:_.*)?$').",
+    )
+
+    # Toolchain options
+    parser.add_argument(
+        "--docker",
+        action="store_true",
+        help="Force execution inside Docker container.",
+    )
+    parser.add_argument(
+        "--no-docker",
+        action="store_true",
+        help="Force native execution on host instead of Docker.",
+    )
+    parser.add_argument(
+        "--toolchain-version",
+        default="v0.7",
+        help="Docker toolchain tag (default: v0.7).",
+    )
+    parser.add_argument(
+        "--toolchain-url",
+        default="ghcr.io/snapdragon-toolchain",
+        help="Docker toolchain registry (default: ghcr.io/snapdragon-toolchain).",
+    )
+    parser.add_argument(
+        "--image-name",
+        default="arm64-android",
+        help="Docker toolchain image name (default: arm64-android).",
+    )
+
+    args = parser.parse_args()
+
+    logging.basicConfig(level=logging.INFO, format="%(message)s")
+
+    repo_root = get_repo_root()
+
+    # Determine target library
+    lib_path = None
+    if args.lib:
+        lib_path = Path(args.lib).resolve()
+        if not lib_path.is_file():
+            logger.error(f"Error: Specified library '{args.lib}' does not exist.")
+            sys.exit(1)
+    else:
+        lib_path = find_default_lib(repo_root, args.arch)
+        if not lib_path:
+            logger.error("Error: No Hexagon library found in build-* or pkg-* directories.")
+            logger.error("Build the project first via ./scripts/snapdragon/build.py --target adb or specify --lib.")
+            sys.exit(1)
+
+    # Initialize toolchain wrapper
+    use_docker = args.docker or (not args.no_docker and platform.system() == "Darwin")
+    try:
+        toolchain = HexagonToolchain(
+            repo_root=repo_root,
+            use_docker=use_docker,
+            image_url=args.toolchain_url,
+            image_name=args.image_name,
+            image_ver=args.toolchain_version,
+        )
+    except Exception as e:
+        logger.error(f"Error initializing toolchain: {e}")
+        sys.exit(1)
+
+    # Dispatch commands
+    if args.addr2line is not None:
+        sys.exit(run_addr2line(toolchain, lib_path, args))
+    elif args.disasm:
+        sys.exit(run_disasm(toolchain, lib_path, args))
+    elif args.promotions:
+        sys.exit(run_promotions(toolchain, lib_path, args))
+    else:
+        # Default action is --spills
+        sys.exit(run_spills(toolchain, lib_path, args))
+
+
+if __name__ == "__main__":
+    logging.basicConfig(level=logging.INFO, format="%(message)s")
+    main()
diff --git a/scripts/snapdragon/run.py b/scripts/snapdragon/run.py
index dc71d4a32..8917febc1 100755
--- a/scripts/snapdragon/run.py
+++ b/scripts/snapdragon/run.py
@@ -21,6 +21,7 @@ MANAGED_ENV_NAMES = (
     "GGML_HEXAGON_NHVX",
     "GGML_HEXAGON_NHMX",
     "GGML_HEXAGON_HOSTBUF",
+    "GGML_HEXAGON_DMA64",
     "GGML_HEXAGON_OPBATCH",
     "GGML_HEXAGON_OPQUEUE",
     "GGML_HEXAGON_OPPOLL",
@@ -155,6 +156,7 @@ def main():
     parser.add_argument("--hex-nhvx", help="Number of HVX units to use (GGML_HEXAGON_NHVX)")
     parser.add_argument("--hex-nhmx", help="Number of HMX units to use. 0 disables HMX power-up (GGML_HEXAGON_NHMX)")
     parser.add_argument("--hex-hostbuf", help="Enable host buffers (GGML_HEXAGON_HOSTBUF)")
+    parser.add_argument("--hex-dma64", nargs="?", const="1", help="Enable (1) or disable (0) 64-bit DMA for model weights (GGML_HEXAGON_DMA64)")
     parser.add_argument("--hex-opbatch", help="Maximum number of operations to batch into a single HTP execution (GGML_HEXAGON_OPBATCH)")
     parser.add_argument("--hex-opqueue", help="Size of the asynchronous NPU operation queue (GGML_HEXAGON_OPQUEUE)")
     parser.add_argument("--hex-oppoll", default="1", help="Enable (1) or Disable (0) polling for NPU opbatch completion (GGML_HEXAGON_OPPOLL) (default: 1)")
@@ -162,7 +164,7 @@ def main():
     parser.add_argument("--hex-opfusion", help="NPU graph node fusion optimization level (0: disabled, 1: enabled) (GGML_HEXAGON_OPFUSION)")
     parser.add_argument("--hex-vmem", help="Maximum NPU VMEM size limit in MB to allocate (GGML_HEXAGON_VMEM)")
     parser.add_argument("--hex-mbuf", help="Maximum host buffer size limit in MB to allocate (GGML_HEXAGON_MBUF)")
-    parser.add_argument("--hex-mm-select", help="Select MUL_MAT and MUL_MAT_ID kernel (GGML_HEXAGON_MM_SELECT) 3:HMX,2:HVX-tiled,1:HVX-flat,0:disable")
+    parser.add_argument("--hex-mm-select", help="Select MUL_MAT and MUL_MAT_ID kernel (GGML_HEXAGON_MM_SELECT) 2:HMX,1:HVX,0:disable")
     parser.add_argument("--hex-fa-select", help="Select Flash Attention kernel (GGML_HEXAGON_FA_SELECT) 2:HMX,1:HVX,0:disable")
     parser.add_argument("--hex-ar-select", help="Select All-Reduce kernel (GGML_HEXAGON_AR_SELECT) 1:enable,0:disable")
     parser.add_argument("--hex-etm", help="Enable Embedded Trace Macrocell hardware tracing / trace logging (GGML_HEXAGON_ETM)")
@@ -294,6 +296,7 @@ def main():
     set_env("GGML_HEXAGON_NHVX", args.hex_nhvx)
     set_env("GGML_HEXAGON_NHMX", args.hex_nhmx)
     set_env("GGML_HEXAGON_HOSTBUF", args.hex_hostbuf)
+    set_env("GGML_HEXAGON_DMA64", args.hex_dma64)
     set_env("GGML_HEXAGON_OPBATCH", args.hex_opbatch)
     set_env("GGML_HEXAGON_OPQUEUE", args.hex_opqueue)
     set_env("GGML_HEXAGON_OPPOLL", args.hex_oppoll)
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index c75cb3c0f..5c7de196d 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -10348,6 +10348,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {200001, 2, 3, 1}, true,   true,  GGML_TYPE_F16, {1, 1}, 0.1f, 8.0f));
     test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {200000, 1, 1, 1}, false,  false, GGML_TYPE_F32, {1, 1}, 1.0f, 0.0f));
     test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {200000, 4, 1, 1}, false,  false, GGML_TYPE_F32, {1, 1}, 1.0f, 0.0f));
+    test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {4,      1, 1, 1}, false,  false, GGML_TYPE_F32, {1, 1}, 1.0f, 0.0f));
+    test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {4,   1023, 1, 1}, false,  false, GGML_TYPE_F32, {1, 1}, 1.0f, 0.0f));
     test_cases.emplace_back(new test_soft_max(GGML_TYPE_F32, {643251, 3, 1, 1}, false,  false, GGML_TYPE_F32, {1, 1}, 1.0f, 0.0f));

     for (float max_bias : {0.0f, 8.0f}) {