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)) {