From 9f4b18ae6e45602ad13417b627bb3d5bf9f330eb Mon Sep 17 00:00:00 2001 From: Yiwei Shao <44545837+njsyw1997@users.noreply.github.com> Date: Wed, 19 Aug 2026 14:42:57 -0700 Subject: [PATCH] hexagon: fix FA HMX queue ordering and pack the rescale D matrices (llama/27042) * hexagon: fix FA HMX queue ordering in the pipelined path * hexagon: double buffer D matrix, store diagonal tile only * format code * align the indentation --- ggml/src/ggml-hexagon/htp/flash-attn-ops.c | 76 ++++++++++++---------- ggml/src/ggml-hexagon/htp/flash-attn-ops.h | 19 ++++-- 2 files changed, 57 insertions(+), 38 deletions(-) diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c index fe78718c6..817656290 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c @@ -132,8 +132,8 @@ struct hmx_fa_context { __fp16 * vtcm_v_tiles[2]; // V tiles (column-major, double-buffered) __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] + __fp16 * vtcm_d_tiles[2]; // Diagonal rescale, g_br/32 packed diagonal tiles (double-buffered) + __fp16 * vtcm_d_inv_l; // Diagonal rescale (1/l), same packed layout 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] @@ -782,13 +782,14 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) { } } - // Initialize vtcm_d_tiles and vtcm_d_inv_l to 0 + // Zero the whole rescale region: vtcm_d_tiles[0], the optional vtcm_d_tiles[1] + // and vtcm_d_inv_l are equal-sized and allocated back to back, so one run covers + // them all. The scatter only ever writes the diagonal, ignore the rest. 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); + hvx_splat_u8_a((char *) factx->vtcm_d_tiles[0] + d_start, 0, d_end - d_start); } } @@ -1432,17 +1433,19 @@ static inline void fa_softmax_impl( const HVX_VectorPred q_32_mask = Q6_Q_vsetq_R(32 * sizeof(__fp16)); HVX_Vector v_exp_m_diff = exp_m_diff_f16; + __fp16 * const d_tiles_out = factx->vtcm_d_tiles[args->buf_idx]; + size_t t0 = r_vec_idx * 2; if (t0 < args->n_row_tiles) { const HVX_Vector v_content = v_exp_m_diff; - __fp16 * out_base = factx->vtcm_d_tiles + t0 * (args->n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + __fp16 * out_base = d_tiles_out + t0 * HMX_FP16_TILE_N_ELMS; Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content); } size_t t1 = r_vec_idx * 2 + 1; if (t1 < args->n_row_tiles) { const HVX_Vector v_content = Q6_V_vror_VR(v_exp_m_diff, 64); - __fp16 * out_base = factx->vtcm_d_tiles + t1 * (args->n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + __fp16 * out_base = d_tiles_out + t1 * HMX_FP16_TILE_N_ELMS; Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content); } } @@ -1506,7 +1509,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_inv_l + i * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + __fp16 * out_base = factx->vtcm_d_inv_l + i * HMX_FP16_TILE_N_ELMS; Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content); } } @@ -1615,7 +1618,7 @@ static void hmx_fa_o_update_worker(void * data) { const size_t o_stride = n_row_tiles_g_br * HMX_FP16_TILE_N_ELMS; const size_t v_stride = n_tiles_per_bc * HMX_FP16_TILE_N_ELMS; for (size_t r = 0; r < n_row_tiles; ++r) { - const __fp16 * d_diag = d_tiles + r * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + const __fp16 * d_diag = d_tiles + r * HMX_FP16_TILE_N_ELMS; const __fp16 * p_tile_in = p_tiles + (r * n_tiles_per_bc) * HMX_FP16_TILE_N_ELMS; const __fp16 * o_rc = o_prev + r * HMX_FP16_TILE_N_ELMS; const __fp16 * v_tile_in = v_tiles; @@ -1654,7 +1657,7 @@ static void hmx_fa_o_norm_worker(void * data) { asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)job->hmx_scales)); const size_t o_stride = n_row_tiles_g_br * HMX_FP16_TILE_N_ELMS; for (size_t r = 0; r < n_row_tiles; ++r) { - const __fp16 * d_diag = d_tiles + r * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS; + const __fp16 * d_diag = d_tiles + r * HMX_FP16_TILE_N_ELMS; const __fp16 * o_rc = o_prev + r * HMX_FP16_TILE_N_ELMS; __fp16 * o_out = o_curr + r * DV_tiles * HMX_FP16_TILE_N_ELMS; @@ -1882,7 +1885,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { 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_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_d_tiles[0]); + factx.vtcm_d_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_d_tiles[1], pipeline); 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); @@ -2039,7 +2043,30 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { } } - // ---- 3. Pop and run K-prep for next block & push next QK-dot ---- + // ---- 3. Start HMX O update for block kv_blk - 1 (reads P[1 - buf_idx], V[1 - buf_idx], D) ---- + // O update relys on the previous block's P and V tiles. + // O update MUST be pushed before the next block's QK-dot: hmx_queue_pop() retires the + // oldest descriptor, so push order alone decides which pop waits for which job. + // If OU went in after QK(i+1), the pop below would retire QK(i+1) and leave + // OU(i-1) in flight into the next iteration, where V-prep overwrites V[prev_buf]. + 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[prev_buf]; + 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])); + } + + // ---- 4. 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); @@ -2059,10 +2086,10 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { 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 ---- + // ---- 5. Wait for current block's QK-dot to finish ---- hmx_queue_pop(hmx_q); - // ---- 5. Phase 2: softmax + build_D ---- + // ---- 6. Phase 2: softmax + build_D ---- fa_softmax_args_t sargs; memset(&sargs, 0, sizeof(sargs)); sargs.factx = &factx; @@ -2085,23 +2112,6 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { 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); @@ -2128,7 +2138,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { 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].d_tiles = factx.vtcm_d_tiles[1 - buf_idx]; 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; @@ -2232,7 +2242,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { ou_job.o_prev = o_tile_prev; 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.d_tiles = factx.vtcm_d_tiles[0]; ou_job.hmx_scales = factx.vtcm_hmx_scales_id; ou_job.n_row_tiles = n_row_tiles; ou_job.n_col_tiles = n_col_tiles; diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h index efe5ce548..c4d190631 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h @@ -109,7 +109,7 @@ struct hmx_fa_vtcm_layout { 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_tiles[2]; size_t off_d_inv_l; size_t off_m_vec; size_t off_l_vec; @@ -125,7 +125,7 @@ struct hmx_fa_vtcm_layout { size_t q_tile_bytes; size_t o_tile_bytes; size_t s_tile_bytes; // S and P tiles (same size) - size_t d_tile_bytes; + size_t d_tile_bytes; // d_tiles[0..1] + d_inv_l, allocated back to back size_t m_line_bytes; // one mask row size_t m_buf_slot_bytes; // one dma_cache slot = align_up(Br * m_line_bytes, 4096) size_t col_vec_bytes; @@ -149,7 +149,12 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, const size_t k_tile_size = hex_align_up(Bc * DK * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); const size_t v_tile_size = hex_align_up(Bc * DV * sizeof(__fp16), HTP_FA_HMX_TILE_SIZE); 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); + + // The rescale matrices are diagonal: the HMX kernels only ever load the g_br/32 + // tiles that sit on the diagonal, so store just those, packed back to back with + // a stride of one tile. The old [g_br, g_br] square layout allocated g_br/32 + // times more than it used, which is also why a second D buffer was unaffordable. + const size_t d_tile_size = (g_br / HMX_FP16_TILE_N_ROWS) * 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); @@ -167,7 +172,8 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, 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_d_tiles, d_tile_size); + VTCM_LAYOUT_ALLOC(off, off_d_tiles[0], d_tile_size); + VTCM_LAYOUT_ALLOC_OPTIONAL(off, off_d_tiles[1], d_tile_size, pipeline); VTCM_LAYOUT_ALLOC(off, off_d_inv_l, d_tile_size); // Group B & C share start offset (Group B tiles must be 2KB aligned) @@ -213,7 +219,10 @@ static inline void hmx_fa_vtcm_layout_build(struct hmx_fa_vtcm_layout * L, L->o_tile_bytes = o_tile_size; L->col_vec_bytes = col_vec_size; L->s_tile_bytes = s_tile_size; - L->d_tile_bytes = d_tile_size; + // Measured from the actual offsets rather than assumed to be N * d_tile_size, so + // that inserting a region between them (or adding padding to VTCM_LAYOUT_ALLOC) + // cannot silently leave the tail of the run unzeroed. + L->d_tile_bytes = (L->off_d_inv_l + d_tile_size) - L->off_d_tiles[0]; L->m_line_bytes = m_line_size; L->m_buf_slot_bytes = m_buf_slot; L->row_buf_stride = row_vec_size / 128;