diff --git a/ggml/src/ggml-hexagon/htp/concat-ops.c b/ggml/src/ggml-hexagon/htp/concat-ops.c index 1fa6ec1bdf..259a2c46d7 100644 --- a/ggml/src/ggml-hexagon/htp/concat-ops.c +++ b/ggml/src/ggml-hexagon/htp/concat-ops.c @@ -20,6 +20,7 @@ struct htp_concat_context { uint32_t nrows; uint32_t elem_start; uint32_t nelems; + uint32_t nplanes; struct fastdiv_values div_ne0; struct fastdiv_values div_ne1; struct fastdiv_values div_ne2; @@ -60,39 +61,47 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void * struct htp_thread_trace * tr = &octx->ctx->trace[ith]; - for (uint32_t i = start_i; i < end_i; i += block_i) { - uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i; + for (uint32_t p = 0; p < cctx->nplanes; p++) { + const uint32_t i3 = p / dst->ne[2]; + const uint32_t i2 = p - i3 * dst->ne[2]; + const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3]; + const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3]; + const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3]; - uint32_t src1_width_bytes = current_block_i * sizeof(float); - const dma_addr_t src1_addr = src1->data + i * src1->nb[1]; - dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0); + for (uint32_t i = start_i; i < end_i; i += block_i) { + uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i; - uint32_t src0_row_bytes = src0_ne0 * sizeof(float); - const dma_addr_t src0_addr = src0->data + i * src0->nb[1]; - dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i); + uint32_t src1_width_bytes = current_block_i * sizeof(float); + const dma_addr_t src1_addr = src1_plane + i * src1->nb[1]; + dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0); - dma_queue_pop(dma_q); // src1 + uint32_t src0_row_bytes = src0_ne0 * sizeof(float); + const dma_addr_t src0_addr = src0_plane + i * src0->nb[1]; + dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i); - HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride); + dma_queue_pop(dma_q); // src1 - htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); - for (uint32_t j = 0; j < src1_ne0_padded; j += 32) { - #pragma unroll(4) - for (uint32_t ii = 0; ii < current_block_i; ii++) { - size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float)); - Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv); - uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float); - hvx_vmemu(dst_ptr) = vtcm_tmp[ii]; + HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride); + + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); + for (uint32_t j = 0; j < src1_ne0_padded; j += 32) { + #pragma unroll(4) + for (uint32_t ii = 0; ii < current_block_i; ii++) { + size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float)); + Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv); + uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float); + hvx_vmemu(dst_ptr) = vtcm_tmp[ii]; + } } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); + + dma_queue_pop(dma_q); // src0 + + const dma_addr_t dst_addr = dst_plane + i * dst->nb[1]; + dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i); + + dma_queue_pop(dma_q); } - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); - - dma_queue_pop(dma_q); // src0 - - const dma_addr_t dst_addr = dst->data + i * dst->nb[1]; - dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i); - - dma_queue_pop(dma_q); } } @@ -131,39 +140,47 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void * struct htp_thread_trace * tr = &octx->ctx->trace[ith]; - for (uint32_t i = start_i; i < end_i; i += block_i) { - uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i; + for (uint32_t p = 0; p < cctx->nplanes; p++) { + const uint32_t i3 = p / dst->ne[2]; + const uint32_t i2 = p - i3 * dst->ne[2]; + const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3]; + const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3]; + const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3]; - uint32_t src1_width_bytes = current_block_i * sizeof(__fp16); - const dma_addr_t src1_addr = src1->data + i * src1->nb[1]; - dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0); + for (uint32_t i = start_i; i < end_i; i += block_i) { + uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i; - uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16); - const dma_addr_t src0_addr = src0->data + i * src0->nb[1]; - dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i); + uint32_t src1_width_bytes = current_block_i * sizeof(__fp16); + const dma_addr_t src1_addr = src1_plane + i * src1->nb[1]; + dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0); - dma_queue_pop(dma_q); // src1 + uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16); + const dma_addr_t src0_addr = src0_plane + i * src0->nb[1]; + dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i); - HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride); + dma_queue_pop(dma_q); // src1 - htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); - for (uint32_t j = 0; j < src1_ne0_padded; j += 64) { - #pragma unroll(4) - for (uint32_t ii = 0; ii < current_block_i; ii++) { - size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16)); - Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv); - uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16); - hvx_vmemu(dst_ptr) = vtcm_tmp[ii]; + HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride); + + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); + for (uint32_t j = 0; j < src1_ne0_padded; j += 64) { + #pragma unroll(4) + for (uint32_t ii = 0; ii < current_block_i; ii++) { + size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16)); + Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv); + uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16); + hvx_vmemu(dst_ptr) = vtcm_tmp[ii]; + } } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); + + dma_queue_pop(dma_q); // src0 + + const dma_addr_t dst_addr = dst_plane + i * dst->nb[1]; + dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i); + + dma_queue_pop(dma_q); } - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i); - - dma_queue_pop(dma_q); // src0 - - const dma_addr_t dst_addr = dst->data + i * dst->nb[1]; - dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i); - - dma_queue_pop(dma_q); } } @@ -237,8 +254,6 @@ int op_concat(struct htp_ops_context * octx) { int dim = octx->op_params[0]; - bool is_2d = dst->ne[2] == 1 && dst->ne[3] == 1; - const uint32_t type_size = (dst->type == HTP_TYPE_F32 || dst->type == HTP_TYPE_I32) ? 4 : 2; bool is_src1_transposed = (src1->nb[0] > src1->nb[1]); bool is_src0_transposed = (src0->nb[0] > src0->nb[1]); @@ -253,7 +268,9 @@ int op_concat(struct htp_ops_context * octx) { void (*worker_func)(unsigned int, unsigned int, void *) = concat_generic; - if (dim == 0 && is_2d && is_src1_transposed && !is_src0_transposed) { + const bool rows_ok = src0->nb[0] == type_size && src1->nb[1] == type_size && dst->nb[0] == type_size; + + if (dim == 0 && is_src1_transposed && !is_src0_transposed && rows_ok) { const uint32_t total_rows = dst->ne[1]; const size_t dst_data_row_size = dst->ne[0] * type_size; uint32_t row_start = 0; @@ -272,6 +289,7 @@ int op_concat(struct htp_ops_context * octx) { cctx.row_start = row_start; cctx.nrows = nrows; + cctx.nplanes = dst->ne[2] * dst->ne[3]; uint32_t block_i = (type_size == 4) ? 32 : 64;