Commit 8345f3339 for llama.cpp

commit 8345f333951c661d166b00e6f9362e553768f292
Author: Max Krasnyansky <maxk@qti.qualcomm.com>
Date:   Mon Oct 5 08:55:21 2026 -0700

    hexagon: matmul and flash-atten scalability updates (#29974)

    * hexagon: head-parallel flash_attn partitioning for row-split multicore

    In row-split mode each core computes its output row shard of every
    MUL_MAT, but flash_attn was previously partitioning by Q tokens
    (flat qrow split) instead of by heads. This forced every core to
    read the full KV cache (all n_kv_heads), negating the memory
    bandwidth benefit of multicore on flash_attn.

    Change both HMX and HVX flash_attn kernels to partition by KV heads
    when n_kv_heads is divisible by n_cores: core i processes heads
    [i*n_kv_heads/N, (i+1)*n_kv_heads/N) exclusively, reading only its
    head shard of the KV cache. Falls back to the original token-block
    split when n_kv_heads % n_cores != 0 (e.g. Gemma-4 with 2 KV heads
    on 4 cores).

    Controlled by GGML_HEXAGON_FA_HEAD_SPLIT (default 1 = on).
    The flag is packed into bit 1 of the existing is_dst_fp32 kparams
    byte to stay within the 128-byte kernel_params blob limit.

    Measured gains at 4c row-split (PP t/s, ubatch=1024):
      Qwen3-0.6B:    6977 -> 11026  (+58%)
      llama-3.2-3B:  3717 ->  5522  (+49%)
      Qwen3.5-4B:    2739 ->  2855   (+4%)
      Gemma-4 MoE:   no change (MoE FFN dominates, fallback path)

    TG is unchanged (flash_attn is a small fraction of decode time
    relative to the matmul+barrier cost per layer).

    * hex-fa: cleanup kern_params and head-split selection

    * hex-fa: add -fa-head-split option to run.py

    * hex-mdev: update matmul solver to account for reduced work in row-split scenarios

    * hex-mmid: better work splitting by expers in multi-dev scenarios

    * hex-fa: update HMX gating based on the model/n-hvx/ctx-len sweep

    * hex-fa: precompute softcap/scale on the host

    * hexagon: flatten matmul into 2d to use HMX in multi-sequence

    * hex-mm: cleanup kparams and use collapse to 3/4D -> 2D mapping

    * hex-mm: fix typo in collapse fallback

    * hex-mm: another pass at consistent naming for act tensors

    * hex-mm: add support for colapsing dims in fused matmuls

    * hex-build: fix WoS build errors

    * hex-mm: make sure to enforce dst stride in can_collapse

    * hex-fa: add a onliner commit for head-split check

    * hex-fa: remove unused local head_split var

    * hex-fa: tighten up can_split checks

    * hex-mm: update unfused paths to use act instead src1

    * hex-mm: make sure to check all dsts for splitting

    * hexagon: fix the second weight chunk address in the batched HMX matmul prologue

    * hexagon: F16 activation and ragged N in the HMX matmul

    * hex-mm: tighten the ragged/split checks in mdev cases

    * hex-mm: enable MM fusion for F16 activations

    * hex-mm: pass tiled sizes to the solver in fused paths

    * hex-mmid: remove scalar divs from expert mapping loops

    * hex-mmid: proper cacheline safety enforcement for mdev splits

    * hex-mm: improve solver for mdev split scanarios and tail handling

    * hex-mm: remove redundant checks

    * hex-mm: fix fused HMX MUL_MAT_NX drops the final partial tile for quantized weights

    * hex-mm: better handling of ragged shapes (removes scalar memset of vtcm)

    ---------

    Co-authored-by: ebateni <ebateni@qti.qualcomm.com>
    Co-authored-by: Jhen-Jie Hong <iainst0409@gmail.com>
    Co-authored-by: Yiwei Shao <yiwei@aizip.ai>

diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 454594472..2282a04a8 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -1,3 +1,5 @@
+#define _USE_MATH_DEFINES
+
 #include <assert.h>
 #include <inttypes.h>
 #include <stdio.h>
@@ -24,6 +26,10 @@
 #include <cmath>
 #include <initializer_list>

+#ifndef M_LOG2E
+#    define M_LOG2E 1.44269504088896340736
+#endif
+
 #ifdef _WIN32
 #    define WIN32_LEAN_AND_MEAN
 #    ifndef NOMINMAX
@@ -101,6 +107,7 @@ static bool   opt_dma64   = false;

 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_fa_head_split = 1; // 1 = partition flash_attn by KV heads in multicore (default on), 0 = token-based (original)
 static int    opt_gdn_select = 2; // 2 = HMX -> HVX, 1 = HVX, 0 = CPU (unsupported)
 static int    opt_ar_select  = 2; // 2 = fused ALLREDUCE+ADD (default), 1 = unfused ALLREDUCE, 0 = fallback to CPY+FENCE
 static int    opt_ar_scatter = 1; // 1 = reduce-scatter the fused ALLREDUCE+ADD (default), 0 = full reduction
@@ -387,7 +394,7 @@ static void ggml_hexagon_precompute_sort_params(
 static void ggml_hexagon_precompute_fused_mmnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
-    const struct ggml_tensor * src1,
+    const struct ggml_tensor * act,
     int32_t n_weights,
     struct htp_mm_kernel_params * kparams
 );
@@ -395,7 +402,7 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
 static void ggml_hexagon_precompute_fused_mmidnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
-    const struct ggml_tensor * src1,
+    const struct ggml_tensor * act,
     const struct ggml_tensor * dst,
     int32_t n_weights,
     struct htp_mm_kernel_params * kparams
@@ -412,6 +419,10 @@ static bool ggml_hexagon_precompute_allreduce_params(
     struct htp_allreduce_kernel_params * kparams
 );

+static bool ggml_hexagon_rows_stride(const int64_t * ne, const size_t * nb, size_t * stride);
+static bool ggml_hexagon_matmul_can_collapse(const struct ggml_tensor * src0, const struct ggml_tensor * src1, const struct ggml_tensor * dst);
+static ggml_tensor ggml_hexagon_tensor_collapse_rows(const struct ggml_tensor * t);
+
 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);
@@ -444,6 +455,13 @@ static inline bool ggml_hexagon_tensors_overlap(const struct ggml_tensor * a, co
     return a0 < b1 && b0 < a1;
 }

+static inline bool ggml_hexagon_can_row_partition(const struct ggml_tensor * t) {
+    if (t->ne[1] > 1 && (t->nb[1] & 127) != 0) return false;
+    if (t->ne[2] > 1 && (t->nb[2] & 127) != 0) return false;
+    if (t->ne[3] > 1 && (t->nb[3] & 127) != 0) return false;
+    return true;
+}
+
 struct htp_opnode;

 struct ggml_hexagon_opbatch;
@@ -3410,6 +3428,10 @@ struct ggml_hexagon_opbatch {
             return false;
         }

+        if (orig_kparams->collapse != kparams.collapse) {
+            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;
@@ -3493,8 +3515,20 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

+            const struct htp_mm_kernel_params * orig_kparams = (const struct htp_mm_kernel_params *) last_node.kernel_params;
+            const bool collapse = orig_kparams->collapse && ggml_hexagon_matmul_can_collapse(w_in, x, d_in);
+            if (orig_kparams->collapse && !collapse) {
+                return false;
+            }
+
             struct htp_mm_kernel_params kparams;
-            ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams);
+            if (collapse) {
+                const ggml_tensor x_collapsed = ggml_hexagon_tensor_collapse_rows(x);
+                ggml_hexagon_precompute_fused_mmnx_params(sess, w0, &x_collapsed, curr_n + 1, &kparams);
+                kparams.collapse = 1;
+            } else {
+                ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams);
+            }
             if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
                 return false;
             }
@@ -3543,9 +3577,23 @@ struct ggml_hexagon_opbatch {
             const ggml_tensor * w0 = last_node.src0();
             const ggml_tensor * x  = last_node.src1();
             const ggml_tensor * w1 = node.src0();
+            const ggml_tensor * dst_0 = last_node.dst();
+            const ggml_tensor * dst_1 = node.dst();
+
+            const struct htp_mm_kernel_params * orig_kparams = (const struct htp_mm_kernel_params *) last_node.kernel_params;
+            const bool collapse = orig_kparams->collapse && ggml_hexagon_matmul_can_collapse(w1, x, dst_1);
+            if (orig_kparams->collapse && !collapse) {
+                return false;
+            }

             struct htp_mm_kernel_params kparams;
-            ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
+            if (collapse) {
+                const ggml_tensor x_collapsed = ggml_hexagon_tensor_collapse_rows(x);
+                ggml_hexagon_precompute_fused_mmnx_params(sess, w0, &x_collapsed, 2, &kparams);
+                kparams.collapse = 1;
+            } else {
+                ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
+            }
             if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
                 return false;
             }
@@ -3559,9 +3607,6 @@ struct ggml_hexagon_opbatch {
                 return false;
             }

-            const ggml_tensor * dst_0 = last_node.dst();
-            const ggml_tensor * dst_1 = node.dst();
-
             last_node.opcode = HTP_OP_MUL_MAT_NX;
             last_node.name   = "MUL_MAT_NX";
             last_node.inputs.clear();
@@ -4831,15 +4876,48 @@ static bool ggml_hexagon_flash_attn_is_hmx_eligible(
         return false;
     }

-    // Fall back to HVX for small token counts if head dimension is small (DK <= 128)
-    const uint32_t neq1 = q->ne[1];
-    if (DK <= 128 && neq1 < 5) {
-        return false;
+    GGML_UNUSED(sinks);
+
+    // Explicit force mode
+    if (opt_fa_select > 2) {
+        return true;
     }

-    return true;
+    const uint32_t M = q->ne[1];

-    GGML_UNUSED(sinks);
+    // Prefill or batched decode
+    if (M > 1) {
+        return true;
+    }
+
+    // Compute-bound head dim
+    if (DK >= 256) {
+        return true;
+    }
+
+    const uint32_t n_head = q->ne[2];
+    const uint32_t n_kv_heads = k->ne[2];
+    const uint32_t G = n_kv_heads > 0 ? n_head / n_kv_heads : 1;
+    const uint32_t S = k->ne[1];
+
+    // Tile alignment for 32-row HMX tiles
+    const bool is_tile_aligned = (G > 0 && (32 % G == 0));
+    if (!is_tile_aligned) {
+        if (sess->n_threads >= 6) {
+            return false;
+        }
+        return S >= 1024;
+    }
+
+    // Context depth crossover
+    uint32_t s_cross = 512;
+    if (DK <= 64) {
+        s_cross = (sess->n_threads >= 8) ? 2048 : ((sess->n_threads >= 6) ? 768 : 512);
+    } else {
+        s_cross = (sess->n_threads >= 8) ? 1024 : 512;
+    }
+
+    return S >= s_cross;
 }

 static bool ggml_hexagon_precompute_flash_attn_params(
@@ -4857,7 +4935,6 @@ static bool ggml_hexagon_precompute_flash_attn_params(
     const struct ggml_tensor * k    = op->src[1];
     const struct ggml_tensor * v    = op->src[2];
     const struct ggml_tensor * mask = op->src[3];
-    const struct ggml_tensor * dst  = op;

     const uint32_t neq0 = q->ne[0];  // head_dim (DK)
     const uint32_t neq1 = q->ne[1];  // n_tokens
@@ -4888,8 +4965,8 @@ static bool ggml_hexagon_precompute_flash_attn_params(
     kparams->max_bias = max_bias;
     kparams->logit_softcap = logit_softcap;

-    kparams->is_q_fp32 = (q->type == GGML_TYPE_F32) ? 1 : 0;
-    kparams->is_dst_fp32 = (dst->type == GGML_TYPE_F32) ? 1 : 0;
+    kparams->head_split = (opt_fa_head_split != 0) ? 1 : 0;
+    kparams->flags      = 0;
     kparams->G = G;

     const uint32_t n_head = q->ne[2];
@@ -4906,9 +4983,16 @@ 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, sinks != nullptr, n_head);
+        int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, (q->type == GGML_TYPE_F32), sinks != nullptr, n_head);
         if (ret == 0) {
             kparams->kernel_type = HTP_FA_KERNEL_HMX;
+            if (logit_softcap == 0.0f) {
+                kparams->scale = scale * (float) M_LOG2E;
+                kparams->logit_softcap = 0.0f;
+            } else {
+                kparams->scale = scale;
+                kparams->logit_softcap = logit_softcap * (float) M_LOG2E;
+            }
             kparams->Br = Br;
             kparams->Bc = Bc;
             kparams->n_kv_blocks = (nek1 + Bc - 1) / Bc;
@@ -4916,7 +5000,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, sinks != nullptr, n_head);
+            kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, (q->type == GGML_TYPE_F32), 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
@@ -4943,11 +5027,11 @@ static bool ggml_hexagon_precompute_flash_attn_params(
     kparams->n_kv_blocks = (k->ne[1] + 64 - 1) / 64;
     kparams->n_threads = sess->n_threads;

-    const size_t size_q_row_padded = hex_round_up(q->ne[0] * (kparams->is_q_fp32 ? 4 : 2), 128);
+    const size_t size_q_row_padded = hex_round_up(q->ne[0] * ((q->type == GGML_TYPE_F32) ? 4 : 2), 128);
     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, sinks != nullptr, n_head, sess->n_threads);
+    kparams->vtcm_size = hvx_fa_compute_vtcm_usage(DK, DV, (q->type == GGML_TYPE_F32), 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;
@@ -5102,7 +5186,7 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
     bool is_matmul_id,
     bool is_batched
 ) {
-    if (src1->type != GGML_TYPE_F32) {
+    if (src1->type != GGML_TYPE_F32 && (src1->type != GGML_TYPE_F16 || is_matmul_id)) {
         return false;
     }

@@ -5111,8 +5195,8 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
     const int ne12  = src1->ne[2];
     const int wtype = src0->type;

-    // HMX weight tile requires N to be 32-aligned.
-    if (ne01_padded % 32 != 0) {
+    // HMX weight tiles accept non-32-aligned N for non-matmul_id.
+    if (ne01_padded % 32 != 0 && is_matmul_id) {
         return false;
     }

@@ -5151,7 +5235,7 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
 static bool ggml_hexagon_precompute_hmx_mm_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
-    const struct ggml_tensor * src1,
+    const struct ggml_tensor * act,
     const struct ggml_tensor * dst,
     int wtype,
     int ne00_padded,
@@ -5167,9 +5251,22 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     struct htp_mm_kernel_params * kparams
 ) {
     const int aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype);
-    const bool pipeline = is_matmul_id ? false : htp_mm_hmx_pipeline(ne11);
     const int n_threads = (int)sess->n_threads;
-    const int ne10 = src1->ne[0];
+    const int ne10 = act->ne[0];
+
+    int m_for_solver = ne11;
+    int m_for_solver_padded = ne11_padded;
+    // matmul_id partitions by expert; regular matmul partitions M rows (ne11) across devices
+    if (!is_matmul_id && sess->mdev.count > 1 && ((uint32_t) ne11 >= sess->mdev.count)) {
+        // when dst is null, padded dims are used for estimate which are 128-byte aligned
+        const bool dst_row_split = dst ? ggml_hexagon_can_row_partition(dst) : true;
+        const bool act_row_split = ggml_hexagon_can_row_partition(act);
+        if (dst_row_split && act_row_split) {
+            m_for_solver = (ne11 + (int) sess->mdev.count - 1) / (int) sess->mdev.count;
+            m_for_solver_padded = hex_round_up(std::max(m_for_solver, 32), 32);
+        }
+    }
+    const bool pipeline = is_matmul_id ? false : htp_mm_hmx_pipeline(m_for_solver);

     const bool is_batched_val = is_matmul_id ? false : is_batched;
     const int group_size = (ne02 > 0 ? ne12 / ne02 : 1);
@@ -5182,15 +5279,25 @@ static bool ggml_hexagon_precompute_hmx_mm_params(

     if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
         // Try grouped path first
-        if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+        if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, m_for_solver, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
             use_grouped = true;
         }
     }

     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, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+        int m_id_rows = 0;
+        if (dst && is_matmul_id) {
+            const int n_experts = ne02 > 0 ? ne02 : 1;
+            const size_t total_expert_rows = (size_t) dst->ne[1] * dst->ne[2];
+            int m_per_expert = (int) ((total_expert_rows + n_experts - 1) / n_experts);
+            if (sess->mdev.count > 1 && ggml_hexagon_can_row_partition(dst)) {
+                m_per_expert = (m_per_expert + (int) sess->mdev.count - 1) / (int) sess->mdev.count;
+            }
+            m_id_rows = hex_round_up(std::max(m_per_expert, 32), 32);
+        }
+        const uint32_t cost_m = is_matmul_id ? (uint32_t) m_id_rows : (uint32_t) m_for_solver;
+        if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, m_for_solver_padded, cost_m, n_threads, pipeline, is_matmul_id, aligned_tile_size, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
             return false;
         }
     }
@@ -5203,14 +5310,14 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
     kparams->n_act_threads = act_threads_selected;
     kparams->tile_size = htp_mm_get_weight_tile_size(wtype);
     kparams->aligned_tile_size = aligned_tile_size;
-    kparams->src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+    kparams->act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
     kparams->vtcm_size = vtcm_size;
     kparams->vtcm_src0_size = 0;
     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;
+    kparams->vtcm_act_size     = 0;
+    kparams->vtcm_bias_size    = (int32_t) src2_size;
+    kparams->vtcm_dst_size     = 0;

     if (is_batched && !is_matmul_id) {
         kparams->kernel_type = HTP_MM_KERNEL_HMX_F16_BATCHED;
@@ -5247,6 +5354,9 @@ static void ggml_hexagon_precompute_hvx_mm_params(
     kparams->n_hmx = 0;
     kparams->n_threads = sess->n_threads;

+    GGML_UNUSED(ne02);
+    GGML_UNUSED(ne03);
+
     const bool is_quant = (wtype != GGML_TYPE_F16 && wtype != GGML_TYPE_F32);
     const int src1_nrows = ne11 * ne12 * ne13;

@@ -5259,7 +5369,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(

         if (is_matmul_id) {
             kparams->kernel_type   = (src1_nrows < (int) sess->n_threads) ? HTP_MM_KERNEL_HVX_QUANT_BLOCK : HTP_MM_KERNEL_HVX_QUANT_ROW;
-            kparams->src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+            kparams->act_row_size  = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);

             struct htp_mm_hvx_vtcm_layout L;
             uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
@@ -5267,7 +5377,7 @@ 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], kparams->src1_row_size, 0, d, true, false
+                    0, src0->nb[1], kparams->act_row_size, 0, d, true, false
                 );
                 if (L.total_bytes <= vtcm_budget) {
                     best_n_prefetch = d;
@@ -5281,13 +5391,14 @@ static void ggml_hexagon_precompute_hvx_mm_params(
             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_act_size  = L.act_bytes;
+            kparams->vtcm_bias_size = 0;
             kparams->vtcm_dst_size  = L.dst_bytes;
             goto done_quant;
         } else {
             bool try_tiled = (k_align && opt_mm_select >= 1);
             if (try_tiled) {
-                kparams->src1_row_size = htp_mm_weight_has_offset(wtype)
+                kparams->act_row_size = htp_mm_weight_has_offset(wtype)
                                        ? htp_mm_q8_1_tiled_row_size(ne10)
                                        : htp_mm_q8_0_tiled_row_size(ne10);
                 if (src1_nrows < (int) sess->n_threads) {
@@ -5319,8 +5430,8 @@ static void ggml_hexagon_precompute_hvx_mm_params(
                     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_src2_size = L.src2_bytes;
+                    kparams->vtcm_act_size  = L.act_bytes;
+                    kparams->vtcm_bias_size = L.bias_bytes;
                     kparams->vtcm_dst_size  = L.dst_bytes;
                     goto done_quant;
                 }
@@ -5341,11 +5452,11 @@ static void ggml_hexagon_precompute_hvx_mm_params(
                 &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->act_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_act_size  = L.act_bytes;
+            kparams->vtcm_bias_size = L.bias_bytes;
             kparams->vtcm_dst_size = L.dst_bytes;
             kparams->n_prefetch = 16;
             return;
@@ -5363,11 +5474,11 @@ static void ggml_hexagon_precompute_hvx_mm_params(
                 &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->act_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_act_size  = L.act_bytes;
+            kparams->vtcm_bias_size = L.bias_bytes;
             kparams->vtcm_dst_size = L.dst_bytes;
             kparams->n_prefetch = 16;
             return;
@@ -5404,6 +5515,9 @@ static void ggml_hexagon_precompute_matmul_params_impl(
     const int ne00_padded = is_repack ? hex_round_up(ne00, 32) : ne00;
     const int ne01_padded = is_repack ? hex_round_up(ne01, 32) : ne01;
     const int ne11_padded = hex_round_up(ne11, 32);
+    // VTCM has to hold whole 32-row weight tiles, so size for the rounded-up N
+    // even when the tensor itself is ragged.
+    const int  ne01_tiled  = hex_round_up(ne01_padded, 32);

     const bool is_matmul_id = (dst->op == GGML_OP_MUL_MAT_ID);
     const bool is_batched   = (ne02 * ne03 > 1 || ne12 * ne13 > 1);
@@ -5413,7 +5527,7 @@ static void ggml_hexagon_precompute_matmul_params_impl(
     // Check HMX eligibility and try precomputing HMX parameters
     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, src2_size, vtcm_budget, kparams)) {
+        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, dst, wtype, ne00_padded, ne01_tiled, ne02, ne11, ne12, ne11_padded, is_matmul_id, is_batched, src2_size, vtcm_budget, kparams)) {
             goto finalize;
         }
     }
@@ -5429,6 +5543,66 @@ finalize:
     kparams->div_ne12     = init_fastdiv_values(ne12);
 }

+// The rows of dims 1..3 can be walked with one stride (size-1 dims skipped); returns that stride
+static bool ggml_hexagon_rows_stride(const int64_t * ne, const size_t * nb, size_t * stride) {
+    size_t s = 0, next = 0;
+    for (int i = 1; i < GGML_MAX_DIMS; i++) {
+        if (ne[i] == 1) continue;
+        if (s == 0) { s = nb[i]; next = s * ne[i]; continue; }
+        if (nb[i] != next) return false;
+        next *= ne[i];
+    }
+    *stride = s ? s : nb[1];
+    return true;
+}
+
+// A 2D weight applied to a batched activation whose rows are evenly strided is the same matmul over ne11 * ne12 * ne13 rows
+static bool ggml_hexagon_matmul_can_collapse(const struct ggml_tensor * src0, const struct ggml_tensor * src1, const struct ggml_tensor * dst) {
+    size_t s1, sd;
+    return (dst->op == GGML_OP_MUL_MAT || dst->op == GGML_OP_ADD) &&
+           src0->ne[2] == 1 && src0->ne[3] == 1 && src1->ne[2] * src1->ne[3] > 1 &&
+           src1->nb[0] == ggml_type_size(src1->type) && ggml_hexagon_rows_stride(src1->ne, src1->nb, &s1) &&
+           dst->nb[0] == ggml_type_size(dst->type) && ggml_hexagon_rows_stride(dst->ne, dst->nb, &sd);
+}
+
+static bool ggml_hexagon_matmul_add_can_collapse(
+    const struct ggml_tensor * src0,
+    const struct ggml_tensor * src1,
+    const struct ggml_tensor * src2,
+    const struct ggml_tensor * dst
+) {
+    if (!ggml_hexagon_matmul_can_collapse(src0, src1, dst)) {
+        return false;
+    }
+    if (!src2) {
+        return true;
+    }
+    const int64_t src2_nrows = src2->ne[1] * src2->ne[2] * src2->ne[3];
+    if (src2_nrows == 1) {
+        return true;
+    }
+    size_t s2;
+    return src2->nb[0] == ggml_type_size(src2->type) &&
+           src2->ne[0] == dst->ne[0] &&
+           src2->ne[1] == dst->ne[1] &&
+           src2->ne[2] == dst->ne[2] &&
+           src2->ne[3] == dst->ne[3] &&
+           ggml_hexagon_rows_stride(src2->ne, src2->nb, &s2);
+}
+
+static ggml_tensor ggml_hexagon_tensor_collapse_rows(const struct ggml_tensor * t) {
+    size_t stride = 0;
+    ggml_hexagon_rows_stride(t->ne, t->nb, &stride);
+    ggml_tensor c = *t;
+    c.ne[1] = t->ne[1] * t->ne[2] * t->ne[3];
+    c.ne[2] = 1;
+    c.ne[3] = 1;
+    c.nb[1] = stride;
+    c.nb[2] = c.nb[1] * c.ne[1];
+    c.nb[3] = c.nb[2];
+    return c;
+}
+
 static void ggml_hexagon_precompute_matmul_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0,
@@ -5436,6 +5610,13 @@ static void ggml_hexagon_precompute_matmul_params(
     const struct ggml_tensor * dst,
     struct htp_mm_kernel_params * kparams
 ) {
+    if (ggml_hexagon_matmul_can_collapse(src0, src1, dst)) {
+        const ggml_tensor src1_collapsed = ggml_hexagon_tensor_collapse_rows(src1);
+        const ggml_tensor dst_collapsed  = ggml_hexagon_tensor_collapse_rows(dst);
+        ggml_hexagon_precompute_matmul_params_impl(sess, src0, &src1_collapsed, &dst_collapsed, 0, 0, kparams);
+        kparams->collapse = 1;
+        return;
+    }
     ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
 }

@@ -5447,6 +5628,18 @@ static void ggml_hexagon_precompute_fused_matmul_add_params(
     const struct ggml_tensor * dst,
     struct htp_mm_kernel_params * kparams
 ) {
+    if (ggml_hexagon_matmul_add_can_collapse(src0, src1, src2, dst)) {
+        const ggml_tensor src1_collapsed = ggml_hexagon_tensor_collapse_rows(src1);
+        const ggml_tensor dst_collapsed  = ggml_hexagon_tensor_collapse_rows(dst);
+        const ggml_tensor src2_collapsed = (src2 && (src2->ne[1] * src2->ne[2] * src2->ne[3] > 1))
+                                           ? ggml_hexagon_tensor_collapse_rows(src2)
+                                           : (src2 ? *src2 : ggml_tensor{});
+        const struct ggml_tensor * p_src2 = src2 ? &src2_collapsed : nullptr;
+        const size_t src2_size = p_src2 ? hex_round_up(ggml_nbytes(p_src2), 128) : 0;
+        ggml_hexagon_precompute_matmul_params_impl(sess, src0, &src1_collapsed, &dst_collapsed, p_src2 ? p_src2->nb[1] : 0, src2_size, kparams);
+        kparams->collapse = 1;
+        return;
+    }
     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);
 }
@@ -6066,7 +6259,7 @@ static void ggml_hexagon_precompute_sort_params(
 static void ggml_hexagon_precompute_fused_mmnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0, // W0
-    const struct ggml_tensor * src1, // x
+    const struct ggml_tensor * act,  // x
     int32_t n_weights,
     struct htp_mm_kernel_params * kparams
 ) {
@@ -6078,23 +6271,24 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
     const int ne02 = src0->ne[2];
     const int ne03 = src0->ne[3];

-    const int ne10 = src1->ne[0];
-    const int ne11 = src1->ne[1];
-    const int ne12 = src1->ne[2];
-    const int ne13 = src1->ne[3];
+    const int ne10 = act->ne[0];
+    const int ne11 = act->ne[1];
+    const int ne12 = act->ne[2];
+    const int ne13 = act->ne[3];

     const int wtype = src0->type;
     const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
     const int ne00_padded = is_repack ? hex_round_up(ne00, 32) : ne00;
     const int ne01_padded = is_repack ? hex_round_up(ne01, 32) : ne01;
     const int ne11_padded = hex_round_up(ne11, 32);
+    const int ne01_tiled  = hex_round_up(ne01_padded, 32);

     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 >= 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, 0, vtcm_budget, kparams)) {
+    if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, act, nullptr, ne01_padded, false, is_batched)) {
+        if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, act, nullptr, wtype, ne00_padded, ne01_tiled, ne02, ne11, ne12, ne11_padded, false, is_batched, 0, vtcm_budget, kparams)) {
             kparams->n_weights = n_weights;
             goto finalize;
         }
@@ -6106,20 +6300,20 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
     }

     {
-        const int src1_nrows = ne11 * ne12 * ne13;
-        const size_t src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+        const int act_nrows = ne11 * ne12 * ne13;
+        const size_t act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
         const size_t src0_row_size = src0->nb[1];

         uint32_t best_n_prefetch = 16;

         if (is_repack) {
-            const uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
+            const uint32_t max_prefetch = (act_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
             best_n_prefetch = 2;
             for (uint32_t d = max_prefetch; d >= 2; d /= 2) {
                 struct htp_mm_hvx_vtcm_layout L;
                 htp_mm_hvx_vtcm_layout_build(
-                    &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads,
-                    0, src0_row_size, src1_row_size, 0, d, false, true
+                    &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, act_nrows, sess->n_threads,
+                    0, src0_row_size, act_row_size, 0, d, false, true
                 );
                 if (L.total_bytes <= sess->vtcm_size) {
                     best_n_prefetch = d;
@@ -6133,14 +6327,16 @@ static void ggml_hexagon_precompute_fused_mmnx_params(

         // Test tiled first
         htp_mm_hvx_vtcm_layout_build(
-            &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads,
-            0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true
+            &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, act_nrows, sess->n_threads,
+            0, src0_row_size, act_row_size, 0, best_n_prefetch, false, true
         );

         if (try_tiled && L.total_bytes <= sess->vtcm_size) {
-            kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
+            kparams->kernel_type    = HTP_MM_KERNEL_HVX_QUANT_ROW;
+            kparams->act_row_size   = act_row_size;
             kparams->vtcm_src0_size = L.src0_bytes;
-            kparams->vtcm_src1_size = L.src1_bytes;
+            kparams->vtcm_act_size  = L.act_bytes;
+            kparams->vtcm_bias_size = 0;
             kparams->vtcm_dst_size  = L.dst_bytes;
             kparams->vtcm_size      = L.total_bytes;
             kparams->n_prefetch     = best_n_prefetch;
@@ -6162,12 +6358,12 @@ finalize:
 static void ggml_hexagon_precompute_fused_mmidnx_params(
     const struct ggml_hexagon_session * sess,
     const struct ggml_tensor * src0, // W0
-    const struct ggml_tensor * src1, // x
+    const struct ggml_tensor * act,  // x
     const struct ggml_tensor * dst,  // dst0
     int32_t n_weights,
     struct htp_mm_kernel_params * kparams
 ) {
-    ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
+    ggml_hexagon_precompute_matmul_params_impl(sess, src0, act, dst, 0, 0, kparams);
     kparams->n_weights = n_weights;
 }

@@ -7137,6 +7333,14 @@ static bool mm_is_hmx_eligible(const ggml_tensor * t) {
     const ggml_tensor * src0 = t->src[0];
     const ggml_tensor * src1 = t->src[1];

+    if (ggml_hexagon_matmul_can_collapse(src0, src1, t)) {
+        const ggml_tensor src1_c = ggml_hexagon_tensor_collapse_rows(src1);
+        const int wtype = src0->type;
+        const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
+        const int ne01_padded = is_repack ? hex_round_up(src0->ne[1], 32) : src0->ne[1];
+        return ggml_hexagon_matmul_is_hmx_eligible(src0, &src1_c, t, ne01_padded, false, false);
+    }
+
     const int wtype = src0->type;
     const bool is_repack    = ggml_hexagon_is_repack_type((ggml_type) wtype);
     const bool is_matmul_id = (t->op == GGML_OP_MUL_MAT_ID);
@@ -7176,13 +7380,15 @@ static bool is_mergeable_mul_mat(const ggml_tensor * t) {

     const ggml_tensor * src0 = t->src[0];
     const ggml_tensor * src1 = t->src[1];
-    if (src1->type != GGML_TYPE_F32) return false;
     if (src0->ne[2] != 1 || src0->ne[3] != 1) return false;

     if (mm_is_hmx_eligible(t)) {
         return ggml_hexagon_is_hmx_weight_type(src0->type);
     }

+    // HVX path requires F32 activations and repacked weights (except Q6_K)
+    if (src1->type != GGML_TYPE_F32) return false;
+
     return ggml_hexagon_is_repack_type(src0->type) && src0->type != GGML_TYPE_Q6_K;
 }

@@ -8683,6 +8889,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
     const char * str_nhmx     = getenv("GGML_HEXAGON_NHMX");
     const char * str_mm_select = getenv("GGML_HEXAGON_MM_SELECT");
     const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT");
+    const char * str_fa_head_split = getenv("GGML_HEXAGON_FA_HEAD_SPLIT");
     const char * str_gdn_select = getenv("GGML_HEXAGON_GDN_SELECT");
     const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT");
     const char * str_ar_scatter = getenv("GGML_HEXAGON_AR_SCATTER");
@@ -8737,6 +8944,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
     opt_nhmx      = str_nhmx     ? atoi(str_nhmx)                         : opt_nhmx;
     opt_mm_select = str_mm_select ? atoi(str_mm_select)                   : opt_mm_select;
     opt_fa_select = str_fa_select ? atoi(str_fa_select)                   : opt_fa_select;
+    opt_fa_head_split = str_fa_head_split ? atoi(str_fa_head_split)        : opt_fa_head_split;
     opt_gdn_select = str_gdn_select ? atoi(str_gdn_select)                 : opt_gdn_select;
     opt_ar_select = str_ar_select ? atoi(str_ar_select)                   : opt_ar_select;
     opt_ar_scatter = str_ar_scatter ? atoi(str_ar_scatter)                : opt_ar_scatter;
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
index f079f7389..2ee55feb9 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
@@ -1869,8 +1869,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     factx.Bc             = kparams->Bc;
     factx.g_br           = kparams->u.hmx.g_br;
     factx.n_kv_blocks    = kparams->n_kv_blocks;
-    factx.is_q_fp32      = (kparams->is_q_fp32 != 0);
-    factx.is_dst_fp32    = (kparams->is_dst_fp32 != 0);
+    factx.is_q_fp32      = (q->type == HTP_TYPE_F32);
+    factx.is_dst_fp32    = (dst->type == HTP_TYPE_F32);
     factx.pipeline       = (kparams->u.hmx.pipeline != 0);
     factx.mask_broadcast = (kparams->u.hmx.mask_broadcast != 0);
     if (mask) {
@@ -1879,13 +1879,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     }

     factx.has_softcap   = (kparams->logit_softcap != 0.0f);
-    if (!factx.has_softcap) {
-        factx.scale = (__fp16) (kparams->scale * EXP_LOG2E_F);  // log2(e)
-    } else {
-        factx.scale = (__fp16) kparams->scale;
-    }
+    factx.scale         = (__fp16) kparams->scale;
     factx.max_bias      = kparams->max_bias;
-    factx.logit_softcap = factx.has_softcap ? (__fp16) (kparams->logit_softcap * EXP_LOG2E_F) : 0;
+    factx.logit_softcap = 0;
+    if (factx.has_softcap) {
+        factx.logit_softcap = (__fp16) kparams->logit_softcap;
+    }

     factx.n_head_log2 = kparams->n_head_log2;
     factx.m0          = kparams->m0;
@@ -1898,22 +1897,36 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     const uint32_t n_threads = factx.n_threads;
     const uint32_t G = factx.G;

-    // Multi-device: split Q blocks across devices
+    // Multi-device: prefer head-parallel partitioning (each core owns a disjoint head
+    // shard), falling back to Q-block (token) split when heads don't divide evenly.
     const uint32_t n_q_blocks = (neq1 + Br - 1) / Br;
-    uint32_t q_start_min = 0;
-    uint32_t q_start_max = neq1;
+    uint32_t q_start_min  = 0;
+    uint32_t q_start_max  = neq1;
+    uint32_t kv_head_min  = 0;
+    uint32_t kv_head_max  = n_kv_heads;

     if (octx->ctx->mdev.count > 1) {
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, htp_tensor_mdev_data_aligned(dst) ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
-        const uint32_t block_start = range.start;
-        const uint32_t block_end   = range.start + range.count;
+        const uint32_t mdev_count = octx->ctx->mdev.count;
+        const uint32_t mdev_idx   = octx->ctx->mdev.idx;
+        const uint32_t dst_e_size = (dst->type == HTP_TYPE_F32) ? sizeof(float) : sizeof(__fp16);
+        const bool can_split      = htp_tensor_can_row_partition(dst, dst_e_size);
+
+        if (kparams->head_split && can_split && n_kv_heads >= mdev_count && n_kv_heads % mdev_count == 0) {
+            const uint32_t kv_per_core = n_kv_heads / mdev_count;
+            kv_head_min = mdev_idx * kv_per_core;
+            kv_head_max = kv_head_min + kv_per_core;
+        } else {
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, can_split ? 1 : 0, mdev_idx, mdev_count, &octx->ctx->mdev.count_div);
+            const uint32_t block_start = range.start;
+            const uint32_t block_end   = range.start + range.count;

-        if (block_start >= block_end) {
-            return HTP_STATUS_OK;
-        }
+            if (block_start >= block_end) {
+                return HTP_STATUS_OK;
+            }

-        q_start_min = block_start * Br;
-        q_start_max = MIN(block_end * Br, neq1);
+            q_start_min = block_start * Br;
+            q_start_max = MIN(block_end * Br, neq1);
+        }
     }

     // ======== VTCM allocation (GQA-aware) ========
@@ -2032,7 +2045,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
             const size_t   g_br_actual = hex_align_up(n_rows_g, HMX_FP16_TILE_N_ROWS);
             const size_t   n_row_tiles = g_br_actual / HMX_FP16_TILE_N_ROWS;

-            for (uint32_t kv_head = 0; kv_head < n_kv_heads; ++kv_head) {
+            for (uint32_t kv_head = kv_head_min; kv_head < kv_head_max; ++kv_head) {
                 const uint32_t ik2 = kv_head;
                 const uint32_t ik3 = fastdiv(ib3, &kparams->broadcast_rk3);
                 const uint32_t iv2 = kv_head;
@@ -2040,7 +2053,7 @@ 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) {
+                if (ib3 == 0 && q_start == q_start_min && kv_head == kv_head_min) {
                     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;
@@ -2358,8 +2371,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                 uint32_t next_kv_head = kv_head + 1;
                 uint32_t next_q_start = q_start;
                 uint32_t next_ib3     = ib3;
-                if (next_kv_head >= n_kv_heads) {
-                    next_kv_head = 0;
+                if (next_kv_head >= kv_head_max) {
+                    next_kv_head = kv_head_min;
                     next_q_start = q_start + Br;
                     if (next_q_start >= q_start_max) {
                         next_q_start = q_start_min;
@@ -2478,7 +2491,7 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
         factx.src3_div3 = kparams->src3_div3;
     }

-    factx.is_q_fp32 = (kparams->is_q_fp32 != 0);
+    factx.is_q_fp32 = (q->type == HTP_TYPE_F32);
     factx.size_q_row_padded = kparams->u.hvx.size_q_row_padded;
     factx.size_k_row_padded = kparams->u.hvx.size_k_row_padded;
     factx.size_v_row_padded = kparams->u.hvx.size_v_row_padded;
@@ -2488,7 +2501,10 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     factx.scale = kparams->scale;
     factx.max_bias = kparams->max_bias;
     factx.has_softcap = (kparams->logit_softcap != 0.0f);
-    factx.logit_softcap = factx.has_softcap ? (__fp16) kparams->logit_softcap : 0;
+    factx.logit_softcap = 0;
+    if (factx.has_softcap) {
+        factx.logit_softcap = (__fp16) kparams->logit_softcap;
+    }

     factx.n_head_log2 = kparams->n_head_log2;
     factx.m0          = kparams->m0;
@@ -2512,10 +2528,25 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     uint32_t qrows      = total_qrows;

     if (octx->ctx->mdev.count > 1) {
-        const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
-        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
-        qrow_start = range.start;
-        qrows      = range.count;
+        const uint32_t mdev_count = octx->ctx->mdev.count;
+        const uint32_t mdev_idx   = octx->ctx->mdev.idx;
+        const uint32_t n_kv_heads = k->ne[2];
+        const uint32_t dst_e_size = (dst->type == HTP_TYPE_F32) ? sizeof(float) : sizeof(__fp16);
+        const bool can_split      = htp_tensor_can_row_partition(dst, dst_e_size);
+
+        // head range is contiguous in flat row space only when neq3 == 1
+        if (kparams->head_split && can_split && neq3 == 1 && n_kv_heads >= mdev_count && n_kv_heads % mdev_count == 0) {
+            const uint32_t G              = kparams->G;
+            const uint32_t kv_per_core    = n_kv_heads / mdev_count;
+            const uint32_t heads_per_core = kv_per_core * G;
+            const uint32_t head_start     = mdev_idx * heads_per_core;
+            qrow_start = head_start * neq1;
+            qrows      = heads_per_core * neq1;
+        } else {
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, mdev_idx, mdev_count, &octx->ctx->mdev.count_div);
+            qrow_start = range.start;
+            qrows      = range.count;
+        }
     }

     if (qrows == 0) {
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
index 22bb8c53d..04538842c 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
@@ -34,8 +34,8 @@ enum htp_fa_kernel_type {

 struct htp_fa_kernel_params {
     uint8_t  kernel_type;        // enum htp_fa_kernel_type
-    uint8_t  is_q_fp32;          // 1 = Q type is F32, 0 = F16
-    uint8_t  is_dst_fp32;        // 1 = dst type is F32, 0 = F16
+    uint8_t  head_split;         // 1 = partition by KV heads in multicore, 0 = token partition
+    uint8_t  flags;              // reserved
     uint8_t  n_threads;          // Number of threads to run

     // Common parameters
diff --git a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
index 7d7455d04..698d6a33b 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
@@ -745,7 +745,7 @@ void convert_f16_weight_to_fp16_tiles_task(
                 const uint8_t *r0 = state->src + row0 * state->row_stride;
                 const uint8_t *r1 = state->src + row1 * state->row_stride;

-                HVX_Vector v0 = hvx_vmemu((const __fp16 *)(r0 + byte_off));
+                HVX_Vector v0 = (row0 < state->n_cols) ? hvx_vmemu((const __fp16 *)(r0 + byte_off)) : Q6_V_vzero();
                 HVX_Vector v1 = (row1 < state->n_cols) ? hvx_vmemu((const __fp16 *)(r1 + byte_off)) : Q6_V_vzero();

                 Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v0);
@@ -788,7 +788,7 @@ void quantize_f32_weight_to_fp16_tiles_task(
                 const uint8_t *r0 = state->src + row0 * state->row_stride;
                 const uint8_t *r1 = state->src + row1 * state->row_stride;

-                HVX_Vector v0_f32 = hvx_vmem((const float *)(r0 + byte_off));
+                HVX_Vector v0_f32 = (row0 < state->n_cols) ? hvx_vmem((const float *)(r0 + byte_off)) : Q6_V_vzero();
                 HVX_Vector v1_f32 = (row1 < state->n_cols) ? hvx_vmem((const float *)(r1 + byte_off)) : Q6_V_vzero();

                 HVX_Vector v_out = hvx_vec_f32_to_f16(v0_f32, v1_f32);
@@ -988,9 +988,7 @@ static void transfer_output_chunk_fp16_to_fp32_col_chunk(
     uint32_t src2_stride,
     uint32_t dst_cols
 ) {
-    assert(c_len % HTP_MM_HMX_TILE_N_COLS == 0);
-    assert(total_n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
-    const size_t tile_row_stride = (total_n_cols / HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;
+    const size_t tile_row_stride = hmx_ceil_div(total_n_cols, HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;

     const HVX_Vector one = hvx_vec_splat_f16(1.0);

@@ -1137,6 +1135,73 @@ static void transfer_activation_row_pair_fp32_to_fp16(
     }
 }

+// F16-input variant of transfer_activation_row_pair_fp32_to_fp16, for F16 activation
+// (src1). Same shape as the F16 Q-prep in hmx-fa-kernels.h: one 128-byte load carries 64 f16
+// = two tile columns, and Q6_W_vshuff_VVR interleaves the two rows straight into the HMX tile
+// layout, so no F32 round-trip is needed. Rows are only 64-byte aligned when k_block is an odd
+// multiple of the tile width, hence the unaligned load type.
+static void transfer_activation_row_pair_f16_to_f16(__fp16 * restrict vtcm_dst,
+                                                    const __fp16 * restrict row0,
+                                                    const __fp16 * restrict row1,
+                                                    uint32_t r,
+                                                    uint32_t k_block,
+                                                    uint32_t k_valid,
+                                                    bool     row0_valid,
+                                                    bool     row1_valid) {
+    uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS;  // tile row index
+    uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS;  // intra-tile row idx
+
+    const uint32_t n_tile_cols = k_block / HTP_MM_HMX_TILE_N_COLS;
+    __fp16 * restrict tile_row = vtcm_dst + (size_t) r0 * n_tile_cols * HTP_MM_HMX_TILE_N_ELMS;
+
+    const HVX_UVector * pv0 = row0_valid ? (const HVX_UVector *) row0 : NULL;
+    const HVX_UVector * pv1 = row1_valid ? (const HVX_UVector *) row1 : NULL;
+
+    uint32_t c = 0;
+    for (; c + 64 <= k_valid; c += 64) {
+        HVX_Vector     v0 = pv0 ? pv0[c / 64] : Q6_V_vzero();
+        HVX_Vector     v1 = pv1 ? pv1[c / 64] : Q6_V_vzero();
+        HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
+
+        uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
+
+        HVX_Vector * tile0 = (HVX_Vector *) (tile_row + (size_t) c0 * HTP_MM_HMX_TILE_N_ELMS);
+        HVX_Vector * tile1 = (HVX_Vector *) (tile_row + (size_t) (c0 + 1) * HTP_MM_HMX_TILE_N_ELMS);
+
+        tile0[r1 / 2] = Q6_V_lo_W(vp);
+        tile1[r1 / 2] = Q6_V_hi_W(vp);
+    }
+    // Tail: fewer than 64 valid columns left, plus the k_valid..k_block padding that HMX will
+    // still multiply, so it has to be written as zeros.
+    for (; c < k_block; c += 64) {
+        HVX_Vector v0 = Q6_V_vzero();
+        HVX_Vector v1 = Q6_V_vzero();
+
+        if (c < k_valid) {
+            uint32_t       rem  = k_valid - c;  // 1..63 valid f16 lanes
+            HVX_VectorPred mask = Q6_Q_vsetq2_R(rem * sizeof(__fp16));
+            if (pv0) {
+                v0 = Q6_V_vmux_QVV(mask, pv0[c / 64], Q6_V_vzero());
+            }
+            if (pv1) {
+                v1 = Q6_V_vmux_QVV(mask, pv1[c / 64], Q6_V_vzero());
+            }
+        }
+
+        HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
+
+        uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
+
+        HVX_Vector * tile0 = (HVX_Vector *) (tile_row + (size_t) c0 * HTP_MM_HMX_TILE_N_ELMS);
+        tile0[r1 / 2]      = Q6_V_lo_W(vp);
+
+        if (c0 + 1 < n_tile_cols) {
+            HVX_Vector * tile1 = (HVX_Vector *) (tile_row + (size_t) (c0 + 1) * HTP_MM_HMX_TILE_N_ELMS);
+            tile1[r1 / 2]      = Q6_V_hi_W(vp);
+        }
+    }
+}
+
 static void transfer_activation_row_pair_fp32_to_fp16_col_chunk(
         __fp16 *restrict vtcm_dst,
         const float *restrict row0, // offset by c_first
diff --git a/ggml/src/ggml-hexagon/htp/hmx-utils.h b/ggml/src/ggml-hexagon/htp/hmx-utils.h
index ad295cb7d..1952aaa2c 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-utils.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-utils.h
@@ -73,19 +73,20 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
         for (uint32_t r = start_row; r < end_row; r += 2) {
             const uint32_t   ct             = r / HMX_FP16_TILE_N_ROWS;
             const uint32_t   local_r        = r % HMX_FP16_TILE_N_ROWS;
+            const bool       row0_valid     = r < n_cols;
             const bool       next_row_valid = (r + 1) < end_row && (r + 1) < n_cols;
             const HVX_Vector v_off0         = Q6_Vw_vadd_VwVw(v_scat_base, Q6_V_vsplat_R(local_r * 4));
             const HVX_Vector v_off1         = Q6_Vw_vadd_VwVw(v_off0, v_scat_step);

             __fp16 * tile_base = vtcm_dst + (size_t) ct * n_k_tiles * HMX_FP16_TILE_N_ELMS;
-            const uint8_t * p0 = (const uint8_t *) (vtcm_src + r * src_stride);
+            const uint8_t * p0 = row0_valid ? (const uint8_t *) (vtcm_src + r * src_stride) : NULL;
             const uint8_t * p1 = next_row_valid ? (const uint8_t *) (vtcm_src + (r + 1) * src_stride) : NULL;

-            assert(hex_is_aligned(p0, 128));
-            assert(hex_is_aligned(p1, 128));
+            assert(!p0 || hex_is_aligned(p0, 128));
+            assert(!p1 || hex_is_aligned(p1, 128));
             assert(c_byte_step % 128 == 0);

-            if (p1) {
+            if (p0 && p1) {
                 for (uint32_t i = 0; i < n_c_iters; ++i) {
                     HVX_Vector v0 = hvx_vmem(p0); p0 += c_byte_step;
                     HVX_Vector v1 = hvx_vmem(p1); p1 += c_byte_step;
@@ -96,9 +97,12 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
             } else {
                 const HVX_Vector vzero = Q6_V_vzero();
                 for (uint32_t i = 0; i < n_c_iters; ++i) {
-                    HVX_Vector v0 = hvx_vmem(p0); p0 += c_byte_step;
+                    HVX_Vector v0 = p0 ? hvx_vmem(p0) : vzero;
+                    if (p0) p0 += c_byte_step;
+                    HVX_Vector v1 = p1 ? hvx_vmem(p1) : vzero;
+                    if (p1) p1 += c_byte_step;
                     Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off0, v0);
-                    Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off1, vzero);
+                    Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off1, v1);
                     tile_base += dst_step;
                 }
             }
@@ -113,15 +117,16 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
         for (uint32_t r = start_row; r < end_row; r += 2) {
             const uint32_t   ct             = r / HMX_FP16_TILE_N_ROWS;
             const uint32_t   local_r        = r % HMX_FP16_TILE_N_ROWS;
+            const bool       row0_valid     = r < n_cols;
             const bool       next_row_valid = (r + 1) < end_row && (r + 1) < n_cols;
             const HVX_Vector v_off0         = Q6_Vw_vadd_VwVw(v_scat_base, Q6_V_vsplat_R(local_r * 4));
             const HVX_Vector v_off1         = Q6_Vw_vadd_VwVw(v_off0, v_scat_step);

             __fp16 * tile_base = vtcm_dst + (size_t) ct * n_k_tiles * HMX_FP16_TILE_N_ELMS;
-            const uint8_t * p0 = (const uint8_t *) (vtcm_src + r * src_stride);
+            const uint8_t * p0 = row0_valid ? (const uint8_t *) (vtcm_src + r * src_stride) : NULL;
             const uint8_t * p1 = next_row_valid ? (const uint8_t *) (vtcm_src + (r + 1) * src_stride) : NULL;

-            if (p1) {
+            if (p0 && p1) {
                 for (uint32_t i = 0; i < n_c_iters; ++i) {
                     HVX_Vector v0 = hvx_vmemu(p0); p0 += c_byte_step;
                     HVX_Vector v1 = hvx_vmemu(p1); p1 += c_byte_step;
@@ -132,9 +137,12 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
             } else {
                 const HVX_Vector vzero = Q6_V_vzero();
                 for (uint32_t i = 0; i < n_c_iters; ++i) {
-                    HVX_Vector v0 = hvx_vmemu(p0); p0 += c_byte_step;
+                    HVX_Vector v0 = p0 ? hvx_vmemu(p0) : vzero;
+                    if (p0) p0 += c_byte_step;
+                    HVX_Vector v1 = p1 ? hvx_vmemu(p1) : vzero;
+                    if (p1) p1 += c_byte_step;
                     Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off0, v0);
-                    Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off1, vzero);
+                    Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off1, v1);
                     tile_base += dst_step;
                 }
             }
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c
index 9dfd35649..a589e0a20 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.c
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c
@@ -27,8 +27,8 @@

 typedef struct {
     float        *dst;
-    dma_addr_t    src2_addr;
-    size_t        src2_bytes;
+    dma_addr_t    bias_addr;
+    size_t        bias_bytes;
     dma_addr_t    act_dma_addr;
     dma_addr_t    weight;
     dma_queue *   weight_dma;
@@ -36,9 +36,10 @@ typedef struct {
     int           k;
     int           n;
     int           act_stride;
+    uint32_t      act_elem_size;  // 4=F32 src1, 2=F16 src1
     int           weight_stride;
     int           dst_stride;
-    uint32_t      src2_stride;
+    uint32_t      bias_stride;
     int           ne02;
     int           ne03;
     int           ne12;
@@ -47,8 +48,8 @@ typedef struct {
     size_t        src0_nb3;
     size_t        act_nb2;
     size_t        act_nb3;
-    size_t        src2_nb2;
-    size_t        src2_nb3;
+    size_t        bias_nb2;
+    size_t        bias_nb3;
     size_t        dst_nb2;
     size_t        dst_nb3;
     int           r2;
@@ -118,24 +119,19 @@ struct htp_mm_context {

     // Dynamic VTCM pointers allocated sequentially
     uint8_t * vtcm_src0;
-    uint8_t * vtcm_src1;
-    uint8_t * vtcm_src2;
-    uint8_t * vtcm_src3;
+    uint8_t * vtcm_act;
+    uint8_t * vtcm_bias;
     uint8_t * vtcm_dst;
     uint8_t * vtcm_act_raw;

     // Cached strides
     uint32_t vtcm_src0_stride;
-    uint32_t vtcm_src1_stride;
-    uint32_t vtcm_src2_stride;
-    uint32_t vtcm_src3_stride;
+    uint32_t vtcm_act_stride;
     uint32_t vtcm_act_raw_stride;

     // Cached thread offsets/sizes
     uint32_t vtcm_src0_size_per_thread;
-    uint32_t vtcm_src1_size_per_thread;
-    uint32_t vtcm_src2_size_per_thread;
-    uint32_t vtcm_src3_size_per_thread;
+    uint32_t vtcm_act_size_per_thread;
     uint32_t vtcm_dst_size_per_thread;
 };

@@ -190,6 +186,7 @@ static const uint8_t __attribute__((aligned(VLEN))) kvalues_mxfp4_lut[] = {
     const struct htp_tensor * restrict src1 = octx->src[1];         \
     const struct htp_tensor * restrict src2 = octx->src[2];         \
     const struct htp_tensor * restrict  dst = octx->dst;            \
+    const struct htp_tensor * restrict  act = src1;                 \
                                                                     \
     const uint32_t ne00 = src0->ne[0];                              \
     const uint32_t ne01 = src0->ne[1];                              \
@@ -247,7 +244,7 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     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 act_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;                                                  \
@@ -260,13 +257,13 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     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 act_row_size = nb11;                                                                                                      \
+    const size_t act_stride = mmctx->vtcm_act_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;                                                                                       \
+    uint8_t * restrict act_data = mmctx->vtcm_act;                                                                                         \
                                                                                                                                            \
     const dma_addr_t src0_row = src0->data;                                                                                                \
                                                                                                                                            \
@@ -301,9 +298,9 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
                                                                                                                                            \
         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);                                    \
+        for (; ir1 + 1 < act_nrows; ir1 += 2) {                                                                                            \
+            const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);                                       \
+            const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_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));                                    \
                                                                                                                                            \
@@ -318,11 +315,11 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
                 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);                             \
+            DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_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);                                         \
+        for (; ir1 < act_nrows; ++ir1) {                                                                                                   \
+            const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);                                            \
             float * restrict dst_row          = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));                              \
             float * dst_ptr = &dst_row[ct * 32];                                                                                           \
                                                                                                                                            \
@@ -331,7 +328,7 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
                 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);                                                                \
+            DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, src2_ptr);                                                                 \
         }                                                                                                                                  \
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                              \
                                                                                                                                            \
@@ -359,18 +356,18 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     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 act_row_size = nb11;                                                                                    \
+    const size_t act_stride = mmctx->vtcm_act_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;                                                                              \
+    uint8_t * act_data = mmctx->vtcm_act;                                                                                \
                                                                                                                          \
     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;                                                     \
+    const uint8_t * restrict act_col = (const uint8_t *) act_data;                                                       \
     float * restrict dst_col          = (float *) dst->data;                                                             \
                                                                                                                          \
     const uint32_t tile_size = TILE_SIZE;                                                                                \
@@ -387,11 +384,11 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     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;                                         \
+            float * vtcm_bias_ptr = (float *) mmctx->vtcm_bias + 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),                                           \
+                dma_queue_push(dma_q, dma_make_data(vtcm_bias_ptr, src2_addr),                                           \
                                slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);   \
                 dma_queue_pop_nowait(dma_q);                                                                             \
             }                                                                                                            \
@@ -414,7 +411,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
         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);                                                      \
+        DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL);                                                       \
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                            \
                                                                                                                          \
         if (push_ct < ct_end) {                                                                                          \
@@ -430,7 +427,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
         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),                       \
+                            (const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row),                       \
                             copy_cnt);                                                                                   \
         } else {                                                                                                         \
             hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt);                            \
@@ -448,11 +445,11 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
                                                                                                                                   \
     const struct htp_tensor * restrict act = octx->src[n_weights]; /* x */                                                        \
     const uint32_t ne10 = act->ne[0];                                                                                             \
-    const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];                                                             \
-    const size_t src1_stride = mmctx->vtcm_src1_stride;                                                                           \
+    const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3];                                                              \
+    const size_t act_stride = mmctx->vtcm_act_stride;                                                                             \
                                                                                                                                   \
     uint8_t * restrict vtcm_weight_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;                               \
-    uint8_t * restrict src1_data       = mmctx->vtcm_src1;                                                                        \
+    uint8_t * restrict act_data        = mmctx->vtcm_act;                                                                         \
                                                                                                                                   \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                        \
     const uint32_t n_prefetch = kparams->n_prefetch;                                                                              \
@@ -511,23 +508,23 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
                                                                                                                                   \
             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);                       \
+            for (; ir1 + 1 < act_nrows; ir1 += 2) {                                                                               \
+                const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);                          \
+                const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_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];                                                                            \
                                                                                                                                   \
-                DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, NULL, NULL);                          \
+                DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_col1, valid_rows, NULL, NULL);                            \
             }                                                                                                                     \
                                                                                                                                   \
-            for (; ir1 < src1_nrows; ++ir1) {                                                                                     \
-                const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);                            \
+            for (; ir1 < act_nrows; ++ir1) {                                                                                      \
+                const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);                               \
                 float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));                                          \
                 float * dst_ptr = &dst_row[ct * 32];                                                                              \
-                DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL);                                                       \
+                DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL);                                                        \
             }                                                                                                                     \
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                 \
                                                                                                                                   \
@@ -550,10 +547,10 @@ MATMUL_2D_REPACKED_IMPL(q2_k,       512,  tiled_vec_dot_q2_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)

-static void hvx_mm_transfer_src1_dma(
+static void hvx_mm_transfer_act_dma(
     struct htp_ops_context * octx,
     const struct htp_mm_kernel_params * kparams,
-    const struct htp_tensor * src1,
+    const struct htp_tensor * act,
     uint8_t * dst_base,
     size_t dst_row_size,
     uint32_t m_start,
@@ -564,22 +561,22 @@ static void hvx_mm_transfer_src1_dma(
     }

     dma_queue * dma_q = octx->ctx->dma[0];
-    const uint32_t ne0 = src1->ne[0];
-    const size_t elem_size = (src1->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+    const uint32_t ne0 = act->ne[0];
+    const size_t elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
     const size_t row_bytes = ne0 * elem_size;
-    const size_t src1_nb1 = src1->nb[1];
-    const dma_addr_t src_base = src1->data;
+    const size_t act_nb1 = act->nb[1];
+    const dma_addr_t act_base = act->data;

-    const bool is_contiguous = (src1->nb[2] == src1->ne[1] * src1_nb1) &&
-                               (src1->nb[3] == src1->ne[2] * src1->nb[2]);
+    const bool is_contiguous = (act->nb[2] == act->ne[1] * act_nb1) &&
+                               (act->nb[3] == act->ne[2] * act->nb[2]);

     if (is_contiguous) {
-        const dma_addr_t src_addr = src_base + m_start * src1_nb1;
+        const dma_addr_t src_addr = act_base + m_start * act_nb1;
         dma_queue_push(dma_q, dma_make_data(dst_base, src_addr),
-                       dst_row_size, src1_nb1, row_bytes, m_rows);
+                       dst_row_size, act_nb1, row_bytes, m_rows);
         dma_queue_pop(dma_q);
     } else {
-        const uint32_t ne12_ne1 = src1->ne[2] * src1->ne[1];
+        const uint32_t ne12_ne1 = act->ne[2] * act->ne[1];
         const bool use_fastdiv = kparams->div_ne12_ne1.mp != 0;
         for (uint32_t ir = 0; ir < m_rows; ++ir) {
             const uint32_t ir1 = m_start + ir;
@@ -588,19 +585,19 @@ static void hvx_mm_transfer_src1_dma(
                 i13 = fastdiv(ir1, &kparams->div_ne12_ne1);
                 const uint32_t rem = ir1 - i13 * ne12_ne1;
                 i12 = fastdiv(rem, &kparams->div_ne1);
-                i11 = rem - i12 * src1->ne[1];
+                i11 = rem - i12 * act->ne[1];
             } else {
                 i13 = ne12_ne1 ? ir1 / ne12_ne1 : 0;
                 const uint32_t rem = ir1 - i13 * ne12_ne1;
-                i12 = src1->ne[1] ? rem / src1->ne[1] : 0;
-                i11 = rem - i12 * src1->ne[1];
+                i12 = act->ne[1] ? rem / act->ne[1] : 0;
+                i11 = rem - i12 * act->ne[1];
             }
-            const dma_addr_t row_src = src_base + (i11 * src1->nb[1] +
-                                                   i12 * src1->nb[2] +
-                                                   i13 * src1->nb[3]);
+            const dma_addr_t row_src = act_base + (i11 * act->nb[1] +
+                                                   i12 * act->nb[2] +
+                                                   i13 * act->nb[3]);
             uint8_t * row_dst = dst_base + ir * dst_row_size;
             dma_queue_push(dma_q, dma_make_data(row_dst, row_src),
-                           dst_row_size, src1_nb1, row_bytes, 1);
+                           dst_row_size, act_nb1, row_bytes, 1);
             dma_queue_pop(dma_q);
         }
     }
@@ -640,7 +637,7 @@ static void name(unsigned int nth, unsigned int ith, void * data) {
     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;                                                             \
+    uint8_t * restrict dst = mmctx->vtcm_act;                                                              \
     const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows);                                      \
     const size_t raw_row_size = mmctx->vtcm_act_raw_stride;                                                \
     const size_t dst_row_size = (dst_row_size_expr);                                                       \
@@ -655,9 +652,9 @@ static void name(unsigned int nth, unsigned int ith, void * data) {
 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_1_s16_tiled, "quantize-f32-q8_1_s16_tiled", quantize_f32_q8_1_s16_tiled_kernel, htp_mm_q8_1_tiled_row_size(ne0))
-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)
+QUANTIZE_IMPL(quantize_f32_f32,        "quantize-f32-f32",        quantize_f32_f32_kernel,        mmctx->vtcm_act_stride)
+QUANTIZE_IMPL(quantize_f32_f16,        "quantize-f32-f16",        quantize_f32_f16_kernel,        mmctx->vtcm_act_stride)
+QUANTIZE_IMPL(quantize_f16_f16,        "quantize-f16-f16",        quantize_f16_f16_kernel,        mmctx->vtcm_act_stride)

 static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, void * data) {
     (void) nth;
@@ -673,7 +670,7 @@ static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, vo

     quantize_f32_q8_0_tiled_block_kernel(
         (const float *) mmctx->vtcm_act_raw,
-        mmctx->vtcm_src1,
+        mmctx->vtcm_act,
         NULL,
         src->ne[0],
         mmctx->quant_ib_first[ith],
@@ -701,7 +698,7 @@ static void quantize_f32_q8_1_tiled_block(unsigned int nth, unsigned int ith, vo

     quantize_f32_q8_1_tiled_block_kernel(
         (const float *) mmctx->vtcm_act_raw,
-        mmctx->vtcm_src1,
+        mmctx->vtcm_act,
         NULL,
         src->ne[0],
         mmctx->quant_ib_first[ith],
@@ -729,7 +726,7 @@ static void quantize_f32_q8_1_s16_tiled_block(unsigned int nth, unsigned int ith

     quantize_f32_q8_1_s16_tiled_block_kernel(
         (const float *) mmctx->vtcm_act_raw,
-        mmctx->vtcm_src1,
+        mmctx->vtcm_act,
         NULL,
         src->ne[0],
         mmctx->quant_ib_first[ith],
@@ -795,10 +792,10 @@ static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     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;                                                                                                     \
+    const size_t act_stride = mmctx->vtcm_act_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;                                                                                                        \
+    uint8_t * restrict act_data = mmctx->vtcm_act;                                                                                                          \
                                                                                                                                                             \
     const uint32_t tile_size = TILE_SIZE;                                                                                                                   \
     const uint32_t aligned_tile_size = hex_align_up(tile_size, 128);                                                                                        \
@@ -874,19 +871,19 @@ static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
                                                                                                                                                             \
                 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);                          \
+                    const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 0) * act_stride);                             \
+                    const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 1) * act_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);                                                \
+                    DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_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);                               \
+                    const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_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);                                                                             \
+                    DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL);                                                                              \
                 }                                                                                                                                           \
             }                                                                                                                                               \
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);                                                                                           \
@@ -923,7 +920,7 @@ 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 = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows;                          // src1 rows
+    const uint32_t act_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows;
     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;
@@ -934,15 +931,15 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {

     const size_t dst_row_size  = nb1;
     const size_t src0_row_size = nb01;
-    const size_t src1_row_size = nb11;
+    const size_t act_row_size = nb11;

     const size_t src0_stride = mmctx->vtcm_src0_stride;
-    const size_t src1_stride = mmctx->vtcm_src1_stride;
+    const size_t act_stride = mmctx->vtcm_act_stride;

     // Per-thread VTCMs for all tensors
     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;
+    uint8_t * restrict act_data      = mmctx->vtcm_act;

     const dma_addr_t src0_row = src0->data;

@@ -968,21 +965,21 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
         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)
+        // Process act columns in pairs (2x2 tiling)
         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);
+        for (; ir1 + 1 < act_nrows; ir1 += 2) {
+            const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);
+            const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_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));
-            mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+            mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_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);
+        // Handle remaining act rows (fallback to 2x1)
+        for (; ir1 < act_nrows; ++ir1) {
+            const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
             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);
+            mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);

@@ -1005,15 +1002,15 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {

         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);
+        for (uint32_t ir1 = 0; ir1 < act_nrows; ++ir1) {
+            const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
             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);
+            mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
     }
     if (src2) {
-        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);
+        hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + act_nrows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
     }
 }

@@ -1029,21 +1026,21 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {

     const size_t dst_row_size  = nb1;
     const size_t src0_row_size = nb01;
-    const size_t src1_row_size = nb11;
+    const size_t act_row_size  = nb11;

     const size_t src0_stride = mmctx->vtcm_src0_stride;
-    const size_t src1_stride = mmctx->vtcm_src1_stride;
+    const size_t act_stride  = mmctx->vtcm_act_stride;

     // Per-thread VTCMs for all tensors
     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;
+    uint8_t * act_data      = mmctx->vtcm_act;

     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 uint8_t * restrict act_col = (const uint8_t *) act_data;
+    float * restrict dst_col         = (float *) dst->data;

     const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);

@@ -1055,11 +1052,11 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     // Prefill vtcm with 2x src0 rows
     if (src0_start_row < src0_end_row) {
         if (src2) {
-            float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row;
+            float * vtcm_bias_ptr = (float *) mmctx->vtcm_bias + 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_q, dma_make_data(vtcm_src2_ptr, src2_addr),
+                dma_queue_push(dma_q, dma_make_data(vtcm_bias_ptr, src2_addr),
                                slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);
                 dma_queue_pop_nowait(dma_q);
             }
@@ -1083,7 +1080,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     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);
-        mmctx->vec_dot_2x1(ne00, &tmp[ir0 - src0_start_row], ss0, ss0 + src0_stride, src1_col);
+        mmctx->vec_dot_2x1(ne00, &tmp[ir0 - src0_start_row], ss0, ss0 + src0_stride, act_col);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);

         // Prefetch next (n + vtcm_nrows) row
@@ -1103,7 +1100,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
                        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);
-        mmctx->vec_dot_1x1(ne00, &tmp[ir0 - src0_start_row], ss0, src1_col);
+        mmctx->vec_dot_1x1(ne00, &tmp[ir0 - src0_start_row], ss0, act_col);
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
     }

@@ -1113,7 +1110,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
         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),
+                            (const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row),
                             copy_cnt);
         } else {
             hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt);
@@ -1143,10 +1140,10 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
     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;
+    const size_t act_stride  = mmctx->vtcm_act_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;
+    uint8_t * restrict act_data      = mmctx->vtcm_act;


     if (src0_start_row >= src0_end_row || cur_m_rows == 0) {
@@ -1208,16 +1205,16 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {

                 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);
+                    const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 0) * act_stride);
+                    const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 1) * act_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);
+                    mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_col1);
                 }
                 for (; ir1 < batch_nrows; ++ir1) {
-                    const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);
+                    const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_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);
+                    mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
                 }
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -1253,9 +1250,9 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
                 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);
+                    const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_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);
+                    mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
                 }
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -1277,7 +1274,7 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
     const struct htp_tensor * restrict ids = octx->src[2];

     const uint32_t src0_nrows      = mmctx->src0_row_end - mmctx->src0_row_start;  // src0 rows per expert
-    const uint32_t src1_nrows      = ne11;
+    const uint32_t act_nrows       = ne11;
     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);

@@ -1299,13 +1296,13 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
     const struct mmid_row_mapping * matrix_rows       = mmctx->matrix_rows;

     const size_t dst_row_size  = nb1;
-    const size_t src1_row_size = htp_mm_q8_0_tiled_row_size(ne10);
+    const size_t act_row_size  = htp_mm_q8_0_tiled_row_size(ne10);

-    const size_t src1_stride = mmctx->vtcm_src1_stride;
+    const size_t act_stride = mmctx->vtcm_act_stride;

     // Per-thread VTCMs for all tensors
     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
-    uint8_t * restrict src1_data = mmctx->vtcm_src1;
+    uint8_t * restrict act_data = mmctx->vtcm_act;

     for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
         const int32_t cne1 = matrix_row_counts[cur_a];
@@ -1343,11 +1340,11 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
                 const int               rm1         = row_mapping.i1;  // expert idx
                 const int               rm2         = row_mapping.i2;  // token idx

-                const uint32_t ir1 = fastmodulo(rm1, ne11, &mmctx->mm_div_ne11);        // src1 row idx
-                const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (ir1 + rm2 * ne11 + 0) * src1_stride);
+                const uint32_t ir1 = fastmodulo(rm1, ne11, &mmctx->mm_div_ne11);        // act row idx
+                const uint8_t * restrict act_col = (const uint8_t *) (act_data + (ir1 + rm2 * ne11 + 0) * act_stride);
                 float * restrict dst_row = (float *) (dst->data + (rm1 * nb1 + rm2 * nb2 + 0));

-                mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+                mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

@@ -1383,14 +1380,14 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
     assert(ne13 % ne03 == 0);

     const size_t dst_row_size  = nb1;
-    const size_t src1_row_size = htp_mm_q8_0_tiled_row_size(ne10);
+    const size_t act_row_size = htp_mm_q8_0_tiled_row_size(ne10);

     const uint32_t n_aids = src2->ne[0];  // num activated experts
     const uint32_t n_ids  = ne02;         // num experts

     // Per-thread VTCMs for all tensors
     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
-    uint8_t * restrict src1_data = mmctx->vtcm_src1;
+    uint8_t * restrict act_data = mmctx->vtcm_act;

     for (uint32_t ie1 = 0; ie1 < n_aids; ++ie1) {  // for each expert
         const int32_t eid = *(const int32_t *) ((const uint8_t *) src2->data + ie1 * src2->nb[0]);
@@ -1400,7 +1397,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
         assert(eid < (int32_t) n_ids);

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

         const uint32_t tile_size = htp_mm_get_weight_tile_size(src0->type);
@@ -1426,7 +1423,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
             valid_rows = MIN(32, MAX(0, valid_rows));

             htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);
-            mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+            mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

             if (push_ct < ct_end) {
@@ -1457,7 +1454,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t n_ids  = src0->ne[2];

     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
-    uint8_t * restrict src1_data = mmctx->vtcm_src1;
+    uint8_t * restrict act_data = mmctx->vtcm_act;

     for (uint32_t ie1 = 0; ie1 < n_aids; ++ie1) {
         const int32_t eid = *(const int32_t *) ((const uint8_t *) ids->data + ie1 * ids->nb[0]);
@@ -1489,7 +1486,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
             if (src0_start_row >= src0_end_row) continue;

             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;
+            const uint8_t * restrict act_col = (const uint8_t *) act_data;
             float * restrict dst_row = (float *) (dst->data + ie1 * dst->nb[1]);

             const uint32_t tile_size = htp_mm_get_weight_tile_size(src_w->type);
@@ -1515,7 +1512,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
                 valid_rows = MIN(32, MAX(0, valid_rows));

                 htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);
-                mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+                mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
                 htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

                 if (push_ct < ct_end) {
@@ -1548,10 +1545,10 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t * matrix_row_counts = mmctx->matrix_row_counts;
     const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows;

-    const size_t src1_stride = mmctx->vtcm_src1_stride;
+    const size_t act_stride = mmctx->vtcm_act_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;
+    uint8_t * restrict act_data = mmctx->vtcm_act;

     for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
         const int32_t cne1 = matrix_row_counts[cur_a];
@@ -1612,10 +1609,10 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
                     const int rm2 = row_mapping.i2;

                     const uint32_t ir1 = fastmodulo(rm1, act->ne[1], &mmctx->mm_div_ne11);
-                    const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (ir1 + rm2 * act->ne[1]) * src1_stride);
+                    const uint8_t * restrict act_col = (const uint8_t *) (act_data + (ir1 + rm2 * act->ne[1]) * act_stride);
                     float * restrict dst_row = (float *) (dst->data + (rm1 * dst->nb[1] + rm2 * dst->nb[2]));

-                    mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+                    mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
                 }
                 htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);

@@ -1682,13 +1679,13 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
     struct htp_mm_context mmctx_struct = {0};
     struct htp_mm_context * mmctx = &mmctx_struct;
     mmctx->octx = octx;
-    mmctx->act = src1;
+    mmctx->act = act;

     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;

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

     uint32_t src0_row_start = 0;
     uint32_t src0_row_end   = src0_nrows;
@@ -1724,10 +1721,9 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     const size_t src0_row_size = nb01;
     const size_t dst_row_size  = nb1;
-    size_t       src1_row_size = nb11;
+    size_t       act_row_size  = nb11;

     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
-    size_t       src1_row_size_padded;

     worker_callback_t quant_task_func;
     worker_callback_t matmul_job_func;
@@ -1751,7 +1747,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
         } else {
             matmul_job_func = hvx_mm_4d;
         }
-    } else if (src1_nrows > 1) {
+    } else if (act_nrows > 1) {
         if (is_repacked) {
             switch (src0->type) {
                 case HTP_TYPE_Q4_0:   matmul_job_func = hvx_mm_2d_repacked_q4_0;   break;
@@ -1793,13 +1789,13 @@ 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 : quantize_f16_f16;
-            need_quant             = (src1->type == HTP_TYPE_F32);
-            mmctx->type            = (src1->type == HTP_TYPE_F32) ? "f32-f16" : "f16-f16";
+            quant_task_func        = (act->type == HTP_TYPE_F32) ? quantize_f32_f16 : quantize_f16_f16;
+            need_quant             = (act->type == HTP_TYPE_F32);
+            mmctx->type            = (act->type == HTP_TYPE_F32) ? "f32-f16" : "f16-f16";
             mmctx->vec_dot_1x1     = vec_dot_f16_f16_aa_1x1;
             mmctx->vec_dot_2x1     = vec_dot_f16_f16_aa_2x1;
             mmctx->vec_dot_2x2     = vec_dot_f16_f16_aa_2x2;
-            src1_row_size          = hex_round_up(ne10 * 2, 128);
+            act_row_size           = kparams->act_row_size;
             break;

         case HTP_MM_KERNEL_HVX_F32_F32_VTCM:
@@ -1809,7 +1805,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             mmctx->vec_dot_1x1     = vec_dot_f32_f32_aa_1x1;
             mmctx->vec_dot_2x1     = vec_dot_f32_f32_aa_2x1;
             mmctx->vec_dot_2x2     = vec_dot_f32_f32_aa_2x2;
-            src1_row_size          = hex_round_up(ne10 * 4, 128);
+            act_row_size           = kparams->act_row_size;
             break;

         case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
@@ -1821,9 +1817,9 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

             const uint32_t qk = QK_Q8_0_TILED;
             const uint32_t nb = (ne10 + qk - 1) / qk;
-            const uint32_t total_nb = src1_nrows * nb;
+            const uint32_t total_nb = act_nrows * nb;

-            if (src1_nrows < octx->n_threads && !is_batched) {
+            if (act_nrows < octx->n_threads && !is_batched) {
                 n_quant_tasks = MIN(total_nb, octx->n_threads);
                 quant_task_func = htp_mm_act_quant_block_func(src0->type);
                 for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -1835,28 +1831,28 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
                     mmctx->quant_c[ith]        = ib_first % nb;
                 }
             } else {
-                n_quant_tasks = MIN(src1_nrows, octx->n_threads);
+                n_quant_tasks = MIN(act_nrows, octx->n_threads);
                 quant_task_func = htp_mm_act_quant_row_func(src0->type);
             }
-            src1_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+            act_row_size = kparams->act_row_size;
             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_chunk = (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < act_nrows)
+                           ? (uint32_t) kparams->m_chunk : act_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, 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);
+                                 dst_row_size, src0_row_size, act_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false);

     if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM ||
         kparams->kernel_type == HTP_MM_KERNEL_HVX_F32_F32_VTCM ||
         kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW ||
         kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
-        mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+        mmctx->vtcm_act_size_per_thread = L.act_bytes;
     } else {
-        mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->n_threads_div);
+        mmctx->vtcm_act_size_per_thread = fastdiv(L.act_bytes, &octx->n_threads_div);
     }

     mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
@@ -1864,12 +1860,12 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     const size_t vtcm_size = L.total_bytes;

-    FARF(HIGH, "matmul-%s : src0-vtcm-size %zu src1-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
-         L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
+    FARF(HIGH, "matmul-%s : src0-vtcm-size %zu act-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
+         L.src0_bytes, L.act_bytes, L.dst_bytes, vtcm_size);

     FARF(HIGH, "matmul-%s : %ux%ux%ux%u * %ux%ux%ux%u-> %ux%ux%ux%u (0x%p, 0x%p, 0x%p)\n", mmctx->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], src0->data, src1->data, dst->data);
+         src0->ne[1], src0->ne[2], src0->ne[3], act->ne[0], act->ne[1], act->ne[2], act->ne[3], dst->ne[0],
+         dst->ne[1], dst->ne[2], dst->ne[3], src0->data, act->data, dst->data);

     if (octx->ctx->vtcm_size < vtcm_size) {
         FARF(ERROR, "matmul-%s : current VTCM reservation %zu is too small, needed %zu\n", mmctx->type,
@@ -1878,18 +1874,14 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
     }

     uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
-    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    mmctx->vtcm_act      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
     mmctx->vtcm_src0     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src2     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2);
+    mmctx->vtcm_bias     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_bias);
     mmctx->vtcm_dst      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
     mmctx->vtcm_act_raw  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);

-    octx->src1_spad.src  = NULL;
-    octx->src0_spad.src  = NULL;
-    octx->dst_spad.src   = NULL;
-
     mmctx->vtcm_src0_stride = src0_row_size_padded;
-    mmctx->vtcm_src1_stride = src1_row_size;
+    mmctx->vtcm_act_stride  = act_row_size;
     if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW) {
         mmctx->vtcm_act_raw_stride = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
     } else if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM) {
@@ -1900,14 +1892,14 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows) {
-        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);
+    if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < act_nrows) {
+        for (uint32_t m_start = 0; m_start < act_nrows; m_start += m_chunk) {
+            const uint32_t cur_m_rows = MIN(act_nrows - m_start, m_chunk);
             mmctx->cur_m_start = m_start;
             mmctx->cur_m_rows  = cur_m_rows;

             if (need_quant) {
-                hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, m_start, cur_m_rows);
+                hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, m_start, cur_m_rows);

                 const uint32_t qk = QK_Q8_0_TILED;
                 const uint32_t nb = (ne10 + qk - 1) / qk;
@@ -1933,22 +1925,22 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
                 mmctx->n_quant_tasks = quant_tasks;
                 work_queue_run(octx->ctx->work_queue, q_func, mmctx, quant_tasks);
             } else {
-                hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, m_start, cur_m_rows);
+                hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act, mmctx->vtcm_act_stride, m_start, cur_m_rows);
             }

             work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, octx->n_threads);
         }
     } else {
         mmctx->cur_m_start = 0;
-        mmctx->cur_m_rows  = src1_nrows;
+        mmctx->cur_m_rows  = act_nrows;

         if (need_quant) {
-            hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, src1_nrows);
-            mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
+            hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+            mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
             mmctx->n_quant_tasks = n_quant_tasks;
             work_queue_run(octx->ctx->work_queue, quant_task_func, mmctx, n_quant_tasks);
         } else {
-            hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, 0, src1_nrows);
+            hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act, mmctx->vtcm_act_stride, 0, act_nrows);
         }

         work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, octx->n_threads);
@@ -1964,11 +1956,11 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t n_weights = kparams->n_weights;

     const struct htp_tensor * restrict act = octx->src[n_weights];
-    const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
-    const size_t src1_stride = mmctx->vtcm_src1_stride;
+    const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3];
+    const size_t act_stride  = mmctx->vtcm_act_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;
+    uint8_t * restrict act_data      = mmctx->vtcm_act;

     const uint32_t n_prefetch = kparams->n_prefetch;
     assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
@@ -2020,17 +2012,17 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
             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) {
-                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);
+            for (; ir1 + 1 < act_nrows; ir1 += 2) {
+                const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);
+                const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_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));
-                mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+                mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_col1);
             }
-            for (; ir1 < src1_nrows; ++ir1) {
-                const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+            for (; ir1 < act_nrows; ++ir1) {
+                const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
                 float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));
-                mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
+                mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);

@@ -2049,10 +2041,10 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
                            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 ir1 = 0; ir1 < src1_nrows; ++ir1) {
-                const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+            for (uint32_t ir1 = 0; ir1 < act_nrows; ++ir1) {
+                const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
                 float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));
-                mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
+                mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
             }
             htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
         }
@@ -2144,6 +2136,7 @@ typedef struct {
     size_t                          vtcm_f32_act_bytes_per_thread;
     uint32_t                        dma_step_rows;
     uint32_t                        dma_step_rows_shift;
+    uint32_t                        act_elem_size;  // 4=F32 src1, 2=F16 src1
 } activation_transfer_task_state_t;

 typedef struct {
@@ -2262,7 +2255,7 @@ static void transfer_activation_chunk_col_chunk_worker_fn(unsigned int n, unsign
     );
 }

-static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
+static void transfer_activation_chunk_to_fp16_dma_pipelined(
         dma_queue *dma_q,
         __fp16 *restrict vtcm_dst,
         dma_addr_t act_dma_addr,
@@ -2270,7 +2263,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
         uint32_t k_block,
         uint32_t k_stride,
         uint32_t k_valid,
-        float *thread_f32_act,
+        uint8_t *thread_act,
+        uint32_t act_elem_size,
         struct htp_thread_trace *tr,
         uint32_t dma_step_rows,
         uint32_t dma_step_rows_shift) {
@@ -2280,38 +2274,56 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(

     const uint32_t n_steps = n_rows_padded >> dma_step_rows_shift;

+    const bool act_is_f16 = (act_elem_size == sizeof(__fp16));
+    const size_t row_bytes = (size_t) k_block * act_elem_size;  // staging row (DMA dst stride)
+    const size_t src_row_bytes = (size_t) k_stride * act_elem_size;  // DDR row (DMA src stride)
+    const size_t width_bytes = (size_t) k_valid * act_elem_size;
+
     // 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_data(thread_f32_act, act_dma_addr),
-                       k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+        dma_queue_push(dma_q, dma_make_data(thread_act, act_dma_addr),
+                       row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
     }
     // Push step 1 (if valid)
     if (n_steps > 1) {
         uint32_t next_r = R * 1;
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
-            float *next_buf = thread_f32_act + 1 * R * k_block;
-            dma_queue_push(dma_q, dma_make_data(next_buf, act_dma_addr + (size_t) next_r * k_stride * sizeof(float)),
-                           k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+            uint8_t *next_buf = thread_act + 1 * R * row_bytes;
+            dma_queue_push(dma_q, dma_make_data(next_buf, act_dma_addr + (size_t) next_r * src_row_bytes),
+                           row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
         }
     }
     for (uint32_t s = 0; s < n_steps; ++s) {
         uint32_t r = s << dma_step_rows_shift;
-        float *curr_buf = thread_f32_act;
+        uint8_t *curr_buf = thread_act;

         if (r < n_rows) {
-            curr_buf = (float *) dma_queue_pop(dma_q).dst;
+            curr_buf = (uint8_t *) dma_queue_pop(dma_q).dst;
         }

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, r);
-        for (uint32_t p = 0; p < (R >> 1); ++p) {
-            uint32_t row_idx = r + (p << 1);
-            float *pair_buf = curr_buf + (p << 1) * k_block;
-            bool r0_valid = ((row_idx + 0) < n_rows);
-            bool r1_valid = ((row_idx + 1) < n_rows);
+        // Two copies of the pair loop so the type is resolved once per step and the
+        // row-pair kernels stay direct (inlinable) calls.
+        if (act_is_f16) {
+            for (uint32_t p = 0; p < (R >> 1); ++p) {
+                uint32_t row_idx = r + (p << 1);
+                const __fp16 *pair_buf = (const __fp16 *) (curr_buf + (p << 1) * row_bytes);
+                bool r0_valid = ((row_idx + 0) < n_rows);
+                bool r1_valid = ((row_idx + 1) < n_rows);
+
+                transfer_activation_row_pair_f16_to_f16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+            }
+        } else {
+            for (uint32_t p = 0; p < (R >> 1); ++p) {
+                uint32_t row_idx = r + (p << 1);
+                const float *pair_buf = (const float *) (curr_buf + (p << 1) * row_bytes);
+                bool r0_valid = ((row_idx + 0) < n_rows);
+                bool r1_valid = ((row_idx + 1) < n_rows);

-            transfer_activation_row_pair_fp32_to_fp16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+                transfer_activation_row_pair_fp32_to_fp16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+            }
         }
         htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, r);

@@ -2320,8 +2332,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
         uint32_t next_r = next_s << dma_step_rows_shift;
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
-            dma_queue_push(dma_q, dma_make_data(curr_buf, act_dma_addr + (size_t) next_r * k_stride * sizeof(float)),
-                           k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+            dma_queue_push(dma_q, dma_make_data(curr_buf, act_dma_addr + (size_t) next_r * src_row_bytes),
+                           row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
         }
     }
 }
@@ -2336,11 +2348,11 @@ static void transfer_activation_chunk_worker_fn(unsigned int n, unsigned int i,
         size_t chunk_size = hex_smin(st->n_tot_chunks - chunk_idx, st->n_chunks_per_task);

         __fp16      *dst = st->dst + chunk_idx * st->k_block;
-        const dma_addr_t act_dma_addr = st->act_dma_addr + (size_t) chunk_idx * st->k_stride * sizeof(float);
+        const dma_addr_t act_dma_addr = st->act_dma_addr + (size_t) chunk_idx * st->k_stride * st->act_elem_size;

-        float *thread_f32_act = (float *)((char *)st->vtcm_f32_act + i * st->vtcm_f32_act_bytes_per_thread);
-        transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
-            st->ctx->dma[i], dst, act_dma_addr, chunk_size, st->k_block, st->k_stride, st->k_valid, thread_f32_act, tr, st->dma_step_rows, st->dma_step_rows_shift
+        uint8_t *thread_act = (uint8_t *) st->vtcm_f32_act + i * st->vtcm_f32_act_bytes_per_thread;
+        transfer_activation_chunk_to_fp16_dma_pipelined(
+            st->ctx->dma[i], dst, act_dma_addr, chunk_size, st->k_block, st->k_stride, st->k_valid, thread_act, st->act_elem_size, tr, st->dma_step_rows, st->dma_step_rows_shift
         );
     }
 }
@@ -2446,10 +2458,9 @@ static void dequantize_tiled_weight_chunk_to_fp16_tiles(
         int n_k_tiles, struct fastdiv_values n_k_tiles_div,
         worker_callback_t dequant_worker_fn, int n_threads) {

-    assert(n_cols  % HTP_MM_HMX_TILE_N_COLS == 0);
     assert(k_block % HTP_MM_HMX_TILE_N_COLS == 0);

-    size_t n_col_tiles = n_cols / HTP_MM_HMX_TILE_N_COLS;
+    size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS);
     size_t n_tot_tiles = n_col_tiles * n_k_tiles;

     size_t n_tiles_per_task = (n_threads == 1) ? n_tot_tiles : hmx_ceil_div(n_tot_tiles, n_threads);
@@ -2499,7 +2510,9 @@ static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned i
     output_transfer_col_chunk_state_t *st = (output_transfer_col_chunk_state_t *) data;
     struct htp_thread_trace * tr = &st->traces[i];

-    uint32_t n_blocks = st->n_cols / 32;
+    // Round up: the last block is partial when N is not 32-aligned. Its pad
+    // columns are dropped by the dst_cols clamp inside the store.
+    uint32_t n_blocks = hmx_ceil_div(st->n_cols, 32);
     uint32_t b_first  = fastdiv(n_blocks * i, &st->n_threads_div);
     uint32_t b_last   = fastdiv(n_blocks * (i + 1), &st->n_threads_div);
     uint32_t c_first  = b_first * 32;
@@ -2527,11 +2540,9 @@ static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned i

 static void transfer_output_chunk_threaded(struct htp_context *ctx, float *dst, const float *src2, const __fp16 *vtcm_src,
                                               int n_rows, int n_cols, int dst_stride, uint32_t src2_stride, int dst_cols, int n_threads) {
-    assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
-
     if (n_rows <= 0) return;

-    uint32_t n_blocks = (uint32_t)n_cols / 32;
+    uint32_t n_blocks = hmx_ceil_div((uint32_t) n_cols, 32);
     if (n_threads > 1 && n_blocks >= (uint32_t)n_threads) {
         struct fastdiv_values n_threads_div = (n_threads == (int)ctx->n_threads) ? ctx->n_threads_div : init_fastdiv_values(n_threads);
         output_transfer_col_chunk_state_t col_state;
@@ -2590,6 +2601,7 @@ struct activation_transfer_params {
     int                           k_valid;
     float *                       vtcm_f32_act;
     size_t                        vtcm_f32_act_bytes;
+    uint32_t                      act_elem_size;  // 4=F32 src1, 2=F16 src1
 };

 static void transfer_activation_chunk_threaded(const struct activation_transfer_params * params) {
@@ -2605,13 +2617,15 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
     int                           k_valid            = params->k_valid;
     float *                       vtcm_f32_act       = params->vtcm_f32_act;
     size_t                        vtcm_f32_act_bytes = params->vtcm_f32_act_bytes;
+    // element size of the activation rows (4 = F32, 2 = F16).
+    const uint32_t                act_elem_size      = params->act_elem_size ? params->act_elem_size : (uint32_t) sizeof(float);

     if (n_rows <= 0) {
         return;
     }

     const size_t n_tasks = (n_rows + 31) >> 5;
-    if (n_threads > 1 && k_block > 32 && n_tasks < (size_t) n_threads) {
+    if (act_elem_size == sizeof(float) && n_threads > 1 && k_block > 32 && n_tasks < (size_t) n_threads) {
         // Calculate step rows parameters for column-chunked dma pipelining
         uint32_t dma_step_rows = 2;
         uint32_t dma_step_rows_shift = 1;
@@ -2662,6 +2676,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
     state.traces             = ctx->trace;
     state.ctx                = ctx;
     state.vtcm_f32_act       = vtcm_f32_act;
+    state.act_elem_size = act_elem_size;

     state.vtcm_f32_act_bytes_per_thread = hex_align_down(fastdiv(vtcm_f32_act_bytes, act_threads_div), 128);

@@ -2729,6 +2744,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                                   dma_addr_t weight,
                                   int m, int k, int n,
                                   int act_stride,
+                                  uint32_t act_elem_size,  // 4=F32 src1, 2=F16 src1
                                   int weight_stride,
                                   int weight_type,
                                   int k_valid,
@@ -2748,7 +2764,12 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
     struct htp_thread_trace * tr = &ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

-    if (k % 32 != 0 || n % 32 != 0) { return -1; }
+    // Quantized weights are repacked and padded to 32, so we it has to be 32-aligned.
+    // Only F16/F32 weights can be non-32-aligned, they will be padded in the following kernel.
+    const bool wtype_is_quant = (weight_type != HTP_TYPE_F16 && weight_type != HTP_TYPE_F32);
+    if (k % 32 != 0 || (wtype_is_quant && n % 32 != 0)) {
+        return -1;
+    }
     if (!hex_is_aligned(dst, VLEN) || (act_dma_addr & (VLEN - 1)) != 0) { return -1; }

     size_t row_stride = htp_mm_get_tiled_row_stride(weight_type, k);
@@ -2817,10 +2838,10 @@ 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);
+    const bool has_bias      = (src2_bytes > 0 && src2_addr != 0);
+    float   *vtcm_bias       = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_bias, has_bias);
+    if (has_bias) {
+        dma_queue_push(weight_dma, dma_make_data(vtcm_bias, src2_addr), hex_align_up(src2_bytes, 128), 0, src2_bytes, 1);
         dma_queue_pop(weight_dma);
     }

@@ -2844,7 +2865,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+                .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,  // byte offset (F16/F32)
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -2854,18 +2875,19 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 .k_valid = k_valid,
                 .vtcm_f32_act = vtcm_f32_act,
                 .vtcm_f32_act_bytes = L.act_f32_bytes,
+                .act_elem_size = act_elem_size,
             };
             transfer_activation_chunk_threaded(&act_params);

             // 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;
+            const uint32_t height_A0 = is_quant ? hmx_ceil_div(n_cols_A0, 32) * n_k_tiles : n_cols_A0;
             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;
+                const uint32_t height_A1 = is_quant ? hmx_ceil_div(n_cols_A1, 32) * n_k_tiles : n_cols_A1;
                 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 +2911,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;
+                    const uint32_t height_p2 = is_quant ? hmx_ceil_div(n_cols_p2, 32) * n_k_tiles : n_cols_p2;
                     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);
                 }
@@ -2907,10 +2929,10 @@ 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 = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_prev) : NULL;
+                    const float *bias_chunk = has_bias ? (vtcm_bias + 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);
+                        transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output_bufs[(i - 1) % 2], n_rows, n_cols_prev, dst_stride, src2_stride, chunk_dst_cols, n_threads);
                     }
                 }
             }
@@ -2920,10 +2942,10 @@ 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 = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_last) : NULL;
+            const float *bias_chunk = has_bias ? (vtcm_bias + 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);
+                transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output_bufs[(n_chunk_cnt - 1) % 2], n_rows, n_cols_last, dst_stride, src2_stride, chunk_dst_cols, n_threads);
             }
         }
     } else {
@@ -2935,7 +2957,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+                .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,  // byte offset (F16/F32)
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -2945,13 +2967,14 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 .k_valid = k_valid,
                 .vtcm_f32_act = vtcm_f32_act,
                 .vtcm_f32_act_bytes = L.act_f32_bytes,
+                .act_elem_size = act_elem_size,
             };
             transfer_activation_chunk_threaded(&act_params);

             // A0: Pre-fetch the first weight chunk (nc = 0)
             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;
+                const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
                 dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
             }

@@ -2973,7 +2996,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                 const size_t nc_next = nc + n_chunk_n_cols;
                 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;
+                    const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_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);
                 }

@@ -2984,10 +3007,10 @@ 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 = has_src2 ? (vtcm_src2 + mr * src2_stride + nc) : NULL;
+                const float *bias_chunk = has_bias ? (vtcm_bias + 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);
+                    transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output, n_rows, n_cols, dst_stride, src2_stride, chunk_dst_cols, n_threads);
                 }
             }
         }
@@ -3013,7 +3036,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     const int k           = (int) act->ne[0];
     const int k_valid     = (int) act->ne[0];
     const int m           = (int) (act->ne[1] * act->ne[2] * act->ne[3]);
-    const int act_stride  = (int) (act->nb[1] / sizeof(float));
+    const uint32_t act_elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+    const int act_stride  = (int) (act->nb[1] / act_elem_size);
     const dma_addr_t act_dma_addr = act->data;

     if (k % 32 != 0) { return HTP_STATUS_NO_SUPPORT; }
@@ -3089,7 +3113,16 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     int m_start = 0;
     int m_rows  = m;
     if (octx->ctx->mdev.count > 1) {
-        const bool can_split = htp_tensor_can_row_partition(octx->dsts[0], sizeof(float));
+        bool can_split = htp_tensor_can_row_partition(act, act_elem_size);
+        if (can_split) {
+            for (uint32_t p = 0; p < n_weights; ++p) {
+                const struct htp_tensor * restrict dst = octx->dsts[p];
+                if (!htp_tensor_can_row_partition(dst, sizeof(float))) {
+                    can_split = false;
+                    break;
+                }
+            }
+        }
         const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
         m_start = (int) range.start;
         m_rows  = (int) range.count;
@@ -3118,7 +3151,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+                .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -3128,6 +3161,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 .k_valid = k_valid,
                 .vtcm_f32_act = vtcm_f32_act,
                 .vtcm_f32_act_bytes = L.act_f32_bytes,
+                .act_elem_size = act_elem_size,
             };
             transfer_activation_chunk_threaded(&act_params);

@@ -3149,13 +3183,13 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 const uint32_t dma_src_stride = is_quant ? tile_size : weight_stride;

                 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;
+                const uint32_t height_A0 = is_quant ? hmx_ceil_div(n_cols_A0, 32) * n_k_tiles : n_cols_A0;
                 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;
+                    const uint32_t height_A1 = is_quant ? hmx_ceil_div(n_cols_A1, 32) * n_k_tiles : n_cols_A1;
                     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);
                 }
@@ -3175,7 +3209,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                         n_k_tiles, n_k_tiles_div, dequant_worker_fn, n_threads);

                     if (i + 2 < n_chunk_cnt) {
-                        const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2;
+                        const uint32_t height_p2 = is_quant ? hmx_ceil_div(n_cols_p2, 32) * n_k_tiles : n_cols_p2;
                         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);
                     }
@@ -3216,7 +3250,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+                .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -3226,6 +3260,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                 .k_valid = k_valid,
                 .vtcm_f32_act = vtcm_f32_act,
                 .vtcm_f32_act_bytes = L.act_f32_bytes,
+                .act_elem_size = act_elem_size,
             };
             transfer_activation_chunk_threaded(&act_params);

@@ -3247,7 +3282,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;
+                    const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
                     dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
                 }

@@ -3266,7 +3301,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
                     const size_t nc_next = nc + n_chunk_n_cols;
                     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;
+                        const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_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);
                     }

@@ -3314,16 +3349,16 @@ 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) {
-            dma_addr_t cur_src2_addr = params->src2_addr ? (params->src2_addr +
-                                       b2 * params->src2_nb2 +
-                                       b3 * params->src2_nb3) : 0;
+            dma_addr_t cur_bias_addr = params->bias_addr ? (params->bias_addr +
+                                       b2 * params->bias_nb2 +
+                                       b3 * params->bias_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,
+                                cur_bias_addr, params->bias_bytes,
                                 hmx_mm_act_batch_addr(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,
+                                params->act_stride, params->act_elem_size, params->weight_stride * (int)sizeof(__fp16),
+                                HTP_TYPE_F16, params->k, params->dst_stride, params->bias_stride, params->n,
                                 m_chunk, n_chunk, pipeline, n_threads, act_threads,
                                 act_threads_div, k_div, 0, 0, vtcm_size);
         }
@@ -3339,7 +3374,8 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
     if (params->act_stride < params->k || params->weight_stride < params->k || params->dst_stride < params->n) { return -1; }
     if (params->ne02 <= 0 || params->ne03 <= 0 || params->ne12 <= 0 || params->ne13 <= 0) { return -1; }
     if (params->ne12 % params->ne02 != 0 || params->ne13 % params->ne03 != 0) { return -1; }
-    if (params->k % 32 != 0 || params->n % 32 != 0) { return -1; }
+    // N (the weight row count) does not have to be 32-aligned:
+    if (params->k % 32 != 0) { return -1; }
     if (!hex_is_aligned(params->dst, VLEN) || (params->act_dma_addr & (VLEN - 1)) != 0) { return -1; }

     const int group_size = params->r2;
@@ -3363,7 +3399,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, false, act_threads, 0, params->src2_bytes);
+    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, false, act_threads, 0, params->bias_bytes);

     if (L.total_bytes > vtcm_budget) {
         FARF(HIGH, "%s: grouped layout overflowed VTCM, falling back to simple batched loop", __func__);
@@ -3380,10 +3416,10 @@ 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(float, base, L.off_act_f32);

-    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);
+    const bool has_bias      = (params->bias_bytes > 0 && params->bias_addr != 0);
+    float   *vtcm_bias       = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_bias, has_bias);
+    if (has_bias) {
+        dma_queue_push(params->weight_dma, dma_make_data(vtcm_bias, params->bias_addr), hex_align_up(params->bias_bytes, 128), 0, params->bias_bytes, 1);
         dma_queue_pop(params->weight_dma);
     }

@@ -3417,7 +3453,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                 // thrashing from HVX loads at large strides.
                 for (int g = 0; g < group_size; ++g) {
                     const dma_addr_t act_dma_addr = hmx_mm_act_batch_addr(params, b2_base + g, b3) +
-                                                      mr * params->act_stride * sizeof(float);
+                                                      (size_t) mr * params->act_stride * params->act_elem_size;
                     __fp16 *vtcm_act_g = vtcm_f16_act + (size_t) g * L.act_head_stride;
                     struct activation_transfer_params act_params = {
                         .ctx = ctx,
@@ -3432,6 +3468,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                         .k_valid = params->k,
                         .vtcm_f32_act = vtcm_f32_act,
                         .vtcm_f32_act_bytes = L.act_f32_bytes,
+                        .act_elem_size = params->act_elem_size,
                     };
                     transfer_activation_chunk_threaded(&act_params);
                 }
@@ -3444,7 +3481,8 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                 }
                 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(weight_dma, dma_make_data(vtcm_scratch1, weight_group + params->weight_stride * sizeof(__fp16)),
+                    const dma_addr_t second_weight_chunk = weight_group + n_chunk_n_cols * params->weight_stride * sizeof(__fp16);
+                    dma_queue_push(weight_dma, dma_make_data(vtcm_scratch1, second_weight_chunk),
                                       fp16_row_bytes, weight_row_bytes, fp16_row_bytes, n_cols_second);
                 }

@@ -3455,7 +3493,9 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                     {
                         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 n_cols_tiled = hex_align_up(n_cols, HTP_MM_HMX_TILE_N_COLS);
+
+                        hmx_interleave_rows_to_tiles(vtcm_weight, (const __fp16 *) curr_raw, (uint32_t) n_cols, params->k, params->k, 0, (uint32_t) n_cols_tiled);

                         const size_t nc_next = nc + n_chunk_n_cols * 2;
                         if (nc_next < (size_t) params->n) {
@@ -3478,11 +3518,11 @@ 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 = has_src2 ? (vtcm_src2 + mr * params->src2_stride + nc) : NULL;
+                            const float *bias_chunk = has_bias ? (vtcm_bias + mr * params->bias_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,
-                                                               params->dst_stride, params->src2_stride, chunk_dst_cols, n_threads);
+                                transfer_output_chunk_threaded(ctx, output, bias_chunk, vtcm_output, (int) n_rows, (int) n_cols,
+                                                               params->dst_stride, params->bias_stride, chunk_dst_cols, n_threads);
                             }
                         }
                     }
@@ -3612,7 +3652,8 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

     const int cne1 = m;
-    const int m_padded = hex_align_up(m, 32);
+    const int m_core = m_end - m_start;
+    const int m_core_padded = hex_align_up(m_core > 0 ? m_core : 1, 32);

     if (k % 32 != 0 || n % 32 != 0) { return -1; }
     if (!hex_is_aligned(dst, VLEN) || !hex_is_aligned(activation, VLEN)) { return -1; }
@@ -3666,10 +3707,10 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
     const size_t overhead = htp_mm_hmx_get_2d_overhead(/*pipeline=*/false, /*is_matmul_id=*/true);
     size_t m_chunk_n_rows = 0, n_chunk_n_cols = 0;
     if (htp_mm_hmx_compute_chunks(vtcm_budget, overhead, size_per_n, size_per_m, size_per_mn,
-                           m_padded, n,
+                           m_core_padded, n,
                            /*m_block_cost=*/(size_t) n * HTP_MM_HMX_COST_W_DEQUANT,
-                           /*n_block_cost=*/(size_t) m_padded * HTP_MM_HMX_COST_A_CONVERT, &m_chunk_n_rows, &n_chunk_n_cols, &vtcm_used)) {
-        FARF(ERROR, "hmx-mm-id-2d: VTCM too small : m %d k %d n %d budget %zu", m_padded, k, n, vtcm_budget);
+                           /*n_block_cost=*/(size_t) m_core_padded * HTP_MM_HMX_COST_A_CONVERT, &m_chunk_n_rows, &n_chunk_n_cols, &vtcm_used)) {
+        FARF(ERROR, "hmx-mm-id-2d: VTCM too small : m %d k %d n %d budget %zu", m_core_padded, k, n, vtcm_budget);
         return -1;
     }

@@ -3709,7 +3750,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
         // A0: Pre-fetch the first weight chunk (nc = 0)
         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;
+            const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
             dma_queue_push(weight_dma, dma_make_data(vtcm_weight, weight),
                            dma_dst_stride, dma_src_stride, dma_width_bytes, height);
         }
@@ -3732,7 +3773,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
             const size_t nc_next = nc + n_chunk_n_cols;
             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;
+                const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_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);
             }
@@ -3759,14 +3800,19 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k

     int k = (int) src0->ne[0];
     int n = (int) src0->ne[1];
-    const int m_total    = (int) src1->ne[1];
-    const int act_stride = (int)(src1->nb[1] / sizeof(float));
+    const int m_total    = (int) act->ne[1];
+    const uint32_t act_elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+    const int act_stride = (int) (act->nb[1] / act_elem_size);
     const int wgt_stride = (int)(src0->nb[1] / sizeof(__fp16));

     int m_start = 0;
     int m_rows  = m_total;
     if (octx->ctx->mdev.count > 1) {
-        const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+        bool can_split = htp_tensor_can_row_partition(dst, sizeof(float)) &&
+                         htp_tensor_can_row_partition(act, act_elem_size);
+        if (src2 && src2->ne[1] > 1 && !htp_tensor_can_row_partition(src2, sizeof(float))) {
+            can_split = false;
+        }
         const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_total, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
         m_start = (int) range.start;
         m_rows  = (int) range.count;
@@ -3784,22 +3830,23 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
     if (src2) {
         src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
         src2_addr = src2->data + m_start * src2_stride * sizeof(float);
-        src2_bytes = (size_t) kparams->vtcm_src2_size;
+        src2_bytes = (size_t) kparams->vtcm_bias_size;
         src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
         src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
     }

     const int dst_stride = (int)(dst->nb[1] / sizeof(float));
     float       * dst_ptr = (float *) dst->data + m_start * dst_stride;
-    const dma_addr_t act_addr = src1->data + m_start * act_stride * sizeof(float);
+    // byte offset, act_stride is in activation elements (F16 or F32)
+    const dma_addr_t act_addr = act->data + (size_t) m_start * act_stride * act_elem_size;

     int ret = -1;
     const int n_threads = kparams->n_threads;
     if (kparams->kernel_type == HTP_MM_KERNEL_HMX_F16_BATCHED) {
         hmx_mm_f16_f32_batched_params_t batch_params = {
             .dst             = dst_ptr,
-            .src2_addr       = src2_addr,
-            .src2_bytes      = src2_bytes,
+            .bias_addr       = src2_addr,
+            .bias_bytes      = src2_bytes,
             .act_dma_addr    = act_addr,
             .weight          = src0->data,
             .weight_dma      = octx->ctx->dma[0],
@@ -3807,21 +3854,22 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
             .k               = k,
             .n               = n,
             .act_stride      = act_stride,
+            .act_elem_size = act_elem_size,
             .weight_stride   = wgt_stride,
             .dst_stride      = dst_stride,
-            .src2_stride     = src2_stride,
+            .bias_stride     = src2_stride,
             .ne02            = ne02,
             .ne03            = ne03,
             .ne12            = ne12,
             .ne13            = ne13,
             .src0_nb2        = src0->nb[2],
             .src0_nb3        = src0->nb[3],
-            .act_nb2         = src1->nb[2],
-            .act_nb3         = src1->nb[3],
+            .act_nb2         = act->nb[2],
+            .act_nb3         = act->nb[3],
             .dst_nb2         = dst->nb[2],
             .dst_nb3         = dst->nb[3],
-            .src2_nb2        = src2_nb2,
-            .src2_nb3        = src2_nb3,
+            .bias_nb2        = src2_nb2,
+            .bias_nb3        = src2_nb3,
             .r2              = (ne02 > 0) ? (ne12 / ne02) : 1,
             .r3              = (ne03 > 0) ? (ne13 / ne03) : 1,
             .div_r2          = kparams->div_r2,
@@ -3838,7 +3886,7 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
         ret = hmx_mm_2d_f32(
             octx->ctx, octx->ctx->dma[0], dst_ptr, src2_addr, src2_bytes,
             act_addr, src0->data,
-            m_rows, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
+            m_rows, k, n, act_stride, act_elem_size, (int) src0->nb[1], (int) src0->type, (int) act->ne[0],
             dst_stride, src2_stride, (int)dst->ne[0],
             kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
             kparams->n_act_threads,
@@ -3855,7 +3903,17 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
     return HTP_STATUS_OK;
 }

-int op_matmul(struct htp_ops_context * octx) {
+static inline void htp_mm_tensor_collapse_rows(struct htp_tensor * c, const struct htp_tensor * t, uint32_t stride) {
+    *c = *t;
+    c->ne[1] = t->ne[1] * t->ne[2] * t->ne[3];
+    c->ne[2] = 1;
+    c->ne[3] = 1;
+    c->nb[1] = stride;
+    c->nb[2] = c->nb[1] * c->ne[1];
+    c->nb[3] = c->nb[2];
+}
+
+static int op_matmul_impl(struct htp_ops_context * octx) {
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;

     const int status = htp_mm_init_context(octx, kparams);
@@ -3870,6 +3928,40 @@ int op_matmul(struct htp_ops_context * octx) {
     return hvx_mm_matmul(octx);
 }

+int op_matmul(struct htp_ops_context * octx) {
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+
+    if (kparams->collapse) {
+        const struct htp_tensor * act  = octx->src[1];
+        const struct htp_tensor * dst  = octx->dst;
+        const uint32_t s_act = (act->ne[1] > 1) ? act->nb[1] : ((act->ne[2] > 1) ? act->nb[2] : act->nb[3]);
+        const uint32_t sd    = (dst->ne[1] > 1) ? dst->nb[1] : ((dst->ne[2] > 1) ? dst->nb[2] : dst->nb[3]);
+        struct htp_tensor act_collapsed, dst_collapsed;
+        htp_mm_tensor_collapse_rows(&act_collapsed, act, s_act);
+        htp_mm_tensor_collapse_rows(&dst_collapsed, dst, sd);
+        octx->src[1] = &act_collapsed;
+        octx->dst    = &dst_collapsed;
+
+        struct htp_tensor src2_collapsed;
+        const struct htp_tensor * src2 = octx->src[2];
+        if (src2 && (src2->ne[1] * src2->ne[2] * src2->ne[3] > 1)) {
+            const uint32_t s2 = (src2->ne[1] > 1) ? src2->nb[1] : ((src2->ne[2] > 1) ? src2->nb[2] : src2->nb[3]);
+            htp_mm_tensor_collapse_rows(&src2_collapsed, src2, s2);
+            octx->src[2] = &src2_collapsed;
+        }
+
+        const int status = op_matmul_impl(octx);
+        octx->src[1] = act;
+        octx->dst    = dst;
+        if (src2) {
+            octx->src[2] = src2;
+        }
+        return status;
+    }
+
+    return op_matmul_impl(octx);
+}
+
 static int hmx_mm_op_matmul_id(
     struct htp_ops_context * octx,
     struct htp_mm_context * mmctx
@@ -3881,21 +3973,40 @@ static int hmx_mm_op_matmul_id(
     const int n_ids = octx->src[2]->ne[0];
     const int n_as  = ne02;

-    for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
+    const bool mdev_split = (octx->ctx->mdev.count > 1) && htp_tensor_can_row_partition(dst, sizeof(float));
+    if (octx->ctx->mdev.count > 1 && !mdev_split && octx->ctx->mdev.idx > 0) {
+        return HTP_STATUS_OK;
+    }
+    uint32_t n_active = 0;
+    if (mdev_split) {
+        for (uint32_t a = 0; a < (uint32_t) n_as; ++a) {
+            if (matrix_row_counts[a] > 0) n_active++;
+        }
+    }
+    const bool expert_split = mdev_split && (n_active >= octx->ctx->mdev.count);
+
+    uint32_t target_dev = 0;
+    for (uint32_t cur_a = 0; cur_a < (uint32_t) n_as; ++cur_a) {
         const int32_t cne1 = matrix_row_counts[cur_a];
         if (cne1 == 0) continue;

         const int m_padded = hex_align_up(cne1, 32);
         int m_start = 0, m_end = m_padded;
-        if (octx->ctx->mdev.count > 1) {
-            const bool can_split = htp_tensor_mdev_data_aligned(dst) && (uint32_t) cne1 >= octx->ctx->mdev.count;
+        if (expert_split) {
+            const bool my_expert = (target_dev == octx->ctx->mdev.idx);
+            if (++target_dev == octx->ctx->mdev.count) {
+                target_dev = 0;
+            }
+            if (!my_expert) continue;
+        } else if (mdev_split) {
+            const bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
             const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
             m_start = (int) range.start;
             m_end   = (int) (range.start + range.count);
         }
         if (m_start >= m_end) continue;

-        int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) src1->data,
+        int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) act->data,
                                    src0->data + cur_a * nb02,
                                    cne1, ne00, ne01,
                                    ne10,
@@ -3951,21 +4062,21 @@ static int hvx_mm_matmul_id(
         n_quant_tasks = MIN(act_nrows, octx->n_threads);
         quant_task_func = htp_mm_act_quant_row_func(src0->type);
     }
-    size_t src1_row_size  = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+    size_t act_row_size   = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);

     struct htp_mm_hvx_vtcm_layout L;
     htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, act_nrows, octx->n_threads,
-                                 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
+                                 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, true, false);

     const size_t vtcm_size = L.total_bytes;

-    FARF(HIGH, "matmul-id-%s : src0-spad-size %zu src1-spad-size %zu src2-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
-         L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
+    FARF(HIGH, "matmul-id-%s : src0-spad-size %zu act-spad-size %zu bias-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
+         L.src0_bytes, L.act_bytes, L.dst_bytes, vtcm_size);

     FARF(HIGH, "matmul-id-%s : %ux%ux%ux%u * %ux%ux%ux%u (%ux%ux%ux%u) -> %ux%ux%ux%u (0x%p, 0x%p, 0x%p)\n", mmctx->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],
+         src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], act->ne[0], act->ne[1], act->ne[2], act->ne[3],
          ids->ne[0], ids->ne[1], ids->ne[2], ids->ne[3], dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], src0->data,
-         src1->data, dst->data);
+         act->data, dst->data);

     // Make sure the reserved vtcm size is sufficient
     if (octx->ctx->vtcm_size < vtcm_size) {
@@ -3974,24 +4085,18 @@ static int hvx_mm_matmul_id(
     }

     uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
-    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    mmctx->vtcm_act      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
     mmctx->vtcm_src0     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src2     = NULL;
+    mmctx->vtcm_bias     = NULL;
     mmctx->vtcm_dst      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
     mmctx->vtcm_act_raw  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);

-    octx->src1_spad.src  = NULL;
-    octx->src0_spad.src  = NULL;
-    octx->src2_spad.src  = NULL;
-    octx->dst_spad.src   = NULL;
-
     mmctx->vtcm_src0_stride    = src0_row_size_padded;
-    mmctx->vtcm_src1_stride    = src1_row_size;
+    mmctx->vtcm_act_stride     = act_row_size;
     mmctx->vtcm_act_raw_stride = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));

     mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
-    mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
-    mmctx->vtcm_src2_size_per_thread = 0;
+    mmctx->vtcm_act_size_per_thread  = L.act_bytes;
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->cur_m_start = 0;
@@ -3999,7 +4104,7 @@ static int hvx_mm_matmul_id(

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+    hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);

     mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->n_quant_tasks = n_quant_tasks;
@@ -4022,18 +4127,37 @@ static int hmx_mm_op_matmul_id_nx(
     const struct htp_tensor * restrict act  = octx->src[n_weights];
     const int n_as = src0->ne[2];

+    bool mdev_split = (octx->ctx->mdev.count > 1);
+    for (uint32_t p = 0; p < n_weights && mdev_split; ++p) {
+        const struct htp_tensor * restrict dst = octx->dsts[p];
+        mdev_split = htp_tensor_can_row_partition(dst, sizeof(float));
+    }
+    if (octx->ctx->mdev.count > 1 && !mdev_split && octx->ctx->mdev.idx > 0) {
+        return HTP_STATUS_OK;
+    }
+    uint32_t n_active = 0;
+    if (mdev_split) {
+        for (uint32_t a = 0; a < (uint32_t) n_as; ++a) {
+            if (matrix_row_counts[a] > 0) n_active++;
+        }
+    }
+    const bool expert_split = mdev_split && (n_active >= octx->ctx->mdev.count);
+
+    uint32_t target_dev = 0;
     for (uint32_t cur_a = 0; cur_a < (uint32_t) n_as; ++cur_a) {
         const int32_t cne1 = matrix_row_counts[cur_a];
         if (cne1 == 0) continue;

         const int m_padded = hex_align_up(cne1, 32);
         int m_start = 0, m_end = m_padded;
-        if (octx->ctx->mdev.count > 1) {
-            bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
-            for (uint32_t p = 0; p < n_weights && can_split; ++p) {
-                const struct htp_tensor * restrict dst = octx->dsts[p];
-                can_split = !dst || htp_tensor_mdev_data_aligned(dst);
+        if (expert_split) {
+            const bool my_expert = (target_dev == octx->ctx->mdev.idx);
+            if (++target_dev == octx->ctx->mdev.count) {
+                target_dev = 0;
             }
+            if (!my_expert) continue;
+        } else if (mdev_split) {
+            const bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
             const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
             m_start = (int) range.start;
             m_end   = (int) (range.start + range.count);
@@ -4104,11 +4228,11 @@ static int hvx_mm_matmul_id_nx(
         n_quant_tasks = MIN(act_nrows, octx->n_threads);
         quant_task_func = htp_mm_act_quant_row_func(src0->type);
     }
-    size_t src1_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]);
+    size_t act_row_size  = htp_mm_weight_has_offset(src0->type) ? 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], act_nrows, octx->n_threads,
-                                 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
+                                 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, true, false);

     const size_t vtcm_size = L.total_bytes;

@@ -4120,35 +4244,29 @@ static int hvx_mm_matmul_id_nx(

     uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
     mmctx->vtcm_src0     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    mmctx->vtcm_act      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
     mmctx->vtcm_dst      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
     mmctx->vtcm_act_raw  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);

-    octx->src0_spad.src = NULL;
-    octx->src1_spad.src = NULL;
-    octx->src2_spad.src = NULL;
-    octx->src3_spad.src = NULL;
-    octx->dst_spad.src  = NULL;
-
     mmctx->vtcm_src0_stride    = 0;
-    mmctx->vtcm_src1_stride    = src1_row_size;
+    mmctx->vtcm_act_stride     = act_row_size;
     mmctx->vtcm_act_raw_stride = hex_round_up(act->ne[0] * sizeof(float), QK_Q8_0_TILED * sizeof(float));

     mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
-    mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+    mmctx->vtcm_act_size_per_thread  = L.act_bytes;
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->cur_m_start = 0;
     mmctx->cur_m_rows  = act_nrows;

-    FARF(HIGH, "matmul-id-nx: src0 %d:%d:%d type %s nrows %u, src1 %d:%d:%d nrows %u, vtcm %zu/%zu, threads %d\n",
+    FARF(HIGH, "matmul-id-nx: src0 %d:%d:%d type %s nrows %u, act %d:%d:%d nrows %u, vtcm %zu/%zu, threads %d\n",
          src0->ne[0], src0->ne[1], src0->ne[2], mmctx->type, src0->ne[1],
          act->ne[0], act->ne[1], act->ne[2], act_nrows,
          L.total_bytes, octx->ctx->vtcm_size, octx->n_threads);

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    hvx_mm_transfer_src1_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+    hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);

     mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->n_quant_tasks           = n_quant_tasks;
@@ -4242,10 +4360,10 @@ int op_matmul_id(struct htp_ops_context * octx) {
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

     mmctx->octx = octx;
-    mmctx->act = src1;
+    mmctx->act = act;

     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)) {
+    if (htp_tensor_is_extended(ids) || htp_tensor_is_extended(act) || htp_tensor_is_extended(dst)) {
         return HTP_STATUS_NO_SUPPORT;
     }

@@ -4255,7 +4373,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);

     const uint32_t src0_nrows = ne01;  // per expert
-    const uint32_t src1_nrows = ne11 * ne12 * ne13;
+    const uint32_t act_nrows  = ne11 * ne12 * ne13;

     // row groups
     const int n_ids = ids->ne[0];  // n_expert_used
@@ -4266,7 +4384,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
     uint32_t * matrix_row_counts = (uint32_t *) mapping_buf;
     struct mmid_row_mapping * matrix_rows = NULL;

-    if (src1_nrows > 1) {
+    if (act_nrows > 1) {
         const size_t matrix_row_counts_size = n_as * sizeof(uint32_t);
         assert(octx->ctx->ddr_spad_size >= matrix_row_counts_size);

@@ -4300,9 +4418,9 @@ int op_matmul_id(struct htp_ops_context * octx) {
     mmctx->mapping_stride       = mapping_stride;
     mmctx->mm_div_ne11          = kparams->div_ne1;
     mmctx->src0_row_size_padded = src0_row_size_padded;
-    mmctx->act_nrows            = src1_nrows;
+    mmctx->act_nrows            = act_nrows;
     mmctx->cur_m_start          = 0;
-    mmctx->cur_m_rows           = src1_nrows;
+    mmctx->cur_m_rows           = act_nrows;

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

@@ -4334,7 +4452,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
         mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);

         if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) {
-            s = hvx_mm_matmul_id(octx, mmctx, src1_nrows > 1 ? hvx_mm_id : hvx_mv_id);
+            s = hvx_mm_matmul_id(octx, mmctx, act_nrows > 1 ? hvx_mm_id : hvx_mv_id);
         } else {
             s = HTP_STATUS_NO_SUPPORT;
         }
@@ -4369,7 +4487,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
         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])) {
+        if (htp_tensor_is_extended(octx->dsts[p])) {
             return HTP_STATUS_NO_SUPPORT;
         }
     }
@@ -4446,7 +4564,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {

     return s;
 }
-int op_matmul_nx(struct htp_ops_context * octx) {
+static int op_matmul_nx_impl(struct htp_ops_context * octx) {
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;

     const int status = htp_mm_init_context(octx, kparams);
@@ -4511,13 +4629,13 @@ int op_matmul_nx(struct htp_ops_context * octx) {
         quant_task_func = htp_mm_act_quant_row_func(src0->type);
     }

-    const size_t src1_row_size = htp_mm_weight_has_offset(src0->type)
+    const size_t act_row_size = htp_mm_weight_has_offset(src0->type)
                                ? 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], act_nrows, octx->n_threads,
-                                 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true);
+                                 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, false, true);

     const size_t vtcm_size = L.total_bytes;

@@ -4529,22 +4647,16 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
     mmctx->vtcm_src0     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    mmctx->vtcm_act      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
     mmctx->vtcm_dst      = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
     mmctx->vtcm_act_raw  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);

-    octx->src0_spad.src  = NULL;
-    octx->src1_spad.src  = NULL;
-    octx->src2_spad.src  = NULL;
-    octx->src3_spad.src  = NULL;
-    octx->dst_spad.src   = NULL;
-
     mmctx->vtcm_src0_stride    = is_repacked ? 0 : src0_row_size_padded;
-    mmctx->vtcm_src1_stride    = src1_row_size;
+    mmctx->vtcm_act_stride     = act_row_size;
     mmctx->vtcm_act_raw_stride = hex_round_up(act->ne[0] * sizeof(float), QK_Q8_0_TILED * sizeof(float));

     mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
-    mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+    mmctx->vtcm_act_size_per_thread  = L.act_bytes;
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     // Run fused matmul
@@ -4571,7 +4683,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    hvx_mm_transfer_src1_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+    hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);

     mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->n_quant_tasks = n_quant_tasks;
@@ -4581,3 +4693,37 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     return HTP_STATUS_OK;
 }
+
+int op_matmul_nx(struct htp_ops_context * octx) {
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+
+    if (kparams->collapse) {
+        const uint32_t n_weights = kparams->n_weights;
+        const struct htp_tensor * act = octx->src[n_weights];
+        const uint32_t s1 = (act->ne[1] > 1) ? act->nb[1] : ((act->ne[2] > 1) ? act->nb[2] : act->nb[3]);
+        struct htp_tensor act_collapsed;
+        struct htp_tensor dsts_collapsed[HTP_OP_MAX_OUTPUTS];
+        const struct htp_tensor * orig_dsts[HTP_OP_MAX_OUTPUTS];
+
+        htp_mm_tensor_collapse_rows(&act_collapsed, act, s1);
+        octx->src[n_weights] = &act_collapsed;
+
+        for (uint32_t p = 0; p < n_weights; p++) {
+            orig_dsts[p] = octx->dsts[p];
+            const struct htp_tensor * d = octx->dsts[p];
+            const uint32_t sd = (d->ne[1] > 1) ? d->nb[1] : ((d->ne[2] > 1) ? d->nb[2] : d->nb[3]);
+            htp_mm_tensor_collapse_rows(&dsts_collapsed[p], d, sd);
+            octx->dsts[p] = &dsts_collapsed[p];
+        }
+
+        const int status = op_matmul_nx_impl(octx);
+
+        octx->src[n_weights] = act;
+        for (uint32_t p = 0; p < n_weights; p++) {
+            octx->dsts[p] = orig_dsts[p];
+        }
+        return status;
+    }
+
+    return op_matmul_nx_impl(octx);
+}
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h
index 386cb3049..1d11a4c12 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.h
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h
@@ -87,24 +87,26 @@ enum htp_mm_kernel_type {

 // Op-specific struct for precomputed matmul params
 struct htp_mm_kernel_params {
-    int32_t  kernel_type;        // enum htp_mm_kernel_type
-    int32_t  pipeline;           // 1 = pipelined execution, 0 = standard
+    uint8_t  kernel_type;        // enum htp_mm_kernel_type
+    uint8_t  pipeline;           // 1 = pipelined execution, 0 = standard
+    uint8_t  collapse;           // 1 = collapse outer dims into 2D, 0 = standard
+    uint8_t  n_hmx;              // 1 = use HMX, 0 = use HVX
+
+    uint8_t  n_threads;          // Number of threads to spawn
+    uint8_t  n_act_threads;      // Number of threads for activation preparation
+    uint8_t  n_prefetch;         // Prefetch lookahead buffers/rows in VTCM
+    uint8_t  n_weights;          // Number of weights for fused NX
+
     int32_t  m_chunk;            // Row chunk size (M chunk)
     int32_t  n_chunk;            // Col chunk size (N chunk)
-    int32_t  n_threads;          // Number of threads to spawn
-    int32_t  n_act_threads;      // Number of threads for activation preparation
-    int32_t  n_hmx;              // 1 = use HMX, 0 = use HVX
-    int32_t  n_prefetch;         // Prefetch lookahead buffers/rows in VTCM
     int32_t  tile_size;          // Weight tile size
     int32_t  aligned_tile_size;  // Aligned weight tile size (padded to 128)
-    int32_t  src1_row_size;      // Row size for quantized activation
+    int32_t  act_row_size;       // Row size for activation scratchpad
     int32_t  vtcm_size;          // Total required scratchpad size in VTCM
     int32_t  vtcm_src0_size;     // src0 scratchpad size in VTCM
-    int32_t  vtcm_src1_size;     // src1 scratchpad size in VTCM
-    int32_t  vtcm_src2_size;     // src2 scratchpad size in VTCM (fused only)
-    int32_t  vtcm_src3_size;     // src3 scratchpad size in VTCM (fused only)
+    int32_t  vtcm_act_size;      // activation scratchpad size in VTCM
+    int32_t  vtcm_bias_size;     // bias scratchpad size in VTCM (fused only)
     int32_t  vtcm_dst_size;      // dst scratchpad size in VTCM
-    int32_t  n_weights;          // Number of weights for fused NX

     // Precomputed division values
     struct fastdiv_values div_ne12_ne1;
@@ -147,6 +149,7 @@ static inline int htp_mm_hmx_compute_chunks(size_t   vtcm_total,
     const size_t usable = vtcm_total - overhead;

     size_t best_cost = SIZE_MAX;
+    size_t best_tail_waste = SIZE_MAX;
     size_t best_mn   = 0;
     size_t best_m = 0, best_n = 0;

@@ -173,12 +176,17 @@ static inline int htp_mm_hmx_compute_chunks(size_t   vtcm_total,
             size_t mblocks = ((size_t) m + mc - 1) / mc;
             size_t nblocks = ((size_t) n + nc - 1) / nc;
             size_t cost    = mblocks * m_block_cost + nblocks * n_block_cost;
+            size_t rem     = n % nc;
+            size_t tail_waste = (rem == 0) ? 0 : (nc - rem);
             size_t mn      = mc * nc;
-            if (cost < best_cost || (cost == best_cost && mn > best_mn)) {
-                best_cost = cost;
-                best_mn   = mn;
-                best_m    = mc;
-                best_n    = nc;
+            if (cost < best_cost ||
+                (cost == best_cost && tail_waste < best_tail_waste) ||
+                (cost == best_cost && tail_waste == best_tail_waste && mn > best_mn)) {
+                best_cost       = cost;
+                best_tail_waste = tail_waste;
+                best_mn         = mn;
+                best_m          = mc;
+                best_n          = nc;
             }
         }

@@ -349,7 +357,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
+    size_t off_bias;          // bias in VTCM

     // Cached sizes of regions for HMX kernel use
     size_t weight_area_bytes;
@@ -358,25 +366,23 @@ 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 bias_bytes;

     size_t total_bytes;
 };

 struct htp_mm_hvx_vtcm_layout {
     // Byte offsets from vtcm_base for each region
-    size_t off_src1;          // vtcm_src1 (activation)
+    size_t off_act;           // vtcm_act (activation)
     size_t off_src0;          // vtcm_src0 (weight/Wk)
-    size_t off_src2;          // vtcm_src2 (Wq / fused only)
-    size_t off_src3;          // vtcm_src3 (Wv / fused only)
+    size_t off_bias;          // vtcm_bias (bias / fused add only)
     size_t off_dst;           // vtcm_dst (output scratch)
     size_t off_act_raw;       // vtcm_act_raw (raw activation DMA staging)

     // Cached sizes
     size_t src0_bytes;
-    size_t src1_bytes;
-    size_t src2_bytes;
-    size_t src3_bytes;
+    size_t act_bytes;
+    size_t bias_bytes;
     size_t dst_bytes;
     size_t act_raw_bytes;

@@ -394,7 +400,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
     bool pipeline,
     uint32_t act_threads,
     uint32_t aligned_tile_size,
-    size_t src2_size
+    size_t bias_size
 ) {
     size_t off = 0;

@@ -411,7 +417,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);
+        VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE), bias_size > 0);

         // Group B: Compute-only buffers (starts at off_group_a)
         size_t off_group_b = off_group_a;
@@ -439,7 +445,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;
+        L->bias_bytes        = bias_size;

         off = off_group_a + hex_smax(group_b_size, group_c_size);
     } else {
@@ -463,7 +469,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);
+        VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE), bias_size > 0);

         // Group B: Compute-only buffers (starts at off_group_a)
         size_t off_group_b = off_group_a;
@@ -491,7 +497,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;
+        L->bias_bytes        = bias_size;

         off = off_group_a + hex_smax(group_b_size, group_c_size);
     }
@@ -504,21 +510,20 @@ static inline void htp_mm_hvx_vtcm_layout_build(
     int kernel_type,
     int wtype,
     uint32_t ne10,       // k
-    uint32_t src1_nrows, // m_total
+    uint32_t act_nrows,  // m_total
     uint32_t n_threads,
     size_t dst_row_size,
     size_t src0_row_size,
-    size_t src1_row_size,
-    size_t src2_row_size,
+    size_t act_row_size,
+    size_t bias_row_size,
     uint32_t n_prefetch,
     bool is_matmul_id,
     bool is_fused_nx
 ) {
-    (void)src1_row_size;
+    (void)act_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;
-    size_t src3_sz    = 0;
+    size_t act_sz     = 0;
+    size_t bias_sz    = bias_row_size > 0 ? htp_mm_round_up(bias_row_size, 128) : 0;
     size_t dst_sz     = 0;
     size_t act_raw_sz = 0;

@@ -544,22 +549,21 @@ static inline void htp_mm_hvx_vtcm_layout_build(
         }

         size_t tiled_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
-        size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);
+        size_t q_act_sz = hex_round_up(tiled_act_row_size * act_nrows, 128);
         size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));

         src0_sz    = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
-        src1_sz    = act_sz;                           // quantized activation buffer
-        src2_sz    = 0;
-        src3_sz    = 0;
+        act_sz     = q_act_sz;                         // quantized activation buffer
+        bias_sz    = 0;
         dst_sz     = 0;
-        act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+        act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
     } else if (is_matmul_id) {
         const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
-        const size_t src1_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
-                                                                                               : htp_mm_q8_0_tiled_row_size(ne10);
+        const size_t act_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
+                                                                          : htp_mm_q8_0_tiled_row_size(ne10);

         size_t src0_sz_per_thread = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
-        src1_sz                   = htp_mm_round_up(src1_row_size_tiled * src1_nrows, 256);
+        act_sz                    = htp_mm_round_up(act_row_size_tiled * act_nrows, 256);

         if (is_repack) {
             const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype);
@@ -573,25 +577,24 @@ static inline void htp_mm_hvx_vtcm_layout_build(

         src0_sz    = src0_sz_per_thread * n_threads;
         dst_sz     = 0;
-        src2_sz    = 0;
-        src3_sz    = 0;
-        act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+        bias_sz    = 0;
+        act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
     } 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;
+        const size_t dst_nrows = (act_nrows > 1) ? 0 : 1;

         switch (kernel_type) {
             case HTP_MM_KERNEL_HVX_F16_F16_VTCM: {
-                size_t f16_src1_row_size = htp_mm_round_up(ne10 * 2, 128);
-                src1_sz    = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
+                size_t f16_act_row_size = htp_mm_round_up(ne10 * 2, 128);
+                act_sz     = htp_mm_round_up(f16_act_row_size * act_nrows, 256);
                 src0_sz    = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
                 dst_sz     = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
-                act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * src1_nrows, 128);
+                act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * act_nrows, 128);
                 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);
+                size_t f32_act_row_size = htp_mm_round_up(ne10 * 4, 128);
+                act_sz     = htp_mm_round_up(f32_act_row_size * act_nrows, 256);
                 src0_sz    = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
                 dst_sz     = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
                 act_raw_sz = 0;
@@ -599,10 +602,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
             }
             case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
             case HTP_MM_KERNEL_HVX_QUANT_ROW: {
-                size_t q_src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+                size_t q_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_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);
+                act_sz  = htp_mm_round_up(q_act_row_size * act_nrows, 256);

                 src0_sz = src0_sz * n_threads;

@@ -614,10 +617,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
                     src0_sz = repacked_vtcm_size * n_threads;
                 }

-                size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
+                size_t dst_slice_per_thread = (dst_nrows > 0 && act_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
                 dst_sz = dst_slice_per_thread * n_threads;
                 size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
-                act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+                act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
                 break;
             }
             default:
@@ -627,9 +630,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(

     // Group A: Persistent buffers across chunk compute
     size_t off_group_a = 0;
-    VTCM_LAYOUT_ALLOC(off_group_a, off_src1, src1_sz);
-    VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, src2_sz, src2_sz > 0);
-    VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src3, src3_sz, src3_sz > 0);
+    VTCM_LAYOUT_ALLOC(off_group_a, off_act, act_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, bias_sz, bias_sz > 0);

     // Group B: Compute-only buffers (starts at off_group_a)
     size_t off_group_b = off_group_a;
@@ -643,9 +645,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
     const size_t group_c_size = off_group_c - off_group_a;

     L->src0_bytes    = src0_sz;
-    L->src1_bytes    = src1_sz;
-    L->src2_bytes    = src2_sz;
-    L->src3_bytes    = src3_sz;
+    L->act_bytes     = act_sz;
+    L->bias_bytes    = bias_sz;
     L->dst_bytes     = dst_sz;
     L->act_raw_bytes = act_raw_sz;
     L->total_bytes   = off_group_a + hex_smax(group_b_size, group_c_size);
@@ -655,12 +656,12 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
     int kernel_type,
     int wtype,
     uint32_t ne10,
-    uint32_t src1_nrows,
+    uint32_t act_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,
+    size_t act_row_size,
+    size_t bias_row_size,
     uint32_t n_prefetch,
     size_t vtcm_budget,
     struct htp_mm_hvx_vtcm_layout * L_out,
@@ -668,17 +669,17 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
 ) {
     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
+        &L, kernel_type, wtype, ne10, act_nrows, n_threads,
+        dst_row_size, src0_row_size, act_row_size, bias_row_size, n_prefetch, false, false
     );

     if (L.total_bytes <= vtcm_budget) {
         *L_out = L;
-        *m_chunk_out = src1_nrows;
+        *m_chunk_out = act_nrows;
         return true;
     }

-    const size_t fixed_bytes = L.src0_bytes + L.src2_bytes + L.dst_bytes;
+    const size_t fixed_bytes = L.src0_bytes + L.bias_bytes + L.dst_bytes;
     if (vtcm_budget <= fixed_bytes) {
         return false;
     }
@@ -707,8 +708,8 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
     if (m_chunk > 1) {
         m_chunk &= ~1U;
     }
-    if (m_chunk > src1_nrows) {
-        m_chunk = src1_nrows;
+    if (m_chunk > act_nrows) {
+        m_chunk = act_nrows;
     }
     if (m_chunk < 1) {
         return false;
@@ -716,14 +717,14 @@ static inline bool htp_mm_hvx_solve_vtcm_params(

     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
+        dst_row_size, src0_row_size, act_row_size, bias_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
+            dst_row_size, src0_row_size, act_row_size, bias_row_size, n_prefetch, false, false
         );
     }

@@ -737,18 +738,18 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
 }

 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, size_t src2_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 bias_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, pipeline, act_threads, aligned_tile_size, src2_size);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, pipeline, act_threads, aligned_tile_size, bias_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 pipeline, uint32_t act_threads, size_t src2_size) {
+    int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool pipeline, uint32_t act_threads, size_t bias_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, false, act_threads, 0, src2_size);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, false, act_threads, 0, bias_size);
     return L.total_bytes;
 }

@@ -760,7 +761,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
     uint32_t group_size,
     int n_threads,
     bool pipeline,
-    size_t src2_size,
+    size_t bias_size,
     size_t vtcm_budget,
     size_t * m_chunk_out,
     size_t * n_chunk_out,
@@ -775,7 +776,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() + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
+        size_t group_overhead = htp_mm_hmx_get_batched_overhead() + (bias_size > 0 ? hex_align_up(bias_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);

@@ -785,8 +786,8 @@ 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, pipeline, act_threads, src2_size);
+                                &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, pipeline, act_threads, bias_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)) {
@@ -826,7 +827,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 bias_size,
     size_t vtcm_budget,
     size_t * m_chunk_out,
     size_t * n_chunk_out,
@@ -843,7 +844,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) + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
+        size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id) + (bias_size > 0 ? hex_align_up(bias_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);

@@ -854,7 +855,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, src2_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, bias_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/scripts/snapdragon/run.py b/scripts/snapdragon/run.py
index 01093ddd8..14e53aee6 100755
--- a/scripts/snapdragon/run.py
+++ b/scripts/snapdragon/run.py
@@ -31,8 +31,10 @@ MANAGED_ENV_NAMES = (
     "GGML_HEXAGON_MBUF",
     "GGML_HEXAGON_MM_SELECT",
     "GGML_HEXAGON_FA_SELECT",
+    "GGML_HEXAGON_FA_HEAD_SPLIT",
     "GGML_HEXAGON_GDN_SELECT",
     "GGML_HEXAGON_AR_SELECT",
+    "GGML_HEXAGON_AR_SCATTER",
     "GGML_HEXAGON_ETM",
     "GGML_HEXAGON_ARCH",
     "GGML_HEXAGON_OPTRACE",
@@ -167,6 +169,7 @@ def main():
     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) 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-fa-head-split", help="Enable (1) or disable (0) head-parallel flash_attn partitioning (GGML_HEXAGON_FA_HEAD_SPLIT)")
     parser.add_argument("--hex-gdn-select", help="Select Gated Delta Net kernel (GGML_HEXAGON_GDN_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-ar-scatter", help="Enable (1) or disable (0) reduce-scatter for fused ALLREDUCE+ADD (GGML_HEXAGON_AR_SCATTER)")
@@ -309,6 +312,7 @@ def main():
     set_env("GGML_HEXAGON_MBUF", args.hex_mbuf)
     set_env("GGML_HEXAGON_MM_SELECT", args.hex_mm_select)
     set_env("GGML_HEXAGON_FA_SELECT", args.hex_fa_select)
+    set_env("GGML_HEXAGON_FA_HEAD_SPLIT", args.hex_fa_head_split)
     set_env("GGML_HEXAGON_GDN_SELECT", args.hex_gdn_select)
     set_env("GGML_HEXAGON_AR_SELECT", args.hex_ar_select)
     set_env("GGML_HEXAGON_AR_SCATTER", args.hex_ar_scatter)