From 83d855c5a6d70487121edbf4020b25c96b7a04e7 Mon Sep 17 00:00:00 2001 From: Aparna M P Date: Fri, 28 Aug 2026 03:08:02 +0530 Subject: [PATCH] hex-unary: fix RMS_NORM_MUL weight-offset bugs for grouped/broadcast norms (#27798) --- ggml/src/ggml-hexagon/htp/unary-ops.c | 32 +++++++++++++++++---------- 1 file changed, 20 insertions(+), 12 deletions(-) diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c index b21415a67d..971f3b5377 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.c +++ b/ggml/src/ggml-hexagon/htp/unary-ops.c @@ -478,6 +478,9 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const uint32_t nb11 = src1 ? src1->nb[1] : 0; \ const uint32_t nb12 = src1 ? src1->nb[2] : 0; \ const uint32_t nb13 = src1 ? src1->nb[3] : 0; \ + const uint32_t nb11_bc = (src1 && src1->ne[1] > 1) ? nb11 : 0; \ + const uint32_t nb12_bc = (src1 && src1->ne[2] > 1) ? nb12 : 0; \ + const uint32_t nb13_bc = (src1 && src1->ne[3] > 1) ? nb13 : 0; \ const bool src1_contig = src1 ? ((nb12 == (size_t)ne01 * nb11) && (nb13 == (size_t)ne02 * nb12)) : false; \ \ uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread); \ @@ -497,8 +500,12 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const struct fastdiv_values * div_ne02 = &uctx->kparams->div_ne02; \ const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012; \ \ - const uint32_t src0_max_block = src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01); \ - const uint32_t dst_max_block = dst_contig ? uctx->block : MIN((uint32_t)uctx->block, ne1); \ + const bool src1_needs_row_clip = (IS_RMS_NORM_MUL) && !uctx->broadcast_weight && !src1_contig; \ + const bool block_src0_contig = src0_contig && !src1_needs_row_clip; \ + const bool block_dst_contig = dst_contig && !src1_needs_row_clip; \ + \ + const uint32_t src0_max_block = block_src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01); \ + const uint32_t dst_max_block = block_dst_contig ? uctx->block : MIN((uint32_t)uctx->block, ne1); \ const uint32_t BLOCK = MIN(src0_max_block, dst_max_block); \ if (BLOCK == 0) { \ FARF(ERROR, "unary-f32 : current VTCM reservation %zu is too small, needed at least %zu\n", \ @@ -515,8 +522,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat } \ \ for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) { \ - const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig, ne01, \ - div_ne01); \ + const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig, \ + ne01, div_ne01); \ \ dma_queue_push(dma_queue, \ dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)), \ @@ -530,7 +537,7 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) { \ const size_t src1_off = src1_contig ? (ir * nb11) : \ - unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11, nb12, nb13); \ + unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, nb13_bc); \ dma_queue_push(dma_queue, \ dma_make_ptr(src1_vtcm_data + (vtcm_idx * src1_vtcm_half_size), data_src1 + src1_off), \ uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, block_size); \ @@ -540,8 +547,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat } \ \ for (uint32_t ir = src0_start_row; ir < src0_end_row; ) { \ - const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig, ne01, \ - div_ne01); \ + const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig, \ + ne01, div_ne01); \ \ float * dst_vtcm = (float *) dma_queue_pop(dma_queue).src; \ float * src0_vtcm = (float *) dma_queue_pop(dma_queue).dst; \ @@ -562,12 +569,12 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ const uint32_t next_ir = ir + block_size; \ if (next_ir < src0_end_row) { \ - const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, src0_contig, dst_contig,\ - ne01, div_ne01); \ + const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, block_src0_contig, \ + block_dst_contig, ne01, div_ne01); \ const uint32_t pref_ir = next_ir + next_block_size; \ if (pref_ir < src0_end_row) { \ - const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, src0_contig, \ - dst_contig, ne01, div_ne01); \ + const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, block_src0_contig, \ + block_dst_contig, ne01, div_ne01); \ const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) : \ unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03); \ dma_queue_push(dma_queue, \ @@ -576,7 +583,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) { \ const size_t src1_pref_off = src1_contig ? (pref_ir * nb11) : \ - unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11, nb12, nb13); \ + unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, \ + nb13_bc); \ dma_queue_push(dma_queue, \ dma_make_ptr(src1_vtcm, data_src1 + src1_pref_off), \ uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, pref_block_size); \