diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 827c772ff..5db8fc84b 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -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 t_map; // tensor ptr to index std::unordered_multimap d_map; // tensor data to index - struct tensor_range { - uint64_t start; - uint64_t end; - int bi; - std::vector tensors; - }; - std::vector 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; diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.c b/ggml/src/ggml-hexagon/htp/dma-queue.c index 8f60a2c9d..4beded1de 100644 --- a/ggml/src/ggml-hexagon/htp/dma-queue.c +++ b/ggml/src/ggml-hexagon/htp/dma-queue.c @@ -101,6 +101,4 @@ void dma_queue_alias_free(dma_queue_t q) { (void) q; } -void dma_queue_flush(dma_queue_t q) { - while (dma_queue_pop(q).dst != NULL) ; -} + diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.h b/ggml/src/ggml-hexagon/htp/dma-queue.h index dba452cfb..264284bda 100644 --- a/ggml/src/ggml-hexagon/htp/dma-queue.h +++ b/ggml/src/ggml-hexagon/htp/dma-queue.h @@ -106,7 +106,7 @@ struct dma_queue_s { bool alias; // When set, dma_queue_delete will not free the ring }; -void dma_queue_flush(dma_queue_t q); + size_t dma_queue_sizeof(size_t capacity); size_t dma_queue_alignof(void); @@ -154,7 +154,6 @@ static inline bool dma_is_vtcm(const dma_queue * q, const void * ptr) { static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t size) { dma_ring * r = q->ring; if (((r->push_idx + 1) & r->idx_mask) == r->pop_idx) { - FARF(HIGH, "dma-push: queue full\n"); return false; } @@ -165,6 +164,8 @@ static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t r->dptr[r->push_idx] = dptr; + htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx); + if (size) { desc->next = NULL; desc->desc_size = 0; // 1D mode @@ -173,7 +174,6 @@ static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t desc->order = 0; desc->done = 0; - htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx); dmlink(r->tail, desc); r->tail = (dma_descriptor_2d *) desc; } else { @@ -188,7 +188,6 @@ static inline bool dma_queue_push_single_1d(dma_queue * q, dma_ptr dptr, size_t static inline bool dma_queue_push_single_2d(dma_queue * q, dma_ptr dptr, size_t dst_stride, size_t src_stride, size_t row_size, size_t nrows) { dma_ring * r = q->ring; if (((r->push_idx + 1) & r->idx_mask) == r->pop_idx) { - FARF(HIGH, "dma-push: queue full\n"); return false; } @@ -224,8 +223,9 @@ static inline bool dma_queue_push_single_2d(dma_queue * q, dma_ptr dptr, size_t r->dptr[r->push_idx] = dptr; + htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx); + if (nrows) { - htp_trace_event_start(r->trace, HTP_TRACE_EVT_DMA, r->push_idx); dmlink(r->tail, desc); r->tail = desc; } else { @@ -252,10 +252,11 @@ static inline dma_ptr dma_queue_pop(dma_queue * q) { dmpoll(); } } - htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx); dptr = r->dptr[r->pop_idx]; + htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx); + r->pop_idx = (r->pop_idx + 1) & r->idx_mask; return dptr; } @@ -270,6 +271,8 @@ static inline dma_ptr dma_queue_pop_nowait(dma_queue * q) { dptr = r->dptr[r->pop_idx]; + htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx); + r->pop_idx = (r->pop_idx + 1) & r->idx_mask; return dptr; } @@ -278,6 +281,10 @@ static inline bool dma_queue_empty(dma_queue * q) { return q->ring->push_idx == q->ring->pop_idx; } +static inline void dma_queue_flush(dma_queue * q) { + while (dma_queue_pop(q).dst != NULL) ; +} + static inline uint32_t dma_queue_depth(dma_queue * q) { return (q->ring->push_idx - q->ring->pop_idx) & q->ring->idx_mask; } @@ -314,14 +321,18 @@ static inline bool dma_queue_push(dma_queue *q, dma_ptr dptr, size_t dst_stride, { const uint8_t *src = (const uint8_t *) dptr.src; uint8_t *dst = (uint8_t *) dptr.dst; - for (size_t r = 0; r < nrows; ++r) { + size_t r = 0; + while (r + 1 < nrows) { dma_ptr p = dma_make_ptr(dst + r * dst_stride, src + r * src_stride); - if (!dma_queue_push_single_1d(q, p, row_size)) - return false; - if (r + 1 < nrows) - dma_queue_pop(q); + if (!dma_queue_push_single_1d(q, p, row_size)) { + dma_queue_flush(q); + } else { + r++; + } } - return true; + dma_queue_flush(q); + dma_ptr p = dma_make_ptr(dst + r * dst_stride, src + r * src_stride); + return dma_queue_push_single_1d(q, p, row_size); } } diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c index a3bd0bf8f..fe78718c6 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c @@ -123,15 +123,17 @@ struct hmx_fa_context { uint32_t g_br; // hex_align_up(G * Br, 32) - actual tile row dim // VTCM buffers (allocated by vtcm_seq_alloc) + __fp16 * vtcm_q_dma; // Q DMA fetch buffer __fp16 * vtcm_q_tiles; // Q tile format [g_br, D] __fp16 * vtcm_o_tiles[2]; // O ping-pong [g_br, D] __fp16 * vtcm_k_fp16[2]; // K DMA double-buffer [Bc, D] __fp16 * vtcm_v_fp16[2]; // V DMA double-buffer [Bc, D] - __fp16 * vtcm_k_tiles; // K tiles (transposed) + __fp16 * vtcm_k_tiles[2]; // K tiles (transposed, double-buffered) __fp16 * vtcm_v_tiles[2]; // V tiles (column-major, double-buffered) - __fp16 * vtcm_s_tiles; // S = QK^T [g_br, Bc] - __fp16 * vtcm_p_tiles; // P = softmax(S) [g_br, Bc] + __fp16 * vtcm_s_tiles[2]; // S = QK^T [g_br, Bc] (double-buffered) + __fp16 * vtcm_p_tiles[2]; // P = softmax(S) [g_br, Bc] __fp16 * vtcm_d_tiles; // Diagonal rescale [g_br, g_br] + __fp16 * vtcm_d_inv_l; // Diagonal rescale (1/l) [g_br, g_br] HVX_Vector * vtcm_m_vec; // Row max [g_br] HVX_Vector * vtcm_l_vec; // Row sum [g_br] HVX_Vector * vtcm_s_rowmax; // Softmax intermediate [g_br] @@ -236,10 +238,6 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * const uint32_t iv3 = fastdiv(iq3, &factx->broadcast_rv3); const uint32_t iv2 = fastdiv(iq2, &factx->broadcast_rv2); - // Fetch Q row - const uint8_t * q_row_ptr = (const uint8_t *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3); - dma_queue_push(dma, dma_make_ptr(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1); - const __fp16 * mp_base = NULL; if (mask) { const uint32_t im2 = fastmodulo(iq2, mask->ne[2], &factx->src3_div2); @@ -247,26 +245,91 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * mp_base = (const __fp16 *) ((const uint8_t *) mask->data + iq1*mask->nb[1] + im2*mask->nb[2] + im3*mask->nb[3]); } - // Prefetch first two blocks - for (uint32_t ib = 0; ib < MIN(factx->n_blocks, 2); ++ib) { - const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE; - const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start); + // Precalculate next row variables if there is a next row + bool has_next_ir = (ir + 1 < ir1); + uint32_t next_ik2 = 0, next_ik3 = 0, next_iv2 = 0, next_iv3 = 0; + const uint8_t * next_q_row_ptr = NULL; + const __fp16 * next_mp_base = NULL; - // K - const uint8_t * k_src = (const uint8_t *) k->data + (ic_start*nbk1 + ik2*nbk2 + ik3*nbk3); - uint8_t * k_dst = spad_k + (ib % 2) * factx->size_k_block; - dma_queue_push(dma, dma_make_ptr(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size); + const uint8_t * next_k_src0 = NULL; + const uint8_t * next_v_src0 = NULL; + const uint8_t * next_m_src0 = NULL; + uint32_t next_block_size0 = 0; - // V - const uint8_t * v_src = (const uint8_t *) v->data + (ic_start*nbv1 + iv2*nbv2 + iv3*nbv3); - uint8_t * v_dst = spad_v + (ib % 2) * factx->size_v_block; - dma_queue_push(dma, dma_make_ptr(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size); + const uint8_t * next_k_src1 = NULL; + const uint8_t * next_v_src1 = NULL; + const uint8_t * next_m_src1 = NULL; + uint32_t next_block_size1 = 0; + + if (has_next_ir) { + const uint32_t next_ir = ir + 1; + const uint32_t next_iq3 = fastdiv(next_ir, &factx->src0_div21); + const uint32_t next_iq2 = fastdiv(next_ir - next_iq3*neq2*neq1, &factx->src0_div1); + const uint32_t next_iq1 = (next_ir - next_iq3*neq2*neq1 - next_iq2 * neq1); + + next_ik3 = fastdiv(next_iq3, &factx->broadcast_rk3); + next_ik2 = fastdiv(next_iq2, &factx->broadcast_rk2); + + next_iv3 = fastdiv(next_iq3, &factx->broadcast_rv3); + next_iv2 = fastdiv(next_iq2, &factx->broadcast_rv2); + + next_q_row_ptr = (const uint8_t *) q->data + (next_iq1*nbq1 + next_iq2*nbq2 + next_iq3*nbq3); - // Mask if (mask) { - const uint8_t * m_src = (const uint8_t *) (mp_base + ic_start); - // Mask is 1D contiguous for this row - dma_cache_push(dma, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1); + const uint32_t next_im2 = fastmodulo(next_iq2, mask->ne[2], &factx->src3_div2); + const uint32_t next_im3 = fastmodulo(next_iq3, mask->ne[3], &factx->src3_div3); + next_mp_base = (const __fp16 *) ((const uint8_t *) mask->data + next_iq1*mask->nb[1] + next_im2*mask->nb[2] + next_im3*mask->nb[3]); + } + + // Precalculate next K/V block 0 source pointers + { + const uint32_t ic_start = 0; + next_block_size0 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start); + next_k_src0 = (const uint8_t *) k->data + (ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3); + next_v_src0 = (const uint8_t *) v->data + (ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3); + if (mask) { + next_m_src0 = (const uint8_t *) (next_mp_base + ic_start); + } + } + + // Precalculate next K/V block 1 source pointers (if n_blocks > 1) + if (factx->n_blocks > 1) { + const uint32_t ic_start = 1 * FLASH_ATTN_BLOCK_SIZE; + next_block_size1 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start); + next_k_src1 = (const uint8_t *) k->data + (ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3); + next_v_src1 = (const uint8_t *) v->data + (ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3); + if (mask) { + next_m_src1 = (const uint8_t *) (next_mp_base + ic_start); + } + } + } + + if (ir == ir0) { + // Fetch Q row + const uint8_t * q_row_ptr = (const uint8_t *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3); + dma_queue_push(dma, dma_make_ptr(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1); + + // Prefetch first two blocks + for (uint32_t ib = 0; ib < MIN(factx->n_blocks, 2); ++ib) { + const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE; + const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start); + + // K + const uint8_t * k_src = (const uint8_t *) k->data + (ic_start*nbk1 + ik2*nbk2 + ik3*nbk3); + uint8_t * k_dst = spad_k + (ib % 2) * factx->size_k_block; + dma_queue_push(dma, dma_make_ptr(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size); + + // V + const uint8_t * v_src = (const uint8_t *) v->data + (ic_start*nbv1 + iv2*nbv2 + iv3*nbv3); + uint8_t * v_dst = spad_v + (ib % 2) * factx->size_v_block; + dma_queue_push(dma, dma_make_ptr(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size); + + // Mask + if (mask) { + const uint8_t * m_src = (const uint8_t *) (mp_base + ic_start); + // Mask is 1D contiguous for this row + dma_cache_push(dma, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1); + } } } @@ -287,6 +350,11 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * const HVX_Vector slope_vec = hvx_vec_splat_f16(slope); const HVX_Vector v_neg_inf = Q6_Vh_vsplat_R(0xfbff); + const HVX_Vector v_cap = (factx->logit_softcap != 0.0f) ? hvx_vec_splat_f16(factx->logit_softcap) : Q6_V_vzero(); + const HVX_Vector vinf = Q6_Vh_vsplat_R(0xFC00); + const HVX_Vector vmin = Q6_Vh_vsplat_R(0xFBFF); + const HVX_Vector v_log2e = hvx_vec_splat_f16(EXP_LOG2E_F); + const uint32_t stride_v2 = factx->size_v_row_padded * 2; for (uint32_t ib = 0; ib < factx->n_blocks; ++ib) { const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE; const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start); @@ -309,7 +377,6 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * // 2. Softcap (in FP16) if (factx->logit_softcap != 0.0f) { - const HVX_Vector v_cap = hvx_vec_splat_f16(factx->logit_softcap); scores_f16 = hvx_vec_tanh_f16(scores_f16); scores_f16 = hvx_vec_mul_f16_f16(scores_f16, v_cap); } @@ -319,8 +386,6 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * // 3. Mask (in FP16) if (mask) { HVX_Vector m_vals_f16 = *(const HVX_UVector *) m_base; - HVX_Vector vinf = Q6_Vh_vsplat_R(0xFC00); - HVX_Vector vmin = Q6_Vh_vsplat_R(0xFBFF); HVX_VectorPred is_inf = Q6_Q_vcmp_eq_VhVh(m_vals_f16, vinf); m_vals_f16 = Q6_V_vmux_QVV(is_inf, vmin, m_vals_f16); @@ -335,10 +400,30 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * HVX_Vector v_max = Q6_V_lo_W(hvx_vec_f16_to_f32(v_max_f16)); // splat block max in FP32 htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_QK, ir); + if (ib + 1 == factx->n_blocks && has_next_ir) { + // Queue next row's Q row! + dma_queue_push(dma, dma_make_ptr(spad_q, next_q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1); + + if (factx->n_blocks % 2 == 0) { + // Queue next row's block 0 (into buffer slot 0) + uint8_t * k_dst = spad_k + 0 * factx->size_k_block; + uint8_t * v_dst = spad_v + 0 * factx->size_v_block; + + // K (block 0 of next row) + dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0); + + // V (block 0 of next row) + dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0); + + // Mask (block 0 of next row) + if (mask) { + dma_cache_push(dma, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1); + } + } + } + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir); { - const HVX_Vector v_log2e = hvx_vec_splat_f16(EXP_LOG2E_F); - // 4. Online Softmax Update HVX_Vector M_new_vec = Q6_Vsf_vmax_VsfVsf(v_max, M_vec); HVX_Vector diff_vec = HVX_OP_SUB_F32(M_vec, M_new_vec); @@ -370,24 +455,20 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * S_vec = HVX_OP_ADD_F32(HVX_OP_MUL_F32(S_vec, ms_vec), p_sum_vec); // 5. Accumulate V (F16 * F16 -> F32 accumulator) - __fp16 __attribute__((aligned(128))) p_arr[VLEN_FP16]; - hvx_vec_store_a(p_arr, 128, P); + const uint8_t * v_ptr = v_base; for (uint32_t j = 0; j < current_block_size; j += 2) { if (j + 1 == current_block_size) { - if (p_arr[j] != 0.0f) { - const uint8_t * v_ptr = v_base + j * factx->size_v_row_padded; - hvx_mad_f32_f16_aa(VKQ32, v_ptr, (p_arr + j), DV); - } + HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2)); + hvx_mad_f32_f16_aa_vec(VKQ32, v_ptr, S0, DV); break; } - if (p_arr[j] == 0.0f && p_arr[j + 1] == 0.0f) { - continue; - } + HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2)); + HVX_Vector S1 = hvx_vec_repl_f16(Q6_V_vror_VR(P, (j + 1) * 2)); - const uint8_t * v_ptr = v_base + j * factx->size_v_row_padded; - hvx_mad_f32_f16_aa_rx2(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, (p_arr + j), (p_arr + j + 1), DV); + hvx_mad_f32_f16_aa_rx2_vec(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, S0, S1, DV); + v_ptr += stride_v2; } } htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir); @@ -414,6 +495,61 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * } } + if (has_next_ir) { + if (factx->n_blocks % 2 == 0) { + // Queue next row's block 1 (into buffer slot 1, if n_blocks > 1) + if (factx->n_blocks > 1) { + uint8_t * k_dst = spad_k + 1 * factx->size_k_block; + uint8_t * v_dst = spad_v + 1 * factx->size_v_block; + + // K (block 1 of next row) + dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1); + + // V (block 1 of next row) + dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1); + + // Mask (block 1 of next row) + if (mask) { + dma_cache_push(dma, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1); + } + } + } else { + // Queue next row's block 0 (into buffer slot 0) + { + uint8_t * k_dst = spad_k + 0 * factx->size_k_block; + uint8_t * v_dst = spad_v + 0 * factx->size_v_block; + + // K (block 0 of next row) + dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0); + + // V (block 0 of next row) + dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0); + + // Mask (block 0 of next row) + if (mask) { + dma_cache_push(dma, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1); + } + } + + // Queue next row's block 1 (into buffer slot 1, if n_blocks > 1) + if (factx->n_blocks > 1) { + uint8_t * k_dst = spad_k + 1 * factx->size_k_block; + uint8_t * v_dst = spad_v + 1 * factx->size_v_block; + + // K (block 1 of next row) + dma_queue_push(dma, dma_make_ptr(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1); + + // V (block 1 of next row) + dma_queue_push(dma, dma_make_ptr(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1); + + // Mask (block 1 of next row) + if (mask) { + dma_cache_push(dma, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1); + } + } + } + } + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, ir); // sinks float M = hvx_vec_get_f32(M_vec); @@ -471,6 +607,7 @@ typedef struct { void * curr_k; uint32_t kv_start; uint32_t rows_per_t; + size_t buf_idx; } fa_k_int_args_t; static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data) { @@ -488,19 +625,19 @@ static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data) struct htp_thread_trace * tr = &factx->octx->ctx->trace[i]; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start)); - hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles, (const __fp16 *) args->curr_k, total_rows, factx->DK, + hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK, args->src_stride, start, end); htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start)); } -static void fa_phase_k_interleave(struct hmx_fa_context * factx, uint32_t kv_rows, size_t src_stride, void * curr_k, uint32_t kv_start) { +static void fa_phase_k_interleave(struct hmx_fa_context * factx, uint32_t kv_rows, size_t src_stride, void * curr_k, uint32_t kv_start, size_t buf_idx) { work_queue_t wp = factx->octx->ctx->work_queue; uint32_t n = 1; if (factx->n_threads > 1 && kv_rows >= factx->n_threads * 2) { n = factx->n_threads; } uint32_t rows_per_t = hex_align_up(hmx_ceil_div(kv_rows, n), 2); - fa_k_int_args_t args = { factx, kv_rows, src_stride, curr_k, kv_start, rows_per_t }; + fa_k_int_args_t args = { factx, kv_rows, src_stride, curr_k, kv_start, rows_per_t, buf_idx }; if (n > 1) { work_queue_run(wp, fa_k_interleave_thread, &args, n); } else { @@ -645,12 +782,13 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) { } } - // Initialize vtcm_d_tiles to 0 + // Initialize vtcm_d_tiles and vtcm_d_inv_l to 0 const size_t d_bytes_per_t = hex_align_up(d_tile_bytes / n, 128); const size_t d_start = i * d_bytes_per_t; const size_t d_end = hex_smin(d_start + d_bytes_per_t, d_tile_bytes); if (d_start < d_tile_bytes) { hvx_splat_u8_a((char *) factx->vtcm_d_tiles + d_start, 0, d_end - d_start); + hvx_splat_u8_a((char *) factx->vtcm_d_inv_l + d_start, 0, d_end - d_start); } } @@ -662,15 +800,14 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) { assert(factx->DK == factx->DV); - const size_t o_tile_bytes = factx->o_tile_bytes; - const bool use_q_dma = (2 * o_tile_bytes >= factx->g_br * DK * (factx->is_q_fp32 ? 4 : 2)); + const bool use_q_dma = (factx->vtcm_q_dma != NULL); __fp16 * q_tiles = factx->vtcm_q_tiles; if (use_q_dma) { const size_t g_rows_end = hex_smin(end, n_rows_g); const uint32_t d_limit = factx->is_q_fp32 ? DK / 32 : DK / 64; - uint8_t * q_flat = (uint8_t *) factx->vtcm_o_tiles[0]; + uint8_t * q_flat = (uint8_t *) factx->vtcm_q_dma; if (factx->is_q_fp32) { switch (d_limit) { case 2: hmx_fa_q_prep_fp32_d2(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break; @@ -781,10 +918,10 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) { const uint32_t kv_head = args->kv_head; const uint32_t ib3 = args->ib3; - for (size_t r = start; r < end; ++r) { - const size_t q_idx = fastdiv(r, &factx->div_G); - const size_t h_idx = fastmodulo(r, G, &factx->div_G); + size_t q_idx = fastdiv(start, &factx->div_G); + size_t h_idx = fastmodulo(start, G, &factx->div_G); + for (size_t r = start; r < end; ++r) { float * out = (float *) ((uint8_t *) dst->data + (kv_head * G + h_idx) * dst->nb[1] + (q_start + q_idx) * dst->nb[2] + ib3 * dst->nb[3]); @@ -801,6 +938,12 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) { *(HVX_UVector *) (out + d * 32) = Q6_V_hi_W(vp); } } + + h_idx++; + if (h_idx == G) { + h_idx = 0; + q_idx++; + } } htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start)); } @@ -829,10 +972,10 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) { const uint32_t kv_head = args->kv_head; const uint32_t ib3 = args->ib3; - for (size_t r = start; r < end; ++r) { - const size_t q_idx = fastdiv(r, &factx->div_G); - const size_t h_idx = fastmodulo(r, G, &factx->div_G); + size_t q_idx = fastdiv(start, &factx->div_G); + size_t h_idx = fastmodulo(start, G, &factx->div_G); + for (size_t r = start; r < end; ++r) { __fp16 * out = (__fp16 *) ((uint8_t *) dst->data + (kv_head * G + h_idx) * dst->nb[1] + (q_start + q_idx) * dst->nb[2] + ib3 * dst->nb[3]); @@ -851,6 +994,12 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) { *(HVX_UVector *) (out + d * 64) = Q6_V_hi_W(vp); } } + + h_idx++; + if (h_idx == G) { + h_idx = 0; + q_idx++; + } } htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start)); } @@ -879,6 +1028,7 @@ static void fa_phase_o_store(struct hmx_fa_context * factx, typedef struct { struct hmx_fa_context * factx; + size_t buf_idx; size_t kv_rows; size_t n_rows_g; size_t n_col_tiles; @@ -960,8 +1110,8 @@ static inline void fa_softmax_impl( uint32_t r0 = r / HMX_FP16_TILE_N_ROWS; uint32_t r1 = r % HMX_FP16_TILE_N_ROWS; - const __fp16 * s_ld_base = factx->vtcm_s_tiles + r0 * HMX_FP16_TILE_N_ROWS * Bc; - __fp16 * p_st_base = factx->vtcm_p_tiles + r0 * HMX_FP16_TILE_N_ROWS * Bc; + const __fp16 * s_ld_base = factx->vtcm_s_tiles[args->buf_idx] + r0 * HMX_FP16_TILE_N_ROWS * Bc; + __fp16 * p_st_base = factx->vtcm_p_tiles[args->buf_idx] + r0 * HMX_FP16_TILE_N_ROWS * Bc; // Decode 2 rows from S tiles into per-thread row buffers if (has_softcap) { @@ -983,7 +1133,26 @@ static inline void fa_softmax_impl( my_row_buf1[ci] = hvx_vec_mul_f16_f16(t1, v_cap); } } else { - for (size_t c = 0; c < kv_rows; c += 64) { + size_t c = 0; + for (; c + 64 < kv_rows; c += 128) { + size_t ci0 = c / 64; + size_t ci1 = ci0 + 1; + const __fp16 * in_dtile0 = s_ld_base + ci0 * HMX_FP16_TILE_N_ELMS * 2; + const __fp16 * in_dtile1 = s_ld_base + ci1 * HMX_FP16_TILE_N_ELMS * 2; + const HVX_Vector * pv_s_in0_0 = ((const HVX_Vector *) in_dtile0) + r1 / 2; + const HVX_Vector * pv_s_in1_0 = pv_s_in0_0 + 16; + const HVX_Vector * pv_s_in0_1 = ((const HVX_Vector *) in_dtile1) + r1 / 2; + const HVX_Vector * pv_s_in1_1 = pv_s_in0_1 + 16; + + HVX_VectorPair vp_s_drow0 = Q6_W_vdeal_VVR(*pv_s_in1_0, *pv_s_in0_0, -2); + my_row_buf0[ci0] = Q6_V_lo_W(vp_s_drow0); + my_row_buf1[ci0] = Q6_V_hi_W(vp_s_drow0); + + HVX_VectorPair vp_s_drow1 = Q6_W_vdeal_VVR(*pv_s_in1_1, *pv_s_in0_1, -2); + my_row_buf0[ci1] = Q6_V_lo_W(vp_s_drow1); + my_row_buf1[ci1] = Q6_V_hi_W(vp_s_drow1); + } + for (; c < kv_rows; c += 64) { size_t ci = c / 64; const __fp16 * in_dtile = s_ld_base + ci * HMX_FP16_TILE_N_ELMS * 2; const HVX_Vector * pv_s_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2; @@ -1007,12 +1176,12 @@ static inline void fa_softmax_impl( HVX_Vector v_s_rowmax0 = v_neg_inf; HVX_Vector v_s_rowmax1 = v_neg_inf; - for (size_t c = 0; c < kv_rows; c += 64) { - size_t ci = c / 64; - const size_t ne = hex_smin(kv_rows - c, 64); - HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(ne * sizeof(__fp16)); + if (has_mask) { + for (size_t c = 0; c < kv_rows; c += 64) { + size_t ci = c / 64; + const size_t ne = hex_smin(kv_rows - c, 64); + HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(ne * sizeof(__fp16)); - if (has_mask) { HVX_Vector v_mask0, v_mask1; if (mask_broadcast) { @@ -1066,15 +1235,31 @@ static inline void fa_softmax_impl( my_row_buf0[ci] = Q6_V_vmux_QVV(q_keep0, hvx_vec_add_f16_f16(my_row_buf0[ci], v_mask0_scaled), v_neg_inf); my_row_buf1[ci] = Q6_V_vmux_QVV(q_keep1, hvx_vec_add_f16_f16(my_row_buf1[ci], v_mask1_scaled), v_neg_inf); } - } else { + + v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci]); + v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci]); + } + } else { + size_t c = 0; + for (; c + 64 < kv_rows; c += 128) { + size_t ci0 = c / 64; + size_t ci1 = ci0 + 1; + v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci0]); + v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci0]); + v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci1]); + v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci1]); + } + for (; c < kv_rows; c += 64) { + size_t ci = c / 64; + const size_t ne = hex_smin(kv_rows - c, 64); + HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(ne * sizeof(__fp16)); if (ne < 64) { my_row_buf0[ci] = Q6_V_vmux_QVV(q_tail_keep, my_row_buf0[ci], v_neg_inf); my_row_buf1[ci] = Q6_V_vmux_QVV(q_tail_keep, my_row_buf1[ci], v_neg_inf); } + v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci]); + v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci]); } - - v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci]); - v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci]); } v_s_rowmax0 = hvx_vec_reduce_max_f16(v_s_rowmax0); @@ -1121,8 +1306,48 @@ static inline void fa_softmax_impl( HVX_Vector v_p_rowsum0 = v_zero; HVX_Vector v_p_rowsum1 = v_zero; - for (size_t c = 0; c < kv_rows; c += 64) { - size_t ci = c / 64; + size_t c = 0; + for (; c + 64 < kv_rows; c += 128) { + size_t ci0 = c / 64; + size_t ci1 = ci0 + 1; + + HVX_Vector v_s_minus_m0_0 = Q6_Vqf16_vsub_VhfVhf(my_row_buf0[ci0], v_dup_m0); + HVX_Vector v_s_minus_m1_0 = Q6_Vqf16_vsub_VhfVhf(my_row_buf1[ci0], v_dup_m1); + HVX_Vector v_s_minus_m0_1 = Q6_Vqf16_vsub_VhfVhf(my_row_buf0[ci1], v_dup_m0); + HVX_Vector v_s_minus_m1_1 = Q6_Vqf16_vsub_VhfVhf(my_row_buf1[ci1], v_dup_m1); + + HVX_Vector v_p_row0_hf_0 = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m0_0)); + HVX_Vector v_p_row1_hf_0 = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m1_0)); + HVX_Vector v_p_row0_hf_1 = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m0_1)); + HVX_Vector v_p_row1_hf_1 = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m1_1)); + + __fp16 * out_dtile0 = p_st_base + ci0 * HMX_FP16_TILE_N_ELMS * 2; + __fp16 * out_dtile1 = p_st_base + ci1 * HMX_FP16_TILE_N_ELMS * 2; + HVX_Vector * pv_p_out0_0 = ((HVX_Vector *) out_dtile0) + r1 / 2; + HVX_Vector * pv_p_out1_0 = pv_p_out0_0 + 16; + HVX_Vector * pv_p_out0_1 = ((HVX_Vector *) out_dtile1) + r1 / 2; + HVX_Vector * pv_p_out1_1 = pv_p_out0_1 + 16; + + HVX_VectorPair vp_p_dual0 = Q6_W_vshuff_VVR(v_p_row1_hf_0, v_p_row0_hf_0, -2); + *pv_p_out0_0 = Q6_V_lo_W(vp_p_dual0); + *pv_p_out1_0 = Q6_V_hi_W(vp_p_dual0); + + HVX_VectorPair vp_p_dual1 = Q6_W_vshuff_VVR(v_p_row1_hf_1, v_p_row0_hf_1, -2); + *pv_p_out0_1 = Q6_V_lo_W(vp_p_dual1); + *pv_p_out1_1 = Q6_V_hi_W(vp_p_dual1); + + HVX_VectorPair vp_p0_0 = hvx_vec_f16_to_f32_shuff(v_p_row0_hf_0); + HVX_VectorPair vp_p1_0 = hvx_vec_f16_to_f32_shuff(v_p_row1_hf_0); + HVX_VectorPair vp_p0_1 = hvx_vec_f16_to_f32_shuff(v_p_row0_hf_1); + HVX_VectorPair vp_p1_1 = hvx_vec_f16_to_f32_shuff(v_p_row1_hf_1); + + v_p_rowsum0 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum0, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p0_0), Q6_V_hi_W(vp_p0_0))); + v_p_rowsum0 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum0, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p0_1), Q6_V_hi_W(vp_p0_1))); + v_p_rowsum1 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum1, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p1_0), Q6_V_hi_W(vp_p1_0))); + v_p_rowsum1 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum1, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p1_1), Q6_V_hi_W(vp_p1_1))); + } + for (size_t c_rem = c; c_rem < kv_rows; c_rem += 64) { + size_t ci = c_rem / 64; HVX_Vector v_s_minus_m0 = Q6_Vqf16_vsub_VhfVhf(my_row_buf0[ci], v_dup_m0); HVX_Vector v_s_minus_m1 = Q6_Vqf16_vsub_VhfVhf(my_row_buf1[ci], v_dup_m1); @@ -1281,7 +1506,7 @@ static __attribute__((noinline)) void fa_build_d_diag_inv_l(struct hmx_fa_contex v_content = Q6_V_vror_VR(v_content, 64); } - __fp16 * out_base = factx->vtcm_d_tiles + i * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + __fp16 * out_base = factx->vtcm_d_inv_l + i * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content); } } @@ -1514,6 +1739,27 @@ static void fa_pop_mask_dma_gqa(dma_queue * dma, uint32_t G) { } } +static inline void fa_prefetch_block(dma_queue * dma, const struct htp_tensor * k, const struct htp_tensor * v, const struct htp_tensor * mask, + uint32_t b, size_t Bc, size_t size_k_row_padded, size_t size_k_row, size_t size_v_row_padded, size_t size_v_row, + uint32_t ik2, uint32_t ik3, uint32_t iv2, uint32_t iv3, uint32_t q_start, uint32_t im3, uint32_t kv_head, uint32_t G, + size_t m_line_bytes, size_t n_rows_q, size_t nek1, size_t prefetch_buf, struct hmx_fa_context * factx) { + const uint32_t prefetch_start = b * Bc; + const uint32_t prefetch_rows = hex_smin(Bc, nek1 - prefetch_start); + const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx->vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows); + const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx->vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows); + + if (mask) { + if (__builtin_expect(factx->mask_broadcast, true)) { + const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + prefetch_start * sizeof(__fp16); + dma_cache_push(dma, &factx->m_cache, ms_src, m_line_bytes, mask->nb[1], prefetch_rows * sizeof(__fp16), n_rows_q); + } else { + fa_push_mask_dma_gqa(dma, mask, q_start, im3, prefetch_start, kv_head, G, m_line_bytes, prefetch_rows, n_rows_q, factx); + } + } +} + // ============================================================================ // Core HMX flash attention algorithm (GQA-merged) // ============================================================================ @@ -1612,7 +1858,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { // Build the VTCM layout once (shared with the host estimator) and place every // scratch buffer at its computed offset. struct hmx_fa_vtcm_layout L; - hmx_fa_vtcm_layout_build(&L, G, DK, DV, Br, Bc, n_threads, pipeline); + hmx_fa_vtcm_layout_build(&L, G, DK, DV, Br, Bc, n_threads, pipeline, factx.is_q_fp32); if (L.total_bytes > ctx->vtcm_size) { return HTP_STATUS_VTCM_TOO_SMALL; @@ -1620,6 +1866,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { uint8_t * const base = ctx->vtcm_base; + factx.vtcm_q_dma = VTCM_LAYOUT_PTR(__fp16, base, L.off_q_dma); factx.vtcm_q_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_q_tiles); factx.vtcm_o_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_o_tiles[0]); factx.vtcm_o_tiles[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_o_tiles[1]); @@ -1627,12 +1874,16 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { factx.vtcm_k_fp16[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_fp16[1]); factx.vtcm_v_fp16[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_fp16[0]); factx.vtcm_v_fp16[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_fp16[1]); - factx.vtcm_k_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_tiles); + factx.vtcm_k_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_tiles[0]); + factx.vtcm_k_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_k_tiles[1], pipeline); factx.vtcm_v_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_tiles[0]); factx.vtcm_v_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_v_tiles[1], pipeline); - factx.vtcm_s_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_s_tiles); - factx.vtcm_p_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_p_tiles); + factx.vtcm_s_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_s_tiles[0]); + factx.vtcm_s_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_s_tiles[1], pipeline); + factx.vtcm_p_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_p_tiles[0]); + factx.vtcm_p_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_p_tiles[1], pipeline); factx.vtcm_d_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_d_tiles); + factx.vtcm_d_inv_l = VTCM_LAYOUT_PTR(__fp16, base, L.off_d_inv_l); factx.vtcm_m_vec = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_m_vec); factx.vtcm_l_vec = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_l_vec); factx.vtcm_s_rowmax = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_s_rowmax); @@ -1670,6 +1921,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { const size_t qo_element_size = factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16); + const bool q_transposed = q->nb[1] < q->nb[2]; + const size_t q_src_stride = q_transposed ? q->nb[2] : q->nb[1]; + const size_t q_row_bytes_untransposed = factx.G * factx.DK * qo_element_size; + const size_t q_row_bytes_trans_factor = factx.DK * qo_element_size; + const uint32_t kv_rows0 = hex_smin(Bc, nek1); + // ======== Reusable job descriptors for pipeline ======== hmx_fa_qk_job_t qk_job; hmx_fa_o_update_job_t ou_job; @@ -1690,34 +1947,34 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { const uint32_t iv2 = kv_head; const uint32_t iv3 = fastdiv(ib3, &kparams->broadcast_rv3); - // 1. Push Q DMA (if Q DMA is used) - const size_t o_tile_bytes = factx.o_tile_bytes; - const bool use_q_dma = (2 * o_tile_bytes >= factx.g_br * factx.DK * (factx.is_q_fp32 ? 4 : 2)); - if (use_q_dma) { - const bool q_transposed = q->nb[1] < q->nb[2]; - const uint8_t * q_ptr = (const uint8_t *) q->data + q_start * q->nb[1] + (kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3]; - const size_t el_size = factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16); - const size_t q_row_bytes = q_transposed ? n_rows_q * factx.DK * el_size : factx.G * factx.DK * el_size; - const size_t src_stride = q_transposed ? q->nb[2] : q->nb[1]; + // 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 == 0 && kv_head == 0) { + const uint8_t * q_ptr = (const uint8_t *) q->data; + const size_t q_row_bytes = q_transposed ? n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed; const size_t n_rows = q_transposed ? factx.G : n_rows_q; - dma_queue_push(dma, dma_make_ptr(factx.vtcm_o_tiles[0], q_ptr), q_row_bytes, hex_smax(src_stride, q_row_bytes), q_row_bytes, n_rows); + dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, q_ptr), q_row_bytes, hex_smax(q_src_stride, q_row_bytes), q_row_bytes, n_rows); + + if (factx.n_kv_blocks > 0) { + const uint8_t * k_src = (const uint8_t *) k->data + ik2 * k->nb[2] + ik3 * k->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0); + + const uint8_t * v_src = (const uint8_t *) v->data + iv2 * v->nb[2] + iv3 * v->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0); + + if (factx.pipeline && mask) { + if (__builtin_expect(factx.mask_broadcast, true)) { + const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + 0; + dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), n_rows_q); + } else { + fa_push_mask_dma_gqa(dma, mask, q_start, im3, 0, kv_head, G, m_line_bytes, kv_rows0, n_rows_q, &factx); + } + } + } } - // 2. Prefetch first KV block - if (factx.n_kv_blocks > 0) { - const uint32_t kv_rows0 = hex_smin(Bc, nek1); - - const uint8_t * k_src = (const uint8_t *) k->data + ik2 * k->nb[2] + ik3 * k->nb[3]; - dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0); - - const uint8_t * v_src = (const uint8_t *) v->data + iv2 * v->nb[2] + iv3 * v->nb[3]; - dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0); - } - - // 3. Pop Q DMA (blocks until Q is loaded) - if (use_q_dma) { - dma_queue_pop(dma); - } + // 2. Pop Q DMA (blocks until Q is loaded) + dma_queue_pop(dma); // ---- Load Q block & Initialize per-block state ---- fa_phase_q_load(&factx, q, q_start, kv_head, ib3, n_rows_g); @@ -1738,76 +1995,40 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { hmx_queue_t hmx_q = ctx->hmx_queue; if (factx.pipeline) { - // Pipeline path + // Double-buffered job structs because HMX queue runs asynchronously + hmx_fa_qk_job_t qk_job[2]; + hmx_fa_o_update_job_t ou_job[2]; + + // Prefetch block 1 early if there are multiple blocks + if (factx.n_kv_blocks > 1) { + fa_prefetch_block(dma, k, v, mask, 1, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row, + ik2, ik3, iv2, iv3, q_start, im3, kv_head, G, m_line_bytes, n_rows_q, nek1, 1, &factx); + } + + // Prep and start QK-dot(0) + void * curr_k0 = dma_queue_pop(dma).dst; + fa_phase_k_interleave(&factx, kv_rows0, k_src_stride, curr_k0, 0, 0); + + qk_job[0].q_tiles = factx.vtcm_q_tiles; + qk_job[0].k_tiles = factx.vtcm_k_tiles[0]; + qk_job[0].s_tiles = factx.vtcm_s_tiles[0]; + qk_job[0].n_row_tiles = n_row_tiles; + qk_job[0].n_col_tiles = hmx_ceil_div(kv_rows0, HMX_FP16_TILE_N_COLS); + qk_job[0].n_dot_tiles = DK / 32; + qk_job[0].n_tiles_per_bc = n_tiles_per_bc; + qk_job[0].hmx_scales = factx.vtcm_hmx_scales_qk; + hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[0])); + for (uint32_t kv_blk = 0; kv_blk < factx.n_kv_blocks; ++kv_blk) { const uint32_t kv_start = kv_blk * Bc; const uint32_t kv_rows = hex_smin(Bc, nek1 - kv_start); const size_t n_col_tiles = hmx_ceil_div(kv_rows, HMX_FP16_TILE_N_COLS); - // Push mask DMA - if (mask) { - if (__builtin_expect(factx.mask_broadcast, true)) { - const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + kv_start * sizeof(__fp16); - dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q); - } else { - fa_push_mask_dma_gqa(dma, mask, q_start, im3, kv_start, kv_head, G, m_line_bytes, kv_rows, n_rows_q, &factx); - } - } - - // Prefetch next KV block early - if (kv_blk + 1 < factx.n_kv_blocks) { - const uint32_t prefetch_start = (kv_blk + 1) * Bc; - const uint32_t prefetch_rows = hex_smin(Bc, nek1 - prefetch_start); - const size_t prefetch_buf = 1 - buf_idx; - const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3]; - dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows); - const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3]; - dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows); - } - - // ---- Phase 1: K_int ---- - if (kv_blk > 0) { - ou_job.o_curr = o_tile_curr; - ou_job.o_prev = o_tile_prev; - ou_job.p_tiles = factx.vtcm_p_tiles; - ou_job.v_tiles = factx.vtcm_v_tiles[1 - buf_idx]; - ou_job.d_tiles = factx.vtcm_d_tiles; - ou_job.hmx_scales = factx.vtcm_hmx_scales_id; - ou_job.n_row_tiles = n_row_tiles; - ou_job.n_col_tiles = hmx_ceil_div(hex_smin(Bc, nek1 - (kv_blk - 1) * Bc), HMX_FP16_TILE_N_COLS); - ou_job.n_row_tiles_g_br = n_row_tiles_g_br; - ou_job.n_tiles_per_bc = n_tiles_per_bc; - ou_job.DV = DV; - hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job)); - } - - // Wait for current K DMA and interleave - void * curr_k = dma_queue_pop(dma).dst; - fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start); - - // ---- Phase 2: qk_dot ---- - qk_job.q_tiles = factx.vtcm_q_tiles; - qk_job.k_tiles = factx.vtcm_k_tiles; - qk_job.s_tiles = factx.vtcm_s_tiles; - qk_job.n_row_tiles = n_row_tiles; - qk_job.n_col_tiles = n_col_tiles; - qk_job.n_dot_tiles = DK / 32; - qk_job.n_tiles_per_bc = n_tiles_per_bc; - qk_job.hmx_scales = factx.vtcm_hmx_scales_qk; - hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job)); - - // Wait for current V DMA and interleave + // ---- 1. Pop and run V-prep for current block ---- void * curr_v = dma_queue_pop(dma).dst; fa_phase_v_interleave(&factx, kv_rows, v_src_stride, curr_v, factx.vtcm_v_tiles[buf_idx], n_tiles_per_bc, kv_start); - if (kv_blk > 0) { - hmx_queue_pop(hmx_q); - hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev); - } - - hmx_queue_pop(hmx_q); - - // ---- Phase 3: softmax + build_D ---- + // ---- 2. Pop and run mask-prep for current block ---- __fp16 * current_mask_vtcm = NULL; if (mask) { if (__builtin_expect(factx.mask_broadcast, true)) { @@ -1818,9 +2039,34 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { } } + // ---- 3. Pop and run K-prep for next block & push next QK-dot ---- + if (kv_blk + 1 < factx.n_kv_blocks) { + const uint32_t next_start = (kv_blk + 1) * Bc; + const uint32_t next_rows = hex_smin(Bc, nek1 - next_start); + const size_t next_buf = 1 - buf_idx; + + void * next_k = dma_queue_pop(dma).dst; + fa_phase_k_interleave(&factx, next_rows, k_src_stride, next_k, next_start, next_buf); + + qk_job[next_buf].q_tiles = factx.vtcm_q_tiles; + qk_job[next_buf].k_tiles = factx.vtcm_k_tiles[next_buf]; + qk_job[next_buf].s_tiles = factx.vtcm_s_tiles[next_buf]; + qk_job[next_buf].n_row_tiles = n_row_tiles; + qk_job[next_buf].n_col_tiles = hmx_ceil_div(next_rows, HMX_FP16_TILE_N_COLS); + qk_job[next_buf].n_dot_tiles = DK / 32; + qk_job[next_buf].n_tiles_per_bc = n_tiles_per_bc; + qk_job[next_buf].hmx_scales = factx.vtcm_hmx_scales_qk; + hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[next_buf])); + } + + // ---- 4. Wait for current block's QK-dot to finish ---- + hmx_queue_pop(hmx_q); + + // ---- 5. Phase 2: softmax + build_D ---- fa_softmax_args_t sargs; memset(&sargs, 0, sizeof(sargs)); sargs.factx = &factx; + sargs.buf_idx = buf_idx; sargs.kv_rows = kv_rows; sargs.n_rows_g = n_rows_g; sargs.n_col_tiles = n_col_tiles; @@ -1838,8 +2084,39 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { sargs.mask_vtcm = current_mask_vtcm; sargs.mask_vtcm_row_stride = factx.mask_buf_row_stride; sargs.slopes = factx.vtcm_slopes; + + // Start HMX O update for block kv_blk - 1 (reads P[1 - buf_idx], V[1 - buf_idx]) + if (kv_blk > 0) { + const size_t prev_buf = 1 - buf_idx; + ou_job[prev_buf].o_curr = o_tile_curr; + ou_job[prev_buf].o_prev = o_tile_prev; + ou_job[prev_buf].p_tiles = factx.vtcm_p_tiles[prev_buf]; + ou_job[prev_buf].v_tiles = factx.vtcm_v_tiles[prev_buf]; + ou_job[prev_buf].d_tiles = factx.vtcm_d_tiles; + ou_job[prev_buf].hmx_scales = factx.vtcm_hmx_scales_id; + ou_job[prev_buf].n_row_tiles = n_row_tiles; + ou_job[prev_buf].n_col_tiles = hmx_ceil_div(hex_smin(Bc, nek1 - (kv_blk - 1) * Bc), HMX_FP16_TILE_N_COLS); + ou_job[prev_buf].n_row_tiles_g_br = n_row_tiles_g_br; + ou_job[prev_buf].n_tiles_per_bc = n_tiles_per_bc; + ou_job[prev_buf].DV = DV; + hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[prev_buf])); + } + + // Run Softmax on HVX (blocking call) fa_phase_softmax_and_build_d(&factx, &sargs, n_row_tiles, n_row_tiles_g_br); + // Wait for HMX O update for block kv_blk - 1 to finish + if (kv_blk > 0) { + hmx_queue_pop(hmx_q); + hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev); + } + + // Prefetch block kv_blk + 2 + if (kv_blk + 2 < factx.n_kv_blocks) { + fa_prefetch_block(dma, k, v, mask, kv_blk + 2, Bc, size_k_row_padded, size_k_row, size_v_row_padded, size_v_row, + ik2, ik3, iv2, iv3, q_start, im3, kv_head, G, m_line_bytes, n_rows_q, nek1, buf_idx, &factx); + } + buf_idx = 1 - buf_idx; } @@ -1847,18 +2124,23 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { if (factx.n_kv_blocks > 0) { const uint32_t last_blk = factx.n_kv_blocks - 1; const size_t last_cols = hmx_ceil_div(hex_smin(Bc, nek1 - last_blk * Bc), HMX_FP16_TILE_N_COLS); - ou_job.o_curr = o_tile_curr; - ou_job.o_prev = o_tile_prev; - ou_job.p_tiles = factx.vtcm_p_tiles; - ou_job.v_tiles = factx.vtcm_v_tiles[1 - buf_idx]; - ou_job.d_tiles = factx.vtcm_d_tiles; - ou_job.hmx_scales = factx.vtcm_hmx_scales_id; - ou_job.n_row_tiles = n_row_tiles; - ou_job.n_col_tiles = last_cols; - ou_job.n_row_tiles_g_br = n_row_tiles_g_br; - ou_job.n_tiles_per_bc = n_tiles_per_bc; - ou_job.DV = DV; - hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job)); + ou_job[0].o_curr = o_tile_curr; + ou_job[0].o_prev = o_tile_prev; + ou_job[0].p_tiles = factx.vtcm_p_tiles[1 - buf_idx]; + ou_job[0].v_tiles = factx.vtcm_v_tiles[1 - buf_idx]; + ou_job[0].d_tiles = factx.vtcm_d_tiles; + ou_job[0].hmx_scales = factx.vtcm_hmx_scales_id; + ou_job[0].n_row_tiles = n_row_tiles; + ou_job[0].n_col_tiles = last_cols; + ou_job[0].n_row_tiles_g_br = n_row_tiles_g_br; + ou_job[0].n_tiles_per_bc = n_tiles_per_bc; + ou_job[0].DV = DV; + hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[0])); + + // Overlapped: run HVX build diag inv L while HMX is busy executing the update + htp_trace_event_start(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); + fa_build_d_diag_inv_l(&factx, n_row_tiles, n_row_tiles_g_br); + htp_trace_event_stop(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); hmx_queue_pop(hmx_q); hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev); @@ -1892,12 +2174,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { // Wait for current K DMA and interleave void * curr_k = dma_queue_pop(dma).dst; - fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start); + fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start, 0); { qk_job.q_tiles = factx.vtcm_q_tiles; - qk_job.k_tiles = factx.vtcm_k_tiles; - qk_job.s_tiles = factx.vtcm_s_tiles; + qk_job.k_tiles = factx.vtcm_k_tiles[0]; + qk_job.s_tiles = factx.vtcm_s_tiles[0]; qk_job.n_row_tiles = n_row_tiles; qk_job.n_col_tiles = n_col_tiles; qk_job.n_dot_tiles = (size_t) (DK / 32); @@ -1948,7 +2230,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { { ou_job.o_curr = o_tile_curr; ou_job.o_prev = o_tile_prev; - ou_job.p_tiles = factx.vtcm_p_tiles; + ou_job.p_tiles = factx.vtcm_p_tiles[0]; ou_job.v_tiles = factx.vtcm_v_tiles[0]; ou_job.d_tiles = factx.vtcm_d_tiles; ou_job.hmx_scales = factx.vtcm_hmx_scales_id; @@ -1959,6 +2241,12 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { ou_job.DV = DV; hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job)); + if (kv_blk + 1 == factx.n_kv_blocks) { + // Overlapped: run HVX build diag inv L while HMX is busy executing the update + htp_trace_event_start(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); + fa_build_d_diag_inv_l(&factx, n_row_tiles, n_row_tiles_g_br); + htp_trace_event_stop(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); + } hmx_queue_pop(ctx->hmx_queue); hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev); @@ -1968,15 +2256,63 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { } } + // Enqueue DMAs for the next iteration early so they overlap with O-PROC + 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; + next_q_start = q_start + Br; + if (next_q_start >= neq1) { + next_q_start = 0; + next_ib3 = ib3 + 1; + } + } + bool has_next = (next_ib3 < neq3); + + if (has_next) { + const uint32_t next_n_rows_q = hex_smin(Br, neq1 - next_q_start); + const uint8_t * next_q_ptr = (const uint8_t *) q->data + next_q_start * q->nb[1] + (next_kv_head * factx.G) * q->nb[2] + next_ib3 * q->nb[3]; + const size_t next_q_row_bytes = q_transposed ? next_n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed; + const size_t next_n_rows = q_transposed ? factx.G : next_n_rows_q; + dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, next_q_ptr), next_q_row_bytes, hex_smax(q_src_stride, next_q_row_bytes), next_q_row_bytes, next_n_rows); + + if (factx.n_kv_blocks > 0) { + const uint32_t next_ik2 = next_kv_head; + const uint32_t next_iv2 = next_kv_head; + uint32_t next_ik3 = ik3; + uint32_t next_iv3 = iv3; + if (next_ib3 != ib3) { + next_ik3 = fastdiv(next_ib3, &kparams->broadcast_rk3); + next_iv3 = fastdiv(next_ib3, &kparams->broadcast_rv3); + } + + const uint8_t * next_k_src = (const uint8_t *) k->data + next_ik2 * k->nb[2] + next_ik3 * k->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], next_k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0); + + const uint8_t * next_v_src = (const uint8_t *) v->data + next_iv2 * v->nb[2] + next_iv3 * v->nb[3]; + dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], next_v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0); + + if (factx.pipeline && mask) { + uint32_t next_im3 = im3; + if (next_ib3 != ib3) { + next_im3 = fastmodulo(next_ib3, mask->ne[3], &factx.src3_div3); + } + if (__builtin_expect(factx.mask_broadcast, true)) { + const uint8_t * ms_src = (const uint8_t *) mask->data + next_q_start * mask->nb[1] + next_im3 * mask->nb[3] + 0; + dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows0 * sizeof(__fp16), next_n_rows_q); + } else { + fa_push_mask_dma_gqa(dma, mask, next_q_start, next_im3, 0, next_kv_head, G, m_line_bytes, kv_rows0, next_n_rows_q, &factx); + } + } + } + } + // ---- Final normalization ---- { - htp_trace_event_start(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); - fa_build_d_diag_inv_l(&factx, n_row_tiles, n_row_tiles_g_br); - htp_trace_event_stop(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start); - on_job.o_curr = o_tile_curr; on_job.o_prev = o_tile_prev; - on_job.d_tiles = factx.vtcm_d_tiles; + on_job.d_tiles = factx.vtcm_d_inv_l; on_job.hmx_scales = factx.vtcm_hmx_scales_id; on_job.n_row_tiles = n_row_tiles; on_job.n_row_tiles_g_br = n_row_tiles_g_br; diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h index 16822f22b..efe5ce548 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h @@ -101,14 +101,16 @@ static_assert(sizeof(struct htp_fa_kernel_params) <= 128, "htp_fa_kernel_params struct hmx_fa_vtcm_layout { // Byte offsets from vtcm_base for each region. size_t off_q_tiles; + size_t off_q_dma; size_t off_o_tiles[2]; size_t off_k_fp16[2]; size_t off_v_fp16[2]; - size_t off_k_tiles; - size_t off_v_tiles[2]; // [1] allocated only when pipeline, else 0 - size_t off_s_tiles; - size_t off_p_tiles; + size_t off_k_tiles[2]; + size_t off_v_tiles[2]; + size_t off_s_tiles[2]; + size_t off_p_tiles[2]; size_t off_d_tiles; + size_t off_d_inv_l; size_t off_m_vec; size_t off_l_vec; size_t off_s_rowmax; @@ -140,7 +142,7 @@ struct hmx_fa_vtcm_layout { static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, size_t gqa_factor, size_t DK, size_t DV, - size_t Br, size_t Bc, size_t n_threads, bool pipeline) { + size_t Br, size_t Bc, size_t n_threads, bool pipeline, bool is_q_fp32) { const size_t g_br = hex_align_up(gqa_factor * Br, HMX_FP16_TILE_N_ROWS); const size_t q_tile_size = hex_align_up(g_br * DK * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); const size_t o_tile_size = hex_align_up(g_br * DV * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); @@ -149,6 +151,7 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, const size_t s_tile_size = hex_align_up(g_br * Bc * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); const size_t d_tile_size = hex_align_up(g_br * g_br * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); + const size_t q_dma_size = hex_align_up(g_br * DK * (is_q_fp32 ? sizeof(float) : sizeof(__fp16)), 128); const size_t k_dma_size = hex_align_up(Bc * hex_round_up(DK * sizeof(__fp16), 128), 128); const size_t v_dma_size = hex_align_up(Bc * hex_round_up(DV * sizeof(__fp16), 128), 128); const size_t col_vec_size = hex_align_up(g_br * sizeof(float), 256); @@ -160,27 +163,47 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, size_t off = 0; - // Section 1: HMX Tiled Buffers (FA_HMX_TILE_SIZE = 2KB Aligned) + // Group A (Part 1 - HMX Tiled buffers) VTCM_LAYOUT_ALLOC(off, off_q_tiles, q_tile_size); VTCM_LAYOUT_ALLOC(off, off_o_tiles[0], o_tile_size); VTCM_LAYOUT_ALLOC(off, off_o_tiles[1], o_tile_size); - VTCM_LAYOUT_ALLOC(off, off_k_tiles, k_tile_size); - VTCM_LAYOUT_ALLOC(off, off_v_tiles[0], v_tile_size); - VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_v_tiles[1], v_tile_size, pipeline); - VTCM_LAYOUT_ALLOC(off, off_s_tiles, s_tile_size); - VTCM_LAYOUT_ALLOC(off, off_p_tiles, s_tile_size); VTCM_LAYOUT_ALLOC(off, off_d_tiles, d_tile_size); + VTCM_LAYOUT_ALLOC(off, off_d_inv_l, d_tile_size); - // Section 2: HVX/DMA flat and vector buffers (128B / 256B Aligned) + // Group B & C share start offset (Group B tiles must be 2KB aligned) + size_t off_group_b_c = hex_align_up(off, HTP_FA_HMX_TILE_SIZE); + + // Group B: Compute-only buffers + size_t off_group_b = off_group_b_c; + VTCM_LAYOUT_ALLOC(off_group_b, off_k_tiles[0], k_tile_size); + VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_k_tiles[1], k_tile_size, pipeline); + VTCM_LAYOUT_ALLOC(off_group_b, off_v_tiles[0], v_tile_size); + VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_v_tiles[1], v_tile_size, pipeline); + VTCM_LAYOUT_ALLOC(off_group_b, off_s_tiles[0], s_tile_size); + VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_s_tiles[1], s_tile_size, pipeline); + VTCM_LAYOUT_ALLOC(off_group_b, off_p_tiles[0], s_tile_size); + VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_p_tiles[1], s_tile_size, pipeline); + VTCM_LAYOUT_ALLOC(off_group_b, off_s_rowmax, col_vec_size); + VTCM_LAYOUT_ALLOC(off_group_b, off_p_rowsum, col_vec_size); + VTCM_LAYOUT_ALLOC(off_group_b, off_row_bufs, row_vec_size * 2 * n_threads); + + const size_t group_b_size = off_group_b - off_group_b_c; + + // Group C: Q fetch DMA buffer + size_t off_group_c = off_group_b_c; + VTCM_LAYOUT_ALLOC(off_group_c, off_q_dma, q_dma_size); + + const size_t group_c_size = off_group_c - off_group_b_c; + + off = off_group_b_c + hex_smax(group_b_size, group_c_size); + + // Group A (Part 2 - remaining non-HMX buffers) VTCM_LAYOUT_ALLOC(off, off_k_fp16[0], k_dma_size); VTCM_LAYOUT_ALLOC(off, off_k_fp16[1], k_dma_size); VTCM_LAYOUT_ALLOC(off, off_v_fp16[0], v_dma_size); VTCM_LAYOUT_ALLOC(off, off_v_fp16[1], v_dma_size); VTCM_LAYOUT_ALLOC(off, off_m_vec, col_vec_size); VTCM_LAYOUT_ALLOC(off, off_l_vec, col_vec_size); - VTCM_LAYOUT_ALLOC(off, off_s_rowmax, col_vec_size); - VTCM_LAYOUT_ALLOC(off, off_p_rowsum, col_vec_size); - VTCM_LAYOUT_ALLOC(off, off_row_bufs, row_vec_size * 2 * n_threads); VTCM_LAYOUT_ALLOC(off, off_hmx_scales_id, 256); VTCM_LAYOUT_ALLOC(off, off_hmx_scales_qk, 256); VTCM_LAYOUT_ALLOC(off, off_mask_buf, m_buf_size); @@ -200,9 +223,9 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, } // Exact VTCM usage for a given (gqa_factor, DK, DV, Br, Bc) configuration. -static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, size_t DV, size_t Br, size_t Bc, size_t n_threads, bool pipeline) { +static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, size_t DV, size_t Br, size_t Bc, size_t n_threads, bool pipeline, bool is_q_fp32) { struct hmx_fa_vtcm_layout L; - hmx_fa_vtcm_layout_build(&L, gqa_factor, DK, DV, Br, Bc, n_threads, pipeline); + hmx_fa_vtcm_layout_build(&L, gqa_factor, DK, DV, Br, Bc, n_threads, pipeline, is_q_fp32); return L.total_bytes; } @@ -239,7 +262,8 @@ static inline int hmx_fa_find_chunk_size(size_t * Br_out, size_t qo_len, size_t kv_len, size_t vtcm_budget, - size_t n_threads) { + size_t n_threads, + bool is_q_fp32) { const size_t T = HMX_FP16_TILE_N_ROWS; // 32 const size_t br_unit = hmx_ceil_div(T, gqa_factor); const size_t bc_unit = HMX_FP16_TILE_N_COLS * 2; // 64 @@ -253,8 +277,9 @@ static inline int hmx_fa_find_chunk_size(size_t * Br_out, const size_t Bc_limit = can_pipeline ? hex_align_down(kv_len / FA_MIN_KV_BLOCKS, bc_unit) : (kv_len >= bc_unit ? hex_align_down(kv_len, bc_unit) : bc_unit); // Cost coefficients calibrated from profiling - const size_t c_q_fixed = 1400; // per-Q-block: q_load + epilogue o_update + o_norm + o_store - const size_t c_iter_fixed = 200; // per-KV-iter: HMX queue push/pop + DMA pop + barriers + const size_t c_q_fixed = 800; // per-Q-block: q_load + epilogue o_update + o_norm + o_store + const size_t c_iter_base = 200; // per-KV-iter base (HMX dot/update + DMA) + const size_t c_softmax = 600; // per 64-row vector chunk on HVX size_t best_cost = SIZE_MAX, best_mn = 0; size_t best_Br = 0, best_Bc = 0; @@ -262,13 +287,20 @@ static inline int hmx_fa_find_chunk_size(size_t * Br_out, for (size_t Br = Br_max; Br >= br_unit; Br -= br_unit) { // Try all Bc candidates from Bc_limit down to bc_unit for (size_t Bc = Bc_limit; Bc >= bc_unit; Bc -= bc_unit) { - size_t vtcm_needed = hmx_fa_compute_vtcm_usage(gqa_factor, DK, DV, Br, Bc, n_threads, can_pipeline); + size_t vtcm_needed = hmx_fa_compute_vtcm_usage(gqa_factor, DK, DV, Br, Bc, n_threads, can_pipeline, is_q_fp32); if (vtcm_needed <= vtcm_budget) { // This Bc fits for this Br! - const size_t q_blocks = (qo_len + Br - 1) / Br; - const size_t kv_blocks = (kv_len + Bc - 1) / Bc; - const size_t cost = q_blocks * (c_q_fixed + kv_blocks * c_iter_fixed); - const size_t mn = Br * Bc; + const size_t q_blocks = (qo_len + Br - 1) / Br; + const size_t kv_blocks = (kv_len + Bc - 1) / Bc; + const size_t actual_threads = (kv_blocks >= 3 && n_threads >= 2) ? n_threads : 1; + const size_t n_rows_g = Br * gqa_factor; + const size_t n_row_vec_cnt = (n_rows_g + 63) / 64; + const size_t n_use = n_row_vec_cnt < actual_threads ? n_row_vec_cnt : actual_threads; + const size_t vecs_per_t = n_use > 0 ? (n_row_vec_cnt + n_use - 1) / n_use : 1; + + const size_t c_iter_actual = c_iter_base + c_softmax * vecs_per_t; + const size_t cost = q_blocks * (c_q_fixed + kv_blocks * c_iter_actual); + const size_t mn = Br * Bc; if (cost < best_cost || (cost == best_cost && mn > best_mn)) { best_cost = cost; 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 a4efb1b00..0011abba5 100644 --- a/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +++ b/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h @@ -767,23 +767,25 @@ static void core_mma_chunk_fp16(__fp16 *restrict c, const __fp16 *restrict a, co // output : fp16 -> f32p -static void transfer_output_chunk_fp16_to_fp32( +static void transfer_output_chunk_fp16_to_fp32_col_chunk( float *restrict dst, const float *restrict src2, const __fp16 *restrict vtcm_src, uint32_t start_row, uint32_t n_rows, - uint32_t n_cols, + uint32_t c_len, + uint32_t total_n_cols, uint32_t dst_stride, uint32_t src2_stride, uint32_t dst_cols ) { - assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0); - const size_t tile_row_stride = (n_cols / HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS; + 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 HVX_Vector one = hvx_vec_splat_f16(1.0); - const size_t limit_c = hex_smin(n_cols, dst_cols); + const size_t limit_c = hex_smin(c_len, dst_cols); const size_t limit_c_aligned = (limit_c & ~31); for (size_t r = 0; r < n_rows; r += 2) { @@ -848,6 +850,22 @@ static void transfer_output_chunk_fp16_to_fp32( } } +static inline void transfer_output_chunk_fp16_to_fp32( + float *restrict dst, + const float *restrict src2, + const __fp16 *restrict vtcm_src, + uint32_t start_row, + uint32_t n_rows, + uint32_t n_cols, + uint32_t dst_stride, + uint32_t src2_stride, + uint32_t dst_cols +) { + transfer_output_chunk_fp16_to_fp32_col_chunk( + dst, src2, vtcm_src, start_row, n_rows, n_cols, n_cols, dst_stride, src2_stride, dst_cols + ); +} + typedef struct { const __fp16 *vtcm_src; float *dst; diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h index c4b60eb26..97b8c7f29 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ctx.h +++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h @@ -19,6 +19,8 @@ #endif #define HTP_MAX_MMAPS 16 +#define HTP_MAX_DIRTY_RANGES 16 + // Memory mapping struct htp_mmap { uint64_t size; @@ -95,7 +97,11 @@ struct htp_context { atomic_bool vtcm_needs_release; uint64_t max_vmem; - uint32_t dirty_map[HTP_OP_MAX_TENSORS / 32]; + struct htp_dirty_range { + uint32_t start; + uint32_t end; + uint32_t bi; + } dirty_ranges[HTP_MAX_DIRTY_RANGES]; // Persistent DDR scratchpad for MUL_MAT_ID mappings void * ddr_spad_base; diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index 262d6110f..cad9a4f54 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -123,7 +123,7 @@ enum htp_tensor_flags { // Tensor descriptor struct htp_tensor { uint32_t data; // Buffer offset in the messages, and data pointer on the NPU - uint32_t alias; // Index of the canonical tensor for this memory buffer + uint32_t reserved; // Reserved for alignment padding (must be multiple of 8) uint32_t size; // Data size in bytes uint32_t flags; // Buffer / tensor flags uint32_t type; // Data type @@ -173,6 +173,7 @@ enum htp_trace_event_id { HTP_TRACE_EVT_DMA = 0, HTP_TRACE_EVT_L2FLUSH = 1, HTP_TRACE_EVT_INIT = 2, + HTP_TRACE_EVT_BUFF = 3, HTP_TRACE_EVT_HVX_COMP = 20, HTP_TRACE_EVT_HVX_A_QUANT = 21, @@ -225,7 +226,10 @@ struct htp_opbatch_rsp { uint32_t n_tensors; // Number of tensors uint32_t n_ops; // Number of op profile descriptors uint32_t n_traces[HTP_MAX_NTHREADS + 1]; - uint8_t pad[8]; // align to 8 bytes + uint32_t usecs; // Number of usec + uint32_t pad; // align to 8 bytes + uint64_t cycles_start; // Start cycle counter + uint64_t cycles_stop; // Stop cycle counter // struct htp_prof_desc profs[]; -- dspqueue buf 0 }; diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.c b/ggml/src/ggml-hexagon/htp/htp-tensor.c index 3f028a7e4..39436e26d 100644 --- a/ggml/src/ggml-hexagon/htp/htp-tensor.c +++ b/ggml/src/ggml-hexagon/htp/htp-tensor.c @@ -2,6 +2,7 @@ #include #include +#include #include "hex-common.h" #include "hex-utils.h" @@ -10,84 +11,6 @@ #include "htp-ctx.h" #include "work-queue.h" -struct l2flush_task { - struct htp_thread_trace * trace; - uint32_t start; - uint32_t end; - uint32_t chunk_size; - uint32_t ti; -}; - -static void l2flush_thread_worker(unsigned int n, unsigned int i, void * data) { - struct l2flush_task * task = (struct l2flush_task *) data; - const uint32_t start = task->start; - const uint32_t end = task->end; - const uint32_t ti = task->ti; - const uint32_t chunk_size = task->chunk_size; - - const uint32_t thread_s = start + i * chunk_size; - if (thread_s >= end) { - return; - } - uint32_t thread_e = thread_s + chunk_size; - if (thread_e > end) { - thread_e = end; - } - - struct htp_thread_trace * tr = &task->trace[i]; - htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, ti); - hex_l2flush((void *) (uintptr_t) thread_s, thread_e - thread_s); - htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, ti); -} - -static void flush_all_dcache(struct htp_context * ctx) { - struct htp_thread_trace * tr = &ctx->trace[0]; - htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, 0); - qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); - hex_l2fetch_block(ctx, ctx->footprint); - htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, 0); - bitmap_reset(ctx->dirty_map, HTP_OP_MAX_TENSORS); -} - -static void flush_tensor_range(struct htp_context * ctx, const struct htp_tensor * t) { - struct htp_thread_trace * tr = &ctx->trace[0]; - - if (t->size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) { - struct l2flush_task task; - task.start = hex_align_down((size_t) t->data, HEX_L2_LINE_SIZE); - task.end = hex_align_up((size_t) t->data + t->size, HEX_L2_LINE_SIZE); - task.ti = t->ti; - task.trace = ctx->trace; - - const uint32_t total_size = task.end - task.start; - const uint32_t n_blocks = (total_size + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE; - const uint32_t blocks_per_thread = fastdiv(n_blocks + ctx->n_threads - 1, &ctx->n_threads_div); - task.chunk_size = blocks_per_thread * HEX_L2_BLOCK_SIZE; - - work_queue_run(ctx->work_queue, l2flush_thread_worker, &task, ctx->n_threads); - } else { - htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, t->ti); - hex_l2flush((void *) t->data, t->size); - htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, t->ti); - } - - htp_tensor_make_clean(t, ctx->dirty_map); -} - -void htp_tensor_flush(struct htp_context * ctx, const struct htp_tensor * t) { - if (!bitmap_test(ctx->dirty_map, t->ti)) { - return; - } - - if (t->size > HEX_L2_FLUSH_ALL_THRESHOLD) { - flush_all_dcache(ctx); - return; - } - - flush_tensor_range(ctx, t); -} - -// One dirty tensor's line-aligned range, placed in the flattened global block space. struct l2flush_range { uint32_t start; // line-aligned start address uint32_t end; // line-aligned end address @@ -103,9 +26,18 @@ struct l2flush_multi_task { uint32_t blocks_per_thread; }; +static void flush_all_dcache(struct htp_context * ctx) { + struct htp_thread_trace * tr = &ctx->trace[0]; + htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, 0); + qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); + hex_l2fetch_block(ctx, ctx->footprint); + htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, 0); + memset(ctx->dirty_ranges, 0, sizeof(ctx->dirty_ranges)); +} + static void l2flush_multi_worker(unsigned int n, unsigned int i, void * data) { - (void) n; struct l2flush_multi_task * task = (struct l2flush_multi_task *) data; + (void) n; const uint32_t gb_first = i * task->blocks_per_thread; uint32_t gb_last = gb_first + task->blocks_per_thread; @@ -141,11 +73,177 @@ static void l2flush_multi_worker(unsigned int n, unsigned int i, void * data) { htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, gb_first); } -void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) { - uint64_t total_dirty = 0; +void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) { + const struct htp_tensor * pending[HTP_OP_MAX_OUTPUTS]; + uint32_t n_pending = 0; + for (uint32_t i = 0; i < n; i++) { const struct htp_tensor * t = tensors[i]; - if (t && bitmap_test(ctx->dirty_map, t->ti)) { + if (!t) continue; + + uint32_t t_start = t->data; + uint32_t t_end = t_start + t->size; + + bool merged = false; + for (uint32_t j = 0; j < HTP_MAX_DIRTY_RANGES; j++) { + struct htp_dirty_range * r = &ctx->dirty_ranges[j]; + if (!r->start) continue; + + if (r->start <= t_end && t_start <= r->end) { + uint32_t new_start = (t_start < r->start) ? t_start : r->start; + uint32_t new_end = (t_end > r->end) ? t_end : r->end; + r->start = new_start; + r->end = new_end; + merged = true; + } + } + + if (!merged) { + pending[n_pending++] = t; + } + } + + if (n_pending == 0) { + return; + } + + uint32_t empty_indices[HTP_MAX_DIRTY_RANGES]; + uint32_t active_indices[HTP_MAX_DIRTY_RANGES]; + uint32_t n_active = 0; + uint32_t n_empty = 0; + for (uint32_t j = 0; j < HTP_MAX_DIRTY_RANGES; j++) { + if (ctx->dirty_ranges[j].start) { + active_indices[n_active++] = j; + } else { + empty_indices[n_empty++] = j; + } + } + + if (n_pending <= n_empty) { + for (uint32_t i = 0; i < n_pending; i++) { + uint32_t idx = empty_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + r->start = pending[i]->data; + r->end = pending[i]->data + pending[i]->size; + r->bi = pending[i]->bi; + } + return; + } + + uint32_t n_evict = n_pending - n_empty; + uint32_t total_evict_size = 0; + for (uint32_t i = 0; i < n_evict; i++) { + uint32_t idx = active_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + total_evict_size += r->end - r->start; + } + + if (total_evict_size > HEX_L2_FLUSH_ALL_THRESHOLD) { + flush_all_dcache(ctx); + for (uint32_t i = 0; i < n_pending; i++) { + struct htp_dirty_range * r = &ctx->dirty_ranges[i]; + r->start = pending[i]->data; + r->end = pending[i]->data + pending[i]->size; + r->bi = pending[i]->bi; + } + return; + } + + if (total_evict_size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1 && n_evict <= HTP_OP_MAX_INPUTS) { + struct l2flush_multi_task task; + task.trace = ctx->trace; + task.n_ranges = n_evict; + + uint32_t block_acc = 0; + for (uint32_t i = 0; i < n_evict; i++) { + uint32_t idx = active_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + + struct l2flush_range * rg = &task.ranges[i]; + rg->start = hex_align_down((size_t) r->start, HEX_L2_LINE_SIZE); + rg->end = hex_align_up((size_t) r->end, HEX_L2_LINE_SIZE); + rg->block_first = block_acc; + rg->n_blocks = (rg->end - rg->start + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE; + block_acc += rg->n_blocks; + } + + task.total_blocks = block_acc; + task.blocks_per_thread = fastdiv(block_acc + ctx->n_threads - 1, &ctx->n_threads_div); + + work_queue_run(ctx->work_queue, l2flush_multi_worker, &task, ctx->n_threads); + } else { + struct htp_thread_trace * tr = &ctx->trace[0]; + htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, 0); + for (uint32_t i = 0; i < n_evict; i++) { + uint32_t idx = active_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + uint32_t size = r->end - r->start; + hex_l2flush((void *) (uintptr_t) r->start, size); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, 0); + } + + for (uint32_t i = 0; i < n_evict; i++) { + uint32_t idx = active_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + r->start = pending[i]->data; + r->end = pending[i]->data + pending[i]->size; + r->bi = pending[i]->bi; + } + + for (uint32_t i = 0; i < n_empty; i++) { + uint32_t idx = empty_indices[i]; + struct htp_dirty_range * r = &ctx->dirty_ranges[idx]; + r->start = pending[n_evict + i]->data; + r->end = pending[n_evict + i]->data + pending[n_evict + i]->size; + r->bi = pending[n_evict + i]->bi; + } +} + +static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor * t) { + uint32_t t_start = t->data; + uint32_t t_end = t_start + t->size; + + for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) { + struct htp_dirty_range * r = &ctx->dirty_ranges[i]; + if (!r->start) continue; + + if (r->start < t_end && t_start < r->end) { + if (t_start <= r->start && r->end <= t_end) { + r->start = 0; + } else if (t_start <= r->start) { + r->start = t_end; + } else if (r->end <= t_end) { + r->end = t_start; + } + } + } +} + +static inline bool is_tensor_dirty(struct htp_context * ctx, const struct htp_tensor * t) { + uint32_t t_start = t->data; + uint32_t t_end = t_start + t->size; + + for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) { + struct htp_dirty_range * r = &ctx->dirty_ranges[i]; + if (!r->start) continue; + + if (r->start < t_end && t_start < r->end) { + return true; + } + } + return false; +} + +void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) { + const struct htp_tensor * dirty_tensors[HTP_OP_MAX_INPUTS]; + uint32_t n_dirty = 0; + uint64_t total_dirty = 0; + + for (uint32_t i = 0; i < n; i++) { + const struct htp_tensor * t = tensors[i]; + if (t && (t->flags & HTP_TENSOR_COMPUTE) && is_tensor_dirty(ctx, t)) { + dirty_tensors[n_dirty++] = t; total_dirty += t->size; } } @@ -159,21 +257,15 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co return; } - // Aggregate is small enough to walk. Thread it across all dirty ranges at once - // when it is worth the dispatch, otherwise flush sequentially. - if (total_dirty > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) { + if (total_dirty >= HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) { struct l2flush_multi_task task; task.trace = ctx->trace; task.n_ranges = 0; uint32_t block_acc = 0; - for (uint32_t i = 0; i < n; i++) { - const struct htp_tensor * t = tensors[i]; - if (!t || !bitmap_test(ctx->dirty_map, t->ti)) { - continue; - } - // Clear as we go: dedups a tensor passed as multiple srcs (e.g. mul(x,x)). - htp_tensor_make_clean(t, ctx->dirty_map); + for (uint32_t i = 0; i < n_dirty; i++) { + const struct htp_tensor * t = dirty_tensors[i]; + make_tensor_clean(ctx, t); struct l2flush_range * rg = &task.ranges[task.n_ranges++]; rg->start = hex_align_down((size_t) t->data, HEX_L2_LINE_SIZE); @@ -191,14 +283,11 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co } struct htp_thread_trace * tr = &ctx->trace[0]; - for (uint32_t i = 0; i < n; i++) { - const struct htp_tensor * t = tensors[i]; - if (!t || !bitmap_test(ctx->dirty_map, t->ti)) { - continue; - } + for (uint32_t i = 0; i < n_dirty; i++) { + const struct htp_tensor * t = dirty_tensors[i]; htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, t->ti); - hex_l2flush((void *) t->data, t->size); + hex_l2flush((void *) (uintptr_t) t->data, t->size); htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, t->ti); - htp_tensor_make_clean(t, ctx->dirty_map); + make_tensor_clean(ctx, t); } } diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.h b/ggml/src/ggml-hexagon/htp/htp-tensor.h index e503bb90f..2c3fc54c7 100644 --- a/ggml/src/ggml-hexagon/htp/htp-tensor.h +++ b/ggml/src/ggml-hexagon/htp/htp-tensor.h @@ -5,10 +5,6 @@ #include "htp-ops.h" #include "hex-bitmap.h" -static inline struct htp_tensor * htp_tensor_alias(const struct htp_tensor * t) { - return (struct htp_tensor *) (uintptr_t) t->alias; -} - static inline void * htp_tensor_data(const struct htp_tensor * t) { return (void *) (uintptr_t) t->data; } @@ -17,20 +13,8 @@ static inline uint32_t * htp_tensor_flags(const struct htp_tensor * t) { return (uint32_t *) &t->flags; } -static inline void htp_tensor_make_dirty(const struct htp_tensor * t, uint32_t * dirty_map) { - struct htp_tensor * curr = (struct htp_tensor *) t; - do { - bitmap_set(dirty_map, curr->ti); - curr = htp_tensor_alias(curr); - } while (curr != t); -} - -static inline void htp_tensor_make_clean(const struct htp_tensor * t, uint32_t * dirty_map) { - bitmap_clear(dirty_map, t->ti); -} - struct htp_context; -void htp_tensor_flush(struct htp_context * ctx, const struct htp_tensor * t); void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n); +void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n); #endif // HTP_TENSOR_H diff --git a/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h b/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h index c05bd0b85..5b18f163c 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +++ b/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h @@ -208,6 +208,77 @@ static inline void hvx_mad_f32_f16_aa_rx2(float * restrict y, const void * restr } } } +static inline void hvx_mad_f32_f16_aa_vec(float * restrict y, const void * restrict x, HVX_Vector S0, uint32_t n) { + const HVX_Vector * restrict vx0 = (const HVX_Vector *) x; + + HVX_VectorPair * restrict vy_p = (HVX_VectorPair *) y; + HVX_Vector * restrict vy = (HVX_Vector *) y; + + uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors + uint32_t nloe = n % VLEN_FP16; // leftover elements + + uint32_t i = 0; + + #pragma unroll(2) + for (i = 0; i < nvec; ++i) { + vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx0[i]), S0); + } + + if (nloe) { + HVX_VectorPair xy_p = vy_p[i]; + xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx0[i]), S0); + + HVX_Vector xy = Q6_V_lo_W(xy_p); + i = 2 * i; // index for vy + + if (nloe >= VLEN_FP32) { + vy[i] = xy; + nloe -= VLEN_FP32; ++i; xy = Q6_V_hi_W(xy_p); + } + + if (nloe) { + hvx_vec_store_a(&vy[i], nloe * 4, xy); + } + } +} + +static inline void hvx_mad_f32_f16_aa_rx2_vec(float * restrict y, const void * restrict x0, const void * restrict x1, + HVX_Vector S0, HVX_Vector S1, uint32_t n) { + const HVX_Vector * restrict vx0 = (const HVX_Vector *) x0; + const HVX_Vector * restrict vx1 = (const HVX_Vector *) x1; + + HVX_VectorPair * restrict vy_p = (HVX_VectorPair *) y; + HVX_Vector * restrict vy = (HVX_Vector *) y; + + uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors + uint32_t nloe = n % VLEN_FP16; // leftover elements + + uint32_t i = 0; + + #pragma unroll(2) + for (i = 0; i < nvec; ++i) { + vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx0[i]), S0); + vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx1[i]), S1); + } + + if (nloe) { + HVX_VectorPair xy_p = vy_p[i]; + xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx0[i]), S0); + xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx1[i]), S1); + + HVX_Vector xy = Q6_V_lo_W(xy_p); + i = 2 * i; // index for vy + + if (nloe >= VLEN_FP32) { + vy[i] = xy; + nloe -= VLEN_FP32; ++i; xy = Q6_V_hi_W(xy_p); + } + + if (nloe) { + hvx_vec_store_a(&vy[i], nloe * 4, xy); + } + } +} static inline void hvx_scale_vec_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const uint32_t n, HVX_Vector vs) { assert((size_t) dst % 128 == 0); diff --git a/ggml/src/ggml-hexagon/htp/hvx-reduce.h b/ggml/src/ggml-hexagon/htp/hvx-reduce.h index 3c0073ef6..76d712dc8 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-reduce.h +++ b/ggml/src/ggml-hexagon/htp/hvx-reduce.h @@ -286,6 +286,46 @@ static inline float hvx_sum_of_squares_f32(const uint8_t * restrict src, const i } } +// Signed 32-bit Integer Max variants + +static inline HVX_Vector hvx_vec_reduce_max_n_i32(HVX_Vector in, unsigned int n) { + unsigned int total = n * 4; // total vec nbytes + unsigned int width = 4; // int32 nbytes + + HVX_Vector max_val = in, max_t; + while (width < total) { + max_t = Q6_V_vror_VR(max_val, width); // rotate right + max_val = Q6_Vw_vmax_VwVw(max_t, max_val); // elementwise signed max + width = width << 1; + } + return max_val; +} + +static inline HVX_Vector hvx_vec_reduce_max_i32(HVX_Vector in) { + return hvx_vec_reduce_max_n_i32(in, 32); +} + +static inline int32_t hvx_reduce_max_i32_a(const uint8_t * restrict src, const int num_elems) { + HVX_Vector init_vec = Q6_V_vsplat_R(((const int32_t *) src)[0]); + HVX_Vector pad_vec = Q6_V_vsplat_R(0x80000000); + assert((uintptr_t) src % 128 == 0); + hvx_reduce_loop_body(HVX_Vector, init_vec, pad_vec, Q6_Vw_vmax_VwVw, hvx_vec_reduce_max_i32, hvx_vec_get_i32); +} + +static inline int32_t hvx_reduce_max_i32_u(const uint8_t * restrict src, const int num_elems) { + HVX_Vector init_vec = Q6_V_vsplat_R(((const int32_t *) src)[0]); + HVX_Vector pad_vec = Q6_V_vsplat_R(0x80000000); + hvx_reduce_loop_body(HVX_UVector, init_vec, pad_vec, Q6_Vw_vmax_VwVw, hvx_vec_reduce_max_i32, hvx_vec_get_i32); +} + +static inline int32_t hvx_reduce_max_i32(const uint8_t * restrict src, const int num_elems) { + if (hex_is_aligned((void *) src, 128)) { + return hvx_reduce_max_i32_a(src, num_elems); + } else { + return hvx_reduce_max_i32_u(src, num_elems); + } +} + #undef hvx_reduce_loop_body #undef HVX_REDUCE_MAX_OP #undef HVX_REDUCE_SUM_OP diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 321b832e1..7d65e4643 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -901,10 +901,8 @@ static void prep_tensor(struct htp_context *ctx, struct htp_buf_desc *bufs, stru uint32_t offset = t->data; uint32_t size = t->size; uint32_t bi = t->bi; - uint32_t alias = t->alias; t->data = (uint32_t) (bufs[bi].base + offset); // update data to the actual pointer - t->alias = (uint32_t) (tens + alias); // update alias to the actual pointer FARF(HIGH, "prep-tensor #%u: bi %u offset %u size %u data %p : %u:%u:%u:%u", idx, t->bi, offset, t->size, (void*) t->data, t->ne[0], t->ne[1], t->ne[3], t->ne[3]); @@ -955,14 +953,14 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u octx->dsts[i] = dst; octx->dst_dma[i] = octx->ctx->dma; // FIXME: ? octx->ctx->dma_cached : octx->ctx->dma; - htp_tensor_make_dirty(dst, octx->ctx->dirty_map); - FARF(HIGH, "prep-dst[%u] #%u: data %p size %u : %u:%u:%u:%u", i, dst_idx, (void*) dst->data, dst->size, dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]); } int status = execute_op(octx); + htp_tensor_dirty_all(octx->ctx, octx->dsts, HTP_OP_MAX_OUTPUTS); + octx->src0_spad.src = NULL; octx->src1_spad.src = NULL; octx->src2_spad.src = NULL; @@ -994,12 +992,6 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r FARF(HIGH, "processing opbatch #%u: n-bufs %u n-tensors %u n-ops %u n-traces %u : m-size %u b-size %u t-size %u o-size %u", req->id, n_bufs, n_tens, n_ops, req->n_traces, dbuf->size, b_size, t_size, o_size); - // Clean cache at the start of the batch - // We cant trace this part because the trace buffer is setup later - qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); - hex_l2fetch_block(ctx, ctx->footprint); - bitmap_reset(ctx->dirty_map, HTP_OP_MAX_TENSORS); - // Setup descriptor pointers uint8_t * m_ptr = dbuf->ptr; struct htp_buf_desc* bufs = (struct htp_buf_desc*) m_ptr; m_ptr += b_size; @@ -1007,13 +999,8 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r struct htp_op_desc* ops = (struct htp_op_desc*) m_ptr; m_ptr += o_size; struct htp_prof_desc* pds = (struct htp_prof_desc*) m_ptr; - prep_op_bufs(ctx, bufs, n_bufs); - prep_tensors(ctx, bufs, tens, n_tens); - - struct htp_ops_context *octx = &ctx->octx; - memset(octx, 0, sizeof(*octx)); - octx->n_threads = ctx->n_threads; - octx->ctx = ctx; + struct profile_data batch_prof; + profile_start(HTP_PROF_BASIC, &batch_prof); memset(ctx->trace, 0, sizeof(ctx->trace)); if (ctx->profiler == HTP_PROF_TRACE) { @@ -1024,6 +1011,24 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r } } + // Clean cache at the start of the batch + htp_trace_event_start(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); + qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); + hex_l2fetch_block(ctx, ctx->footprint); + memset(ctx->dirty_ranges, 0, sizeof(ctx->dirty_ranges)); + htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); + + htp_trace_event_start(&ctx->trace[0], HTP_TRACE_EVT_BUFF, 0); + prep_op_bufs(ctx, bufs, n_bufs); + htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_BUFF, 0); + + prep_tensors(ctx, bufs, tens, n_tens); + + struct htp_ops_context *octx = &ctx->octx; + memset(octx, 0, sizeof(*octx)); + octx->n_threads = ctx->n_threads; + octx->ctx = ctx; + work_queue_wakeup(ctx->work_queue); if (ctx->hmx_queue) { hmx_queue_wakeup(ctx->hmx_queue); @@ -1056,13 +1061,23 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r } work_queue_suspend(ctx->work_queue); + // Flush remaining dirty tensors at the end of the batch + htp_trace_event_start(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); + qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); + htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); + + profile_stop(HTP_PROF_BASIC, &batch_prof); + struct htp_opbatch_rsp rsp; memset(&rsp, 0, sizeof(rsp)); - rsp.id = req->id; - rsp.status = op_status; - rsp.n_bufs = n_bufs; - rsp.n_tensors = n_tens; - rsp.n_ops = n_ops; + rsp.id = req->id; + rsp.status = op_status; + rsp.n_bufs = n_bufs; + rsp.n_tensors = n_tens; + rsp.n_ops = n_ops; + rsp.usecs = batch_prof.usecs; + rsp.cycles_start = batch_prof.cycles_start; + rsp.cycles_stop = batch_prof.cycles_stop; if (ctx->profiler == HTP_PROF_TRACE) { for (int t = 0; t <= HTP_MAX_NTHREADS; t++) { @@ -1073,11 +1088,6 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r struct dspqueue_buffer write_dbuf = *dbuf; write_dbuf.flags = DSPQUEUE_BUFFER_FLAG_FLUSH_SENDER | DSPQUEUE_BUFFER_FLAG_INVALIDATE_RECIPIENT; - // Flush remaining dirty tensors at the end of the batch - htp_trace_event_start(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); - qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); - htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0); - err = dspqueue_write(queue, 0, 1, &write_dbuf, sizeof(rsp), (const uint8_t *) &rsp, DSPQUEUE_TIMEOUT_NONE); if (err != 0) { FARF(ERROR, "dspqueue_write failed: 0x%08x", (unsigned) err); diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c index 0611aa982..9d385469a 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.c +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c @@ -14,6 +14,8 @@ #include "hex-dma.h" #include "hvx-utils.h" #include "hvx-dump.h" +#include "hvx-arith.h" +#include "hvx-reduce.h" #define GGML_COMMON_DECL_C #include "ggml-common.h" @@ -82,6 +84,8 @@ struct htp_mm_context { // Precomputed values uint32_t src0_nrows_per_thread; + uint32_t src0_row_size_padded; + uint32_t src1_nrows; struct fastdiv_values mm_div_ne12_ne1; struct fastdiv_values mm_div_ne1; @@ -103,6 +107,7 @@ struct htp_mm_context { // Fields for scattered mapping & HMX support in MUL_MAT_ID const uint32_t * matrix_row_counts; const struct mmid_row_mapping * matrix_rows; + uint32_t mapping_stride; // Dynamic VTCM pointers allocated sequentially uint8_t * vtcm_src0; @@ -154,8 +159,6 @@ static const uint8_t __attribute__((aligned(VLEN))) kvalues_mxfp4_lut[] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, }; - - #define htp_matmul_tensors_preamble \ const struct htp_tensor * restrict src0 = octx->src[0]; \ const struct htp_tensor * restrict src1 = octx->src[1]; \ @@ -444,6 +447,16 @@ 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; \ + const float * src2_ptr = (const float *) src2->data + src0_start_row; \ + int slice_size = (int)MIN(src0_end_row, ne0) - (int)src0_start_row; \ + if (slice_size > 0) { \ + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr, src2_ptr), \ + slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1); \ + dma_queue_pop_nowait(dma_queue); \ + } \ + } \ for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) { \ dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, \ src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ @@ -465,7 +478,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void \ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \ - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ \ if (push_ct < ct_end) { \ dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride), \ @@ -476,24 +489,16 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void \ int copy_cnt = (int)MIN(src0_end_row, ne0) - (int)src0_start_row; \ if (copy_cnt > 0) { \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct_end); \ if (src2) { \ - float * dst_ptr = &dst_col[src0_start_row]; \ - const float * src2_ptr = (const float *) src2->data + src0_start_row; \ - float * tmp_ptr = tmp; \ - int remaining = copy_cnt; \ - while (remaining > 0) { \ - int n = MIN(remaining, 32); \ - HVX_Vector v_out = hvx_vmemu(tmp_ptr); \ - HVX_Vector v_z = hvx_vmemu(src2_ptr); \ - hvx_vec_store_u(dst_ptr, n * sizeof(float), hvx_vec_add_f32_f32(v_out, v_z)); \ - dst_ptr += n; \ - src2_ptr += n; \ - tmp_ptr += n; \ - remaining -= n; \ - } \ + 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), \ + copy_cnt); \ } else { \ hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt); \ } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct_end); \ } \ } @@ -1069,6 +1074,16 @@ 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; + const float * src2_ptr = (const float *) src2->data + src0_start_row; + int slice_size = (int)src0_end_row - (int)src0_start_row; + if (slice_size > 0) { + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr, src2_ptr), + slice_size * sizeof(float), slice_size * sizeof(float), slice_size * sizeof(float), 1); + dma_queue_pop_nowait(dma_queue); + } + } for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { const uint32_t is0 = (ir0 - src0_start_row); if (is0 >= n_prefetch) { @@ -1114,27 +1129,21 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) { } int copy_cnt = src0_end_row - src0_start_row; - if (src2) { - float * dst_ptr = &dst_col[src0_start_row]; - const float * src2_ptr = (const float *) src2->data + src0_start_row; - float * tmp_ptr = tmp; - int remaining = copy_cnt; - while (remaining > 0) { - int n = MIN(remaining, 32); - HVX_Vector v_out = hvx_vmemu(tmp_ptr); - HVX_Vector v_z = hvx_vmemu(src2_ptr); - hvx_vec_store_u(dst_ptr, n * sizeof(float), hvx_vec_add_f32_f32(v_out, v_z)); - dst_ptr += n; - src2_ptr += n; - tmp_ptr += n; - remaining -= n; + if (copy_cnt > 0) { + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_end_row); + 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), + copy_cnt); + } else { + hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt); } - } else { - hvx_copy_f32_ua((uint8_t *) &dst_col[src0_start_row], (uint8_t *) tmp, copy_cnt); + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, src0_end_row); } } -#define MMID_MATRIX_ROW(row_id, i1) matrix_rows[(row_id) * ids->ne[0] * ids->ne[1] + (i1)] +#define MMID_MATRIX_ROW(row_id, i1) matrix_rows[(row_id) * mmctx->mapping_stride + (i1)] static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) { htp_matmul_preamble; @@ -1519,7 +1528,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads, - dst_row_size, src0_row_size, src1_row_size, kparams->n_prefetch, false, false, false); + dst_row_size, src0_row_size, src1_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false, false); if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM || kparams->kernel_type == HTP_MM_KERNEL_HVX_F32_F32_VTCM || @@ -1551,6 +1560,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base; mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1); mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0); + mmctx->vtcm_src2 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2); mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst); octx->src1_spad.src = NULL; @@ -2346,12 +2356,77 @@ static void dequantize_tiled_weight_chunk_to_fp16_tiles( } } +typedef struct { + float *dst; + const float *src2; + const __fp16 *vtcm_src; + uint32_t n_rows; + uint32_t n_cols; + uint32_t dst_stride; + uint32_t src2_stride; + uint32_t dst_cols; + struct fastdiv_values n_threads_div; + struct htp_thread_trace *traces; + struct htp_context *ctx; +} output_transfer_col_chunk_state_t; + +static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned int i, void *data) { + (void) n; + 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; + 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; + uint32_t c_last = b_last * 32; + uint32_t c_len = c_last - c_first; + + if (c_len == 0) return; + + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, c_first); + + float *dst = st->dst + c_first; + const float *src2 = st->src2 ? (st->src2 + c_first) : NULL; + const __fp16 *vtcm_src = st->vtcm_src + b_first * HTP_MM_HMX_TILE_N_ELMS; + + int chunk_dst_cols = (int)st->dst_cols - (int)c_first; + if (chunk_dst_cols > 0) { + transfer_output_chunk_fp16_to_fp32_col_chunk( + dst, src2, vtcm_src, 0, st->n_rows, c_len, st->n_cols, + st->dst_stride, st->src2_stride, (uint32_t)chunk_dst_cols + ); + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, c_first); +} + 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; + if (n_threads > 1 && n_blocks >= (uint32_t)n_threads) { + struct fastdiv_values n_threads_div = init_fastdiv_values(n_threads); + output_transfer_col_chunk_state_t col_state; + col_state.dst = dst; + col_state.src2 = src2; + col_state.vtcm_src = vtcm_src; + col_state.n_rows = (uint32_t)n_rows; + col_state.n_cols = (uint32_t)n_cols; + col_state.dst_stride = (uint32_t)dst_stride; + col_state.src2_stride = src2_stride; + col_state.dst_cols = (uint32_t)dst_cols; + col_state.n_threads_div = n_threads_div; + col_state.traces = ctx->trace; + col_state.ctx = ctx; + + worker_pool_run_func(ctx->worker_pool, transfer_output_chunk_col_chunk_worker_fn, &col_state, n_threads); + return; + } + size_t n_tot_chunks = n_rows; size_t n_chunks_per_task = (n_threads == 1) ? n_tot_chunks : hmx_ceil_div(n_rows, n_threads); n_chunks_per_task = hex_align_up(n_chunks_per_task, 2); @@ -3338,12 +3413,10 @@ int op_matmul(struct htp_ops_context * octx) { static int hmx_mm_op_matmul_id( struct htp_ops_context * octx, - struct htp_mm_context * mmctx, - const uint32_t * matrix_row_counts, - const struct mmid_row_mapping * matrix_rows, - void * mapping_buf, - bool must_free_mapping + struct htp_mm_context * mmctx ) { + const uint32_t * matrix_row_counts = mmctx->matrix_row_counts; + const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows; htp_matmul_tensors_preamble; const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; const int n_ids = octx->src[2]->ne[0]; @@ -3361,28 +3434,24 @@ static int hmx_mm_op_matmul_id( nb11, nb12, nb1, nb2, (int) src0->nb[1], (int) src0->type, - matrix_rows, cur_a, n_ids * octx->src[2]->ne[1]); + matrix_rows, cur_a, mmctx->mapping_stride); if (ret != 0) { FARF(ERROR, "HMX matmul failed for expert %u, error %d\n", cur_a, ret); - if (must_free_mapping) free(mapping_buf); return HTP_STATUS_NO_SUPPORT; } } - if (must_free_mapping) free(mapping_buf); return HTP_STATUS_OK; } static int hvx_mm_matmul_id( struct htp_ops_context * octx, struct htp_mm_context * mmctx, - size_t src0_row_size_padded, - uint32_t src1_nrows, - worker_callback_t matmul_id_job_func, - void * mapping_buf, - bool must_free_mapping + work_queue_func_t hvx_mmid_task_func ) { htp_matmul_tensors_preamble; + const uint32_t src0_row_size_padded = mmctx->src0_row_size_padded; + const uint32_t src1_nrows = mmctx->src1_nrows; struct htp_thread_trace * tr = &octx->ctx->trace[0]; htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); @@ -3395,7 +3464,7 @@ static int hvx_mm_matmul_id( const uint32_t nb = (ne10 + qk - 1) / qk; const uint32_t total_nb = src1_nrows * nb; - worker_callback_t quant_task_func; + work_queue_func_t quant_task_func; uint32_t n_quant_tasks = 1; if (src1_nrows < octx->n_threads) { n_quant_tasks = MIN(total_nb, octx->n_threads); @@ -3416,7 +3485,7 @@ static int hvx_mm_matmul_id( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, kparams->n_prefetch, true, false, false); + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false, false); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; @@ -3431,7 +3500,6 @@ static int hvx_mm_matmul_id( // Make sure the reserved vtcm size is sufficient if (octx->ctx->vtcm_size < vtcm_size) { FARF(ERROR, "matmul-id-%s : current VTCM reservation %zu is too small, needed %zu\n", mmctx->type, octx->ctx->vtcm_size, vtcm_size); - if (must_free_mapping) free(mapping_buf); return HTP_STATUS_VTCM_TOO_SMALL; } @@ -3461,12 +3529,78 @@ static int hvx_mm_matmul_id( htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); - worker_pool_run_func(octx->ctx->worker_pool, matmul_id_job_func, mmctx, octx->n_threads); + worker_pool_run_func(octx->ctx->worker_pool, hvx_mmid_task_func, mmctx, octx->n_threads); - if (must_free_mapping) free(mapping_buf); return HTP_STATUS_OK; } +static inline void scan_expert_ids_n( + const struct htp_tensor * ids, + const uint32_t n_ids, + uint32_t n_as, + uint32_t * counts, + struct mmid_row_mapping * matrix_rows, + uint32_t mapping_stride +) { + const size_t ids_nb1 = ids->nb[1]; + const uint8_t * ids_data = (const uint8_t *) ids->data; + + for (uint32_t iid1 = 0; iid1 < ids->ne[1]; ++iid1) { + const int32_t * row_ptr = (const int32_t *) (ids_data + iid1 * ids_nb1); + for (uint32_t id = 0; id < n_ids; ++id) { + const int32_t i02 = row_ptr[id]; + if (i02 < 0) { + continue; + } + assert(i02 < n_as); + + if (matrix_rows) { + matrix_rows[i02 * mapping_stride + counts[i02]] = (struct mmid_row_mapping) { id, iid1 }; + } + counts[i02] += 1; + } + } +} + +static inline void scan_expert_ids( + const struct htp_tensor * ids, + uint32_t n_ids, + uint32_t n_as, + uint32_t * counts, + struct mmid_row_mapping * matrix_rows, + uint32_t mapping_stride +) { + const size_t ids_nb0 = ids->nb[0]; + + if (ids_nb0 == 4) { + switch (n_ids) { + case 8: scan_expert_ids_n(ids, 8, n_as, counts, matrix_rows, mapping_stride); break; + case 4: scan_expert_ids_n(ids, 4, n_as, counts, matrix_rows, mapping_stride); break; + case 2: scan_expert_ids_n(ids, 2, n_as, counts, matrix_rows, mapping_stride); break; + default: scan_expert_ids_n(ids, n_ids, n_as, counts, matrix_rows, mapping_stride); break; + } + } else { + // Strided fallback + const size_t ids_nb1 = ids->nb[1]; + const uint8_t * ids_data = (const uint8_t *) ids->data; + for (uint32_t iid1 = 0; iid1 < ids->ne[1]; ++iid1) { + const int32_t * row_ptr = (const int32_t *) (ids_data + iid1 * ids_nb1); + for (uint32_t id = 0; id < n_ids; ++id) { + const int32_t i02 = *(const int32_t *) ((const uint8_t *) row_ptr + id * ids_nb0); + if (i02 < 0) { + continue; + } + assert(i02 < n_as); + + if (matrix_rows) { + matrix_rows[i02 * mapping_stride + counts[i02]] = (struct mmid_row_mapping) { id, iid1 }; + } + counts[i02] += 1; + } + } + } +} + int op_matmul_id(struct htp_ops_context * octx) { htp_matmul_tensors_preamble; @@ -3489,74 +3623,72 @@ int op_matmul_id(struct htp_ops_context * octx) { const uint32_t src0_nrows = ne01; // per expert const uint32_t src1_nrows = ne11 * ne12 * ne13; - worker_callback_t quant_task_func; - worker_callback_t matmul_id_job_func = src1_nrows > 1 ? hvx_mm_id : hvx_mv_id; - - // Compute src0_nrows_per_thread - mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; - mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); + mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; + mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); // row groups const int n_ids = ids->ne[0]; // n_expert_used const int n_as = ne02; // n_expert - size_t matrix_row_counts_size = n_as * sizeof(uint32_t); - size_t matrix_row_map_size = n_as * ids->ne[0] * ids->ne[1] * sizeof(struct mmid_row_mapping); - const size_t total_map_size = matrix_row_counts_size + matrix_row_map_size; - - void * mapping_buf = NULL; - bool must_free_mapping = false; - - if (octx->ctx->ddr_spad_base && total_map_size <= octx->ctx->ddr_spad_size) { - mapping_buf = octx->ctx->ddr_spad_base; - } else { - mapping_buf = memalign(128, total_map_size); - if (mapping_buf) { - must_free_mapping = true; - } else { - return HTP_STATUS_INTERNAL_ERR; - } - } - - uint32_t * matrix_row_counts = (uint32_t *) mapping_buf; - struct mmid_row_mapping * matrix_rows = (struct mmid_row_mapping *) ((uint8_t *) mapping_buf + matrix_row_counts_size); - - mmctx->matrix_row_counts = matrix_row_counts; - mmctx->matrix_rows = matrix_rows; - mmctx->mm_div_ne11 = kparams->div_ne11; - - if (hvx_mm_init_vec_dot(mmctx, src0->type) != 0) { - if (must_free_mapping) free(mapping_buf); - return HTP_STATUS_NO_SUPPORT; - } + uint8_t * mapping_buf = octx->ctx->ddr_spad_base; + uint32_t mapping_stride = 1; + uint32_t * matrix_row_counts = (uint32_t *) mapping_buf; + struct mmid_row_mapping * matrix_rows = NULL; if (src1_nrows > 1) { - // initialize matrix_row_counts and map - memset(matrix_row_counts, 0, n_as * sizeof(uint32_t)); + const size_t matrix_row_counts_size = n_as * sizeof(uint32_t); + assert(octx->ctx->ddr_spad_size >= matrix_row_counts_size); - // group rows by src0 matrix - for (uint32_t iid1 = 0; iid1 < ids->ne[1]; ++iid1) { // token idx - for (uint32_t id = 0; id < n_ids; ++id) { // expert idx - const int32_t i02 = *(const int32_t *) ((const uint8_t *) ids->data + iid1 * ids->nb[1] + id * ids->nb[0]); + hex_l2fetch_block((const void *) ids->data, ids->ne[1] * ids->nb[1]); - if (i02 < 0) { - continue; - } - assert(i02 < n_as); + memset(matrix_row_counts, 0, matrix_row_counts_size); + scan_expert_ids(ids, n_ids, n_as, matrix_row_counts, NULL, 0); - matrix_rows[i02 * n_ids * ids->ne[1] + matrix_row_counts[i02]] = (struct mmid_row_mapping) { id, iid1 }; - matrix_row_counts[i02] += 1; + uint32_t max_count = hvx_reduce_max_i32((const uint8_t *) matrix_row_counts, n_as); + mapping_stride = max_count > 0 ? max_count : 1; + + size_t matrix_row_map_size = n_as * mapping_stride * sizeof(struct mmid_row_mapping); + const size_t total_map_size = matrix_row_counts_size + matrix_row_map_size; + + if (total_map_size > octx->ctx->ddr_spad_size) { + mapping_buf = memalign(128, total_map_size); + if (!mapping_buf) { + return HTP_STATUS_INTERNAL_ERR; } } + + matrix_row_counts = (uint32_t *) mapping_buf; + matrix_rows = (struct mmid_row_mapping *) (mapping_buf + matrix_row_counts_size); + + memset(matrix_row_counts, 0, n_as * sizeof(uint32_t)); + scan_expert_ids(ids, n_ids, n_as, matrix_row_counts, matrix_rows, mapping_stride); } + mmctx->matrix_row_counts = matrix_row_counts; + mmctx->matrix_rows = matrix_rows; + mmctx->mapping_stride = mapping_stride; + mmctx->mm_div_ne11 = kparams->div_ne11; + mmctx->src0_row_size_padded = src0_row_size_padded; + mmctx->src1_nrows = src1_nrows; + htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); + int s; if (kparams->n_hmx) { - return hmx_mm_op_matmul_id(octx, mmctx, matrix_row_counts, matrix_rows, mapping_buf, must_free_mapping); + s = hmx_mm_op_matmul_id(octx, mmctx); + } else { + 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); + } else { + s = HTP_STATUS_NO_SUPPORT; + } } - return hvx_mm_matmul_id(octx, mmctx, src0_row_size_padded, src1_nrows, matmul_id_job_func, mapping_buf, must_free_mapping); + if (mapping_buf != octx->ctx->ddr_spad_base) { + free(mapping_buf); + } + + return s; } int op_matmul_qkv(struct htp_ops_context * octx) { @@ -3633,7 +3765,7 @@ int op_matmul_qkv(struct htp_ops_context * octx) { struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, src1->ne[0], src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, kparams->n_prefetch, false, true, false); + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true, false); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; @@ -3778,7 +3910,7 @@ int op_matmul_ffn(struct htp_ops_context * octx) { struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, src1->ne[0], src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, kparams->n_prefetch, false, false, true); + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, false, true); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h index 445044660..6c393664c 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.h +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h @@ -460,6 +460,7 @@ static inline void htp_mm_hvx_vtcm_layout_build( size_t dst_row_size, size_t src0_row_size, size_t src1_row_size, + size_t src2_row_size, uint32_t n_prefetch, bool is_matmul_id, bool is_fused_qkv, @@ -467,7 +468,7 @@ static inline void htp_mm_hvx_vtcm_layout_build( ) { size_t src0_sz = 0; size_t src1_sz = 0; - size_t src2_sz = 0; + size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0; size_t src3_sz = 0; size_t dst_sz = 0; diff --git a/scripts/snapdragon/ggml-hexagon-profile.py b/scripts/snapdragon/ggml-hexagon-profile.py index 2e096aaa9..97a3acd26 100755 --- a/scripts/snapdragon/ggml-hexagon-profile.py +++ b/scripts/snapdragon/ggml-hexagon-profile.py @@ -6,6 +6,7 @@ import re import argparse import statistics import logging +import bisect from typing import Any, Dict, List, Optional from collections import defaultdict @@ -30,7 +31,7 @@ op_pattern = re.compile( ) trace_pattern = re.compile( - r"trace-op\s+(?P[A-Z_0-9+]+):\s+thread\s+(?P\d+)\s+event\s+(?P[A-Z_0-9\-]+)\s+info\s+(?P\d+)\s+(?Pstart|stop)\s+(?P\d+)" + r"trace-evt\s+(?P[A-Z_0-9\-]+):\s+thread\s+(?P\d+)\s+info\s+(?P\d+)\s+(?Pstart|stop)\s+(?P\d+)" ) logger = logging.getLogger("ggml-hexagon-profile") @@ -50,9 +51,13 @@ def normalize_event_name(evt_type): class CycleUnwrapper: - def __init__(self): - self.last_raw = None - self.high_part = 0 + def __init__(self, initial_val=None): + if initial_val is not None: + self.last_raw = initial_val & 0xFFFFFFFF + self.high_part = initial_val & 0xFFFFFFFF00000000 + else: + self.last_raw = None + self.high_part = 0 def unwrap(self, raw): if self.last_raw is None: @@ -78,10 +83,12 @@ def parse_log(file_path, pmu_index=None): sys.exit(1) all_ops: List[Dict[str, Any]] = [] + all_traces: List[Dict[str, Any]] = [] current_op: Optional[Dict[str, Any]] = None timestamp_pattern = re.compile(r"^(?P\d+)\.(?P\d+)\.(?P\d+)\.(?P\d+)\s+[A-Z]\s+") - unwrapper = CycleUnwrapper() + unwrapper = None + trace_unwrapper = None for line in f: ts_match = timestamp_pattern.match(line) @@ -100,6 +107,7 @@ def parse_log(file_path, pmu_index=None): if not prefix_match: continue + names = parts[1] if len(parts) == 7: dims, types, timings = parts[2], parts[3], parts[6] elif len(parts) == 6: @@ -120,6 +128,7 @@ def parse_log(file_path, pmu_index=None): op_match = op_pattern.search(line) if op_match: op_name = op_match.group('op_name') + names = "" dims = op_match.group('dims').strip() types = op_match.group('types').strip() else: @@ -136,24 +145,31 @@ def parse_log(file_path, pmu_index=None): except (ValueError, IndexError): pmu_val = None - evt_raw = op_match.group('evt') if 'evt' in op_match.groupdict() else None evt_val = None - if evt_raw: + evt_val = None + if types.startswith("evt-cnt "): try: - evt_val = [int(x.strip()) for x in evt_raw.split(',')] + evt_val = [int(x.strip()) for x in types[8:].split(',')] except ValueError: evt_val = None cycles_start_raw = op_match.group('start') unwrapped_cycles_start = None - if cycles_start_raw: - unwrapped_cycles_start = unwrapper.unwrap(int(cycles_start_raw)) + if op_name == "OPBATCH": + if cycles_start_raw: + unwrapped_cycles_start = int(cycles_start_raw) + unwrapper = CycleUnwrapper(unwrapped_cycles_start) + trace_unwrapper = CycleUnwrapper(unwrapped_cycles_start) + else: + if cycles_start_raw and unwrapper is not None: + unwrapped_cycles_start = unwrapper.unwrap(int(cycles_start_raw)) idx = line.find("profile-op ") op_text = line[idx + 11:].strip() if idx != -1 else line.strip() current_op = { 'name': op_name, + 'names': names, 'dims': dims, 'types': types, 'op_text': op_text, @@ -170,110 +186,239 @@ def parse_log(file_path, pmu_index=None): continue trace_match = trace_pattern.search(line) - if trace_match and current_op: - if trace_match.group('op_name') == current_op['name']: - raw_cyc = int(trace_match.group('cycles')) - current_op['trace_events'].append({ - 'thread': int(trace_match.group('thread')), - 'event': trace_match.group('event'), - 'info': int(trace_match.group('info')), - 'cycles': raw_cyc, - 'unwrapped_cycles': unwrapper.unwrap(raw_cyc), - 'state': trace_match.group('state') - }) + if trace_match: + raw_cyc = int(trace_match.group('cycles')) + unwrapped_cyc = None + if trace_unwrapper is not None: + unwrapped_cyc = trace_unwrapper.unwrap(raw_cyc) + all_traces.append({ + 'thread': int(trace_match.group('thread')), + 'event': trace_match.group('event'), + 'info': int(trace_match.group('info')), + 'cycles': raw_cyc, + 'unwrapped_cycles': unwrapped_cyc, + 'state': trace_match.group('state') + }) f.close() + + # Assign start/end cycles to all ops + for op in all_ops: + op['start_cycles'] = op['unwrapped_cycles_start'] + op['end_cycles'] = op['start_cycles'] + op['cycles'] if op['start_cycles'] is not None else None + + # Filter ops with valid start_cycles + valid_ops = [op for op in all_ops if op['start_cycles'] is not None and op['end_cycles'] is not None] + + # Separate OPBATCH ops from other ops + opbatch_ops = [op for op in valid_ops if op['name'] == "OPBATCH"] + other_ops = [op for op in valid_ops if op['name'] != "OPBATCH"] + + # Sort them by start_cycles to enable binary search + opbatch_ops.sort(key=lambda op: op['start_cycles']) + other_ops.sort(key=lambda op: op['start_cycles']) + + opbatch_starts = [op['start_cycles'] for op in opbatch_ops] + other_starts = [op['start_cycles'] for op in other_ops] + + # Map trace events to any operator whose cycles contain them + for e in all_traces: + cyc = e['unwrapped_cycles'] + if cyc is None: + continue + + # Map to OPBATCH + idx = bisect.bisect_right(opbatch_starts, cyc) - 1 + if idx >= 0: + op = opbatch_ops[idx] + if op['start_cycles'] <= cyc <= op['end_cycles']: + op['trace_events'].append(e) + + # Map to other ops + idx = bisect.bisect_right(other_starts, cyc) - 1 + if idx >= 0: + op = other_ops[idx] + if op['start_cycles'] <= cyc <= op['end_cycles']: + op['trace_events'].append(e) + return all_ops -def print_ascii_timeline(op_name, dims, types, usec, cycles, events, evt_val=None): - evt_str = "" - if evt_val: - evt_str = " - evt [" + ",".join(str(x) for x in evt_val) + "]" +def print_bubbles_timeline(op): + op_name = op['name'] + dims = op['dims'] + types = op['types'] + usec = op['usec'] + cycles = op['cycles'] + events = op['trace_events'] logger.info("=" * 100) - logger.info(f"{op_name} ({dims} : {types}) - {usec} usec {cycles} cycles{evt_str}") + logger.info(f"{op_name} ({dims} : {types}) - {usec} usec {cycles} cycles") logger.info("=" * 100) - events = sorted(events, key=lambda e: e['cycles']) if not events: logger.info(" No trace events recorded.") return - min_cycles = events[0]['cycles'] + # Identify start and end cycles for this operator + op_start = op['start_cycles'] + op_end = op['end_cycles'] + if op_start is None or op_end is None: + logger.info(" Cannot analyze bubbles: missing start/end cycle counts.") + return - logger.info("Cycles %-30s" % "EventDetails" + " ".join(f"T{i:<2}" for i in range(10)) + " HMX") - logger.info("-" * 100) - - thread_stacks = [[] for _ in range(11)] + batch_duration = op_end - op_start + if batch_duration <= 0: + logger.info(" Cannot analyze bubbles: batch duration is 0.") + return + # Group events by (thread, track_type) + tracks = defaultdict(list) for e in events: t = e['thread'] - if t < 0 or t > 10: - continue + is_dma = (normalize_event_name(e['event']) == 'DMA') + track_type = 'dma' if is_dma else 'compute' + tracks[(t, track_type)].append(e) - if e['cycles'] >= min_cycles: - rel_cycles = e['cycles'] - min_cycles - else: - rel_cycles = (e['cycles'] + 0x100000000) - min_cycles + active_threads = sorted(list(set(t for (t, track_type) in tracks.keys()))) + if not active_threads: + logger.info(" No active threads in trace.") + return - state = e['state'] - evt_type = e['event'] + bubble_threshold = 10000 # 10k cycles - # Determine char representing the event - norm_evt = normalize_event_name(evt_type) - char = '?' - if norm_evt == 'V-COMP': - char = 'V' - elif norm_evt == 'M-COMP': - char = 'H' - elif norm_evt == 'A-QUANT': - char = 'Q' - elif norm_evt == 'A-PREP': - char = 'A' - elif norm_evt == 'Q-PREP': - char = 'q' - elif norm_evt == 'K-PREP': - char = 'k' - elif norm_evt == 'V-PREP': - char = 'v' - elif norm_evt == 'W-DEQUANT': - char = 'D' - elif norm_evt == 'O-PROC': - char = 'O' - elif norm_evt == 'W-PREP': - char = 'P' - elif norm_evt == 'DMA': - char = 'M' + thread_stats = {} + for t in active_threads: + thread_stats[t] = { + 'compute_idle_cycles': batch_duration, + 'compute_idle_pct': 100.0, + 'compute_bubbles': [], - if state == 'start': - thread_stacks[t].append(char) - elif state == 'stop': - if thread_stacks[t]: - if thread_stacks[t][-1] == char: - thread_stacks[t].pop() - elif char in thread_stacks[t]: - thread_stacks[t].remove(char) - else: - thread_stacks[t].pop() + 'dma_idle_cycles': batch_duration, + 'dma_idle_pct': 100.0, + 'dma_bubbles': [] + } - cols = [] - for i in range(11): - if thread_stacks[i]: - cols.append(f"[{thread_stacks[i][-1]}]") + total_compute_idle_pct = 0.0 + total_dma_idle_pct = 0.0 + + for t in active_threads: + for track_type in ['compute', 'dma']: + key = (t, track_type) + track_events = tracks.get(key, []) + + if not track_events: + gaps = [(op_start, op_end)] + idle_cycles = batch_duration else: - cols.append(" | ") + track_events = sorted(track_events, key=lambda e: e.get('unwrapped_cycles') or e['cycles']) - evt_desc = f"T{t}: {evt_type} {state} ({e['info']})" - logger.info(f"{rel_cycles:10d} %-30s" % evt_desc + " ".join(cols[:10]) + " " + cols[10]) + active_intervals = [] + active_count = 0 + curr_start = None + + for e in track_events: + cyc = e.get('unwrapped_cycles') or e['cycles'] + cyc = max(op_start, min(op_end, cyc)) + state = e['state'] + + if state == 'start': + if active_count == 0: + curr_start = cyc + active_count += 1 + elif state == 'stop': + if active_count > 0: + active_count -= 1 + if active_count == 0: + active_intervals.append((curr_start, cyc)) + else: + active_intervals.append((op_start, cyc)) + + if active_count > 0 and curr_start is not None: + active_intervals.append((curr_start, op_end)) + + # Merge intervals + active_intervals.sort(key=lambda x: x[0]) + merged_intervals = [] + for start, end in active_intervals: + if not merged_intervals: + merged_intervals.append([start, end]) + else: + last_start, last_end = merged_intervals[-1] + if start <= last_end: + merged_intervals[-1][1] = max(last_end, end) + else: + merged_intervals.append([start, end]) + + # Calculate gaps + gaps = [] + curr_time = op_start + for start, end in merged_intervals: + if start > curr_time: + gaps.append((curr_time, start)) + curr_time = max(curr_time, end) + if curr_time < op_end: + gaps.append((curr_time, op_end)) + + idle_cycles = sum(end - start for start, end in gaps) + + idle_pct = (idle_cycles / batch_duration) * 100.0 + + bubbles = [] + for start, end in gaps: + dur = end - start + if dur >= bubble_threshold: + bubbles.append((start, end, dur)) + + if track_type == 'compute': + thread_stats[t]['compute_idle_cycles'] = idle_cycles + thread_stats[t]['compute_idle_pct'] = idle_pct + thread_stats[t]['compute_bubbles'] = bubbles + total_compute_idle_pct += idle_pct + else: + thread_stats[t]['dma_idle_cycles'] = idle_cycles + thread_stats[t]['dma_idle_pct'] = idle_pct + thread_stats[t]['dma_bubbles'] = bubbles + total_dma_idle_pct += idle_pct + + avg_compute_idle = total_compute_idle_pct / len(active_threads) + avg_dma_idle = total_dma_idle_pct / len(active_threads) + + logger.info(" Combined Idle Statistics:") + logger.info(f" Active Threads : {', '.join(str(t) for t in active_threads)}") + logger.info(f" Avg Thread Compute IDLE : {avg_compute_idle:.1f}%") + logger.info(f" Avg Thread DMA IDLE : {avg_dma_idle:.1f}%") logger.info("-" * 100) + logger.info(" Per-Thread Idle Analysis:") + for t in active_threads: + stats = thread_stats[t] + thread_name = f"Thread {t:<2} (HVX)" if t != 10 else "Thread 10 (HMX)" + logger.info(f" {thread_name} -> Compute Idle: {stats['compute_idle_pct']:.1f}% | DMA Idle: {stats['dma_idle_pct']:.1f}%") -def print_ascii_summary(op_name, dims, types, usec, cycles, events, evt_val=None): - evt_str = "" - if evt_val: - evt_str = " - evt [" + ",".join(str(x) for x in evt_val) + "]" + all_bubbles = [] + for t in active_threads: + stats = thread_stats[t] + for start, end, dur in stats['compute_bubbles']: + pct = (dur / batch_duration) * 100.0 + all_bubbles.append((dur, f"Thread {t} Compute: bubble of {dur} cycles ({pct:.1f}%) at {start - op_start} to {end - op_start}")) + for start, end, dur in stats['dma_bubbles']: + pct = (dur / batch_duration) * 100.0 + all_bubbles.append((dur, f"Thread {t} DMA : bubble of {dur} cycles ({pct:.1f}%) at {start - op_start} to {end - op_start}")) + + if all_bubbles: + logger.info("-" * 100) + logger.info(f" Significant Bubbles (>= {bubble_threshold} cycles):") + all_bubbles.sort(key=lambda x: x[0], reverse=True) + for dur, desc in all_bubbles[:15]: + logger.info(f" {desc}") + else: + logger.info("-" * 100) + logger.info(f" No significant bubbles detected (all idle gaps < {bubble_threshold} cycles).") + + +def print_ascii_summary(op_name, dims, types, usec, cycles, events): logger.info("=" * 100) - logger.info(f"{op_name} ({dims} : {types}) - {usec} usec {cycles} cycles{evt_str}") + logger.info(f"{op_name} ({dims} : {types}) - {usec} usec {cycles} cycles") logger.info("=" * 100) events = sorted(events, key=lambda e: e['cycles']) @@ -415,8 +560,8 @@ def main(): parser.add_argument("--pmu-index", type=int) parser.add_argument("--pmu-name", type=str) parser.add_argument("--width", action='append', default=['dims:40'], help="Override column width, e.g. --width dims:50") - parser.add_argument("--timeline", type=str, nargs='?', const='summary', choices=["summary", "diagram"], - help="Output ASCII art event summary or timing diagram (default: summary)") + parser.add_argument("--timeline", type=str, nargs='?', const='summary', choices=["summary", "bubbles"], + help="Output ASCII art event summary or thread idle bubble analysis (default: summary)") parser.add_argument("--filter", type=str, help="Regex filter matching against the original profile-op line") group = parser.add_mutually_exclusive_group() @@ -457,12 +602,11 @@ def main(): ops = ops[-args.tail:] if args.timeline: - logger.info(f"\n# ASCII Timing {args.timeline.capitalize()}\n") for op in ops: if args.timeline == "summary": - print_ascii_summary(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events'], op.get('evt_val')) - elif args.timeline == "diagram": - print_ascii_timeline(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events'], op.get('evt_val')) + print_ascii_summary(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events']) + elif args.timeline == "bubbles": + print_bubbles_timeline(op) else: generate_report(ops, args.top, overrides, args.sort, pmu_name=final_pmu_name) diff --git a/scripts/snapdragon/ggml-hexagon-trace.py b/scripts/snapdragon/ggml-hexagon-trace.py index 37f137a9e..4755adfa1 100755 --- a/scripts/snapdragon/ggml-hexagon-trace.py +++ b/scripts/snapdragon/ggml-hexagon-trace.py @@ -6,6 +6,7 @@ import re import argparse import statistics import logging +import bisect from typing import Any, Dict, List, Optional from collections import defaultdict @@ -16,11 +17,11 @@ op_pattern = re.compile( ) trace_pattern = re.compile( - r"trace-op\s+(?P[A-Z_0-9+]+):\s+thread\s+(?P\d+)\s+event\s+(?P[A-Z_0-9\-]+)\s+info\s+(?P\d+)\s+(?Pstart|stop)\s+(?P\d+)" + r"trace-evt\s+(?P[A-Z_0-9\-]+):\s+thread\s+(?P\d+)\s+info\s+(?P\d+)\s+(?Pstart|stop)\s+(?P\d+)" ) -def normalize_event_name(evt_type): +def normalize_event_name(evt_type, info=0): if evt_type == "HVX_COMP": return "V-COMP" if evt_type == "HMX_COMP": @@ -32,9 +33,13 @@ def normalize_event_name(evt_type): class CycleUnwrapper: - def __init__(self): - self.last_raw = None - self.high_part = 0 + def __init__(self, initial_val=None): + if initial_val is not None: + self.last_raw = initial_val & 0xFFFFFFFF + self.high_part = initial_val & 0xFFFFFFFF00000000 + else: + self.last_raw = None + self.high_part = 0 def unwrap(self, raw): if self.last_raw is None: @@ -60,8 +65,10 @@ def parse_log(file_path): sys.exit(1) all_ops: List[Dict[str, Any]] = [] + all_traces: List[Dict[str, Any]] = [] current_op: Optional[Dict[str, Any]] = None - unwrapper = CycleUnwrapper() + unwrapper = None + trace_unwrapper = None line_idx = 0 for line in f: @@ -73,6 +80,7 @@ def parse_log(file_path): if not prefix_match: continue + names = parts[1] if len(parts) == 7: dims, types, strides, params, timings = parts[2], parts[3], parts[4], parts[5], parts[6] elif len(parts) == 6: @@ -93,6 +101,7 @@ def parse_log(file_path): op_match = op_pattern.search(line) if op_match: op_name = op_match.group('op_name') + names = "" dims = op_match.group('dims').strip() if op_match.group('dims') else '' types = op_match.group('types').strip() if op_match.group('types') else '' strides = op_match.group('strides').strip() if op_match.group('strides') else '' @@ -103,18 +112,30 @@ def parse_log(file_path): if op_match: cycles_start_raw = op_match.group('start') unwrapped_cycles_start = None - if cycles_start_raw: - unwrapped_cycles_start = unwrapper.unwrap(int(cycles_start_raw)) + if op_name == "OPBATCH": + if cycles_start_raw: + unwrapped_cycles_start = int(cycles_start_raw) + unwrapper = CycleUnwrapper(unwrapped_cycles_start) + trace_unwrapper = CycleUnwrapper(unwrapped_cycles_start) + else: + if cycles_start_raw and unwrapper is not None: + unwrapped_cycles_start = unwrapper.unwrap(int(cycles_start_raw)) idx = line.find("profile-op ") op_text = line[idx + 11:].strip() if idx != -1 else line.strip() + evt_str = None + if types.startswith("evt-cnt "): + evt_str = types[8:].strip() + current_op = { 'name': op_name, + 'names': names, 'dims': dims, 'types': types, 'strides': strides, 'params': params, + 'evt': evt_str, 'op_text': op_text, 'usec': int(op_match.group('usec')), 'cycles': int(op_match.group('cycles')), @@ -127,20 +148,22 @@ def parse_log(file_path): continue trace_match = trace_pattern.search(line) - if trace_match and current_op: - if trace_match.group('op_name') == current_op['name']: - raw_cyc = int(trace_match.group('cycles')) - current_op['trace_events'].append({ - 'thread': int(trace_match.group('thread')), - 'event': trace_match.group('event'), - 'info': int(trace_match.group('info')), - 'cycles': raw_cyc, - 'unwrapped_cycles': unwrapper.unwrap(raw_cyc), - 'state': trace_match.group('state') - }) + if trace_match: + raw_cyc = int(trace_match.group('cycles')) + unwrapped_cyc = None + if trace_unwrapper is not None: + unwrapped_cyc = trace_unwrapper.unwrap(raw_cyc) + all_traces.append({ + 'thread': int(trace_match.group('thread')), + 'event': trace_match.group('event'), + 'info': int(trace_match.group('info')), + 'cycles': raw_cyc, + 'unwrapped_cycles': unwrapped_cyc, + 'state': trace_match.group('state') + }) f.close() - return all_ops + return all_ops, all_traces # --- Simple protobuf encoder --- @@ -246,7 +269,7 @@ def write_trace_packet_to_file(f, packet_bytes): # --- End Protobuf Encoder --- -def generate_perfetto_trace(filtered_ops, output_path): +def generate_perfetto_trace(filtered_ops, trace_events, output_path): if not filtered_ops: logger.warning("No operators found after filtering.") return @@ -269,14 +292,12 @@ def generate_perfetto_trace(filtered_ops, output_path): # Process events completed_events = [] - for op in filtered_ops: - events = op['trace_events'] - if not events: - continue - events = sorted(events, key=lambda e: e['unwrapped_cycles']) + if trace_events: + trace_events = sorted(trace_events, key=lambda e: e['unwrapped_cycles']) + one_usec_cycles = max(avg_freq_mhz, 1.0) active_starts = {} - for e in events: + for e in trace_events: t = e['thread'] evt = e['event'] info = e['info'] @@ -285,6 +306,17 @@ def generate_perfetto_trace(filtered_ops, output_path): key = (t, evt, info) if state == 'start': + # Handle missing stop (start followed by another start) + if key in active_starts: + prev_start = active_starts[key] + completed_events.append({ + 'thread': t, + 'event': evt, + 'info': info, + 'start_cyc': prev_start, + 'end_cyc': prev_start + one_usec_cycles, + 'missing_stop': True, + }) active_starts[key] = cyc elif state == 'stop': if key in active_starts: @@ -296,8 +328,29 @@ def generate_perfetto_trace(filtered_ops, output_path): 'info': info, 'start_cyc': start_cyc, 'end_cyc': cyc, - 'op_name': op['name'] }) + else: + # Handle missing start (stop without start) + completed_events.append({ + 'thread': t, + 'event': evt, + 'info': info, + 'start_cyc': cyc - one_usec_cycles, + 'end_cyc': cyc, + 'missing_start': True, + }) + + # Clear remaining unmatched starts + for key, start_cyc in active_starts.items(): + t, evt, info = key + completed_events.append({ + 'thread': t, + 'event': evt, + 'info': info, + 'start_cyc': start_cyc, + 'end_cyc': start_cyc + one_usec_cycles, + 'missing_stop': True, + }) completed_events.sort(key=lambda e: e['start_cyc']) @@ -316,7 +369,7 @@ def generate_perfetto_trace(filtered_ops, output_path): ts = e['ts_ns'] dur = e['dur_ns'] - norm_evt = normalize_event_name(evt) + norm_evt = normalize_event_name(evt, e['info']) if norm_evt == "DMA": track_key = (t, "DMA") elif t == 10: @@ -343,7 +396,7 @@ def generate_perfetto_trace(filtered_ops, output_path): evt = e['event'] slot = e['slot'] - norm_evt = normalize_event_name(evt) + norm_evt = normalize_event_name(evt, e['info']) if norm_evt == "DMA": track_evt = "DMA" evt_id = 1 @@ -421,18 +474,26 @@ def generate_perfetto_trace(filtered_ops, output_path): for op in filtered_ops: op_start_ns = int(round(((op['start_cycles'] - global_min_cyc) / avg_freq_mhz) * 1000)) op_dur_ns = int(round((op['cycles'] / avg_freq_mhz) * 1000)) - if op_start_ns < last_op_end_ns: - op_start_ns = last_op_end_ns - clamped_dur = max(op_dur_ns, 100) # Clamp to 100ns (0.1us) + if op['name'] != "OPBATCH": + if op_start_ns < last_op_end_ns: + op_start_ns = last_op_end_ns + clamped_dur = max(op_dur_ns, 100) # Clamp to 100ns (0.1us) + last_op_end_ns = op_start_ns + clamped_dur + else: + clamped_dur = max(op_dur_ns, 100) # Debug annotations for Ops debug_annots = [] if 'line_num' in op: debug_annots.append(make_debug_annotation("line", int_val=op['line_num'])) - if 'strides' in op and op['strides']: + if 'names' in op and op['names'] and op['names'] != '----': + debug_annots.append(make_debug_annotation("names", string_val=op['names'])) + if 'strides' in op and op['strides'] and op['strides'] != '----': debug_annots.append(make_debug_annotation("strides", string_val=op['strides'])) if 'params' in op and op['params'] and op['params'] != '----': debug_annots.append(make_debug_annotation("params", string_val=op['params'])) + if 'evt' in op and op['evt']: + debug_annots.append(make_debug_annotation("evt", string_val=op['evt'])) # Slice Begin evt_begin = make_track_event(1, 2, name=f"{op['name']} ({op['dims']})", category="operator", debug_annotations=debug_annots) @@ -444,15 +505,21 @@ def generate_perfetto_trace(filtered_ops, output_path): packet_end = make_trace_packet(op_start_ns + clamped_dur, track_event=evt_end) write_trace_packet_to_file(f, packet_end) - last_op_end_ns = op_start_ns + clamped_dur - # Emit Thread Trace Events for e in completed_events: - norm_name = normalize_event_name(e['event']) + norm_name = normalize_event_name(e['event'], e['info']) name = f"DMA {e['info']}" if norm_name == "DMA" else norm_name + if e.get('missing_start') or e.get('missing_stop'): + name += "!" + + debug_annots = [] + if e.get('missing_start'): + debug_annots.append(make_debug_annotation("missing_start", string_val="true")) + if e.get('missing_stop'): + debug_annots.append(make_debug_annotation("missing_stop", string_val="true")) # Slice Begin - evt_begin = make_track_event(1, e['uuid'], name=name, category="trace") + evt_begin = make_track_event(1, e['uuid'], name=name, category="trace", debug_annotations=debug_annots if debug_annots else None) packet_begin = make_trace_packet(e['ts_ns'], track_event=evt_begin) write_trace_packet_to_file(f, packet_begin) @@ -477,7 +544,7 @@ def main(): args = parser.parse_args() logging.basicConfig(level=logging.INFO, format='%(message)s') - ops = parse_log(args.logfile) + ops, traces = parse_log(args.logfile) if args.filter: try: @@ -492,7 +559,30 @@ def main(): elif args.tail is not None: ops = ops[-args.tail:] - generate_perfetto_trace(ops, args.output) + if args.filter or args.head is not None or args.tail is not None: + valid_ranges = [] + for op in ops: + start_cyc = op['unwrapped_cycles_start'] + end_cyc = start_cyc + op['cycles'] if start_cyc is not None else None + if start_cyc is not None and end_cyc is not None: + valid_ranges.append((start_cyc, end_cyc)) + + valid_ranges.sort(key=lambda r: r[0]) + range_starts = [r[0] for r in valid_ranges] + + filtered_traces = [] + for e in traces: + cyc = e['unwrapped_cycles'] + if cyc is None: + continue + idx = bisect.bisect_right(range_starts, cyc) - 1 + if idx >= 0: + start, end = valid_ranges[idx] + if start <= cyc <= end: + filtered_traces.append(e) + traces = filtered_traces + + generate_perfetto_trace(ops, traces, args.output) if __name__ == "__main__":