Commit 07fc586e3 for llama.cpp

commit 07fc586e385fd1f6b79366130ff77f1b53d957ec
Author: kurquhar <kurquhar@qti.qualcomm.com>
Date:   Thu Sep 24 11:39:28 2026 -0700

    hexagon: dynamic quantizer improvements (#29395)

    * hexagon: fix accuracy issue in Q8_0 N=1 MUL_MAT

    * hex-quant: fix register spills

    * hex-mm: use dma for all dyn.quant paths

    Co-authored-by: Aparna M P <aparmp@qti.qualcomm.com>

    * hex-mm: remove obsolete run_quant_task

    * hex-mm: update tracing to properly wrap the events

    * hex-mm: use act for activation data in all paths

    * hex-mm: use act_ instead of src1_ to avoid confusion in fused kernels

    * hex-mm: remove/reroute the rest of the non-DMA act (aka src1) logic

    * hex-dma64: yet another pass at cleaning up the dma_addr_t casts

    * Update ggml/src/ggml-hexagon/htp/matmul-ops.h

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

    * Update ggml/src/ggml-hexagon/htp/matmul-ops.c

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

    * Update ggml/src/ggml-hexagon/htp/matmul-ops.c

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

    * Update ggml/src/ggml-hexagon/htp/matmul-ops.c

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

    ---------

    Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>
    Co-authored-by: Aparna M P <aparmp@qti.qualcomm.com>
    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index cc62b0f71..9d3c97a99 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -4494,8 +4494,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(

     if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
         // Try grouped path first
-        const bool use_dma_activation = (src1->nb[1]/sizeof(float) > (size_t)ne00_padded);
-        if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, use_dma_activation, n_threads, pipeline, 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, ne11, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
             use_grouped = true;
         }
     }
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 d6d40586c..b75f601cc 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
@@ -914,98 +914,6 @@ typedef struct {

 // activations : fp32 -> fp16

-static void transfer_activation_chunk_fp32_to_fp16(__fp16 *restrict vtcm_dst, const float *restrict src, uint32_t n_rows, uint32_t k_block, uint32_t k_stride, uint32_t k_valid) {
-    const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
-    const uint32_t n_rows_tiled  = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
-
-    uint32_t r = 0;
-
-    #pragma unroll(2)
-    for (r = 0; r < n_rows_tiled; r += 2) {
-        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 float *ptr_in0 = src + (r + 0) * k_stride;
-        const float *ptr_in1 = src + (r + 1) * k_stride;
-
-        uint32_t c = 0;
-        for (; c + 32 <= k_valid; c += 32) {
-            HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
-            HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
-            HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
-
-            uint32_t c0       = c / HTP_MM_HMX_TILE_N_COLS;  // tile column index
-            uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
-
-            HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
-            tile[r1 / 2]     = v_out;
-        }
-        if (c < k_block) {
-            HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
-            HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
-
-            uint32_t rem = k_valid - c;
-            HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
-            v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
-            v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
-
-            HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
-
-            uint32_t c0       = c / HTP_MM_HMX_TILE_N_COLS;  // tile column index
-            uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
-
-            HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
-            tile[r1 / 2]     = v_out;
-        }
-    }
-
-    for (; r < n_rows_padded; r += 2) {
-        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 bool row0_valid = r       < n_rows;
-        const bool row1_valid = (r + 1) < n_rows;
-
-        const float *ptr_in0 = row0_valid ? (src + (r + 0) * k_stride) : NULL;
-        const float *ptr_in1 = row1_valid ? (src + (r + 1) * k_stride) : NULL;
-
-        uint32_t c = 0;
-        for (; c + 32 <= k_valid; c += 32) {
-            HVX_Vector v0 = Q6_V_vzero();
-            HVX_Vector v1 = Q6_V_vzero();
-            if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
-            if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
-
-            HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
-
-            uint32_t c0       = c / HTP_MM_HMX_TILE_N_COLS;  // tile column index
-            uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
-
-            HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
-            tile[r1 / 2]     = v_out;
-        }
-        if (c < k_block) {
-            HVX_Vector v0 = Q6_V_vzero();
-            HVX_Vector v1 = Q6_V_vzero();
-            if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
-            if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
-
-            uint32_t rem = k_valid - c;
-            HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
-            v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
-            v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
-
-            HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
-
-            uint32_t c0       = c / HTP_MM_HMX_TILE_N_COLS;  // tile column index
-            uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
-
-            HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
-            tile[r1 / 2]     = v_out;
-        }
-    }
-}
-
 static void transfer_activation_row_pair_fp32_to_fp16(
         __fp16 *restrict vtcm_dst,
         const float *restrict row0,
diff --git a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
index 4d6110ffa..5706259e1 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
@@ -121,29 +121,43 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
     HVX_Vector * vx = (HVX_Vector *) x;
     HVX_Vector zero   = Q6_V_vzero();

+    HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
+    HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
+    HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
+    HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
+
     HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
     HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
     HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
     HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);

+    HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
+    HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
+    HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
+    HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
+
+    HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
+    HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
+
     HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
     HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));

-    HVX_Vector vmax_hf = hvx_vec_reduce_max_f16(hvx_vec_abs_f16(vx01_hf));
-    vmax_hf            = hvx_vec_reduce_max2_f16(hvx_vec_abs_f16(vx23_hf), vmax_hf);
-
-    HVX_Vector vd_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax_hf, Q6_Vh_vsplat_R(0x2008));
-    HVX_Vector vd_hf   = Q6_Vhf_equals_Vqf16(vd_qf16);
+    HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008));
+    HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008));
+    HVX_Vector vd01_hf   = Q6_Vhf_equals_Vqf16(vd01_qf16);
+    HVX_Vector vd23_hf   = Q6_Vhf_equals_Vqf16(vd23_qf16);

-    HVX_Vector vd_inv_hf = hvx_vec_inverse_f16(vd_hf);
-    vx01_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd_inv_hf));
-    vx23_hf              = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd_inv_hf));
+    HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
+    HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
+    vx01_hf                = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
+    vx23_hf                = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));

     HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
     HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
     HVX_Vector vx_i8    = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);

-    HVX_Vector r_scale = hvx_vec_repl_f16(vd_hf);
+    HVX_VectorPair vp01 = Q6_W_vshuff_VVR(vd01_hf, vd01_hf, -64);
+    HVX_VectorPair vp23 = Q6_W_vshuff_VVR(vd23_hf, vd23_hf, -64);

     static const uint8_t __attribute__((aligned(128))) repl[128] = {
         0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
@@ -157,8 +171,19 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
     };
     HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;

+    #pragma unroll
     for (int b = 0; b < 4; b++) {
         HVX_Vector v_act = Q6_V_vror_VR(vx_i8, b * 32);
+        HVX_Vector r_scale;
+        if (b == 0) {
+            r_scale = Q6_V_lo_W(vp01);
+        } else if (b == 1) {
+            r_scale = Q6_V_hi_W(vp01);
+        } else if (b == 2) {
+            r_scale = Q6_V_lo_W(vp23);
+        } else {
+            r_scale = Q6_V_hi_W(vp23);
+        }

         HVX_Vector r0 = Q6_V_vdelta_VV(v_act, v_repl_ctrl);
         HVX_Vector r1 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 4),  v_repl_ctrl);
@@ -941,14 +966,9 @@ static inline void quantize_f32_q8_0_tiled_kernel(
     size_t src_row_size,
     size_t dst_row_size
 ) {
-    const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
-    hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
-
+    (void) tmp_data;
     for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_row_size, 2);
-        hvx_copy_f32_aa(tmp_data, src_data, ne0);
-
-        quantize_row_f32_q8_0_tiled((float *) tmp_data, dst_data, ne0);
+        quantize_row_f32_q8_0_tiled((float *) src_data, dst_data, ne0);
         dst_data += dst_row_size;
         src_data += src_row_size;
     }
@@ -963,14 +983,9 @@ static inline void quantize_f32_q8_1_tiled_kernel(
     size_t src_row_size,
     size_t dst_row_size
 ) {
-    const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
-    hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
-
+    (void) tmp_data;
     for (uint32_t i = 0; i < nrows; ++i) {
-        hex_l2fetch(src_data, src_row_size, src_row_size, 2);
-        hvx_copy_f32_aa(tmp_data, src_data, ne0);
-
-        quantize_row_f32_q8_1_tiled((float *) tmp_data, dst_data, ne0);
+        quantize_row_f32_q8_1_tiled((float *) src_data, dst_data, ne0);
         dst_data += dst_row_size;
         src_data += src_row_size;
     }
@@ -988,24 +1003,15 @@ static inline void quantize_f32_q8_0_tiled_block_kernel(
     uint32_t r,
     uint32_t c
 ) {
+    (void) tmp_data;
     const uint32_t qk = QK_Q8_0_TILED;
     const uint32_t nb = (ne0 + qk - 1) / qk;

     for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
-        const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
+        const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
         uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1152;

-        hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
-
-        if (c == nb - 1) {
-            uint32_t active_elements = ne0 - c * qk;
-            hvx_splat_f32_a(tmp_data, 0.0f, qk);
-            hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
-        } else {
-            hvx_copy_f32_aa(tmp_data, src_ptr, qk);
-        }
-
-        quantize_block_f32_q8_0_tiled((float *) tmp_data, dst_ptr);
+        quantize_block_f32_q8_0_tiled((float *) src_ptr, dst_ptr);

         c++;
         if (c == nb) {
@@ -1027,24 +1033,15 @@ static inline void quantize_f32_q8_1_tiled_block_kernel(
     uint32_t r,
     uint32_t c
 ) {
+    (void) tmp_data;
     const uint32_t qk = QK_Q8_0_TILED;
     const uint32_t nb = (ne0 + qk - 1) / qk;

     for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
-        const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
+        const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
         uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1280;

-        hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
-
-        if (c == nb - 1) {
-            uint32_t active_elements = ne0 - c * qk;
-            hvx_splat_f32_a(tmp_data, 0.0f, qk);
-            hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
-        } else {
-            hvx_copy_f32_aa(tmp_data, src_ptr, qk);
-        }
-
-        quantize_block_f32_q8_1_tiled((float *) tmp_data, dst_ptr);
+        quantize_block_f32_q8_1_tiled((float *) src_ptr, dst_ptr);

         c++;
         if (c == nb) {
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c
index a09bc7a28..2b45112b8 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.c
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c
@@ -29,7 +29,7 @@ typedef struct {
     float        *dst;
     dma_addr_t    src2_addr;
     size_t        src2_bytes;
-    const float  *activation;
+    dma_addr_t    act_dma_addr;
     dma_addr_t    weight;
     dma_queue *   weight_dma;
     int           m;
@@ -45,8 +45,8 @@ typedef struct {
     int           ne13;
     size_t        src0_nb2;
     size_t        src0_nb3;
-    size_t        src1_nb2;
-    size_t        src1_nb3;
+    size_t        act_nb2;
+    size_t        act_nb3;
     size_t        src2_nb2;
     size_t        src2_nb3;
     size_t        dst_nb2;
@@ -93,7 +93,7 @@ struct htp_mm_context {
     uint32_t src0_row_start;
     uint32_t src0_row_end;
     uint32_t src0_row_size_padded;
-    uint32_t src1_nrows;
+    uint32_t act_nrows;
     uint32_t cur_m_start;
     uint32_t cur_m_rows;

@@ -103,16 +103,13 @@ struct htp_mm_context {
     struct fastdiv_values mm_div_r3;
     struct fastdiv_values mm_div_ne11;

-    // Per thread quant tasks
     // Precomputed block-parallel quantization values
-    worker_callback_t quant_task_func;
     uint32_t          quant_ib_first[WORK_QUEUE_MAX_N_THREADS];
     uint32_t          quant_ib_last[WORK_QUEUE_MAX_N_THREADS];
     uint32_t          quant_r[WORK_QUEUE_MAX_N_THREADS];
     uint32_t          quant_c[WORK_QUEUE_MAX_N_THREADS];
     uint32_t          n_quant_tasks;
     uint32_t          n_quant_rows_per_thread;
-    atomic_uint       quant_barrier;

     // Fields for scattered mapping & HMX support in MUL_MAT_ID
     const uint32_t * matrix_row_counts;
@@ -125,12 +122,14 @@ struct htp_mm_context {
     uint8_t * vtcm_src2;
     uint8_t * vtcm_src3;
     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_raw_stride;

     // Cached thread offsets/sizes
     uint32_t vtcm_src0_size_per_thread;
@@ -234,17 +233,6 @@ static const uint8_t __attribute__((aligned(VLEN))) kvalues_mxfp4_lut[] = {
     uint32_t src0_nrows_per_thread = mmctx->src0_nrows_per_thread;  \
     htp_matmul_tensors_preamble;

-static inline void hvx_mm_run_quant_task(struct htp_mm_context * mmctx, unsigned int ith) {
-    if (mmctx->quant_task_func) {
-        if (ith < mmctx->n_quant_tasks) {
-            mmctx->quant_task_func(mmctx->n_quant_tasks, ith, mmctx);
-            atomic_fetch_sub(&mmctx->quant_barrier, 1);
-        }
-        while (atomic_load(&mmctx->quant_barrier) > 0) {
-            // spin
-        }
-    }
-}



@@ -301,8 +289,6 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
         }                                                                                                                                  \
     }                                                                                                                                      \
                                                                                                                                            \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                                     \
-                                                                                                                                           \
     if (src0_start_row >= src0_end_row) {                                                                                                  \
         return;                                                                                                                            \
     }                                                                                                                                      \
@@ -416,8 +402,6 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
         }                                                                                                                \
     }                                                                                                                    \
                                                                                                                          \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                   \
-                                                                                                                         \
     if (src0_start_row >= src0_end_row) {                                                                                \
         return;                                                                                                          \
     }                                                                                                                    \
@@ -479,8 +463,6 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
     uint32_t n_k_tiles_a = ne10 / 32;                                                                                             \
     uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size;                                                    \
                                                                                                                                   \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                            \
-                                                                                                                                  \
     for (uint32_t widx = 0; widx < n_weights; widx++) {                                                                           \
         const struct htp_tensor * restrict src_w = octx->src[widx];                                                               \
         const struct htp_tensor * restrict dst   = octx->dsts[widx];                                                              \
@@ -565,60 +547,120 @@ MATMUL_2D_REPACKED_IMPL(q6_k,       896,  tiled_vec_dot_q6_k_32x2,  tiled_vec_do
 MATMUL_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
 MATMUL_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)

-#define QUANTIZE_IMPL(name, log_name, kernel_fn, dst_row_size_expr)                                                                                               \
-static void name(unsigned int nth, unsigned int ith, void * data) {                                                                                               \
-    struct htp_mm_context * mmctx = data;                                                                                                                         \
-    struct htp_ops_context * octx = mmctx->octx;                                                                                                                  \
-    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;                                                      \
-    const struct htp_tensor * src = mmctx->act;                                                                                                                   \
-    const uint32_t ne0 = src->ne[0];                                                                                                                              \
-    const uint32_t nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->src1_nrows;                                                                             \
-    const uint32_t nrows_per_thread = mmctx->n_quant_rows_per_thread;                                                                                             \
-                                                                                                                                                                  \
-    const uint32_t ir_first = nrows_per_thread * ith;                                                                                                             \
-    if (ir_first >= nrows) {                                                                                                                                      \
-        return;                                                                                                                                                   \
-    }                                                                                                                                                             \
-                                                                                                                                                                  \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                                                        \
-    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                                                                               \
-                                                                                                                                                                  \
-    uint8_t * restrict dst = mmctx->vtcm_src1;                                                                                                                    \
-    const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows);                                                                                             \
-    const size_t src_row_size = src->nb[1];                                                                                                                       \
-    const size_t dst_row_size = (dst_row_size_expr);                                                                                                              \
-    uint8_t * restrict tmp_data = (uint8_t *) mmctx->vtcm_dst + (mmctx->vtcm_dst_size_per_thread * ith);                                                          \
-                                                                                                                                                                  \
-    const bool is_contiguous = (src->nb[2] == src->ne[1] * src->nb[1]) && (src->nb[3] == src->ne[2] * src->nb[2]);                                                \
-    if (is_contiguous) {                                                                                                                                          \
-        const uint8_t * restrict src_data = (const uint8_t *) src->data + (src_row_size * (mmctx->cur_m_start + ir_first));                                       \
-        uint8_t * restrict dst_data = (uint8_t *) dst + (dst_row_size * ir_first);                                                                                \
-        kernel_fn(src_data, dst_data, tmp_data, ne0, ir_last - ir_first, src_row_size, dst_row_size);                                                             \
-    } else {                                                                                                                                                      \
-        const uint32_t ne12_ne1 = src->ne[2] * src->ne[1];                                                                                                        \
-        for (uint32_t ir = ir_first; ir < ir_last; ++ir) {                                                                                                        \
-            const uint32_t ir1 = mmctx->cur_m_start + ir;                                                                                                         \
-            const uint32_t i13 = fastdiv(ir1, &kparams->div_ne12_ne1);                                                                                            \
-            const uint32_t rem = ir1 - i13 * ne12_ne1;                                                                                                            \
-            const uint32_t i12 = fastdiv(rem, &kparams->div_ne1);                                                                                                 \
-            const uint32_t i11 = rem - i12 * src->ne[1];                                                                                                          \
-            const uint8_t * restrict row_src = (const uint8_t *) src->data + ((size_t) i11 * src->nb[1] + (size_t) i12 * src->nb[2] + (size_t) i13 * src->nb[3]); \
-            uint8_t * restrict row_dst = dst + (dst_row_size * ir);                                                                                               \
-            kernel_fn(row_src, row_dst, tmp_data, ne0, 1, src_row_size, dst_row_size);                                                                            \
-        }                                                                                                                                                         \
-    }                                                                                                                                                             \
-                                                                                                                                                                  \
-    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                                                                                \
+static void hvx_mm_transfer_src1_dma(
+    struct htp_ops_context * octx,
+    const struct htp_mm_kernel_params * kparams,
+    const struct htp_tensor * src1,
+    uint8_t * dst_base,
+    size_t dst_row_size,
+    uint32_t m_start,
+    uint32_t m_rows
+) {
+    if (m_rows == 0) {
+        return;
+    }
+
+    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 size_t row_bytes = ne0 * elem_size;
+    const size_t src1_nb1 = src1->nb[1];
+    const dma_addr_t src_base = src1->data;
+
+    const bool is_contiguous = (src1->nb[2] == src1->ne[1] * src1_nb1) &&
+                               (src1->nb[3] == src1->ne[2] * src1->nb[2]);
+
+    if (is_contiguous) {
+        const dma_addr_t src_addr = src_base + m_start * src1_nb1;
+        dma_queue_push(dma_q, dma_make_data(dst_base, src_addr),
+                       dst_row_size, src1_nb1, row_bytes, m_rows);
+        dma_queue_pop(dma_q);
+    } else {
+        const uint32_t ne12_ne1 = src1->ne[2] * src1->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;
+            uint32_t i11, i12, i13;
+            if (use_fastdiv) {
+                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];
+            } 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];
+            }
+            const dma_addr_t row_src = src_base + (i11 * src1->nb[1] +
+                                                   i12 * src1->nb[2] +
+                                                   i13 * src1->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);
+            dma_queue_pop(dma_q);
+        }
+    }
+
+    if (dst_row_size > row_bytes) {
+        struct htp_thread_trace * tr = &octx->ctx->trace[0];
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) m_start);
+        const uint32_t pad_elems = (dst_row_size - row_bytes) / elem_size;
+        if (elem_size == sizeof(float)) {
+            for (uint32_t ir = 0; ir < m_rows; ++ir) {
+                hvx_splat_f32_u(dst_base + ir * dst_row_size + row_bytes, 0.0f, pad_elems);
+            }
+        } else {
+            for (uint32_t ir = 0; ir < m_rows; ++ir) {
+                hvx_splat_f16_u(dst_base + ir * dst_row_size + row_bytes, (_Float16) 0.0f, pad_elems);
+            }
+        }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) m_start);
+    }
+}
+
+#define QUANTIZE_IMPL(name, log_name, kernel_fn, dst_row_size_expr)                                        \
+static void name(unsigned int nth, unsigned int ith, void * data) {                                        \
+    (void) nth;                                                                                            \
+    struct htp_mm_context * mmctx = data;                                                                  \
+    struct htp_ops_context * octx = mmctx->octx;                                                           \
+    const struct htp_tensor * src = mmctx->act;                                                            \
+    const uint32_t ne0 = src->ne[0];                                                                       \
+    const uint32_t nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows;                       \
+    const uint32_t nrows_per_thread = mmctx->n_quant_rows_per_thread;                                      \
+                                                                                                           \
+    const uint32_t ir_first = nrows_per_thread * ith;                                                      \
+    if (ir_first >= nrows) {                                                                               \
+        return;                                                                                            \
+    }                                                                                                      \
+                                                                                                           \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                 \
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                        \
+                                                                                                           \
+    uint8_t * restrict dst = mmctx->vtcm_src1;                                                             \
+    const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows);                                      \
+    const size_t raw_row_size = mmctx->vtcm_act_raw_stride;                                                \
+    const size_t dst_row_size = (dst_row_size_expr);                                                       \
+                                                                                                           \
+    const uint8_t * restrict src_data = (const uint8_t *) mmctx->vtcm_act_raw + (raw_row_size * ir_first); \
+    uint8_t * restrict dst_data = (uint8_t *) dst + (dst_row_size * ir_first);                             \
+    kernel_fn(src_data, dst_data, NULL, ne0, ir_last - ir_first, raw_row_size, dst_row_size);              \
+                                                                                                           \
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first);                                         \
 }

 QUANTIZE_IMPL(quantize_f32_q8_0_tiled, "quantize-f32-q8_0_tiled", quantize_f32_q8_0_tiled_kernel, htp_mm_q8_0_tiled_row_size(ne0))
 QUANTIZE_IMPL(quantize_f32_q8_1_tiled, "quantize-f32-q8_1_tiled", quantize_f32_q8_1_tiled_kernel, htp_mm_q8_1_tiled_row_size(ne0))
-QUANTIZE_IMPL(quantize_f32_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_src1_stride)
+QUANTIZE_IMPL(quantize_f32_f16,        "quantize-f32-f16",        quantize_f32_f16_kernel,        mmctx->vtcm_src1_stride)
+QUANTIZE_IMPL(quantize_f16_f16,        "quantize-f16-f16",        quantize_f16_f16_kernel,        mmctx->vtcm_src1_stride)

 static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
     struct htp_mm_context * mmctx = data;
+    if (mmctx->quant_ib_first[ith] >= mmctx->quant_ib_last[ith]) {
+        return;
+    }
     struct htp_ops_context * octx = mmctx->octx;
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
     htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, mmctx->quant_ib_first[ith]);
@@ -626,13 +668,13 @@ static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, vo
     const struct htp_tensor * src = mmctx->act;

     quantize_f32_q8_0_tiled_block_kernel(
-        (const float *) src->data,
+        (const float *) mmctx->vtcm_act_raw,
         mmctx->vtcm_src1,
-        (uint8_t *) mmctx->vtcm_dst + (mmctx->vtcm_dst_size_per_thread * ith),
+        NULL,
         src->ne[0],
         mmctx->quant_ib_first[ith],
         mmctx->quant_ib_last[ith],
-        src->nb[1],
+        mmctx->vtcm_act_raw_stride,
         htp_mm_q8_0_tiled_row_size(src->ne[0]),
         mmctx->quant_r[ith],
         mmctx->quant_c[ith]
@@ -642,7 +684,11 @@ static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, vo
 }

 static void quantize_f32_q8_1_tiled_block(unsigned int nth, unsigned int ith, void * data) {
+    (void) nth;
     struct htp_mm_context * mmctx = data;
+    if (mmctx->quant_ib_first[ith] >= mmctx->quant_ib_last[ith]) {
+        return;
+    }
     struct htp_ops_context * octx = mmctx->octx;
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
     htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, mmctx->quant_ib_first[ith]);
@@ -650,13 +696,13 @@ static void quantize_f32_q8_1_tiled_block(unsigned int nth, unsigned int ith, vo
     const struct htp_tensor * src = mmctx->act;

     quantize_f32_q8_1_tiled_block_kernel(
-        (const float *) src->data,
+        (const float *) mmctx->vtcm_act_raw,
         mmctx->vtcm_src1,
-        (uint8_t *) mmctx->vtcm_dst + (mmctx->vtcm_dst_size_per_thread * ith),
+        NULL,
         src->ne[0],
         mmctx->quant_ib_first[ith],
         mmctx->quant_ib_last[ith],
-        src->nb[1],
+        mmctx->vtcm_act_raw_stride,
         htp_mm_q8_1_tiled_row_size(src->ne[0]),
         mmctx->quant_r[ith],
         mmctx->quant_c[ith]
@@ -672,11 +718,11 @@ MATVEC_2D_REPACKED_IMPL(q6_k,       896,  tiled_vec_dot_q6_k_32x1)
 MATVEC_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x1)
 MATVEC_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x1)

-MATMUL_NX_2D_REPACKED_IMPL(q4_0,       576,  tiled_vec_dot_q4_0_32x2,  tiled_vec_dot_q4_0_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(q4_1,       640,  tiled_vec_dot_q4_1_32x2,  tiled_vec_dot_q4_1_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(q8_0,       1088, tiled_vec_dot_q8_0_32x2,  tiled_vec_dot_q8_0_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(iq4nl,      576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
-MATMUL_NX_2D_REPACKED_IMPL(mxfp4,      544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)
+MATMUL_NX_2D_REPACKED_IMPL(q4_0,    576,  tiled_vec_dot_q4_0_32x2,  tiled_vec_dot_q4_0_32x1)
+MATMUL_NX_2D_REPACKED_IMPL(q4_1,    640,  tiled_vec_dot_q4_1_32x2,  tiled_vec_dot_q4_1_32x1)
+MATMUL_NX_2D_REPACKED_IMPL(q8_0,    1088, tiled_vec_dot_q8_0_32x2,  tiled_vec_dot_q8_0_32x1)
+MATMUL_NX_2D_REPACKED_IMPL(iq4nl,   576,  tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
+MATMUL_NX_2D_REPACKED_IMPL(mxfp4,   544,  tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)

 #define MATMUL_4D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1)                                                                                        \
 static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                                                  \
@@ -713,7 +759,6 @@ static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
     const uint32_t ct_start = src0_start_row / 32;                                                                                                          \
     const uint32_t ct_end   = (src0_end_row + 31) / 32;                                                                                                     \
                                                                                                                                                             \
-    hvx_mm_run_quant_task(mmctx, ith);                                                                                                                      \
                                                                                                                                                             \
     if (src0_start_row >= src0_end_row || cur_m_rows == 0) {                                                                                                \
         return;                                                                                                                                             \
@@ -822,7 +867,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->src1_nrows;                          // src1 rows
+    const uint32_t src1_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows;                          // src1 rows
     const uint32_t cur_m_start = mmctx->cur_m_start;

     const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
@@ -845,6 +890,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {

     const dma_addr_t src0_row = src0->data;

+
     // Prefill vtcm with src0 rows
     if (src0_start_row < src0_end_row) {
         for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
@@ -857,8 +903,6 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
         }
     }

-    hvx_mm_run_quant_task(mmctx, ith);
-
     if (src0_start_row >= src0_end_row) {
         return;
     }
@@ -974,7 +1018,6 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
         }
     }

-    hvx_mm_run_quant_task(mmctx, ith);

     if (src0_start_row >= src0_end_row) {
         return;
@@ -1049,7 +1092,6 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
     uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
     uint8_t * restrict src1_data     = mmctx->vtcm_src1;

-    hvx_mm_run_quant_task(mmctx, ith);

     if (src0_start_row >= src0_end_row || cur_m_rows == 0) {
         return;
@@ -1183,7 +1225,6 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
     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);

-    hvx_mm_run_quant_task(mmctx, ith);

     if (src0_start_row >= src0_end_row) {
         return;
@@ -1272,7 +1313,6 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
     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);

-    hvx_mm_run_quant_task(mmctx, ith);

     if (src0_start_row >= src0_end_row) {
         return;
@@ -1351,7 +1391,6 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
     const struct htp_tensor * restrict act  = octx->src[n_weights];
     const struct htp_tensor * restrict ids  = octx->src[n_weights + 1];

-    hvx_mm_run_quant_task(mmctx, ith);

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

@@ -1442,7 +1481,6 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
     const struct htp_tensor * restrict act  = octx->src[n_weights];
     const struct htp_tensor * restrict ids  = octx->src[n_weights + 1];

-    hvx_mm_run_quant_task(mmctx, ith);

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

@@ -1582,7 +1620,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

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

     uint32_t src0_row_start = 0;
     uint32_t src0_row_end   = src0_nrows;
@@ -1678,7 +1716,8 @@ 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;
-            mmctx->type            = "f16-f16";
+            need_quant             = (src1->type == HTP_TYPE_F32);
+            mmctx->type            = (src1->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;
@@ -1686,7 +1725,8 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
             break;

         case HTP_MM_KERNEL_HVX_F32_F32_VTCM:
-            quant_task_func        = quantize_f32_f32;
+            quant_task_func        = NULL;
+            need_quant             = false;
             mmctx->type            = "f32-f32";
             mmctx->vec_dot_1x1     = vec_dot_f32_f32_aa_1x1;
             mmctx->vec_dot_2x1     = vec_dot_f32_f32_aa_2x1;
@@ -1760,10 +1800,11 @@ 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_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src2 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2);
-    mmctx->vtcm_dst  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
+    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    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_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;
@@ -1771,46 +1812,68 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     mmctx->vtcm_src0_stride = src0_row_size_padded;
     mmctx->vtcm_src1_stride = src1_row_size;
+    if (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) {
+        mmctx->vtcm_act_raw_stride = hex_round_up(ne10 * sizeof(float), 128);
+    } else {
+        mmctx->vtcm_act_raw_stride = 0;
+    }

-    if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows) {
-        atomic_init(&mmctx->quant_barrier, 0);
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
+    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);
             mmctx->cur_m_start = m_start;
             mmctx->cur_m_rows  = cur_m_rows;

             if (need_quant) {
-                const uint32_t quant_tasks = MIN(cur_m_rows, octx->n_threads);
-                mmctx->n_quant_rows_per_thread = (cur_m_rows + quant_tasks - 1) / quant_tasks;
+                hvx_mm_transfer_src1_dma(octx, kparams, src1, 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;
+                const uint32_t total_nb = cur_m_rows * nb;
+                uint32_t quant_tasks;
+                work_queue_func_t q_func;
+                if (cur_m_rows < octx->n_threads && (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW)) {
+                    quant_tasks = MIN(total_nb, octx->n_threads);
+                    q_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
+                    for (uint32_t ith = 0; ith < quant_tasks; ++ith) {
+                        uint32_t ib_first = (total_nb * ith) / quant_tasks;
+                        uint32_t ib_last  = (total_nb * (ith + 1)) / quant_tasks;
+                        mmctx->quant_ib_first[ith] = ib_first;
+                        mmctx->quant_ib_last[ith]  = ib_last;
+                        mmctx->quant_r[ith]        = ib_first / nb;
+                        mmctx->quant_c[ith]        = ib_first % nb;
+                    }
+                } else {
+                    quant_tasks = MIN(cur_m_rows, octx->n_threads);
+                    q_func = quant_task_func;
+                    mmctx->n_quant_rows_per_thread = (cur_m_rows + quant_tasks - 1) / quant_tasks;
+                }
                 mmctx->n_quant_tasks = quant_tasks;
-                atomic_store(&mmctx->quant_barrier, quant_tasks);
-                mmctx->quant_task_func = quant_task_func;
+                work_queue_run(octx->ctx->work_queue, q_func, mmctx, quant_tasks);
             } else {
-                mmctx->quant_task_func = NULL;
-                mmctx->n_quant_tasks = 0;
+                hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, m_start, cur_m_rows);
             }

-            worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, octx->n_threads);
+            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;

         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;
-            mmctx->quant_task_func = quant_task_func;
             mmctx->n_quant_tasks = n_quant_tasks;
-            atomic_init(&mmctx->quant_barrier, n_quant_tasks);
+            work_queue_run(octx->ctx->work_queue, quant_task_func, mmctx, n_quant_tasks);
         } else {
-            mmctx->quant_task_func = NULL;
-            mmctx->n_quant_tasks = 0;
+            hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, 0, src1_nrows);
         }

-        htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
-
-        worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, octx->n_threads);
+        work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, octx->n_threads);
     }

     return HTP_STATUS_OK;
@@ -1835,7 +1898,6 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {

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

-    hvx_mm_run_quant_task(mmctx, ith);

     for (uint32_t widx = 0; widx < n_weights; widx++) {
         const struct htp_tensor * restrict src_w = octx->src[widx];
@@ -1990,8 +2052,7 @@ typedef struct {
     struct htp_context            * ctx;
     struct htp_thread_trace       * traces;
     __fp16                        * dst;
-    const float                   * src;
-    const struct mmid_row_mapping * matrix_rows;
+    dma_addr_t                      act_dma_addr;
     float                         * vtcm_f32_act;
     uint32_t                        n_tasks;
     uint32_t                        n_tot_chunks;
@@ -2008,7 +2069,7 @@ typedef struct {
     struct htp_context            * ctx;
     struct htp_thread_trace       * traces;
     __fp16                        * dst;
-    const float                   * src;
+    dma_addr_t                      act_dma_addr;
     float                         * vtcm_f32_act;
     uint32_t                        n_rows;
     uint32_t                        k_block;
@@ -2024,7 +2085,7 @@ typedef struct {
 static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
         dma_queue *dma_q,
         __fp16 *restrict vtcm_dst,
-        const float *restrict src,
+        dma_addr_t act_dma_addr,
         uint32_t n_rows,
         uint32_t k_block,
         uint32_t k_stride,
@@ -2044,7 +2105,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
     // Push step 0
     if (n_steps > 0 && n_rows > 0) {
         uint32_t nrows_to_fetch = hex_smin(n_rows, R);
-        dma_queue_push(dma_q, dma_make_data(thread_f32_act, src + c_first),
+        dma_queue_push(dma_q, dma_make_data(thread_f32_act, act_dma_addr + (size_t) c_first * sizeof(float)),
                        c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
     }
     // Push step 1
@@ -2052,9 +2113,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
         uint32_t next_r = R * 1;
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
-            const float *next_src = src + next_r * k_stride + c_first;
             float *next_buf = thread_f32_act + 1 * R * c_len;
-            dma_queue_push(dma_q, dma_make_data(next_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(next_buf, act_dma_addr + ((size_t) next_r * k_stride + c_first) * sizeof(float)),
                            c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
         }
     }
@@ -2084,50 +2144,12 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
         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);
-            const float *next_src = src + next_r * k_stride + c_first;
-            dma_queue_push(dma_q, dma_make_data(curr_buf, next_src),
+            dma_queue_push(dma_q, dma_make_data(curr_buf, act_dma_addr + ((size_t) next_r * k_stride + c_first) * sizeof(float)),
                            c_len * sizeof(float), k_stride * sizeof(float), k_chunk_valid * sizeof(float), nrows_to_fetch);
         }
     }
 }

-static void transfer_activation_chunk_fp32_to_fp16_col_chunk(
-        __fp16 *restrict vtcm_dst,
-        const float *restrict src,
-        uint32_t n_rows,
-        uint32_t k_block,
-        uint32_t k_stride,
-        uint32_t c_first,
-        uint32_t c_len,
-        uint32_t k_chunk_valid) {
-    const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
-    const uint32_t n_rows_tiled  = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
-
-    uint32_t r = 0;
-
-    #pragma unroll(2)
-    for (r = 0; r < n_rows_tiled; r += 2) {
-        const float *ptr_in0 = src + (r + 0) * k_stride + c_first;
-        const float *ptr_in1 = src + (r + 1) * k_stride + c_first;
-
-        transfer_activation_row_pair_fp32_to_fp16_col_chunk(
-            vtcm_dst, ptr_in0, ptr_in1, r, k_block, c_first, c_len, k_chunk_valid, true, true
-        );
-    }
-
-    for (; r < n_rows_padded; r += 2) {
-        const bool row0_valid = r       < n_rows;
-        const bool row1_valid = (r + 1) < n_rows;
-
-        const float *ptr_in0 = row0_valid ? (src + (r + 0) * k_stride + c_first) : NULL;
-        const float *ptr_in1 = row1_valid ? (src + (r + 1) * k_stride + c_first) : NULL;
-
-        transfer_activation_row_pair_fp32_to_fp16_col_chunk(
-            vtcm_dst, ptr_in0, ptr_in1, r, k_block, c_first, c_len, k_chunk_valid, row0_valid, row1_valid
-        );
-    }
-}
-
 static void transfer_activation_chunk_col_chunk_worker_fn(unsigned int n, unsigned int i, void *data) {
     activation_transfer_col_chunk_state_t *st = (activation_transfer_col_chunk_state_t *) data;
     struct htp_thread_trace * tr = &st->traces[i];
@@ -2149,29 +2171,20 @@ static void transfer_activation_chunk_col_chunk_worker_fn(unsigned int n, unsign
     }

     __fp16 *dst = st->dst;
-    const float *src = st->src;

-    if (st->vtcm_f32_act) {
-        size_t thread_scratch_bytes = hex_align_down(fastdiv(st->vtcm_f32_act_bytes, &st->n_threads_div), 128);
-        float *thread_f32_act = (float *)((char *)st->vtcm_f32_act + i * thread_scratch_bytes);
+    size_t thread_scratch_bytes = hex_align_down(fastdiv(st->vtcm_f32_act_bytes, &st->n_threads_div), 128);
+    float *thread_f32_act = (float *)((char *)st->vtcm_f32_act + i * thread_scratch_bytes);

-        transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
-            st->ctx->dma[i], dst, src, st->n_rows, st->k_block, st->k_stride, k_chunk_valid,
-            c_first, c_len, thread_f32_act, tr, st->dma_step_rows, st->dma_step_rows_shift
-        );
-    } else {
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, c_first);
-        transfer_activation_chunk_fp32_to_fp16_col_chunk(
-            dst, src, st->n_rows, st->k_block, st->k_stride, c_first, c_len, k_chunk_valid
-        );
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, c_first);
-    }
+    transfer_activation_chunk_fp32_to_fp16_dma_pipelined_col_chunk(
+        st->ctx->dma[i], dst, st->act_dma_addr, st->n_rows, st->k_block, st->k_stride, k_chunk_valid,
+        c_first, c_len, thread_f32_act, tr, st->dma_step_rows, st->dma_step_rows_shift
+    );
 }

 static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
         dma_queue *dma_q,
         __fp16 *restrict vtcm_dst,
-        const float *restrict src,
+        dma_addr_t act_dma_addr,
         uint32_t n_rows,
         uint32_t k_block,
         uint32_t k_stride,
@@ -2189,7 +2202,7 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
     // Push step 0
     if (n_steps > 0 && n_rows > 0) {
         uint32_t nrows_to_fetch = hex_smin(n_rows, R);
-        dma_queue_push(dma_q, dma_make_data(thread_f32_act, src),
+        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);
     }
     // Push step 1 (if valid)
@@ -2197,9 +2210,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
         uint32_t next_r = R * 1;
         if (next_r < n_rows) {
             uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
-            const float *next_src = src + next_r * k_stride;
             float *next_buf = thread_f32_act + 1 * R * k_block;
-            dma_queue_push(dma_q, dma_make_data(next_buf, next_src),
+            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);
         }
     }
@@ -2227,8 +2239,7 @@ 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);
-            const float *next_src = src + next_r * k_stride;
-            dma_queue_push(dma_q, dma_make_data(curr_buf, next_src),
+            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);
         }
     }
@@ -2244,18 +2255,12 @@ 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 float *src = st->src + chunk_idx * st->k_stride;
+        const dma_addr_t act_dma_addr = st->act_dma_addr + (size_t) chunk_idx * st->k_stride * sizeof(float);

-        if (st->vtcm_f32_act) {
-            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, src, chunk_size, st->k_block, st->k_stride, st->k_valid, thread_f32_act, tr, st->dma_step_rows, st->dma_step_rows_shift
-            );
-        } else {
-            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, chunk_idx);
-            transfer_activation_chunk_fp32_to_fp16(dst, src, chunk_size, st->k_block, st->k_stride, st->k_valid);
-            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, chunk_idx);
-        }
+        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
+        );
     }
 }

@@ -2494,7 +2499,7 @@ static void transfer_output_chunk_threaded(struct htp_context *ctx, float *dst,
 struct activation_transfer_params {
     struct htp_context *          ctx;
     __fp16 *                      dst;
-    const float *                 src;
+    dma_addr_t                    act_dma_addr;
     int                           n_rows;
     int                           k_block;
     int                           k_stride;
@@ -2509,7 +2514,7 @@ struct activation_transfer_params {
 static void transfer_activation_chunk_threaded(const struct activation_transfer_params * params) {
     struct htp_context *          ctx                = params->ctx;
     __fp16 *                      dst                = params->dst;
-    const float *                 src                = params->src;
+    const dma_addr_t              act_dma_addr       = params->act_dma_addr;
     int                           n_rows             = params->n_rows;
     int                           k_block            = params->k_block;
     int                           k_stride           = params->k_stride;
@@ -2529,7 +2534,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
         // Calculate step rows parameters for column-chunked dma pipelining
         uint32_t dma_step_rows = 2;
         uint32_t dma_step_rows_shift = 1;
-        if (vtcm_f32_act && vtcm_f32_act_bytes > 0 && k_block > 0) {
+        if (vtcm_f32_act_bytes > 0) {
             size_t thread_scratch_bytes = hex_align_down(fastdiv(vtcm_f32_act_bytes, act_threads_div), 128);
             size_t thread_scratch_elements = thread_scratch_bytes / sizeof(float);
             size_t dma_step_rows_max = fastdiv(thread_scratch_elements / 2, k_div);
@@ -2541,7 +2546,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_

         activation_transfer_col_chunk_state_t col_state;
         col_state.dst = dst;
-        col_state.src = src;
+        col_state.act_dma_addr = act_dma_addr;
         col_state.n_rows = n_rows;
         col_state.k_block = k_block;
         col_state.k_stride = k_stride;
@@ -2569,7 +2574,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
     state.n_tot_chunks       = n_tot_chunks;
     state.n_chunks_per_task  = n_chunks_per_task;
     state.dst                = dst;
-    state.src                = src;
+    state.act_dma_addr       = act_dma_addr;
     state.k_block            = k_block;
     state.k_stride           = k_stride;
     state.k_valid            = k_valid;
@@ -2581,7 +2586,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_

     uint32_t dma_step_rows = 2;
     uint32_t dma_step_rows_shift = 1;
-    if (vtcm_f32_act && state.vtcm_f32_act_bytes_per_thread > 0 && k_block > 0) {
+    if (state.vtcm_f32_act_bytes_per_thread > 0) {
         size_t thread_scratch_elements = state.vtcm_f32_act_bytes_per_thread / sizeof(float);
         size_t dma_step_rows_max = fastdiv(thread_scratch_elements / 2, k_div);
         if (dma_step_rows_max >= 4) {
@@ -2639,7 +2644,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
                                   float *restrict dst,
                                   dma_addr_t src2_addr,
                                   size_t src2_bytes,
-                                  const float *activation,
+                                  dma_addr_t act_dma_addr,
                                   dma_addr_t weight,
                                   int m, int k, int n,
                                   int act_stride,
@@ -2663,7 +2668,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

     if (k % 32 != 0 || n % 32 != 0) { return -1; }
-    if (!hex_is_aligned(dst, VLEN) || !hex_is_aligned(activation, VLEN)) { 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);
     if (row_stride == 0) {
@@ -2703,7 +2708,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
     const size_t qweight_row_stride = is_quant ? (size_t)(n_k_tiles * aligned_tile_size) / 32 : 0;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size, src2_bytes);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, pipeline, act_threads, aligned_tile_size, src2_bytes);

     vtcm_used = L.total_bytes;
     if (vtcm_used > vtcm_budget) {
@@ -2755,7 +2760,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .src = activation + mr * act_stride,
+                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -2846,7 +2851,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
             struct activation_transfer_params act_params = {
                 .ctx = ctx,
                 .dst = vtcm_f16_act,
-                .src = activation + mr * act_stride,
+                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -2925,10 +2930,10 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     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 float * activation = (const float *) act->data;
+    const dma_addr_t act_dma_addr = act->data;

     if (k % 32 != 0) { return HTP_STATUS_NO_SUPPORT; }
-    if (!hex_is_aligned(activation, VLEN)) { return HTP_STATUS_NO_SUPPORT; }
+    if ((act_dma_addr & (VLEN - 1)) != 0) { return HTP_STATUS_NO_SUPPORT; }

     size_t row_stride = htp_mm_get_tiled_row_stride(weight_type, k);
     if (row_stride == 0) {
@@ -2970,7 +2975,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     const uint32_t dma_width_bytes = is_quant ? tile_size : row_stride;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size, 0);
+    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, pipeline, act_threads, aligned_tile_size, 0);

     if (L.total_bytes > vtcm_budget) {
         FARF(ERROR, "hmx-mm-nx-2d: VTCM overflow: used %zu budget %zu, m %d k %d mc %d nc %d",
@@ -3026,7 +3031,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,
-                .src = activation + mr * act_stride,
+                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -3124,7 +3129,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,
-                .src = activation + mr * act_stride,
+                .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
                 .n_rows = (int) n_rows,
                 .k_block = k,
                 .k_stride = act_stride,
@@ -3202,11 +3207,10 @@ static inline dma_addr_t hmx_mm_weight_batch_data(const hmx_mm_f16_f32_batched_p
     return params->weight + b2_idx * params->src0_nb2 + b3_idx * params->src0_nb3;
 }

-static inline const float *hmx_mm_activation_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params,
-                                                           int dst_b2, int dst_b3) {
-    return (const float *) ((const uint8_t *) params->activation +
-                            (size_t) dst_b2 * params->src1_nb2 +
-                            (size_t) dst_b3 * params->src1_nb3);
+static inline dma_addr_t hmx_mm_act_batch_addr(const hmx_mm_f16_f32_batched_params_t *params,
+                                                int dst_b2, int dst_b3) {
+    return params->act_dma_addr + dst_b2 * params->act_nb2 +
+           dst_b3 * params->act_nb3;
 }

 static inline float *hmx_mm_dst_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params,
@@ -3224,11 +3228,11 @@ static int hmx_mm_f16_f32_batched_simple(struct htp_context *ctx,
     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 +
-                                       (dma_addr_t) b2 * params->src2_nb2 +
-                                       (dma_addr_t) b3 * params->src2_nb3) : 0;
+                                       b2 * params->src2_nb2 +
+                                       b3 * params->src2_nb3) : 0;
             ret = hmx_mm_2d_f32(ctx, params->weight_dma, hmx_mm_dst_batch_ptr(params, b2, b3),
                                 cur_src2_addr, params->src2_bytes,
-                                hmx_mm_activation_batch_ptr(params, b2, b3),
+                                hmx_mm_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),
@@ -3249,7 +3253,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
     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; }
-    if (!hex_is_aligned(params->dst, VLEN) || !hex_is_aligned(params->activation, VLEN)) { return -1; }
+    if (!hex_is_aligned(params->dst, VLEN) || (params->act_dma_addr & (VLEN - 1)) != 0) { return -1; }

     const int group_size = params->r2;
     const size_t vtcm_budget  = ctx->vtcm_size;
@@ -3267,16 +3271,12 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_

     const size_t vec_dot_size = params->k * sizeof(__fp16);

-    const bool use_dma_activation = (params->act_stride > params->k);
-    const size_t f32_scratch_size = use_dma_activation
-        ? hex_align_up((size_t)act_threads * HTP_MM_DMA_ACT_MULTIPLIER * (size_t) params->k * sizeof(float), HTP_MM_HMX_TILE_SIZE) : 0;
-
     size_t m_chunk_n_rows = m_chunk;
     size_t n_chunk_n_cols = n_chunk;
     size_t vtcm_used = vtcm_size;

     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, HTP_TYPE_F16, params->k, m_chunk_n_rows, n_chunk_n_cols, group_size, use_dma_activation, false, act_threads, 0, 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->src2_bytes);

     if (L.total_bytes > vtcm_budget) {
         FARF(HIGH, "%s: grouped layout overflowed VTCM, falling back to simple batched loop", __func__);
@@ -3291,7 +3291,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
     void    *vtcm_scratch0   = VTCM_LAYOUT_PTR(void, base, L.off_scratch[0]);
     void    *vtcm_scratch1   = VTCM_LAYOUT_PTR(void, base, L.off_scratch[1]);
     __fp16  *vtcm_scales     = VTCM_LAYOUT_PTR(__fp16, base, L.off_scales);
-    float   *vtcm_f32_act    = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_act_f32, use_dma_activation);
+    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);
@@ -3329,12 +3329,13 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                 // converts from the contiguous VTCM buffer.  This avoids L2 cache
                 // thrashing from HVX loads at large strides.
                 for (int g = 0; g < group_size; ++g) {
-                    const float *activation_chunk = hmx_mm_activation_batch_ptr(params, b2_base + g, b3) + mr * params->act_stride;
+                    const dma_addr_t act_dma_addr = hmx_mm_act_batch_addr(params, b2_base + g, b3) +
+                                                      mr * params->act_stride * sizeof(float);
                     __fp16 *vtcm_act_g = vtcm_f16_act + (size_t) g * L.act_head_stride;
                     struct activation_transfer_params act_params = {
                         .ctx = ctx,
                         .dst = vtcm_act_g,
-                        .src = activation_chunk,
+                        .act_dma_addr = act_dma_addr,
                         .n_rows = (int) n_rows,
                         .k_block = params->k,
                         .k_stride = params->act_stride,
@@ -3692,15 +3693,15 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
     size_t src2_nb3 = 0;
     if (src2) {
         src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
-        src2_addr = src2->data + (dma_addr_t) m_start * src2_stride * sizeof(float);
+        src2_addr = src2->data + m_start * src2_stride * sizeof(float);
         src2_bytes = (size_t) kparams->vtcm_src2_size;
         src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
         src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
     }

     const int dst_stride = (int)(dst->nb[1] / sizeof(float));
-    float       * dst_ptr = (float *)       dst->data  + m_start * dst_stride;
-    const float * act_ptr = (const float *) src1->data + m_start * act_stride;
+    float       * dst_ptr = (float *) dst->data + m_start * dst_stride;
+    const dma_addr_t act_addr = src1->data + m_start * act_stride * sizeof(float);

     int ret = -1;
     const int n_threads = kparams->n_threads;
@@ -3709,7 +3710,7 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
             .dst             = dst_ptr,
             .src2_addr       = src2_addr,
             .src2_bytes      = src2_bytes,
-            .activation      = act_ptr,
+            .act_dma_addr    = act_addr,
             .weight          = src0->data,
             .weight_dma      = octx->ctx->dma[0],
             .m               = m_rows,
@@ -3725,8 +3726,8 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
             .ne13            = ne13,
             .src0_nb2        = src0->nb[2],
             .src0_nb3        = src0->nb[3],
-            .src1_nb2        = src1->nb[2],
-            .src1_nb3        = src1->nb[3],
+            .act_nb2         = src1->nb[2],
+            .act_nb3         = src1->nb[3],
             .dst_nb2         = dst->nb[2],
             .dst_nb3         = dst->nb[3],
             .src2_nb2        = src2_nb2,
@@ -3745,7 +3746,8 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
                                      kparams->vtcm_size);
     } else {
         ret = hmx_mm_2d_f32(
-            octx->ctx, octx->ctx->dma[0], dst_ptr, src2_addr, src2_bytes, act_ptr, src0->data,
+            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],
             dst_stride, src2_stride, (int)dst->ne[0],
             kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
@@ -3829,7 +3831,7 @@ static int hvx_mm_matmul_id(
 ) {
     htp_matmul_tensors_preamble;
     const uint32_t src0_row_size_padded = mmctx->src0_row_size_padded;
-    const uint32_t src1_nrows           = mmctx->src1_nrows;
+    const uint32_t act_nrows            = mmctx->act_nrows;

     struct htp_thread_trace * tr = &octx->ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
@@ -3840,11 +3842,11 @@ static int hvx_mm_matmul_id(

     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;

     work_queue_func_t quant_task_func;
     uint32_t n_quant_tasks = 1;
-    if (src1_nrows < octx->n_threads) {
+    if (act_nrows < octx->n_threads) {
         n_quant_tasks = MIN(total_nb, octx->n_threads);
         quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
         for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -3856,13 +3858,13 @@ static int hvx_mm_matmul_id(
             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 = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled;
     }
     size_t src1_row_size  = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? 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, src1_nrows, octx->n_threads,
+    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);

     const size_t vtcm_size = L.total_bytes;
@@ -3882,18 +3884,20 @@ 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_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
-    mmctx->vtcm_src2 = NULL;
-    mmctx->vtcm_dst  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
+    mmctx->vtcm_src1     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+    mmctx->vtcm_src0     = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
+    mmctx->vtcm_src2     = 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_src0_stride    = src0_row_size_padded;
+    mmctx->vtcm_src1_stride    = src1_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;
@@ -3901,16 +3905,17 @@ static int hvx_mm_matmul_id(
     mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->cur_m_start = 0;
-    mmctx->cur_m_rows  = src1_nrows;
-
-    mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
-    mmctx->quant_task_func = quant_task_func;
-    mmctx->n_quant_tasks = n_quant_tasks;
-    atomic_init(&mmctx->quant_barrier, n_quant_tasks);
+    mmctx->cur_m_rows  = act_nrows;

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    worker_pool_run_func(octx->ctx->worker_pool, hvx_mmid_task_func, mmctx, octx->n_threads);
+    hvx_mm_transfer_src1_dma(octx, kparams, src1, 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);
+
+    work_queue_run(octx->ctx->work_queue, hvx_mmid_task_func, mmctx, octx->n_threads);

     return HTP_STATUS_OK;
 }
@@ -3976,7 +3981,7 @@ static int hvx_mm_matmul_id_nx(
     work_queue_func_t hvx_mmid_task_func
 ) {
     const uint32_t src0_row_size_padded = mmctx->src0_row_size_padded;
-    const uint32_t src1_nrows           = mmctx->src1_nrows;
+    const uint32_t act_nrows            = mmctx->act_nrows;

     struct htp_thread_trace * tr = &octx->ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
@@ -3990,11 +3995,11 @@ static int hvx_mm_matmul_id_nx(

     const uint32_t qk = QK_Q8_0_TILED;
     const uint32_t nb = (act->ne[0] + qk - 1) / qk;
-    const uint32_t total_nb = src1_nrows * nb;
+    const uint32_t total_nb = act_nrows * nb;

     work_queue_func_t quant_task_func;
     uint32_t n_quant_tasks = 1;
-    if (src1_nrows < octx->n_threads) {
+    if (act_nrows < octx->n_threads) {
         n_quant_tasks = MIN(total_nb, octx->n_threads);
         quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
         for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -4006,13 +4011,13 @@ static int hvx_mm_matmul_id_nx(
             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 = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled;
     }
     size_t src1_row_size = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]);

     struct htp_mm_hvx_vtcm_layout L;
-    htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
+    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);

     const size_t vtcm_size = L.total_bytes;
@@ -4024,9 +4029,10 @@ 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_dst  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
+    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_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;
@@ -4034,29 +4040,31 @@ static int hvx_mm_matmul_id_nx(
     octx->src3_spad.src = NULL;
     octx->dst_spad.src  = NULL;

-    mmctx->vtcm_src0_stride = 0;
-    mmctx->vtcm_src1_stride = src1_row_size;
+    mmctx->vtcm_src0_stride    = 0;
+    mmctx->vtcm_src1_stride    = src1_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_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->cur_m_start = 0;
-    mmctx->cur_m_rows  = src1_nrows;
-
-    mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
-    mmctx->quant_task_func         = quant_task_func;
-    mmctx->n_quant_tasks           = n_quant_tasks;
-    atomic_init(&mmctx->quant_barrier, n_quant_tasks);
+    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",
          src0->ne[0], src0->ne[1], src0->ne[2], mmctx->type, src0->ne[1],
-         act->ne[0], act->ne[1], act->ne[2], src1_nrows,
+         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);

-    worker_pool_run_func(octx->ctx->worker_pool, hvx_mmid_task_func, mmctx, octx->n_threads);
+    hvx_mm_transfer_src1_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);
+
+    work_queue_run(octx->ctx->work_queue, hvx_mmid_task_func, mmctx, octx->n_threads);

     return HTP_STATUS_OK;
 }
@@ -4202,7 +4210,7 @@ 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->src1_nrows           = src1_nrows;
+    mmctx->act_nrows            = src1_nrows;
     mmctx->cur_m_start          = 0;
     mmctx->cur_m_rows           = src1_nrows;

@@ -4281,7 +4289,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
     const size_t src0_row_size = src0->nb[1];
     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);

-    const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
+    const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3];

     const int n_ids = ids->ne[0];
     const int n_as  = src0->ne[2];
@@ -4291,7 +4299,7 @@ int op_matmul_id_nx(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);

@@ -4325,9 +4333,9 @@ int op_matmul_id_nx(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->src1_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);

@@ -4336,7 +4344,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
         s = hmx_mm_op_matmul_id_nx(octx, mmctx);
     } else {
         if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) {
-            s = hvx_mm_matmul_id_nx(octx, mmctx, src1_nrows > 1 ? hvx_mm_id_nx : hvx_mv_id_nx);
+            s = hvx_mm_matmul_id_nx(octx, mmctx, act_nrows > 1 ? hvx_mm_id_nx : hvx_mv_id_nx);
         } else {
             s = HTP_STATUS_NO_SUPPORT;
         }
@@ -4377,10 +4385,10 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     mmctx->octx = octx;
     mmctx->act  = act;

-    const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
-    mmctx->src1_nrows  = src1_nrows;
+    const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3];
+    mmctx->act_nrows   = act_nrows;
     mmctx->cur_m_start = 0;
-    mmctx->cur_m_rows  = src1_nrows;
+    mmctx->cur_m_rows  = act_nrows;

     const size_t src0_row_size = src0->nb[1];
     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
@@ -4391,11 +4399,11 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     const uint32_t qk = QK_Q8_0_TILED;
     const uint32_t nb = (act->ne[0] + qk - 1) / qk;
-    const uint32_t total_nb = src1_nrows * nb;
+    const uint32_t total_nb = act_nrows * nb;

     worker_callback_t quant_task_func;
     uint32_t n_quant_tasks = 1;
-    if (src1_nrows < octx->n_threads) {
+    if (act_nrows < octx->n_threads) {
         n_quant_tasks = MIN(total_nb, octx->n_threads);
         quant_task_func = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block;
         for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -4407,7 +4415,7 @@ int op_matmul_nx(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 = (src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q4_K) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled;
     }

@@ -4416,7 +4424,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {
                                : htp_mm_q8_0_tiled_row_size(act->ne[0]);

     struct htp_mm_hvx_vtcm_layout L;
-    htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
+    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);

     const size_t vtcm_size = L.total_bytes;
@@ -4428,9 +4436,10 @@ 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_dst  = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
+    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_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;
@@ -4438,18 +4447,14 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     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_src0_stride    = is_repacked ? 0 : src0_row_size_padded;
+    mmctx->vtcm_src1_stride    = src1_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_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

-    mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
-    mmctx->quant_task_func = quant_task_func;
-    mmctx->n_quant_tasks = n_quant_tasks;
-    atomic_init(&mmctx->quant_barrier, n_quant_tasks);
-
     // Run fused matmul
     const uint32_t n_matmul_jobs = octx->n_threads;
     worker_callback_t matmul_job_func;
@@ -4461,7 +4466,9 @@ int op_matmul_nx(struct htp_ops_context * octx) {
             case HTP_TYPE_Q8_0:   matmul_job_func = hvx_mm_nx_2d_repacked_q8_0;   break;
             case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl;  break;
             case HTP_TYPE_MXFP4:  matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4;  break;
-            default:              return HTP_STATUS_NO_SUPPORT;
+            default:
+                htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
+                return HTP_STATUS_NO_SUPPORT;
         }
     } else {
         matmul_job_func = hvx_mm_nx_2d;
@@ -4469,7 +4476,13 @@ int op_matmul_nx(struct htp_ops_context * octx) {

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

-    worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, n_matmul_jobs);
+    hvx_mm_transfer_src1_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);
+
+    work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, n_matmul_jobs);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h
index fe9dbb61c..49a76839a 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.h
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h
@@ -332,6 +332,7 @@ struct htp_mm_hvx_vtcm_layout {
     size_t off_src2;          // vtcm_src2 (Wq / fused only)
     size_t off_src3;          // vtcm_src3 (Wv / fused 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;
@@ -339,6 +340,7 @@ struct htp_mm_hvx_vtcm_layout {
     size_t src2_bytes;
     size_t src3_bytes;
     size_t dst_bytes;
+    size_t act_raw_bytes;

     size_t total_bytes;
 };
@@ -351,7 +353,6 @@ static inline void htp_mm_hmx_vtcm_layout_build(
     size_t mc,
     size_t nc,
     uint32_t group_size,
-    bool use_dma_activation,
     bool pipeline,
     uint32_t act_threads,
     uint32_t aligned_tile_size,
@@ -366,8 +367,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
         const size_t activation_area_size = hex_align_up(group_size * act_head_stride * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
         const size_t output_area_size  = hex_align_up(group_size * mc * nc * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
         const size_t scratch_area_size = hex_align_up(nc * vec_dot_size, HTP_MM_HMX_TILE_SIZE);
-        const size_t min_f32_size = use_dma_activation
-            ? hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128) : 0;
+        const size_t min_f32_size = hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128);

         // Group A: Permanent activation tiles and scales
         size_t off_group_a = 0;
@@ -388,10 +388,9 @@ static inline void htp_mm_hmx_vtcm_layout_build(

         // Group C: Activation prep temporary buffer (overlaps Group B, starting at off_group_a)
         const size_t max_f32_size = act_threads * 64 * k * sizeof(float);
-        const size_t act_f32_size = use_dma_activation
-            ? hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128) : 0;
+        const size_t act_f32_size = hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128);
         size_t off_group_c = off_group_a;
-        VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_f32, act_f32_size, use_dma_activation);
+        VTCM_LAYOUT_ALLOC(off_group_c, off_act_f32, act_f32_size);

         const size_t group_c_size = off_group_c - off_group_a;

@@ -478,11 +477,12 @@ static inline void htp_mm_hvx_vtcm_layout_build(
     bool is_fused_nx
 ) {
     (void)src1_row_size;
-    size_t src0_sz = 0;
-    size_t src1_sz = 0;
-    size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
-    size_t src3_sz = 0;
-    size_t dst_sz  = 0;
+    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 dst_sz     = 0;
+    size_t act_raw_sz = 0;

     const bool is_repack = (wtype == HTP_TYPE_Q4_0 || wtype == HTP_TYPE_Q4_1 ||
                             wtype == HTP_TYPE_Q8_0 || wtype == HTP_TYPE_IQ4_NL ||
@@ -491,7 +491,6 @@ static inline void htp_mm_hvx_vtcm_layout_build(

     if (is_fused_nx) {
         const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
-        const size_t quant_scratch_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;

         size_t weight_sz_per_thread = 0;

@@ -507,12 +506,14 @@ static inline void htp_mm_hvx_vtcm_layout_build(

         size_t tiled_act_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
         size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);
-
-        src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
-        src1_sz = act_sz;                           // quantized activation buffer
-        src2_sz = 0;
-        src3_sz = 0;
-        dst_sz  = quant_scratch_size;
+        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;
+        dst_sz     = 0;
+        act_raw_sz = hex_round_up(raw_row_size * src1_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 = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10)
@@ -529,10 +530,13 @@ static inline void htp_mm_hvx_vtcm_layout_build(
             src0_sz_per_thread               = repacked_vtcm_size;
         }

-        src0_sz = src0_sz_per_thread * n_threads;
-        dst_sz  = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
-        src2_sz = 0;
-        src3_sz = 0;
+        size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
+
+        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);
     } 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;
@@ -540,16 +544,18 @@ static inline void htp_mm_hvx_vtcm_layout_build(
         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);
-                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;
+                src1_sz    = htp_mm_round_up(f16_src1_row_size * src1_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);
                 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);
-                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;
+                src1_sz    = htp_mm_round_up(f32_src1_row_size * src1_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;
                 break;
             }
             case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
@@ -569,10 +575,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
                     src0_sz = repacked_vtcm_size * n_threads;
                 }

-                size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
                 size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
-                size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
-                dst_sz = dst_size_per_thread * n_threads;
+                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);
                 break;
             }
             default:
@@ -580,19 +586,30 @@ static inline void htp_mm_hvx_vtcm_layout_build(
         }
     }

-    size_t off = 0;
-    VTCM_LAYOUT_ALLOC(off, off_src0, src0_sz);
-    VTCM_LAYOUT_ALLOC(off, off_src1, src1_sz);
-    VTCM_LAYOUT_ALLOC(off, off_src2, src2_sz);
-    VTCM_LAYOUT_ALLOC(off, off_src3, src3_sz);
-    VTCM_LAYOUT_ALLOC(off, off_dst,  dst_sz);
-
-    L->src0_bytes = src0_sz;
-    L->src1_bytes = src1_sz;
-    L->src2_bytes = src2_sz;
-    L->src3_bytes = src3_sz;
-    L->dst_bytes  = dst_sz;
-    L->total_bytes = off;
+    // 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);
+
+    // Group B: Compute-only buffers (starts at off_group_a)
+    size_t off_group_b = off_group_a;
+    VTCM_LAYOUT_ALLOC(off_group_b, off_src0, src0_sz);
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_dst, dst_sz, dst_sz > 0);
+    const size_t group_b_size = off_group_b - off_group_a;
+
+    // Group C: Raw activation staging buffer (overlaps Group B, starts at off_group_a)
+    size_t off_group_c = off_group_a;
+    VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_raw, act_raw_sz, act_raw_sz > 0);
+    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->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);
 }

 static inline bool htp_mm_hvx_solve_vtcm_params(
@@ -642,7 +659,12 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
         return false;
     }

-    uint32_t m_chunk = (uint32_t) (avail_act / row_size);
+    size_t eff_row_size = row_size;
+    if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
+        eff_row_size += hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
+    }
+
+    uint32_t m_chunk = (uint32_t) (avail_act / eff_row_size);
     if (m_chunk > 1) {
         m_chunk &= ~1U;
     }
@@ -679,15 +701,15 @@ 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
 ) {
     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, false, pipeline, act_threads, aligned_tile_size, 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, src2_size);
     return L.total_bytes;
 }

 static inline size_t htp_mm_hmx_get_batched_vtcm_size(
-    int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool use_dma_activation, bool pipeline, uint32_t act_threads, 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 src2_size) {
     (void)pipeline;
     struct htp_mm_hmx_vtcm_layout L;
-    htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, use_dma_activation, false, act_threads, 0, 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, src2_size);
     return L.total_bytes;
 }

@@ -697,7 +719,6 @@ static inline bool htp_mm_hmx_solve_batched_params(
     uint32_t ne01_padded,
     uint32_t ne11,
     uint32_t group_size,
-    bool use_dma_activation,
     int n_threads,
     bool pipeline,
     size_t src2_size,
@@ -726,7 +747,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
         if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
                                (size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
                                &m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
-            size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads, src2_size);
+            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);
             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)) {