hexagon: L2 cache handling rework (dirty bit tracking with lazy flushing) and more MUL_MAT updates (#25762)
* hex-mm: fix artificial limit in the solver that restricted number of act-prep threads * hex-mm: fix warning * hex-prof: do not apply --top to the timeline report * hmx-mm: add suport for tiled act-processing to better distribute hvx work * hex-l2: add tracing for l2flush events * workqueue: redo the legacy workpool api to match hmx-queue and dma-queue * hmx-mm: fix f32 activation buffer alignmnet for nhvx=5,6,7 * hex-work: minor cleanup for work-queue apis * hex-work: further cleanup of the work-queue api * hex-l2: optimize l2flushes at the opbatch level * hex-work: remove unused mask * hex-work: no need to drop hvx ctx in the work-queue * hex-work: add explicit wakeup/suspend and make threads spin * hex-bufs: mark any non-weight tensor as compute * hex-dma: dma-queue support for alias queues and cached dma * hex-l2: track tensor aliases and delay or skip flushes as much as possible * hex-l2: simplify tensor alias handling * hex-l2: handle overlapping views as a circular list of aliases * hex-tens: add flags helper * hex-l2: add helper for marking tensors clearn/dirty * hex-l2: mark binary and rope outputs as l2-clean and keep the rest as is for now * hex-l2: proper support for handling all tensor overlap scenarios * hex-trace: instrument matmul init code and cleanup trace checks * hex-thread: introduce dedicated main thread with explicit stack and priority * hex-l2: track dirty state as bitmap and introduce threaded flush * hex-trace: remove redundant checks for ctx != null * hex-l2: allocate entire context as one buffer and l2fetch it after big flushes * hex-l2: disable tensor clearing in binary and rope for now seems to cause issues with fusion * hmx-mm: update act proc to use fastdivs and fix DMA overflow * hmx-mm: make MUL_MAT_ID kernels robust to multi-chunk cases (start_row>0) * hex-queue: remove obsolete queue interfaces and flush hmx-queue at the end of the op-batch * hex-queue: dont use early wakeup for small op-batches * hex-tensors: properly cap max_tensors in op-batches and dirty_map * hex-l2: make sure threaded l2flush does proper rounding * hex-l2: factor out htp_tensor_flush for reuse (if needed) * hex-l2: optimize tensor flushes by coalescing flush-all * hex-l2: optimize multi-threaded flush * hex-drv: futureproof version checks * hexagon: fix errors and warnings on windows * hex-main: update main thread to only use dspqueue_read, dspqueue_peek is not available on some platforms * hex-main: add fallback mode for dspqueue with callbacks * hex-main: introduce fallback mode for using dspqueue callbacks for full op processing * hex-main: remove early wakeup, not helping and seems to cause some errors with certain batch sizes * hex-l2: make sure to use invalidate version of flushall * hex-l2: dont try to trace early l2flush at the start of op-batch * hex-main: remove offset_ctx that must be zero anyway * hex-hmx: fix hmx_queue_depth to use idx_write - idx_read * hex-hmx: use atomic_load for idx_read/write * hex-main: add static assert to make sure n_threads are aligned
This commit is contained in:
@@ -95,6 +95,8 @@ struct htp_mm_kernel_params {
|
||||
struct fastdiv_values div_r2;
|
||||
struct fastdiv_values div_r3;
|
||||
struct fastdiv_values div_ne11;
|
||||
struct fastdiv_values div_n_act_threads;
|
||||
struct fastdiv_values div_ne00_padded;
|
||||
};
|
||||
|
||||
#if defined(__cplusplus)
|
||||
@@ -643,6 +645,136 @@ static inline size_t htp_mm_hmx_get_batched_vtcm_size(
|
||||
return L.total_bytes;
|
||||
}
|
||||
|
||||
static inline bool htp_mm_hmx_solve_batched_params(
|
||||
int wtype,
|
||||
uint32_t k,
|
||||
uint32_t ne01_padded,
|
||||
uint32_t ne11,
|
||||
uint32_t group_size,
|
||||
bool use_dma_activation,
|
||||
int n_threads,
|
||||
bool pipeline,
|
||||
size_t vtcm_budget,
|
||||
size_t * m_chunk_out,
|
||||
size_t * n_chunk_out,
|
||||
int * act_threads_out,
|
||||
size_t * vtcm_size_out
|
||||
) {
|
||||
size_t best_mblocks = SIZE_MAX;
|
||||
int best_act_threads = 0;
|
||||
size_t best_m_chunk = 0;
|
||||
size_t best_n_chunk = 0;
|
||||
size_t best_vtcm_size = 0;
|
||||
|
||||
int act_threads = n_threads;
|
||||
while (act_threads >= 1) {
|
||||
size_t group_overhead = 256;
|
||||
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);
|
||||
|
||||
size_t m_chunk_candidate = 0;
|
||||
size_t n_chunk_candidate = 0;
|
||||
size_t vtcm_size_candidate = 0;
|
||||
|
||||
if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
|
||||
(size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
|
||||
&m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
|
||||
size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads);
|
||||
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)) {
|
||||
best_mblocks = mblocks;
|
||||
best_act_threads = act_threads;
|
||||
best_m_chunk = m_chunk_candidate;
|
||||
best_n_chunk = n_chunk_candidate;
|
||||
best_vtcm_size = exact_size;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (act_threads == 1) {
|
||||
act_threads = 0;
|
||||
} else {
|
||||
act_threads /= 2;
|
||||
}
|
||||
}
|
||||
|
||||
if (best_act_threads > 0) {
|
||||
*m_chunk_out = best_m_chunk;
|
||||
*n_chunk_out = best_n_chunk;
|
||||
*vtcm_size_out = best_vtcm_size;
|
||||
*act_threads_out = best_act_threads;
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
static inline bool htp_mm_hmx_solve_2d_params(
|
||||
int wtype,
|
||||
uint32_t k,
|
||||
uint32_t m_id_rows,
|
||||
uint32_t ne01_padded,
|
||||
uint32_t ne11_padded,
|
||||
uint32_t m_for_cost,
|
||||
int n_threads,
|
||||
bool pipeline,
|
||||
bool is_matmul_id,
|
||||
uint32_t aligned_tile_size,
|
||||
size_t vtcm_budget,
|
||||
size_t * m_chunk_out,
|
||||
size_t * n_chunk_out,
|
||||
int * act_threads_out,
|
||||
size_t * vtcm_size_out
|
||||
) {
|
||||
size_t best_mblocks = SIZE_MAX;
|
||||
int best_act_threads = 0;
|
||||
size_t best_m_chunk = 0;
|
||||
size_t best_n_chunk = 0;
|
||||
size_t best_vtcm_size = 0;
|
||||
|
||||
const int m_for_chunks = is_matmul_id ? hex_align_up(m_id_rows, 32) : ne11_padded;
|
||||
|
||||
int act_threads = n_threads;
|
||||
while (act_threads >= 1) {
|
||||
size_t simple_2d_overhead = 256;
|
||||
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);
|
||||
|
||||
size_t m_chunk_candidate = 0;
|
||||
size_t n_chunk_candidate = 0;
|
||||
size_t vtcm_size_candidate = 0;
|
||||
|
||||
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);
|
||||
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)) {
|
||||
best_mblocks = mblocks;
|
||||
best_act_threads = act_threads;
|
||||
best_m_chunk = m_chunk_candidate;
|
||||
best_n_chunk = n_chunk_candidate;
|
||||
best_vtcm_size = exact_size;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (act_threads == 1) {
|
||||
act_threads = 0;
|
||||
} else {
|
||||
act_threads /= 2;
|
||||
}
|
||||
}
|
||||
|
||||
if (best_act_threads > 0) {
|
||||
*m_chunk_out = best_m_chunk;
|
||||
*n_chunk_out = best_n_chunk;
|
||||
*vtcm_size_out = best_vtcm_size;
|
||||
*act_threads_out = best_act_threads;
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user