Commit a4d880fd5 for llama.cpp
commit a4d880fd5c7f88713ded6db9f0111893bd78afa6
Author: Ehsan Bateni <ebateni@qti.qualcomm.com>
Date: Wed Sep 30 15:56:44 2026 -0400
Hexagon: optimize ALLREDUCE with support for safe scatter mode (#29757)
* hex-allreduce: add support for safe scatter mode
* hex-allreduce: pare down excessive comments
* hex-allreduce: re-write to remove register spills
---------
Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>
diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 4b913947d..ad0905f6f 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -99,10 +99,11 @@ static int opt_profile = 0; // profiling mode (0-disabled, 1-basic, 2-pmu)
static bool opt_hostbuf = false;
static bool opt_dma64 = false;
-static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
-static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
+static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
+static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_gdn_select = 2; // 2 = HMX -> HVX, 1 = HVX, 0 = CPU (unsupported)
-static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = unfused ALLREDUCE (DMA), 0 = fallback to CPY+FENCE
+static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (default), 1 = unfused ALLREDUCE, 0 = fallback to CPY+FENCE
+static int opt_ar_scatter = 1; // 1 = reduce-scatter the fused ALLREDUCE+ADD (default), 0 = full reduction
// Default PMU events, if profiling with PMU (mode=2) is enabled
// See https://docs.qualcomm.com/doc/80-N2040-60/topic/pmu-events.html
@@ -400,6 +401,7 @@ static bool ggml_hexagon_precompute_allreduce_params(
uint32_t n_ranks,
bool has_add,
bool is_row_bcast,
+ bool is_shard_ok,
struct htp_allreduce_kernel_params * kparams
);
@@ -2795,22 +2797,33 @@ struct ggml_hexagon_opbatch {
return false;
}
- for (uint32_t r = 0; r < n_ranks; r++) {
- const ggml_tensor * ar_src = last_node.inputs[r];
- if (ggml_hexagon_tensors_overlap(add_dst, ar_src)) {
- HEX_VERBOSE("ggml-hex: %s skip ALLREDUCE_ADD fusion: dst overlaps allreduce src %u\n", sess->c_name(), r);
- return false;
- }
- }
+ // scatter is only valid for in-place add within max outputs
+ const bool is_shard_ok = n_ranks <= HTP_OP_MAX_OUTPUTS && add_dst->data == ar_local->data;
struct htp_allreduce_kernel_params new_kparams;
if (!ggml_hexagon_precompute_allreduce_params(
- sess, add_dst, (uint32_t) ar_kparams->rank, (uint32_t) ar_kparams->n_ranks, true, is_row_bcast, &new_kparams
+ sess, add_dst, (uint32_t) ar_kparams->rank, (uint32_t) ar_kparams->n_ranks, true, is_row_bcast, is_shard_ok, &new_kparams
)) {
HEX_VERBOSE("ggml-hex: %s skip ALLREDUCE_ADD fusion: solver failed\n", sess->c_name());
return false;
}
+ const bool scatter_ok = new_kparams.n_dsts > 1;
+ new_kparams.mode = scatter_ok ? HTP_ALLREDUCE_SHARDED_FANOUT : HTP_ALLREDUCE_FULL;
+
+ for (uint32_t r = 0; r < n_ranks; r++) {
+ const ggml_tensor * ar_src = last_node.inputs[r];
+ if (!ggml_hexagon_tensors_overlap(add_dst, ar_src)) {
+ continue;
+ }
+ // in-place aliasing is safe under reduce-scatter since each rank writes disjoint shards
+ if (scatter_ok && r == rank) {
+ continue;
+ }
+ HEX_VERBOSE("ggml-hex: %s skip ALLREDUCE_ADD fusion: dst overlaps allreduce src %u\n", sess->c_name(), r);
+ return false;
+ }
+
if (!try_fuse_common(res_tensor, add_dst)) {
return false;
}
@@ -2828,12 +2841,25 @@ struct ggml_hexagon_opbatch {
memcpy(o.kernel_params, &new_kparams, sizeof(new_kparams));
o.src[2 * n_ranks] = add_tensor(res_tensor);
- o.dst[0] = add_tensor(add_dst);
- for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) {
- o.dst[d] = 0xffff;
+ if (new_kparams.mode == HTP_ALLREDUCE_SHARDED_FANOUT) {
+ // fan out shard to all per-rank partial buffers
+ GGML_ASSERT((uint32_t) new_kparams.n_dsts == n_ranks);
+ for (uint32_t d = 0; d < n_ranks; d++) {
+ GGML_ASSERT(o.src[d] != 0xffff);
+ o.dst[d] = o.src[d];
+ }
+ for (uint32_t d = n_ranks; d < HTP_OP_MAX_OUTPUTS; d++) {
+ o.dst[d] = 0xffff;
+ }
+ } else {
+ o.dst[0] = add_tensor(add_dst);
+ for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) {
+ o.dst[d] = 0xffff;
+ }
}
- HEX_VERBOSE("ggml-hex: %s fused ALLREDUCE+ADD (#%u)\n", sess->c_name(), n_ops - 1);
+ HEX_VERBOSE("ggml-hex: %s fused ALLREDUCE+ADD (#%u) mode=%d n_dsts=%d\n",
+ sess->c_name(), n_ops - 1, (int) new_kparams.mode, (int) new_kparams.n_dsts);
return true;
}
@@ -3774,6 +3800,7 @@ static bool ggml_hexagon_precompute_allreduce_params(
uint32_t n_ranks,
bool has_add,
bool is_row_bcast,
+ bool is_shard_ok,
struct htp_allreduce_kernel_params * kparams
) {
memset(kparams, 0, sizeof(*kparams));
@@ -3793,11 +3820,20 @@ static bool ggml_hexagon_precompute_allreduce_params(
const bool use_1d = is_contiguous && !(has_add && is_row_bcast && ne1 > 1);
if (has_add) {
- kparams->n_dsts = 1;
- if (use_1d) {
+ // sharded reduce-scatter for contiguous in-place add
+ if (opt_ar_scatter && is_shard_ok && use_1d && !is_row_bcast) {
+ const uint32_t rank_chunk_elems = hex_round_up((nelem + n_ranks - 1) / n_ranks, 128);
+ const uint32_t rank_elem_start = (std::min)(rank * rank_chunk_elems, nelem);
+ const uint32_t rank_elem_end = (std::min)(rank_elem_start + rank_chunk_elems, nelem);
+ kparams->n_dsts = (int32_t) n_ranks;
+ kparams->rank_elem_start = (int32_t) rank_elem_start;
+ kparams->rank_nelem = (int32_t) (rank_elem_end - rank_elem_start);
+ } else if (use_1d) {
+ kparams->n_dsts = 1;
kparams->rank_elem_start = 0;
kparams->rank_nelem = (int32_t) nelem;
} else {
+ kparams->n_dsts = 1;
kparams->rank_elem_start = 0;
kparams->rank_nelem = (int32_t) ne1;
}
@@ -3820,6 +3856,8 @@ static bool ggml_hexagon_precompute_allreduce_params(
}
}
+ kparams->mode = (kparams->n_dsts > 1) ? HTP_ALLREDUCE_SHARDED_FANOUT : HTP_ALLREDUCE_FULL;
+
if (use_1d) {
const uint32_t rank_nelem = (uint32_t) kparams->rank_nelem;
const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, (std::max)(1u, rank_nelem / 128));
@@ -3925,9 +3963,8 @@ void ggml_hexagon_session::enqueue_allreduce(
}
ggml_hexagon_precompute_allreduce_params(
- this, dst, rank, n_ranks, false, false,
- (struct htp_allreduce_kernel_params *) ar_node.kernel_params
- );
+ this, dst, rank, n_ranks, false, false, /*is_shard_ok=*/ false,
+ (struct htp_allreduce_kernel_params *) ar_node.kernel_params);
ar_node.name = "ALLREDUCE";
this->enqueue_op(ar_node);
@@ -7999,7 +8036,7 @@ static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct
for (size_t r = 0; r < n_backends; r++) {
auto sess = static_cast<ggml_hexagon_session *>(comm_ctx->backends[r]->context);
struct htp_allreduce_kernel_params kparams;
- if (!ggml_hexagon_precompute_allreduce_params(sess, tensors[r], (uint32_t) r, (uint32_t) n_backends, false, false, &kparams)) {
+ if (!ggml_hexagon_precompute_allreduce_params(sess, tensors[r], (uint32_t) r, (uint32_t) n_backends, false, false, /*is_shard_ok=*/ false, &kparams)) {
return false;
}
}
@@ -8208,6 +8245,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT");
const char * str_gdn_select = getenv("GGML_HEXAGON_GDN_SELECT");
const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT");
+ const char * str_ar_scatter = getenv("GGML_HEXAGON_AR_SCATTER");
const char * str_ndev = getenv("GGML_HEXAGON_NDEV");
const char * str_arch = getenv("GGML_HEXAGON_ARCH");
const char * str_vmem = getenv("GGML_HEXAGON_VMEM");
@@ -8261,6 +8299,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select;
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
+ opt_ar_scatter = str_ar_scatter ? atoi(str_ar_scatter) : opt_ar_scatter;
opt_mbuf = str_mbuf ? strtoul(str_mbuf, NULL, 0) * MiB : opt_mbuf;
opt_vmem = str_vmem ? strtoul(str_vmem, NULL, 0) * MiB : opt_vmem;
opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) != 0 : opt_hostbuf;
diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.c b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
index 7b577befb..7fcc4e866 100644
--- a/ggml/src/ggml-hexagon/htp/allreduce-ops.c
+++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
@@ -38,97 +38,103 @@ struct htp_allreduce_context {
uint8_t * res_spad_base;
};
-#define DEFINE_ALLREDUCE_THREAD_DMA_1D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD) \
-static void allreduce_thread_dma_1d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \
- struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data; \
- struct htp_ops_context * octx = actx->octx; \
- \
- const uint32_t n_ranks = actx->n_ranks; \
- const uint32_t n_dsts = actx->n_dsts; \
- const uint32_t block_elems = actx->block_elems; \
- \
- const uint32_t dr = actx->elems_per_thread; \
- const uint32_t ir0 = actx->rank_elem_start + dr * ith; \
- const uint32_t ir1 = MIN(ir0 + dr, actx->rank_elem_start + actx->rank_nelem); \
- if (ir0 >= ir1) return; \
- \
- struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
- dma_queue * dma_q = octx->ctx->dma[ith]; \
- \
- uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS]; \
- for (uint32_t s = 0; s < n_ranks; s++) { \
- src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread); \
- } \
- uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread); \
- uint8_t * res_spad_base = HAS_ADD ? (actx->res_spad_base + (ith * actx->vtcm_size_per_thread)) : NULL; \
- \
- const size_t spad_half = actx->vtcm_size_per_thread / 2; \
- uint32_t ir_prefetch = ir0; \
- int spad_idx = 0; \
- \
- for (int k = 0; k < 2 && ir_prefetch < ir1; k++) { \
- uint32_t cur_elems = MIN(block_elems, ir1 - ir_prefetch); \
- size_t cur_bytes = cur_elems * sizeof(TYPE); \
- uint8_t * d_spad = dst_spad_base + spad_idx * spad_half; \
- for (uint32_t d = 0; d < n_dsts; d++) { \
- dma_addr_t d_ddr = octx->dsts[d]->data + ir_prefetch * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 0); \
- } \
- for (uint32_t s = 0; s < n_ranks; s++) { \
- uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half; \
- const dma_addr_t s_ddr = octx->src[s]->data + ir_prefetch * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(s_spad, s_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \
- } \
- if (HAS_ADD) { \
- uint8_t * r_spad = res_spad_base + spad_idx * spad_half; \
- const dma_addr_t r_ddr = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(r_spad, r_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \
- } \
- ir_prefetch += cur_elems; \
- spad_idx ^= 1; \
- } \
- \
- for (uint32_t ir = ir0; ir < ir1; ) { \
- uint32_t cur_elems = MIN(block_elems, ir1 - ir); \
- size_t cur_bytes = cur_elems * sizeof(TYPE); \
- uint8_t * d_spad = NULL; \
- for (uint32_t d = 0; d < n_dsts; d++) { \
- d_spad = (uint8_t *) dma_queue_pop(dma_q).src; \
- } \
- uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS]; \
- for (uint32_t s = 0; s < n_ranks; s++) { \
- s_spad[s] = (uint8_t *) dma_queue_pop(dma_q).dst; \
- } \
- uint8_t * r_spad = HAS_ADD ? (uint8_t *) dma_queue_pop(dma_q).dst : NULL; \
- htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \
- HVX_ADD_FN(d_spad, s_spad[0], s_spad[1], cur_elems); \
- for (uint32_t s = 2; s < n_ranks; s++) { \
- HVX_ADD_FN(d_spad, d_spad, s_spad[s], cur_elems); \
- } \
- if (HAS_ADD) { \
- HVX_ADD_FN(d_spad, d_spad, r_spad, cur_elems); \
- } \
- htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \
- for (uint32_t d = 0; d < n_dsts; d++) { \
- dma_addr_t d_ddr = octx->dsts[d]->data + ir * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 1); \
- } \
- if (ir_prefetch < ir1) { \
- uint32_t next_elems = MIN(block_elems, ir1 - ir_prefetch); \
- size_t next_bytes = next_elems * sizeof(TYPE); \
- for (uint32_t s = 0; s < n_ranks; s++) { \
- const dma_addr_t s_next = octx->src[s]->data + ir_prefetch * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(s_spad[s], s_next), next_bytes, next_bytes, next_bytes, 1); \
- } \
- if (HAS_ADD) { \
- const dma_addr_t r_next = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \
- dma_queue_push(dma_q, dma_make_data(r_spad, r_next), next_bytes, next_bytes, next_bytes, 1); \
- } \
- ir_prefetch += next_elems; \
- } \
- ir += cur_elems; \
- } \
- dma_queue_flush(dma_q); \
+#define DEFINE_ALLREDUCE_THREAD_DMA_1D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD) \
+static void allreduce_thread_dma_1d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \
+ struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data; \
+ struct htp_ops_context * octx = actx->octx; \
+ \
+ const uint32_t n_ranks = actx->n_ranks; \
+ const uint32_t n_dsts = actx->n_dsts; \
+ const uint32_t block_elems = actx->block_elems; \
+ \
+ const uint32_t dr = actx->elems_per_thread; \
+ const uint32_t ir0 = actx->rank_elem_start + dr * ith; \
+ const uint32_t ir1 = MIN(ir0 + dr, actx->rank_elem_start + actx->rank_nelem); \
+ if (ir0 >= ir1) return; \
+ \
+ struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
+ dma_queue * dma_q = octx->ctx->dma[ith]; \
+ \
+ const size_t vtcm_thread_offset = ith * actx->vtcm_size_per_thread; \
+ uint8_t * dst_spad_base = actx->dst_spad_base + vtcm_thread_offset; \
+ uint8_t * res_spad_base = HAS_ADD ? (actx->res_spad_base + vtcm_thread_offset) : NULL; \
+ \
+ const size_t spad_half = actx->vtcm_size_per_thread / 2; \
+ uint32_t ir_prefetch = ir0; \
+ int spad_idx = 0; \
+ \
+ for (int k = 0; k < 2 && ir_prefetch < ir1; k++) { \
+ uint32_t cur_elems = MIN(block_elems, ir1 - ir_prefetch); \
+ size_t cur_bytes = cur_elems * sizeof(TYPE); \
+ uint8_t * d_spad = dst_spad_base + spad_idx * spad_half; \
+ for (uint32_t d = 0; d < n_dsts; d++) { \
+ dma_addr_t d_ddr = octx->dsts[d]->data + ir_prefetch * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 0); \
+ } \
+ for (uint32_t s = 0; s < n_ranks; s++) { \
+ uint8_t * s_spad = actx->src_spad_base[s] + vtcm_thread_offset + spad_idx * spad_half; \
+ const dma_addr_t s_ddr = octx->src[s]->data + ir_prefetch * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(s_spad, s_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \
+ } \
+ if (HAS_ADD) { \
+ uint8_t * r_spad = res_spad_base + spad_idx * spad_half; \
+ const dma_addr_t r_ddr = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(r_spad, r_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \
+ } \
+ ir_prefetch += cur_elems; \
+ spad_idx ^= 1; \
+ } \
+ \
+ int comp_spad_idx = 0; \
+ for (uint32_t ir = ir0; ir < ir1; ) { \
+ uint32_t cur_elems = MIN(block_elems, ir1 - ir); \
+ size_t cur_bytes = cur_elems * sizeof(TYPE); \
+ for (uint32_t d = 0; d < n_dsts; d++) { \
+ dma_queue_pop(dma_q); \
+ } \
+ for (uint32_t s = 0; s < n_ranks; s++) { \
+ dma_queue_pop(dma_q); \
+ } \
+ if (HAS_ADD) { \
+ dma_queue_pop(dma_q); \
+ } \
+ uint8_t * d_spad = dst_spad_base + comp_spad_idx * spad_half; \
+ uint8_t * s0_spad = actx->src_spad_base[0] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ uint8_t * s1_spad = actx->src_spad_base[1] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ uint8_t * r_spad = HAS_ADD ? (res_spad_base + comp_spad_idx * spad_half) : NULL; \
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \
+ HVX_ADD_FN(d_spad, s0_spad, s1_spad, cur_elems); \
+ for (uint32_t s = 2; s < n_ranks; s++) { \
+ uint8_t * ss_spad = actx->src_spad_base[s] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ HVX_ADD_FN(d_spad, d_spad, ss_spad, cur_elems); \
+ } \
+ if (HAS_ADD) { \
+ HVX_ADD_FN(d_spad, d_spad, r_spad, cur_elems); \
+ } \
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \
+ for (uint32_t d = 0; d < n_dsts; d++) { \
+ dma_addr_t d_ddr = octx->dsts[d]->data + ir * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 1); \
+ } \
+ if (ir_prefetch < ir1) { \
+ uint32_t next_elems = MIN(block_elems, ir1 - ir_prefetch); \
+ size_t next_bytes = next_elems * sizeof(TYPE); \
+ for (uint32_t s = 0; s < n_ranks; s++) { \
+ uint8_t * s_spad = actx->src_spad_base[s] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ const dma_addr_t s_next = octx->src[s]->data + ir_prefetch * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(s_spad, s_next), next_bytes, next_bytes, next_bytes, 1); \
+ } \
+ if (HAS_ADD) { \
+ uint8_t * r_spad_next = res_spad_base + comp_spad_idx * spad_half; \
+ const dma_addr_t r_next = octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \
+ dma_queue_push(dma_q, dma_make_data(r_spad_next, r_next), next_bytes, next_bytes, next_bytes, 1); \
+ } \
+ ir_prefetch += next_elems; \
+ } \
+ comp_spad_idx ^= 1; \
+ ir += cur_elems; \
+ } \
+ dma_queue_flush(dma_q); \
}
DEFINE_ALLREDUCE_THREAD_DMA_1D(f16, __fp16, hvx_add_f16_aaa, 0)
@@ -156,12 +162,9 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
dma_queue * dma_q = octx->ctx->dma[ith]; \
\
- uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS]; \
- for (uint32_t s = 0; s < n_ranks; s++) { \
- src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread); \
- } \
- uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread); \
- uint8_t * res_spad_base = HAS_ADD ? (IS_ROW_BCAST ? actx->res_spad_base : (actx->res_spad_base + (ith * actx->vtcm_size_per_thread))) : NULL; \
+ const size_t vtcm_thread_offset = ith * actx->vtcm_size_per_thread; \
+ uint8_t * dst_spad_base = actx->dst_spad_base + vtcm_thread_offset; \
+ uint8_t * res_spad_base = HAS_ADD ? (IS_ROW_BCAST ? actx->res_spad_base : (actx->res_spad_base + vtcm_thread_offset)) : NULL; \
\
const size_t spad_half = actx->vtcm_size_per_thread / 2; \
uint32_t r_prefetch = r0; \
@@ -175,7 +178,7 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
dma_queue_push(dma_q, dma_make_data(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, 0); \
} \
for (uint32_t s = 0; s < n_ranks; s++) { \
- uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half; \
+ uint8_t * s_spad = actx->src_spad_base[s] + vtcm_thread_offset + spad_idx * spad_half; \
const dma_addr_t s_ddr = octx->src[s]->data + r_prefetch * octx->src[s]->nb[1]; \
dma_queue_push(dma_q, dma_make_data(s_spad, s_ddr), row_size_aligned, octx->src[s]->nb[1], row_bytes, cur_rows); \
} \
@@ -188,25 +191,30 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
spad_idx ^= 1; \
} \
\
+ int comp_spad_idx = 0; \
for (uint32_t r = r0; r < r1; ) { \
uint32_t cur_rows = MIN(block_rows, r1 - r); \
- uint8_t * d_spad = NULL; \
for (uint32_t d = 0; d < n_dsts; d++) { \
- d_spad = (uint8_t *) dma_queue_pop(dma_q).src; \
+ dma_queue_pop(dma_q); \
} \
- uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS]; \
for (uint32_t s = 0; s < n_ranks; s++) { \
- s_spad[s] = (uint8_t *) dma_queue_pop(dma_q).dst; \
+ dma_queue_pop(dma_q); \
+ } \
+ if (HAS_ADD && !IS_ROW_BCAST) { \
+ dma_queue_pop(dma_q); \
} \
- uint8_t * r_spad = (HAS_ADD && !IS_ROW_BCAST) ? (uint8_t *) dma_queue_pop(dma_q).dst : NULL; \
+ uint8_t * d_spad = dst_spad_base + comp_spad_idx * spad_half; \
+ uint8_t * s0_spad = actx->src_spad_base[0] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ uint8_t * s1_spad = actx->src_spad_base[1] + vtcm_thread_offset + comp_spad_idx * spad_half; \
+ uint8_t * r_spad = (HAS_ADD && !IS_ROW_BCAST) ? (res_spad_base + comp_spad_idx * spad_half) : NULL; \
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r); \
for (uint32_t row = 0; row < cur_rows; row++) { \
uint8_t * d_row = d_spad + row * row_size_aligned; \
- const uint8_t * s0_row = s_spad[0] + row * row_size_aligned; \
- const uint8_t * s1_row = s_spad[1] + row * row_size_aligned; \
+ const uint8_t * s0_row = s0_spad + row * row_size_aligned; \
+ const uint8_t * s1_row = s1_spad + row * row_size_aligned; \
HVX_ADD_FN(d_row, s0_row, s1_row, ne0); \
for (uint32_t s = 2; s < n_ranks; s++) { \
- const uint8_t * ss_row = s_spad[s] + row * row_size_aligned; \
+ const uint8_t * ss_row = actx->src_spad_base[s] + vtcm_thread_offset + comp_spad_idx * spad_half + row * row_size_aligned; \
HVX_ADD_FN(d_row, d_row, ss_row, ne0); \
} \
if (HAS_ADD) { \
@@ -222,15 +230,18 @@ static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith,
if (r_prefetch < r1) { \
uint32_t next_rows = MIN(block_rows, r1 - r_prefetch); \
for (uint32_t s = 0; s < n_ranks; s++) { \
+ uint8_t * s_spad = actx->src_spad_base[s] + vtcm_thread_offset + comp_spad_idx * spad_half; \
const dma_addr_t s_next = octx->src[s]->data + r_prefetch * octx->src[s]->nb[1]; \
- dma_queue_push(dma_q, dma_make_data(s_spad[s], s_next), row_size_aligned, octx->src[s]->nb[1], row_bytes, next_rows); \
+ dma_queue_push(dma_q, dma_make_data(s_spad, s_next), row_size_aligned, octx->src[s]->nb[1], row_bytes, next_rows); \
} \
if (HAS_ADD && !IS_ROW_BCAST) { \
+ uint8_t * r_spad_next = res_spad_base + comp_spad_idx * spad_half; \
const dma_addr_t r_next = octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1]; \
- dma_queue_push(dma_q, dma_make_data(r_spad, r_next), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, next_rows); \
+ dma_queue_push(dma_q, dma_make_data(r_spad_next, r_next), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, next_rows); \
} \
r_prefetch += next_rows; \
} \
+ comp_spad_idx ^= 1; \
r += cur_rows; \
} \
dma_queue_flush(dma_q); \
@@ -248,9 +259,10 @@ static int validate_allreduce(
const struct htp_allreduce_kernel_params * kparams,
uint32_t n_ranks
) {
- if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
+ if (kparams->n_threads == 0 || (uint32_t) kparams->n_threads > octx->ctx->n_threads) {
return HTP_STATUS_INVAL_PARAMS;
}
+ octx->n_threads = (uint32_t) kparams->n_threads;
if (kparams->vtcm_size_per_thread <= 0 || kparams->vtcm_size <= 0) {
return HTP_STATUS_INVAL_PARAMS;
@@ -306,6 +318,7 @@ int op_allreduce(struct htp_ops_context * octx) {
const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
const uint32_t nelem = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
+ const int32_t mode = kparams->mode;
// 1. Entry Barrier: Synchronize all ranks before reading
struct htp_thread_trace * tr0 = &octx->ctx->trace[0];
@@ -419,6 +432,11 @@ int op_allreduce(struct htp_ops_context * octx) {
// 4. Exit Barrier: Synchronize all ranks after writing
htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
+ // drain fan-out DMA writes before exit barrier
+ if (mode == HTP_ALLREDUCE_SHARDED_FANOUT) {
+ asm volatile ("syncht" : : : "memory");
+ }
+
htp_fence_write(my_fence, fence_seq_exit, octx->status);
for (uint32_t j = 0; j < n_ranks; j++) {
diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.h b/ggml/src/ggml-hexagon/htp/allreduce-ops.h
index 0aed2b8b7..54ef60378 100644
--- a/ggml/src/ggml-hexagon/htp/allreduce-ops.h
+++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.h
@@ -17,6 +17,11 @@ enum htp_allreduce_kernel_type {
HTP_ALLREDUCE_KERNEL_DMA_2D,
};
+enum htp_allreduce_mode {
+ HTP_ALLREDUCE_FULL = 0, // all ranks reduce full tensor
+ HTP_ALLREDUCE_SHARDED_FANOUT = 1, // each rank reduces 1/N shard and fans out to all buffers
+};
+
static inline size_t htp_allreduce_vtcm_buffer_count(
uint32_t n_ranks,
uint32_t n_threads,
@@ -42,6 +47,7 @@ struct htp_allreduce_kernel_params {
int32_t rank_nelem;
int32_t n_dsts;
int32_t is_row_bcast;
+ int32_t mode; // enum htp_allreduce_mode: FULL or SHARDED_FANOUT
};
#ifdef __cplusplus
diff --git a/scripts/snapdragon/run.py b/scripts/snapdragon/run.py
index 6d845c341..01093ddd8 100755
--- a/scripts/snapdragon/run.py
+++ b/scripts/snapdragon/run.py
@@ -169,6 +169,7 @@ def main():
parser.add_argument("--hex-fa-select", help="Select Flash Attention kernel (GGML_HEXAGON_FA_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-gdn-select", help="Select Gated Delta Net kernel (GGML_HEXAGON_GDN_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-ar-select", help="Select All-Reduce kernel (GGML_HEXAGON_AR_SELECT) 1:enable,0:disable")
+ parser.add_argument("--hex-ar-scatter", help="Enable (1) or disable (0) reduce-scatter for fused ALLREDUCE+ADD (GGML_HEXAGON_AR_SCATTER)")
parser.add_argument("--hex-etm", help="Enable Embedded Trace Macrocell hardware tracing / trace logging (GGML_HEXAGON_ETM)")
parser.add_argument("--hex-arch", help="Target Hexagon NPU architecture version override (v73, v75, v79, v81, etc.) (GGML_HEXAGON_ARCH)")
parser.add_argument("--hex-optrace", help="Trace buffer size in number of records (GGML_HEXAGON_OPTRACE)")
@@ -310,6 +311,7 @@ def main():
set_env("GGML_HEXAGON_FA_SELECT", args.hex_fa_select)
set_env("GGML_HEXAGON_GDN_SELECT", args.hex_gdn_select)
set_env("GGML_HEXAGON_AR_SELECT", args.hex_ar_select)
+ set_env("GGML_HEXAGON_AR_SCATTER", args.hex_ar_scatter)
set_env("GGML_HEXAGON_ETM", args.hex_etm)
set_env("GGML_HEXAGON_ARCH", args.hex_arch)
set_env("GGML_HEXAGON_OPTRACE", args.hex_optrace)