hexagon: further improved pipeline of the core bits (L2, DMA, MM, FA) (#26049)

* hex-l2: use dirty ranges for flushing

* hex-l2: simplify range based flush logic

* hex-l2: optimize dirty range scans

* hex-hvx: support for reduce_max_i32

* hex-mm: optimize fused MUL_MAT+ADD to use vtcm for bias when it fits

* hex-mmid: optimize mmid row-mapping generation

* hex-mmid: optimize mmid row-mapping generation

* hex-mmid: optimize mmid row-mapping generation (round2)

* hmx-mm: optimize output proc by tiling (col-chunking)

* hex-fa: start the next q dmas a bit earlier

* hex-fa: prefetch Q even earlier

* hvx-fa: optimize softmax to keep things in hvx registers

* hex-fa: hoist const register init in softmax loop

* hmx-fa: kick off next-qkv DMAs before o-proc

* hmx-fa: hoist various checks out of the inner loop

* hmx-fa: adjust the cost model to better balance softmax work across hvx threads

* hmx-fa: overlap diag rescale build with last HMX task

* hmx-fa: optimize idx update in output proc

* hmx-fa: unroll the softmax loops for improved perf

* hmx-fa: overlap qk-dot with softmax, double-buffer p and s tiles

* hex-trace: double the default number of trace entries

* hex-trace: add trace events for opbatch and buffer mgmt

* hex-trace: overhaul tracing to simplify runtime event handling and support opbatch stats

* hex-trace: replace ascii timeline diagram with pipeline bubbles detector

* hex-trace: handle missing start/stop events

* hex-dma: always log stop/start trace events even for dummy dmas

* hex-scripts: fix flake warnings
This commit is contained in:
Max Krasnyansky
2026-07-23 19:13:03 -07:00
committed by GitHub
parent c0bc8591e8
commit 0a50d9909a
17 changed files with 1666 additions and 761 deletions
+92 -153
View File
@@ -143,12 +143,12 @@ static const char * htp_event_name(uint16_t id) {
case HTP_TRACE_EVT_HMX_COMP: return "HMX_COMP";
case HTP_TRACE_EVT_L2FLUSH: return "L2FLUSH";
case HTP_TRACE_EVT_INIT: return "INIT";
case HTP_TRACE_EVT_BUFF: return "BUFF";
default: return "UNKNOWN";
}
}
static void ggml_hexagon_dump_op_prof(const std::string &sess_name, const htp_opnode & node,
const htp_prof_desc & pd) {
static void ggml_hexagon_dump_op_prof(const std::string &sess_name, const htp_opnode & node, const htp_prof_desc & pd) {
if (!opt_profile) return;
uint32_t op_usec = pd.usecs;
@@ -168,6 +168,43 @@ static void ggml_hexagon_dump_op_prof(const std::string &sess_name, const htp_op
node.op_name().c_str(), fmt.names, fmt.dims, fmt.types, fmt.strides, fmt.kparams, op_usec, op_cycles, pd.cycles_start, mhz, pmu_str);
}
static void ggml_hexagon_dump_batch_prof(const std::string & sess_name, const htp_opbatch_rsp & rsp) {
uint64_t batch_cycles = rsp.cycles_stop - rsp.cycles_start;
float batch_mhz = rsp.usecs > 0 ? (float) batch_cycles / rsp.usecs : 0.0f;
char evt_str[256] = "----";
if (opt_profile == 3) {
snprintf(evt_str, sizeof(evt_str), "evt-cnt %u,%u,%u,%u,%u,%u,%u,%u,%u,%u,%u",
rsp.n_traces[0], rsp.n_traces[1], rsp.n_traces[2], rsp.n_traces[3],
rsp.n_traces[4], rsp.n_traces[5], rsp.n_traces[6], rsp.n_traces[7],
rsp.n_traces[8], rsp.n_traces[9], rsp.n_traces[10]);
}
GGML_LOG_DEBUG("ggml-hex: %s profile-op OPBATCH|----|n-ops %u|%s|----|----|usec %u cycles %llu start %llu mhz %.1f\n",
sess_name.c_str(), rsp.n_ops, evt_str, rsp.usecs, (unsigned long long) batch_cycles, (unsigned long long) rsp.cycles_start, batch_mhz);
}
static void ggml_hexagon_dump_trace_events(const std::string & sess_name, const htp_opbatch_rsp & rsp,
const htp_trace_desc * trace_events, uint32_t n_traces) {
if (opt_profile == 3 && trace_events) {
uint32_t valid_cnt[HTP_MAX_NTHREADS + 1] = {0};
for (uint32_t t = 0; t <= HTP_MAX_NTHREADS; t++) {
uint32_t count = rsp.n_traces[t];
valid_cnt[t] = count > n_traces ? n_traces : count;
}
for (uint32_t t = 0; t <= HTP_MAX_NTHREADS; t++) {
for (uint32_t idx = 0; idx < valid_cnt[t]; idx++) {
const auto & e = trace_events[t * n_traces + idx];
bool is_stop = (e.info & 0x8000) != 0;
uint16_t info = e.info & 0x7FFF;
GGML_LOG_DEBUG("ggml-hex: %s trace-evt %s: thread %u info %u %s %u\n",
sess_name.c_str(), htp_event_name(e.id), t, info, is_stop ? "stop" : "start", e.cycles);
}
}
}
}
// **
static inline bool ggml_hexagon_is_repack_type(enum ggml_type type) {
@@ -1128,13 +1165,7 @@ struct ggml_hexagon_opbatch {
std::unordered_map<const ggml_tensor*, int> t_map; // tensor ptr to index
std::unordered_multimap<void*, int> d_map; // tensor data to index
struct tensor_range {
uint64_t start;
uint64_t end;
int bi;
std::vector<int> tensors;
};
std::vector<tensor_range> ranges;
unsigned int n_bufs; // num buffers in the batch
unsigned int n_tens; // num tensors ...
@@ -1155,7 +1186,6 @@ struct ggml_hexagon_opbatch {
b_map.clear();
t_map.clear();
d_map.clear();
ranges.clear();
}
ggml_hexagon_opbatch(ggml_hexagon_session *sess, size_t batch_size, size_t max_vmem) {
@@ -1209,70 +1239,7 @@ struct ggml_hexagon_opbatch {
return bi;
}
void add_range(const htp_tensor * h, int ti) {
uint64_t t_start = h->data;
uint64_t t_end = t_start + h->size;
int bi = h->bi;
int first_match = -1;
int unused_idx = -1;
for (size_t i = 0; i < ranges.size(); i++) {
if (ranges[i].bi == -1) {
unused_idx = i;
continue;
}
if (ranges[i].bi != bi) {
continue;
}
if (ranges[i].start >= t_end || ranges[i].end <= t_start) {
continue;
}
if (first_match == -1) {
first_match = i;
HEX_VERBOSE("ggml-hex: %s range-grow #%d : bi %d [%p, %p) + #%d [%p, %p) -> [%p, %p)\n",
sess->c_name(), (int) i, ranges[i].bi,
(void *) (h_bufs[ranges[i].bi].base + ranges[i].start),
(void *) (h_bufs[ranges[i].bi].base + ranges[i].end),
ti,
(void *) (h_bufs[bi].base + t_start),
(void *) (h_bufs[bi].base + t_end),
(void *) (h_bufs[ranges[i].bi].base + std::min(ranges[i].start, t_start)),
(void *) (h_bufs[ranges[i].bi].base + std::max(ranges[i].end, t_end)));
ranges[i].start = std::min(ranges[i].start, t_start);
ranges[i].end = std::max(ranges[i].end, t_end);
ranges[i].tensors.push_back(ti);
} else {
HEX_VERBOSE("ggml-hex: %s range-merge #%d [%p, %p) + #%d [%p, %p) -> [%p, %p)\n",
sess->c_name(), first_match,
(void *) (h_bufs[bi].base + ranges[first_match].start),
(void *) (h_bufs[bi].base + ranges[first_match].end),
(int) i,
(void *) (h_bufs[bi].base + ranges[i].start),
(void *) (h_bufs[bi].base + ranges[i].end),
(void *) (h_bufs[bi].base + std::min(ranges[first_match].start, ranges[i].start)),
(void *) (h_bufs[bi].base + std::max(ranges[first_match].end, ranges[i].end)));
ranges[first_match].start = std::min(ranges[first_match].start, ranges[i].start);
ranges[first_match].end = std::max(ranges[first_match].end, ranges[i].end);
ranges[first_match].tensors.insert(
ranges[first_match].tensors.end(),
ranges[i].tensors.begin(),
ranges[i].tensors.end()
);
ranges[i].bi = -1;
}
}
if (first_match == -1) {
if (unused_idx != -1) {
ranges[unused_idx] = {t_start, t_end, bi, {ti}};
} else {
ranges.push_back({t_start, t_end, bi, {ti}});
}
}
}
bool same_shape(const htp_tensor * h, const ggml_tensor * t) const {
int64_t ne0 = t->ne[0];
@@ -1341,8 +1308,7 @@ struct ggml_hexagon_opbatch {
h.nb[0] = t->nb[0]; h.nb[1] = t->nb[1]; h.nb[2] = t->nb[2]; h.nb[3] = t->nb[3];
}
h.alias = ti;
add_range(&h, ti);
h.flags = 0;
if (ggml_backend_buffer_get_usage(t->buffer) != GGML_BACKEND_BUFFER_USAGE_WEIGHTS) {
@@ -1424,14 +1390,6 @@ struct ggml_hexagon_opbatch {
}
void finalize_ranges() {
for (const auto & r : ranges) {
if (r.bi == -1) {
continue;
}
for (size_t i = 0; i < r.tensors.size(); i++) {
h_tens[r.tensors[i]].alias = r.tensors[(i + 1) % r.tensors.size()];
}
}
}
};
@@ -1582,9 +1540,6 @@ struct ggml_hexagon_opqueue {
if (opt_profile && rsp.n_ops > 0) {
auto & ops = op_cache[rsp.id];
uint64_t batch_usec = ggml_time_us() - start_usec[rsp.id];
uint32_t htp_usec = 0;
GGML_ASSERT(rsp.n_ops <= ops.size());
const htp_prof_desc * pd = (const htp_prof_desc *) p_ptr;
@@ -1595,55 +1550,13 @@ struct ggml_hexagon_opqueue {
trace_events = (const htp_trace_desc *) (p_ptr + p_size);
}
uint32_t trace_idx[HTP_MAX_NTHREADS + 1] = {0};
uint32_t valid_cnt[HTP_MAX_NTHREADS + 1] = {0};
if (opt_profile == 3) {
for (uint32_t t = 0; t <= HTP_MAX_NTHREADS; t++) {
uint32_t count = rsp.n_traces[t];
valid_cnt[t] = count > n_traces ? n_traces : count;
}
}
ggml_hexagon_dump_batch_prof(shm_buf->sess->name, rsp);
for (uint32_t i = 0; i < rsp.n_ops; i++) {
htp_usec += pd[i].usecs;
ggml_hexagon_dump_op_prof(shm_buf->sess->name, ops[i], pd[i]);
if (opt_profile == 3) {
uint32_t op_duration = pd[i].cycles_stop - pd[i].cycles_start;
for (uint32_t t = 0; t <= HTP_MAX_NTHREADS; t++) {
while (trace_idx[t] < valid_cnt[t]) {
const auto & e = trace_events[t * n_traces + trace_idx[t]];
uint32_t offset = e.cycles - pd[i].cycles_start;
if (offset >= 0x80000000) {
trace_idx[t]++;
continue;
}
if (offset > op_duration) {
break;
}
bool is_stop = (e.info & 0x8000) != 0;
uint16_t info = e.info & 0x7FFF;
GGML_LOG_DEBUG("ggml-hex: %s trace-op %s: thread %u event %s info %u %s %u\n",
shm_buf->sess->c_name(), ops[i].op_name().c_str(), t, htp_event_name(e.id), info, is_stop ? "stop" : "start", e.cycles);
trace_idx[t]++;
}
}
}
}
char evt_str[256] = "";
if (opt_profile == 3) {
snprintf(evt_str, sizeof(evt_str), " evt [%u,%u,%u,%u,%u,%u,%u,%u,%u,%u,%u]",
rsp.n_traces[0], rsp.n_traces[1], rsp.n_traces[2], rsp.n_traces[3],
rsp.n_traces[4], rsp.n_traces[5], rsp.n_traces[6], rsp.n_traces[7],
rsp.n_traces[8], rsp.n_traces[9], rsp.n_traces[10]);
}
GGML_LOG_DEBUG("ggml-hex: %s profile-batch n-ops %u batch-dur-usec %lld htp-ops-usec %u%s\n",
shm_buf->sess->c_name(), rsp.n_ops, (long long) batch_usec, htp_usec, evt_str);
ggml_hexagon_dump_trace_events(shm_buf->sess->name, rsp, trace_events, n_traces);
}
}
};
@@ -2114,7 +2027,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(
const struct ggml_tensor * sinks = op->src[4];
if (ggml_hexagon_flash_attn_is_hmx_eligible(sess, q, k, v, sinks)) {
size_t Br = 0, Bc = 0;
int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK, DV, neq1, nek1, sess->vtcm_size, sess->n_threads);
int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK, DV, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0);
if (ret == 0) {
kparams->kernel_type = HTP_FA_KERNEL_HMX;
kparams->Br = Br;
@@ -2124,7 +2037,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, DV, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0);
kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK, DV, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0);
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
@@ -2413,6 +2326,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
int ne12,
int ne13,
bool is_matmul_id,
const size_t src2_row_size,
size_t vtcm_budget,
struct htp_mm_kernel_params * kparams
) {
@@ -2438,7 +2352,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], 0, d, true, false, false
0, src0->nb[1], 0, src2_row_size, d, true, false, false
);
if (L.total_bytes <= vtcm_budget) {
best_n_prefetch = d;
@@ -2448,7 +2362,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) {
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
0, src0->nb[1], 0, 2, true, false, false
0, src0->nb[1], 0, src2_row_size, 2, true, false, false
);
}
kparams->n_prefetch = best_n_prefetch;
@@ -2472,7 +2386,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,
dst->nb[1], src0->nb[1], src1->nb[1], d, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, d, false, false, false
);
if (L.total_bytes <= vtcm_budget) {
best_n_prefetch = d;
@@ -2482,7 +2396,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) {
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 2, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 2, false, false, false
);
}
@@ -2506,7 +2420,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 16, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false
);
kparams->n_prefetch = 16;
@@ -2526,7 +2440,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(
&L, HTP_MM_KERNEL_HVX_F16_F16_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 16, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false
);
if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) {
@@ -2546,7 +2460,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
kparams->src1_row_size = src1->nb[1];
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 16, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false
);
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
@@ -2562,7 +2476,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
struct htp_mm_hvx_vtcm_layout L;
htp_mm_hvx_vtcm_layout_build(
&L, HTP_MM_KERNEL_HVX_F32_F32_VTCM, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 16, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false
);
if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) {
@@ -2578,7 +2492,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
kparams->src1_row_size = src1->nb[1];
htp_mm_hvx_vtcm_layout_build(
&L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads,
dst->nb[1], src0->nb[1], src1->nb[1], 16, false, false, false
dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false
);
kparams->vtcm_size = L.total_bytes;
kparams->vtcm_src0_size = L.src0_bytes;
@@ -2589,11 +2503,12 @@ static void ggml_hexagon_precompute_hvx_mm_params(
}
}
static void ggml_hexagon_precompute_matmul_params(
static void ggml_hexagon_precompute_matmul_params_impl(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
const struct ggml_tensor * src1,
const struct ggml_tensor * dst,
const size_t src2_row_size,
struct htp_mm_kernel_params * kparams
) {
memset(kparams, 0, sizeof(*kparams));
@@ -2628,7 +2543,7 @@ static void ggml_hexagon_precompute_matmul_params(
}
// Fallback to HVX parameter computation
ggml_hexagon_precompute_hvx_mm_params(sess, src0, src1, dst, wtype, ne02, ne03, ne10, ne11, ne12, ne13, is_matmul_id, vtcm_budget, kparams);
ggml_hexagon_precompute_hvx_mm_params(sess, src0, src1, dst, wtype, ne02, ne03, ne10, ne11, ne12, ne13, is_matmul_id, src2_row_size, vtcm_budget, kparams);
finalize:
kparams->div_ne12_ne1 = init_fastdiv_values(ne12 * ne11);
@@ -2638,6 +2553,27 @@ finalize:
kparams->div_ne11 = init_fastdiv_values(ne11);
}
static void ggml_hexagon_precompute_matmul_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
const struct ggml_tensor * src1,
const struct ggml_tensor * dst,
struct htp_mm_kernel_params * kparams
) {
ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, kparams);
}
static void ggml_hexagon_precompute_fused_matmul_add_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
const struct ggml_tensor * src1,
const struct ggml_tensor * src2,
const struct ggml_tensor * dst,
struct htp_mm_kernel_params * kparams
) {
ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, src2->nb[1], kparams);
}
static void ggml_hexagon_precompute_unary_params(
const struct ggml_hexagon_session * sess,
uint32_t op,
@@ -2731,7 +2667,7 @@ static void ggml_hexagon_precompute_fused_qkv_params(
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, d, false, true, false
0, src0_row_size, src1_row_size, 0, d, false, true, false
);
if (L.total_bytes <= sess->vtcm_size) {
best_n_prefetch = d;
@@ -2746,7 +2682,7 @@ static void ggml_hexagon_precompute_fused_qkv_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, best_n_prefetch, false, true, false
0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true, false
);
if (try_tiled && L.total_bytes <= sess->vtcm_size) {
@@ -2764,7 +2700,7 @@ static void ggml_hexagon_precompute_fused_qkv_params(
htp_mm_hvx_vtcm_layout_build(
&L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads,
0, src0_row_size, flat_src1_row_size, best_n_prefetch, false, true, false
0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true, false
);
kparams->vtcm_src0_size = L.src0_bytes;
kparams->vtcm_src1_size = L.src1_bytes;
@@ -2801,7 +2737,7 @@ static void ggml_hexagon_precompute_fused_ffn_params(
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, d, false, false, true
0, src0_row_size, src1_row_size, 0, d, false, false, true
);
if (L.total_bytes <= sess->vtcm_size) {
best_n_prefetch = d;
@@ -2816,7 +2752,7 @@ static void ggml_hexagon_precompute_fused_ffn_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, best_n_prefetch, false, false, true
0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, false, true
);
if (try_tiled && L.total_bytes <= sess->vtcm_size) {
@@ -2833,7 +2769,7 @@ static void ggml_hexagon_precompute_fused_ffn_params(
htp_mm_hvx_vtcm_layout_build(
&L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads,
0, src0_row_size, flat_src1_row_size, best_n_prefetch, false, false, true
0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, false, true
);
kparams->vtcm_src0_size = L.src0_bytes;
kparams->vtcm_src1_size = L.src1_bytes;
@@ -3656,16 +3592,19 @@ static bool try_fuse_node(const ggml_hexagon_session * sess, const ggml_cgraph *
if (n->op == GGML_OP_MUL_MAT && next_node) {
if (next_node->op == GGML_OP_ADD && op_is_compute(next_node) && ggml_can_fuse(graph, i, { GGML_OP_MUL_MAT, GGML_OP_ADD })) {
if (next_node->src[0] == n || next_node->src[1] == n) {
const struct ggml_tensor * src2 = (next_node->src[0] == n) ? next_node->src[1] : next_node->src[0];
struct htp_mm_kernel_params kparams;
ggml_hexagon_precompute_matmul_params(sess, n->src[0], n->src[1], next_node, &kparams);
if ((size_t)kparams.vtcm_size <= sess->vtcm_size) {
ggml_hexagon_precompute_fused_matmul_add_params(sess, n->src[0], n->src[1], src2, next_node, &kparams);
const int src1_nrows = n->src[1]->ne[1] * n->src[1]->ne[2] * n->src[1]->ne[3];
const bool can_fuse = (kparams.n_hmx > 0) || (src1_nrows == 1);
if (can_fuse && (size_t)kparams.vtcm_size <= sess->vtcm_size) {
htp_opnode node(n, {}, HTP_OP_MUL_MAT_ADD);
node.add_fused(next_node);
memcpy(node.kernel_params, &kparams, sizeof(kparams));
nodes.push_back(std::move(node));
i += 1;
return true;
} else {
} else if (can_fuse) {
HEX_VERBOSE("ggml-hex: skip MUL_MAT_ADD fusion because VTCM needed (%d) > budget (%zu)\n",
kparams.vtcm_size, sess->vtcm_size);
}
@@ -4455,7 +4394,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
opt_opstage = str_opstage ? strtoul(str_opstage, NULL, 0) : opt_opstage;
opt_opbatch = str_opbatch ? strtoul(str_opbatch, NULL, 0) : opt_opbatch;
opt_opqueue = str_opqueue ? strtoul(str_opqueue, NULL, 0) : opt_opqueue;
opt_optrace = str_optrace ? strtoul(str_optrace, NULL, 0) : (opt_opbatch * 128);
opt_optrace = str_optrace ? strtoul(str_optrace, NULL, 0) : (opt_opbatch * 256);
opt_oppoll = str_oppoll ? strtoul(str_oppoll, NULL, 0) : opt_oppoll;
opt_opfusion = str_opfusion ? atoi(str_opfusion) : opt_opfusion;
opt_profile = str_profile ? atoi(str_profile) : 0;