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)