hexagon: MUL_MAT and MUL_MAT_ID fusion and fixes (#28202)
* hex-mm: fuse QKV and FFN matmuls that land on HMX * hex-mm: remove hardcoded ne[1] < 32K restriction * hex-get-rows: explicitly reject repacked Q8_0 just in case somebody decided to add an override * hex-mm: correct overhead sizing to make sure we dont exceed vtcm budget for large dims * hex-mm: fuse MUL_MAT_ID into MUL_MAT_ID_NX (2x,3x,...) where possible * hex-fusion: update opbatch and opqueue sizing to acount for new fusion and reduce overhead for trace buffer alloc * hex-bufs: sort buffers while finalizing opbatch, helps avoid va space fragmentation * hex-bufs: add simple va defrag to make sure we dont abort just because the va space is fragmented * hex-mm: replaced more scalar divs with fastdiv and minor cleanup * hex-mm: tighten up supported fusion checks to exactly match supported kernels
This commit is contained in:
@@ -98,12 +98,26 @@ static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 =
|
||||
// https://docs.qualcomm.com/doc/80-N2040-61/topic/hvx-pmu-events.html
|
||||
static u32vec opt_pmu_evt { 0x3, 0x111, 0x100, 0x105, 0x240, 0x256, 0x7D, 0x8C };
|
||||
|
||||
static int opt_opbatch = 1024; // max number of ops in a batch
|
||||
static int opt_opqueue = 64; // max number of pending batches
|
||||
static int opt_opbatch = 1280; // max number of ops in a batch
|
||||
static int opt_opqueue = 32; // max number of pending batches
|
||||
static int opt_optrace = 0; // trace buffer size per thread (0 means default)
|
||||
static int opt_oppoll = 0; // polling for batch completions
|
||||
static int opt_opfusion = 1; // enable/disable op fusion
|
||||
|
||||
enum ggml_hexagon_fusion_flags {
|
||||
GGML_HEXAGON_FUSE_ALLREDUCE_ADD = (1 << 1), // 2
|
||||
GGML_HEXAGON_FUSE_RMS_NORM_MUL = (1 << 2), // 4
|
||||
GGML_HEXAGON_FUSE_MUL_MAT_ADD = (1 << 3), // 8
|
||||
GGML_HEXAGON_FUSE_MUL_MAT_NX = (1 << 4), // 16
|
||||
GGML_HEXAGON_FUSE_MUL_MAT_ID_NX = (1 << 5), // 32
|
||||
};
|
||||
|
||||
static inline bool ggml_hexagon_is_fusion_enabled(int flag) {
|
||||
if (opt_opfusion <= 0) return false;
|
||||
if (opt_opfusion == 1) return true; // 1 enables all
|
||||
return (opt_opfusion & flag) != 0;
|
||||
}
|
||||
|
||||
static std::regex* opt_opfilter = NULL; // regex of ops to not claim
|
||||
|
||||
#define HEX_VERBOSE(...) \
|
||||
@@ -293,6 +307,15 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
|
||||
struct htp_mm_kernel_params * kparams
|
||||
);
|
||||
|
||||
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 * dst,
|
||||
int32_t n_weights,
|
||||
struct htp_mm_kernel_params * kparams
|
||||
);
|
||||
|
||||
static bool ggml_hexagon_precompute_allreduce_params(
|
||||
const struct ggml_hexagon_session * sess,
|
||||
const struct ggml_tensor * dst,
|
||||
@@ -304,8 +327,12 @@ static bool ggml_hexagon_precompute_allreduce_params(
|
||||
);
|
||||
|
||||
static bool mm_is_hmx_eligible(const ggml_tensor * t);
|
||||
static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams);
|
||||
static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams);
|
||||
static bool is_mergeable_mul_mat(const ggml_tensor * t);
|
||||
static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2);
|
||||
static bool is_mergeable_mul_mat_id(const ggml_tensor * t);
|
||||
static bool is_mergeable_mul_mat_id_pair(const ggml_tensor * n1, const ggml_tensor * n2);
|
||||
|
||||
// ** backend sessions
|
||||
|
||||
@@ -1832,6 +1859,42 @@ struct ggml_hexagon_opbatch {
|
||||
}
|
||||
}
|
||||
|
||||
void sort_buffers() {
|
||||
if (n_bufs <= 1) return;
|
||||
|
||||
std::vector<int> order(n_bufs);
|
||||
for (unsigned int i = 0; i < n_bufs; i++) { order[i] = (int) i; }
|
||||
|
||||
std::stable_sort(order.begin(), order.end(), [&](int a, int b) {
|
||||
return h_bufs[a].size > h_bufs[b].size;
|
||||
});
|
||||
|
||||
bool already_sorted = true;
|
||||
for (unsigned int i = 0; i < n_bufs; i++) {
|
||||
if (order[i] != (int) i) {
|
||||
already_sorted = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (already_sorted) return;
|
||||
|
||||
std::vector<uint16_t> remap(n_bufs);
|
||||
std::vector<htp_buf_desc> sorted_bufs(n_bufs);
|
||||
for (unsigned int new_bi = 0; new_bi < n_bufs; new_bi++) {
|
||||
int old_bi = order[new_bi];
|
||||
remap[old_bi] = (uint16_t) new_bi;
|
||||
sorted_bufs[new_bi] = h_bufs[old_bi];
|
||||
}
|
||||
|
||||
for (unsigned int i = 0; i < n_bufs; i++) {
|
||||
h_bufs[i] = sorted_bufs[i];
|
||||
}
|
||||
|
||||
for (unsigned int i = 0; i < n_tens; i++) {
|
||||
h_tens[i].bi = remap[h_tens[i].bi];
|
||||
}
|
||||
}
|
||||
|
||||
bool try_fuse_allreduce_add(const htp_opnode & node) {
|
||||
if (n_ops == 0 || opt_ar_select != 2) return false;
|
||||
if (node.opcode != HTP_OP_ADD) return false;
|
||||
@@ -2144,9 +2207,15 @@ struct ggml_hexagon_opbatch {
|
||||
if (x_in != x || w_in->type != w0->type || w_in->ne[0] != w0->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
if (!last_node.fused.empty() && (mm_is_hmx_eligible(last_node.fused[0]) != mm_is_hmx_eligible(node.node))) {
|
||||
return false;
|
||||
}
|
||||
|
||||
struct htp_mm_kernel_params kparams;
|
||||
ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams);
|
||||
if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
|
||||
return false;
|
||||
}
|
||||
if ((size_t) kparams.vtcm_size > sess->vtcm_size) {
|
||||
HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n",
|
||||
sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
|
||||
@@ -2210,6 +2279,9 @@ struct ggml_hexagon_opbatch {
|
||||
|
||||
struct htp_mm_kernel_params kparams;
|
||||
ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
|
||||
if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) {
|
||||
return false;
|
||||
}
|
||||
if ((size_t) kparams.vtcm_size > sess->vtcm_size) {
|
||||
HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n",
|
||||
sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
|
||||
@@ -2272,18 +2344,172 @@ struct ggml_hexagon_opbatch {
|
||||
return false;
|
||||
}
|
||||
|
||||
enum ggml_hexagon_fusion_flags {
|
||||
GGML_HEXAGON_FUSE_ALLREDUCE_ADD = (1 << 1), // 2
|
||||
GGML_HEXAGON_FUSE_RMS_NORM_MUL = (1 << 2), // 4
|
||||
GGML_HEXAGON_FUSE_MUL_MAT_ADD = (1 << 3), // 8
|
||||
GGML_HEXAGON_FUSE_MUL_MAT_NX = (1 << 4), // 16
|
||||
};
|
||||
bool try_fuse_mul_mat_id_nx(const htp_opnode & node) {
|
||||
if (n_ops == 0 || node.opcode != HTP_OP_MUL_MAT_ID) return false;
|
||||
if (!is_mergeable_mul_mat_id(node.node)) return false;
|
||||
|
||||
static inline bool ggml_hexagon_is_fusion_enabled(int flag) {
|
||||
if (opt_opfusion <= 0) return false;
|
||||
if (opt_opfusion == 1) return true; // 1 enables all
|
||||
return (opt_opfusion & flag) != 0;
|
||||
}
|
||||
const ggml_tensor * w_in = node.src0();
|
||||
const ggml_tensor * x_in = node.src1();
|
||||
const ggml_tensor * ids_in = node.node->src[2];
|
||||
const ggml_tensor * d_in = node.dst();
|
||||
if (!w_in || !x_in || !ids_in || !d_in) return false;
|
||||
|
||||
htp_opnode & last_node = ops[n_ops - 1];
|
||||
|
||||
// Case 1: last_node is already MUL_MAT_ID_NX
|
||||
if (last_node.opcode == HTP_OP_MUL_MAT_ID_NX) {
|
||||
const uint32_t curr_n = (uint32_t) last_node.outputs.size();
|
||||
if (curr_n >= HTP_OP_MAX_OUTPUTS || curr_n + 2 >= HTP_OP_MAX_INPUTS) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * w0 = last_node.inputs[0];
|
||||
const ggml_tensor * x = last_node.inputs[curr_n];
|
||||
const ggml_tensor * ids = last_node.inputs[curr_n + 1];
|
||||
|
||||
if (x_in != x || ids_in != ids || w_in->type != w0->type || w_in->ne[0] != w0->ne[0] || w_in->ne[2] != w0->ne[2]) {
|
||||
return false;
|
||||
}
|
||||
if (!last_node.fused.empty() && (mm_is_hmx_eligible(last_node.fused[0]) != mm_is_hmx_eligible(node.node))) {
|
||||
return false;
|
||||
}
|
||||
|
||||
struct htp_mm_kernel_params kparams;
|
||||
ggml_hexagon_precompute_fused_mmidnx_params(sess, w0, x, d_in, curr_n + 1, &kparams);
|
||||
if (!is_supported_mul_mat_id_nx_kernel(w0, &kparams)) {
|
||||
return false;
|
||||
}
|
||||
if ((size_t) kparams.vtcm_size > sess->vtcm_size) {
|
||||
HEX_VERBOSE("ggml-hex: %s skip ID NX fusion: VTCM needed (%d) > budget (%zu)\n",
|
||||
sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
|
||||
return false;
|
||||
}
|
||||
|
||||
size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
|
||||
auto fit_t = [&](const ggml_tensor * t) {
|
||||
if (!t) return;
|
||||
if (!t_map.count(t)) {
|
||||
extra_tens++;
|
||||
auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
|
||||
if (!b_map.count(sbuf->fd())) {
|
||||
extra_vmem += sbuf->size();
|
||||
extra_bufs += 1;
|
||||
}
|
||||
}
|
||||
};
|
||||
fit_t(w_in);
|
||||
fit_t(d_in);
|
||||
if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
|
||||
return false;
|
||||
}
|
||||
|
||||
last_node.inputs[curr_n] = w_in;
|
||||
last_node.inputs[curr_n + 1] = x;
|
||||
last_node.inputs.push_back(ids);
|
||||
last_node.outputs.push_back(d_in);
|
||||
last_node.fused.push_back(node.node);
|
||||
memcpy(last_node.kernel_params, &kparams, sizeof(kparams));
|
||||
|
||||
htp_op_desc & o = h_ops[n_ops - 1];
|
||||
memcpy(o.kernel_params, &kparams, sizeof(kparams));
|
||||
|
||||
for (uint32_t s = 0; s <= curr_n + 2; s++) {
|
||||
o.src[s] = add_tensor(last_node.inputs[s]);
|
||||
}
|
||||
for (uint32_t s = curr_n + 3; s < HTP_OP_MAX_INPUTS; s++) {
|
||||
o.src[s] = 0xffff;
|
||||
}
|
||||
for (uint32_t d = 0; d <= curr_n; d++) {
|
||||
o.dst[d] = add_tensor(last_node.outputs[d]);
|
||||
}
|
||||
for (uint32_t d = curr_n + 1; d < HTP_OP_MAX_OUTPUTS; d++) {
|
||||
o.dst[d] = 0xffff;
|
||||
}
|
||||
|
||||
HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_ID_NX (N=%u, #%u)\n", sess->c_name(), curr_n + 1, n_ops - 1);
|
||||
return true;
|
||||
}
|
||||
|
||||
// Case 2: last_node is single MUL_MAT_ID
|
||||
if (last_node.opcode == HTP_OP_MUL_MAT_ID) {
|
||||
if (!is_mergeable_mul_mat_id_pair(last_node.node, node.node)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * w0 = last_node.src0();
|
||||
const ggml_tensor * x = last_node.src1();
|
||||
const ggml_tensor * ids = last_node.node->src[2];
|
||||
const ggml_tensor * w1 = node.src0();
|
||||
if (!w0 || !x || !ids || !w1) return false;
|
||||
|
||||
struct htp_mm_kernel_params kparams;
|
||||
ggml_hexagon_precompute_fused_mmidnx_params(sess, w0, x, node.dst(), 2, &kparams);
|
||||
if (!is_supported_mul_mat_id_nx_kernel(w0, &kparams)) {
|
||||
return false;
|
||||
}
|
||||
if ((size_t) kparams.vtcm_size > sess->vtcm_size) {
|
||||
HEX_VERBOSE("ggml-hex: %s skip ID NX fusion: VTCM needed (%d) > budget (%zu)\n",
|
||||
sess->c_name(), kparams.vtcm_size, sess->vtcm_size);
|
||||
return false;
|
||||
}
|
||||
|
||||
size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
|
||||
auto fit_t = [&](const ggml_tensor * t) {
|
||||
if (!t) return;
|
||||
if (!t_map.count(t)) {
|
||||
extra_tens++;
|
||||
auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
|
||||
if (!b_map.count(sbuf->fd())) {
|
||||
extra_vmem += sbuf->size();
|
||||
extra_bufs += 1;
|
||||
}
|
||||
}
|
||||
};
|
||||
fit_t(w1);
|
||||
fit_t(node.dst());
|
||||
if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * dst_0 = last_node.dst();
|
||||
const ggml_tensor * dst_1 = node.dst();
|
||||
|
||||
last_node.opcode = HTP_OP_MUL_MAT_ID_NX;
|
||||
last_node.name = "MUL_MAT_ID_NX";
|
||||
last_node.inputs.clear();
|
||||
last_node.inputs.push_back(w0);
|
||||
last_node.inputs.push_back(w1);
|
||||
last_node.inputs.push_back(x);
|
||||
last_node.inputs.push_back(ids);
|
||||
last_node.outputs.clear();
|
||||
last_node.outputs.push_back(dst_0);
|
||||
last_node.outputs.push_back(dst_1);
|
||||
last_node.fused.push_back(node.node);
|
||||
memcpy(last_node.kernel_params, &kparams, sizeof(kparams));
|
||||
|
||||
htp_op_desc & o = h_ops[n_ops - 1];
|
||||
o.opcode = HTP_OP_MUL_MAT_ID_NX;
|
||||
memcpy(o.kernel_params, &kparams, sizeof(kparams));
|
||||
|
||||
o.src[0] = add_tensor(w0);
|
||||
o.src[1] = add_tensor(w1);
|
||||
o.src[2] = add_tensor(x);
|
||||
o.src[3] = add_tensor(ids);
|
||||
for (uint32_t s = 4; s < HTP_OP_MAX_INPUTS; s++) {
|
||||
o.src[s] = 0xffff;
|
||||
}
|
||||
o.dst[0] = add_tensor(dst_0);
|
||||
o.dst[1] = add_tensor(dst_1);
|
||||
for (uint32_t d = 2; d < HTP_OP_MAX_OUTPUTS; d++) {
|
||||
o.dst[d] = 0xffff;
|
||||
}
|
||||
|
||||
HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_ID_NX (N=2, #%u)\n", sess->c_name(), n_ops - 1);
|
||||
return true;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
bool try_fuse(const htp_opnode & node) {
|
||||
if (!opt_opfusion) return false;
|
||||
@@ -2291,6 +2517,7 @@ static inline bool ggml_hexagon_is_fusion_enabled(int flag) {
|
||||
if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_RMS_NORM_MUL) && try_fuse_rms_norm_mul(node)) return true;
|
||||
if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ADD) && try_fuse_mul_mat_add(node)) return true;
|
||||
if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_NX) && try_fuse_mul_mat_nx(node)) return true;
|
||||
if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ID_NX) && try_fuse_mul_mat_id_nx(node)) return true;
|
||||
return false;
|
||||
}
|
||||
};
|
||||
@@ -2350,6 +2577,8 @@ struct ggml_hexagon_opqueue {
|
||||
delete shm_buf;
|
||||
}
|
||||
|
||||
size_t shm_size() const { return shm_buf ? shm_buf->size() : 0; }
|
||||
|
||||
// push new batch
|
||||
bool push(htp_opbatch_req& req, dspqueue_buffer& dbuf, ggml_hexagon_opbatch* op_batch) {
|
||||
static_assert(sizeof(htp_opbatch_req) % 8 == 0, "sizeof(htp_opbatch_req) must be multiple of 8");
|
||||
@@ -2396,6 +2625,8 @@ struct ggml_hexagon_opqueue {
|
||||
uint8_t * t_ptr = m_ptr; m_ptr += t_size;
|
||||
uint8_t * o_ptr = m_ptr;
|
||||
|
||||
op_batch->sort_buffers();
|
||||
|
||||
memcpy(b_ptr, (void *) op_batch->h_bufs.data(), b_size);
|
||||
memcpy(t_ptr, (void *) op_batch->h_tens.data(), t_size);
|
||||
memcpy(o_ptr, (void *) op_batch->h_ops.data(), o_size);
|
||||
@@ -3018,7 +3249,8 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
|
||||
opt_vmem = ggml_hexagon_measure_max_vmem(this);
|
||||
GGML_LOG_INFO("ggml-hex: %s measured max vmem %zu\n", this->c_name(), opt_vmem);
|
||||
}
|
||||
this->max_vmem = opt_vmem;
|
||||
const size_t shm_size = this->op_queue->shm_size();
|
||||
this->max_vmem = (opt_vmem > shm_size) ? (opt_vmem - shm_size) : opt_vmem;
|
||||
|
||||
this->op_batch = new ggml_hexagon_opbatch(this, opt_opbatch, this->max_vmem);
|
||||
|
||||
@@ -3378,6 +3610,10 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
|
||||
bool is_matmul_id,
|
||||
bool is_batched
|
||||
) {
|
||||
if (src1->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const int ne00 = src0->ne[0];
|
||||
const int ne11 = src1->ne[1];
|
||||
const int ne12 = src1->ne[2];
|
||||
@@ -3408,7 +3644,8 @@ static bool ggml_hexagon_matmul_is_hmx_eligible(
|
||||
return false;
|
||||
}
|
||||
|
||||
// M alignment: Use HMX when M > HTP_MM_HMX_MIN_NROWS
|
||||
// M alignment: Use HMX when M > HTP_MM_HMX_MIN_NROWS.
|
||||
// For MUL_MAT_ID, src1 shape is [K, n_expert_used, n_tokens, 1], so n_tokens is ne12.
|
||||
const int m = is_matmul_id ? ne12 : ne11;
|
||||
if (m <= HTP_MM_HMX_MIN_NROWS) {
|
||||
return false;
|
||||
@@ -3460,7 +3697,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
|
||||
|
||||
if (!use_grouped) {
|
||||
// Fallback to simple 2D path (group_size = 1)
|
||||
const int m_id_rows = (int) ((size_t) dst->ne[1] * dst->ne[2]);
|
||||
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, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
|
||||
return false;
|
||||
}
|
||||
@@ -3918,64 +4155,113 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
|
||||
) {
|
||||
memset(kparams, 0, sizeof(*kparams));
|
||||
|
||||
const int wtype = src0->type;
|
||||
const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype);
|
||||
const int ne00 = src0->ne[0];
|
||||
const int ne01 = src0->ne[1];
|
||||
const int ne02 = src0->ne[2];
|
||||
const int ne03 = src0->ne[3];
|
||||
|
||||
const int ne10 = src1->ne[0];
|
||||
const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3];
|
||||
const size_t src1_row_size = (wtype == GGML_TYPE_Q4_1) ? 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];
|
||||
const int ne11 = src1->ne[1];
|
||||
const int ne12 = src1->ne[2];
|
||||
const int ne13 = src1->ne[3];
|
||||
|
||||
uint32_t best_n_prefetch = 16;
|
||||
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);
|
||||
|
||||
if (is_repack) {
|
||||
const uint32_t max_prefetch = (src1_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
|
||||
);
|
||||
if (L.total_bytes <= sess->vtcm_size) {
|
||||
best_n_prefetch = d;
|
||||
break;
|
||||
}
|
||||
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 >= 3);
|
||||
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, vtcm_budget, kparams)) {
|
||||
kparams->n_weights = n_weights;
|
||||
goto finalize;
|
||||
}
|
||||
}
|
||||
|
||||
struct htp_mm_hvx_vtcm_layout L;
|
||||
bool try_tiled = (opt_mm_select >= 2);
|
||||
|
||||
// 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
|
||||
);
|
||||
|
||||
if (try_tiled && L.total_bytes <= sess->vtcm_size) {
|
||||
kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
|
||||
kparams->vtcm_src0_size = L.src0_bytes;
|
||||
kparams->vtcm_src1_size = L.src1_bytes;
|
||||
kparams->vtcm_dst_size = L.dst_bytes;
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
kparams->n_prefetch = best_n_prefetch;
|
||||
kparams->n_weights = n_weights;
|
||||
} else {
|
||||
kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
|
||||
size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
|
||||
|
||||
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, 0, best_n_prefetch, false, true
|
||||
);
|
||||
kparams->vtcm_src0_size = L.src0_bytes;
|
||||
kparams->vtcm_src1_size = L.src1_bytes;
|
||||
kparams->vtcm_dst_size = L.dst_bytes;
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
kparams->n_prefetch = best_n_prefetch;
|
||||
kparams->n_weights = n_weights;
|
||||
if (!is_repack) {
|
||||
kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED;
|
||||
return;
|
||||
}
|
||||
|
||||
{
|
||||
const int src1_nrows = ne11 * ne12 * ne13;
|
||||
const size_t src1_row_size = (wtype == GGML_TYPE_Q4_1) ? 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;
|
||||
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
|
||||
);
|
||||
if (L.total_bytes <= sess->vtcm_size) {
|
||||
best_n_prefetch = d;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct htp_mm_hvx_vtcm_layout L;
|
||||
bool try_tiled = (opt_mm_select >= 2);
|
||||
|
||||
// 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
|
||||
);
|
||||
|
||||
if (try_tiled && L.total_bytes <= sess->vtcm_size) {
|
||||
kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW;
|
||||
kparams->vtcm_src0_size = L.src0_bytes;
|
||||
kparams->vtcm_src1_size = L.src1_bytes;
|
||||
kparams->vtcm_dst_size = L.dst_bytes;
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
kparams->n_prefetch = best_n_prefetch;
|
||||
kparams->n_weights = n_weights;
|
||||
} else {
|
||||
kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
|
||||
size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10);
|
||||
|
||||
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, 0, best_n_prefetch, false, true
|
||||
);
|
||||
kparams->vtcm_src0_size = L.src0_bytes;
|
||||
kparams->vtcm_src1_size = L.src1_bytes;
|
||||
kparams->vtcm_dst_size = L.dst_bytes;
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
kparams->n_prefetch = best_n_prefetch;
|
||||
kparams->n_weights = n_weights;
|
||||
}
|
||||
}
|
||||
|
||||
finalize:
|
||||
kparams->div_ne12_ne1 = init_fastdiv_values(ne12 * ne11);
|
||||
kparams->div_ne1 = init_fastdiv_values(ne11);
|
||||
kparams->div_r2 = init_fastdiv_values(ne02 > 0 ? ne12 / ne02 : 1);
|
||||
kparams->div_r3 = init_fastdiv_values(ne03 > 0 ? ne13 / ne03 : 1);
|
||||
kparams->div_ne11 = init_fastdiv_values(ne11);
|
||||
}
|
||||
|
||||
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 * dst, // dst0
|
||||
int32_t n_weights,
|
||||
struct htp_mm_kernel_params * kparams
|
||||
) {
|
||||
ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, kparams);
|
||||
kparams->n_weights = n_weights;
|
||||
}
|
||||
|
||||
static bool ggml_hexagon_tensor_is_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) {
|
||||
@@ -4010,11 +4296,6 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
|
||||
return false;
|
||||
}
|
||||
|
||||
// hardcoded limit to refuse the lm-head for now
|
||||
if (src0->ne[1] > 32768) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src1->ne[2] != 1 || src1->ne[3] != 1) {
|
||||
return false; // no broadcasting (for now)
|
||||
}
|
||||
@@ -4348,6 +4629,13 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
|
||||
const struct ggml_tensor * src1 = op->src[1]; // indices
|
||||
const struct ggml_tensor * dst = op;
|
||||
|
||||
if (src0->extra) {
|
||||
const auto * extra = (const ggml_hexagon_tensor_extra *) src0->extra;
|
||||
if (extra->flags & GGML_HEXAGON_TENSOR_REPACK) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 && src0->ne[0] < 32) {
|
||||
return false;
|
||||
}
|
||||
@@ -4734,10 +5022,43 @@ static bool mm_is_hmx_eligible(const ggml_tensor * t) {
|
||||
return ggml_hexagon_matmul_is_hmx_eligible(src0, src1, t, ne01_padded, is_matmul_id, is_batched);
|
||||
}
|
||||
|
||||
static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams) {
|
||||
if (kparams->n_hmx) {
|
||||
return kparams->kernel_type == HTP_MM_KERNEL_HMX_2D;
|
||||
}
|
||||
|
||||
if (!ggml_hexagon_is_repack_type(src0->type)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT;
|
||||
}
|
||||
|
||||
static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams) {
|
||||
if (kparams->n_hmx) {
|
||||
return kparams->kernel_type == HTP_MM_KERNEL_HMX_2D;
|
||||
}
|
||||
|
||||
if (!ggml_hexagon_is_repack_type(src0->type)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK;
|
||||
}
|
||||
|
||||
static bool is_mergeable_mul_mat(const ggml_tensor * t) {
|
||||
if (!t || t->op != GGML_OP_MUL_MAT) return false;
|
||||
if (t->src[1]->type != GGML_TYPE_F32) return false;
|
||||
return ggml_is_quantized(t->src[0]->type) && !mm_is_hmx_eligible(t);
|
||||
if (!t || t->op != GGML_OP_MUL_MAT) return false;
|
||||
|
||||
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);
|
||||
}
|
||||
|
||||
return ggml_hexagon_is_repack_type(src0->type);
|
||||
}
|
||||
|
||||
static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2) {
|
||||
@@ -4753,6 +5074,41 @@ static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor
|
||||
if (n1->src[0]->type != n2->src[0]->type) {
|
||||
return false;
|
||||
}
|
||||
if (mm_is_hmx_eligible(n1) != mm_is_hmx_eligible(n2)) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool is_mergeable_mul_mat_id(const ggml_tensor * t) {
|
||||
if (!t || t->op != GGML_OP_MUL_MAT_ID) return false;
|
||||
|
||||
const ggml_tensor * src0 = t->src[0];
|
||||
return ggml_hexagon_is_repack_type(src0->type);
|
||||
}
|
||||
|
||||
static bool is_mergeable_mul_mat_id_pair(const ggml_tensor * n1, const ggml_tensor * n2) {
|
||||
if (!is_mergeable_mul_mat_id(n1) || !is_mergeable_mul_mat_id(n2)) {
|
||||
return false;
|
||||
}
|
||||
if (n1->src[1] != n2->src[1]) {
|
||||
return false;
|
||||
}
|
||||
if (n1->src[2] != n2->src[2]) {
|
||||
return false;
|
||||
}
|
||||
if (n1->src[0]->ne[0] != n2->src[0]->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
if (n1->src[0]->ne[2] != n2->src[0]->ne[2]) {
|
||||
return false;
|
||||
}
|
||||
if (n1->src[0]->type != n2->src[0]->type) {
|
||||
return false;
|
||||
}
|
||||
if (mm_is_hmx_eligible(n1) != mm_is_hmx_eligible(n2)) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -4776,8 +5132,8 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
|
||||
|
||||
if (graph->nodes[i]->op == GGML_OP_RMS_NORM && ggml_can_fuse(graph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) {
|
||||
extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE;
|
||||
} else if (graph->nodes[i]->op == GGML_OP_MUL_MAT) {
|
||||
if ((i + 1 < graph->n_nodes && graph->nodes[i + 1]->op == GGML_OP_ADD && ggml_can_fuse(graph, i, { GGML_OP_MUL_MAT, GGML_OP_ADD })) ||
|
||||
} else if (graph->nodes[i]->op == GGML_OP_MUL_MAT || graph->nodes[i]->op == GGML_OP_MUL_MAT_ID) {
|
||||
if ((i + 1 < graph->n_nodes && graph->nodes[i + 1]->op == GGML_OP_ADD && ggml_can_fuse(graph, i, { graph->nodes[i]->op, GGML_OP_ADD })) ||
|
||||
ggml_node_has_n_uses(graph, i, 1)) {
|
||||
extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE;
|
||||
}
|
||||
|
||||
@@ -315,7 +315,8 @@ struct htp_opformat {
|
||||
}
|
||||
void format_kernel_params(char * str, size_t max_size, const htp_opnode & node) {
|
||||
if (node.opcode == HTP_OP_MUL_MAT || node.opcode == HTP_OP_MUL_MAT_ID ||
|
||||
node.opcode == HTP_OP_MUL_MAT_NX || node.opcode == HTP_OP_MUL_MAT_ADD) {
|
||||
node.opcode == HTP_OP_MUL_MAT_NX || node.opcode == HTP_OP_MUL_MAT_ID_NX ||
|
||||
node.opcode == HTP_OP_MUL_MAT_ADD) {
|
||||
const auto * kparams = (const struct htp_mm_kernel_params *) node.kernel_params;
|
||||
const char * path = "unknown";
|
||||
int32_t type = kparams->kernel_type;
|
||||
|
||||
@@ -118,6 +118,7 @@ struct htp_context {
|
||||
int op_matmul(struct htp_ops_context * octx);
|
||||
int op_matmul_id(struct htp_ops_context * octx);
|
||||
int op_matmul_nx(struct htp_ops_context * octx);
|
||||
int op_matmul_id_nx(struct htp_ops_context * octx);
|
||||
int op_binary(struct htp_ops_context * octx);
|
||||
int op_unary(struct htp_ops_context * octx);
|
||||
int op_sum_rows(struct htp_ops_context * octx);
|
||||
|
||||
@@ -52,6 +52,7 @@ enum htp_op_code {
|
||||
HTP_OP_MUL_MAT,
|
||||
HTP_OP_MUL_MAT_ID,
|
||||
HTP_OP_MUL_MAT_NX,
|
||||
HTP_OP_MUL_MAT_ID_NX,
|
||||
HTP_OP_MUL_MAT_ADD,
|
||||
HTP_OP_RMS_NORM,
|
||||
HTP_OP_RMS_NORM_MUL,
|
||||
|
||||
@@ -753,6 +753,9 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_MUL_MAT_ID:
|
||||
return op_matmul_id(octx);
|
||||
|
||||
case HTP_OP_MUL_MAT_ID_NX:
|
||||
return op_matmul_id_nx(octx);
|
||||
|
||||
case HTP_OP_MUL_MAT_NX:
|
||||
return op_matmul_nx(octx);
|
||||
|
||||
@@ -878,8 +881,8 @@ static inline void drop_mmap(struct htp_context *ctx, struct htp_mmap *m) {
|
||||
}
|
||||
}
|
||||
|
||||
static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) {
|
||||
if (b->base) return; // already mapped
|
||||
static inline bool mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) {
|
||||
if (b->base) return true; // already mapped
|
||||
|
||||
// find unused mapping
|
||||
for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) {
|
||||
@@ -887,8 +890,8 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) {
|
||||
if (!m->size) {
|
||||
void *va = htp_mmap(b->fd, b->size);
|
||||
if (va == NULL) {
|
||||
FARF(ERROR, "mmap failed : fd %u size %u", b->fd, (uint32_t) b->size);
|
||||
abort(); // can't do much else at this point
|
||||
FARF(HIGH, "mmap failed (will attempt defrag) : fd %u size %u", b->fd, (uint32_t) b->size);
|
||||
return false;
|
||||
}
|
||||
|
||||
m->base = b->base = (uint64_t) va;
|
||||
@@ -896,12 +899,12 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) {
|
||||
m->size = b->size;
|
||||
|
||||
FARF(ALWAYS, "mmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size);
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
FARF(ERROR, "mmap failed : exceeded mapping capacity limit of %u", HTP_MAX_MMAPS);
|
||||
abort();
|
||||
return false;
|
||||
}
|
||||
|
||||
static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uint32_t n_bufs) {
|
||||
@@ -934,12 +937,32 @@ static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uin
|
||||
}
|
||||
}
|
||||
|
||||
// Create missing mappings
|
||||
// Create missing mappings (pass 1)
|
||||
bool mmap_ok = true;
|
||||
for (uint32_t i=0; i < n_bufs; i++) {
|
||||
struct htp_buf_desc *b = bufs + i;
|
||||
mmap_buf(ctx, b);
|
||||
if (!mmap_buf(ctx, b)) {
|
||||
mmap_ok = false;
|
||||
break;
|
||||
}
|
||||
FARF(HIGH, "prep-buf #%u : pass1 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags);
|
||||
}
|
||||
|
||||
if (!mmap_ok) {
|
||||
// Attempt clean defragmentation: drop all mappings and remap (pass 2)
|
||||
FARF(HIGH, "prep-bufs : dropping all mappings to defragment address space");
|
||||
for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) { drop_mmap(ctx, ctx->mmap + i); }
|
||||
|
||||
for (uint32_t i=0; i < n_bufs; i++) {
|
||||
struct htp_buf_desc *b = bufs + i;
|
||||
b->base = 0;
|
||||
if (!mmap_buf(ctx, b)) {
|
||||
FARF(ERROR, "prep-bufs : mmap failed after defragmentation (fd %u size %u)", b->fd, (uint32_t) b->size);
|
||||
abort();
|
||||
}
|
||||
FARF(HIGH, "prep-buf #%u : pass2 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static void prep_tensor(struct htp_context *ctx, struct htp_buf_desc *bufs, struct htp_tensor *tens, uint32_t idx, struct htp_tensor *t) {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -134,7 +134,8 @@ static inline int htp_mm_hmx_compute_chunks(size_t vtcm_total,
|
||||
size_t best_mn = 0;
|
||||
size_t best_m = 0, best_n = 0;
|
||||
|
||||
const size_t n_max = hex_align_down((size_t)n, HTP_MM_HMX_TILE_N_COLS);
|
||||
const size_t max_nc_budget = (usable / per_n_cost);
|
||||
const size_t n_max = hex_align_down(hex_smin((size_t)n, max_nc_budget), HTP_MM_HMX_TILE_N_COLS);
|
||||
for (size_t nc = n_max; nc >= HTP_MM_HMX_TILE_N_COLS; nc -= HTP_MM_HMX_TILE_N_COLS) {
|
||||
size_t n_fixed = 0, ncmn = 0, mc_denom = 0;
|
||||
if (hex_mul_overflow(nc, per_n_cost, &n_fixed)) continue;
|
||||
@@ -299,6 +300,15 @@ static inline void htp_mm_hmx_get_batched_chunk_costs(
|
||||
*size_per_mn_out = sizeof(uint16_t);
|
||||
}
|
||||
|
||||
static inline size_t htp_mm_hmx_get_2d_overhead(bool pipeline, bool is_matmul_id) {
|
||||
size_t num_regions = pipeline ? 7 : (is_matmul_id ? 4 : 5);
|
||||
return num_regions * HTP_MM_HMX_TILE_SIZE + 256;
|
||||
}
|
||||
|
||||
static inline size_t htp_mm_hmx_get_batched_overhead(void) {
|
||||
return 5 * HTP_MM_HMX_TILE_SIZE + 256;
|
||||
}
|
||||
|
||||
struct htp_mm_hmx_vtcm_layout {
|
||||
// Byte offsets from vtcm_base for each region
|
||||
size_t off_weight[2]; // [1] is only used when pipelined
|
||||
@@ -568,10 +578,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
}
|
||||
|
||||
size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
size_t dst_size_per_thread = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) : 0;
|
||||
if (dst_size_per_thread < quant_scratch_size_per_thread) {
|
||||
dst_size_per_thread = quant_scratch_size_per_thread;
|
||||
}
|
||||
size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
|
||||
size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
|
||||
dst_sz = dst_size_per_thread * n_threads;
|
||||
break;
|
||||
}
|
||||
@@ -592,10 +600,8 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
}
|
||||
|
||||
size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
size_t dst_size_per_thread = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) : 0;
|
||||
if (dst_size_per_thread < quant_scratch_size_per_thread) {
|
||||
dst_size_per_thread = quant_scratch_size_per_thread;
|
||||
}
|
||||
size_t dst_slice_per_thread = dst_nrows > 0 ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
|
||||
size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
|
||||
dst_sz = dst_size_per_thread * n_threads;
|
||||
break;
|
||||
}
|
||||
@@ -658,7 +664,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
|
||||
|
||||
int act_threads = n_threads;
|
||||
while (act_threads >= 1) {
|
||||
size_t group_overhead = 256;
|
||||
size_t group_overhead = htp_mm_hmx_get_batched_overhead();
|
||||
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);
|
||||
|
||||
@@ -725,7 +731,7 @@ static inline bool htp_mm_hmx_solve_2d_params(
|
||||
|
||||
int act_threads = n_threads;
|
||||
while (act_threads >= 1) {
|
||||
size_t simple_2d_overhead = 256;
|
||||
size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id);
|
||||
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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user