Commit 8345f3339 for llama.cpp
commit 8345f333951c661d166b00e6f9362e553768f292
Author: Max Krasnyansky <maxk@qti.qualcomm.com>
Date: Mon Oct 5 08:55:21 2026 -0700
hexagon: matmul and flash-atten scalability updates (#29974)
* hexagon: head-parallel flash_attn partitioning for row-split multicore
In row-split mode each core computes its output row shard of every
MUL_MAT, but flash_attn was previously partitioning by Q tokens
(flat qrow split) instead of by heads. This forced every core to
read the full KV cache (all n_kv_heads), negating the memory
bandwidth benefit of multicore on flash_attn.
Change both HMX and HVX flash_attn kernels to partition by KV heads
when n_kv_heads is divisible by n_cores: core i processes heads
[i*n_kv_heads/N, (i+1)*n_kv_heads/N) exclusively, reading only its
head shard of the KV cache. Falls back to the original token-block
split when n_kv_heads % n_cores != 0 (e.g. Gemma-4 with 2 KV heads
on 4 cores).
Controlled by GGML_HEXAGON_FA_HEAD_SPLIT (default 1 = on).
The flag is packed into bit 1 of the existing is_dst_fp32 kparams
byte to stay within the 128-byte kernel_params blob limit.
Measured gains at 4c row-split (PP t/s, ubatch=1024):
Qwen3-0.6B: 6977 -> 11026 (+58%)
llama-3.2-3B: 3717 -> 5522 (+49%)
Qwen3.5-4B: 2739 -> 2855 (+4%)
Gemma-4 MoE: no change (MoE FFN dominates, fallback path)
TG is unchanged (flash_attn is a small fraction of decode time
relative to the matmul+barrier cost per layer).
* hex-fa: cleanup kern_params and head-split selection
* hex-fa: add -fa-head-split option to run.py
* hex-mdev: update matmul solver to account for reduced work in row-split scenarios
* hex-mmid: better work splitting by expers in multi-dev scenarios
* hex-fa: update HMX gating based on the model/n-hvx/ctx-len sweep
* hex-fa: precompute softcap/scale on the host
* hexagon: flatten matmul into 2d to use HMX in multi-sequence
* hex-mm: cleanup kparams and use collapse to 3/4D -> 2D mapping
* hex-mm: fix typo in collapse fallback
* hex-mm: another pass at consistent naming for act tensors
* hex-mm: add support for colapsing dims in fused matmuls
* hex-build: fix WoS build errors
* hex-mm: make sure to enforce dst stride in can_collapse
* hex-fa: add a onliner commit for head-split check
* hex-fa: remove unused local head_split var
* hex-fa: tighten up can_split checks
* hex-mm: update unfused paths to use act instead src1
* hex-mm: make sure to check all dsts for splitting
* hexagon: fix the second weight chunk address in the batched HMX matmul prologue
* hexagon: F16 activation and ragged N in the HMX matmul
* hex-mm: tighten the ragged/split checks in mdev cases
* hex-mm: enable MM fusion for F16 activations
* hex-mm: pass tiled sizes to the solver in fused paths
* hex-mmid: remove scalar divs from expert mapping loops
* hex-mmid: proper cacheline safety enforcement for mdev splits
* hex-mm: improve solver for mdev split scanarios and tail handling
* hex-mm: remove redundant checks
* hex-mm: fix fused HMX MUL_MAT_NX drops the final partial tile for quantized weights
* hex-mm: better handling of ragged shapes (removes scalar memset of vtcm)
---------
Co-authored-by: ebateni <ebateni@qti.qualcomm.com>
Co-authored-by: Jhen-Jie Hong <iainst0409@gmail.com>
Co-authored-by: Yiwei Shao <yiwei@aizip.ai>
diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 454594472..2282a04a8 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -1,3 +1,5 @@
+#define _USE_MATH_DEFINES
+
#include <assert.h>
#include <inttypes.h>
#include <stdio.h>
@@ -24,6 +26,10 @@
#include <cmath>
#include <initializer_list>
+#ifndef M_LOG2E
+# define M_LOG2E 1.44269504088896340736
+#endif
+
#ifdef _WIN32
# define WIN32_LEAN_AND_MEAN
# ifndef NOMINMAX
@@ -101,6 +107,7 @@ static bool opt_dma64 = false;
static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
+static int opt_fa_head_split = 1; // 1 = partition flash_attn by KV heads in multicore (default on), 0 = token-based (original)
static int opt_gdn_select = 2; // 2 = HMX -> HVX, 1 = HVX, 0 = CPU (unsupported)
static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (default), 1 = unfused ALLREDUCE, 0 = fallback to CPY+FENCE
static int opt_ar_scatter = 1; // 1 = reduce-scatter the fused ALLREDUCE+ADD (default), 0 = full reduction
@@ -387,7 +394,7 @@ static void ggml_hexagon_precompute_sort_params(
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
- const struct ggml_tensor * src1,
+ const struct ggml_tensor * act,
int32_t n_weights,
struct htp_mm_kernel_params * kparams
);
@@ -395,7 +402,7 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
static void ggml_hexagon_precompute_fused_mmidnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
- const struct ggml_tensor * src1,
+ const struct ggml_tensor * act,
const struct ggml_tensor * dst,
int32_t n_weights,
struct htp_mm_kernel_params * kparams
@@ -412,6 +419,10 @@ static bool ggml_hexagon_precompute_allreduce_params(
struct htp_allreduce_kernel_params * kparams
);
+static bool ggml_hexagon_rows_stride(const int64_t * ne, const size_t * nb, size_t * stride);
+static bool ggml_hexagon_matmul_can_collapse(const struct ggml_tensor * src0, const struct ggml_tensor * src1, const struct ggml_tensor * dst);
+static ggml_tensor ggml_hexagon_tensor_collapse_rows(const struct ggml_tensor * t);
+
static bool mm_is_hmx_eligible(const ggml_tensor * t);
static htp_op_code op_remap_to_htp(const ggml_tensor * t);
static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams);
@@ -444,6 +455,13 @@ static inline bool ggml_hexagon_tensors_overlap(const struct ggml_tensor * a, co
return a0 < b1 && b0 < a1;
}
+static inline bool ggml_hexagon_can_row_partition(const struct ggml_tensor * t) {
+ if (t->ne[1] > 1 && (t->nb[1] & 127) != 0) return false;
+ if (t->ne[2] > 1 && (t->nb[2] & 127) != 0) return false;
+ if (t->ne[3] > 1 && (t->nb[3] & 127) != 0) return false;
+ return true;
+}
+
struct htp_opnode;
struct ggml_hexagon_opbatch;
@@ -3410,6 +3428,10 @@ struct ggml_hexagon_opbatch {
return false;
}
+ if (orig_kparams->collapse != kparams.collapse) {
+ return false;
+ }
+
const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3];
const bool can_fuse = (kparams.n_hmx > 0) || (src1_nrows == 1);
if (!can_fuse) return false;
@@ -3493,8 +3515,20 @@ struct ggml_hexagon_opbatch {
return false;
}
+ const struct htp_mm_kernel_params * orig_kparams = (const struct htp_mm_kernel_params *) last_node.kernel_params;
+ const bool collapse = orig_kparams->collapse && ggml_hexagon_matmul_can_collapse(w_in, x, d_in);
+ if (orig_kparams->collapse && !collapse) {
+ return false;
+ }
+
struct htp_mm_kernel_params kparams;
- ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams);
+ if (collapse) {
+ const ggml_tensor x_collapsed = ggml_hexagon_tensor_collapse_rows(x);
+ ggml_hexagon_precompute_fused_mmnx_params(sess, w0, &x_collapsed, curr_n + 1, &kparams);
+ kparams.collapse = 1;
+ } else {
+ ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams);
+ }
if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
return false;
}
@@ -3543,9 +3577,23 @@ struct ggml_hexagon_opbatch {
const ggml_tensor * w0 = last_node.src0();
const ggml_tensor * x = last_node.src1();
const ggml_tensor * w1 = node.src0();
+ const ggml_tensor * dst_0 = last_node.dst();
+ const ggml_tensor * dst_1 = node.dst();
+
+ const struct htp_mm_kernel_params * orig_kparams = (const struct htp_mm_kernel_params *) last_node.kernel_params;
+ const bool collapse = orig_kparams->collapse && ggml_hexagon_matmul_can_collapse(w1, x, dst_1);
+ if (orig_kparams->collapse && !collapse) {
+ return false;
+ }
struct htp_mm_kernel_params kparams;
- ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
+ if (collapse) {
+ const ggml_tensor x_collapsed = ggml_hexagon_tensor_collapse_rows(x);
+ ggml_hexagon_precompute_fused_mmnx_params(sess, w0, &x_collapsed, 2, &kparams);
+ kparams.collapse = 1;
+ } else {
+ ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
+ }
if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
return false;
}
@@ -3559,9 +3607,6 @@ struct ggml_hexagon_opbatch {
return false;
}
- const ggml_tensor * dst_0 = last_node.dst();
- const ggml_tensor * dst_1 = node.dst();
-
last_node.opcode = HTP_OP_MUL_MAT_NX;
last_node.name = "MUL_MAT_NX";
last_node.inputs.clear();
@@ -4831,15 +4876,48 @@ static bool ggml_hexagon_flash_attn_is_hmx_eligible(
return false;
}
- // Fall back to HVX for small token counts if head dimension is small (DK <= 128)
- const uint32_t neq1 = q->ne[1];
- if (DK <= 128 && neq1 < 5) {
- return false;
+ GGML_UNUSED(sinks);
+
+ // Explicit force mode
+ if (opt_fa_select > 2) {
+ return true;
}
- return true;
+ const uint32_t M = q->ne[1];
- GGML_UNUSED(sinks);
+ // Prefill or batched decode
+ if (M > 1) {
+ return true;
+ }
+
+ // Compute-bound head dim
+ if (DK >= 256) {
+ return true;
+ }
+
+ const uint32_t n_head = q->ne[2];
+ const uint32_t n_kv_heads = k->ne[2];
+ const uint32_t G = n_kv_heads > 0 ? n_head / n_kv_heads : 1;
+ const uint32_t S = k->ne[1];
+
+ // Tile alignment for 32-row HMX tiles
+ const bool is_tile_aligned = (G > 0 && (32 % G == 0));
+ if (!is_tile_aligned) {
+ if (sess->n_threads >= 6) {
+ return false;
+ }
+ return S >= 1024;
+ }
+
+ // Context depth crossover
+ uint32_t s_cross = 512;
+ if (DK <= 64) {
+ s_cross = (sess->n_threads >= 8) ? 2048 : ((sess->n_threads >= 6) ? 768 : 512);
+ } else {
+ s_cross = (sess->n_threads >= 8) ? 1024 : 512;
+ }
+
+ return S >= s_cross;
}
static bool ggml_hexagon_precompute_flash_attn_params(
@@ -4857,7 +4935,6 @@ static bool ggml_hexagon_precompute_flash_attn_params(
const struct ggml_tensor * k = op->src[1];
const struct ggml_tensor * v = op->src[2];
const struct ggml_tensor * mask = op->src[3];
- const struct ggml_tensor * dst = op;
const uint32_t neq0 = q->ne[0]; // head_dim (DK)
const uint32_t neq1 = q->ne[1]; // n_tokens
@@ -4888,8 +4965,8 @@ static bool ggml_hexagon_precompute_flash_attn_params(
kparams->max_bias = max_bias;
kparams->logit_softcap = logit_softcap;
- kparams->is_q_fp32 = (q->type == GGML_TYPE_F32) ? 1 : 0;
- kparams->is_dst_fp32 = (dst->type == GGML_TYPE_F32) ? 1 : 0;
+ kparams->head_split = (opt_fa_head_split != 0) ? 1 : 0;
+ kparams->flags = 0;
kparams->G = G;
const uint32_t n_head = q->ne[2];
@@ -4906,9 +4983,16 @@ static bool ggml_hexagon_precompute_flash_attn_params(
const uint32_t DK_pad = hex_round_up(DK, 64);
const uint32_t DV_pad = hex_round_up(DV, 64);
size_t Br = 0, Bc = 0;
- int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0, sinks != nullptr, n_head);
+ int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, (q->type == GGML_TYPE_F32), sinks != nullptr, n_head);
if (ret == 0) {
kparams->kernel_type = HTP_FA_KERNEL_HMX;
+ if (logit_softcap == 0.0f) {
+ kparams->scale = scale * (float) M_LOG2E;
+ kparams->logit_softcap = 0.0f;
+ } else {
+ kparams->scale = scale;
+ kparams->logit_softcap = logit_softcap * (float) M_LOG2E;
+ }
kparams->Br = Br;
kparams->Bc = Bc;
kparams->n_kv_blocks = (nek1 + Bc - 1) / Bc;
@@ -4916,7 +5000,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(
kparams->u.hmx.g_br = hex_align_up(G * Br, 32);
kparams->u.hmx.pipeline = (kparams->n_kv_blocks >= 3 && sess->n_threads >= 2) ? 1 : 0;
- kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0, sinks != nullptr, n_head);
+ kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, (q->type == GGML_TYPE_F32), sinks != nullptr, n_head);
const size_t row_vec_bytes = hex_align_up(Bc * sizeof(uint16_t), 256);
kparams->u.hmx.row_buf_stride = row_vec_bytes / 128; // HVX vector is 128 bytes
@@ -4943,11 +5027,11 @@ static bool ggml_hexagon_precompute_flash_attn_params(
kparams->n_kv_blocks = (k->ne[1] + 64 - 1) / 64;
kparams->n_threads = sess->n_threads;
- const size_t size_q_row_padded = hex_round_up(q->ne[0] * (kparams->is_q_fp32 ? 4 : 2), 128);
+ const size_t size_q_row_padded = hex_round_up(q->ne[0] * ((q->type == GGML_TYPE_F32) ? 4 : 2), 128);
const size_t size_k_row_padded = hex_round_up(k->ne[0] * 2, 128);
const size_t size_v_row_padded = hex_round_up(v->ne[0] * 2, 128);
- kparams->vtcm_size = hvx_fa_compute_vtcm_usage(DK, DV, kparams->is_q_fp32 != 0, mask != nullptr, sinks != nullptr, n_head, sess->n_threads);
+ kparams->vtcm_size = hvx_fa_compute_vtcm_usage(DK, DV, (q->type == GGML_TYPE_F32), mask != nullptr, sinks != nullptr, n_head, sess->n_threads);
kparams->u.hvx.size_q_row_padded = size_q_row_padded;
kparams->u.hvx.size_k_row_padded = size_k_row_padded;
@@ -5102,7 +5186,7 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
bool is_matmul_id,
bool is_batched
) {
- if (src1->type != GGML_TYPE_F32) {
+ if (src1->type != GGML_TYPE_F32 && (src1->type != GGML_TYPE_F16 || is_matmul_id)) {
return false;
}
@@ -5111,8 +5195,8 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
const int ne12 = src1->ne[2];
const int wtype = src0->type;
- // HMX weight tile requires N to be 32-aligned.
- if (ne01_padded % 32 != 0) {
+ // HMX weight tiles accept non-32-aligned N for non-matmul_id.
+ if (ne01_padded % 32 != 0 && is_matmul_id) {
return false;
}
@@ -5151,7 +5235,7 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
static bool ggml_hexagon_precompute_hmx_mm_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
- const struct ggml_tensor * src1,
+ const struct ggml_tensor * act,
const struct ggml_tensor * dst,
int wtype,
int ne00_padded,
@@ -5167,9 +5251,22 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
struct htp_mm_kernel_params * kparams
) {
const int aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype);
- const bool pipeline = is_matmul_id ? false : htp_mm_hmx_pipeline(ne11);
const int n_threads = (int)sess->n_threads;
- const int ne10 = src1->ne[0];
+ const int ne10 = act->ne[0];
+
+ int m_for_solver = ne11;
+ int m_for_solver_padded = ne11_padded;
+ // matmul_id partitions by expert; regular matmul partitions M rows (ne11) across devices
+ if (!is_matmul_id && sess->mdev.count > 1 && ((uint32_t) ne11 >= sess->mdev.count)) {
+ // when dst is null, padded dims are used for estimate which are 128-byte aligned
+ const bool dst_row_split = dst ? ggml_hexagon_can_row_partition(dst) : true;
+ const bool act_row_split = ggml_hexagon_can_row_partition(act);
+ if (dst_row_split && act_row_split) {
+ m_for_solver = (ne11 + (int) sess->mdev.count - 1) / (int) sess->mdev.count;
+ m_for_solver_padded = hex_round_up(std::max(m_for_solver, 32), 32);
+ }
+ }
+ const bool pipeline = is_matmul_id ? false : htp_mm_hmx_pipeline(m_for_solver);
const bool is_batched_val = is_matmul_id ? false : is_batched;
const int group_size = (ne02 > 0 ? ne12 / ne02 : 1);
@@ -5182,15 +5279,25 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
// Try grouped path first
- if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+ if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, m_for_solver, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
use_grouped = true;
}
}
if (!use_grouped) {
// Fallback to simple 2D path (group_size = 1)
- const int m_id_rows = (dst && is_matmul_id) ? (int) ((size_t) dst->ne[1] * dst->ne[2]) : 0;
- if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, ne11_padded, ne11, n_threads, pipeline, is_matmul_id, aligned_tile_size, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
+ int m_id_rows = 0;
+ if (dst && is_matmul_id) {
+ const int n_experts = ne02 > 0 ? ne02 : 1;
+ const size_t total_expert_rows = (size_t) dst->ne[1] * dst->ne[2];
+ int m_per_expert = (int) ((total_expert_rows + n_experts - 1) / n_experts);
+ if (sess->mdev.count > 1 && ggml_hexagon_can_row_partition(dst)) {
+ m_per_expert = (m_per_expert + (int) sess->mdev.count - 1) / (int) sess->mdev.count;
+ }
+ m_id_rows = hex_round_up(std::max(m_per_expert, 32), 32);
+ }
+ const uint32_t cost_m = is_matmul_id ? (uint32_t) m_id_rows : (uint32_t) m_for_solver;
+ if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, m_for_solver_padded, cost_m, n_threads, pipeline, is_matmul_id, aligned_tile_size, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
return false;
}
}
@@ -5203,14 +5310,14 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
kparams->n_act_threads = act_threads_selected;
kparams->tile_size = htp_mm_get_weight_tile_size(wtype);
kparams->aligned_tile_size = aligned_tile_size;
- kparams->src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ kparams->act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
kparams->vtcm_size = vtcm_size;
kparams->vtcm_src0_size = 0;
kparams->div_n_act_threads = init_fastdiv_values(act_threads_selected);
kparams->div_ne00_padded = init_fastdiv_values(ne00_padded);
- kparams->vtcm_src1_size = 0;
- kparams->vtcm_src2_size = (int32_t) src2_size;
- kparams->vtcm_dst_size = 0;
+ kparams->vtcm_act_size = 0;
+ kparams->vtcm_bias_size = (int32_t) src2_size;
+ kparams->vtcm_dst_size = 0;
if (is_batched && !is_matmul_id) {
kparams->kernel_type = HTP_MM_KERNEL_HMX_F16_BATCHED;
@@ -5247,6 +5354,9 @@ static void ggml_hexagon_precompute_hvx_mm_params(
kparams->n_hmx = 0;
kparams->n_threads = sess->n_threads;
+ GGML_UNUSED(ne02);
+ GGML_UNUSED(ne03);
+
const bool is_quant = (wtype != GGML_TYPE_F16 && wtype != GGML_TYPE_F32);
const int src1_nrows = ne11 * ne12 * ne13;
@@ -5259,7 +5369,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
if (is_matmul_id) {
kparams->kernel_type = (src1_nrows < (int) sess->n_threads) ? HTP_MM_KERNEL_HVX_QUANT_BLOCK : HTP_MM_KERNEL_HVX_QUANT_ROW;
- kparams->src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ kparams->act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
struct htp_mm_hvx_vtcm_layout L;
uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
@@ -5267,7 +5377,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
for (uint32_t d = max_prefetch; d >= 2; d /= 2) {
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
- 0, src0->nb[1], kparams->src1_row_size, 0, d, true, false
+ 0, src0->nb[1], kparams->act_row_size, 0, d, true, false
);
if (L.total_bytes <= vtcm_budget) {
best_n_prefetch = d;
@@ -5281,13 +5391,14 @@ static void ggml_hexagon_precompute_hvx_mm_params(
kparams->n_prefetch = best_n_prefetch;
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
- kparams->vtcm_src1_size = L.src1_bytes;
+ kparams->vtcm_act_size = L.act_bytes;
+ kparams->vtcm_bias_size = 0;
kparams->vtcm_dst_size = L.dst_bytes;
goto done_quant;
} else {
bool try_tiled = (k_align && opt_mm_select >= 1);
if (try_tiled) {
- kparams->src1_row_size = htp_mm_weight_has_offset(wtype)
+ kparams->act_row_size = htp_mm_weight_has_offset(wtype)
? htp_mm_q8_1_tiled_row_size(ne10)
: htp_mm_q8_0_tiled_row_size(ne10);
if (src1_nrows < (int) sess->n_threads) {
@@ -5319,8 +5430,8 @@ static void ggml_hexagon_precompute_hvx_mm_params(
kparams->m_chunk = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
- kparams->vtcm_src1_size = L.src1_bytes;
- kparams->vtcm_src2_size = L.src2_bytes;
+ kparams->vtcm_act_size = L.act_bytes;
+ kparams->vtcm_bias_size = L.bias_bytes;
kparams->vtcm_dst_size = L.dst_bytes;
goto done_quant;
}
@@ -5341,11 +5452,11 @@ static void ggml_hexagon_precompute_hvx_mm_params(
&L, &m_chunk)) {
kparams->kernel_type = HTP_MM_KERNEL_HVX_F16_F16_VTCM;
kparams->m_chunk = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
- kparams->src1_row_size = hex_round_up(ne10 * 2, 128);
+ kparams->act_row_size = hex_round_up(ne10 * 2, 128);
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
- kparams->vtcm_src1_size = L.src1_bytes;
- kparams->vtcm_src2_size = L.src2_bytes;
+ kparams->vtcm_act_size = L.act_bytes;
+ kparams->vtcm_bias_size = L.bias_bytes;
kparams->vtcm_dst_size = L.dst_bytes;
kparams->n_prefetch = 16;
return;
@@ -5363,11 +5474,11 @@ static void ggml_hexagon_precompute_hvx_mm_params(
&L, &m_chunk)) {
kparams->kernel_type = HTP_MM_KERNEL_HVX_F32_F32_VTCM;
kparams->m_chunk = (m_chunk < (uint32_t) src1_nrows) ? m_chunk : 0;
- kparams->src1_row_size = hex_round_up(ne10 * 4, 128);
+ kparams->act_row_size = hex_round_up(ne10 * 4, 128);
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
- kparams->vtcm_src1_size = L.src1_bytes;
- kparams->vtcm_src2_size = L.src2_bytes;
+ kparams->vtcm_act_size = L.act_bytes;
+ kparams->vtcm_bias_size = L.bias_bytes;
kparams->vtcm_dst_size = L.dst_bytes;
kparams->n_prefetch = 16;
return;
@@ -5404,6 +5515,9 @@ static void ggml_hexagon_precompute_matmul_params_impl(
const int ne00_padded = is_repack ? hex_round_up(ne00, 32) : ne00;
const int ne01_padded = is_repack ? hex_round_up(ne01, 32) : ne01;
const int ne11_padded = hex_round_up(ne11, 32);
+ // VTCM has to hold whole 32-row weight tiles, so size for the rounded-up N
+ // even when the tensor itself is ragged.
+ const int ne01_tiled = hex_round_up(ne01_padded, 32);
const bool is_matmul_id = (dst->op == GGML_OP_MUL_MAT_ID);
const bool is_batched = (ne02 * ne03 > 1 || ne12 * ne13 > 1);
@@ -5413,7 +5527,7 @@ static void ggml_hexagon_precompute_matmul_params_impl(
// Check HMX eligibility and try precomputing HMX parameters
bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 2);
if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, src1, dst, ne01_padded, is_matmul_id, is_batched)) {
- if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, dst, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, is_matmul_id, is_batched, src2_size, vtcm_budget, kparams)) {
+ if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, dst, wtype, ne00_padded, ne01_tiled, ne02, ne11, ne12, ne11_padded, is_matmul_id, is_batched, src2_size, vtcm_budget, kparams)) {
goto finalize;
}
}
@@ -5429,6 +5543,66 @@ finalize:
kparams->div_ne12 = init_fastdiv_values(ne12);
}
+// The rows of dims 1..3 can be walked with one stride (size-1 dims skipped); returns that stride
+static bool ggml_hexagon_rows_stride(const int64_t * ne, const size_t * nb, size_t * stride) {
+ size_t s = 0, next = 0;
+ for (int i = 1; i < GGML_MAX_DIMS; i++) {
+ if (ne[i] == 1) continue;
+ if (s == 0) { s = nb[i]; next = s * ne[i]; continue; }
+ if (nb[i] != next) return false;
+ next *= ne[i];
+ }
+ *stride = s ? s : nb[1];
+ return true;
+}
+
+// A 2D weight applied to a batched activation whose rows are evenly strided is the same matmul over ne11 * ne12 * ne13 rows
+static bool ggml_hexagon_matmul_can_collapse(const struct ggml_tensor * src0, const struct ggml_tensor * src1, const struct ggml_tensor * dst) {
+ size_t s1, sd;
+ return (dst->op == GGML_OP_MUL_MAT || dst->op == GGML_OP_ADD) &&
+ src0->ne[2] == 1 && src0->ne[3] == 1 && src1->ne[2] * src1->ne[3] > 1 &&
+ src1->nb[0] == ggml_type_size(src1->type) && ggml_hexagon_rows_stride(src1->ne, src1->nb, &s1) &&
+ dst->nb[0] == ggml_type_size(dst->type) && ggml_hexagon_rows_stride(dst->ne, dst->nb, &sd);
+}
+
+static bool ggml_hexagon_matmul_add_can_collapse(
+ const struct ggml_tensor * src0,
+ const struct ggml_tensor * src1,
+ const struct ggml_tensor * src2,
+ const struct ggml_tensor * dst
+) {
+ if (!ggml_hexagon_matmul_can_collapse(src0, src1, dst)) {
+ return false;
+ }
+ if (!src2) {
+ return true;
+ }
+ const int64_t src2_nrows = src2->ne[1] * src2->ne[2] * src2->ne[3];
+ if (src2_nrows == 1) {
+ return true;
+ }
+ size_t s2;
+ return src2->nb[0] == ggml_type_size(src2->type) &&
+ src2->ne[0] == dst->ne[0] &&
+ src2->ne[1] == dst->ne[1] &&
+ src2->ne[2] == dst->ne[2] &&
+ src2->ne[3] == dst->ne[3] &&
+ ggml_hexagon_rows_stride(src2->ne, src2->nb, &s2);
+}
+
+static ggml_tensor ggml_hexagon_tensor_collapse_rows(const struct ggml_tensor * t) {
+ size_t stride = 0;
+ ggml_hexagon_rows_stride(t->ne, t->nb, &stride);
+ ggml_tensor c = *t;
+ c.ne[1] = t->ne[1] * t->ne[2] * t->ne[3];
+ c.ne[2] = 1;
+ c.ne[3] = 1;
+ c.nb[1] = stride;
+ c.nb[2] = c.nb[1] * c.ne[1];
+ c.nb[3] = c.nb[2];
+ return c;
+}
+
static void ggml_hexagon_precompute_matmul_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
@@ -5436,6 +5610,13 @@ static void ggml_hexagon_precompute_matmul_params(
const struct ggml_tensor * dst,
struct htp_mm_kernel_params * kparams
) {
+ if (ggml_hexagon_matmul_can_collapse(src0, src1, dst)) {
+ const ggml_tensor src1_collapsed = ggml_hexagon_tensor_collapse_rows(src1);
+ const ggml_tensor dst_collapsed = ggml_hexagon_tensor_collapse_rows(dst);
+ ggml_hexagon_precompute_matmul_params_impl(sess, src0, &src1_collapsed, &dst_collapsed, 0, 0, kparams);
+ kparams->collapse = 1;
+ return;
+ }
ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
}
@@ -5447,6 +5628,18 @@ static void ggml_hexagon_precompute_fused_matmul_add_params(
const struct ggml_tensor * dst,
struct htp_mm_kernel_params * kparams
) {
+ if (ggml_hexagon_matmul_add_can_collapse(src0, src1, src2, dst)) {
+ const ggml_tensor src1_collapsed = ggml_hexagon_tensor_collapse_rows(src1);
+ const ggml_tensor dst_collapsed = ggml_hexagon_tensor_collapse_rows(dst);
+ const ggml_tensor src2_collapsed = (src2 && (src2->ne[1] * src2->ne[2] * src2->ne[3] > 1))
+ ? ggml_hexagon_tensor_collapse_rows(src2)
+ : (src2 ? *src2 : ggml_tensor{});
+ const struct ggml_tensor * p_src2 = src2 ? &src2_collapsed : nullptr;
+ const size_t src2_size = p_src2 ? hex_round_up(ggml_nbytes(p_src2), 128) : 0;
+ ggml_hexagon_precompute_matmul_params_impl(sess, src0, &src1_collapsed, &dst_collapsed, p_src2 ? p_src2->nb[1] : 0, src2_size, kparams);
+ kparams->collapse = 1;
+ return;
+ }
const size_t src2_size = src2 ? hex_round_up(ggml_nbytes(src2), 128) : 0;
ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, src2 ? src2->nb[1] : 0, src2_size, kparams);
}
@@ -6066,7 +6259,7 @@ static void ggml_hexagon_precompute_sort_params(
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0, // W0
- const struct ggml_tensor * src1, // x
+ const struct ggml_tensor * act, // x
int32_t n_weights,
struct htp_mm_kernel_params * kparams
) {
@@ -6078,23 +6271,24 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
const int ne02 = src0->ne[2];
const int ne03 = src0->ne[3];
- const int ne10 = src1->ne[0];
- const int ne11 = src1->ne[1];
- const int ne12 = src1->ne[2];
- const int ne13 = src1->ne[3];
+ const int ne10 = act->ne[0];
+ const int ne11 = act->ne[1];
+ const int ne12 = act->ne[2];
+ const int ne13 = act->ne[3];
const int wtype = src0->type;
const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
const int ne00_padded = is_repack ? hex_round_up(ne00, 32) : ne00;
const int ne01_padded = is_repack ? hex_round_up(ne01, 32) : ne01;
const int ne11_padded = hex_round_up(ne11, 32);
+ const int ne01_tiled = hex_round_up(ne01_padded, 32);
const size_t vtcm_budget = sess->vtcm_size;
const bool is_batched = (ne02 * ne03 > 1 || ne12 * ne13 > 1);
bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 2);
- if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, src1, nullptr, ne01_padded, false, is_batched)) {
- if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, nullptr, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, false, is_batched, 0, vtcm_budget, kparams)) {
+ if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, act, nullptr, ne01_padded, false, is_batched)) {
+ if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, act, nullptr, wtype, ne00_padded, ne01_tiled, ne02, ne11, ne12, ne11_padded, false, is_batched, 0, vtcm_budget, kparams)) {
kparams->n_weights = n_weights;
goto finalize;
}
@@ -6106,20 +6300,20 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
}
{
- const int src1_nrows = ne11 * ne12 * ne13;
- const size_t src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ const int act_nrows = ne11 * ne12 * ne13;
+ const size_t act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
const size_t src0_row_size = src0->nb[1];
uint32_t best_n_prefetch = 16;
if (is_repack) {
- const uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
+ const uint32_t max_prefetch = (act_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16;
best_n_prefetch = 2;
for (uint32_t d = max_prefetch; d >= 2; d /= 2) {
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(
- &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads,
- 0, src0_row_size, src1_row_size, 0, d, false, true
+ &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, act_nrows, sess->n_threads,
+ 0, src0_row_size, act_row_size, 0, d, false, true
);
if (L.total_bytes <= sess->vtcm_size) {
best_n_prefetch = d;
@@ -6133,14 +6327,16 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
// Test tiled first
htp_mm_hvx_vtcm_layout_build(
- &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads,
- 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true
+ &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, act_nrows, sess->n_threads,
+ 0, src0_row_size, act_row_size, 0, best_n_prefetch, false, true
);
if (try_tiled && L.total_bytes <= sess->vtcm_size) {
- kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
+ kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
+ kparams->act_row_size = act_row_size;
kparams->vtcm_src0_size = L.src0_bytes;
- kparams->vtcm_src1_size = L.src1_bytes;
+ kparams->vtcm_act_size = L.act_bytes;
+ kparams->vtcm_bias_size = 0;
kparams->vtcm_dst_size = L.dst_bytes;
kparams->vtcm_size = L.total_bytes;
kparams->n_prefetch = best_n_prefetch;
@@ -6162,12 +6358,12 @@ finalize:
static void ggml_hexagon_precompute_fused_mmidnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0, // W0
- const struct ggml_tensor * src1, // x
+ const struct ggml_tensor * act, // x
const struct ggml_tensor * dst, // dst0
int32_t n_weights,
struct htp_mm_kernel_params * kparams
) {
- ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, 0, kparams);
+ ggml_hexagon_precompute_matmul_params_impl(sess, src0, act, dst, 0, 0, kparams);
kparams->n_weights = n_weights;
}
@@ -7137,6 +7333,14 @@ static bool mm_is_hmx_eligible(const ggml_tensor * t) {
const ggml_tensor * src0 = t->src[0];
const ggml_tensor * src1 = t->src[1];
+ if (ggml_hexagon_matmul_can_collapse(src0, src1, t)) {
+ const ggml_tensor src1_c = ggml_hexagon_tensor_collapse_rows(src1);
+ const int wtype = src0->type;
+ const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
+ const int ne01_padded = is_repack ? hex_round_up(src0->ne[1], 32) : src0->ne[1];
+ return ggml_hexagon_matmul_is_hmx_eligible(src0, &src1_c, t, ne01_padded, false, false);
+ }
+
const int wtype = src0->type;
const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
const bool is_matmul_id = (t->op == GGML_OP_MUL_MAT_ID);
@@ -7176,13 +7380,15 @@ static bool is_mergeable_mul_mat(const ggml_tensor * t) {
const ggml_tensor * src0 = t->src[0];
const ggml_tensor * src1 = t->src[1];
- if (src1->type != GGML_TYPE_F32) return false;
if (src0->ne[2] != 1 || src0->ne[3] != 1) return false;
if (mm_is_hmx_eligible(t)) {
return ggml_hexagon_is_hmx_weight_type(src0->type);
}
+ // HVX path requires F32 activations and repacked weights (except Q6_K)
+ if (src1->type != GGML_TYPE_F32) return false;
+
return ggml_hexagon_is_repack_type(src0->type) && src0->type != GGML_TYPE_Q6_K;
}
@@ -8683,6 +8889,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
const char * str_nhmx = getenv("GGML_HEXAGON_NHMX");
const char * str_mm_select = getenv("GGML_HEXAGON_MM_SELECT");
const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT");
+ const char * str_fa_head_split = getenv("GGML_HEXAGON_FA_HEAD_SPLIT");
const char * str_gdn_select = getenv("GGML_HEXAGON_GDN_SELECT");
const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT");
const char * str_ar_scatter = getenv("GGML_HEXAGON_AR_SCATTER");
@@ -8737,6 +8944,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
opt_nhmx = str_nhmx ? atoi(str_nhmx) : opt_nhmx;
opt_mm_select = str_mm_select ? atoi(str_mm_select) : opt_mm_select;
opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select;
+ opt_fa_head_split = str_fa_head_split ? atoi(str_fa_head_split) : opt_fa_head_split;
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
opt_ar_scatter = str_ar_scatter ? atoi(str_ar_scatter) : opt_ar_scatter;
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
index f079f7389..2ee55feb9 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
@@ -1869,8 +1869,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
factx.Bc = kparams->Bc;
factx.g_br = kparams->u.hmx.g_br;
factx.n_kv_blocks = kparams->n_kv_blocks;
- factx.is_q_fp32 = (kparams->is_q_fp32 != 0);
- factx.is_dst_fp32 = (kparams->is_dst_fp32 != 0);
+ factx.is_q_fp32 = (q->type == HTP_TYPE_F32);
+ factx.is_dst_fp32 = (dst->type == HTP_TYPE_F32);
factx.pipeline = (kparams->u.hmx.pipeline != 0);
factx.mask_broadcast = (kparams->u.hmx.mask_broadcast != 0);
if (mask) {
@@ -1879,13 +1879,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
}
factx.has_softcap = (kparams->logit_softcap != 0.0f);
- if (!factx.has_softcap) {
- factx.scale = (__fp16) (kparams->scale * EXP_LOG2E_F); // log2(e)
- } else {
- factx.scale = (__fp16) kparams->scale;
- }
+ factx.scale = (__fp16) kparams->scale;
factx.max_bias = kparams->max_bias;
- factx.logit_softcap = factx.has_softcap ? (__fp16) (kparams->logit_softcap * EXP_LOG2E_F) : 0;
+ factx.logit_softcap = 0;
+ if (factx.has_softcap) {
+ factx.logit_softcap = (__fp16) kparams->logit_softcap;
+ }
factx.n_head_log2 = kparams->n_head_log2;
factx.m0 = kparams->m0;
@@ -1898,22 +1897,36 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
const uint32_t n_threads = factx.n_threads;
const uint32_t G = factx.G;
- // Multi-device: split Q blocks across devices
+ // Multi-device: prefer head-parallel partitioning (each core owns a disjoint head
+ // shard), falling back to Q-block (token) split when heads don't divide evenly.
const uint32_t n_q_blocks = (neq1 + Br - 1) / Br;
- uint32_t q_start_min = 0;
- uint32_t q_start_max = neq1;
+ uint32_t q_start_min = 0;
+ uint32_t q_start_max = neq1;
+ uint32_t kv_head_min = 0;
+ uint32_t kv_head_max = n_kv_heads;
if (octx->ctx->mdev.count > 1) {
- const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, htp_tensor_mdev_data_aligned(dst) ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
- const uint32_t block_start = range.start;
- const uint32_t block_end = range.start + range.count;
+ const uint32_t mdev_count = octx->ctx->mdev.count;
+ const uint32_t mdev_idx = octx->ctx->mdev.idx;
+ const uint32_t dst_e_size = (dst->type == HTP_TYPE_F32) ? sizeof(float) : sizeof(__fp16);
+ const bool can_split = htp_tensor_can_row_partition(dst, dst_e_size);
+
+ if (kparams->head_split && can_split && n_kv_heads >= mdev_count && n_kv_heads % mdev_count == 0) {
+ const uint32_t kv_per_core = n_kv_heads / mdev_count;
+ kv_head_min = mdev_idx * kv_per_core;
+ kv_head_max = kv_head_min + kv_per_core;
+ } else {
+ const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, can_split ? 1 : 0, mdev_idx, mdev_count, &octx->ctx->mdev.count_div);
+ const uint32_t block_start = range.start;
+ const uint32_t block_end = range.start + range.count;
- if (block_start >= block_end) {
- return HTP_STATUS_OK;
- }
+ if (block_start >= block_end) {
+ return HTP_STATUS_OK;
+ }
- q_start_min = block_start * Br;
- q_start_max = MIN(block_end * Br, neq1);
+ q_start_min = block_start * Br;
+ q_start_max = MIN(block_end * Br, neq1);
+ }
}
// ======== VTCM allocation (GQA-aware) ========
@@ -2032,7 +2045,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
const size_t g_br_actual = hex_align_up(n_rows_g, HMX_FP16_TILE_N_ROWS);
const size_t n_row_tiles = g_br_actual / HMX_FP16_TILE_N_ROWS;
- for (uint32_t kv_head = 0; kv_head < n_kv_heads; ++kv_head) {
+ for (uint32_t kv_head = kv_head_min; kv_head < kv_head_max; ++kv_head) {
const uint32_t ik2 = kv_head;
const uint32_t ik3 = fastdiv(ib3, &kparams->broadcast_rk3);
const uint32_t iv2 = kv_head;
@@ -2040,7 +2053,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
// 1. Push Q and KV DMAs for the very first iteration.
// Subsequent iterations are enqueued early at the end of the previous iteration.
- if (ib3 == 0 && q_start == q_start_min && kv_head == 0) {
+ if (ib3 == 0 && q_start == q_start_min && kv_head == kv_head_min) {
const dma_addr_t q_ptr = q->data + q_start * q->nb[1] +
(kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3];
const size_t q_row_bytes = q_transposed ? n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed;
@@ -2358,8 +2371,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
uint32_t next_kv_head = kv_head + 1;
uint32_t next_q_start = q_start;
uint32_t next_ib3 = ib3;
- if (next_kv_head >= n_kv_heads) {
- next_kv_head = 0;
+ if (next_kv_head >= kv_head_max) {
+ next_kv_head = kv_head_min;
next_q_start = q_start + Br;
if (next_q_start >= q_start_max) {
next_q_start = q_start_min;
@@ -2478,7 +2491,7 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
factx.src3_div3 = kparams->src3_div3;
}
- factx.is_q_fp32 = (kparams->is_q_fp32 != 0);
+ factx.is_q_fp32 = (q->type == HTP_TYPE_F32);
factx.size_q_row_padded = kparams->u.hvx.size_q_row_padded;
factx.size_k_row_padded = kparams->u.hvx.size_k_row_padded;
factx.size_v_row_padded = kparams->u.hvx.size_v_row_padded;
@@ -2488,7 +2501,10 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
factx.scale = kparams->scale;
factx.max_bias = kparams->max_bias;
factx.has_softcap = (kparams->logit_softcap != 0.0f);
- factx.logit_softcap = factx.has_softcap ? (__fp16) kparams->logit_softcap : 0;
+ factx.logit_softcap = 0;
+ if (factx.has_softcap) {
+ factx.logit_softcap = (__fp16) kparams->logit_softcap;
+ }
factx.n_head_log2 = kparams->n_head_log2;
factx.m0 = kparams->m0;
@@ -2512,10 +2528,25 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
uint32_t qrows = total_qrows;
if (octx->ctx->mdev.count > 1) {
- const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
- const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
- qrow_start = range.start;
- qrows = range.count;
+ const uint32_t mdev_count = octx->ctx->mdev.count;
+ const uint32_t mdev_idx = octx->ctx->mdev.idx;
+ const uint32_t n_kv_heads = k->ne[2];
+ const uint32_t dst_e_size = (dst->type == HTP_TYPE_F32) ? sizeof(float) : sizeof(__fp16);
+ const bool can_split = htp_tensor_can_row_partition(dst, dst_e_size);
+
+ // head range is contiguous in flat row space only when neq3 == 1
+ if (kparams->head_split && can_split && neq3 == 1 && n_kv_heads >= mdev_count && n_kv_heads % mdev_count == 0) {
+ const uint32_t G = kparams->G;
+ const uint32_t kv_per_core = n_kv_heads / mdev_count;
+ const uint32_t heads_per_core = kv_per_core * G;
+ const uint32_t head_start = mdev_idx * heads_per_core;
+ qrow_start = head_start * neq1;
+ qrows = heads_per_core * neq1;
+ } else {
+ const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, mdev_idx, mdev_count, &octx->ctx->mdev.count_div);
+ qrow_start = range.start;
+ qrows = range.count;
+ }
}
if (qrows == 0) {
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
index 22bb8c53d..04538842c 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
@@ -34,8 +34,8 @@ enum htp_fa_kernel_type {
struct htp_fa_kernel_params {
uint8_t kernel_type; // enum htp_fa_kernel_type
- uint8_t is_q_fp32; // 1 = Q type is F32, 0 = F16
- uint8_t is_dst_fp32; // 1 = dst type is F32, 0 = F16
+ uint8_t head_split; // 1 = partition by KV heads in multicore, 0 = token partition
+ uint8_t flags; // reserved
uint8_t n_threads; // Number of threads to run
// Common parameters
diff --git a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
index 7d7455d04..698d6a33b 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h
@@ -745,7 +745,7 @@ void convert_f16_weight_to_fp16_tiles_task(
const uint8_t *r0 = state->src + row0 * state->row_stride;
const uint8_t *r1 = state->src + row1 * state->row_stride;
- HVX_Vector v0 = hvx_vmemu((const __fp16 *)(r0 + byte_off));
+ HVX_Vector v0 = (row0 < state->n_cols) ? hvx_vmemu((const __fp16 *)(r0 + byte_off)) : Q6_V_vzero();
HVX_Vector v1 = (row1 < state->n_cols) ? hvx_vmemu((const __fp16 *)(r1 + byte_off)) : Q6_V_vzero();
Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v0);
@@ -788,7 +788,7 @@ void quantize_f32_weight_to_fp16_tiles_task(
const uint8_t *r0 = state->src + row0 * state->row_stride;
const uint8_t *r1 = state->src + row1 * state->row_stride;
- HVX_Vector v0_f32 = hvx_vmem((const float *)(r0 + byte_off));
+ HVX_Vector v0_f32 = (row0 < state->n_cols) ? hvx_vmem((const float *)(r0 + byte_off)) : Q6_V_vzero();
HVX_Vector v1_f32 = (row1 < state->n_cols) ? hvx_vmem((const float *)(r1 + byte_off)) : Q6_V_vzero();
HVX_Vector v_out = hvx_vec_f32_to_f16(v0_f32, v1_f32);
@@ -988,9 +988,7 @@ static void transfer_output_chunk_fp16_to_fp32_col_chunk(
uint32_t src2_stride,
uint32_t dst_cols
) {
- assert(c_len % HTP_MM_HMX_TILE_N_COLS == 0);
- assert(total_n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
- const size_t tile_row_stride = (total_n_cols / HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;
+ const size_t tile_row_stride = hmx_ceil_div(total_n_cols, HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;
const HVX_Vector one = hvx_vec_splat_f16(1.0);
@@ -1137,6 +1135,73 @@ static void transfer_activation_row_pair_fp32_to_fp16(
}
}
+// F16-input variant of transfer_activation_row_pair_fp32_to_fp16, for F16 activation
+// (src1). Same shape as the F16 Q-prep in hmx-fa-kernels.h: one 128-byte load carries 64 f16
+// = two tile columns, and Q6_W_vshuff_VVR interleaves the two rows straight into the HMX tile
+// layout, so no F32 round-trip is needed. Rows are only 64-byte aligned when k_block is an odd
+// multiple of the tile width, hence the unaligned load type.
+static void transfer_activation_row_pair_f16_to_f16(__fp16 * restrict vtcm_dst,
+ const __fp16 * restrict row0,
+ const __fp16 * restrict row1,
+ uint32_t r,
+ uint32_t k_block,
+ uint32_t k_valid,
+ bool row0_valid,
+ bool row1_valid) {
+ uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
+ uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
+
+ const uint32_t n_tile_cols = k_block / HTP_MM_HMX_TILE_N_COLS;
+ __fp16 * restrict tile_row = vtcm_dst + (size_t) r0 * n_tile_cols * HTP_MM_HMX_TILE_N_ELMS;
+
+ const HVX_UVector * pv0 = row0_valid ? (const HVX_UVector *) row0 : NULL;
+ const HVX_UVector * pv1 = row1_valid ? (const HVX_UVector *) row1 : NULL;
+
+ uint32_t c = 0;
+ for (; c + 64 <= k_valid; c += 64) {
+ HVX_Vector v0 = pv0 ? pv0[c / 64] : Q6_V_vzero();
+ HVX_Vector v1 = pv1 ? pv1[c / 64] : Q6_V_vzero();
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
+
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
+
+ HVX_Vector * tile0 = (HVX_Vector *) (tile_row + (size_t) c0 * HTP_MM_HMX_TILE_N_ELMS);
+ HVX_Vector * tile1 = (HVX_Vector *) (tile_row + (size_t) (c0 + 1) * HTP_MM_HMX_TILE_N_ELMS);
+
+ tile0[r1 / 2] = Q6_V_lo_W(vp);
+ tile1[r1 / 2] = Q6_V_hi_W(vp);
+ }
+ // Tail: fewer than 64 valid columns left, plus the k_valid..k_block padding that HMX will
+ // still multiply, so it has to be written as zeros.
+ for (; c < k_block; c += 64) {
+ HVX_Vector v0 = Q6_V_vzero();
+ HVX_Vector v1 = Q6_V_vzero();
+
+ if (c < k_valid) {
+ uint32_t rem = k_valid - c; // 1..63 valid f16 lanes
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem * sizeof(__fp16));
+ if (pv0) {
+ v0 = Q6_V_vmux_QVV(mask, pv0[c / 64], Q6_V_vzero());
+ }
+ if (pv1) {
+ v1 = Q6_V_vmux_QVV(mask, pv1[c / 64], Q6_V_vzero());
+ }
+ }
+
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
+
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
+
+ HVX_Vector * tile0 = (HVX_Vector *) (tile_row + (size_t) c0 * HTP_MM_HMX_TILE_N_ELMS);
+ tile0[r1 / 2] = Q6_V_lo_W(vp);
+
+ if (c0 + 1 < n_tile_cols) {
+ HVX_Vector * tile1 = (HVX_Vector *) (tile_row + (size_t) (c0 + 1) * HTP_MM_HMX_TILE_N_ELMS);
+ tile1[r1 / 2] = Q6_V_hi_W(vp);
+ }
+ }
+}
+
static void transfer_activation_row_pair_fp32_to_fp16_col_chunk(
__fp16 *restrict vtcm_dst,
const float *restrict row0, // offset by c_first
diff --git a/ggml/src/ggml-hexagon/htp/hmx-utils.h b/ggml/src/ggml-hexagon/htp/hmx-utils.h
index ad295cb7d..1952aaa2c 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-utils.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-utils.h
@@ -73,19 +73,20 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
for (uint32_t r = start_row; r < end_row; r += 2) {
const uint32_t ct = r / HMX_FP16_TILE_N_ROWS;
const uint32_t local_r = r % HMX_FP16_TILE_N_ROWS;
+ const bool row0_valid = r < n_cols;
const bool next_row_valid = (r + 1) < end_row && (r + 1) < n_cols;
const HVX_Vector v_off0 = Q6_Vw_vadd_VwVw(v_scat_base, Q6_V_vsplat_R(local_r * 4));
const HVX_Vector v_off1 = Q6_Vw_vadd_VwVw(v_off0, v_scat_step);
__fp16 * tile_base = vtcm_dst + (size_t) ct * n_k_tiles * HMX_FP16_TILE_N_ELMS;
- const uint8_t * p0 = (const uint8_t *) (vtcm_src + r * src_stride);
+ const uint8_t * p0 = row0_valid ? (const uint8_t *) (vtcm_src + r * src_stride) : NULL;
const uint8_t * p1 = next_row_valid ? (const uint8_t *) (vtcm_src + (r + 1) * src_stride) : NULL;
- assert(hex_is_aligned(p0, 128));
- assert(hex_is_aligned(p1, 128));
+ assert(!p0 || hex_is_aligned(p0, 128));
+ assert(!p1 || hex_is_aligned(p1, 128));
assert(c_byte_step % 128 == 0);
- if (p1) {
+ if (p0 && p1) {
for (uint32_t i = 0; i < n_c_iters; ++i) {
HVX_Vector v0 = hvx_vmem(p0); p0 += c_byte_step;
HVX_Vector v1 = hvx_vmem(p1); p1 += c_byte_step;
@@ -96,9 +97,12 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
} else {
const HVX_Vector vzero = Q6_V_vzero();
for (uint32_t i = 0; i < n_c_iters; ++i) {
- HVX_Vector v0 = hvx_vmem(p0); p0 += c_byte_step;
+ HVX_Vector v0 = p0 ? hvx_vmem(p0) : vzero;
+ if (p0) p0 += c_byte_step;
+ HVX_Vector v1 = p1 ? hvx_vmem(p1) : vzero;
+ if (p1) p1 += c_byte_step;
Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off0, v0);
- Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off1, vzero);
+ Q6_vscatter_RMVwV((size_t) tile_base, pair_region, v_off1, v1);
tile_base += dst_step;
}
}
@@ -113,15 +117,16 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
for (uint32_t r = start_row; r < end_row; r += 2) {
const uint32_t ct = r / HMX_FP16_TILE_N_ROWS;
const uint32_t local_r = r % HMX_FP16_TILE_N_ROWS;
+ const bool row0_valid = r < n_cols;
const bool next_row_valid = (r + 1) < end_row && (r + 1) < n_cols;
const HVX_Vector v_off0 = Q6_Vw_vadd_VwVw(v_scat_base, Q6_V_vsplat_R(local_r * 4));
const HVX_Vector v_off1 = Q6_Vw_vadd_VwVw(v_off0, v_scat_step);
__fp16 * tile_base = vtcm_dst + (size_t) ct * n_k_tiles * HMX_FP16_TILE_N_ELMS;
- const uint8_t * p0 = (const uint8_t *) (vtcm_src + r * src_stride);
+ const uint8_t * p0 = row0_valid ? (const uint8_t *) (vtcm_src + r * src_stride) : NULL;
const uint8_t * p1 = next_row_valid ? (const uint8_t *) (vtcm_src + (r + 1) * src_stride) : NULL;
- if (p1) {
+ if (p0 && p1) {
for (uint32_t i = 0; i < n_c_iters; ++i) {
HVX_Vector v0 = hvx_vmemu(p0); p0 += c_byte_step;
HVX_Vector v1 = hvx_vmemu(p1); p1 += c_byte_step;
@@ -132,9 +137,12 @@ static inline void hmx_interleave_rows_to_tiles(__fp16 * restrict vtcm_dst,
} else {
const HVX_Vector vzero = Q6_V_vzero();
for (uint32_t i = 0; i < n_c_iters; ++i) {
- HVX_Vector v0 = hvx_vmemu(p0); p0 += c_byte_step;
+ HVX_Vector v0 = p0 ? hvx_vmemu(p0) : vzero;
+ if (p0) p0 += c_byte_step;
+ HVX_Vector v1 = p1 ? hvx_vmemu(p1) : vzero;
+ if (p1) p1 += c_byte_step;
Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off0, v0);
- Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off1, vzero);
+ Q6_vscatter_QRMVwV(q_mask64, (size_t) tile_base, single_region, v_off1, v1);
tile_base += dst_step;
}
}
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c
index 9dfd35649..a589e0a20 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.c
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c
@@ -27,8 +27,8 @@
typedef struct {
float *dst;
- dma_addr_t src2_addr;
- size_t src2_bytes;
+ dma_addr_t bias_addr;
+ size_t bias_bytes;
dma_addr_t act_dma_addr;
dma_addr_t weight;
dma_queue * weight_dma;
@@ -36,9 +36,10 @@ typedef struct {
int k;
int n;
int act_stride;
+ uint32_t act_elem_size; // 4=F32 src1, 2=F16 src1
int weight_stride;
int dst_stride;
- uint32_t src2_stride;
+ uint32_t bias_stride;
int ne02;
int ne03;
int ne12;
@@ -47,8 +48,8 @@ typedef struct {
size_t src0_nb3;
size_t act_nb2;
size_t act_nb3;
- size_t src2_nb2;
- size_t src2_nb3;
+ size_t bias_nb2;
+ size_t bias_nb3;
size_t dst_nb2;
size_t dst_nb3;
int r2;
@@ -118,24 +119,19 @@ struct htp_mm_context {
// Dynamic VTCM pointers allocated sequentially
uint8_t * vtcm_src0;
- uint8_t * vtcm_src1;
- uint8_t * vtcm_src2;
- uint8_t * vtcm_src3;
+ uint8_t * vtcm_act;
+ uint8_t * vtcm_bias;
uint8_t * vtcm_dst;
uint8_t * vtcm_act_raw;
// Cached strides
uint32_t vtcm_src0_stride;
- uint32_t vtcm_src1_stride;
- uint32_t vtcm_src2_stride;
- uint32_t vtcm_src3_stride;
+ uint32_t vtcm_act_stride;
uint32_t vtcm_act_raw_stride;
// Cached thread offsets/sizes
uint32_t vtcm_src0_size_per_thread;
- uint32_t vtcm_src1_size_per_thread;
- uint32_t vtcm_src2_size_per_thread;
- uint32_t vtcm_src3_size_per_thread;
+ uint32_t vtcm_act_size_per_thread;
uint32_t vtcm_dst_size_per_thread;
};
@@ -190,6 +186,7 @@ static const uint8_t __attribute__((aligned(VLEN))) kvalues_mxfp4_lut[] = {
const struct htp_tensor * restrict src1 = octx->src[1]; \
const struct htp_tensor * restrict src2 = octx->src[2]; \
const struct htp_tensor * restrict dst = octx->dst; \
+ const struct htp_tensor * restrict act = src1; \
\
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
@@ -247,7 +244,7 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
htp_matmul_preamble; \
\
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; \
- const uint32_t src1_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : (ne11 * ne12 * ne13); \
+ const uint32_t act_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : (ne11 * ne12 * ne13); \
const uint32_t cur_m_start = mmctx->cur_m_start; \
\
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith; \
@@ -260,13 +257,13 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); \
\
const size_t dst_row_size = nb1; \
- const size_t src1_row_size = nb11; \
- const size_t src1_stride = mmctx->vtcm_src1_stride; \
+ const size_t act_row_size = nb11; \
+ const size_t act_stride = mmctx->vtcm_act_stride; \
const size_t src2_stride = src2 ? ((src2->ne[1] == 1) ? 0 : src2->nb[1]) : 0; \
\
uint8_t * restrict vtcm_dst_ptr = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith; \
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \
- uint8_t * restrict src1_data = mmctx->vtcm_src1; \
+ uint8_t * restrict act_data = mmctx->vtcm_act; \
\
const dma_addr_t src0_row = src0->data; \
\
@@ -301,9 +298,9 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
uint32_t ir1 = 0; \
- for (; ir1 + 1 < src1_nrows; ir1 += 2) { \
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); \
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); \
+ for (; ir1 + 1 < act_nrows; ir1 += 2) { \
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride); \
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_stride); \
float * restrict dst_row0 = (float *) (dst->data + ((cur_m_start + ir1+0) * dst_row_size)); \
float * restrict dst_row1 = (float *) (dst->data + ((cur_m_start + ir1+1) * dst_row_size)); \
\
@@ -318,11 +315,11 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
src2_ptr0 = &src2_row0[ct * 32]; \
src2_ptr1 = &src2_row1[ct * 32]; \
} \
- DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, src2_ptr0, src2_ptr1); \
+ DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_col1, valid_rows, src2_ptr0, src2_ptr1); \
} \
\
- for (; ir1 < src1_nrows; ++ir1) { \
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); \
+ for (; ir1 < act_nrows; ++ir1) { \
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride); \
float * restrict dst_row = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size)); \
float * dst_ptr = &dst_row[ct * 32]; \
\
@@ -331,7 +328,7 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
const float * restrict src2_row = (const float *) ((const uint8_t *) src2->data + ((cur_m_start + ir1) * src2_stride)); \
src2_ptr = &src2_row[ct * 32]; \
} \
- DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, src2_ptr); \
+ DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, src2_ptr); \
} \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
\
@@ -359,18 +356,18 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); \
\
const size_t dst_row_size = nb1; \
- const size_t src1_row_size = nb11; \
- const size_t src1_stride = mmctx->vtcm_src1_stride; \
+ const size_t act_row_size = nb11; \
+ const size_t act_stride = mmctx->vtcm_act_stride; \
\
uint8_t * vtcm_dst_ptr = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith; \
uint8_t * vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \
- uint8_t * src1_data = mmctx->vtcm_src1; \
+ uint8_t * act_data = mmctx->vtcm_act; \
\
float * tmp = (float *) vtcm_dst_ptr; \
\
const dma_addr_t src0_row = src0->data; \
\
- const uint8_t * restrict src1_col = (const uint8_t *) src1_data; \
+ const uint8_t * restrict act_col = (const uint8_t *) act_data; \
float * restrict dst_col = (float *) dst->data; \
\
const uint32_t tile_size = TILE_SIZE; \
@@ -387,11 +384,11 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
uint32_t push_ct = ct_start; \
if (src0_start_row < src0_end_row) { \
if (src2) { \
- float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row; \
+ float * vtcm_bias_ptr = (float *) mmctx->vtcm_bias + src0_start_row; \
const dma_addr_t src2_addr = src2->data + src0_start_row * sizeof(float); \
int slice_size = (int)MIN(src0_end_row, ne0) - (int)src0_start_row; \
if (slice_size > 0) { \
- dma_queue_push(dma_q, dma_make_data(vtcm_src2_ptr, src2_addr), \
+ dma_queue_push(dma_q, dma_make_data(vtcm_bias_ptr, src2_addr), \
slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1); \
dma_queue_pop_nowait(dma_q); \
} \
@@ -414,7 +411,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
valid_rows = MIN(32, MAX(0, valid_rows)); \
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
- DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \
+ DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL); \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
\
if (push_ct < ct_end) { \
@@ -430,7 +427,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
if (src2) { \
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row], \
(const uint8_t *) tmp, \
- (const uint8_t *) ((const float *) mmctx->vtcm_src2 + src0_start_row), \
+ (const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row), \
copy_cnt); \
} else { \
hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt); \
@@ -448,11 +445,11 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
\
const struct htp_tensor * restrict act = octx->src[n_weights]; /* x */ \
const uint32_t ne10 = act->ne[0]; \
- const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3]; \
- const size_t src1_stride = mmctx->vtcm_src1_stride; \
+ const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3]; \
+ const size_t act_stride = mmctx->vtcm_act_stride; \
\
uint8_t * restrict vtcm_weight_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \
- uint8_t * restrict src1_data = mmctx->vtcm_src1; \
+ uint8_t * restrict act_data = mmctx->vtcm_act; \
\
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
const uint32_t n_prefetch = kparams->n_prefetch; \
@@ -511,23 +508,23 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
uint32_t ir1 = 0; \
- for (; ir1 + 1 < src1_nrows; ir1 += 2) { \
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); \
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); \
+ for (; ir1 + 1 < act_nrows; ir1 += 2) { \
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride); \
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_stride); \
\
float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size)); \
float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size)); \
float * dst_ptr0 = &dst_row0[ct * 32]; \
float * dst_ptr1 = &dst_row1[ct * 32]; \
\
- DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, NULL, NULL); \
+ DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_col1, valid_rows, NULL, NULL); \
} \
\
- for (; ir1 < src1_nrows; ++ir1) { \
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); \
+ for (; ir1 < act_nrows; ++ir1) { \
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride); \
float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size)); \
float * dst_ptr = &dst_row[ct * 32]; \
- DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \
+ DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL); \
} \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
\
@@ -550,10 +547,10 @@ MATMUL_2D_REPACKED_IMPL(q2_k, 512, tiled_vec_dot_q2_k_32x2, tiled_vec_do
MATMUL_2D_REPACKED_IMPL(iq4nl, 576, tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1)
MATMUL_2D_REPACKED_IMPL(mxfp4, 544, tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1)
-static void hvx_mm_transfer_src1_dma(
+static void hvx_mm_transfer_act_dma(
struct htp_ops_context * octx,
const struct htp_mm_kernel_params * kparams,
- const struct htp_tensor * src1,
+ const struct htp_tensor * act,
uint8_t * dst_base,
size_t dst_row_size,
uint32_t m_start,
@@ -564,22 +561,22 @@ static void hvx_mm_transfer_src1_dma(
}
dma_queue * dma_q = octx->ctx->dma[0];
- const uint32_t ne0 = src1->ne[0];
- const size_t elem_size = (src1->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+ const uint32_t ne0 = act->ne[0];
+ const size_t elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
const size_t row_bytes = ne0 * elem_size;
- const size_t src1_nb1 = src1->nb[1];
- const dma_addr_t src_base = src1->data;
+ const size_t act_nb1 = act->nb[1];
+ const dma_addr_t act_base = act->data;
- const bool is_contiguous = (src1->nb[2] == src1->ne[1] * src1_nb1) &&
- (src1->nb[3] == src1->ne[2] * src1->nb[2]);
+ const bool is_contiguous = (act->nb[2] == act->ne[1] * act_nb1) &&
+ (act->nb[3] == act->ne[2] * act->nb[2]);
if (is_contiguous) {
- const dma_addr_t src_addr = src_base + m_start * src1_nb1;
+ const dma_addr_t src_addr = act_base + m_start * act_nb1;
dma_queue_push(dma_q, dma_make_data(dst_base, src_addr),
- dst_row_size, src1_nb1, row_bytes, m_rows);
+ dst_row_size, act_nb1, row_bytes, m_rows);
dma_queue_pop(dma_q);
} else {
- const uint32_t ne12_ne1 = src1->ne[2] * src1->ne[1];
+ const uint32_t ne12_ne1 = act->ne[2] * act->ne[1];
const bool use_fastdiv = kparams->div_ne12_ne1.mp != 0;
for (uint32_t ir = 0; ir < m_rows; ++ir) {
const uint32_t ir1 = m_start + ir;
@@ -588,19 +585,19 @@ static void hvx_mm_transfer_src1_dma(
i13 = fastdiv(ir1, &kparams->div_ne12_ne1);
const uint32_t rem = ir1 - i13 * ne12_ne1;
i12 = fastdiv(rem, &kparams->div_ne1);
- i11 = rem - i12 * src1->ne[1];
+ i11 = rem - i12 * act->ne[1];
} else {
i13 = ne12_ne1 ? ir1 / ne12_ne1 : 0;
const uint32_t rem = ir1 - i13 * ne12_ne1;
- i12 = src1->ne[1] ? rem / src1->ne[1] : 0;
- i11 = rem - i12 * src1->ne[1];
+ i12 = act->ne[1] ? rem / act->ne[1] : 0;
+ i11 = rem - i12 * act->ne[1];
}
- const dma_addr_t row_src = src_base + (i11 * src1->nb[1] +
- i12 * src1->nb[2] +
- i13 * src1->nb[3]);
+ const dma_addr_t row_src = act_base + (i11 * act->nb[1] +
+ i12 * act->nb[2] +
+ i13 * act->nb[3]);
uint8_t * row_dst = dst_base + ir * dst_row_size;
dma_queue_push(dma_q, dma_make_data(row_dst, row_src),
- dst_row_size, src1_nb1, row_bytes, 1);
+ dst_row_size, act_nb1, row_bytes, 1);
dma_queue_pop(dma_q);
}
}
@@ -640,7 +637,7 @@ static void name(unsigned int nth, unsigned int ith, void * data) {
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, ir_first); \
\
- uint8_t * restrict dst = mmctx->vtcm_src1; \
+ uint8_t * restrict dst = mmctx->vtcm_act; \
const uint32_t ir_last = MIN(ir_first + nrows_per_thread, nrows); \
const size_t raw_row_size = mmctx->vtcm_act_raw_stride; \
const size_t dst_row_size = (dst_row_size_expr); \
@@ -655,9 +652,9 @@ static void name(unsigned int nth, unsigned int ith, void * data) {
QUANTIZE_IMPL(quantize_f32_q8_0_tiled, "quantize-f32-q8_0_tiled", quantize_f32_q8_0_tiled_kernel, htp_mm_q8_0_tiled_row_size(ne0))
QUANTIZE_IMPL(quantize_f32_q8_1_tiled, "quantize-f32-q8_1_tiled", quantize_f32_q8_1_tiled_kernel, htp_mm_q8_1_tiled_row_size(ne0))
QUANTIZE_IMPL(quantize_f32_q8_1_s16_tiled, "quantize-f32-q8_1_s16_tiled", quantize_f32_q8_1_s16_tiled_kernel, htp_mm_q8_1_tiled_row_size(ne0))
-QUANTIZE_IMPL(quantize_f32_f32, "quantize-f32-f32", quantize_f32_f32_kernel, mmctx->vtcm_src1_stride)
-QUANTIZE_IMPL(quantize_f32_f16, "quantize-f32-f16", quantize_f32_f16_kernel, mmctx->vtcm_src1_stride)
-QUANTIZE_IMPL(quantize_f16_f16, "quantize-f16-f16", quantize_f16_f16_kernel, mmctx->vtcm_src1_stride)
+QUANTIZE_IMPL(quantize_f32_f32, "quantize-f32-f32", quantize_f32_f32_kernel, mmctx->vtcm_act_stride)
+QUANTIZE_IMPL(quantize_f32_f16, "quantize-f32-f16", quantize_f32_f16_kernel, mmctx->vtcm_act_stride)
+QUANTIZE_IMPL(quantize_f16_f16, "quantize-f16-f16", quantize_f16_f16_kernel, mmctx->vtcm_act_stride)
static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, void * data) {
(void) nth;
@@ -673,7 +670,7 @@ static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, vo
quantize_f32_q8_0_tiled_block_kernel(
(const float *) mmctx->vtcm_act_raw,
- mmctx->vtcm_src1,
+ mmctx->vtcm_act,
NULL,
src->ne[0],
mmctx->quant_ib_first[ith],
@@ -701,7 +698,7 @@ static void quantize_f32_q8_1_tiled_block(unsigned int nth, unsigned int ith, vo
quantize_f32_q8_1_tiled_block_kernel(
(const float *) mmctx->vtcm_act_raw,
- mmctx->vtcm_src1,
+ mmctx->vtcm_act,
NULL,
src->ne[0],
mmctx->quant_ib_first[ith],
@@ -729,7 +726,7 @@ static void quantize_f32_q8_1_s16_tiled_block(unsigned int nth, unsigned int ith
quantize_f32_q8_1_s16_tiled_block_kernel(
(const float *) mmctx->vtcm_act_raw,
- mmctx->vtcm_src1,
+ mmctx->vtcm_act,
NULL,
src->ne[0],
mmctx->quant_ib_first[ith],
@@ -795,10 +792,10 @@ static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); \
\
const size_t dst_row_size = nb1; \
- const size_t src1_stride = mmctx->vtcm_src1_stride; \
+ const size_t act_stride = mmctx->vtcm_act_stride; \
\
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \
- uint8_t * restrict src1_data = mmctx->vtcm_src1; \
+ uint8_t * restrict act_data = mmctx->vtcm_act; \
\
const uint32_t tile_size = TILE_SIZE; \
const uint32_t aligned_tile_size = hex_align_up(tile_size, 128); \
@@ -874,19 +871,19 @@ static void hvx_mm_4d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
\
uint32_t ir1 = 0; \
for (; ir1 + 1 < batch_nrows; ir1 += 2) { \
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 0) * src1_stride); \
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 1) * src1_stride); \
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 0) * act_stride); \
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 1) * act_stride); \
float * restrict dst_row0 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 0) * dst_row_size); \
float * restrict dst_row1 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 1) * dst_row_size); \
float * dst_ptr0 = &dst_row0[ct * 32]; \
float * dst_ptr1 = &dst_row1[ct * 32]; \
- DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, NULL, NULL); \
+ DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, act_col0, act_col1, valid_rows, NULL, NULL); \
} \
for (; ir1 < batch_nrows; ++ir1) { \
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride); \
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_stride); \
float * restrict dst_row = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size); \
float * dst_ptr = &dst_row[ct * 32]; \
- DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \
+ DOT_2X1(ne10, dst_ptr, w_tile, act_col, valid_rows, NULL); \
} \
} \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \
@@ -923,7 +920,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
const uint32_t prefetch_mask = n_prefetch - 1;
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; // src0 rows
- const uint32_t src1_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows; // src1 rows
+ const uint32_t act_nrows = mmctx->cur_m_rows ? mmctx->cur_m_rows : mmctx->act_nrows;
const uint32_t cur_m_start = mmctx->cur_m_start;
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
@@ -934,15 +931,15 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
const size_t dst_row_size = nb1;
const size_t src0_row_size = nb01;
- const size_t src1_row_size = nb11;
+ const size_t act_row_size = nb11;
const size_t src0_stride = mmctx->vtcm_src0_stride;
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const size_t act_stride = mmctx->vtcm_act_stride;
// Per-thread VTCMs for all tensors
uint8_t * restrict vtcm_dst_ptr = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith;
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
const dma_addr_t src0_row = src0->data;
@@ -968,21 +965,21 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
- // Process src1 columns in pairs (2x2 tiling)
+ // Process act columns in pairs (2x2 tiling)
uint32_t ir1 = 0;
- for (; ir1 + 1 < src1_nrows; ir1 += 2) {
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride);
+ for (; ir1 + 1 < act_nrows; ir1 += 2) {
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_stride);
float * restrict dst_row0 = (float *) (dst->data + ((cur_m_start + ir1+0) * dst_row_size));
float * restrict dst_row1 = (float *) (dst->data + ((cur_m_start + ir1+1) * dst_row_size));
- mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+ mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_col1);
}
- // Handle remaining src1 rows (fallback to 2x1)
- for (; ir1 < src1_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+ // Handle remaining act rows (fallback to 2x1)
+ for (; ir1 < act_nrows; ++ir1) {
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
float * restrict dst_row = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));
- mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
+ mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -1005,15 +1002,15 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
#pragma unroll(2)
- for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+ for (uint32_t ir1 = 0; ir1 < act_nrows; ++ir1) {
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
float * restrict dst_row = (float *) (dst->data + ((cur_m_start + ir1) * dst_row_size));
- mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
+ mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
}
if (src2) {
- hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + src1_nrows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
+ hvx_tensor_add_f32_grid(dst, src2, cur_m_start, cur_m_start + act_nrows, src0_start_row, src0_end_row, &kparams->div_ne12_ne1, &kparams->div_ne1);
}
}
@@ -1029,21 +1026,21 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
const size_t dst_row_size = nb1;
const size_t src0_row_size = nb01;
- const size_t src1_row_size = nb11;
+ const size_t act_row_size = nb11;
const size_t src0_stride = mmctx->vtcm_src0_stride;
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const size_t act_stride = mmctx->vtcm_act_stride;
// Per-thread VTCMs for all tensors
uint8_t * vtcm_dst_ptr = mmctx->vtcm_dst + mmctx->vtcm_dst_size_per_thread * ith;
uint8_t * vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * src1_data = mmctx->vtcm_src1;
+ uint8_t * act_data = mmctx->vtcm_act;
float * tmp = (float *) vtcm_dst_ptr;
const dma_addr_t src0_row = src0->data;
- const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
- float * restrict dst_col = (float *) dst->data;
+ const uint8_t * restrict act_col = (const uint8_t *) act_data;
+ float * restrict dst_col = (float *) dst->data;
const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);
@@ -1055,11 +1052,11 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
// Prefill vtcm with 2x src0 rows
if (src0_start_row < src0_end_row) {
if (src2) {
- float * vtcm_src2_ptr = (float *) mmctx->vtcm_src2 + src0_start_row;
+ float * vtcm_bias_ptr = (float *) mmctx->vtcm_bias + src0_start_row;
const dma_addr_t src2_addr = src2->data + src0_start_row * sizeof(float);
int slice_size = (int)src0_end_row - (int)src0_start_row;
if (slice_size > 0) {
- dma_queue_push(dma_q, dma_make_data(vtcm_src2_ptr, src2_addr),
+ dma_queue_push(dma_q, dma_make_data(vtcm_bias_ptr, src2_addr),
slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1);
dma_queue_pop_nowait(dma_q);
}
@@ -1083,7 +1080,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) {
const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
- mmctx->vec_dot_2x1(ne00, &tmp[ir0 - src0_start_row], ss0, ss0 + src0_stride, src1_col);
+ mmctx->vec_dot_2x1(ne00, &tmp[ir0 - src0_start_row], ss0, ss0 + src0_stride, act_col);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
// Prefetch next (n + vtcm_nrows) row
@@ -1103,7 +1100,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
src0_stride, src0_row_size, src0_row_size, 1);
const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
- mmctx->vec_dot_1x1(ne00, &tmp[ir0 - src0_start_row], ss0, src1_col);
+ mmctx->vec_dot_1x1(ne00, &tmp[ir0 - src0_start_row], ss0, act_col);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
}
@@ -1113,7 +1110,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
if (src2) {
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row],
(const uint8_t *) tmp,
- (const uint8_t *) ((const float *) mmctx->vtcm_src2 + src0_start_row),
+ (const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row),
copy_cnt);
} else {
hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt);
@@ -1143,10 +1140,10 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
const size_t dst_row_size = nb1;
const size_t src0_row_size = nb01;
const size_t src0_stride = mmctx->vtcm_src0_stride;
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const size_t act_stride = mmctx->vtcm_act_stride;
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
if (src0_start_row >= src0_end_row || cur_m_rows == 0) {
@@ -1208,16 +1205,16 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
uint32_t ir1 = 0;
for (; ir1 + 1 < batch_nrows; ir1 += 2) {
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 0) * src1_stride);
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (chunk_m_offset + ir1 + 1) * src1_stride);
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 0) * act_stride);
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (chunk_m_offset + ir1 + 1) * act_stride);
float * restrict dst_row0 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 0) * dst_row_size);
float * restrict dst_row1 = (float *) (dst_batch_base + (dst_m_offset + ir1 + 1) * dst_row_size);
- mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+ mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_col1);
}
for (; ir1 < batch_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_stride);
float * restrict dst_row = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size);
- mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
+ mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -1253,9 +1250,9 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
const uint32_t batch_nrows = m_last - m_first;
for (uint32_t ir1 = 0; ir1 < batch_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (chunk_m_offset + ir1) * src1_stride);
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + (chunk_m_offset + ir1) * act_stride);
float * restrict dst_row = (float *) (dst_batch_base + (dst_m_offset + ir1) * dst_row_size);
- mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
+ mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -1277,7 +1274,7 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
const struct htp_tensor * restrict ids = octx->src[2];
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; // src0 rows per expert
- const uint32_t src1_nrows = ne11;
+ const uint32_t act_nrows = ne11;
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
@@ -1299,13 +1296,13 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows;
const size_t dst_row_size = nb1;
- const size_t src1_row_size = htp_mm_q8_0_tiled_row_size(ne10);
+ const size_t act_row_size = htp_mm_q8_0_tiled_row_size(ne10);
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const size_t act_stride = mmctx->vtcm_act_stride;
// Per-thread VTCMs for all tensors
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
const int32_t cne1 = matrix_row_counts[cur_a];
@@ -1343,11 +1340,11 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
const int rm1 = row_mapping.i1; // expert idx
const int rm2 = row_mapping.i2; // token idx
- const uint32_t ir1 = fastmodulo(rm1, ne11, &mmctx->mm_div_ne11); // src1 row idx
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (ir1 + rm2 * ne11 + 0) * src1_stride);
+ const uint32_t ir1 = fastmodulo(rm1, ne11, &mmctx->mm_div_ne11); // act row idx
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + (ir1 + rm2 * ne11 + 0) * act_stride);
float * restrict dst_row = (float *) (dst->data + (rm1 * nb1 + rm2 * nb2 + 0));
- mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+ mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);
@@ -1383,14 +1380,14 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
assert(ne13 % ne03 == 0);
const size_t dst_row_size = nb1;
- const size_t src1_row_size = htp_mm_q8_0_tiled_row_size(ne10);
+ const size_t act_row_size = htp_mm_q8_0_tiled_row_size(ne10);
const uint32_t n_aids = src2->ne[0]; // num activated experts
const uint32_t n_ids = ne02; // num experts
// Per-thread VTCMs for all tensors
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
for (uint32_t ie1 = 0; ie1 < n_aids; ++ie1) { // for each expert
const int32_t eid = *(const int32_t *) ((const uint8_t *) src2->data + ie1 * src2->nb[0]);
@@ -1400,7 +1397,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
assert(eid < (int32_t) n_ids);
const dma_addr_t src0_row = src0->data + eid * nb02;
- const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
+ const uint8_t * restrict act_col = (const uint8_t *) act_data;
float * restrict dst_row = (float *) (dst->data + ie1 * nb1);
const uint32_t tile_size = htp_mm_get_weight_tile_size(src0->type);
@@ -1426,7 +1423,7 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
valid_rows = MIN(32, MAX(0, valid_rows));
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);
- mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+ mmctx->vec_dot_32x1(ne10, &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);
if (push_ct < ct_end) {
@@ -1457,7 +1454,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
const uint32_t n_ids = src0->ne[2];
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
for (uint32_t ie1 = 0; ie1 < n_aids; ++ie1) {
const int32_t eid = *(const int32_t *) ((const uint8_t *) ids->data + ie1 * ids->nb[0]);
@@ -1489,7 +1486,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
if (src0_start_row >= src0_end_row) continue;
const dma_addr_t src0_row = src_w->data + eid * src_w->nb[2];
- const uint8_t * restrict src1_col = (const uint8_t *) src1_data;
+ const uint8_t * restrict act_col = (const uint8_t *) act_data;
float * restrict dst_row = (float *) (dst->data + ie1 * dst->nb[1]);
const uint32_t tile_size = htp_mm_get_weight_tile_size(src_w->type);
@@ -1515,7 +1512,7 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
valid_rows = MIN(32, MAX(0, valid_rows));
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct);
- mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+ mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);
if (push_ct < ct_end) {
@@ -1548,10 +1545,10 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
const uint32_t * matrix_row_counts = mmctx->matrix_row_counts;
const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows;
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const size_t act_stride = mmctx->vtcm_act_stride;
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
const int32_t cne1 = matrix_row_counts[cur_a];
@@ -1612,10 +1609,10 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
const int rm2 = row_mapping.i2;
const uint32_t ir1 = fastmodulo(rm1, act->ne[1], &mmctx->mm_div_ne11);
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (ir1 + rm2 * act->ne[1]) * src1_stride);
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + (ir1 + rm2 * act->ne[1]) * act_stride);
float * restrict dst_row = (float *) (dst->data + (rm1 * dst->nb[1] + rm2 * dst->nb[2]));
- mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL);
+ mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, act_col, valid_rows, NULL);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct);
@@ -1682,13 +1679,13 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
struct htp_mm_context mmctx_struct = {0};
struct htp_mm_context * mmctx = &mmctx_struct;
mmctx->octx = octx;
- mmctx->act = src1;
+ mmctx->act = act;
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const uint32_t src0_nrows = ne01;
- const uint32_t src1_nrows = ne11 * ne12 * ne13;
- mmctx->act_nrows = src1_nrows;
+ const uint32_t act_nrows = ne11 * ne12 * ne13;
+ mmctx->act_nrows = act_nrows;
uint32_t src0_row_start = 0;
uint32_t src0_row_end = src0_nrows;
@@ -1724,10 +1721,9 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
const size_t src0_row_size = nb01;
const size_t dst_row_size = nb1;
- size_t src1_row_size = nb11;
+ size_t act_row_size = nb11;
const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
- size_t src1_row_size_padded;
worker_callback_t quant_task_func;
worker_callback_t matmul_job_func;
@@ -1751,7 +1747,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
} else {
matmul_job_func = hvx_mm_4d;
}
- } else if (src1_nrows > 1) {
+ } else if (act_nrows > 1) {
if (is_repacked) {
switch (src0->type) {
case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_2d_repacked_q4_0; break;
@@ -1793,13 +1789,13 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
switch (kparams->kernel_type) {
case HTP_MM_KERNEL_HVX_F16_F16_VTCM:
- quant_task_func = (src1->type == HTP_TYPE_F32) ? quantize_f32_f16 : quantize_f16_f16;
- need_quant = (src1->type == HTP_TYPE_F32);
- mmctx->type = (src1->type == HTP_TYPE_F32) ? "f32-f16" : "f16-f16";
+ quant_task_func = (act->type == HTP_TYPE_F32) ? quantize_f32_f16 : quantize_f16_f16;
+ need_quant = (act->type == HTP_TYPE_F32);
+ mmctx->type = (act->type == HTP_TYPE_F32) ? "f32-f16" : "f16-f16";
mmctx->vec_dot_1x1 = vec_dot_f16_f16_aa_1x1;
mmctx->vec_dot_2x1 = vec_dot_f16_f16_aa_2x1;
mmctx->vec_dot_2x2 = vec_dot_f16_f16_aa_2x2;
- src1_row_size = hex_round_up(ne10 * 2, 128);
+ act_row_size = kparams->act_row_size;
break;
case HTP_MM_KERNEL_HVX_F32_F32_VTCM:
@@ -1809,7 +1805,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
mmctx->vec_dot_1x1 = vec_dot_f32_f32_aa_1x1;
mmctx->vec_dot_2x1 = vec_dot_f32_f32_aa_2x1;
mmctx->vec_dot_2x2 = vec_dot_f32_f32_aa_2x2;
- src1_row_size = hex_round_up(ne10 * 4, 128);
+ act_row_size = kparams->act_row_size;
break;
case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
@@ -1821,9 +1817,9 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
const uint32_t qk = QK_Q8_0_TILED;
const uint32_t nb = (ne10 + qk - 1) / qk;
- const uint32_t total_nb = src1_nrows * nb;
+ const uint32_t total_nb = act_nrows * nb;
- if (src1_nrows < octx->n_threads && !is_batched) {
+ if (act_nrows < octx->n_threads && !is_batched) {
n_quant_tasks = MIN(total_nb, octx->n_threads);
quant_task_func = htp_mm_act_quant_block_func(src0->type);
for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) {
@@ -1835,28 +1831,28 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
mmctx->quant_c[ith] = ib_first % nb;
}
} else {
- n_quant_tasks = MIN(src1_nrows, octx->n_threads);
+ n_quant_tasks = MIN(act_nrows, octx->n_threads);
quant_task_func = htp_mm_act_quant_row_func(src0->type);
}
- src1_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ act_row_size = kparams->act_row_size;
break;
}
- const uint32_t m_chunk = (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows)
- ? (uint32_t) kparams->m_chunk : src1_nrows;
+ const uint32_t m_chunk = (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < act_nrows)
+ ? (uint32_t) kparams->m_chunk : act_nrows;
const uint32_t m_layout_rows = m_chunk;
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, m_layout_rows, octx->n_threads,
- dst_row_size, src0_row_size, src1_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false);
+ dst_row_size, src0_row_size, act_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false);
if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM ||
kparams->kernel_type == HTP_MM_KERNEL_HVX_F32_F32_VTCM ||
kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW ||
kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
- mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+ mmctx->vtcm_act_size_per_thread = L.act_bytes;
} else {
- mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->n_threads_div);
+ mmctx->vtcm_act_size_per_thread = fastdiv(L.act_bytes, &octx->n_threads_div);
}
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
@@ -1864,12 +1860,12 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
const size_t vtcm_size = L.total_bytes;
- FARF(HIGH, "matmul-%s : src0-vtcm-size %zu src1-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
- L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
+ FARF(HIGH, "matmul-%s : src0-vtcm-size %zu act-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
+ L.src0_bytes, L.act_bytes, L.dst_bytes, vtcm_size);
FARF(HIGH, "matmul-%s : %ux%ux%ux%u * %ux%ux%ux%u-> %ux%ux%ux%u (0x%p, 0x%p, 0x%p)\n", mmctx->type, src0->ne[0],
- src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
- dst->ne[1], dst->ne[2], dst->ne[3], src0->data, src1->data, dst->data);
+ src0->ne[1], src0->ne[2], src0->ne[3], act->ne[0], act->ne[1], act->ne[2], act->ne[3], dst->ne[0],
+ dst->ne[1], dst->ne[2], dst->ne[3], src0->data, act->data, dst->data);
if (octx->ctx->vtcm_size < vtcm_size) {
FARF(ERROR, "matmul-%s : current VTCM reservation %zu is too small, needed %zu\n", mmctx->type,
@@ -1878,18 +1874,14 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
}
uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
- mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+ mmctx->vtcm_act = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
- mmctx->vtcm_src2 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2);
+ mmctx->vtcm_bias = VTCM_LAYOUT_PTR(uint8_t, base, L.off_bias);
mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
mmctx->vtcm_act_raw = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);
- octx->src1_spad.src = NULL;
- octx->src0_spad.src = NULL;
- octx->dst_spad.src = NULL;
-
mmctx->vtcm_src0_stride = src0_row_size_padded;
- mmctx->vtcm_src1_stride = src1_row_size;
+ mmctx->vtcm_act_stride = act_row_size;
if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW) {
mmctx->vtcm_act_raw_stride = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
} else if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM) {
@@ -1900,14 +1892,14 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
- if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < src1_nrows) {
- for (uint32_t m_start = 0; m_start < src1_nrows; m_start += m_chunk) {
- const uint32_t cur_m_rows = MIN(src1_nrows - m_start, m_chunk);
+ if (kparams->m_chunk > 0 && (uint32_t) kparams->m_chunk < act_nrows) {
+ for (uint32_t m_start = 0; m_start < act_nrows; m_start += m_chunk) {
+ const uint32_t cur_m_rows = MIN(act_nrows - m_start, m_chunk);
mmctx->cur_m_start = m_start;
mmctx->cur_m_rows = cur_m_rows;
if (need_quant) {
- hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, m_start, cur_m_rows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, m_start, cur_m_rows);
const uint32_t qk = QK_Q8_0_TILED;
const uint32_t nb = (ne10 + qk - 1) / qk;
@@ -1933,22 +1925,22 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
mmctx->n_quant_tasks = quant_tasks;
work_queue_run(octx->ctx->work_queue, q_func, mmctx, quant_tasks);
} else {
- hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, m_start, cur_m_rows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act, mmctx->vtcm_act_stride, m_start, cur_m_rows);
}
work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, octx->n_threads);
}
} else {
mmctx->cur_m_start = 0;
- mmctx->cur_m_rows = src1_nrows;
+ mmctx->cur_m_rows = act_nrows;
if (need_quant) {
- hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, src1_nrows);
- mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+ mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->n_quant_tasks = n_quant_tasks;
work_queue_run(octx->ctx->work_queue, quant_task_func, mmctx, n_quant_tasks);
} else {
- hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_src1, mmctx->vtcm_src1_stride, 0, src1_nrows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act, mmctx->vtcm_act_stride, 0, act_nrows);
}
work_queue_run(octx->ctx->work_queue, matmul_job_func, mmctx, octx->n_threads);
@@ -1964,11 +1956,11 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
const uint32_t n_weights = kparams->n_weights;
const struct htp_tensor * restrict act = octx->src[n_weights];
- const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
- const size_t src1_stride = mmctx->vtcm_src1_stride;
+ const uint32_t act_nrows = act->ne[1] * act->ne[2] * act->ne[3];
+ const size_t act_stride = mmctx->vtcm_act_stride;
uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith;
- uint8_t * restrict src1_data = mmctx->vtcm_src1;
+ uint8_t * restrict act_data = mmctx->vtcm_act;
const uint32_t n_prefetch = kparams->n_prefetch;
assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
@@ -2020,17 +2012,17 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
uint32_t ir1 = 0;
- for (; ir1 + 1 < src1_nrows; ir1 += 2) {
- const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);
- const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride);
+ for (; ir1 + 1 < act_nrows; ir1 += 2) {
+ const uint8_t * restrict act_col0 = (const uint8_t *) (act_data + (ir1+0) * act_stride);
+ const uint8_t * restrict act_col1 = (const uint8_t *) (act_data + (ir1+1) * act_stride);
float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size));
float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size));
- mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
+ mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, act_col0, act_col1);
}
- for (; ir1 < src1_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+ for (; ir1 < act_nrows; ++ir1) {
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));
- mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col);
+ mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, act_col);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
@@ -2049,10 +2041,10 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
src0_stride, src0_row_size, src0_row_size, 1);
const uint8_t * ss0 = (void *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
- for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) {
- const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
+ for (uint32_t ir1 = 0; ir1 < act_nrows; ++ir1) {
+ const uint8_t * restrict act_col = (const uint8_t *) (act_data + ir1 * act_stride);
float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));
- mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col);
+ mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, act_col);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
}
@@ -2144,6 +2136,7 @@ typedef struct {
size_t vtcm_f32_act_bytes_per_thread;
uint32_t dma_step_rows;
uint32_t dma_step_rows_shift;
+ uint32_t act_elem_size; // 4=F32 src1, 2=F16 src1
} activation_transfer_task_state_t;
typedef struct {
@@ -2262,7 +2255,7 @@ static void transfer_activation_chunk_col_chunk_worker_fn(unsigned int n, unsign
);
}
-static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
+static void transfer_activation_chunk_to_fp16_dma_pipelined(
dma_queue *dma_q,
__fp16 *restrict vtcm_dst,
dma_addr_t act_dma_addr,
@@ -2270,7 +2263,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
uint32_t k_block,
uint32_t k_stride,
uint32_t k_valid,
- float *thread_f32_act,
+ uint8_t *thread_act,
+ uint32_t act_elem_size,
struct htp_thread_trace *tr,
uint32_t dma_step_rows,
uint32_t dma_step_rows_shift) {
@@ -2280,38 +2274,56 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
const uint32_t n_steps = n_rows_padded >> dma_step_rows_shift;
+ const bool act_is_f16 = (act_elem_size == sizeof(__fp16));
+ const size_t row_bytes = (size_t) k_block * act_elem_size; // staging row (DMA dst stride)
+ const size_t src_row_bytes = (size_t) k_stride * act_elem_size; // DDR row (DMA src stride)
+ const size_t width_bytes = (size_t) k_valid * act_elem_size;
+
// Push step 0
if (n_steps > 0 && n_rows > 0) {
uint32_t nrows_to_fetch = hex_smin(n_rows, R);
- dma_queue_push(dma_q, dma_make_data(thread_f32_act, act_dma_addr),
- k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+ dma_queue_push(dma_q, dma_make_data(thread_act, act_dma_addr),
+ row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
}
// Push step 1 (if valid)
if (n_steps > 1) {
uint32_t next_r = R * 1;
if (next_r < n_rows) {
uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
- float *next_buf = thread_f32_act + 1 * R * k_block;
- dma_queue_push(dma_q, dma_make_data(next_buf, act_dma_addr + (size_t) next_r * k_stride * sizeof(float)),
- k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+ uint8_t *next_buf = thread_act + 1 * R * row_bytes;
+ dma_queue_push(dma_q, dma_make_data(next_buf, act_dma_addr + (size_t) next_r * src_row_bytes),
+ row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
}
}
for (uint32_t s = 0; s < n_steps; ++s) {
uint32_t r = s << dma_step_rows_shift;
- float *curr_buf = thread_f32_act;
+ uint8_t *curr_buf = thread_act;
if (r < n_rows) {
- curr_buf = (float *) dma_queue_pop(dma_q).dst;
+ curr_buf = (uint8_t *) dma_queue_pop(dma_q).dst;
}
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, r);
- for (uint32_t p = 0; p < (R >> 1); ++p) {
- uint32_t row_idx = r + (p << 1);
- float *pair_buf = curr_buf + (p << 1) * k_block;
- bool r0_valid = ((row_idx + 0) < n_rows);
- bool r1_valid = ((row_idx + 1) < n_rows);
+ // Two copies of the pair loop so the type is resolved once per step and the
+ // row-pair kernels stay direct (inlinable) calls.
+ if (act_is_f16) {
+ for (uint32_t p = 0; p < (R >> 1); ++p) {
+ uint32_t row_idx = r + (p << 1);
+ const __fp16 *pair_buf = (const __fp16 *) (curr_buf + (p << 1) * row_bytes);
+ bool r0_valid = ((row_idx + 0) < n_rows);
+ bool r1_valid = ((row_idx + 1) < n_rows);
+
+ transfer_activation_row_pair_f16_to_f16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+ }
+ } else {
+ for (uint32_t p = 0; p < (R >> 1); ++p) {
+ uint32_t row_idx = r + (p << 1);
+ const float *pair_buf = (const float *) (curr_buf + (p << 1) * row_bytes);
+ bool r0_valid = ((row_idx + 0) < n_rows);
+ bool r1_valid = ((row_idx + 1) < n_rows);
- transfer_activation_row_pair_fp32_to_fp16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+ transfer_activation_row_pair_fp32_to_fp16(vtcm_dst, pair_buf, pair_buf + k_block, row_idx, k_block, k_valid, r0_valid, r1_valid);
+ }
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_A_PREP, r);
@@ -2320,8 +2332,8 @@ static void transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
uint32_t next_r = next_s << dma_step_rows_shift;
if (next_r < n_rows) {
uint32_t nrows_to_fetch = hex_smin(n_rows - next_r, R);
- dma_queue_push(dma_q, dma_make_data(curr_buf, act_dma_addr + (size_t) next_r * k_stride * sizeof(float)),
- k_block * sizeof(float), k_stride * sizeof(float), k_valid * sizeof(float), nrows_to_fetch);
+ dma_queue_push(dma_q, dma_make_data(curr_buf, act_dma_addr + (size_t) next_r * src_row_bytes),
+ row_bytes, src_row_bytes, width_bytes, nrows_to_fetch);
}
}
}
@@ -2336,11 +2348,11 @@ static void transfer_activation_chunk_worker_fn(unsigned int n, unsigned int i,
size_t chunk_size = hex_smin(st->n_tot_chunks - chunk_idx, st->n_chunks_per_task);
__fp16 *dst = st->dst + chunk_idx * st->k_block;
- const dma_addr_t act_dma_addr = st->act_dma_addr + (size_t) chunk_idx * st->k_stride * sizeof(float);
+ const dma_addr_t act_dma_addr = st->act_dma_addr + (size_t) chunk_idx * st->k_stride * st->act_elem_size;
- float *thread_f32_act = (float *)((char *)st->vtcm_f32_act + i * st->vtcm_f32_act_bytes_per_thread);
- transfer_activation_chunk_fp32_to_fp16_dma_pipelined(
- st->ctx->dma[i], dst, act_dma_addr, chunk_size, st->k_block, st->k_stride, st->k_valid, thread_f32_act, tr, st->dma_step_rows, st->dma_step_rows_shift
+ uint8_t *thread_act = (uint8_t *) st->vtcm_f32_act + i * st->vtcm_f32_act_bytes_per_thread;
+ transfer_activation_chunk_to_fp16_dma_pipelined(
+ st->ctx->dma[i], dst, act_dma_addr, chunk_size, st->k_block, st->k_stride, st->k_valid, thread_act, st->act_elem_size, tr, st->dma_step_rows, st->dma_step_rows_shift
);
}
}
@@ -2446,10 +2458,9 @@ static void dequantize_tiled_weight_chunk_to_fp16_tiles(
int n_k_tiles, struct fastdiv_values n_k_tiles_div,
worker_callback_t dequant_worker_fn, int n_threads) {
- assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
assert(k_block % HTP_MM_HMX_TILE_N_COLS == 0);
- size_t n_col_tiles = n_cols / HTP_MM_HMX_TILE_N_COLS;
+ size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS);
size_t n_tot_tiles = n_col_tiles * n_k_tiles;
size_t n_tiles_per_task = (n_threads == 1) ? n_tot_tiles : hmx_ceil_div(n_tot_tiles, n_threads);
@@ -2499,7 +2510,9 @@ static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned i
output_transfer_col_chunk_state_t *st = (output_transfer_col_chunk_state_t *) data;
struct htp_thread_trace * tr = &st->traces[i];
- uint32_t n_blocks = st->n_cols / 32;
+ // Round up: the last block is partial when N is not 32-aligned. Its pad
+ // columns are dropped by the dst_cols clamp inside the store.
+ uint32_t n_blocks = hmx_ceil_div(st->n_cols, 32);
uint32_t b_first = fastdiv(n_blocks * i, &st->n_threads_div);
uint32_t b_last = fastdiv(n_blocks * (i + 1), &st->n_threads_div);
uint32_t c_first = b_first * 32;
@@ -2527,11 +2540,9 @@ static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned i
static void transfer_output_chunk_threaded(struct htp_context *ctx, float *dst, const float *src2, const __fp16 *vtcm_src,
int n_rows, int n_cols, int dst_stride, uint32_t src2_stride, int dst_cols, int n_threads) {
- assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
-
if (n_rows <= 0) return;
- uint32_t n_blocks = (uint32_t)n_cols / 32;
+ uint32_t n_blocks = hmx_ceil_div((uint32_t) n_cols, 32);
if (n_threads > 1 && n_blocks >= (uint32_t)n_threads) {
struct fastdiv_values n_threads_div = (n_threads == (int)ctx->n_threads) ? ctx->n_threads_div : init_fastdiv_values(n_threads);
output_transfer_col_chunk_state_t col_state;
@@ -2590,6 +2601,7 @@ struct activation_transfer_params {
int k_valid;
float * vtcm_f32_act;
size_t vtcm_f32_act_bytes;
+ uint32_t act_elem_size; // 4=F32 src1, 2=F16 src1
};
static void transfer_activation_chunk_threaded(const struct activation_transfer_params * params) {
@@ -2605,13 +2617,15 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
int k_valid = params->k_valid;
float * vtcm_f32_act = params->vtcm_f32_act;
size_t vtcm_f32_act_bytes = params->vtcm_f32_act_bytes;
+ // element size of the activation rows (4 = F32, 2 = F16).
+ const uint32_t act_elem_size = params->act_elem_size ? params->act_elem_size : (uint32_t) sizeof(float);
if (n_rows <= 0) {
return;
}
const size_t n_tasks = (n_rows + 31) >> 5;
- if (n_threads > 1 && k_block > 32 && n_tasks < (size_t) n_threads) {
+ if (act_elem_size == sizeof(float) && n_threads > 1 && k_block > 32 && n_tasks < (size_t) n_threads) {
// Calculate step rows parameters for column-chunked dma pipelining
uint32_t dma_step_rows = 2;
uint32_t dma_step_rows_shift = 1;
@@ -2662,6 +2676,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_
state.traces = ctx->trace;
state.ctx = ctx;
state.vtcm_f32_act = vtcm_f32_act;
+ state.act_elem_size = act_elem_size;
state.vtcm_f32_act_bytes_per_thread = hex_align_down(fastdiv(vtcm_f32_act_bytes, act_threads_div), 128);
@@ -2729,6 +2744,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
dma_addr_t weight,
int m, int k, int n,
int act_stride,
+ uint32_t act_elem_size, // 4=F32 src1, 2=F16 src1
int weight_stride,
int weight_type,
int k_valid,
@@ -2748,7 +2764,12 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
struct htp_thread_trace * tr = &ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
- if (k % 32 != 0 || n % 32 != 0) { return -1; }
+ // Quantized weights are repacked and padded to 32, so we it has to be 32-aligned.
+ // Only F16/F32 weights can be non-32-aligned, they will be padded in the following kernel.
+ const bool wtype_is_quant = (weight_type != HTP_TYPE_F16 && weight_type != HTP_TYPE_F32);
+ if (k % 32 != 0 || (wtype_is_quant && n % 32 != 0)) {
+ return -1;
+ }
if (!hex_is_aligned(dst, VLEN) || (act_dma_addr & (VLEN - 1)) != 0) { return -1; }
size_t row_stride = htp_mm_get_tiled_row_stride(weight_type, k);
@@ -2817,10 +2838,10 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00)); // scale: 1.0, bias: 0.0 in FP16
- const bool has_src2 = (src2_bytes > 0 && src2_addr != 0);
- float *vtcm_src2 = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_src2, has_src2);
- if (has_src2) {
- dma_queue_push(weight_dma, dma_make_data(vtcm_src2, src2_addr), hex_align_up(src2_bytes, 128), 0, src2_bytes, 1);
+ const bool has_bias = (src2_bytes > 0 && src2_addr != 0);
+ float *vtcm_bias = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_bias, has_bias);
+ if (has_bias) {
+ dma_queue_push(weight_dma, dma_make_data(vtcm_bias, src2_addr), hex_align_up(src2_bytes, 128), 0, src2_bytes, 1);
dma_queue_pop(weight_dma);
}
@@ -2844,7 +2865,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
struct activation_transfer_params act_params = {
.ctx = ctx,
.dst = vtcm_f16_act,
- .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+ .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size, // byte offset (F16/F32)
.n_rows = (int) n_rows,
.k_block = k,
.k_stride = act_stride,
@@ -2854,18 +2875,19 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
.k_valid = k_valid,
.vtcm_f32_act = vtcm_f32_act,
.vtcm_f32_act_bytes = L.act_f32_bytes,
+ .act_elem_size = act_elem_size,
};
transfer_activation_chunk_threaded(&act_params);
// Prologue: push A0 and optionally A1 (if n_chunk_cnt > 1)
const size_t n_cols_A0 = hex_smin(n - 0 * n_chunk_n_cols, n_chunk_n_cols);
- const uint32_t height_A0 = is_quant ? (n_cols_A0 / 32) * n_k_tiles : n_cols_A0;
+ const uint32_t height_A0 = is_quant ? hmx_ceil_div(n_cols_A0, 32) * n_k_tiles : n_cols_A0;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_A0);
if (1 < n_chunk_cnt) {
const size_t n_cols_A1 = hex_smin(n - 1 * n_chunk_n_cols, n_chunk_n_cols);
- const uint32_t height_A1 = is_quant ? (n_cols_A1 / 32) * n_k_tiles : n_cols_A1;
+ const uint32_t height_A1 = is_quant ? hmx_ceil_div(n_cols_A1, 32) * n_k_tiles : n_cols_A1;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_A1);
}
@@ -2889,7 +2911,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
// 3. push A_{i+2} (if i+2 < n_chunk_cnt)
if (i + 2 < n_chunk_cnt) {
- const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2;
+ const uint32_t height_p2 = is_quant ? hmx_ceil_div(n_cols_p2, 32) * n_k_tiles : n_cols_p2;
dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_p2 * weight_stride),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_p2);
}
@@ -2907,10 +2929,10 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
const size_t nc_prev = (i - 1) * n_chunk_n_cols;
const size_t n_cols_prev = hex_smin(n - nc_prev, n_chunk_n_cols);
float *output_chunk = dst + (mr * dst_stride + nc_prev);
- const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_prev) : NULL;
+ const float *bias_chunk = has_bias ? (vtcm_bias + mr * src2_stride + nc_prev) : NULL;
int chunk_dst_cols = dst_cols - (int)nc_prev;
if (chunk_dst_cols > 0) {
- transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output_bufs[(i - 1) % 2], n_rows, n_cols_prev, dst_stride, src2_stride, chunk_dst_cols, n_threads);
+ transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output_bufs[(i - 1) % 2], n_rows, n_cols_prev, dst_stride, src2_stride, chunk_dst_cols, n_threads);
}
}
}
@@ -2920,10 +2942,10 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
const size_t nc_last = (n_chunk_cnt - 1) * n_chunk_n_cols;
const size_t n_cols_last = hex_smin(n - nc_last, n_chunk_n_cols);
float *output_chunk = dst + (mr * dst_stride + nc_last);
- const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc_last) : NULL;
+ const float *bias_chunk = has_bias ? (vtcm_bias + mr * src2_stride + nc_last) : NULL;
int chunk_dst_cols = dst_cols - (int)nc_last;
if (chunk_dst_cols > 0) {
- transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output_bufs[(n_chunk_cnt - 1) % 2], n_rows, n_cols_last, dst_stride, src2_stride, chunk_dst_cols, n_threads);
+ transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output_bufs[(n_chunk_cnt - 1) % 2], n_rows, n_cols_last, dst_stride, src2_stride, chunk_dst_cols, n_threads);
}
}
} else {
@@ -2935,7 +2957,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
struct activation_transfer_params act_params = {
.ctx = ctx,
.dst = vtcm_f16_act,
- .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+ .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size, // byte offset (F16/F32)
.n_rows = (int) n_rows,
.k_block = k,
.k_stride = act_stride,
@@ -2945,13 +2967,14 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
.k_valid = k_valid,
.vtcm_f32_act = vtcm_f32_act,
.vtcm_f32_act_bytes = L.act_f32_bytes,
+ .act_elem_size = act_elem_size,
};
transfer_activation_chunk_threaded(&act_params);
// A0: Pre-fetch the first weight chunk (nc = 0)
if (n > 0) {
const size_t n_cols = hex_smin(n, n_chunk_n_cols);
- const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
+ const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
}
@@ -2973,7 +2996,7 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
const size_t nc_next = nc + n_chunk_n_cols;
if (nc_next < n) {
const size_t n_cols_next = hex_smin(n - nc_next, n_chunk_n_cols);
- const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
+ const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_next;
dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
}
@@ -2984,10 +3007,10 @@ static int hmx_mm_2d_f32(struct htp_context *ctx,
// D: Output Store
float *output_chunk = dst + (mr * dst_stride + nc);
- const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * src2_stride + nc) : NULL;
+ const float *bias_chunk = has_bias ? (vtcm_bias + mr * src2_stride + nc) : NULL;
int chunk_dst_cols = dst_cols - (int)nc;
if (chunk_dst_cols > 0) {
- transfer_output_chunk_threaded(ctx, output_chunk, src2_chunk, vtcm_output, n_rows, n_cols, dst_stride, src2_stride, chunk_dst_cols, n_threads);
+ transfer_output_chunk_threaded(ctx, output_chunk, bias_chunk, vtcm_output, n_rows, n_cols, dst_stride, src2_stride, chunk_dst_cols, n_threads);
}
}
}
@@ -3013,7 +3036,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
const int k = (int) act->ne[0];
const int k_valid = (int) act->ne[0];
const int m = (int) (act->ne[1] * act->ne[2] * act->ne[3]);
- const int act_stride = (int) (act->nb[1] / sizeof(float));
+ const uint32_t act_elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+ const int act_stride = (int) (act->nb[1] / act_elem_size);
const dma_addr_t act_dma_addr = act->data;
if (k % 32 != 0) { return HTP_STATUS_NO_SUPPORT; }
@@ -3089,7 +3113,16 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
int m_start = 0;
int m_rows = m;
if (octx->ctx->mdev.count > 1) {
- const bool can_split = htp_tensor_can_row_partition(octx->dsts[0], sizeof(float));
+ bool can_split = htp_tensor_can_row_partition(act, act_elem_size);
+ if (can_split) {
+ for (uint32_t p = 0; p < n_weights; ++p) {
+ const struct htp_tensor * restrict dst = octx->dsts[p];
+ if (!htp_tensor_can_row_partition(dst, sizeof(float))) {
+ can_split = false;
+ break;
+ }
+ }
+ }
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_rows = (int) range.count;
@@ -3118,7 +3151,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
struct activation_transfer_params act_params = {
.ctx = ctx,
.dst = vtcm_f16_act,
- .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+ .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,
.n_rows = (int) n_rows,
.k_block = k,
.k_stride = act_stride,
@@ -3128,6 +3161,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
.k_valid = k_valid,
.vtcm_f32_act = vtcm_f32_act,
.vtcm_f32_act_bytes = L.act_f32_bytes,
+ .act_elem_size = act_elem_size,
};
transfer_activation_chunk_threaded(&act_params);
@@ -3149,13 +3183,13 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
const uint32_t dma_src_stride = is_quant ? tile_size : weight_stride;
const size_t n_cols_A0 = hex_smin(n - 0 * n_chunk_n_cols, n_chunk_n_cols);
- const uint32_t height_A0 = is_quant ? (n_cols_A0 / 32) * n_k_tiles : n_cols_A0;
+ const uint32_t height_A0 = is_quant ? hmx_ceil_div(n_cols_A0, 32) * n_k_tiles : n_cols_A0;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_A0);
if (1 < n_chunk_cnt) {
const size_t n_cols_A1 = hex_smin(n - 1 * n_chunk_n_cols, n_chunk_n_cols);
- const uint32_t height_A1 = is_quant ? (n_cols_A1 / 32) * n_k_tiles : n_cols_A1;
+ const uint32_t height_A1 = is_quant ? hmx_ceil_div(n_cols_A1, 32) * n_k_tiles : n_cols_A1;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_A1);
}
@@ -3175,7 +3209,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
n_k_tiles, n_k_tiles_div, dequant_worker_fn, n_threads);
if (i + 2 < n_chunk_cnt) {
- const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2;
+ const uint32_t height_p2 = is_quant ? hmx_ceil_div(n_cols_p2, 32) * n_k_tiles : n_cols_p2;
dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_p2 * weight_stride),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_p2);
}
@@ -3216,7 +3250,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
struct activation_transfer_params act_params = {
.ctx = ctx,
.dst = vtcm_f16_act,
- .act_dma_addr = act_dma_addr + mr * act_stride * sizeof(float),
+ .act_dma_addr = act_dma_addr + (size_t) mr * act_stride * act_elem_size,
.n_rows = (int) n_rows,
.k_block = k,
.k_stride = act_stride,
@@ -3226,6 +3260,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
.k_valid = k_valid,
.vtcm_f32_act = vtcm_f32_act,
.vtcm_f32_act_bytes = L.act_f32_bytes,
+ .act_elem_size = act_elem_size,
};
transfer_activation_chunk_threaded(&act_params);
@@ -3247,7 +3282,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
if (n > 0) {
const size_t n_cols = hex_smin(n, n_chunk_n_cols);
- const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
+ const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height);
}
@@ -3266,7 +3301,7 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
const size_t nc_next = nc + n_chunk_n_cols;
if (nc_next < n) {
const size_t n_cols_next = hex_smin(n - nc_next, n_chunk_n_cols);
- const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
+ const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_next;
dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
}
@@ -3314,16 +3349,16 @@ static int hmx_mm_f16_f32_batched_simple(struct htp_context *ctx,
int ret = 0;
for (int b3 = 0; b3 < params->ne13 && ret == 0; ++b3) {
for (int b2 = 0; b2 < params->ne12 && ret == 0; ++b2) {
- dma_addr_t cur_src2_addr = params->src2_addr ? (params->src2_addr +
- b2 * params->src2_nb2 +
- b3 * params->src2_nb3) : 0;
+ dma_addr_t cur_bias_addr = params->bias_addr ? (params->bias_addr +
+ b2 * params->bias_nb2 +
+ b3 * params->bias_nb3) : 0;
ret = hmx_mm_2d_f32(ctx, params->weight_dma, hmx_mm_dst_batch_ptr(params, b2, b3),
- cur_src2_addr, params->src2_bytes,
+ cur_bias_addr, params->bias_bytes,
hmx_mm_act_batch_addr(params, b2, b3),
hmx_mm_weight_batch_data(params, b2, b3),
params->m, params->k, params->n,
- params->act_stride, params->weight_stride * (int)sizeof(__fp16),
- HTP_TYPE_F16, params->k, params->dst_stride, params->src2_stride, params->n,
+ params->act_stride, params->act_elem_size, params->weight_stride * (int)sizeof(__fp16),
+ HTP_TYPE_F16, params->k, params->dst_stride, params->bias_stride, params->n,
m_chunk, n_chunk, pipeline, n_threads, act_threads,
act_threads_div, k_div, 0, 0, vtcm_size);
}
@@ -3339,7 +3374,8 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
if (params->act_stride < params->k || params->weight_stride < params->k || params->dst_stride < params->n) { return -1; }
if (params->ne02 <= 0 || params->ne03 <= 0 || params->ne12 <= 0 || params->ne13 <= 0) { return -1; }
if (params->ne12 % params->ne02 != 0 || params->ne13 % params->ne03 != 0) { return -1; }
- if (params->k % 32 != 0 || params->n % 32 != 0) { return -1; }
+ // N (the weight row count) does not have to be 32-aligned:
+ if (params->k % 32 != 0) { return -1; }
if (!hex_is_aligned(params->dst, VLEN) || (params->act_dma_addr & (VLEN - 1)) != 0) { return -1; }
const int group_size = params->r2;
@@ -3363,7 +3399,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
size_t vtcm_used = vtcm_size;
struct htp_mm_hmx_vtcm_layout L;
- htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, HTP_TYPE_F16, params->k, m_chunk_n_rows, n_chunk_n_cols, group_size, false, act_threads, 0, params->src2_bytes);
+ htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, HTP_TYPE_F16, params->k, m_chunk_n_rows, n_chunk_n_cols, group_size, false, act_threads, 0, params->bias_bytes);
if (L.total_bytes > vtcm_budget) {
FARF(HIGH, "%s: grouped layout overflowed VTCM, falling back to simple batched loop", __func__);
@@ -3380,10 +3416,10 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
__fp16 *vtcm_scales = VTCM_LAYOUT_PTR(__fp16, base, L.off_scales);
float *vtcm_f32_act = VTCM_LAYOUT_PTR(float, base, L.off_act_f32);
- const bool has_src2 = (params->src2_bytes > 0 && params->src2_addr != 0);
- float *vtcm_src2 = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_src2, has_src2);
- if (has_src2) {
- dma_queue_push(params->weight_dma, dma_make_data(vtcm_src2, params->src2_addr), hex_align_up(params->src2_bytes, 128), 0, params->src2_bytes, 1);
+ const bool has_bias = (params->bias_bytes > 0 && params->bias_addr != 0);
+ float *vtcm_bias = VTCM_LAYOUT_PTR_OPTIONAL(float, base, L.off_bias, has_bias);
+ if (has_bias) {
+ dma_queue_push(params->weight_dma, dma_make_data(vtcm_bias, params->bias_addr), hex_align_up(params->bias_bytes, 128), 0, params->bias_bytes, 1);
dma_queue_pop(params->weight_dma);
}
@@ -3417,7 +3453,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
// thrashing from HVX loads at large strides.
for (int g = 0; g < group_size; ++g) {
const dma_addr_t act_dma_addr = hmx_mm_act_batch_addr(params, b2_base + g, b3) +
- mr * params->act_stride * sizeof(float);
+ (size_t) mr * params->act_stride * params->act_elem_size;
__fp16 *vtcm_act_g = vtcm_f16_act + (size_t) g * L.act_head_stride;
struct activation_transfer_params act_params = {
.ctx = ctx,
@@ -3432,6 +3468,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
.k_valid = params->k,
.vtcm_f32_act = vtcm_f32_act,
.vtcm_f32_act_bytes = L.act_f32_bytes,
+ .act_elem_size = params->act_elem_size,
};
transfer_activation_chunk_threaded(&act_params);
}
@@ -3444,7 +3481,8 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
}
if (n_chunk_n_cols < (size_t) params->n) {
const size_t n_cols_second = hex_smin((size_t) params->n - n_chunk_n_cols, n_chunk_n_cols);
- dma_queue_push(weight_dma, dma_make_data(vtcm_scratch1, weight_group + params->weight_stride * sizeof(__fp16)),
+ const dma_addr_t second_weight_chunk = weight_group + n_chunk_n_cols * params->weight_stride * sizeof(__fp16);
+ dma_queue_push(weight_dma, dma_make_data(vtcm_scratch1, second_weight_chunk),
fp16_row_bytes, weight_row_bytes, fp16_row_bytes, n_cols_second);
}
@@ -3455,7 +3493,9 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
{
void * curr_raw = (void *) dma_queue_pop(weight_dma).dst;
- hmx_interleave_rows_to_tiles(vtcm_weight, (const __fp16 *) curr_raw, n_cols, params->k, params->k, 0, n_cols);
+ const size_t n_cols_tiled = hex_align_up(n_cols, HTP_MM_HMX_TILE_N_COLS);
+
+ hmx_interleave_rows_to_tiles(vtcm_weight, (const __fp16 *) curr_raw, (uint32_t) n_cols, params->k, params->k, 0, (uint32_t) n_cols_tiled);
const size_t nc_next = nc + n_chunk_n_cols * 2;
if (nc_next < (size_t) params->n) {
@@ -3478,11 +3518,11 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
{
float *output = hmx_mm_dst_batch_ptr(params, b2_base + g, b3) + mr * params->dst_stride + nc;
- const float *src2_chunk = has_src2 ? (vtcm_src2 + mr * params->src2_stride + nc) : NULL;
+ const float *bias_chunk = has_bias ? (vtcm_bias + mr * params->bias_stride + nc) : NULL;
int chunk_dst_cols = params->n - (int)nc;
if (chunk_dst_cols > 0) {
- transfer_output_chunk_threaded(ctx, output, src2_chunk, vtcm_output, (int) n_rows, (int) n_cols,
- params->dst_stride, params->src2_stride, chunk_dst_cols, n_threads);
+ transfer_output_chunk_threaded(ctx, output, bias_chunk, vtcm_output, (int) n_rows, (int) n_cols,
+ params->dst_stride, params->bias_stride, chunk_dst_cols, n_threads);
}
}
}
@@ -3612,7 +3652,8 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
const int cne1 = m;
- const int m_padded = hex_align_up(m, 32);
+ const int m_core = m_end - m_start;
+ const int m_core_padded = hex_align_up(m_core > 0 ? m_core : 1, 32);
if (k % 32 != 0 || n % 32 != 0) { return -1; }
if (!hex_is_aligned(dst, VLEN) || !hex_is_aligned(activation, VLEN)) { return -1; }
@@ -3666,10 +3707,10 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
const size_t overhead = htp_mm_hmx_get_2d_overhead(/*pipeline=*/false, /*is_matmul_id=*/true);
size_t m_chunk_n_rows = 0, n_chunk_n_cols = 0;
if (htp_mm_hmx_compute_chunks(vtcm_budget, overhead, size_per_n, size_per_m, size_per_mn,
- m_padded, n,
+ m_core_padded, n,
/*m_block_cost=*/(size_t) n * HTP_MM_HMX_COST_W_DEQUANT,
- /*n_block_cost=*/(size_t) m_padded * HTP_MM_HMX_COST_A_CONVERT, &m_chunk_n_rows, &n_chunk_n_cols, &vtcm_used)) {
- FARF(ERROR, "hmx-mm-id-2d: VTCM too small : m %d k %d n %d budget %zu", m_padded, k, n, vtcm_budget);
+ /*n_block_cost=*/(size_t) m_core_padded * HTP_MM_HMX_COST_A_CONVERT, &m_chunk_n_rows, &n_chunk_n_cols, &vtcm_used)) {
+ FARF(ERROR, "hmx-mm-id-2d: VTCM too small : m %d k %d n %d budget %zu", m_core_padded, k, n, vtcm_budget);
return -1;
}
@@ -3709,7 +3750,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
// A0: Pre-fetch the first weight chunk (nc = 0)
if (n > 0) {
const size_t n_cols = hex_smin((size_t) n, n_chunk_n_cols);
- const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols;
+ const uint32_t height = is_quant ? hmx_ceil_div(n_cols, 32) * n_k_tiles : n_cols;
dma_queue_push(weight_dma, dma_make_data(vtcm_weight, weight),
dma_dst_stride, dma_src_stride, dma_width_bytes, height);
}
@@ -3732,7 +3773,7 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
const size_t nc_next = nc + n_chunk_n_cols;
if (nc_next < (size_t) n) {
const size_t n_cols_next = hex_smin((size_t) n - nc_next, n_chunk_n_cols);
- const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next;
+ const uint32_t height_next = is_quant ? hmx_ceil_div(n_cols_next, 32) * n_k_tiles : n_cols_next;
dma_queue_push(weight_dma, dma_make_data(curr_raw, weight + nc_next * weight_stride),
dma_dst_stride, dma_src_stride, dma_width_bytes, height_next);
}
@@ -3759,14 +3800,19 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
int k = (int) src0->ne[0];
int n = (int) src0->ne[1];
- const int m_total = (int) src1->ne[1];
- const int act_stride = (int)(src1->nb[1] / sizeof(float));
+ const int m_total = (int) act->ne[1];
+ const uint32_t act_elem_size = (act->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
+ const int act_stride = (int) (act->nb[1] / act_elem_size);
const int wgt_stride = (int)(src0->nb[1] / sizeof(__fp16));
int m_start = 0;
int m_rows = m_total;
if (octx->ctx->mdev.count > 1) {
- const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+ bool can_split = htp_tensor_can_row_partition(dst, sizeof(float)) &&
+ htp_tensor_can_row_partition(act, act_elem_size);
+ if (src2 && src2->ne[1] > 1 && !htp_tensor_can_row_partition(src2, sizeof(float))) {
+ can_split = false;
+ }
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_total, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_rows = (int) range.count;
@@ -3784,22 +3830,23 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
if (src2) {
src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
src2_addr = src2->data + m_start * src2_stride * sizeof(float);
- src2_bytes = (size_t) kparams->vtcm_src2_size;
+ src2_bytes = (size_t) kparams->vtcm_bias_size;
src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
}
const int dst_stride = (int)(dst->nb[1] / sizeof(float));
float * dst_ptr = (float *) dst->data + m_start * dst_stride;
- const dma_addr_t act_addr = src1->data + m_start * act_stride * sizeof(float);
+ // byte offset, act_stride is in activation elements (F16 or F32)
+ const dma_addr_t act_addr = act->data + (size_t) m_start * act_stride * act_elem_size;
int ret = -1;
const int n_threads = kparams->n_threads;
if (kparams->kernel_type == HTP_MM_KERNEL_HMX_F16_BATCHED) {
hmx_mm_f16_f32_batched_params_t batch_params = {
.dst = dst_ptr,
- .src2_addr = src2_addr,
- .src2_bytes = src2_bytes,
+ .bias_addr = src2_addr,
+ .bias_bytes = src2_bytes,
.act_dma_addr = act_addr,
.weight = src0->data,
.weight_dma = octx->ctx->dma[0],
@@ -3807,21 +3854,22 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
.k = k,
.n = n,
.act_stride = act_stride,
+ .act_elem_size = act_elem_size,
.weight_stride = wgt_stride,
.dst_stride = dst_stride,
- .src2_stride = src2_stride,
+ .bias_stride = src2_stride,
.ne02 = ne02,
.ne03 = ne03,
.ne12 = ne12,
.ne13 = ne13,
.src0_nb2 = src0->nb[2],
.src0_nb3 = src0->nb[3],
- .act_nb2 = src1->nb[2],
- .act_nb3 = src1->nb[3],
+ .act_nb2 = act->nb[2],
+ .act_nb3 = act->nb[3],
.dst_nb2 = dst->nb[2],
.dst_nb3 = dst->nb[3],
- .src2_nb2 = src2_nb2,
- .src2_nb3 = src2_nb3,
+ .bias_nb2 = src2_nb2,
+ .bias_nb3 = src2_nb3,
.r2 = (ne02 > 0) ? (ne12 / ne02) : 1,
.r3 = (ne03 > 0) ? (ne13 / ne03) : 1,
.div_r2 = kparams->div_r2,
@@ -3838,7 +3886,7 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
ret = hmx_mm_2d_f32(
octx->ctx, octx->ctx->dma[0], dst_ptr, src2_addr, src2_bytes,
act_addr, src0->data,
- m_rows, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
+ m_rows, k, n, act_stride, act_elem_size, (int) src0->nb[1], (int) src0->type, (int) act->ne[0],
dst_stride, src2_stride, (int)dst->ne[0],
kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
kparams->n_act_threads,
@@ -3855,7 +3903,17 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
return HTP_STATUS_OK;
}
-int op_matmul(struct htp_ops_context * octx) {
+static inline void htp_mm_tensor_collapse_rows(struct htp_tensor * c, const struct htp_tensor * t, uint32_t stride) {
+ *c = *t;
+ c->ne[1] = t->ne[1] * t->ne[2] * t->ne[3];
+ c->ne[2] = 1;
+ c->ne[3] = 1;
+ c->nb[1] = stride;
+ c->nb[2] = c->nb[1] * c->ne[1];
+ c->nb[3] = c->nb[2];
+}
+
+static int op_matmul_impl(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const int status = htp_mm_init_context(octx, kparams);
@@ -3870,6 +3928,40 @@ int op_matmul(struct htp_ops_context * octx) {
return hvx_mm_matmul(octx);
}
+int op_matmul(struct htp_ops_context * octx) {
+ const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+
+ if (kparams->collapse) {
+ const struct htp_tensor * act = octx->src[1];
+ const struct htp_tensor * dst = octx->dst;
+ const uint32_t s_act = (act->ne[1] > 1) ? act->nb[1] : ((act->ne[2] > 1) ? act->nb[2] : act->nb[3]);
+ const uint32_t sd = (dst->ne[1] > 1) ? dst->nb[1] : ((dst->ne[2] > 1) ? dst->nb[2] : dst->nb[3]);
+ struct htp_tensor act_collapsed, dst_collapsed;
+ htp_mm_tensor_collapse_rows(&act_collapsed, act, s_act);
+ htp_mm_tensor_collapse_rows(&dst_collapsed, dst, sd);
+ octx->src[1] = &act_collapsed;
+ octx->dst = &dst_collapsed;
+
+ struct htp_tensor src2_collapsed;
+ const struct htp_tensor * src2 = octx->src[2];
+ if (src2 && (src2->ne[1] * src2->ne[2] * src2->ne[3] > 1)) {
+ const uint32_t s2 = (src2->ne[1] > 1) ? src2->nb[1] : ((src2->ne[2] > 1) ? src2->nb[2] : src2->nb[3]);
+ htp_mm_tensor_collapse_rows(&src2_collapsed, src2, s2);
+ octx->src[2] = &src2_collapsed;
+ }
+
+ const int status = op_matmul_impl(octx);
+ octx->src[1] = act;
+ octx->dst = dst;
+ if (src2) {
+ octx->src[2] = src2;
+ }
+ return status;
+ }
+
+ return op_matmul_impl(octx);
+}
+
static int hmx_mm_op_matmul_id(
struct htp_ops_context * octx,
struct htp_mm_context * mmctx
@@ -3881,21 +3973,40 @@ static int hmx_mm_op_matmul_id(
const int n_ids = octx->src[2]->ne[0];
const int n_as = ne02;
- for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) {
+ const bool mdev_split = (octx->ctx->mdev.count > 1) && htp_tensor_can_row_partition(dst, sizeof(float));
+ if (octx->ctx->mdev.count > 1 && !mdev_split && octx->ctx->mdev.idx > 0) {
+ return HTP_STATUS_OK;
+ }
+ uint32_t n_active = 0;
+ if (mdev_split) {
+ for (uint32_t a = 0; a < (uint32_t) n_as; ++a) {
+ if (matrix_row_counts[a] > 0) n_active++;
+ }
+ }
+ const bool expert_split = mdev_split && (n_active >= octx->ctx->mdev.count);
+
+ uint32_t target_dev = 0;
+ for (uint32_t cur_a = 0; cur_a < (uint32_t) n_as; ++cur_a) {
const int32_t cne1 = matrix_row_counts[cur_a];
if (cne1 == 0) continue;
const int m_padded = hex_align_up(cne1, 32);
int m_start = 0, m_end = m_padded;
- if (octx->ctx->mdev.count > 1) {
- const bool can_split = htp_tensor_mdev_data_aligned(dst) && (uint32_t) cne1 >= octx->ctx->mdev.count;
+ if (expert_split) {
+ const bool my_expert = (target_dev == octx->ctx->mdev.idx);
+ if (++target_dev == octx->ctx->mdev.count) {
+ target_dev = 0;
+ }
+ if (!my_expert) continue;
+ } else if (mdev_split) {
+ const bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_end = (int) (range.start + range.count);
}
if (m_start >= m_end) continue;
- int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) src1->data,
+ int ret = hmx_mm_id_2d_f32(octx->ctx, octx->ctx->dma[0], (float*) dst->data, (float*) act->data,
src0->data + cur_a * nb02,
cne1, ne00, ne01,
ne10,
@@ -3951,21 +4062,21 @@ static int hvx_mm_matmul_id(
n_quant_tasks = MIN(act_nrows, octx->n_threads);
quant_task_func = htp_mm_act_quant_row_func(src0->type);
}
- size_t src1_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ size_t act_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, act_nrows, octx->n_threads,
- 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
+ 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, true, false);
const size_t vtcm_size = L.total_bytes;
- FARF(HIGH, "matmul-id-%s : src0-spad-size %zu src1-spad-size %zu src2-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
- L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
+ FARF(HIGH, "matmul-id-%s : src0-spad-size %zu act-spad-size %zu bias-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
+ L.src0_bytes, L.act_bytes, L.dst_bytes, vtcm_size);
FARF(HIGH, "matmul-id-%s : %ux%ux%ux%u * %ux%ux%ux%u (%ux%ux%ux%u) -> %ux%ux%ux%u (0x%p, 0x%p, 0x%p)\n", mmctx->type,
- src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3],
+ src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], act->ne[0], act->ne[1], act->ne[2], act->ne[3],
ids->ne[0], ids->ne[1], ids->ne[2], ids->ne[3], dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], src0->data,
- src1->data, dst->data);
+ act->data, dst->data);
// Make sure the reserved vtcm size is sufficient
if (octx->ctx->vtcm_size < vtcm_size) {
@@ -3974,24 +4085,18 @@ static int hvx_mm_matmul_id(
}
uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
- mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+ mmctx->vtcm_act = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
- mmctx->vtcm_src2 = NULL;
+ mmctx->vtcm_bias = NULL;
mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
mmctx->vtcm_act_raw = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);
- octx->src1_spad.src = NULL;
- octx->src0_spad.src = NULL;
- octx->src2_spad.src = NULL;
- octx->dst_spad.src = NULL;
-
mmctx->vtcm_src0_stride = src0_row_size_padded;
- mmctx->vtcm_src1_stride = src1_row_size;
+ mmctx->vtcm_act_stride = act_row_size;
mmctx->vtcm_act_raw_stride = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
- mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
- mmctx->vtcm_src2_size_per_thread = 0;
+ mmctx->vtcm_act_size_per_thread = L.act_bytes;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
mmctx->cur_m_start = 0;
@@ -3999,7 +4104,7 @@ static int hvx_mm_matmul_id(
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
- hvx_mm_transfer_src1_dma(octx, kparams, src1, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->n_quant_tasks = n_quant_tasks;
@@ -4022,18 +4127,37 @@ static int hmx_mm_op_matmul_id_nx(
const struct htp_tensor * restrict act = octx->src[n_weights];
const int n_as = src0->ne[2];
+ bool mdev_split = (octx->ctx->mdev.count > 1);
+ for (uint32_t p = 0; p < n_weights && mdev_split; ++p) {
+ const struct htp_tensor * restrict dst = octx->dsts[p];
+ mdev_split = htp_tensor_can_row_partition(dst, sizeof(float));
+ }
+ if (octx->ctx->mdev.count > 1 && !mdev_split && octx->ctx->mdev.idx > 0) {
+ return HTP_STATUS_OK;
+ }
+ uint32_t n_active = 0;
+ if (mdev_split) {
+ for (uint32_t a = 0; a < (uint32_t) n_as; ++a) {
+ if (matrix_row_counts[a] > 0) n_active++;
+ }
+ }
+ const bool expert_split = mdev_split && (n_active >= octx->ctx->mdev.count);
+
+ uint32_t target_dev = 0;
for (uint32_t cur_a = 0; cur_a < (uint32_t) n_as; ++cur_a) {
const int32_t cne1 = matrix_row_counts[cur_a];
if (cne1 == 0) continue;
const int m_padded = hex_align_up(cne1, 32);
int m_start = 0, m_end = m_padded;
- if (octx->ctx->mdev.count > 1) {
- bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
- for (uint32_t p = 0; p < n_weights && can_split; ++p) {
- const struct htp_tensor * restrict dst = octx->dsts[p];
- can_split = !dst || htp_tensor_mdev_data_aligned(dst);
+ if (expert_split) {
+ const bool my_expert = (target_dev == octx->ctx->mdev.idx);
+ if (++target_dev == octx->ctx->mdev.count) {
+ target_dev = 0;
}
+ if (!my_expert) continue;
+ } else if (mdev_split) {
+ const bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_end = (int) (range.start + range.count);
@@ -4104,11 +4228,11 @@ static int hvx_mm_matmul_id_nx(
n_quant_tasks = MIN(act_nrows, octx->n_threads);
quant_task_func = htp_mm_act_quant_row_func(src0->type);
}
- size_t src1_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]);
+ size_t act_row_size = htp_mm_weight_has_offset(src0->type) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]);
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], act_nrows, octx->n_threads,
- 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
+ 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, true, false);
const size_t vtcm_size = L.total_bytes;
@@ -4120,35 +4244,29 @@ static int hvx_mm_matmul_id_nx(
uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
- mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+ mmctx->vtcm_act = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
mmctx->vtcm_act_raw = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);
- octx->src0_spad.src = NULL;
- octx->src1_spad.src = NULL;
- octx->src2_spad.src = NULL;
- octx->src3_spad.src = NULL;
- octx->dst_spad.src = NULL;
-
mmctx->vtcm_src0_stride = 0;
- mmctx->vtcm_src1_stride = src1_row_size;
+ mmctx->vtcm_act_stride = act_row_size;
mmctx->vtcm_act_raw_stride = hex_round_up(act->ne[0] * sizeof(float), QK_Q8_0_TILED * sizeof(float));
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
- mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+ mmctx->vtcm_act_size_per_thread = L.act_bytes;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
mmctx->cur_m_start = 0;
mmctx->cur_m_rows = act_nrows;
- FARF(HIGH, "matmul-id-nx: src0 %d:%d:%d type %s nrows %u, src1 %d:%d:%d nrows %u, vtcm %zu/%zu, threads %d\n",
+ FARF(HIGH, "matmul-id-nx: src0 %d:%d:%d type %s nrows %u, act %d:%d:%d nrows %u, vtcm %zu/%zu, threads %d\n",
src0->ne[0], src0->ne[1], src0->ne[2], mmctx->type, src0->ne[1],
act->ne[0], act->ne[1], act->ne[2], act_nrows,
L.total_bytes, octx->ctx->vtcm_size, octx->n_threads);
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
- hvx_mm_transfer_src1_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->n_quant_tasks = n_quant_tasks;
@@ -4242,10 +4360,10 @@ int op_matmul_id(struct htp_ops_context * octx) {
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
mmctx->octx = octx;
- mmctx->act = src1;
+ mmctx->act = act;
const struct htp_tensor * restrict ids = octx->src[2];
- if (htp_tensor_is_extended(ids) || htp_tensor_is_extended(src1) || htp_tensor_is_extended(dst)) {
+ if (htp_tensor_is_extended(ids) || htp_tensor_is_extended(act) || htp_tensor_is_extended(dst)) {
return HTP_STATUS_NO_SUPPORT;
}
@@ -4255,7 +4373,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
const uint32_t src0_nrows = ne01; // per expert
- const uint32_t src1_nrows = ne11 * ne12 * ne13;
+ const uint32_t act_nrows = ne11 * ne12 * ne13;
// row groups
const int n_ids = ids->ne[0]; // n_expert_used
@@ -4266,7 +4384,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
uint32_t * matrix_row_counts = (uint32_t *) mapping_buf;
struct mmid_row_mapping * matrix_rows = NULL;
- if (src1_nrows > 1) {
+ if (act_nrows > 1) {
const size_t matrix_row_counts_size = n_as * sizeof(uint32_t);
assert(octx->ctx->ddr_spad_size >= matrix_row_counts_size);
@@ -4300,9 +4418,9 @@ int op_matmul_id(struct htp_ops_context * octx) {
mmctx->mapping_stride = mapping_stride;
mmctx->mm_div_ne11 = kparams->div_ne1;
mmctx->src0_row_size_padded = src0_row_size_padded;
- mmctx->act_nrows = src1_nrows;
+ mmctx->act_nrows = act_nrows;
mmctx->cur_m_start = 0;
- mmctx->cur_m_rows = src1_nrows;
+ mmctx->cur_m_rows = act_nrows;
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
@@ -4334,7 +4452,7 @@ int op_matmul_id(struct htp_ops_context * octx) {
mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) {
- s = hvx_mm_matmul_id(octx, mmctx, src1_nrows > 1 ? hvx_mm_id : hvx_mv_id);
+ s = hvx_mm_matmul_id(octx, mmctx, act_nrows > 1 ? hvx_mm_id : hvx_mv_id);
} else {
s = HTP_STATUS_NO_SUPPORT;
}
@@ -4369,7 +4487,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
for (uint32_t p = 0; p < n_weights; p++) {
- if (octx->dsts[p] && htp_tensor_is_extended(octx->dsts[p])) {
+ if (htp_tensor_is_extended(octx->dsts[p])) {
return HTP_STATUS_NO_SUPPORT;
}
}
@@ -4446,7 +4564,7 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
return s;
}
-int op_matmul_nx(struct htp_ops_context * octx) {
+static int op_matmul_nx_impl(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const int status = htp_mm_init_context(octx, kparams);
@@ -4511,13 +4629,13 @@ int op_matmul_nx(struct htp_ops_context * octx) {
quant_task_func = htp_mm_act_quant_row_func(src0->type);
}
- const size_t src1_row_size = htp_mm_weight_has_offset(src0->type)
+ const size_t act_row_size = htp_mm_weight_has_offset(src0->type)
? htp_mm_q8_1_tiled_row_size(act->ne[0])
: htp_mm_q8_0_tiled_row_size(act->ne[0]);
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], act_nrows, octx->n_threads,
- 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true);
+ 0, src0_row_size, act_row_size, 0, kparams->n_prefetch, false, true);
const size_t vtcm_size = L.total_bytes;
@@ -4529,22 +4647,16 @@ int op_matmul_nx(struct htp_ops_context * octx) {
uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base;
mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0);
- mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1);
+ mmctx->vtcm_act = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act);
mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
mmctx->vtcm_act_raw = VTCM_LAYOUT_PTR(uint8_t, base, L.off_act_raw);
- octx->src0_spad.src = NULL;
- octx->src1_spad.src = NULL;
- octx->src2_spad.src = NULL;
- octx->src3_spad.src = NULL;
- octx->dst_spad.src = NULL;
-
mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded;
- mmctx->vtcm_src1_stride = src1_row_size;
+ mmctx->vtcm_act_stride = act_row_size;
mmctx->vtcm_act_raw_stride = hex_round_up(act->ne[0] * sizeof(float), QK_Q8_0_TILED * sizeof(float));
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
- mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
+ mmctx->vtcm_act_size_per_thread = L.act_bytes;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
// Run fused matmul
@@ -4571,7 +4683,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
- hvx_mm_transfer_src1_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
+ hvx_mm_transfer_act_dma(octx, kparams, act, mmctx->vtcm_act_raw, mmctx->vtcm_act_raw_stride, 0, act_nrows);
mmctx->n_quant_rows_per_thread = (act_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->n_quant_tasks = n_quant_tasks;
@@ -4581,3 +4693,37 @@ int op_matmul_nx(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
+
+int op_matmul_nx(struct htp_ops_context * octx) {
+ const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+
+ if (kparams->collapse) {
+ const uint32_t n_weights = kparams->n_weights;
+ const struct htp_tensor * act = octx->src[n_weights];
+ const uint32_t s1 = (act->ne[1] > 1) ? act->nb[1] : ((act->ne[2] > 1) ? act->nb[2] : act->nb[3]);
+ struct htp_tensor act_collapsed;
+ struct htp_tensor dsts_collapsed[HTP_OP_MAX_OUTPUTS];
+ const struct htp_tensor * orig_dsts[HTP_OP_MAX_OUTPUTS];
+
+ htp_mm_tensor_collapse_rows(&act_collapsed, act, s1);
+ octx->src[n_weights] = &act_collapsed;
+
+ for (uint32_t p = 0; p < n_weights; p++) {
+ orig_dsts[p] = octx->dsts[p];
+ const struct htp_tensor * d = octx->dsts[p];
+ const uint32_t sd = (d->ne[1] > 1) ? d->nb[1] : ((d->ne[2] > 1) ? d->nb[2] : d->nb[3]);
+ htp_mm_tensor_collapse_rows(&dsts_collapsed[p], d, sd);
+ octx->dsts[p] = &dsts_collapsed[p];
+ }
+
+ const int status = op_matmul_nx_impl(octx);
+
+ octx->src[n_weights] = act;
+ for (uint32_t p = 0; p < n_weights; p++) {
+ octx->dsts[p] = orig_dsts[p];
+ }
+ return status;
+ }
+
+ return op_matmul_nx_impl(octx);
+}
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h
index 386cb3049..1d11a4c12 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.h
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h
@@ -87,24 +87,26 @@ enum htp_mm_kernel_type {
// Op-specific struct for precomputed matmul params
struct htp_mm_kernel_params {
- int32_t kernel_type; // enum htp_mm_kernel_type
- int32_t pipeline; // 1 = pipelined execution, 0 = standard
+ uint8_t kernel_type; // enum htp_mm_kernel_type
+ uint8_t pipeline; // 1 = pipelined execution, 0 = standard
+ uint8_t collapse; // 1 = collapse outer dims into 2D, 0 = standard
+ uint8_t n_hmx; // 1 = use HMX, 0 = use HVX
+
+ uint8_t n_threads; // Number of threads to spawn
+ uint8_t n_act_threads; // Number of threads for activation preparation
+ uint8_t n_prefetch; // Prefetch lookahead buffers/rows in VTCM
+ uint8_t n_weights; // Number of weights for fused NX
+
int32_t m_chunk; // Row chunk size (M chunk)
int32_t n_chunk; // Col chunk size (N chunk)
- int32_t n_threads; // Number of threads to spawn
- int32_t n_act_threads; // Number of threads for activation preparation
- int32_t n_hmx; // 1 = use HMX, 0 = use HVX
- int32_t n_prefetch; // Prefetch lookahead buffers/rows in VTCM
int32_t tile_size; // Weight tile size
int32_t aligned_tile_size; // Aligned weight tile size (padded to 128)
- int32_t src1_row_size; // Row size for quantized activation
+ int32_t act_row_size; // Row size for activation scratchpad
int32_t vtcm_size; // Total required scratchpad size in VTCM
int32_t vtcm_src0_size; // src0 scratchpad size in VTCM
- int32_t vtcm_src1_size; // src1 scratchpad size in VTCM
- int32_t vtcm_src2_size; // src2 scratchpad size in VTCM (fused only)
- int32_t vtcm_src3_size; // src3 scratchpad size in VTCM (fused only)
+ int32_t vtcm_act_size; // activation scratchpad size in VTCM
+ int32_t vtcm_bias_size; // bias scratchpad size in VTCM (fused only)
int32_t vtcm_dst_size; // dst scratchpad size in VTCM
- int32_t n_weights; // Number of weights for fused NX
// Precomputed division values
struct fastdiv_values div_ne12_ne1;
@@ -147,6 +149,7 @@ static inline int htp_mm_hmx_compute_chunks(size_t vtcm_total,
const size_t usable = vtcm_total - overhead;
size_t best_cost = SIZE_MAX;
+ size_t best_tail_waste = SIZE_MAX;
size_t best_mn = 0;
size_t best_m = 0, best_n = 0;
@@ -173,12 +176,17 @@ static inline int htp_mm_hmx_compute_chunks(size_t vtcm_total,
size_t mblocks = ((size_t) m + mc - 1) / mc;
size_t nblocks = ((size_t) n + nc - 1) / nc;
size_t cost = mblocks * m_block_cost + nblocks * n_block_cost;
+ size_t rem = n % nc;
+ size_t tail_waste = (rem == 0) ? 0 : (nc - rem);
size_t mn = mc * nc;
- if (cost < best_cost || (cost == best_cost && mn > best_mn)) {
- best_cost = cost;
- best_mn = mn;
- best_m = mc;
- best_n = nc;
+ if (cost < best_cost ||
+ (cost == best_cost && tail_waste < best_tail_waste) ||
+ (cost == best_cost && tail_waste == best_tail_waste && mn > best_mn)) {
+ best_cost = cost;
+ best_tail_waste = tail_waste;
+ best_mn = mn;
+ best_m = mc;
+ best_n = nc;
}
}
@@ -349,7 +357,7 @@ struct htp_mm_hmx_vtcm_layout {
size_t off_dst[2]; // [1] is only used when pipelined
size_t off_scratch[2]; // dequantization scratch pads
size_t off_scales; // HMX scales (256 bytes)
- size_t off_src2; // src2 bias in VTCM
+ size_t off_bias; // bias in VTCM
// Cached sizes of regions for HMX kernel use
size_t weight_area_bytes;
@@ -358,25 +366,23 @@ struct htp_mm_hmx_vtcm_layout {
size_t output_area_bytes;
size_t scratch_bytes[2];
size_t act_head_stride;
- size_t src2_bytes;
+ size_t bias_bytes;
size_t total_bytes;
};
struct htp_mm_hvx_vtcm_layout {
// Byte offsets from vtcm_base for each region
- size_t off_src1; // vtcm_src1 (activation)
+ size_t off_act; // vtcm_act (activation)
size_t off_src0; // vtcm_src0 (weight/Wk)
- size_t off_src2; // vtcm_src2 (Wq / fused only)
- size_t off_src3; // vtcm_src3 (Wv / fused only)
+ size_t off_bias; // vtcm_bias (bias / fused add only)
size_t off_dst; // vtcm_dst (output scratch)
size_t off_act_raw; // vtcm_act_raw (raw activation DMA staging)
// Cached sizes
size_t src0_bytes;
- size_t src1_bytes;
- size_t src2_bytes;
- size_t src3_bytes;
+ size_t act_bytes;
+ size_t bias_bytes;
size_t dst_bytes;
size_t act_raw_bytes;
@@ -394,7 +400,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
bool pipeline,
uint32_t act_threads,
uint32_t aligned_tile_size,
- size_t src2_size
+ size_t bias_size
) {
size_t off = 0;
@@ -411,7 +417,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
size_t off_group_a = 0;
VTCM_LAYOUT_ALLOC(off_group_a, off_act, activation_area_size);
VTCM_LAYOUT_ALLOC(off_group_a, off_scales, HTP_MM_HMX_TILE_SIZE); // Padded to 2K for alignment and future persistent data
- VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE), src2_size > 0);
+ VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE), bias_size > 0);
// Group B: Compute-only buffers (starts at off_group_a)
size_t off_group_b = off_group_a;
@@ -439,7 +445,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
L->scratch_bytes[0] = scratch_area_size;
L->scratch_bytes[1] = scratch_area_size;
L->act_head_stride = act_head_stride;
- L->src2_bytes = src2_size;
+ L->bias_bytes = bias_size;
off = off_group_a + hex_smax(group_b_size, group_c_size);
} else {
@@ -463,7 +469,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
size_t off_group_a = 0;
VTCM_LAYOUT_ALLOC(off_group_a, off_scales, HTP_MM_HMX_TILE_SIZE); // Padded to 2K for alignment and future persistent data
VTCM_LAYOUT_ALLOC(off_group_a, off_act, act_area_size);
- VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE), src2_size > 0);
+ VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE), bias_size > 0);
// Group B: Compute-only buffers (starts at off_group_a)
size_t off_group_b = off_group_a;
@@ -491,7 +497,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
L->scratch_bytes[0] = scratch0_size;
L->scratch_bytes[1] = scratch1_size;
L->act_head_stride = 0;
- L->src2_bytes = src2_size;
+ L->bias_bytes = bias_size;
off = off_group_a + hex_smax(group_b_size, group_c_size);
}
@@ -504,21 +510,20 @@ static inline void htp_mm_hvx_vtcm_layout_build(
int kernel_type,
int wtype,
uint32_t ne10, // k
- uint32_t src1_nrows, // m_total
+ uint32_t act_nrows, // m_total
uint32_t n_threads,
size_t dst_row_size,
size_t src0_row_size,
- size_t src1_row_size,
- size_t src2_row_size,
+ size_t act_row_size,
+ size_t bias_row_size,
uint32_t n_prefetch,
bool is_matmul_id,
bool is_fused_nx
) {
- (void)src1_row_size;
+ (void)act_row_size;
size_t src0_sz = 0;
- size_t src1_sz = 0;
- size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
- size_t src3_sz = 0;
+ size_t act_sz = 0;
+ size_t bias_sz = bias_row_size > 0 ? htp_mm_round_up(bias_row_size, 128) : 0;
size_t dst_sz = 0;
size_t act_raw_sz = 0;
@@ -544,22 +549,21 @@ static inline void htp_mm_hvx_vtcm_layout_build(
}
size_t tiled_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
- size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);
+ size_t q_act_sz = hex_round_up(tiled_act_row_size * act_nrows, 128);
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
- src1_sz = act_sz; // quantized activation buffer
- src2_sz = 0;
- src3_sz = 0;
+ act_sz = q_act_sz; // quantized activation buffer
+ bias_sz = 0;
dst_sz = 0;
- act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+ act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
} else if (is_matmul_id) {
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
- const size_t src1_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
- : htp_mm_q8_0_tiled_row_size(ne10);
+ const size_t act_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
+ : htp_mm_q8_0_tiled_row_size(ne10);
size_t src0_sz_per_thread = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
- src1_sz = htp_mm_round_up(src1_row_size_tiled * src1_nrows, 256);
+ act_sz = htp_mm_round_up(act_row_size_tiled * act_nrows, 256);
if (is_repack) {
const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype);
@@ -573,25 +577,24 @@ static inline void htp_mm_hvx_vtcm_layout_build(
src0_sz = src0_sz_per_thread * n_threads;
dst_sz = 0;
- src2_sz = 0;
- src3_sz = 0;
- act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+ bias_sz = 0;
+ act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
} else {
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
- const size_t dst_nrows = (src1_nrows > 1) ? 0 : 1;
+ const size_t dst_nrows = (act_nrows > 1) ? 0 : 1;
switch (kernel_type) {
case HTP_MM_KERNEL_HVX_F16_F16_VTCM: {
- size_t f16_src1_row_size = htp_mm_round_up(ne10 * 2, 128);
- src1_sz = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
+ size_t f16_act_row_size = htp_mm_round_up(ne10 * 2, 128);
+ act_sz = htp_mm_round_up(f16_act_row_size * act_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
- act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * src1_nrows, 128);
+ act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * act_nrows, 128);
break;
}
case HTP_MM_KERNEL_HVX_F32_F32_VTCM: {
- size_t f32_src1_row_size = htp_mm_round_up(ne10 * 4, 128);
- src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
+ size_t f32_act_row_size = htp_mm_round_up(ne10 * 4, 128);
+ act_sz = htp_mm_round_up(f32_act_row_size * act_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
act_raw_sz = 0;
@@ -599,10 +602,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
}
case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
case HTP_MM_KERNEL_HVX_QUANT_ROW: {
- size_t q_src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
+ size_t q_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
- src1_sz = htp_mm_round_up(q_src1_row_size * src1_nrows, 256);
+ act_sz = htp_mm_round_up(q_act_row_size * act_nrows, 256);
src0_sz = src0_sz * n_threads;
@@ -614,10 +617,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
src0_sz = repacked_vtcm_size * n_threads;
}
- size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
+ size_t dst_slice_per_thread = (dst_nrows > 0 && act_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
dst_sz = dst_slice_per_thread * n_threads;
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
- act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
+ act_raw_sz = hex_round_up(raw_row_size * act_nrows, 128);
break;
}
default:
@@ -627,9 +630,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
// Group A: Persistent buffers across chunk compute
size_t off_group_a = 0;
- VTCM_LAYOUT_ALLOC(off_group_a, off_src1, src1_sz);
- VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, src2_sz, src2_sz > 0);
- VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src3, src3_sz, src3_sz > 0);
+ VTCM_LAYOUT_ALLOC(off_group_a, off_act, act_sz);
+ VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_bias, bias_sz, bias_sz > 0);
// Group B: Compute-only buffers (starts at off_group_a)
size_t off_group_b = off_group_a;
@@ -643,9 +645,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
const size_t group_c_size = off_group_c - off_group_a;
L->src0_bytes = src0_sz;
- L->src1_bytes = src1_sz;
- L->src2_bytes = src2_sz;
- L->src3_bytes = src3_sz;
+ L->act_bytes = act_sz;
+ L->bias_bytes = bias_sz;
L->dst_bytes = dst_sz;
L->act_raw_bytes = act_raw_sz;
L->total_bytes = off_group_a + hex_smax(group_b_size, group_c_size);
@@ -655,12 +656,12 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
int kernel_type,
int wtype,
uint32_t ne10,
- uint32_t src1_nrows,
+ uint32_t act_nrows,
uint32_t n_threads,
size_t dst_row_size,
size_t src0_row_size,
- size_t src1_row_size,
- size_t src2_row_size,
+ size_t act_row_size,
+ size_t bias_row_size,
uint32_t n_prefetch,
size_t vtcm_budget,
struct htp_mm_hvx_vtcm_layout * L_out,
@@ -668,17 +669,17 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
) {
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(
- &L, kernel_type, wtype, ne10, src1_nrows, n_threads,
- dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+ &L, kernel_type, wtype, ne10, act_nrows, n_threads,
+ dst_row_size, src0_row_size, act_row_size, bias_row_size, n_prefetch, false, false
);
if (L.total_bytes <= vtcm_budget) {
*L_out = L;
- *m_chunk_out = src1_nrows;
+ *m_chunk_out = act_nrows;
return true;
}
- const size_t fixed_bytes = L.src0_bytes + L.src2_bytes + L.dst_bytes;
+ const size_t fixed_bytes = L.src0_bytes + L.bias_bytes + L.dst_bytes;
if (vtcm_budget <= fixed_bytes) {
return false;
}
@@ -707,8 +708,8 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
if (m_chunk > 1) {
m_chunk &= ~1U;
}
- if (m_chunk > src1_nrows) {
- m_chunk = src1_nrows;
+ if (m_chunk > act_nrows) {
+ m_chunk = act_nrows;
}
if (m_chunk < 1) {
return false;
@@ -716,14 +717,14 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
htp_mm_hvx_vtcm_layout_build(
&L, kernel_type, wtype, ne10, m_chunk, n_threads,
- dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+ dst_row_size, src0_row_size, act_row_size, bias_row_size, n_prefetch, false, false
);
while (m_chunk > 2 && L.total_bytes > vtcm_budget) {
m_chunk -= 2;
htp_mm_hvx_vtcm_layout_build(
&L, kernel_type, wtype, ne10, m_chunk, n_threads,
- dst_row_size, src0_row_size, src1_row_size, src2_row_size, n_prefetch, false, false
+ dst_row_size, src0_row_size, act_row_size, bias_row_size, n_prefetch, false, false
);
}
@@ -737,18 +738,18 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
}
static inline size_t htp_mm_hmx_get_2d_vtcm_size(
- int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size, size_t src2_size
+ int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size, size_t bias_size
) {
struct htp_mm_hmx_vtcm_layout L;
- htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, pipeline, act_threads, aligned_tile_size, src2_size);
+ htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, pipeline, act_threads, aligned_tile_size, bias_size);
return L.total_bytes;
}
static inline size_t htp_mm_hmx_get_batched_vtcm_size(
- int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool pipeline, uint32_t act_threads, size_t src2_size) {
+ int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool pipeline, uint32_t act_threads, size_t bias_size) {
(void)pipeline;
struct htp_mm_hmx_vtcm_layout L;
- htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, false, act_threads, 0, src2_size);
+ htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, false, act_threads, 0, bias_size);
return L.total_bytes;
}
@@ -760,7 +761,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
uint32_t group_size,
int n_threads,
bool pipeline,
- size_t src2_size,
+ size_t bias_size,
size_t vtcm_budget,
size_t * m_chunk_out,
size_t * n_chunk_out,
@@ -775,7 +776,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
int act_threads = n_threads;
while (act_threads >= 1) {
- size_t group_overhead = htp_mm_hmx_get_batched_overhead() + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
+ size_t group_overhead = htp_mm_hmx_get_batched_overhead() + (bias_size > 0 ? hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE) : 0);
size_t group_size_per_n, group_size_per_m, group_size_per_mn;
htp_mm_hmx_get_batched_chunk_costs(k, group_size, &group_size_per_n, &group_size_per_m, &group_size_per_mn);
@@ -785,8 +786,8 @@ static inline bool htp_mm_hmx_solve_batched_params(
if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
(size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
- &m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
- size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, pipeline, act_threads, src2_size);
+ &m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
+ size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, pipeline, act_threads, bias_size);
if (exact_size <= vtcm_budget) {
size_t mblocks = ((size_t) ne11 + m_chunk_candidate - 1) / m_chunk_candidate;
if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
@@ -826,7 +827,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
bool pipeline,
bool is_matmul_id,
uint32_t aligned_tile_size,
- size_t src2_size,
+ size_t bias_size,
size_t vtcm_budget,
size_t * m_chunk_out,
size_t * n_chunk_out,
@@ -843,7 +844,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
int act_threads = n_threads;
while (act_threads >= 1) {
- size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id) + (src2_size > 0 ? hex_align_up(src2_size, HTP_MM_HMX_TILE_SIZE) : 0);
+ size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id) + (bias_size > 0 ? hex_align_up(bias_size, HTP_MM_HMX_TILE_SIZE) : 0);
size_t simple_2d_size_per_n, simple_2d_size_per_m, simple_2d_size_per_mn;
htp_mm_hmx_get_2d_chunk_costs(wtype, k, pipeline, aligned_tile_size, &simple_2d_size_per_n, &simple_2d_size_per_m, &simple_2d_size_per_mn);
@@ -854,7 +855,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
if (htp_mm_hmx_compute_chunks(vtcm_budget, simple_2d_overhead, simple_2d_size_per_n, simple_2d_size_per_m, simple_2d_size_per_mn, m_for_chunks, ne01_padded,
(size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) m_for_cost * HTP_MM_HMX_COST_A_CONVERT,
&m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
- size_t exact_size = htp_mm_hmx_get_2d_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, pipeline, is_matmul_id ? 0 : act_threads, aligned_tile_size, src2_size);
+ size_t exact_size = htp_mm_hmx_get_2d_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, pipeline, is_matmul_id ? 0 : act_threads, aligned_tile_size, bias_size);
if (exact_size <= vtcm_budget) {
size_t mblocks = ((size_t) m_for_cost + m_chunk_candidate - 1) / m_chunk_candidate;
if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
diff --git a/scripts/snapdragon/run.py b/scripts/snapdragon/run.py
index 01093ddd8..14e53aee6 100755
--- a/scripts/snapdragon/run.py
+++ b/scripts/snapdragon/run.py
@@ -31,8 +31,10 @@ MANAGED_ENV_NAMES = (
"GGML_HEXAGON_MBUF",
"GGML_HEXAGON_MM_SELECT",
"GGML_HEXAGON_FA_SELECT",
+ "GGML_HEXAGON_FA_HEAD_SPLIT",
"GGML_HEXAGON_GDN_SELECT",
"GGML_HEXAGON_AR_SELECT",
+ "GGML_HEXAGON_AR_SCATTER",
"GGML_HEXAGON_ETM",
"GGML_HEXAGON_ARCH",
"GGML_HEXAGON_OPTRACE",
@@ -167,6 +169,7 @@ def main():
parser.add_argument("--hex-mbuf", help="Maximum host buffer size limit in MB to allocate (GGML_HEXAGON_MBUF)")
parser.add_argument("--hex-mm-select", help="Select MUL_MAT and MUL_MAT_ID kernel (GGML_HEXAGON_MM_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-fa-select", help="Select Flash Attention kernel (GGML_HEXAGON_FA_SELECT) 2:HMX,1:HVX,0:disable")
+ parser.add_argument("--hex-fa-head-split", help="Enable (1) or disable (0) head-parallel flash_attn partitioning (GGML_HEXAGON_FA_HEAD_SPLIT)")
parser.add_argument("--hex-gdn-select", help="Select Gated Delta Net kernel (GGML_HEXAGON_GDN_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-ar-select", help="Select All-Reduce kernel (GGML_HEXAGON_AR_SELECT) 1:enable,0:disable")
parser.add_argument("--hex-ar-scatter", help="Enable (1) or disable (0) reduce-scatter for fused ALLREDUCE+ADD (GGML_HEXAGON_AR_SCATTER)")
@@ -309,6 +312,7 @@ def main():
set_env("GGML_HEXAGON_MBUF", args.hex_mbuf)
set_env("GGML_HEXAGON_MM_SELECT", args.hex_mm_select)
set_env("GGML_HEXAGON_FA_SELECT", args.hex_fa_select)
+ set_env("GGML_HEXAGON_FA_HEAD_SPLIT", args.hex_fa_head_split)
set_env("GGML_HEXAGON_GDN_SELECT", args.hex_gdn_select)
set_env("GGML_HEXAGON_AR_SELECT", args.hex_ar_select)
set_env("GGML_HEXAGON_AR_SCATTER", args.hex_ar_scatter)