Commit eafe15a5e for llama.cpp

commit eafe15a5e3d87dd68ae33acf6a7cbd9415a0ac5e
Author: Max Krasnyansky <maxk@qti.qualcomm.com>
Date:   Fri Sep 11 20:46:51 2026 -0700

    hexagon: support for multi-device model split (aka row-split) (#28589)

    * hex-row-split: add support for multi-device row spliting

    Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>

    * hex-mdev: add work splitting to fused kernels

    * hex-mdev: use mdev_ prefix for all multi-device state

    * hex-mdev: make device configuration more expressive to support device groups

    * hex-mdev: fix mdev session init

    * hex-mdev: fused nx (2x,3x) matmuls must update row counts for each w/o

    * hex-mdev: fix MUL_MAT work partitioning bugs introduced by mdev

    * hex-cont: fix crashes with new tests due to wrong striding

    * hex-mdev: move fences after l2flushes

    * hex-cont: fix work splitting for mnpu -- align chunks to cachelines

    * hex-mdev: fix CPY tests with multi-dev

    * hex-mmid: fix work partitioning with mnpu

    * hex-mm: fix test failures with mdev

    * hex-binary: fix work partitioning for mdev

    * hex-argsort: fix mdev partitioning

    * hex-mdev: fix work partitioning and general updates for all simple ops

    * hex-fa: fix mdev work splitting issues

    * hex-mdev: fixing more failing ops test

    * hex-mdev: update the rest of the ops

    * hex-mdev: refactor all mdev splitting logic to be contained within if (mdev_count > 1) {...}

    * hex-mdev: fix macros

    * hex-mdev: simplify session flush logic

    * hex-sync: fix recursion in session flush

    * hex-mdev: factor out fence buffer and allocator

    * hex-fence: make fence allocation more robust with reserved slots for mdev

    * hex-mdev: keep all mdev state in htp_mdev_group

    * hex-mdev: further cleanup mdev group handling at the host

    * hex-mdev: update group idx in the opbatch before serializing

    * hex-batch: remove separate op_pending and use batch_req/rsp_seq

    * hex-async: workaround another missing tensor_init in ggml-meta

    * hex-fence: cleanup and robustify fences and error handling in multi-device scenarios

    * hex-ar: improve ALLREDUCE error handling

    * hex-async: robust error handling for op_cpy_fence

    * hex-async: use seq0 from allreduce context to allocate fence_seq

    * hex-mdev: fix remaining issues with fence and barrier clearing in CPY_FENCE

    * hex-misc: realign macros and fix misplaces trace events

    * hex-misc: align macros

    * hex-mdev: fix unclone buffer re-entrancy

    * hex-glu: fix mdev partitioning logic

    * hex-mdev: make buffer uncloning/cleanup work with tensor-split scenarios

    * hex-mdev: tighten up the can_split check in act-ops

    * hex-mdev: factor out common bits of the partitioning logic

    * hex-mm: minor realignment of the macros

    * hex-bufs: fix incorrectly placed assert for MAX_BUFS

    * hex-pad: tighten up gating checks for PAD

    * hex-kparams: make sure all kernels properly use kparams->n_threads

    * hex-docs: update user and developer docs with new features and detailed guide for ops development

    * hex-scripts: update run script to properly parse dev groups

    * hex-misc: formatting

    * hex-sess: minor cleanup for session init

    * hex-ar: fix vtcm size calc in allreduce kparams

    * hex-scripts: fix flake8 warnings

    * hex-rope: update ROPE to support mdev work split

    * hex-ops: remove redunant checks and minor reformat

    * hex-dev-guide: update dev-guide to avoid redundant null checks

    * hex-async: improve event_wait, event_sync and fence implementations

    * hex-async: remove synchronous flush from event_sync

    * hex-async: symplify fence recovery protocol and make sync more robust

    * hex-async: futher simplify error recovery for fences

    * hex-err: return status instead of just -1

    * hex-async: print all seq nums in hex

    * hex-async: make sure fences flush dirty ranges

    * hex-async: add dirty ranges merging to reduce fence flushes

    * hex-async: properly sync before freeing the event

    * hex-async: make sure fence owner session is not overriden

    * hex-async: more fence write order more robust

    * hex-async: make sure not to fuse ALLREDUCE+ADD if their dsts overlap

    * hex-fusion: cleanup redundant checks

    ---------

    Co-authored-by: Alexander Lu <alexlu@qti.qualcomm.com>

diff --git a/docs/backend/snapdragon/README.md b/docs/backend/snapdragon/README.md
index 391c8bf23..5d32a5877 100644
--- a/docs/backend/snapdragon/README.md
+++ b/docs/backend/snapdragon/README.md
@@ -188,7 +188,7 @@ llama_memory_breakdown_print: |   - Host               |                  439 =
 Op test for MUL_MAT:

 ```
-~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --hex-hostbuf 0 --devices HTP0:0 -- test-backend-ops -b HTP0:0 -o MUL_MAT
+~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --devices HTP0:0 -- test-backend-ops -b HTP0:0 -o MUL_MAT
 ...
 Backend 2/3: HTP0:0
 Device description: Hexagon
@@ -213,14 +213,109 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
 | llama 1B Q4_0  | 729.75 MiB | 1.24 B | HTP        |  99 |       4 |     128 |    0 |  tg64 |  51.54 ± 1.13 |
 ```

+## Multi-Device Execution Modes
+
+The Hexagon backend supports multiple execution and partitioning modes to accommodate different model sizes, memory
+constraints, and single- or multi-NPU hardware topologies:
+
+### 1. Single-Device Mode with Dynamic Buffer Mapping
+
+Runs the model on a single NPU session (e.g. `HTP0` or `HTP0:0`).
+
+A single NPU session provides ~3.5GB of available virtual address space. For models larger than 3.5GB, the backend
+automatically maps and unmaps weight buffers during graph execution. This allows large models to run on a single NPU
+without manual configuration:
+
+```bash
+./scripts/snapdragon/run.py --target adb --devices HTP0:0 -- \
+    llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf -ngl 99 -p "Hello"
+```
+
+### 2. Layer-Split Mode across Virtual Sessions (`HTP0,HTP1,...` or `HTP0:0,HTP0:1,...`)
+
+Partitions model layers at load time across multiple virtual sessions hosted on a single physical NPU.
+
+Each virtual session acts as an independent backend device from llama.cpp's perspective (similar to multiple GPUs).
+Because layers are permanently distributed across sessions, each session's allocated weights remain within its private 3.5GB
+address space window, eliminating runtime buffer re-mapping overhead.
+
+Here is an example of running the GPT-OSS-20B model on a Snapdragon device using 4 virtual sessions on a single NPU:
+
+```bash
+./scripts/snapdragon/run.py --target adb \
+    --devices HTP0:0,HTP0:1,HTP0:2,HTP0:3 -- \
+    llama-cli --load-mode none -m /data/local/tmp/gguf/gpt-oss-20b-Q4_0.gguf -t 4 \
+    --ctx-size 8192 --batch-size 128 -ctk q8_0 -ctv q8_0 -fa on -ngl 99 -no-cnv -f surfing.txt
+```
+
+Log output snippet:
+
+```
+...
+llama_model_loader: - type  f32:  289 tensors
+llama_model_loader: - type q4_0:   96 tensors
+llama_model_loader: - type q8_0:    2 tensors
+llama_model_loader: - type mxfp4:  72 tensors
+...
+load_tensors: offloaded 25/25 layers to GPU
+load_tensors:          CPU model buffer size =  1182.09 MiB
+load_tensors:       HTP0:1 model buffer size =  2512.58 MiB
+load_tensors:       HTP0:3 model buffer size =  2093.83 MiB
+load_tensors:       HTP0:0 model buffer size =  2931.34 MiB
+load_tensors:       HTP0:2 model buffer size =  2512.58 MiB
+...
+llama_perf_context_print: prompt eval time =    3843.67 ms /   197 tokens ( 19.51 ms per token, 51.25 tokens per second)
+llama_perf_context_print:        eval time =    1686.13 ms /    31 runs   ( 54.39 ms per token, 18.39 tokens per second)
+llama_perf_context_print:       total time =    6266.30 ms /   228 tokens
+llama_memory_breakdown_print: | memory breakdown [MiB] | total   free    self   model   context   compute    unaccounted |
+llama_memory_breakdown_print: |   - HTP0:0 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
+llama_memory_breakdown_print: |   - HTP0:1 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
+llama_memory_breakdown_print: |   - HTP0:2 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
+llama_memory_breakdown_print: |   - HTP0:3 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
+llama_memory_breakdown_print: |   - Host               |                 1476 =  1208 +     105 +     162                |
+```
+
+### 3. Tensor-Split Mode across Physical Devices (`HTP0:0,HTP1:0,...`)
+
+Distributes model tensors across distinct physical NPU hardware cores using llama.cpp's tensor parallelism
+(`--split-mode tensor`).
+
+Tensors are partitioned across physical NPUs for parallel execution (proportions are distributed equally by default without
+needing an explicit `--tensor-split` option):
+
+```bash
+./scripts/snapdragon/run.py --target adb \
+    --devices HTP0:0,HTP1:0 -- \
+    llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf --split-mode tensor -ngl 99 -p "Hello"
+```
+
+### 4. Row-Split Multi-Device Mode via Device Grouping (`HTP0[0-1]`)
+
+Groups multiple physical NPU cores into a single logical device using bracket notation (`HTP0[0-1]` or `HTP0[0,1]`).
+
+Unlike host-level tensor-splitting, row-splitting is executed entirely inside the Hexagon backend:
+
+```bash
+./scripts/snapdragon/run.py --target adb \
+    --devices 'HTP0[0-1]' -- \
+    llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf -ngl 99 -p "Hello"
+```
+
+You can also combine row-splitting with layer-splitting across multiple grouped devices (e.g. `--devices 'HTP0[0-1],HTP1[2-3]'`
+on 4 physical NPUs, or `--devices 'HTP0[0-1:0],HTP1[0-1:1]'` on 2 physical NPUs using virtual sessions 0 and 1).
+
 ## Environment variables

 - `GGML_HEXAGON_DEVICES` (default: not set, defaults to HTP0 session)
-  Controls which NPU devices and sessions to allocate. Can be configured as:
-  - A single integer `N`: Allocates `N` sessions named `HTP0`, `HTP1`, ..., `HTP<N-1>` (behaves identically to `GGML_HEXAGON_NDEV=N`).
-  - A comma-separated list of device names in `HTP<physical_idx>:<virtual_idx>` format (or legacy `HTP<idx>` format). For example, `HTP0:0,HTP0:1` creates two virtual
-    sessions on the first physical NPU (useful for memory limits). `HTP0:0,HTP1:0` allocates one session on each of the two physical NPUs
-    on a dual-NPU device.
+  Controls which NPU devices and sessions to allocate. Configurable via `--devices` in `run.py`:
+  - `N` (single integer): Allocates `N` virtual sessions named `HTP0`, `HTP1`, ..., `HTP<N-1>` on physical NPU 0.
+  - `HTP<phys>:<virt>,...`: Comma-separated list of individual devices specifying physical and virtual index:
+    - `HTP0:0,HTP0:1`: Two virtual sessions on physical NPU 0 (layer-split on single NPU).
+    - `HTP0:0,HTP1:0`: One session on physical NPU 0 and one on physical NPU 1 (tensor-split across physical cores).
+  - `HTP<name>[<phys_spec>]`: Device grouping syntax for row-split multi-device execution:
+    - `HTP0[0-1]`: A single logical device `HTP0` that groups physical cores 0 and 1.
+    - `HTP0[0-1],HTP1[2-3]`: Two layer-split devices across 4 physical NPUs (cores 0-1 and 2-3).
+    - `HTP0[0-1:0],HTP1[0-1:1]`: Two layer-split devices across 2 physical NPUs using virtual sessions 0 and 1.

 - `GGML_HEXAGON_NDEV` (deprecated)
   Replaced by `GGML_HEXAGON_DEVICES`. Controls the number of virtual sessions to allocate on physical NPU `0`.
@@ -229,9 +324,8 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
 - `GGML_HEXAGON_NHVX=0`
   Controls the number of HVX hardware threads to use. The default is all (actual number varies depending on the hardware version).

-- `GGML_HEXAGON_HOSTBUF=1`
-  Controls whether the Hexagon backend allocates host buffers. By default, all buffers except for REPACK are host buffers.
-  This option is required for testing Ops that require REPACK buffers (MUL_MAT and MUL_MAT_ID).
+- `GGML_HEXAGON_HOSTBUF=1` (default: 0, disabled)
+  Enables allocating host buffers for debugging. By default, host buffers are disabled.

 - `GGML_HEXAGON_VERBOSE=1`
   Enables verbose logging of Ops from the backend. Example output:
@@ -246,23 +340,26 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
   ```

 - `GGML_HEXAGON_PROFILE=1`
-  Enables Op profiling:
+  Enables Op profiling (configurable via `--hex-profile` in `run.py`):

-  - `1` Basic profile with per-op `usecs` and `cycles` counters
-  - `2` Extended profile with per-op `usecs`, `cycles` and default PMU counter data
-  - `0x1,...,0x8` Extended profile with per-op `usecs`, `cycles` and custom PMU counter data
+  - `1`: Basic profile with per-op `usecs` and `cycles` counters
+  - `2`: Extended profile with per-op `usecs`, `cycles` and default PMU counter data
+  - `0x1,...,0x8`: Extended profile with per-op `usecs`, `cycles` and custom PMU counter data

-  The logging output can be either saved into a file for post-processing or it can be piped directly into the post-processing tool
-  to generate the report.
-  Examples:
+  The logging output can be saved to a file or piped directly into the post-processing script:

-      `GGML_HEXAGON_PROFILE=1 ./scripts/snapdragon/run.py --target adb -- llama-cli ... |& ./scripts/snapdragon/ggml-hexagon-profile.py -`
+  ```bash
+  ./scripts/snapdragon/run.py --target adb --hex-profile 1 -- llama-cli ... |& \
+      ./scripts/snapdragon/ggml-hexagon-profile.py -
+  ```

 - `GGML_HEXAGON_OPFILTER=regex`
-  Allows filtering (disabling) Ops that match the regex pattern:
+  Filters (disables) Ops matching the regex pattern (configurable via `--hex-opfilter` in `run.py`):

-  Examples:
-
-      `GGML_HEXAGON_OPFILTER="FLASH_ATTN_EXT" ./scripts/snapdragon/run.py --target adb -- llama-cli ...` - Disable Flash Attention on Hexagon (falls back to CPU or GPU)
-      `GGML_HEXAGON_OPFILTER="ADD\|SUB" ./scripts/snapdragon/run.py --target adb -- llama-cli ...` - Disable ADD and SUB on Hexagon (fall back to CPU or GPU)
+  ```bash
+  # Disable Flash Attention on Hexagon (falls back to CPU or GPU)
+  ./scripts/snapdragon/run.py --target adb --hex-opfilter "FLASH_ATTN_EXT" -- llama-cli ...

+  # Disable ADD and SUB on Hexagon (fall back to CPU or GPU)
+  ./scripts/snapdragon/run.py --target adb --hex-opfilter "ADD|SUB" -- llama-cli ...
+  ```
diff --git a/docs/backend/snapdragon/developer.md b/docs/backend/snapdragon/developer.md
index d7d9f2a27..633643c16 100644
--- a/docs/backend/snapdragon/developer.md
+++ b/docs/backend/snapdragon/developer.md
@@ -2,16 +2,16 @@

 ## Backend libraries

-The Hexagon backend consist of two parts:
+The Hexagon backend consists of two parts:

   - `libggml-hexagon`
-    This is the regular CPU-side GGML backend library, either shared or statically linked
+    This is the regular CPU-side GGML backend library, either shared or statically linked.

   - `libggml-htp-vNN`
     This is the NPU-side (HTP stands for Hexagon Tensor Processor) shared library that contains the Op dispatcher and kernels.
     The correct library is selected automatically at runtime based on the HW version.

-Here is an example of the build artifacts
+Here is an example of the build artifacts:

 ```
 ~/src/llama.cpp$ ls -l pkg-adb/llama.cpp/lib/libggml*
@@ -26,75 +26,307 @@ pkg-adb/llama.cpp/lib/libggml-htp-v81.so

 ## Memory buffers

-Hexagon NPU backend takes advantage of the Snapdragon's unified memory model where all buffers are fully accessible by the CPU and GPU.
-The NPU does have a dedicated tightly-coupled memory called VTCM but that memory is used only for intermediate data (e.g. dynamically
-quantized tensors) or temporary data (chunks of the weight tensors fetched via DMA).
-
-Please note that currently the Hexagon backend does not implement SET/GET_ROWS Ops because there is no advantage in offloading those
-to the NPU at this point.
-
-The backend does allocates non-host buffers for the tensors with datatypes that require repacking: Q4_0, Q8_0, MXFP4.
-From the MMU perspective these buffers are still regular buffers (normal access by the CPU) they are marked as non-host simply to force
-the repacking.
+The Hexagon NPU backend takes advantage of Snapdragon unified memory where all DDR buffers are accessible by CPU, GPU, and NPU.
+The NPU has dedicated tightly-coupled memory called VTCM (Vector Tightly-Coupled Memory). VTCM is used for intermediate data (such as
+dynamically quantized activations) and streaming buffers (chunks of weight and activation tensors fetched via DMA).

 ## Large model handling

-Hexagon NPU sessions (aka Process Domains (PD) in the Hexagon SDK) are limited to a maximum memory mapping window of around 3.5GB.
+Hexagon NPU sessions have a 32-bit virtual address space window of around 3.5GB.
 In llama.cpp/GGML, each Hexagon session is mapped to a single GGML backend device (e.g., `HTP0:0`, `HTP0:1`, etc. when using
 `GGML_HEXAGON_DEVICES`, or `HTP0`, `HTP1` in legacy mode).

-To support running models larger than 3.5GB on a single device, the Hexagon backend dynamically maps and unmaps execution buffers
-during the graph execution cycle to stay within the Process Domain window. This enables large models to run successfully on a single
-NPU device.
+To support running models larger than 3.5GB on a single device, the Hexagon backend dynamically maps and unmaps buffers:
+- Buffers are allocated in shared DDR (RPCMEM) via file descriptors (`fastrpc_mmap` using `FASTRPC_MAP_FD_DELAYED`).
+- Pinned buffers (such as KV cache and active compute buffers) remain mapped throughout execution.
+- Inactive weight buffers are dynamically mapped into the NPU session via `HAP_mmap()` during batch buffer preparation
+  (`prep_op_bufs()` in `htp/main.c`) and unmapped via `htp_iface_munmap()` when no longer needed by the active batch.
+- This dynamic sliding window allows a single NPU session to execute models that exceed the 3.5GB window.
+
+Alternatively, users can partition and split the model across multiple virtual sessions or physical NPUs using layer-splitting,
+tensor-splitting, or row-splitting modes. For user-facing execution modes and examples, see the
+[Snapdragon user guide](README.md#multi-device-execution-modes).
+
+## Op and Kernel Development Guidelines
+
+Writing high-performance operators for Hexagon requires following specific guidelines.
+
+### DDR -> DMA -> VTCM Execution Pipeline
+
+- Strongly prefer the `DDR -> DMA -> VTCM -> compute (HVX/HMX) -> VTCM -> DMA -> DDR` data flow.
+- Direct HVX reads/writes from/to DDR are less efficient and should only be used as a fallback.
+- The DMA queue is a strict FIFO where operations must be pushed and popped in strict order.
+- Follow the pipelined multi-buffering sequence properly (typically 2x to 16x buffering) so every push has a corresponding pop:
+
+  1. In the prologue, push initial DDR -> VTCM transfers to prime the pipeline.
+  2. In the loop body, wait for buffer N via DMA pop, launch HVX/HMX compute on buffer N, push VTCM -> DDR writeback of result N,
+     and push DDR -> VTCM prefetch of buffer N+2.
+  3. In the epilogue, pop all remaining in-flight transfers to drain the pipeline.
+
+- Because every push must be matched by a pop, `dma_queue_flush()` is not required when the pipeline sequence is followed
+  properly. Flushing is only used in rare exceptions where a batch of operations is pushed without individual pops.
+- Use the DMA queue interface from [`dma-queue.h`](../../../ggml/src/ggml-hexagon/htp/dma-queue.h)
+  (`dma_queue_push_ddr_to_vtcm`, `dma_queue_pop`, `dma_queue_push_vtcm_to_ddr`).
+  See [`cumsum-ops.c`](../../../ggml/src/ggml-hexagon/htp/cumsum-ops.c) and
+  [`act-ops.c`](../../../ggml/src/ggml-hexagon/htp/act-ops.c) for reference implementations.
+
+### Avoid Scalar Reads and Writes to VTCM
+
+- Access VTCM data using DMA transfers or HVX/HMX vector instructions rather than scalar reads and writes.
+
+### Avoid Scalar Division in Inner Loops
+
+- Hexagon cores do not have hardware division instructions.
+- For recurring divisions across iterations or threads, use `fastdiv` from
+  [`hex-fastdiv.h`](../../../ggml/src/ggml-hexagon/htp/hex-fastdiv.h) with precomputed divisors (such as
+  `octx->ctx->mdev.count_div` or `octx->n_threads_div`).
+- Do not call `init_fastdiv_values()` for single-use divisions; use standard compiler division (`/`) instead.
+
+### Host-Side Precomputation via `kernel_params`
+
+- Precompute tensor shapes, strides, scale conversions, tiling layouts, and validation checks on the host CPU during graph
+  preparation in [`ggml-hexagon.cpp`](../../../ggml/src/ggml-hexagon/ggml-hexagon.cpp).
+- Pack precomputed parameters into the operator's fixed `kernel_params` structure in `htp_op_node` (such as
+  `htp_mm_kernel_params`, `htp_unary_kernel_params`, `htp_fa_kernel_params`, `htp_get_rows_kernel_params`).
+- The NPU executes directly using `octx->kernel_params` without redundant runtime metadata extraction or validation.
+- **Strict Host-Kernel Alignment**:
+  - Verify that parameters calculated by the host CPU are strictly honored by the NPU kernel.
+  - Ensure the kernel does not ignore host-computed fields (for example, falling back to `octx->n_threads` instead of
+    using `kparams->n_threads`, or ignoring precomputed `tasks_per_thread` and chunk counts).
+  - Both human developers and coding agents must audit both sides of the interface: ensure fields populated in `kernel_params`
+    in [`ggml-hexagon.cpp`](../../../ggml/src/ggml-hexagon/ggml-hexagon.cpp) are actively and consistently utilized by the
+    corresponding operator entry point and worker threads in `htp/*-ops.c`.
+
+### Tracing Instrumentation
+
+- All kernels must include trace events for performance profiling and timeline visualization in Perfetto
+  ([`hex-profile.h`](../../../ggml/src/ggml-hexagon/htp/hex-profile.h)).
+- Surround compute sections with `htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) info)` and
+  `htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) info)`.
+- Use specific event types for major phases:
+  - `HTP_TRACE_EVT_HVX_COMP`: Vector compute execution.
+  - `HTP_TRACE_EVT_DMA`: DMA transfer wait or poll cycles.
+  - `HTP_TRACE_EVT_FENCE`: Multi-device fence barrier synchronization.
+  - `HTP_TRACE_EVT_L2FLUSH`: L2 cache cleaning operations.
+- Pass meaningful progress metrics (such as row index, chunk index, or token index) in the 16-bit `info` parameter.
+
+### Work Queue and Threading
+
+- Distribute parallel work across NPU worker threads using the thread pool work queue:
+
+  ```c
+  work_queue_run(ctx->work_queue, worker_func, &op_ctx, n_threads);
+  ```
+
+- Keep worker functions independent and re-entrant. Worker threads should only operate on their designated chunk of rows or elements.
+
+### Avoid Redundant Defensive NULL Checks
+
+- Do not add defensive NULL checks or assertions for internal framework pointers or required graph operands and outputs.
+  Internal pointers include `ctx`, `octx`, local context structs like `*ctx`, `kparams`, and worker callback `data`.
+- These pointers are architectural invariants during kernel execution and host-side graph preparation.
+  Graph compute receives allocated nodes with valid required `node->src[N]` and `node->data` pointers.
+- Do not turn an invariant violation into an unsupported operation or missed fusion.
+  Checks such as `if (!octx || !octx->ctx)` clutter the code, obscure intent, and hide upstream errors.
+- **Distinction**: `octx->src[N]` pointers *can* be NULL by design and must be checked when optional.
+  Examples include attention masks, optional bias or weights in fused kernels, and frequency factors.
+
+### Multiline Macro Formatting
+
+- Keep trailing backslashes in multiline `#define` macros cleanly aligned to a consistent column.
+- Avoid trailing whitespace after macro backslashes.
+- Use [`scripts/snapdragon/ggml-hexagon-align-macros.py`](../../../scripts/snapdragon/ggml-hexagon-align-macros.py) to inspect, diff,
+  or automatically align macro definitions across Hexagon kernel sources:
+
+  ```bash
+  # Check for misaligned macros
+  python3 scripts/snapdragon/ggml-hexagon-align-macros.py ggml/src/ggml-hexagon/htp/
+
+  # Fix misaligned macros in-place
+  python3 scripts/snapdragon/ggml-hexagon-align-macros.py --fix ggml/src/ggml-hexagon/htp/
+  ```
+
+## Multi-Device Partitioning (mdev)
+
+Multi-device (mdev) mode enables row-level tensor parallel execution across multiple physical NPU cores or virtual NPU
+sessions.
+
+### 128-Byte Cache Line Alignment
+
+- Shared tensor buffers reside in DDR (RPCMEM) with a 128-byte cache line granularity
+  (`HEX_L2_LINE_SIZE` = 128 bytes, `HTP_TENSOR_MDEV_LINE_SIZE`).
+- **Rule**: Multi-device work partitions must align destination write regions to 128-byte cache line boundaries so distinct
+  devices never share or overwrite the same cache line.
+
+### Partitioning Helpers in `htp-tensor.h`
+
+Common partitioning logic is factored into reusable inline helpers in
+[`htp-tensor.h`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h):
+
+1. [`htp_tensor_mdev_rows_per_chunk`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L67):
+   Determines the minimum number of rows per chunk so that the chunk byte size is a multiple of 128 bytes:
+
+   ```
+   rows_per_chunk = 128 / hex_gcd_u32(row_size, 128)
+   ```
+
+   If row stride `nb[1]` is already a multiple of 128 bytes, `rows_per_chunk = 1`.
+   Returns `false` if the tensor cannot be safely row-partitioned (such as unaligned base pointer, permuted layout,
+   or non-128-byte aligned outer strides).

-Alternatively, users can choose to use standard llama.cpp/GGML layer-splitting mode to partition and split the model across
-multiple Hexagon devices or virtual sessions (which behave like multiple GPUs from the offload and splitting perspective).
+2. [`htp_tensor_mdev_partition`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L94):
+   Calculates the per-device work range `struct htp_tensor_mdev_range { uint32_t start; uint32_t count; }` given
+   `total_units`, `units_per_chunk`, `mdev_idx`, `mdev_count`, and the precomputed `mdev_count_div`.
+   Handles chunk distribution across devices, assigns remainder units to the last device, and automatically triggers
+   single-device fallback when partitioning is unsafe.

-Here is an example of running GPT-OSS-20B model on a Snapdragon device using 4 virtual sessions on a single NPU (physical index 0).
+### Row-Partitioned Operators

+For row-wise operators
+(such as activations in [`act-ops.c`](../../../ggml/src/ggml-hexagon/htp/act-ops.c),
+binary ops in [`binary-ops.c`](../../../ggml/src/ggml-hexagon/htp/binary-ops.c),
+unary ops in [`unary-ops.c`](../../../ggml/src/ggml-hexagon/htp/unary-ops.c), and
+sameshape copies in [`cpy-ops.c`](../../../ggml/src/ggml-hexagon/htp/cpy-ops.c)):
+
+```c
+const uint32_t total_rows   = ne01 * ne02 * ne03;
+const size_t   dst_row_size = dst->ne[0] * elem_size;
+
+uint32_t row_start = 0;
+uint32_t nrows     = total_rows;
+
+if (octx->ctx->mdev.count > 1) {
+    uint32_t rows_per_chunk = 0;
+    htp_tensor_mdev_rows_per_chunk(dst, elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
+    const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
+        total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+    row_start = range.start;
+    nrows     = range.count;
+}
+
+if (nrows == 0) {
+    return HTP_STATUS_OK;
+}
 ```
-~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --devices HTP0:0,HTP0:1,HTP0:2,HTP0:3 -- llama-cli --load-mode none -m /data/local/tmp/gguf/gpt-oss-20b-Q4_0.gguf -t 4 --ctx-size 8192 --batch-size 128 -ctk q8_0 -ctv q8_0 -fa on -ngl 99 -no-cnv -f surfing.txt
-...
-llama_model_loader: - type  f32:  289 tensors
-llama_model_loader: - type q4_0:   96 tensors
-llama_model_loader: - type q8_0:    2 tensors
-llama_model_loader: - type mxfp4:  72 tensors
-...
-load_tensors: offloaded 25/25 layers to GPU
-load_tensors:          CPU model buffer size =  1182.09 MiB
-load_tensors:       HTP0:1 model buffer size =  2512.58 MiB
-load_tensors:       HTP0:3 model buffer size =  2093.83 MiB
-load_tensors:       HTP0:0 model buffer size =  2931.34 MiB
-load_tensors:       HTP0:2 model buffer size =  2512.58 MiB
-...
-llama_context: n_ctx_per_seq (8192) < n_ctx_train (131072) -- the full capacity of the model will not be utilized
-llama_context:        CPU  output buffer size =     0.77 MiB
-llama_kv_cache_iswa: creating non-SWA KV cache, size = 8192 cells
-llama_kv_cache:     HTP0:1 KV buffer size =    25.50 MiB
-llama_kv_cache:     HTP0:3 KV buffer size =    25.50 MiB
-llama_kv_cache:     HTP0:0 KV buffer size =    25.50 MiB
-llama_kv_cache:     HTP0:2 KV buffer size =    25.50 MiB
-llama_kv_cache: size =  102.00 MiB (  8192 cells,  12 layers,  1/1 seqs), K (q8_0):   51.00 MiB, V (q8_0):   51.00 MiB
-llama_kv_cache_iswa: creating     SWA KV cache, size = 256 cells
-llama_kv_cache:     HTP0:1 KV buffer size =     0.80 MiB
-llama_kv_cache:     HTP0:3 KV buffer size =     0.53 MiB
-llama_kv_cache:     HTP0:0 KV buffer size =     1.06 MiB
-llama_kv_cache:     HTP0:2 KV buffer size =     0.80 MiB
-llama_kv_cache: size =    3.19 MiB (   256 cells,  12 layers,  1/1 seqs), K (q8_0):    1.59 MiB, V (q8_0):    1.59 MiB
-llama_context:     HTP0:0 compute buffer size =    16.06 MiB
-llama_context:     HTP0:1 compute buffer size =    16.06 MiB
-llama_context:     HTP0:2 compute buffer size =    16.06 MiB
-llama_context:     HTP0:3 compute buffer size =    16.06 MiB
-llama_context:        CPU compute buffer size =    98.19 MiB
-...
-llama_perf_context_print: prompt eval time =    3843.67 ms /   197 tokens ( 19.51 ms per token, 51.25 tokens per second)
-llama_perf_context_print:        eval time =    1686.13 ms /    31 runs   ( 54.39 ms per token, 18.39 tokens per second)
-llama_perf_context_print:       total time =    6266.30 ms /   228 tokens
-llama_perf_context_print:    graphs reused =         30
-llama_memory_breakdown_print: | memory breakdown [MiB] | total   free    self   model   context   compute    unaccounted |
-llama_memory_breakdown_print: |   - HTP0:0 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
-llama_memory_breakdown_print: |   - HTP0:1 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
-llama_memory_breakdown_print: |   - HTP0:2 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
-llama_memory_breakdown_print: |   - HTP0:3 (Hexagon)   |  2048 = 2048 + (   0 =     0 +       0 +       0) +           0 |
-llama_memory_breakdown_print: |   - Host               |                 1476 =  1208 +     105 +     162                |
+
+### Element-Partitioned Operators
+
+For flat element-wise operations (such as reshape copies in
+[`cpy-ops.c`](../../../ggml/src/ggml-hexagon/htp/cpy-ops.c)):
+- Partition total linear elements N = ne0 * ne1 * ne2 * ne3 in 128-byte cache line chunks (`elems_per_line = (elem_size == 4) ? 32 : 64`).
+- Requires strict 1D contiguity:
+  [`htp_tensor_is_contiguous(dst, elem_size)`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L28)
+  and 128-byte aligned destination pointer
+  [`htp_tensor_mdev_data_aligned(dst)`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L47).
+- If contiguous and aligned, pass `elems_per_line` to
+  [`htp_tensor_mdev_partition`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L94);
+  otherwise pass 0 to trigger Device 0 fallback.
+
+### Single-Device Fallback (Device 0)
+
+- Fallback to Device 0 (`mdev.idx == 0`) when partitioning would cause cache line tearing or when work cannot be evenly distributed.
+- Triggers:
+  1. Destination tensor cannot be safely partitioned (`rows_per_chunk == 0` or non-contiguous/unaligned buffer).
+  2. Total aligned chunks < `mdev_count`.
+- Device 0 processes the entire tensor `[0, total_units)`.
+- Devices 1 ... N-1 receive `count = 0` and return `HTP_STATUS_OK` immediately.
+
+### Flatten Outer Dimensions Globally
+
+- **Never partition solely on `ne01` (dimension 1).**
+- Partitioning only on `ne01` repeats the device boundary across every 2D slice (`ne02`, `ne03`). If each 2D slice is small,
+  false sharing occurs repeatedly throughout the tensor.
+- Always flatten outer dimensions globally: `total_rows = ne01 * ne02 * ne03` and partition once across the combined row space.
+
+### Stateless Starting Coordinates
+
+- Do not use incremental state variables across slices that assume the thread or device starts at index 0.
+- Precompute starting multidimensional coordinates at `r = row_start` (or `e = elem_start`) once using `fastdiv`.
+- In inner loops, step base pointers directly (`ptr += stride`) or reset/wrap coordinates explicitly (`if (++i01 == ne01) { ... }`).
+
+### Clean Range Encapsulation
+
+- Initialize single-device default ranges at declaration:
+
+  ```c
+  uint32_t row_start = 0;
+  uint32_t nrows     = total_rows;
+  ```
+
+- Encapsulate all multi-device logic inside `if (octx->ctx->mdev.count > 1)`. If the block is omitted or compiled out,
+  the operator runs standard single-device execution untouched.
+- Do not propagate `mdev_` prefixes to worker functions or context structs. Worker threads are device-agnostic and
+  should only receive standard range parameters (`ctx.row_start`, `ctx.nrows`).
+- In worker threads, calculate row intervals using standard arithmetic:
+
+  ```c
+  const uint32_t ir0 = ctx->row_start + dr * ith;
+  const uint32_t ir1 = MIN(ir0 + dr, ctx->row_start + ctx->nrows);
+  ```
+
+  In single-device mode (`row_start == 0`), this naturally simplifies to `dr * ith` and `MIN(ir0 + dr, ctx->nrows)` with zero overhead.
+
+## Multi-Device Synchronization
+
+Multi-device execution synchronizes worker sessions across devices using explicit barriers and tensor cache flushing.
+
+### Synchronization Fence Protocol
+
+Multi-device execution synchronizes worker sessions through atomic fence slots and barriers defined in
+[`htp-fence.h`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h):
+
 ```
+[NPU Session 0]                              [NPU Session 1]
+       |                                            |
+  (Input Prep)                                 (Input Prep)
+       |                                            |
+  Pre-Op Barrier ----------------------------- Pre-Op Barrier
+  (mdev_sync_fence)                            (mdev_sync_fence)
+       |                                            |
+  Kernel Execution                             Kernel Execution
+  (Output Slice 0)                             (Output Slice 1)
+       |                                            |
+  Tensor Cache Flush                           Tensor Cache Flush
+  (htp_tensor_flush_all)                       (htp_tensor_flush_all)
+       |                                            |
+  Post-Op/Batch Barrier ---------------------- Post-Op/Batch Barrier
+  (htp_mdev_group_barrier)                     (htp_mdev_group_barrier)
+       |                                            |
+  Return Response to Host                      Return Response to Host
+```
+
+### Atomic Fence Slots and Cache Invalidation
+
+- Fence synchronization operates on dedicated RPCMEM shared memory mapped across all participating sessions (`ctx->mdev.fence_base`).
+- Each device owns a dedicated 128-byte cache-line aligned fence slot:
+
+  ```c
+  atomic_uint * my_fence = htp_mdev_fence_slot(fence_base, mdev_idx);
+  ```
+
+- **Writing to fence ([`htp_fence_write`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h#L18))**:
+  Stores `seq` and `status`, issues a `syncht` thread synchronization barrier, and flushes/invalidates the line
+  using `Q6_dccleaninva_A(fence)`.
+- **Reading from peer fence ([`htp_fence_read`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h#L26))**:
+  Executes `Q6_dccleaninva_A(fence)` and `syncht` before reading atomic values to ensure fresh data from DDR.
+
+### Deterministic Monotonic Sequence Numbers
+
+- Barrier fences use monotonically increasing sequence numbers:
+
+  ```c
+  const uint32_t seq = ++ctx->mdev.fence_seq;
+  ```
+
+- Comparing sequence numbers with signed arithmetic `(int32_t)(peer_seq - seq) >= 0` prevents race conditions or
+  misaligned barrier arrivals across iterations.
+- If any peer reports an error status (`peer_status > HTP_STATUS_OK`), the barrier propagates the error and unblocks immediately.
+
+### Tensor Cache Flush and Pipeline Completion
+
+- In the kernel, ensure all pushed DMA operations have been popped in strict FIFO order to drain the queue.
+- Use [`htp_tensor_flush_all()`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h) to flush specific dirty tensors back to DDR:
+  - [`htp_tensor_flush_all()`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h) flushes only modified tensor address ranges,
+    ensuring peer devices and the host CPU observe consistent data in DDR.
+- Never signal completion before all DMA transfers are drained and dirty tensor flushes have completed.
+
diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
index 112e9bae6..ec7801388 100644
--- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp
+++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp
@@ -66,7 +66,6 @@ using u32vec  = std::vector<uint32_t>;

 #define GGML_HEXAGON_MAX_SESSIONS          16

-#define GGML_HEXAGON_FENCE_BUFFER_SIZE     8192
 #define GGML_HEXAGON_FENCE_SLOT_SIZE       128

 struct ggml_hexagon_device_config {
@@ -75,6 +74,8 @@ struct ggml_hexagon_device_config {
     int         domain_id    = 0;
     std::string domain_name;
     std::string name;
+
+    std::vector<ggml_hexagon_device_config> mdev_group;
 };

 static ggml_hexagon_device_config opt_device_configs[GGML_HEXAGON_MAX_SESSIONS];
@@ -350,27 +351,48 @@ struct ggml_hexagon_tensor_extra {
 };

 static inline bool ggml_hexagon_tensor_is_fuseable(const struct ggml_tensor * t) {
-    if (!t || !t->extra) return false;
+    if (!t->extra) return false;
     auto extra = (const struct ggml_hexagon_tensor_extra *) t->extra;
     return (extra->flags & GGML_HEXAGON_TENSOR_FUSEABLE) != 0;
 }

+static inline bool ggml_hexagon_tensors_overlap(const struct ggml_tensor * a, const struct ggml_tensor * b) {
+    const uintptr_t a0 = (uintptr_t) a->data;
+    const uintptr_t b0 = (uintptr_t) b->data;
+    const uintptr_t a1 = a0 + ggml_nbytes(a);
+    const uintptr_t b1 = b0 + ggml_nbytes(b);
+
+    return a0 < b1 && b0 < a1;
+}
+
 struct htp_opnode;

 struct ggml_hexagon_opbatch;
 struct ggml_hexagon_opqueue;
 struct ggml_hexagon_shared_buffer;
+struct ggml_hexagon_fence_buffer;
 struct ggml_hexagon_session;
+struct ggml_backend_hexagon_device_context;
+
+struct ggml_hexagon_mdev_group {
+    uint32_t idx   = 0;
+    uint32_t count = 1;
+    std::vector<std::unique_ptr<ggml_hexagon_session>> sessions;
+};

 struct ggml_backend_hexagon_comm_context {
     std::vector<ggml_backend_t> backends;
     size_t                      n_backends = 0;
-    uint32_t                    fence_seq  = 0;
+    volatile uint32_t *         fence_slots[GGML_HEXAGON_MAX_SESSIONS] = {};
+    ggml_tensor                 fence_tensors[GGML_HEXAGON_MAX_SESSIONS] = {};
 };

 struct ggml_hexagon_event {
-    ggml_hexagon_session * sess = nullptr;
-    uint64_t               seq  = 0;
+    ggml_hexagon_session * sess         = nullptr;
+    ggml_hexagon_session * fence_sess   = nullptr;
+    volatile uint32_t *    fence_slot   = nullptr;
+    ggml_tensor            fence_tensor = {};
+    uint32_t               seq          = 0;
 };

 struct ggml_hexagon_session {
@@ -387,12 +409,12 @@ struct ggml_hexagon_session {
     bool             valid_queue;
     bool             valid_iface;

-    std::atomic<int>      op_pending;
     ggml_hexagon_opbatch* op_batch;
     ggml_hexagon_opqueue* op_queue;

     std::unordered_map<int, std::unique_ptr<ggml_hexagon_shared_buffer>> cloned_buffers;
-    std::unordered_set<ggml_hexagon_session *>                           sync_peers;
+    std::unordered_set<ggml_hexagon_session *>                           virt_peers;
+    std::unordered_set<ggml_hexagon_session *>                           phys_peers;

     uint32_t n_threads   = 0;
     uint32_t n_hvx       = 0;
@@ -400,14 +422,23 @@ struct ggml_hexagon_session {
     uint64_t vtcm_size   = 0;
     size_t   max_vmem    = 0;
     size_t   max_bufsize = 0;
-    uint32_t fence_seq;
+    uint32_t fence_seq   = 0;
+
+    std::atomic<uint64_t> batch_req_seq{0};
+    std::atomic<uint64_t> batch_rsp_seq{0};
+    std::atomic<uint32_t> last_error{HTP_STATUS_OK};

     uint64_t                cached_uid = 0;
     std::vector<htp_opnode> cached_nodes;

     mutable std::unordered_set<const ggml_tensor *> needs_repack;

-    ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev = nullptr) noexcept(false);
+    ggml_hexagon_mdev_group                mdev;
+    ggml_backend_dev_t                     dev       = nullptr;
+    ggml_backend_hexagon_device_context *  dev_ctx   = nullptr;
+    ggml_hexagon_fence_buffer *            fence_buf = nullptr;
+
+    ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev = nullptr, uint32_t mdev_idx = 0, uint32_t mdev_count = 0) noexcept(false);
     ~ggml_hexagon_session() noexcept(true);

     const char* c_name() const { return name.c_str(); }
@@ -415,31 +446,36 @@ struct ggml_hexagon_session {
     void allocate(const ggml_hexagon_device_config & config) noexcept(false);
     void release() noexcept(true);

+    uint8_t * alloc_fence(uint32_t n_slots = 1);
+    void      free_fence(void * ptr, uint32_t n_slots = 1);
+
+    uint8_t *                                         mdev_fence_slot = nullptr;
+    std::unordered_map<uint64_t, volatile uint32_t *> cpy_fence_slots;
+
+    void enqueue_mdev_group();
     void enqueue_op(const htp_opnode & node);
     void enqueue_cpy(const ggml_tensor * src, ggml_tensor * dst, const ggml_tensor * sync_tensor = nullptr, uint32_t fence_seq = 0);
-    void enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq = 0);
-    void enqueue_allreduce(const ggml_tensor * dst, const std::vector<const ggml_tensor *> & src_tensors, const std::vector<const ggml_tensor *> & sync_tensors, uint32_t rank, uint32_t n_ranks, uint32_t fence_seq_entry = 0, uint32_t fence_seq_exit = 0);
+    void enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq = 0, bool wait = true);
+    void enqueue_allreduce(const ggml_tensor * dst, const std::vector<const ggml_tensor *> & src_tensors,
+                           const std::vector<const ggml_tensor *> & sync_tensors, uint32_t rank, uint32_t n_ranks,
+                           uint32_t fence_seq_entry = 0, uint32_t fence_seq_exit = 0);

-    void flush(bool all = true);
-    void flush_pending(bool all = false);
+    void flush_sync(bool all = true);
+    void flush_async();
     void flush_batch(size_t min_ops = 1);
-
-    uint64_t record_event();
-    void     wait_event(uint64_t seq);
+    void flush_peers();
+    void flush_pending(bool all = true);

     bool clone_buffer(const ggml_hexagon_shared_buffer*);
+    void release_buffer(const ggml_hexagon_shared_buffer*);
+    void unclone_buffer(const ggml_hexagon_shared_buffer*);

-    void add_sync_peer(ggml_hexagon_session * peer) {
-        sync_peers.insert(peer);
-    }
-
-    void flush_sync_peers() {
-        if (sync_peers.empty()) return;
-
-        for (auto * peer : sync_peers) {
-            peer->flush_batch();
+    void add_peer(ggml_hexagon_session * peer) {
+        if (this->phys_idx == peer->phys_idx) {
+            virt_peers.insert(peer);
+        } else {
+            phys_peers.insert(peer);
         }
-        sync_peers.clear();
     }
 };

@@ -451,8 +487,9 @@ struct ggml_backend_hexagon_device_context {
     ggml_backend_dev_t         dev = nullptr;
     size_t                     max_bufsize = 0;

-    ggml_backend_buffer_type buffer_type      = {};
-    ggml_backend_buffer_type host_buffer_type = {};
+    ggml_backend_buffer_type buffer_type       = {};
+    ggml_backend_buffer_type host_buffer_type  = {};
+    ggml_backend_buffer_type fence_buffer_type = {};

     std::unique_ptr<ggml_hexagon_session> sess;

@@ -484,6 +521,8 @@ struct ggml_hexagon_rpcmem_block {
     int       fd   = -1;
     size_t    size = 0;

+    std::unordered_set<ggml_hexagon_session *> mapped_clones;
+
     ggml_hexagon_rpcmem_block(size_t size) {
         base = (uint8_t *) rpcmem_alloc2(RPCMEM_HEAP_ID_SYSTEM, RPCMEM_DEFAULT_FLAGS, size);
         if (!base) {
@@ -508,8 +547,6 @@ struct ggml_hexagon_shared_buffer {
     ggml_hexagon_session *                     sess;
     std::shared_ptr<ggml_hexagon_rpcmem_block> mem;
     std::vector<ggml_hexagon_tensor_extra *>   tensor_extra;
-    uint32_t fence_head = 0;
-    size_t   fences_size = 0;
     bool     mapped;
     bool     pinned;

@@ -518,16 +555,6 @@ struct ggml_hexagon_shared_buffer {
     size_t       size()   const { return mem ? mem->size : 0;  }
     int          fd()     const { return mem ? mem->fd   : -1; }

-    uint8_t * alloc_fence() {
-        if (fences_size == 0) return nullptr;
-        int max_slots = fences_size / GGML_HEXAGON_FENCE_SLOT_SIZE;
-        uint32_t slot = (fence_head++) % max_slots;
-
-        size_t guard_offset = size() - fences_size;
-        uint8_t * fence_ptr = base() + guard_offset + (size_t)slot * GGML_HEXAGON_FENCE_SLOT_SIZE;
-        return fence_ptr;
-    }
-
     void mmap() {
         if (!this->mem) return;
         fastrpc_map_flags flags = this->pinned ? FASTRPC_MAP_FD : FASTRPC_MAP_FD_DELAYED;
@@ -581,29 +608,24 @@ struct ggml_hexagon_shared_buffer {
         this->mem  = nullptr;
     }

-    ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, size_t size, bool pinned = false, size_t fence_size = 0) {
-        this->sess        = sess;
-        this->mapped      = false;
-        this->pinned      = pinned;
-        this->fences_size = fence_size;
+    ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, size_t size, bool pinned = false) {
+        this->sess   = sess;
+        this->mapped = false;
+        this->pinned = pinned;

-        // Size adjustment inside the buffer class
+        // Size adjustment inside the buffer class: 4K aligned data size + 4K guard page
         size_t guard_offset = (size + 4095) & ~4095;
-        size_t total_size = guard_offset;
-        if (fence_size > 0) {
-            total_size += 4096 + fence_size;
-        }
+        size_t total_size   = guard_offset + 4096;

         alloc(total_size);
     }

     // Clone constructor for cross-session mapping
     ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, const ggml_hexagon_shared_buffer & other) {
-        this->sess        = sess;
-        this->mem         = other.mem;
-        this->mapped      = false;
-        this->pinned      = other.pinned;
-        this->fences_size = other.fences_size;
+        this->sess   = sess;
+        this->mem    = other.mem;
+        this->mapped = false;
+        this->pinned = other.pinned;
     }

     ~ggml_hexagon_shared_buffer() {
@@ -614,6 +636,59 @@ struct ggml_hexagon_shared_buffer {
     }
 };

+struct ggml_hexagon_fence_buffer : public ggml_hexagon_shared_buffer {
+    uint32_t              slot_count = 0;
+    uint32_t              slot_head  = 0;
+    std::vector<uint32_t> free_slots;
+    ggml_backend_buffer   backend_buffer{};
+
+    ggml_hexagon_fence_buffer(ggml_hexagon_session * sess, ggml_backend_buffer_type_t buft, size_t size)
+        : ggml_hexagon_shared_buffer(sess, size, false /* pinned */),
+          slot_count(size / GGML_HEXAGON_FENCE_SLOT_SIZE),
+          slot_head(0) {
+        backend_buffer.buft    = buft;
+        backend_buffer.context = static_cast<ggml_hexagon_shared_buffer *>(this);
+        backend_buffer.size    = size;
+    }
+
+    uint8_t * alloc_slot(uint32_t n_slots = 1) {
+        uint8_t * ptr = nullptr;
+        if (n_slots == 1 && !free_slots.empty()) {
+            uint32_t slot = free_slots.back();
+            free_slots.pop_back();
+            ptr = base() + (size_t) slot * GGML_HEXAGON_FENCE_SLOT_SIZE;
+        } else if (slot_head + n_slots <= slot_count) {
+            uint32_t slot = slot_head;
+            slot_head += n_slots;
+            ptr = base() + (size_t) slot * GGML_HEXAGON_FENCE_SLOT_SIZE;
+        }
+        if (ptr) {
+            memset(ptr, 0, (size_t) n_slots * GGML_HEXAGON_FENCE_SLOT_SIZE);
+        }
+        return ptr;
+    }
+
+    void free_slot(void * ptr, uint32_t n_slots = 1) {
+        if (!ptr) return;
+        uint32_t slot = ((uint8_t *) ptr - base()) / GGML_HEXAGON_FENCE_SLOT_SIZE;
+        for (uint32_t i = 0; i < n_slots; i++) {
+            free_slots.push_back(slot + i);
+        }
+    }
+};
+
+inline uint8_t * ggml_hexagon_session::alloc_fence(uint32_t n_slots) {
+    uint8_t * ptr = fence_buf->alloc_slot(n_slots);
+    GGML_ASSERT(ptr);
+    return ptr;
+}
+
+inline void ggml_hexagon_session::free_fence(void * ptr, uint32_t n_slots) {
+    if (fence_buf) {
+        fence_buf->free_slot(ptr, n_slots);
+    }
+}
+
 static ggml_hexagon_session * ggml_backend_hexagon_buffer_get_sess(ggml_backend_buffer_t buffer) {
     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(buffer->context);
     return sbuf->sess;
@@ -621,6 +696,7 @@ static ggml_hexagon_session * ggml_backend_hexagon_buffer_get_sess(ggml_backend_

 static void ggml_backend_hexagon_buffer_free_buffer(ggml_backend_buffer_t buffer) {
     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(buffer->context);
+    sbuf->sess->unclone_buffer(sbuf);
     delete sbuf;
 }

@@ -1537,7 +1613,7 @@ static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
     auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
     auto sess    = dev_ctx->session();
     try {
-        ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE);
+        ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false);
         return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_buffer_interface, sbuf, size);
     } catch (const std::exception & exc) {
         GGML_LOG_ERROR("ggml-hex: %s failed to allocate device buffer context: %s\n", dev_ctx->c_name(), exc.what());
@@ -1550,7 +1626,7 @@ static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer(
     auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
     auto sess    = dev_ctx->session();
     try {
-        ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE);
+        ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false);
         return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_host_buffer_interface, sbuf, size);
     } catch (const std::exception & exc) {
         GGML_LOG_ERROR("ggml-hex: %s failed to allocate host buffer context: %s\n", dev_ctx->c_name(), exc.what());
@@ -1618,11 +1694,16 @@ ggml_backend_hexagon_device_context::ggml_backend_hexagon_device_context(int dev
     host_buffer_type.device  = dev;
     host_buffer_type.iface   = ggml_backend_hexagon_host_buffer_type_interface;
     host_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name + "-HOST", this);
+
+    fence_buffer_type.device  = dev;
+    fence_buffer_type.iface   = ggml_backend_hexagon_buffer_type_interface;
+    fence_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name + "-FENCE", this);
 }

 ggml_backend_hexagon_device_context::~ggml_backend_hexagon_device_context() {
     delete static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type.context);
     delete static_cast<ggml_backend_hexagon_buffer_type_context *>(host_buffer_type.context);
+    delete static_cast<ggml_backend_hexagon_buffer_type_context *>(fence_buffer_type.context);
 }

 static bool ggml_backend_buffer_is_hexagon(const struct ggml_backend_buffer * b) {
@@ -1698,8 +1779,8 @@ struct ggml_hexagon_opbatch {
         if (it != b_map.end()) { return it->second; }

         // Add new buffer to the batch
-        int bi = n_bufs++;
         GGML_ASSERT(n_bufs < HTP_OP_MAX_BUFS);
+        int bi = n_bufs++;

         b_map.insert({sbuf->fd(), bi});

@@ -1902,6 +1983,12 @@ struct ggml_hexagon_opbatch {
         }
     }

+    void update_mdev_group(uint32_t mdev_idx) {
+        if (n_ops > 0 && h_ops[0].opcode == HTP_OP_MDEV_GROUP) {
+            h_ops[0].params[0] = (int32_t) mdev_idx;
+        }
+    }
+
     bool try_fuse_allreduce_add(const htp_opnode & node) {
         if (n_ops == 0 || opt_ar_select != 2) return false;
         if (node.opcode != HTP_OP_ADD) return false;
@@ -1910,15 +1997,16 @@ struct ggml_hexagon_opbatch {
         if (last_node.opcode != HTP_OP_ALLREDUCE) return false;

         auto * ar_kparams = (struct htp_allreduce_kernel_params *) last_node.kernel_params;
-        const uint32_t rank = (uint32_t) ar_kparams->rank;
-        const ggml_tensor * ar_local = (rank < last_node.inputs.size()) ? last_node.inputs[rank] : nullptr;
+        const uint32_t rank    = (uint32_t) ar_kparams->rank;
+        const uint32_t n_ranks = (uint32_t) ar_kparams->n_ranks;
+        const ggml_tensor * ar_local = last_node.inputs[rank];
         const ggml_tensor * add_src0 = node.src0();
         const ggml_tensor * add_src1 = node.src1();
+        const ggml_tensor * add_dst  = node.dst();

-        if (!add_src0 || !add_src1 || !ar_local) return false;
         if (!ggml_hexagon_tensor_is_fuseable(ar_local)) return false;

-        const ggml_tensor * res_tensor = nullptr;
+        const ggml_tensor * res_tensor;
         if (add_src0 == ar_local || add_src0->data == ar_local->data) {
             res_tensor = add_src1;
         } else if (add_src1 == ar_local || add_src1->data == ar_local->data) {
@@ -1927,14 +2015,12 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        if (!res_tensor || !res_tensor->data) return false;
-
         if (ar_local->type != res_tensor->type) return false;

         const bool is_same_shape = (ar_local->ne[0] == res_tensor->ne[0] && ar_local->ne[1] == res_tensor->ne[1] &&
                                     ar_local->ne[2] == res_tensor->ne[2] && ar_local->ne[3] == res_tensor->ne[3]);
-        const bool is_row_bcast  = (ar_local->ne[0] == res_tensor->ne[0] &&
-                                    res_tensor->ne[1] == 1 && res_tensor->ne[2] == 1 && res_tensor->ne[3] == 1);
+        const bool is_row_bcast  = !is_same_shape && (ar_local->ne[0] == res_tensor->ne[0] && res_tensor->ne[1] == 1 &&
+                                                      res_tensor->ne[2] == 1 && res_tensor->ne[3] == 1);

         if (!is_same_shape && !is_row_bcast) return false;

@@ -1947,13 +2033,21 @@ struct ggml_hexagon_opbatch {
                 return false;
             }
         }
-        if (ggml_is_contiguous(ar_local) != ggml_is_contiguous(node.dst())) {
+        if (ggml_is_contiguous(ar_local) != ggml_is_contiguous(add_dst)) {
             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;
+            }
+        }
+
         struct htp_allreduce_kernel_params new_kparams;
         if (!ggml_hexagon_precompute_allreduce_params(
-            sess, node.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, &new_kparams
         )) {
             HEX_VERBOSE("ggml-hex: %s skip ALLREDUCE_ADD fusion: solver failed\n", sess->c_name());
             return false;
@@ -1961,7 +2055,6 @@ struct ggml_hexagon_opbatch {

         size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
         auto fit_t = [&](const ggml_tensor * t) {
-            if (!t) return;
             if (!t_map.count(t)) {
                 extra_tens++;
                 auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -1972,7 +2065,7 @@ struct ggml_hexagon_opbatch {
             }
         };
         fit_t(res_tensor);
-        fit_t(node.dst());
+        fit_t(add_dst);
         if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) {
             return false;
         }
@@ -1981,7 +2074,7 @@ struct ggml_hexagon_opbatch {
         last_node.name   = "ALLREDUCE+ADD";
         last_node.inputs.push_back(res_tensor);
         last_node.outputs.clear();
-        last_node.outputs.push_back(node.dst());
+        last_node.outputs.push_back(add_dst);
         last_node.fused.push_back(node.node);
         memcpy(last_node.kernel_params, &new_kparams, sizeof(new_kparams));

@@ -1989,9 +2082,8 @@ struct ggml_hexagon_opbatch {
         o.opcode = HTP_OP_ALLREDUCE_ADD;
         memcpy(o.kernel_params, &new_kparams, sizeof(new_kparams));

-        const uint32_t n_ranks = (uint32_t) ar_kparams->n_ranks;
         o.src[2 * n_ranks] = add_tensor(res_tensor);
-        o.dst[0]           = add_tensor(node.dst());
+        o.dst[0]           = add_tensor(add_dst);
         for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) {
             o.dst[d] = 0xffff;
         }
@@ -2011,10 +2103,9 @@ struct ggml_hexagon_opbatch {
         const ggml_tensor * mul_src1 = node.src1();
         const ggml_tensor * rms_out  = last_node.dst();

-        if (!mul_src0 || !mul_src1 || !rms_out) return false;
         if (!ggml_hexagon_tensor_is_fuseable(rms_out)) return false;

-        const ggml_tensor * weight = nullptr;
+        const ggml_tensor * weight;
         if (mul_src0 == rms_out || mul_src0->data == rms_out->data) {
             weight = mul_src1;
         } else if (mul_src1 == rms_out || mul_src1->data == rms_out->data) {
@@ -2023,10 +2114,7 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        if (!weight || !weight->data) return false;
-
         const ggml_tensor * src0 = last_node.src0();
-        if (!src0 || !src0->data) return false;

         if (src0->ne[0] != weight->ne[0] || src0->ne[0] != node.dst()->ne[0]) {
             return false;
@@ -2057,7 +2145,6 @@ struct ggml_hexagon_opbatch {

         size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
         auto fit_t = [&](const ggml_tensor * t) {
-            if (!t) return;
             if (!t_map.count(t)) {
                 extra_tens++;
                 auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2112,10 +2199,9 @@ struct ggml_hexagon_opbatch {
         const ggml_tensor * add_src1 = node.src1();
         const ggml_tensor * mm_out   = last_node.dst();

-        if (!add_src0 || !add_src1 || !mm_out) return false;
         if (!ggml_hexagon_tensor_is_fuseable(mm_out)) return false;

-        const ggml_tensor * src2 = nullptr;
+        const ggml_tensor * src2;
         if (add_src0 == mm_out || add_src0->data == mm_out->data) {
             src2 = add_src1;
         } else if (add_src1 == mm_out || add_src1->data == mm_out->data) {
@@ -2124,11 +2210,8 @@ struct ggml_hexagon_opbatch {
             return false;
         }

-        if (!src2 || !src2->data) return false;
-
         const ggml_tensor * src0 = last_node.src0();
         const ggml_tensor * src1 = last_node.src1();
-        if (!src0 || !src1) return false;

         struct htp_mm_kernel_params kparams;
         ggml_hexagon_precompute_fused_matmul_add_params(sess, src0, src1, src2, node.dst(), &kparams);
@@ -2144,7 +2227,6 @@ struct ggml_hexagon_opbatch {

         size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
         auto fit_t = [&](const ggml_tensor * t) {
-            if (!t) return;
             if (!t_map.count(t)) {
                 extra_tens++;
                 auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2197,7 +2279,6 @@ struct ggml_hexagon_opbatch {
         const ggml_tensor * w_in = node.src0();
         const ggml_tensor * x_in = node.src1();
         const ggml_tensor * d_in = node.dst();
-        if (!w_in || !x_in || !d_in) return false;

         htp_opnode & last_node = ops[n_ops - 1];

@@ -2231,7 +2312,6 @@ struct ggml_hexagon_opbatch {

             size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
             auto fit_t = [&](const ggml_tensor * t) {
-                if (!t) return;
                 if (!t_map.count(t)) {
                     extra_tens++;
                     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2282,7 +2362,6 @@ struct ggml_hexagon_opbatch {
             const ggml_tensor * w0 = last_node.src0();
             const ggml_tensor * x  = last_node.src1();
             const ggml_tensor * w1 = node.src0();
-            if (!w0 || !x || !w1) return false;

             struct htp_mm_kernel_params kparams;
             ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams);
@@ -2297,7 +2376,6 @@ struct ggml_hexagon_opbatch {

             size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
             auto fit_t = [&](const ggml_tensor * t) {
-                if (!t) return;
                 if (!t_map.count(t)) {
                     extra_tens++;
                     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2359,7 +2437,6 @@ struct ggml_hexagon_opbatch {
         const ggml_tensor * x_in   = node.src1();
         const ggml_tensor * ids_in = node.node->src[2];
         const ggml_tensor * d_in   = node.dst();
-        if (!w_in || !x_in || !ids_in || !d_in) return false;

         htp_opnode & last_node = ops[n_ops - 1];

@@ -2394,7 +2471,6 @@ struct ggml_hexagon_opbatch {

             size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
             auto fit_t = [&](const ggml_tensor * t) {
-                if (!t) return;
                 if (!t_map.count(t)) {
                     extra_tens++;
                     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2447,7 +2523,6 @@ struct ggml_hexagon_opbatch {
             const ggml_tensor * x   = last_node.src1();
             const ggml_tensor * ids = last_node.node->src[2];
             const ggml_tensor * w1  = node.src0();
-            if (!w0 || !x || !ids || !w1) return false;

             struct htp_mm_kernel_params kparams;
             ggml_hexagon_precompute_fused_mmidnx_params(sess, w0, x, node.dst(), 2, &kparams);
@@ -2462,7 +2537,6 @@ struct ggml_hexagon_opbatch {

             size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0;
             auto fit_t = [&](const ggml_tensor * t) {
-                if (!t) return;
                 if (!t_map.count(t)) {
                     extra_tens++;
                     auto sbuf = static_cast<ggml_hexagon_shared_buffer *>(t->buffer->context);
@@ -2540,17 +2614,14 @@ struct ggml_hexagon_opqueue {
     // Shared buffer for storing batches
     ggml_hexagon_shared_buffer *shm_buf;
     size_t                      shm_blk_size;
-
-    uint64_t req_seq = 0;
-    uint64_t rsp_seq = 0;
+    size_t                      depth;

     using opvec = std::vector<htp_opnode>;

-    std::queue<unsigned int>    done;           // completed batch ids
     std::vector<opvec>          op_cache;       // per batch op cache
     std::vector<uint64_t>       start_usec;     // per batch start time

-    ggml_hexagon_opqueue(ggml_hexagon_session *sess, size_t batch_size, size_t depth) {
+    ggml_hexagon_opqueue(ggml_hexagon_session *sess, size_t batch_size, size_t depth) : depth(depth) {
         size_t n_bufs    = HTP_OP_MAX_BUFS;
         size_t n_ops     = batch_size;
         size_t n_tensors = n_ops * HTP_OP_MAX_OUTPUTS + n_ops * HTP_OP_MAX_INPUTS;
@@ -2571,9 +2642,6 @@ struct ggml_hexagon_opqueue {
         op_cache.resize(depth);
         start_usec.resize(depth, 0);

-        // init done queue
-        for (unsigned int i = 0; i < depth; i++) { done.push(i); }
-
         if (opt_verbose) {
             GGML_LOG_INFO("ggml-hex: %s allocated opqueue : batch-size %zu depth %zu shm-size %zu shm-block-size %zu\n",
                     sess->c_name(), batch_size, depth, shm_buf->size(), shm_blk_size);
@@ -2587,7 +2655,7 @@ struct ggml_hexagon_opqueue {
     size_t shm_size() const { return shm_buf ? shm_buf->size() : 0; }

     // push new batch
-    bool push(htp_opbatch_req& req, dspqueue_buffer& dbuf, ggml_hexagon_opbatch* op_batch) {
+    bool push(htp_opbatch_req& req, dspqueue_buffer& dbuf, const ggml_hexagon_opbatch* op_batch, uint64_t seq) {
         static_assert(sizeof(htp_opbatch_req) % 8 == 0, "sizeof(htp_opbatch_req) must be multiple of 8");
         static_assert(sizeof(htp_opbatch_rsp) % 8 == 0, "sizeof(htp_opbatch_rsp) must be multiple of 8");
         static_assert(sizeof(htp_buf_desc)    % 8 == 0, "sizeof(htp_buf_desc) must be multiple of 8");
@@ -2595,16 +2663,17 @@ struct ggml_hexagon_opqueue {
         static_assert(sizeof(htp_op_desc)     % 8 == 0, "sizeof(htp_op_desc) must be multiple of 8");
         static_assert(sizeof(htp_prof_desc)   % 8 == 0, "sizeof(htp_prof_desc) must be multiple of 8");

-        if (done.empty()) { return false; }
+        if (seq - shm_buf->sess->batch_rsp_seq > depth) { return false; }

-        req.id        = done.front(); done.pop(); // batch id
+        const uint32_t slot = (uint32_t) ((seq - 1) % depth);
+
+        req.seq       = seq;
         req.n_bufs    = op_batch->n_bufs;
         req.n_tensors = op_batch->n_tens;
         req.n_ops     = op_batch->n_ops;
-        req.seq       = ++req_seq;

-        op_cache[req.id]   = std::move(op_batch->ops);
-        start_usec[req.id] = ggml_time_us();
+        op_cache[slot]   = op_batch->ops;
+        start_usec[slot] = ggml_time_us();

         const size_t b_size = sizeof(htp_buf_desc)  * req.n_bufs;
         const size_t t_size = sizeof(htp_tensor)    * req.n_tensors;
@@ -2619,7 +2688,7 @@ struct ggml_hexagon_opqueue {
             req.n_traces = 0;
         }

-        dbuf.ptr      = shm_buf->base() + (req.id * shm_blk_size);
+        dbuf.ptr      = shm_buf->base() + ((size_t) slot * shm_blk_size);
         dbuf.fd       = shm_buf->fd();
         dbuf.flags    = DSPQUEUE_BUFFER_FLAG_FLUSH_SENDER | DSPQUEUE_BUFFER_FLAG_INVALIDATE_RECIPIENT;
         dbuf.offset   = (uint8_t*) dbuf.ptr - (uint8_t*) shm_buf->base();
@@ -2632,18 +2701,14 @@ struct ggml_hexagon_opqueue {
         uint8_t * t_ptr = m_ptr; m_ptr += t_size;
         uint8_t * o_ptr = m_ptr;

-        op_batch->sort_buffers();
-
         memcpy(b_ptr, (void *) op_batch->h_bufs.data(), b_size);
         memcpy(t_ptr, (void *) op_batch->h_tens.data(), t_size);
         memcpy(o_ptr, (void *) op_batch->h_ops.data(),  o_size);

-        HEX_VERBOSE("ggml-hex: %s opqueue-push batch #%u : n-bufs %u n-tensors %u n-ops %u vmem %zu : b-size %zu t-size %zu o-size %zu m-size %zu\n",
-                shm_buf->sess->c_name(), req.id, req.n_bufs, req.n_tensors, req.n_ops, op_batch->b_vmem,
+        HEX_VERBOSE("ggml-hex: %s opqueue-push batch #%llu : n-bufs %u n-tensors %u n-ops %u vmem %zu : b-size %zu t-size %zu o-size %zu m-size %zu\n",
+                shm_buf->sess->c_name(), (unsigned long long) req.seq, req.n_bufs, req.n_tensors, req.n_ops, op_batch->b_vmem,
                 b_size, t_size, o_size, (size_t) dbuf.size);

-        op_batch->reset();
-
         if (opt_verbose > 1) {
             htp_buf_desc *b = (htp_buf_desc*) b_ptr;
             for (unsigned int i=0; i < req.n_bufs; i++) {
@@ -2662,9 +2727,7 @@ struct ggml_hexagon_opqueue {
     }

     void pop(htp_opbatch_rsp rsp, dspqueue_buffer dbuf) {
-        GGML_ASSERT(rsp.id < op_cache.size());
-
-        done.push(rsp.id);
+        const uint32_t slot = (uint32_t) ((rsp.seq - 1) % depth);

         const size_t b_size = sizeof(htp_buf_desc)  * rsp.n_bufs;
         const size_t t_size = sizeof(htp_tensor)    * rsp.n_tensors;
@@ -2681,15 +2744,15 @@ struct ggml_hexagon_opqueue {
         const size_t m_size = b_size + t_size + o_size + p_size + tr_size;
         GGML_ASSERT(m_size <= shm_blk_size);

-        HEX_VERBOSE("ggml-hex: %s opqueue-pop batch #%u : n-bufs %u n-tensors %u n-ops %u : m-size %zu b-size %zu t-size %zu o-size %zu\n",
-                shm_buf->sess->c_name(), rsp.id, rsp.n_bufs, rsp.n_tensors, rsp.n_ops,
+        HEX_VERBOSE("ggml-hex: %s opqueue-pop batch #%llu : n-bufs %u n-tensors %u n-ops %u : m-size %zu b-size %zu t-size %zu o-size %zu\n",
+                shm_buf->sess->c_name(), (unsigned long long) rsp.seq, rsp.n_bufs, rsp.n_tensors, rsp.n_ops,
                 (size_t) dbuf.size, b_size, t_size, o_size);

         uint8_t * m_ptr = (uint8_t*) dbuf.ptr;
         uint8_t * p_ptr = m_ptr + (b_size + t_size + o_size);

         if (rsp.n_ops > 0) {
-            auto & ops = op_cache[rsp.id];
+            auto & ops = op_cache[slot];
             GGML_ASSERT(rsp.n_ops <= ops.size());

             const htp_prof_desc * pd = (const htp_prof_desc *) p_ptr;
@@ -2712,16 +2775,41 @@ struct ggml_hexagon_opqueue {
                 ggml_hexagon_dump_trace_events(shm_buf->sess->name, rsp, trace_events, n_traces);
             }
         }
-
-        if (rsp.seq > rsp_seq) {
-            rsp_seq = rsp.seq;
-        }
     }
 };

-// Flush HTP response queue i.e wait for all outstanding requests to complete
+void ggml_hexagon_session::flush_peers() {
+    auto vpeers = std::move(virt_peers);
+    virt_peers.clear();
+    for (auto * peer : vpeers) {
+        peer->flush_sync();
+    }
+
+    auto ppeers = std::move(phys_peers);
+    phys_peers.clear();
+    for (auto * peer : ppeers) {
+        peer->flush_async();
+    }
+
+    for (auto & sub : this->mdev.sessions) {
+        sub->flush_peers();
+    }
+}
+
+void ggml_hexagon_session::flush_async() {
+    flush_peers();
+    flush_batch();
+}
+
 void ggml_hexagon_session::flush_pending(bool all) {
-    while (this->op_pending) {
+    for (auto & sub : this->mdev.sessions) {
+        sub->flush_pending(all);
+        if (sub->last_error > HTP_STATUS_OK) {
+            this->last_error = sub->last_error.load();
+        }
+    }
+
+    while (this->batch_rsp_seq < this->batch_req_seq) {
         struct htp_opbatch_rsp rsp;
         uint32_t               rsp_size;
         uint32_t               flags;
@@ -2746,32 +2834,64 @@ void ggml_hexagon_session::flush_pending(bool all) {
             GGML_ABORT("ggml-hex: %s dspcall : bad response : size %u dspbufs %u\n", this->c_name(), rsp_size, n_dbufs);
         }

-        if (rsp.status != HTP_STATUS_OK) {
-            GGML_LOG_ERROR("ggml-hex: %s dspcall : dsp-rsp: %s\n", this->c_name(), status_to_str(rsp.status));
-            // TODO: handle errors
+        if (rsp.status > HTP_STATUS_OK) {
+            GGML_LOG_ERROR("ggml-hex: %s dspcall : dsp-rsp %s\n", this->c_name(), status_to_str(rsp.status));
+            this->last_error = rsp.status;
+            for (auto & sub : this->mdev.sessions) {
+                sub->last_error = rsp.status;
+            }
         }

         op_queue->pop(rsp, dbuf);

-        this->op_pending--;  // atomic dec
+        GGML_ASSERT(rsp.seq == this->batch_rsp_seq + 1);
+        this->batch_rsp_seq = rsp.seq;

         if (!all) break;
     }
 }

+void ggml_hexagon_session::flush_sync(bool all) {
+    flush_async();
+    flush_pending(all);
+}
+
 void ggml_hexagon_session::flush_batch(size_t min_ops) {
     if (op_batch->n_ops < min_ops) { return; }

+    op_batch->sort_buffers();
+
     htp_opbatch_req req {};
     dspqueue_buffer dbuf{};

-    if (!op_queue->push(req, dbuf, op_batch)) {
+    const uint64_t seq = ++this->batch_req_seq;
+
+    op_batch->update_mdev_group(this->mdev.idx);
+
+    if (!op_queue->push(req, dbuf, op_batch, seq)) {
         flush_pending(false);
-        op_queue->push(req, dbuf, op_batch);
+        op_queue->push(req, dbuf, op_batch, seq);
     }

-    // Bump pending flag (cleared in the session::flush once we get the response)
-    this->op_pending++;  // atomic inc
+    for (auto & sub : this->mdev.sessions) {
+        htp_opbatch_req sub_req {};
+        dspqueue_buffer sub_dbuf{};
+
+        sub->batch_req_seq = seq;
+        op_batch->update_mdev_group(sub->mdev.idx);
+
+        if (!sub->op_queue->push(sub_req, sub_dbuf, op_batch, seq)) {
+            sub->flush_pending(false);
+            sub->op_queue->push(sub_req, sub_dbuf, op_batch, seq);
+        }
+
+        HEX_VERBOSE("ggml-hex: %s queue-opbatch: %p size %u\n", sub->c_name(), sub_dbuf.ptr, sub_dbuf.size);
+
+        int err = dspqueue_write(sub->queue, 0, 1, &sub_dbuf, sizeof(sub_req), (const uint8_t*) &sub_req, DSPQUEUE_TIMEOUT);
+        if (err != 0) {
+            GGML_ABORT("ggml-hex: %s dspqueue_write failed: 0x%08x\n", sub->c_name(), (unsigned) err);
+        }
+    }

     HEX_VERBOSE("ggml-hex: %s queue-opbatch: %p size %u\n", this->c_name(), dbuf.ptr, dbuf.size);

@@ -2779,28 +2899,28 @@ void ggml_hexagon_session::flush_batch(size_t min_ops) {
     if (err != 0) {
         GGML_ABORT("ggml-hex: %s dspqueue_write failed: 0x%08x\n", this->c_name(), (unsigned) err);
     }
-}

-void ggml_hexagon_session::flush(bool all) {
-    flush_sync_peers();
-    flush_batch();
-    flush_pending(all);
+    op_batch->reset();
 }

 void ggml_hexagon_session::enqueue_op(const htp_opnode & node) {
-    for (auto t : node.get_inputs()) {
+    auto clone_tensor_buffer = [this](const ggml_tensor * t) {
         if (t && t->buffer && ggml_backend_buffer_is_hexagon(t->buffer)) {
+            auto sbuf = static_cast<const ggml_hexagon_shared_buffer *>(t->buffer->context);
             if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != this) {
-                this->clone_buffer(static_cast<const ggml_hexagon_shared_buffer *>(t->buffer->context));
+                this->clone_buffer(sbuf);
+            }
+            for (auto & sub : this->mdev.sessions) {
+                sub->clone_buffer(sbuf);
             }
         }
+    };
+
+    for (auto t : node.get_inputs()) {
+        clone_tensor_buffer(t);
     }
     for (auto t : node.get_outputs()) {
-        if (t && t->buffer && ggml_backend_buffer_is_hexagon(t->buffer)) {
-            if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != this) {
-                this->clone_buffer(static_cast<const ggml_hexagon_shared_buffer *>(t->buffer->context));
-            }
-        }
+        clone_tensor_buffer(t);
     }

     if (opt_opfusion && op_batch->try_fuse(node)) {
@@ -2808,39 +2928,84 @@ void ggml_hexagon_session::enqueue_op(const htp_opnode & node) {
     }

     if (!op_batch->fit_op(node)) {
-        flush_batch();
+        flush_async();
     }
+
+    if (this->mdev.count > 1 && op_batch->n_ops == 0) {
+        enqueue_mdev_group();
+    }
+
     op_batch->add_op(node);
 }

+void ggml_hexagon_session::enqueue_mdev_group() {
+    htp_opnode group_node(HTP_OP_MDEV_GROUP);
+
+    uint8_t * fence_slot = this->mdev_fence_slot;
+
+    static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };
+    ggml_tensor dummy_t {};
+    dummy_t.buffer = &this->fence_buf->backend_buffer;
+    dummy_t.extra  = &fence_extra;
+    dummy_t.data   = (void *) fence_slot;
+    dummy_t.type   = GGML_TYPE_I8;
+    dummy_t.ne[0]  = HTP_FENCE_SLOT_SIZE;
+    dummy_t.ne[1]  = (int64_t) this->mdev.count;
+    dummy_t.ne[2]  = 1;
+    dummy_t.ne[3]  = 1;
+    dummy_t.nb[0]  = 1;
+    dummy_t.nb[1]  = HTP_FENCE_SLOT_SIZE;
+    dummy_t.nb[2]  = dummy_t.nb[1] * dummy_t.ne[1];
+    dummy_t.nb[3]  = dummy_t.nb[2];
+    dummy_t.op     = GGML_OP_NONE;
+    dummy_t.op_params[0] = (int32_t) this->mdev.idx;
+
+    ggml_tensor * node = group_node.add_dummy(dummy_t);
+    node->src[0] = node;
+    group_node.init(node);
+    group_node.outputs.clear();
+    group_node.name = "MDEV_GROUP";
+
+    if (this->fence_buf->sess != this) {
+        this->clone_buffer(this->fence_buf);
+    }
+    for (auto & sub : this->mdev.sessions) {
+        sub->clone_buffer(this->fence_buf);
+    }
+
+    op_batch->add_op(group_node);
+}
+
 void ggml_hexagon_session::enqueue_cpy(const ggml_tensor * src, ggml_tensor * dst, const ggml_tensor * sync_tensor, uint32_t fence_seq) {
-    htp_opnode cpy_node(HTP_OP_CPY);
+    const bool with_fence = sync_tensor != nullptr;
+    htp_opnode cpy_node(with_fence ? HTP_OP_CPY_FENCE : HTP_OP_CPY);

     ggml_tensor* node = cpy_node.add_dummy(*dst);
     node->op     = GGML_OP_CPY;
     node->src[0] = const_cast<ggml_tensor *>(src);
-    node->src[1] = sync_tensor ? cpy_node.add_dummy(*sync_tensor) : nullptr;
-    if (sync_tensor) {
+    node->src[1] = with_fence ? cpy_node.add_dummy(*sync_tensor) : nullptr;
+    if (with_fence) {
         node->op_params[0] = (int32_t) fence_seq;
     }

     cpy_node.init(node);
-    if (sync_tensor) {
+    if (with_fence) {
         cpy_node.name = "CPY+FENCE";
     }
     this->enqueue_op(cpy_node);
 }

-void ggml_hexagon_session::enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq) {
+void ggml_hexagon_session::enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq, bool wait) {
     htp_opnode sync_node(HTP_OP_FENCE);

     ggml_tensor* node = sync_node.add_dummy(*sync_tensor);
     node->op           = GGML_OP_NONE;
     node->src[0]       = node;
     node->op_params[0] = (int32_t) fence_seq;
+    node->op_params[1] = wait ? 0 : 1;

     sync_node.init(node);
-    sync_node.name = "FENCE";
+    sync_node.name = wait ? "FENCE_WAIT" : "FENCE_SIGNAL";
     this->enqueue_op(sync_node);
 }

@@ -2858,7 +3023,6 @@ static bool ggml_hexagon_precompute_allreduce_params(
     kparams->n_ranks      = (int32_t) n_ranks;
     kparams->is_row_bcast = (has_add && is_row_bcast) ? 1 : 0;

-    const uint32_t n_bufs    = n_ranks + 1 + (has_add ? 1 : 0);
     const uint32_t nelem     = (uint32_t) ggml_nelements(dst);
     const uint32_t elem_size = (dst->type == GGML_TYPE_F16) ? sizeof(ggml_fp16_t) : sizeof(float);
     const bool is_contiguous = ggml_is_contiguous(dst);
@@ -2902,6 +3066,7 @@ static bool ggml_hexagon_precompute_allreduce_params(
         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));
         kparams->n_threads = n_threads;
+        const size_t n_vtcm_buffers = htp_allreduce_vtcm_buffer_count(n_ranks, n_threads, has_add, is_row_bcast);

         uint32_t block_elems = 65536;
         if (block_elems > rank_nelem / n_threads && rank_nelem / n_threads > 128) {
@@ -2911,15 +3076,15 @@ static bool ggml_hexagon_precompute_allreduce_params(

         kparams->block_elems          = block_elems;
         kparams->vtcm_size_per_thread = 2 * block_elems * elem_size;
-        kparams->vtcm_size            = n_threads * n_bufs * kparams->vtcm_size_per_thread;
+        kparams->vtcm_size            = n_vtcm_buffers * kparams->vtcm_size_per_thread;

         while ((size_t) kparams->vtcm_size > sess->vtcm_size && block_elems > 128) {
-            const size_t max_bytes_per_buf = sess->vtcm_size / (n_threads * n_bufs * 2);
+            const size_t max_bytes_per_buf = sess->vtcm_size / (n_vtcm_buffers * 2);
             block_elems = (uint32_t) hex_align_down((size_t) (max_bytes_per_buf / elem_size), 128);
             if (block_elems < 128) break;
             kparams->block_elems          = block_elems;
             kparams->vtcm_size_per_thread = 2 * block_elems * elem_size;
-            kparams->vtcm_size            = n_threads * n_bufs * kparams->vtcm_size_per_thread;
+            kparams->vtcm_size            = n_vtcm_buffers * kparams->vtcm_size_per_thread;
         }

         if (sess->vtcm_size < (size_t) kparams->vtcm_size || block_elems < 128) {
@@ -2935,6 +3100,7 @@ static bool ggml_hexagon_precompute_allreduce_params(
         const uint32_t rank_nrows = (uint32_t) kparams->rank_nelem;
         const uint32_t n_threads  = (std::min)((uint32_t) sess->n_threads, (std::max)(1u, rank_nrows));
         kparams->n_threads = n_threads;
+        const size_t n_vtcm_buffers = htp_allreduce_vtcm_buffer_count(n_ranks, n_threads, has_add, is_row_bcast);

         const uint32_t row_bytes = ne0 * elem_size;
         const uint32_t row_size_aligned = (uint32_t) hex_align_up(row_bytes, 128);
@@ -2946,14 +3112,14 @@ static bool ggml_hexagon_precompute_allreduce_params(
         kparams->block_elems = block_rows;

         kparams->vtcm_size_per_thread = 2 * (block_rows * row_size_aligned);
-        kparams->vtcm_size            = n_threads * n_bufs * kparams->vtcm_size_per_thread;
+        kparams->vtcm_size            = n_vtcm_buffers * kparams->vtcm_size_per_thread;

         while ((size_t) kparams->vtcm_size > sess->vtcm_size && block_rows > 1) {
-            const size_t max_rows_per_buf = sess->vtcm_size / (n_threads * n_bufs * 2 * row_size_aligned);
+            const size_t max_rows_per_buf = sess->vtcm_size / (n_vtcm_buffers * 2 * row_size_aligned);
             block_rows = (std::max)(1u, (uint32_t) max_rows_per_buf);
             kparams->block_elems          = block_rows;
             kparams->vtcm_size_per_thread = 2 * (block_rows * row_size_aligned);
-            kparams->vtcm_size            = n_threads * n_bufs * kparams->vtcm_size_per_thread;
+            kparams->vtcm_size            = n_vtcm_buffers * kparams->vtcm_size_per_thread;
             if (max_rows_per_buf == 0) break;
         }

@@ -3009,28 +3175,20 @@ void ggml_hexagon_session::enqueue_allreduce(
     this->enqueue_op(ar_node);
 }

-void ggml_hexagon_session::wait_event(uint64_t seq) {
-    flush_sync_peers();
-    HEX_VERBOSE("ggml-hex: %s opqueue-wait start: seq %llu, current rsp-seq %llu, pending %d\n",
-                this->name.c_str(), (unsigned long long)seq, (unsigned long long)op_queue->rsp_seq, (int)this->op_pending);
-    while (op_queue->rsp_seq < seq && this->op_pending > 0) {
-        this->flush_pending(false);
-    }
-    HEX_VERBOSE("ggml-hex: %s opqueue-wait end: seq %llu, current rsp-seq %llu, pending %d\n",
-                this->name.c_str(), (unsigned long long)seq, (unsigned long long)op_queue->rsp_seq, (int)this->op_pending);
-}
-
-uint64_t ggml_hexagon_session::record_event() {
-    flush_batch();
-    return op_queue->req_seq;
-}
-
 bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer *sbuf)
 {
-    if (this->cloned_buffers.find(sbuf->fd()) != this->cloned_buffers.end()) return true;
+    GGML_ASSERT(sbuf && sbuf->mem);
+    if (sbuf->sess == this) return true;
+
+    auto mem = sbuf->mem;
+    int   fd = mem->fd;
+
+    GGML_ASSERT(fd >= 0);
+
+    if (this->cloned_buffers.find(fd) != this->cloned_buffers.end()) return true;

     HEX_VERBOSE("ggml-hex: %s clone-buffer: %s base %p size %zu fd %d\n", this->name.c_str(),
-                sbuf->c_name(), sbuf->base(), sbuf->size(), sbuf->fd());
+                sbuf->c_name(), sbuf->base(), sbuf->size(), fd);

     auto clone = std::make_unique<ggml_hexagon_shared_buffer>(this, *sbuf);
     try {
@@ -3040,10 +3198,38 @@ bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer *sbuf)
         return false;
     }

-    this->cloned_buffers[sbuf->fd()] = std::move(clone);
+    this->cloned_buffers[fd] = std::move(clone);
+    mem->mapped_clones.insert(this);
     return true;
 }

+void ggml_hexagon_session::release_buffer(const ggml_hexagon_shared_buffer * sbuf) {
+    GGML_ASSERT(sbuf && sbuf->mem);
+
+    auto mem = sbuf->mem;
+    int   fd = mem->fd;
+
+    GGML_ASSERT(fd >= 0);
+
+    auto it = this->cloned_buffers.find(fd);
+    if (it != this->cloned_buffers.end()) {
+        auto clone = std::move(it->second);
+        this->cloned_buffers.erase(it);
+    }
+    mem->mapped_clones.erase(this);
+}
+
+void ggml_hexagon_session::unclone_buffer(const ggml_hexagon_shared_buffer * sbuf) {
+    GGML_ASSERT(sbuf && sbuf->mem);
+
+    auto mem = sbuf->mem;
+    std::vector<ggml_hexagon_session *> sessions(mem->mapped_clones.begin(), mem->mapped_clones.end());
+
+    for (auto * sess : sessions) {
+        sess->release_buffer(sbuf);
+    }
+}
+
 static size_t ggml_hexagon_measure_max_vmem(ggml_hexagon_session *sess) {
     // Allocate a bunch pinned buffers till failure.
     // This is kind of expensive but handy for figuring out exactly how much we can mmap on a specific device.
@@ -3082,14 +3268,16 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
     this->valid_queue   = false;
     this->valid_iface   = false;

-    this->phys_idx   = phys_idx;
-    this->virt_idx   = virt_idx;
-    this->domain_id  = config.domain_id;
-    this->session_id = 0;
-    this->name       = config.name;
-    this->op_pending = 0;
+    this->name          = config.name;
+    this->phys_idx      = phys_idx;
+    this->virt_idx      = virt_idx;
+    this->domain_id     = config.domain_id;
+    this->session_id    = 0;
+    this->batch_req_seq = 0;
+    this->batch_rsp_seq = 0;
+    this->last_error    = HTP_STATUS_OK;

-    GGML_LOG_DEBUG("ggml-hex: %s allocating new session\n", this->name.c_str());
+    GGML_LOG_DEBUG("ggml-hex: %s allocating new session : domain %u phys-idx %u virt-idx %u\n", this->name.c_str(), this->domain_id, phys_idx, virt_idx);

     if (config.domain_id < 0 || config.domain_name.empty()) {
         GGML_LOG_ERROR("ggml-hex: %s: invalid physical CDSP core %d\n", config.name.c_str(), config.physical_idx);
@@ -3098,25 +3286,14 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n

     const std::string & dom_name = config.domain_name;

-    // Enable Unsigned PD for all domains
-    {
-        struct remote_rpc_control_unsigned_module u;
-        u.domain = -1;
-        u.enable = 1;
-        int err  = remote_session_control(DSPRPC_CONTROL_UNSIGNED_MODULE, (void *) &u, sizeof(u));
-        if (err != AEE_SUCCESS) {
-            GGML_LOG_ERROR("ggml-hex: %s failed to enable unsigned PD : error 0x%x\n", this->c_name(), err);
-            throw std::runtime_error("ggml-hex: remote_session_control(unsign) failed (see log for details)");
-        }
-    }
-
     // Create new session if virtual_idx > 0
     if (virt_idx > 0) {
-        struct remote_rpc_reserve_new_session n;
+        struct remote_rpc_reserve_new_session n {};
         n.domain_name_len  = dom_name.size();
         n.domain_name      = const_cast<char *>(dom_name.c_str());
         n.session_name     = const_cast<char *>(this->name.c_str());
         n.session_name_len = this->name.size();
+        n.session_id       = virt_idx;

         int err = remote_session_control(FASTRPC_RESERVE_NEW_SESSION, (void *) &n, sizeof(n));
         if (err != AEE_SUCCESS) {
@@ -3130,7 +3307,7 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
         this->domain_id     = n.effective_domain_id;
         this->valid_session = true;
     } else {
-        struct remote_rpc_effective_domain_id eff = {};
+        struct remote_rpc_effective_domain_id eff {};
         eff.domain_name     = const_cast<char *>(dom_name.c_str());
         eff.domain_name_len = dom_name.size();
         eff.session_id      = 0;
@@ -3144,6 +3321,18 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
         }
     }

+    // Enable unsigned modules
+    {
+        struct remote_rpc_control_unsigned_module u;
+        u.domain = this->domain_id;
+        u.enable = 1;
+        int err  = remote_session_control(DSPRPC_CONTROL_UNSIGNED_MODULE, (void *) &u, sizeof(u));
+        if (err != AEE_SUCCESS) {
+            GGML_LOG_ERROR("ggml-hex: %s failed to enable unsigned PD : error 0x%x\n", this->c_name(), err);
+            throw std::runtime_error("ggml-hex: remote_session_control(unsign) failed (see log for details)");
+        }
+    }
+
     char session_uri[256];
     {
         char htp_uri[256];
@@ -3171,7 +3360,7 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
     // Open session
     int err = htp_iface_open(session_uri, &this->handle);
     if (err != AEE_SUCCESS) {
-        GGML_LOG_ERROR("ggml-hex: %s failed to open session : error 0x%x\n", this->c_name(), err);
+        GGML_LOG_ERROR("ggml-hex: %s failed to open session : uri %s error 0x%x\n", this->c_name(), session_uri, err);
         throw std::runtime_error("ggml-hex: failed to open session (see log for details)");
     }

@@ -3186,8 +3375,9 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
         unsigned long long hw_vtcm_size = 0;
         int hw_err = htp_iface_hwinfo(this->handle, &hw_n_threads, &hw_n_hvx, &hw_n_hmx, &hw_vtcm_size);
         if (hw_err == 0) {
-            this->n_threads = opt_nhvx > 0 ? (uint32_t)opt_nhvx : (uint32_t)hw_n_threads;
-            this->n_hvx     = opt_nhvx > 0 ? (uint32_t)opt_nhvx : (uint32_t)hw_n_hvx;
+            const uint32_t max_n_threads = (std::min)((uint32_t) HTP_MAX_NTHREADS, (uint32_t) hw_n_threads);
+            this->n_threads = opt_nhvx > 0 ? (uint32_t) (std::min)(opt_nhvx, (size_t) max_n_threads) : max_n_threads;
+            this->n_hvx     = this->n_threads;
             this->n_hmx     = (opt_nhmx != 0) ? (uint32_t)hw_n_hmx : 0;
             this->vtcm_size = (uint64_t)hw_vtcm_size;
             GGML_LOG_INFO("ggml-hex: %s hwinfo: threads %u, hvx %u, hmx %u, vtcm %llu MB\n",
@@ -3195,8 +3385,9 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
                           (unsigned long long)(this->vtcm_size / (1024 * 1024)));
         } else {
             GGML_LOG_WARN("ggml-hex: %s failed to query hwinfo (0x%x), using defaults\n", this->c_name(), hw_err);
-            this->n_threads = opt_nhvx > 0 ? (uint32_t)opt_nhvx : 8;
-            this->n_hvx     = opt_nhvx > 0 ? (uint32_t)opt_nhvx : 8;
+            const uint32_t default_n_threads = (std::min)(8u, (uint32_t) HTP_MAX_NTHREADS);
+            this->n_threads = opt_nhvx > 0 ? (uint32_t) (std::min)(opt_nhvx, (size_t) HTP_MAX_NTHREADS) : default_n_threads;
+            this->n_hvx     = this->n_threads;
             this->n_hmx     = (opt_nhmx != 0) ? 1 : 0;
             this->vtcm_size = 8 * 1024 * 1024;
         }
@@ -3252,6 +3443,11 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
     // Allocate buffers and state for op batching
     this->op_queue = new ggml_hexagon_opqueue(this, opt_opbatch, opt_opqueue);

+    this->fence_buf = new ggml_hexagon_fence_buffer(this, &dev_ctx->fence_buffer_type, 64 * 1024);
+    if (this->mdev.count > 1) {
+        this->mdev_fence_slot = this->alloc_fence(this->mdev.count);
+    }
+
     if (!opt_vmem) {
         opt_vmem = ggml_hexagon_measure_max_vmem(this);
         GGML_LOG_INFO("ggml-hex: %s measured max vmem %zu\n", this->c_name(), opt_vmem);
@@ -3262,7 +3458,7 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
     this->op_batch = new ggml_hexagon_opbatch(this, opt_opbatch, this->max_vmem);

     // Start dspqueue/opbatch processing
-    err = htp_iface_start(this->handle, this->session_id, this->queue_id, opt_nhvx, opt_nhmx, this->max_vmem);
+    err = htp_iface_start(this->handle, this->session_id, this->queue_id, this->n_threads, opt_nhmx, this->max_vmem);
     if (err != 0) {
         GGML_LOG_ERROR("ggml-hex: %s failed to start session: 0x%08x\n", this->c_name(), (unsigned) err);
         throw std::runtime_error("ggml-hex: iface start failed (see log for details)");
@@ -3283,6 +3479,8 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
 void ggml_hexagon_session::release() noexcept(true) {
     GGML_LOG_INFO("ggml-hex: releasing session: %s\n", this->name.c_str());

+    this->mdev.sessions.clear();
+
     int err;

     if (this->valid_iface) {
@@ -3295,6 +3493,19 @@ void ggml_hexagon_session::release() noexcept(true) {

     delete this->op_batch;
     delete this->op_queue;
+    for (auto & it : this->cpy_fence_slots) {
+        free_fence((void *) it.second, 1);
+    }
+    this->cpy_fence_slots.clear();
+
+    if (this->fence_buf) {
+        unclone_buffer(this->fence_buf);
+        delete this->fence_buf;
+        this->fence_buf = nullptr;
+    }
+    while (!this->cloned_buffers.empty()) {
+        release_buffer(this->cloned_buffers.begin()->second.get());
+    }

     if (opt_etm) {
         err = htp_iface_etm(this->handle, 0);
@@ -3321,23 +3532,30 @@ void ggml_hexagon_session::release() noexcept(true) {
     if (this->valid_handle) {
         htp_iface_close(this->handle);
     }
-
-    this->cloned_buffers.clear();
 }

-ggml_hexagon_session::ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev) noexcept(false) {
-    op_batch = nullptr;
-    op_queue = nullptr;
-    fence_seq = ((uintptr_t)this) & 0xFFFF;
+ggml_hexagon_session::ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev, uint32_t mdev_idx, uint32_t mdev_count) noexcept(false) {
+    this->dev        = dev;
+    this->dev_ctx    = static_cast<ggml_backend_hexagon_device_context *>(dev->context);
+    this->mdev.idx   = mdev_idx;
+    this->mdev.count = mdev_count > 0 ? mdev_count : (uint32_t) (1 + config.mdev_group.size());
+    op_batch         = nullptr;
+    op_queue         = nullptr;
+    fence_buf        = nullptr;
+    fence_seq        = ((uintptr_t)this) & 0xFFFF;

     try {
         allocate(config);
+        if (this->mdev.idx == 0 && !config.mdev_group.empty()) {
+            for (size_t i = 0; i < config.mdev_group.size(); i++) {
+                this->mdev.sessions.push_back(std::make_unique<ggml_hexagon_session>(
+                    config.mdev_group[i], this->dev, (uint32_t) (i + 1), this->mdev.count));
+            }
+        }
     } catch (const std::exception & exc) {
         release();
         throw;
     }
-
-    GGML_UNUSED(dev);
 }

 ggml_hexagon_session::~ggml_hexagon_session() noexcept(true) {
@@ -3563,10 +3781,6 @@ static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_ses
     const struct ggml_tensor * state = op->src[5];
     const struct ggml_tensor * dst   = op;

-    if (!q || !k || !v || !g || !beta || !state) {
-        return false;
-    }
-
     if (q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F32 || v->type != GGML_TYPE_F32 ||
         g->type != GGML_TYPE_F32 || beta->type != GGML_TYPE_F32 || state->type != GGML_TYPE_F32 ||
         dst->type != GGML_TYPE_F32) {
@@ -3754,6 +3968,7 @@ static void ggml_hexagon_precompute_hvx_mm_params(
     struct htp_mm_kernel_params * kparams
 ) {
     kparams->n_hmx = 0;
+    kparams->n_threads = sess->n_threads;

     const bool is_quant = (wtype != GGML_TYPE_F16 && wtype != GGML_TYPE_F32);
     const int src1_nrows = ne11 * ne12 * ne13;
@@ -4193,6 +4408,7 @@ static void ggml_hexagon_precompute_fused_mmnx_params(
     struct htp_mm_kernel_params * kparams
 ) {
     memset(kparams, 0, sizeof(*kparams));
+    kparams->n_threads = sess->n_threads;

     const int ne00 = src0->ne[0];
     const int ne01 = src0->ne[1];
@@ -4921,6 +5137,14 @@ static bool ggml_hexagon_supported_pad(const struct ggml_hexagon_session * sess,
         return false;
     }

+    const int32_t lp0 = ((const int32_t *) op->op_params)[0];
+    const int32_t rp0 = ((const int32_t *) op->op_params)[1];
+    const int32_t circular = ((const int32_t *) op->op_params)[8];
+
+    if (circular && (lp0 > src0->ne[0] || rp0 > src0->ne[0])) {
+        return false;
+    }
+
     return true;

     GGML_UNUSED(sess);
@@ -4972,10 +5196,6 @@ static bool ggml_hexagon_supported_solve_tri(const struct ggml_hexagon_session *
     const struct ggml_tensor * src1 = op->src[1]; // B
     const struct ggml_tensor * dst  = op;         // X

-    if (!src0 || !src1) {
-        return false;
-    }
-
     if (src0->type != GGML_TYPE_F32 || src1->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
         return false;
     }
@@ -5145,7 +5365,7 @@ static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const st
 }

 static bool is_mergeable_mul_mat(const ggml_tensor * t) {
-    if (!t || t->op != GGML_OP_MUL_MAT) return false;
+    if (t->op != GGML_OP_MUL_MAT) return false;

     const ggml_tensor * src0 = t->src[0];
     const ggml_tensor * src1 = t->src[1];
@@ -5179,7 +5399,7 @@ static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor
 }

 static bool is_mergeable_mul_mat_id(const ggml_tensor * t) {
-    if (!t || t->op != GGML_OP_MUL_MAT_ID) return false;
+    if (t->op != GGML_OP_MUL_MAT_ID) return false;

     const ggml_tensor * src0 = t->src[0];
     return ggml_hexagon_is_repack_type(src0->type);
@@ -5213,6 +5433,10 @@ static bool is_mergeable_mul_mat_id_pair(const ggml_tensor * n1, const ggml_tens
 static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, ggml_cgraph * graph) {
     auto sess = static_cast<ggml_hexagon_session *>(backend->context);

+    if (sess->last_error > HTP_STATUS_OK) {
+        return GGML_STATUS_FAILED;
+    }
+
     HEX_VERBOSE("ggml-hex: %s graph-compute n_nodes %d\n", sess->c_name(), graph->n_nodes);

     const std::vector<htp_opnode> * nodes_ptr = nullptr;
@@ -5228,6 +5452,8 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
             auto * extra = (ggml_hexagon_tensor_extra *) graph->nodes[i]->extra;
             if (!extra) continue;

+            extra->flags &= ~GGML_HEXAGON_TENSOR_FUSEABLE;
+
             if (graph->nodes[i]->op == GGML_OP_RMS_NORM && ggml_can_fuse(graph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) {
                 extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE;
             } else if (graph->nodes[i]->op == GGML_OP_MUL_MAT || graph->nodes[i]->op == GGML_OP_MUL_MAT_ID) {
@@ -5299,6 +5525,10 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
         sess->enqueue_op(node);
     }

+    if (sess->last_error > HTP_STATUS_OK) {
+        return GGML_STATUS_FAILED;
+    }
+
     return GGML_STATUS_SUCCESS;
 }

@@ -5308,7 +5538,10 @@ static void ggml_backend_hexagon_synchronize(ggml_backend_t backend) {
     HEX_VERBOSE("ggml-hex: %s synchronize\n", sess->c_name());

     // Wait until all pending ops complete
-    sess->flush();
+    sess->flush_sync();
+    if (sess->last_error > HTP_STATUS_OK) {
+        GGML_ABORT("ggml-hex: %s synchronize failed : dsp-error %s\n", sess->c_name(), status_to_str(sess->last_error));
+    }
 }

 enum ggml_hexagon_mem_range_type {
@@ -5543,27 +5776,38 @@ static void ggml_backend_hexagon_graph_optimize(ggml_backend_t backend, ggml_cgr
     GGML_UNUSED(backend);
 }

+static uint64_t ggml_hexagon_session_key(const ggml_hexagon_session * sess) {
+    return ((uint64_t) (uint32_t) sess->phys_idx << 32) | (uint32_t) sess->virt_idx;
+}
+
 static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) {
     auto sess_src = static_cast<ggml_hexagon_session *>(backend_src->context);
     auto sess_dst = static_cast<ggml_hexagon_session *>(backend_dst->context);
     auto sbuf_dst = (ggml_hexagon_shared_buffer *) dst->buffer->context;

-    if (sess_dst->fence_seq == 0) sess_dst->fence_seq = 1;
-    uint32_t fence_seq = sess_dst->fence_seq++;
-    if (sess_dst->fence_seq == 0) sess_dst->fence_seq = 1;
+    if (!sess_src->clone_buffer(sbuf_dst)) { return false; }

-    volatile uint32_t * fence = (volatile uint32_t *) sbuf_dst->alloc_fence();
+    const uint64_t src_key = ggml_hexagon_session_key(sess_src);
+    auto & fence_slot = sess_dst->cpy_fence_slots[src_key];
+    if (!fence_slot) {
+        fence_slot = (volatile uint32_t *) sess_dst->alloc_fence(1);
+    }
+
+    if (!sess_src->clone_buffer(sess_dst->fence_buf)) { return false; }
+
+    if (++sess_dst->fence_seq == 0) sess_dst->fence_seq = 1;
+    uint32_t fence_seq = sess_dst->fence_seq;

-    HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu : seq %u\n",
+    HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu : seq 0x%x\n",
                 sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src), fence_seq);

-    // dummy extra (must be static)
+    // dummy fence extra (must be static)
     static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };

     ggml_tensor fence_tensor {};
-    fence_tensor.buffer = dst->buffer;
+    fence_tensor.buffer = &sess_dst->fence_buf->backend_buffer;
     fence_tensor.extra  = &fence_extra;
-    fence_tensor.data   = (void *) fence;
+    fence_tensor.data   = (void *) fence_slot;
     fence_tensor.type   = GGML_TYPE_I32;
     fence_tensor.ne[0]  = 1;
     fence_tensor.ne[1]  = 1;
@@ -5576,9 +5820,9 @@ static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_
     fence_tensor.op     = GGML_OP_NONE;

     sess_src->enqueue_cpy(src, dst, &fence_tensor, fence_seq);
-    sess_dst->enqueue_fence(&fence_tensor, fence_seq);
+    sess_dst->enqueue_fence(&fence_tensor, fence_seq, /* wait = */ true);

-    sess_dst->add_sync_peer(sess_src);
+    sess_dst->add_peer(sess_src);

     return true;
 }
@@ -5586,15 +5830,15 @@ static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_
 static bool ggml_hexagon_cpy_tensor_async_virt(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) {
     auto sess_src = static_cast<ggml_hexagon_session *>(backend_src->context);
     auto sess_dst = static_cast<ggml_hexagon_session *>(backend_dst->context);
-    auto sbuf_dst = (ggml_hexagon_shared_buffer *) dst->buffer->context;
+    auto sbuf_src = (ggml_hexagon_shared_buffer *) src->buffer->context;

-    if (!sess_src->clone_buffer(sbuf_dst)) { return false; }
+    if (!sess_dst->clone_buffer(sbuf_src)) { return false; }

     HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu\n",
                 sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src));

-    sess_src->enqueue_cpy(src, dst);
-    sess_src->flush(true);
+    sess_dst->enqueue_cpy(src, dst);
+    sess_dst->add_peer(sess_src);

     return true;
 }
@@ -5604,7 +5848,14 @@ static bool ggml_backend_hexagon_cpy_tensor_async(ggml_backend_t backend_src, gg
         return false;
     }

-    *(ggml_hexagon_tensor_extra *) dst->extra = *(const ggml_hexagon_tensor_extra *) src->extra;
+    // FIXME: ggml-meta needs to call init_tensor on auxiliary tensors
+    if (!dst->extra) {
+        ggml_backend_buffer_init_tensor(dst->buffer, dst);
+    }
+
+    auto * dst_extra = static_cast<ggml_hexagon_tensor_extra *>(dst->extra);
+    const auto * src_extra = static_cast<const ggml_hexagon_tensor_extra *>(src->extra);
+    dst_extra->flags = src_extra->flags & ~GGML_HEXAGON_TENSOR_FUSEABLE;

     auto sess_src = static_cast<ggml_hexagon_session *>(backend_src->context);
     auto sess_dst = static_cast<ggml_hexagon_session *>(backend_dst->context);
@@ -5612,7 +5863,6 @@ static bool ggml_backend_hexagon_cpy_tensor_async(ggml_backend_t backend_src, gg
     if (sess_src == sess_dst) {
         HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu\n", sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src));
         sess_src->enqueue_cpy(src, dst);
-        sess_src->flush_batch();
         return true;
     }

@@ -5623,8 +5873,30 @@ static bool ggml_backend_hexagon_cpy_tensor_async(ggml_backend_t backend_src, gg
 }

 static ggml_backend_event_t ggml_backend_hexagon_device_event_new(ggml_backend_dev_t dev) {
+    auto dev_ctx = static_cast<ggml_backend_hexagon_device_context *>(dev->context);
+    auto sess    = dev_ctx->session();
+
     ggml_hexagon_event * hex_event = new ggml_hexagon_event();
-    HEX_VERBOSE("ggml-hex: %s event-new : event %p\n", ggml_backend_dev_name(dev), (void *)hex_event);
+    hex_event->fence_sess = sess;
+    hex_event->sess       = sess;
+    hex_event->fence_slot = (volatile uint32_t *) sess->alloc_fence(1);
+
+    static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };
+    hex_event->fence_tensor.buffer = &sess->fence_buf->backend_buffer;
+    hex_event->fence_tensor.extra  = &fence_extra;
+    hex_event->fence_tensor.data   = (void *) hex_event->fence_slot;
+    hex_event->fence_tensor.type   = GGML_TYPE_I32;
+    hex_event->fence_tensor.ne[0]  = 1;
+    hex_event->fence_tensor.ne[1]  = 1;
+    hex_event->fence_tensor.ne[2]  = 1;
+    hex_event->fence_tensor.ne[3]  = 1;
+    hex_event->fence_tensor.nb[0]  = sizeof(int32_t);
+    hex_event->fence_tensor.nb[1]  = sizeof(int32_t);
+    hex_event->fence_tensor.nb[2]  = sizeof(int32_t);
+    hex_event->fence_tensor.nb[3]  = sizeof(int32_t);
+    hex_event->fence_tensor.op     = GGML_OP_NONE;
+
+    HEX_VERBOSE("ggml-hex: %s event-new : event %p fence %p\n", ggml_backend_dev_name(dev), (void *)hex_event, (void *)hex_event->fence_slot);

     return new ggml_backend_event {
         /* .device  = */ dev,
@@ -5632,49 +5904,83 @@ static ggml_backend_event_t ggml_backend_hexagon_device_event_new(ggml_backend_d
     };
 }

-static void ggml_backend_hexagon_device_event_free(ggml_backend_dev_t dev, ggml_backend_event_t event) {
-    GGML_UNUSED(dev);
-
-    if (event == nullptr) {
+static void ggml_hexagon_event_synchronize(ggml_backend_dev_t dev, ggml_hexagon_event * hex_event) {
+    if (hex_event->seq == 0) {
         return;
     }

-    ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context;
+    HEX_VERBOSE("ggml-hex: %s event-synchronize : event %p seq 0x%x fence %p\n",
+                ggml_backend_dev_name(dev), (void *)hex_event, hex_event->seq, (void *)hex_event->fence_slot);
+
+    auto * fence = reinterpret_cast<const volatile std::atomic<uint32_t> *>(hex_event->fence_slot);
+
+    if ((int32_t)(fence[0].load(std::memory_order_relaxed) - hex_event->seq) < 0) {
+        hex_event->sess->flush_async();
+    }
+
+    while (true) {
+        if ((int32_t)(fence[0].load(std::memory_order_acquire) - hex_event->seq) >= 0) {
+            uint32_t status = fence[1].load(std::memory_order_acquire);
+            if (status > HTP_STATUS_OK) {
+                GGML_ABORT("ggml-hex: %s event-synchronize failed : dsp-error %s\n",
+                           hex_event->sess->c_name(), status_to_str(status));
+            }
+            break;
+        }
+        std::this_thread::yield();
+    }
+}
+
+static void ggml_backend_hexagon_device_event_free(ggml_backend_dev_t dev, ggml_backend_event_t event) {
+    auto * hex_event = static_cast<ggml_hexagon_event *>(event->context);
+    ggml_hexagon_event_synchronize(dev, hex_event);
     HEX_VERBOSE("ggml-hex: %s event-free : event %p\n", ggml_backend_dev_name(dev), (void *)hex_event);
+    hex_event->fence_sess->free_fence((void *) hex_event->fence_slot, 1);
     delete hex_event;
     delete event;
 }

 static void ggml_backend_hexagon_device_event_synchronize(ggml_backend_dev_t dev, ggml_backend_event_t event) {
-    GGML_UNUSED(dev);
-
-    ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context;
-    HEX_VERBOSE("ggml-hex: %s event-synchronize : event %p seq %llu\n",
-                ggml_backend_dev_name(dev), (void *)hex_event, (unsigned long long)hex_event->seq);
-    if (hex_event->sess != nullptr) {
-        hex_event->sess->wait_event(hex_event->seq);
-    }
+    auto * hex_event = static_cast<ggml_hexagon_event *>(event->context);
+    ggml_hexagon_event_synchronize(dev, hex_event);
 }

 static void ggml_backend_hexagon_event_record(ggml_backend_t backend, ggml_backend_event_t event) {
     auto sess = static_cast<ggml_hexagon_session *>(backend->context);
-    ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context;
+    auto hex_event = static_cast<ggml_hexagon_event *>(event->context);

+    if (++sess->fence_seq == 0) sess->fence_seq = 1;
     hex_event->sess = sess;
-    hex_event->seq  = sess->record_event();
-    HEX_VERBOSE("ggml-hex: %s event-record : event %p seq %llu\n",
-                sess->c_name(), (void *)hex_event, (unsigned long long)hex_event->seq);
+    hex_event->seq  = sess->fence_seq;
+
+    sess->enqueue_fence(&hex_event->fence_tensor, hex_event->seq, /* wait = */ false);
+
+    HEX_VERBOSE("ggml-hex: %s event-record : event %p seq 0x%x fence %p\n",
+                sess->c_name(), (void *)hex_event, hex_event->seq, (void *)hex_event->fence_slot);
 }

 static void ggml_backend_hexagon_event_wait(ggml_backend_t backend, ggml_backend_event_t event) {
-    GGML_UNUSED(backend);
+    auto sess = static_cast<ggml_hexagon_session *>(backend->context);
+    auto hex_event = static_cast<ggml_hexagon_event *>(event->context);
+
+    if (hex_event->seq == 0) {
+        return;
+    }
+
+    HEX_VERBOSE("ggml-hex: %s event-wait : event %p seq 0x%x fence %p\n",
+                sess->c_name(), (void *)hex_event, hex_event->seq, (void *)hex_event->fence_slot);

-    ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context;
-    if (hex_event->sess != nullptr) {
-        HEX_VERBOSE("ggml-hex: %s event-wait : event %p seq %llu\n",
-                    hex_event->sess->c_name(), (void *)hex_event, (unsigned long long)hex_event->seq);
-        hex_event->sess->wait_event(hex_event->seq);
+    // same physical NPU runs sequentially in FIFO order
+    if (sess->phys_idx == hex_event->sess->phys_idx) {
+        if (sess != hex_event->sess) {
+            sess->add_peer(hex_event->sess);
+        }
+        return;
     }
+
+    sess->clone_buffer(hex_event->fence_sess->fence_buf);
+    sess->add_peer(hex_event->sess);
+    sess->enqueue_fence(&hex_event->fence_tensor, hex_event->seq, /* wait = */ true);
 }

 static void ggml_backend_hexagon_set_tensor_async(ggml_backend_t backend, struct ggml_tensor * tensor, const void * data, size_t offset, size_t size) {
@@ -5688,7 +5994,10 @@ static void ggml_backend_hexagon_get_tensor_async(ggml_backend_t backend, const
     auto sess = static_cast<ggml_hexagon_session *>(backend->context);
     HEX_VERBOSE("ggml-hex: %s get-tensor-async %s : data %p offset %zu size %zu usage %d\n",
                 sess->c_name(), tensor->name, data, offset, size, tensor->buffer ? (int) tensor->buffer->usage : -1);
-    sess->flush(true);
+    sess->flush_sync();
+    if (sess->last_error > HTP_STATUS_OK) {
+        GGML_ABORT("ggml-hex: %s get-tensor-async failed : dsp-error %s\n", sess->c_name(), status_to_str(sess->last_error));
+    }
     ggml_backend_tensor_get(tensor, data, offset, size);
 }

@@ -5717,7 +6026,10 @@ static void ggml_backend_hexagon_get_tensor_2d_async(ggml_backend_t backend,
     auto sess = static_cast<ggml_hexagon_session *>(backend->context);
     HEX_VERBOSE("ggml-hex: %s get-tensor-2d-async %s : data %p offset %zu size %zu n_copies %zu stride_tensor %zu stride_data %zu usage %d\n",
                 sess->c_name(), tensor->name, data, offset, size, n_copies, stride_tensor, stride_data, tensor->buffer ? (int) tensor->buffer->usage : -1);
-    sess->flush(true);
+    sess->flush_sync();
+    if (sess->last_error > HTP_STATUS_OK) {
+        GGML_ABORT("ggml-hex: %s get-tensor-2d-async failed : dsp-error %s\n", sess->c_name(), status_to_str(sess->last_error));
+    }
     ggml_backend_tensor_get_2d(tensor, data, offset, size, n_copies, stride_tensor, stride_data);
 }

@@ -6083,13 +6395,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
 static bool ggml_backend_hexagon_device_supports_buft(ggml_backend_dev_t dev, ggml_backend_buffer_type_t buft) {
     auto dev_ctx = static_cast<ggml_backend_hexagon_device_context *>(dev->context);

-    // Technically we can clone hexagon buffers from any session but for some reason the output is garbled with layer-split,
-    // tensor-split works correctly, so it needs mode debugging and investigation. For now accept only our own buffers.
-#if 0
-    bool supp = (buft->iface.get_alignment == ggml_backend_hexagon_buffer_type_get_alignment);
-#else
     bool supp = (buft == &dev_ctx->host_buffer_type) || (buft == &dev_ctx->buffer_type);
-#endif

     HEX_VERBOSE("ggml-hex: %s device-supports-buft %s %s\n", dev_ctx->c_name(), ggml_backend_buft_name(buft), supp ? "yes" : "no");
     return supp;
@@ -6122,6 +6428,19 @@ ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) {

     // Create devices
     for (size_t i = 0; i < opt_ndev; i++) {
+        const auto & cfg = opt_device_configs[i];
+        if (cfg.mdev_group.empty()) {
+            GGML_LOG_INFO("ggml-hex: device %zu: %s (phys=%d, virt=%d, domain=%s:%d)\n",
+                          i, cfg.name.c_str(), cfg.physical_idx, cfg.virtual_idx, cfg.domain_name.c_str(), cfg.domain_id);
+        } else {
+            std::string peers_str;
+            for (const auto & p : cfg.mdev_group) {
+                if (!peers_str.empty()) peers_str += ", ";
+                peers_str += p.name + " (phys=" + std::to_string(p.physical_idx) + ")";
+            }
+            GGML_LOG_INFO("ggml-hex: device %zu: %s (phys=%d, virt=%d, domain=%s:%d) [mdev peers: %s]\n",
+                          i, cfg.name.c_str(), cfg.physical_idx, cfg.virtual_idx, cfg.domain_name.c_str(), cfg.domain_id, peers_str.c_str());
+        }
         devices[i].iface   = ggml_backend_hexagon_device_i;
         devices[i].reg     = reg;
         devices[i].context = new ggml_backend_hexagon_device_context(i, opt_device_configs[i], &devices[i]);
@@ -6172,17 +6491,51 @@ static void * ggml_backend_hexagon_comm_init(ggml_backend_t * backends, size_t n
         }
     }

+    for (size_t i = 0; i < n_backends; i++) {
+        auto sess_i = static_cast<ggml_hexagon_session *>(backends[i]->context);
+        for (size_t j = i + 1; j < n_backends; j++) {
+            auto sess_j = static_cast<ggml_hexagon_session *>(backends[j]->context);
+            if (sess_i->phys_idx == sess_j->phys_idx) {
+                return nullptr;
+            }
+        }
+    }
+
     auto * ctx = new ggml_backend_hexagon_comm_context();
     ctx->backends.assign(backends, backends + n_backends);
     ctx->n_backends = n_backends;
-    ctx->fence_seq  = (((uintptr_t) ctx) & 0xFFFF) | 1;
+
+    static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };
+    for (size_t i = 0; i < n_backends; i++) {
+        auto sess_i = static_cast<ggml_hexagon_session *>(backends[i]->context);
+        ctx->fence_slots[i] = (volatile uint32_t *) sess_i->alloc_fence(1);
+        ctx->fence_tensors[i] = {};
+        ctx->fence_tensors[i].buffer = &sess_i->fence_buf->backend_buffer;
+        ctx->fence_tensors[i].extra  = &fence_extra;
+        ctx->fence_tensors[i].data   = (void *) ctx->fence_slots[i];
+        ctx->fence_tensors[i].type   = GGML_TYPE_I32;
+        ctx->fence_tensors[i].ne[0]  = 4;
+        ctx->fence_tensors[i].ne[1]  = 1;
+        ctx->fence_tensors[i].ne[2]  = 1;
+        ctx->fence_tensors[i].ne[3]  = 1;
+        ctx->fence_tensors[i].nb[0]  = sizeof(int32_t);
+        ctx->fence_tensors[i].nb[1]  = sizeof(int32_t);
+        ctx->fence_tensors[i].nb[2]  = sizeof(int32_t);
+        ctx->fence_tensors[i].nb[3]  = sizeof(int32_t);
+        ctx->fence_tensors[i].op     = GGML_OP_NONE;
+    }

     return ctx;
 }

 static void ggml_backend_hexagon_comm_free(void * comm_ctx_v) {
     if (!comm_ctx_v) return;
-    delete static_cast<ggml_backend_hexagon_comm_context *>(comm_ctx_v);
+    auto * ctx = static_cast<ggml_backend_hexagon_comm_context *>(comm_ctx_v);
+    for (size_t i = 0; i < ctx->n_backends; i++) {
+        auto sess_i = static_cast<ggml_hexagon_session *>(ctx->backends[i]->context);
+        sess_i->free_fence((void *) ctx->fence_slots[i], 1);
+    }
+    delete ctx;
 }

 static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct ggml_tensor ** tensors) {
@@ -6192,6 +6545,16 @@ static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct

     if (n_backends < 2 || n_backends > 4) return false;

+    for (size_t i = 0; i < n_backends; i++) {
+        auto sess_i = static_cast<ggml_hexagon_session *>(comm_ctx->backends[i]->context);
+        for (size_t j = i + 1; j < n_backends; j++) {
+            auto sess_j = static_cast<ggml_hexagon_session *>(comm_ctx->backends[j]->context);
+            if (sess_i->phys_idx == sess_j->phys_idx) {
+                return false;
+            }
+        }
+    }
+
     for (size_t i = 0; i < n_backends; i++) {
         if (!tensors[i] || !tensors[i]->buffer || !ggml_backend_buffer_is_hexagon(tensors[i]->buffer)) {
             return false;
@@ -6219,42 +6582,28 @@ static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct
         }
     }

-    if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1;
-    uint32_t fence_seq_entry = comm_ctx->fence_seq++;
-    if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1;
-    uint32_t fence_seq_exit  = comm_ctx->fence_seq++;
-    if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1;
-
-    volatile uint32_t * fences[GGML_HEXAGON_MAX_SESSIONS];
-    for (size_t i = 0; i < n_backends; i++) {
-        auto sbuf = (ggml_hexagon_shared_buffer *) tensors[i]->buffer->context;
-        fences[i] = (volatile uint32_t *) sbuf->alloc_fence();
+    uint32_t max_seq = static_cast<ggml_hexagon_session *>(comm_ctx->backends[0]->context)->fence_seq;
+    for (size_t i = 1; i < n_backends; i++) {
+        auto sess_i = static_cast<ggml_hexagon_session *>(comm_ctx->backends[i]->context);
+        if ((int32_t)(sess_i->fence_seq - max_seq) > 0) {
+            max_seq = sess_i->fence_seq;
+        }
     }
+    if (++max_seq == 0) max_seq = 1;
+    uint32_t fence_seq_entry = max_seq;
+    if (++max_seq == 0) max_seq = 1;
+    uint32_t fence_seq_exit  = max_seq;

-    static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE };
-    ggml_tensor fence_tensors[GGML_HEXAGON_MAX_SESSIONS];
     for (size_t i = 0; i < n_backends; i++) {
-        fence_tensors[i] = {};
-        fence_tensors[i].buffer = tensors[i]->buffer;
-        fence_tensors[i].extra  = &fence_extra;
-        fence_tensors[i].data   = (void *) fences[i];
-        fence_tensors[i].type   = GGML_TYPE_I32;
-        fence_tensors[i].ne[0]  = 4;
-        fence_tensors[i].ne[1]  = 1;
-        fence_tensors[i].ne[2]  = 1;
-        fence_tensors[i].ne[3]  = 1;
-        fence_tensors[i].nb[0]  = sizeof(int32_t);
-        fence_tensors[i].nb[1]  = sizeof(int32_t);
-        fence_tensors[i].nb[2]  = sizeof(int32_t);
-        fence_tensors[i].nb[3]  = sizeof(int32_t);
-        fence_tensors[i].op     = GGML_OP_NONE;
+        auto sess_i = static_cast<ggml_hexagon_session *>(comm_ctx->backends[i]->context);
+        sess_i->fence_seq = max_seq;
     }

     std::vector<const ggml_tensor *> data_tensors(n_backends);
     std::vector<const ggml_tensor *> sync_tensors(n_backends);
     for (size_t i = 0; i < n_backends; i++) {
         data_tensors[i] = tensors[i];
-        sync_tensors[i] = &fence_tensors[i];
+        sync_tensors[i] = &comm_ctx->fence_tensors[i];
     }

     for (size_t r = 0; r < n_backends; r++) {
@@ -6262,7 +6611,7 @@ static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct
         sess->enqueue_allreduce(tensors[r], data_tensors, sync_tensors, (uint32_t) r, (uint32_t) n_backends, fence_seq_entry, fence_seq_exit);
         for (size_t j = 0; j < n_backends; j++) {
             if (r != j) {
-                sess->add_sync_peer(static_cast<ggml_hexagon_session *>(comm_ctx->backends[j]->context));
+                sess->add_peer(static_cast<ggml_hexagon_session *>(comm_ctx->backends[j]->context));
             }
         }
     }
@@ -6270,8 +6619,23 @@ static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct
     return true;
 }

+static ggml_backend_buffer_type_t ggml_backend_hexagon_split_buffer_type(int main_device, const float * tensor_split) {
+    GGML_UNUSED(tensor_split);
+    auto reg = ggml_backend_hexagon_reg();
+    auto dev = ggml_backend_reg_dev_get(reg, main_device);
+    if (!dev) {
+        dev = ggml_backend_reg_dev_get(reg, 0);
+    }
+    if (!dev) return nullptr;
+    auto dev_ctx = static_cast<ggml_backend_hexagon_device_context *>(dev->context);
+    return &dev_ctx->buffer_type;
+}
+
 static void * ggml_backend_hexagon_get_proc_address(ggml_backend_reg_t reg, const char * name) {
     GGML_UNUSED(reg);
+    if (strcmp(name, "ggml_backend_split_buffer_type") == 0) {
+        return (void *) ggml_backend_hexagon_split_buffer_type;
+    }
     if (strcmp(name, "ggml_backend_comm_init") == 0) {
         return (void *) ggml_backend_hexagon_comm_init;
     }
@@ -6304,6 +6668,41 @@ template<typename T, int BASE=10> std::string vec_to_str(std::vector<T> v) {
     return str;
 }

+static void ggml_hexagon_resolve_device_domain(ggml_hexagon_device_config & cfg, bool discovery_supported, const std::unordered_map<int, fastrpc_domain> & cdsp_map) {
+    if (discovery_supported) {
+        auto it = cdsp_map.find(cfg.physical_idx);
+        if (it != cdsp_map.end()) {
+            cfg.domain_id   = it->second.id;
+            cfg.domain_name = it->second.name;
+        } else {
+            GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not found on device (%zu CDSP core(s) available)\n",
+                           cfg.physical_idx, cdsp_map.size());
+            cfg.domain_id   = -1;
+            cfg.domain_name = "";
+        }
+    } else {
+        switch (cfg.physical_idx) {
+            case 0:
+                cfg.domain_id   = 3;
+                cfg.domain_name = CDSP_DOMAIN_NAME;
+                break;
+            case 1:
+                cfg.domain_id   = 4;
+                cfg.domain_name = "cdsp1";
+                break;
+            default:
+                GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not supported without dynamic discovery\n",
+                               cfg.physical_idx);
+                cfg.domain_id   = -1;
+                cfg.domain_name = "";
+                break;
+        }
+    }
+    for (auto & sub_cfg : cfg.mdev_group) {
+        ggml_hexagon_resolve_device_domain(sub_cfg, discovery_supported, cdsp_map);
+    }
+}
+
 // Enumerate NPU (aka CDSP) domains via FASTRPC_GET_DOMAINS if supported,
 // and populate domain_id and domain_name for all configured devices.
 static void ggml_hexagon_discover_devices() {
@@ -6350,36 +6749,7 @@ static void ggml_hexagon_discover_devices() {

     // Populate domain IDs and names for all configured devices
     for (size_t i = 0; i < opt_ndev; i++) {
-        auto & cfg = opt_device_configs[i];
-        if (discovery_supported) {
-            auto it = cdsp_map.find(cfg.physical_idx);
-            if (it != cdsp_map.end()) {
-                cfg.domain_id   = it->second.id;
-                cfg.domain_name = it->second.name;
-            } else {
-                GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not found on device (%zu CDSP core(s) available)\n",
-                               cfg.physical_idx, cdsp_map.size());
-                cfg.domain_id   = -1;
-                cfg.domain_name = "";
-            }
-        } else {
-            switch (cfg.physical_idx) {
-                case 0:
-                    cfg.domain_id   = 3;
-                    cfg.domain_name = CDSP_DOMAIN_NAME;
-                    break;
-                case 1:
-                    cfg.domain_id   = 4;
-                    cfg.domain_name = "cdsp1";
-                    break;
-                default:
-                    GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not supported without dynamic discovery\n",
-                                   cfg.physical_idx);
-                    cfg.domain_id   = -1;
-                    cfg.domain_name = "";
-                    break;
-            }
-        }
+        ggml_hexagon_resolve_device_domain(opt_device_configs[i], discovery_supported, cdsp_map);
     }
 }

@@ -6487,21 +6857,126 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
                 opt_device_configs[i].physical_idx = 0;
                 opt_device_configs[i].virtual_idx  = (int)i;
                 opt_device_configs[i].name         = "HTP" + std::to_string(i);
+                opt_device_configs[i].mdev_group.clear();
             }
         } else {
             std::string s_devices(str_devices);
-            std::stringstream ss(s_devices);
-            std::string item;
-            opt_ndev = 0;
-            while (std::getline(ss, item, ',')) {
-                size_t start = item.find_first_not_of(" \t\r\n");
-                size_t end = item.find_last_not_of(" \t\r\n");
-                if (start == std::string::npos) {
-                    continue;
+            std::vector<std::string> items;
+            std::string curr_item;
+            int bracket_depth = 0;
+            for (char ch : s_devices) {
+                if (ch == '[') {
+                    bracket_depth++;
+                    curr_item += ch;
+                } else if (ch == ']') {
+                    if (bracket_depth > 0) bracket_depth--;
+                    curr_item += ch;
+                } else if (ch == ',' && bracket_depth == 0) {
+                    size_t s = curr_item.find_first_not_of(" \t\r\n");
+                    size_t e = curr_item.find_last_not_of(" \t\r\n");
+                    if (s != std::string::npos) {
+                        items.push_back(curr_item.substr(s, e - s + 1));
+                    }
+                    curr_item.clear();
+                } else {
+                    curr_item += ch;
                 }
-                item = item.substr(start, end - start + 1);
+            }
+            size_t s = curr_item.find_first_not_of(" \t\r\n");
+            size_t e = curr_item.find_last_not_of(" \t\r\n");
+            if (s != std::string::npos) {
+                items.push_back(curr_item.substr(s, e - s + 1));
+            }
+
+            opt_ndev = 0;
+            for (const auto & item : items) {
+                size_t b_open  = item.find('[');
+                size_t b_close = item.rfind(']');
+
+                if (b_open != std::string::npos && b_close != std::string::npos && b_close > b_open) {
+                    // Grouped / composite syntax: Name[phys_spec:virt] or Name[phys_spec]
+                    std::string dev_name = item.substr(0, b_open);
+                    std::string content  = item.substr(b_open + 1, b_close - b_open - 1);

-                if (item.rfind("HTP", 0) == 0) {
+                    int virt = 0;
+                    std::string phys_spec = content;
+                    size_t colon_pos = content.find(':');
+                    if (colon_pos != std::string::npos) {
+                        phys_spec = content.substr(0, colon_pos);
+                        try {
+                            virt = std::stoi(content.substr(colon_pos + 1));
+                        } catch (...) {
+                            virt = 0;
+                        }
+                    } else {
+                        size_t dev_colon = dev_name.find(':');
+                        if (dev_colon != std::string::npos) {
+                            try {
+                                virt = std::stoi(dev_name.substr(dev_colon + 1));
+                            } catch (...) {
+                                virt = 0;
+                            }
+                        }
+                    }
+
+                    // Parse physical indices from phys_spec (e.g. 0-1, 0,1, 0-3, etc.)
+                    std::vector<int> phys_list;
+                    std::stringstream pss(phys_spec);
+                    std::string p_part;
+                    while (std::getline(pss, p_part, ',')) {
+                        size_t ps = p_part.find_first_not_of(" \t\r\n");
+                        size_t pe = p_part.find_last_not_of(" \t\r\n");
+                        if (ps == std::string::npos) continue;
+                        p_part = p_part.substr(ps, pe - ps + 1);
+
+                        size_t dash_pos = p_part.find('-');
+                        if (dash_pos != std::string::npos) {
+                            try {
+                                int p_start = std::stoi(p_part.substr(0, dash_pos));
+                                int p_end   = std::stoi(p_part.substr(dash_pos + 1));
+                                for (int p = p_start; p <= p_end; p++) {
+                                    if (std::find(phys_list.begin(), phys_list.end(), p) == phys_list.end()) {
+                                        phys_list.push_back(p);
+                                    }
+                                }
+                            } catch (...) {
+                                GGML_LOG_WARN("ggml-hex: failed to parse physical range in '%s'\n", p_part.c_str());
+                            }
+                        } else {
+                            try {
+                                int p = std::stoi(p_part);
+                                if (std::find(phys_list.begin(), phys_list.end(), p) == phys_list.end()) {
+                                    phys_list.push_back(p);
+                                }
+                            } catch (...) {
+                                GGML_LOG_WARN("ggml-hex: failed to parse physical index in '%s'\n", p_part.c_str());
+                            }
+                        }
+                    }
+
+                    if (phys_list.empty()) {
+                        phys_list.push_back(0);
+                    }
+
+                    if (opt_ndev < GGML_HEXAGON_MAX_SESSIONS) {
+                        auto & cfg = opt_device_configs[opt_ndev];
+                        cfg.name         = dev_name;
+                        cfg.physical_idx = phys_list[0];
+                        cfg.virtual_idx  = virt;
+                        cfg.mdev_group.clear();
+
+                        for (size_t k = 1; k < phys_list.size(); k++) {
+                            ggml_hexagon_device_config sub_cfg;
+                            sub_cfg.physical_idx = phys_list[k];
+                            sub_cfg.virtual_idx  = virt;
+                            sub_cfg.name         = "HTP" + std::to_string(phys_list[k]) + ":" + std::to_string(virt);
+                            cfg.mdev_group.push_back(sub_cfg);
+                        }
+                        opt_ndev++;
+                    } else {
+                        GGML_LOG_WARN("ggml-hex: max sessions limit reached (%d), ignoring device %s\n", GGML_HEXAGON_MAX_SESSIONS, item.c_str());
+                    }
+                } else if (item.rfind("HTP", 0) == 0) {
                     std::string rest = item.substr(3);
                     size_t colon_pos = rest.find(':');
                     int phys = 0;
@@ -6525,6 +7000,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
                         opt_device_configs[opt_ndev].name         = colon_pos == std::string::npos
                             ? "HTP" + std::to_string(phys)
                             : "HTP" + std::to_string(phys) + ":" + std::to_string(virt);
+                        opt_device_configs[opt_ndev].mdev_group.clear();
                         opt_ndev++;
                     } else {
                         GGML_LOG_WARN("ggml-hex: max sessions limit reached (%d), ignoring device %s\n", GGML_HEXAGON_MAX_SESSIONS, item.c_str());
@@ -6539,6 +7015,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
         opt_device_configs[0].physical_idx = 0;
         opt_device_configs[0].virtual_idx  = 0;
         opt_device_configs[0].name         = "HTP0";
+        opt_device_configs[0].mdev_group.clear();
     }

 #if defined(__ANDROID__)
diff --git a/ggml/src/ggml-hexagon/htp-opnode.h b/ggml/src/ggml-hexagon/htp-opnode.h
index b083e2671..ef7b5184f 100644
--- a/ggml/src/ggml-hexagon/htp-opnode.h
+++ b/ggml/src/ggml-hexagon/htp-opnode.h
@@ -344,6 +344,12 @@ struct htp_opformat {
         } else if (htp_op_is_unary(node.opcode)) {
             const auto * kparams = (const struct htp_unary_kernel_params *) node.kernel_params;
             snprintf(str, max_size, "%s vtcm %d", kparams->col_tile ? "wide-row" : "row-block", (int) kparams->vtcm_size);
+        } else if (node.opcode == HTP_OP_MDEV_GROUP && node.node) {
+            snprintf(str, max_size, "idx %d count %d", (int) node.node->op_params[0], (int) node.dst()->ne[1]);
+        } else if ((node.opcode == HTP_OP_FENCE || node.opcode == HTP_OP_CPY_FENCE) && node.node) {
+            snprintf(str, max_size, "seq 0x%x", (uint32_t) node.node->op_params[0]);
+        } else if (node.opcode == HTP_OP_ALLREDUCE && node.node) {
+            snprintf(str, max_size, "seq 0x%x -> 0x%x", (uint32_t) node.node->op_params[0], (uint32_t) node.node->op_params[1]);
         } else {
             snprintf(str, max_size, "----");
         }
diff --git a/ggml/src/ggml-hexagon/htp/act-ops.c b/ggml/src/ggml-hexagon/htp/act-ops.c
index ac00b447d..5fff372f2 100644
--- a/ggml/src/ggml-hexagon/htp/act-ops.c
+++ b/ggml/src/ggml-hexagon/htp/act-ops.c
@@ -3,7 +3,6 @@
 #pragma clang diagnostic ignored "-Wunused-but-set-variable"

 #include <HAP_farf.h>
-#include <HAP_perf.h>

 #include <math.h>
 #include <string.h>
@@ -15,7 +14,7 @@
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "hex-common.h"
 #include "htp-tensor.h"
 #include "htp-vtcm.h"

@@ -80,6 +79,7 @@ struct htp_act_context {
     uint32_t                 block;
     uint32_t                 src0_nrows;
     uint32_t                 src0_nrows_per_thread;
+    uint32_t                 row_start;
     int                      nc;

     uint8_t *                vtcm_src0;
@@ -329,104 +329,104 @@ static void geglu_f32(const float * restrict src0,
     }
 }

-#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR)                                                                 \
-    static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) {                         \
-        struct htp_act_context * actx = (struct htp_act_context *) data;                                               \
-        htp_act_preamble;                                                                                              \
-                                                                                                                       \
-        struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL;                          \
-                                                                                                                       \
-        size_t src0_row_size = actx->src0_row_size;                                                                    \
-        size_t src1_row_size = actx->src1_row_size;                                                                    \
-        size_t dst_row_size  = actx->dst_row_size;                                                                     \
-                                                                                                                       \
-        size_t src0_row_stride = actx->src0_row_stride;                                                                \
-        size_t src1_row_stride = actx->src1_row_stride;                                                                \
-                                                                                                                       \
-        const uint32_t src0_nrows            = actx->src0_nrows;                                                       \
-        const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread;                                            \
-                                                                                                                       \
-        const uint32_t src0_start_row = src0_nrows_per_thread * ith;                                                   \
-        const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);                       \
-                                                                                                                       \
-        /* no work for this thread */                                                                                  \
-        if (src0_start_row >= src0_end_row) {                                                                          \
-            return;                                                                                                    \
-        }                                                                                                              \
-                                                                                                                       \
-        const uint8_t * restrict data_src0 = actx->data_src0;                                                          \
-        const uint8_t * restrict data_src1 = actx->data_src1;                                                          \
-        uint8_t * restrict data_dst        = actx->data_dst;                                                           \
-                                                                                                                       \
-        const size_t src0_row_size_aligned = actx->src0_row_size_aligned;                                              \
-        const size_t src1_row_size_aligned = actx->src1_row_size_aligned;                                              \
-        const size_t dst_row_size_aligned  = actx->dst_row_size_aligned;                                               \
-                                                                                                                       \
-        uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread);                 \
-        uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread);                 \
-        uint8_t * restrict dst_spad_data  = actx->vtcm_dst  + (ith * actx->vtcm_dst_size_per_thread);                  \
-                                                                                                                       \
-        size_t src0_spad_half_size = actx->src0_spad_half_size;                                                        \
-        size_t src1_spad_half_size = actx->src1_spad_half_size;                                                        \
-        size_t dst_spad_half_size  = actx->dst_spad_half_size;                                                         \
-                                                                                                                       \
-        const int BLOCK = actx->block;                                                                                 \
-        if (BLOCK == 0) {                                                                                              \
-            FARF(ERROR,                                                                                                \
-                 OP_STR                                                                                                \
-                 " : current VTCM reservation %zu is too small for even 1 row per thread, needed at least %zu\n",      \
-                 actx->vtcm_src0_size_per_thread, src0_row_size_aligned);                                              \
-            return;                                                                                                    \
-        }                                                                                                              \
-                                                                                                                       \
-        dma_queue * dma_queue = actx->octx->ctx->dma[ith];                                                             \
-                                                                                                                       \
-        /* See discussion: https://github.com/ggml-org/llama.cpp/pull/18151#issuecomment-3678235379 */                 \
-        for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) { \
-            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                 \
-                                                                                                                       \
-            /* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */                                   \
-            dma_queue_push_vtcm_to_ddr(dma_queue,                                                                      \
-                                       dma_make_ptr(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)),        \
-                                       dst_row_size, dst_row_size_aligned, 0);                                         \
-                                                                                                                       \
-            dma_queue_push(                                                                                            \
-                dma_queue,                                                                                             \
-                dma_make_ptr(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)),   \
-                src0_row_size_aligned, src0_row_stride, src0_row_size, block_size);                                    \
-            dma_queue_push(                                                                                            \
-                dma_queue,                                                                                             \
-                dma_make_ptr(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)),   \
-                src1_row_size_aligned, src1_row_stride, src1_row_size, block_size);                                    \
-        }                                                                                                              \
-                                                                                                                       \
-        for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) {                                           \
-            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                 \
-                                                                                                                       \
-            float * dst_spad  = (float *) dma_queue_pop(dma_queue).src;                                                \
-            float * src0_spad = (float *) dma_queue_pop(dma_queue).dst;                                                \
-            float * src1_spad = (float *) dma_queue_pop(dma_queue).dst;                                                \
-                                                                                                                       \
-            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                     \
-            CORE_EXPR;                                                                                                 \
-            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                      \
-                                                                                                                       \
-            dma_queue_push_vtcm_to_ddr(dma_queue, dma_make_ptr(data_dst + (ir * dst_row_size), dst_spad),              \
-                                       dst_row_size, dst_row_size_aligned, block_size);                                \
-                                                                                                                       \
-            /* prefetch N+2 loop iteration if any */                                                                   \
-            const uint32_t pref_block = (ir + BLOCK * 2);                                                              \
-            if (pref_block < src0_end_row) {                                                                           \
-                const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block);                                \
-                dma_queue_push(dma_queue, dma_make_ptr(src0_spad, data_src0 + (pref_block * src0_row_stride)),         \
-                               src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size);                \
-                dma_queue_push(dma_queue, dma_make_ptr(src1_spad, data_src1 + (pref_block * src1_row_stride)),         \
-                               src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size);                \
-            }                                                                                                          \
-        }                                                                                                              \
-                                                                                                                       \
-        dma_queue_flush(dma_queue);                                                                                    \
-                                                                                                                       \
+#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR)                                                                   \
+    static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) {                           \
+        struct htp_act_context * actx = (struct htp_act_context *) data;                                                 \
+        htp_act_preamble;                                                                                                \
+                                                                                                                         \
+        struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL;                            \
+                                                                                                                         \
+        size_t src0_row_size = actx->src0_row_size;                                                                      \
+        size_t src1_row_size = actx->src1_row_size;                                                                      \
+        size_t dst_row_size  = actx->dst_row_size;                                                                       \
+                                                                                                                         \
+        size_t src0_row_stride = actx->src0_row_stride;                                                                  \
+        size_t src1_row_stride = actx->src1_row_stride;                                                                  \
+                                                                                                                         \
+        const uint32_t src0_nrows            = actx->src0_nrows;                                                         \
+        const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread;                                              \
+                                                                                                                         \
+        const uint32_t src0_start_row = actx->row_start + src0_nrows_per_thread * ith;                                   \
+        const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, actx->row_start + src0_nrows);       \
+                                                                                                                         \
+        /* no work for this thread */                                                                                    \
+        if (src0_start_row >= src0_end_row) {                                                                            \
+            return;                                                                                                      \
+        }                                                                                                                \
+                                                                                                                         \
+        const uint8_t * restrict data_src0 = actx->data_src0;                                                            \
+        const uint8_t * restrict data_src1 = actx->data_src1;                                                            \
+        uint8_t * restrict data_dst        = actx->data_dst;                                                             \
+                                                                                                                         \
+        const size_t src0_row_size_aligned = actx->src0_row_size_aligned;                                                \
+        const size_t src1_row_size_aligned = actx->src1_row_size_aligned;                                                \
+        const size_t dst_row_size_aligned  = actx->dst_row_size_aligned;                                                 \
+                                                                                                                         \
+        uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread);                   \
+        uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread);                   \
+        uint8_t * restrict dst_spad_data  = actx->vtcm_dst  + (ith * actx->vtcm_dst_size_per_thread);                    \
+                                                                                                                         \
+        size_t src0_spad_half_size = actx->src0_spad_half_size;                                                          \
+        size_t src1_spad_half_size = actx->src1_spad_half_size;                                                          \
+        size_t dst_spad_half_size  = actx->dst_spad_half_size;                                                           \
+                                                                                                                         \
+        const int BLOCK = actx->block;                                                                                   \
+        if (BLOCK == 0) {                                                                                                \
+            FARF(ERROR,                                                                                                  \
+                 OP_STR                                                                                                  \
+                 " : current VTCM reservation %zu is too small for even 1 row per thread, needed at least %zu\n",        \
+                 actx->vtcm_src0_size_per_thread, src0_row_size_aligned);                                                \
+            return;                                                                                                      \
+        }                                                                                                                \
+                                                                                                                         \
+        dma_queue * dma_queue = actx->octx->ctx->dma[ith];                                                               \
+                                                                                                                         \
+        /* See discussion: https://github.com/ggml-org/llama.cpp/pull/18151#issuecomment-3678235379 */                   \
+        for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) {   \
+            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                   \
+                                                                                                                         \
+            /* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */                                     \
+            dma_queue_push_vtcm_to_ddr(dma_queue,                                                                        \
+                                       dma_make_ptr(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)),          \
+                                       dst_row_size, dst_row_size_aligned, 0);                                           \
+                                                                                                                         \
+            dma_queue_push(                                                                                              \
+                dma_queue,                                                                                               \
+                dma_make_ptr(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)),     \
+                src0_row_size_aligned, src0_row_stride, src0_row_size, block_size);                                      \
+            dma_queue_push(                                                                                              \
+                dma_queue,                                                                                               \
+                dma_make_ptr(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)),     \
+                src1_row_size_aligned, src1_row_stride, src1_row_size, block_size);                                      \
+        }                                                                                                                \
+                                                                                                                         \
+        for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) {                                             \
+            const uint32_t block_size = MIN(BLOCK, src0_end_row - ir);                                                   \
+                                                                                                                         \
+            float * dst_spad  = (float *) dma_queue_pop(dma_queue).src;                                                  \
+            float * src0_spad = (float *) dma_queue_pop(dma_queue).dst;                                                  \
+            float * src1_spad = (float *) dma_queue_pop(dma_queue).dst;                                                  \
+                                                                                                                         \
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                       \
+            CORE_EXPR;                                                                                                   \
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir);                                                        \
+                                                                                                                         \
+            dma_queue_push_vtcm_to_ddr(dma_queue, dma_make_ptr(data_dst + (ir * dst_row_size), dst_spad),                \
+                                       dst_row_size, dst_row_size_aligned, block_size);                                  \
+                                                                                                                         \
+            /* prefetch N+2 loop iteration if any */                                                                     \
+            const uint32_t pref_block = (ir + BLOCK * 2);                                                                \
+            if (pref_block < src0_end_row) {                                                                             \
+                const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block);                                  \
+                dma_queue_push(dma_queue, dma_make_ptr(src0_spad, data_src0 + (pref_block * src0_row_stride)),           \
+                               src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size);                  \
+                dma_queue_push(dma_queue, dma_make_ptr(src1_spad, data_src1 + (pref_block * src1_row_stride)),           \
+                               src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size);                  \
+            }                                                                                                            \
+        }                                                                                                                \
+                                                                                                                         \
+        dma_queue_flush(dma_queue);                                                                                      \
+                                                                                                                         \
     }

 DEFINE_GLU_PER_THREAD(swiglu, "swiglu-f32", swiglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
@@ -473,14 +473,30 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
     }

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads  = MIN(octx->n_threads, src0_nrows);
+    const size_t dst_row_size = dst->ne[0] * SIZEOF_FP32;
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;

     // row_size   = bytes of useful data per row (what the kernel touches / what DMA copies).
     // row_stride = bytes between successive rows in DDR (may exceed row_size for non-contig src).
-    const size_t nc_bytes    = dst->ne[0] * SIZEOF_FP32;
-    const size_t src0_row_size = nc_bytes;
-    const size_t src1_row_size = nc_bytes;
-    const size_t dst_row_size  = nc_bytes;
+    const size_t nc_bytes        = dst_row_size;
+    const size_t src0_row_size   = nc_bytes;
+    const size_t src1_row_size   = nc_bytes;
     const size_t src0_row_stride = src0->nb[1];
     const size_t src1_row_stride = src1 ? src1->nb[1] : src0->nb[1];

@@ -518,7 +534,7 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
     struct htp_act_context actx;
     actx.octx = octx;

-    actx.src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
+    actx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);

     actx.src0_row_size = src0_row_size;
     actx.src1_row_size = src1_row_size;
@@ -545,7 +561,8 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
     actx.dst_spad_half_size  = L.dst_bytes_per_thread / 2;

     actx.block = actx.src0_spad_half_size / actx.src0_row_size_aligned;
-    actx.src0_nrows = src0_nrows;
+    actx.src0_nrows = nrows;
+    actx.row_start  = row_start;

     actx.nc = dst->ne[0];

@@ -570,7 +587,7 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
     actx.data_src1 = data_src1;
     actx.data_dst  = (uint8_t *) dst->data;

-    worker_pool_run_func(octx->ctx->worker_pool, act_op_func, &actx, n_threads);
+    work_queue_run(octx->ctx->work_queue, act_op_func, &actx, n_threads);
     return HTP_STATUS_OK;
 }

diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.c b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
index d35f685a6..d6e7f0d10 100644
--- a/ggml/src/ggml-hexagon/htp/allreduce-ops.c
+++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.c
@@ -17,6 +17,7 @@
 #include "hex-dma.h"
 #include "hex-profile.h"
 #include "allreduce-ops.h"
+#include "htp-fence.h"

 struct htp_allreduce_context {
     struct htp_ops_context * octx;
@@ -242,7 +243,42 @@ DEFINE_ALLREDUCE_THREAD_DMA_2D(add_f32,       float,  hvx_add_f32_aaa, 1, 0)
 DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f16, __fp16, hvx_add_f16_aaa, 1, 1)
 DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f32, float,  hvx_add_f32_aaa, 1, 1)

+static int validate_allreduce(
+    struct htp_ops_context * octx,
+    const struct htp_allreduce_kernel_params * kparams,
+    uint32_t n_ranks
+) {
+    if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    if (kparams->vtcm_size_per_thread <= 0 || kparams->vtcm_size <= 0) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
+    const size_t n_vtcm_buffers = htp_allreduce_vtcm_buffer_count(
+        n_ranks, octx->n_threads, has_add, kparams->is_row_bcast != 0);
+    const size_t vtcm_size = n_vtcm_buffers * (size_t) kparams->vtcm_size_per_thread;
+    if (vtcm_size != (size_t) kparams->vtcm_size) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+    if (vtcm_size > octx->ctx->vtcm_size) {
+        return HTP_STATUS_VTCM_TOO_SMALL;
+    }
+
+    if (octx->dst->type != HTP_TYPE_F16 && octx->dst->type != HTP_TYPE_F32) {
+        return HTP_STATUS_NO_SUPPORT;
+    }
+
+    return HTP_STATUS_OK;
+}
+
 int op_allreduce(struct htp_ops_context * octx) {
+    if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
+        return HTP_STATUS_OK;
+    }
+
     const struct htp_allreduce_kernel_params * kparams = (const struct htp_allreduce_kernel_params *) octx->kernel_params;
     const struct htp_tensor * dst = octx->dst;

@@ -253,38 +289,53 @@ int op_allreduce(struct htp_ops_context * octx) {
         return HTP_STATUS_INVAL_PARAMS;
     }

-    if (dst->type != HTP_TYPE_F16 && dst->type != HTP_TYPE_F32) {
-        return HTP_STATUS_NO_SUPPORT;
+    const uint32_t fence_seq_entry = (uint32_t) octx->op_params[0];
+    const uint32_t fence_seq_exit  = (uint32_t) octx->op_params[1];
+
+    const struct htp_tensor * my_sync = octx->src[n_ranks + rank];
+    atomic_uint * my_fence = (atomic_uint *) (uintptr_t) my_sync->data;
+
+    const int status = validate_allreduce(octx, kparams, n_ranks);
+    if (status != HTP_STATUS_OK) {
+        if (status == HTP_STATUS_NO_SUPPORT) {
+            FARF(ERROR, "ggml-hex: allreduce unsupported type %d : rank %u\n", dst->type, rank);
+        }
+        htp_fence_write(my_fence, fence_seq_exit, status);
+        return status;
     }

+    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 uint32_t fence_seq_entry = (uint32_t) octx->op_params[0];
-    const uint32_t fence_seq_exit  = (uint32_t) octx->op_params[1];

     // 1. Entry Barrier: Synchronize all ranks before reading
     struct htp_thread_trace * tr0 = &octx->ctx->trace[0];
     htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);

-    const struct htp_tensor * my_sync = octx->src[n_ranks + rank];
-    atomic_uint * my_fence = (atomic_uint *) my_sync->data;
-
-    atomic_store(&my_fence[0], fence_seq_entry);
-    asm volatile ("syncht" : : : "memory");
-    Q6_dccleaninva_A((void *) my_fence);
+    htp_fence_write(my_fence, fence_seq_entry, octx->status);

     for (uint32_t j = 0; j < n_ranks; j++) {
         if (j == rank) continue;
         const struct htp_tensor * peer_sync = octx->src[n_ranks + j];
-        atomic_uint * peer_fence = (atomic_uint *) peer_sync->data;
+        atomic_uint * peer_fence = (atomic_uint *) (uintptr_t) peer_sync->data;
         uint64_t spins = 0;
         while (1) {
-            Q6_dccleaninva_A((void *) peer_fence);
-            uint32_t val = atomic_load(&peer_fence[0]);
-            if (val == fence_seq_entry || val == fence_seq_exit) {
+            uint32_t peer_seq;
+            uint32_t peer_status;
+            htp_fence_read(peer_fence, &peer_seq, &peer_status);
+            if ((int32_t)(peer_seq - fence_seq_entry) >= 0) {
+                if (peer_status > HTP_STATUS_OK) {
+                    FARF(ERROR, "ggml-hex: allreduce entry peer %u failed with status %u\n", j, peer_status);
+                    htp_fence_write(my_fence, fence_seq_exit, peer_status);
+                    htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);
+                    return peer_status;
+                }
                 break;
             }
             if (++spins > HTP_FENCE_TIMEOUT) {
-                FARF(ERROR, "ggml-hex: allreduce entry fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_entry);
+                FARF(ERROR, "ggml-hex: allreduce entry fence-wait TIMEOUT : rank %u waiting on %u fence %p seq 0x%x peer-seq 0x%x\n",
+                     rank, j, peer_fence, fence_seq_entry, peer_seq);
+                htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_INTERNAL_ERR);
+                htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);
                 return HTP_STATUS_INTERNAL_ERR;
             }
             hex_pause();
@@ -301,8 +352,6 @@ int op_allreduce(struct htp_ops_context * octx) {
         const uint32_t elems_per_thread     = (uint32_t) kparams->elems_per_thread;
         const uint32_t vtcm_size_per_thread = (uint32_t) kparams->vtcm_size_per_thread;

-        const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
-
         struct htp_allreduce_context actx;
         actx.octx                 = octx;
         actx.n_ranks              = n_ranks;
@@ -339,6 +388,8 @@ int op_allreduce(struct htp_ops_context * octx) {
                 }
                 break;
             default:
+                FARF(ERROR, "ggml-hex: allreduce unsupported kernel %d : rank %u\n", kparams->kernel_type, rank);
+                htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_NO_SUPPORT);
                 return HTP_STATUS_NO_SUPPORT;
         }

@@ -368,23 +419,31 @@ 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);

-    atomic_store(&my_fence[0], fence_seq_exit);
-    asm volatile ("syncht" : : : "memory");
-    Q6_dccleaninva_A((void *) my_fence);
+    htp_fence_write(my_fence, fence_seq_exit, octx->status);

     for (uint32_t j = 0; j < n_ranks; j++) {
         if (j == rank) continue;
         const struct htp_tensor * peer_sync = octx->src[n_ranks + j];
-        atomic_uint * peer_fence = (atomic_uint *) peer_sync->data;
+        atomic_uint * peer_fence = (atomic_uint *) (uintptr_t) peer_sync->data;
         uint64_t spins = 0;
         while (1) {
-            Q6_dccleaninva_A((void *) peer_fence);
-            uint32_t val = atomic_load(&peer_fence[0]);
-            if (val == fence_seq_exit) {
+            uint32_t peer_seq;
+            uint32_t peer_status;
+            htp_fence_read(peer_fence, &peer_seq, &peer_status);
+            if ((int32_t)(peer_seq - fence_seq_exit) >= 0) {
+                if (peer_status > HTP_STATUS_OK) {
+                    FARF(ERROR, "ggml-hex: allreduce exit peer %u failed with status %u\n", j, peer_status);
+                    htp_fence_write(my_fence, fence_seq_exit, peer_status);
+                    htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
+                    return peer_status;
+                }
                 break;
             }
             if (++spins > HTP_FENCE_TIMEOUT) {
-                FARF(ERROR, "ggml-hex: allreduce exit fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_exit);
+                FARF(ERROR, "ggml-hex: allreduce exit fence-wait TIMEOUT : rank %u waiting on %u fence %p seq 0x%x peer-seq 0x%x\n",
+                     rank, j, peer_fence, fence_seq_exit, peer_seq);
+                htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_INTERNAL_ERR);
+                htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
                 return HTP_STATUS_INTERNAL_ERR;
             }
             hex_pause();
@@ -394,5 +453,5 @@ int op_allreduce(struct htp_ops_context * octx) {

     htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);

-    return HTP_STATUS_OK;
+    return octx->status;
 }
diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.h b/ggml/src/ggml-hexagon/htp/allreduce-ops.h
index de447d87e..0aed2b8b7 100644
--- a/ggml/src/ggml-hexagon/htp/allreduce-ops.h
+++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.h
@@ -2,6 +2,8 @@
 #define ALLREDUCE_OPS_H

 #include <stdint.h>
+#include <stddef.h>
+#include <stdbool.h>

 #define HTP_ALLREDUCE_MAX_RANKS 4

@@ -15,6 +17,15 @@ enum htp_allreduce_kernel_type {
     HTP_ALLREDUCE_KERNEL_DMA_2D,
 };

+static inline size_t htp_allreduce_vtcm_buffer_count(
+    uint32_t n_ranks,
+    uint32_t n_threads,
+    bool has_add,
+    bool is_row_bcast
+) {
+    return (size_t) (n_ranks + 1) * n_threads + (has_add ? (is_row_bcast ? 1 : n_threads) : 0);
+}
+
 struct htp_allreduce_kernel_params {
     int32_t rank;
     int32_t n_ranks;
diff --git a/ggml/src/ggml-hexagon/htp/argsort-ops.c b/ggml/src/ggml-hexagon/htp/argsort-ops.c
index 774faef5f..e3c49e763 100644
--- a/ggml/src/ggml-hexagon/htp/argsort-ops.c
+++ b/ggml/src/ggml-hexagon/htp/argsort-ops.c
@@ -11,9 +11,10 @@
 #include "hvx-utils.h"
 #include "hex-dma.h"

+#include "hex-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "htp-tensor.h"

 #ifndef MIN
 #define MIN(a, b) ((a) < (b) ? (a) : (b))
@@ -22,6 +23,9 @@
 struct htp_argsort_context {
     struct htp_ops_context * octx;
     uint32_t                 nrows_per_thread;
+    uint32_t                 total_rows;
+    uint32_t                 row_start;
+    uint32_t                 row_end;
     uint8_t *                vtcm_base;
     size_t                   vtcm_per_thread;
 };
@@ -336,10 +340,9 @@ static void htp_argsort_f32_##ne00##_##order_name(unsigned int n, unsigned int i
     const struct htp_tensor * src0 = octx->src[0];                                                             \
     const struct htp_tensor * dst = octx->dst;                                                                 \
     uint8_t * spad = actx->vtcm_base + actx->vtcm_per_thread * i;                                              \
-    uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];                                             \
     uint32_t rows_per_thread = actx->nrows_per_thread;                                                         \
-    uint32_t start_row = rows_per_thread * i;                                                                  \
-    uint32_t end_row = MIN(start_row + rows_per_thread, total_rows);                                           \
+    uint32_t start_row = actx->row_start + rows_per_thread * i;                                                \
+    uint32_t end_row = MIN(start_row + rows_per_thread, actx->row_end);                                        \
     size_t values_size = hex_round_up(ne00 * sizeof(float), 128);                                              \
     float * values_buf = (float *) spad;                                                                       \
     int32_t * indices_buf = (int32_t *) (spad + values_size);                                                  \
@@ -386,9 +389,6 @@ static void htp_argsort_f32_fallback(unsigned int n, unsigned int i, void * data

     // Dimensions
     uint32_t ne00 = src0->ne[0];
-    uint32_t ne01 = src0->ne[1];
-    uint32_t ne02 = src0->ne[2];
-    uint32_t ne03 = src0->ne[3];

     uint32_t nb01 = src0->nb[1];

@@ -398,10 +398,9 @@ static void htp_argsort_f32_fallback(unsigned int n, unsigned int i, void * data
     enum ggml_sort_order order = (enum ggml_sort_order) octx->op_params[0];

     // Rows to process
-    uint32_t total_rows = ne01 * ne02 * ne03;
     uint32_t rows_per_thread = actx->nrows_per_thread;
-    uint32_t start_row = rows_per_thread * i;
-    uint32_t end_row = MIN(start_row + rows_per_thread, total_rows);
+    uint32_t start_row = actx->row_start + rows_per_thread * i;
+    uint32_t end_row = MIN(start_row + rows_per_thread, actx->row_end);

     size_t values_size = hex_round_up(ne00 * sizeof(float), 128);
     uint32_t num_vec_ind_values = hmx_ceil_div(ne00, VLEN/(sizeof(int32_t)));
@@ -451,8 +450,28 @@ int op_argsort(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    const uint32_t total_rows = octx->src[0]->ne[1] * octx->src[0]->ne[2] * octx->src[0]->ne[3];
-    const uint32_t n_threads = MIN(total_rows, octx->n_threads);
+    const struct htp_tensor * src0 = octx->src[0];
+    const struct htp_tensor * dst  = octx->dst;
+
+    const uint32_t total_rows  = src0->ne[1] * src0->ne[2] * src0->ne[3];
+    const size_t dst_row_size  = dst->ne[0]  * sizeof(int32_t);
+
+    uint32_t row_start = 0;
+    uint32_t row_end   = total_rows;
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, sizeof(int32_t), (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        row_end   = range.start + range.count;
+    }
+
+    const uint32_t nrows = row_end - row_start;
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;

     // Allocate scratchpad
     // We need 1 row of float + 1 row of int32 per thread.
@@ -478,7 +497,10 @@ int op_argsort(struct htp_ops_context * octx) {

     struct htp_argsort_context actx;
     actx.octx = octx;
-    actx.nrows_per_thread = (total_rows + n_threads - 1) / n_threads;
+    actx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+    actx.total_rows       = nrows;
+    actx.row_start        = row_start;
+    actx.row_end          = row_end;
     actx.vtcm_base = (uint8_t *) octx->ctx->vtcm_base;
     actx.vtcm_per_thread = spad_per_thread;

@@ -508,7 +530,7 @@ int op_argsort(struct htp_ops_context * octx) {
     }

     // Run jobs
-    worker_pool_run_func(octx->ctx->worker_pool, job_func, &actx, n_threads);
+    work_queue_run(octx->ctx->work_queue, job_func, &actx, n_threads);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/binary-ops.c b/ggml/src/ggml-hexagon/htp/binary-ops.c
index db6177963..bfa849e0e 100644
--- a/ggml/src/ggml-hexagon/htp/binary-ops.c
+++ b/ggml/src/ggml-hexagon/htp/binary-ops.c
@@ -13,9 +13,10 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
 #include "htp-tensor.h"

 #ifndef MIN
@@ -36,6 +37,8 @@ struct htp_binary_context {

     uint32_t block_max;
     uint32_t nrows_per_thread;
+    uint32_t total_rows;
+    uint32_t row_start;
     size_t   src0_row_size_aligned;
     size_t   src1_row_size_aligned;
     size_t   dst_row_size_aligned;
@@ -48,27 +51,27 @@ struct htp_binary_context {
     const struct htp_tensor * src0 = octx->src[0]; \
     const struct htp_tensor * src1 = octx->src[1]; \
     const struct htp_tensor * dst  = octx->dst;    \
-                                       \
-    const uint32_t ne00 = src0->ne[0]; \
-    const uint32_t ne01 = src0->ne[1]; \
-    const uint32_t ne02 = src0->ne[2]; \
-    const uint32_t ne03 = src0->ne[3]; \
-                                       \
-    const uint32_t ne10 = src1->ne[0]; \
-    const uint32_t ne11 = src1->ne[1]; \
-    const uint32_t ne12 = src1->ne[2]; \
-    const uint32_t ne13 = src1->ne[3]; \
-                                       \
-    const uint32_t nb01 = src0->nb[1]; \
-    const uint32_t nb02 = src0->nb[2]; \
-    const uint32_t nb03 = src0->nb[3]; \
-                                       \
-    const uint32_t nb11 = src1->nb[1]; \
-    const uint32_t nb12 = src1->nb[2]; \
-    const uint32_t nb13 = src1->nb[3]; \
-                                       \
-    const uint32_t nb1 = dst->nb[1];   \
-    const uint32_t nb2 = dst->nb[2];   \
+                                                   \
+    const uint32_t ne00 = src0->ne[0];             \
+    const uint32_t ne01 = src0->ne[1];             \
+    const uint32_t ne02 = src0->ne[2];             \
+    const uint32_t ne03 = src0->ne[3];             \
+                                                   \
+    const uint32_t ne10 = src1->ne[0];             \
+    const uint32_t ne11 = src1->ne[1];             \
+    const uint32_t ne12 = src1->ne[2];             \
+    const uint32_t ne13 = src1->ne[3];             \
+                                                   \
+    const uint32_t nb01 = src0->nb[1];             \
+    const uint32_t nb02 = src0->nb[2];             \
+    const uint32_t nb03 = src0->nb[3];             \
+                                                   \
+    const uint32_t nb11 = src1->nb[1];             \
+    const uint32_t nb12 = src1->nb[2];             \
+    const uint32_t nb13 = src1->nb[3];             \
+                                                   \
+    const uint32_t nb1 = dst->nb[1];               \
+    const uint32_t nb2 = dst->nb[2];               \
     const uint32_t nb3 = dst->nb[3];

 static inline uint32_t calc_block_size(struct htp_binary_context * bctx, uint32_t ir, uint32_t end_row, uint32_t ne01, uint32_t ne02) {
@@ -93,87 +96,87 @@ static inline uint32_t calc_block_size(struct htp_binary_context * bctx, uint32_
 }

 // Macro for scalar op switch
-#define COMPUTE_SCALAR_OP(DST, SRC, VAL, TYPE, N) \
-    if(TYPE == HTP_TYPE_F32) { \
-        switch (octx->op) { \
-            case HTP_OP_ADD: hvx_add_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
-            case HTP_OP_SUB: hvx_sub_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
-            case HTP_OP_MUL: hvx_mul_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
+#define COMPUTE_SCALAR_OP(DST, SRC, VAL, TYPE, N)                                               \
+    if(TYPE == HTP_TYPE_F32) {                                                                  \
+        switch (octx->op) {                                                                     \
+            case HTP_OP_ADD: hvx_add_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
+            case HTP_OP_SUB: hvx_sub_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
+            case HTP_OP_MUL: hvx_mul_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break;          \
             case HTP_OP_DIV: hvx_mul_scalar_f32_aa(DST, SRC, 1.0f / (*(float *)VAL), N); break; \
-            default: break; \
-        } \
-    } \
-    else { \
-        switch (octx->op) { \
-            case HTP_OP_ADD: hvx_add_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
-            case HTP_OP_SUB: hvx_sub_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
-            case HTP_OP_MUL: hvx_mul_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
-            case HTP_OP_DIV: hvx_div_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
-            default: break; \
-        } \
+            default: break;                                                                     \
+        }                                                                                       \
+    }                                                                                           \
+    else {                                                                                      \
+        switch (octx->op) {                                                                     \
+            case HTP_OP_ADD: hvx_add_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
+            case HTP_OP_SUB: hvx_sub_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
+            case HTP_OP_MUL: hvx_mul_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
+            case HTP_OP_DIV: hvx_div_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break;       \
+            default: break;                                                                     \
+        }                                                                                       \
     }

 // Macro for vector op switch (All Aligned)
-#define COMPUTE_VECTOR_OP_AAA(DST, SRC0, SRC1, TYPE, N) \
-    if(TYPE == HTP_TYPE_F32) { \
-        switch (octx->op) { \
+#define COMPUTE_VECTOR_OP_AAA(DST, SRC0, SRC1, TYPE, N)                  \
+    if(TYPE == HTP_TYPE_F32) {                                           \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f32_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f32_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f32_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f32_aaa(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
-    } \
-    else { \
-        switch (octx->op) { \
+            default: break;                                              \
+        }                                                                \
+    }                                                                    \
+    else {                                                               \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f16_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f16_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f16_aaa(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f16_aaa(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
+            default: break;                                              \
+        }                                                                \
     }

 // Macro for vector op switch (Dst Aligned, Src0 Aligned, Src1 Unaligned)
-#define COMPUTE_VECTOR_OP_AAU(DST, SRC0, SRC1, TYPE, N) \
-    if(TYPE == HTP_TYPE_F32) { \
-        switch (octx->op) { \
+#define COMPUTE_VECTOR_OP_AAU(DST, SRC0, SRC1, TYPE, N)                  \
+    if(TYPE == HTP_TYPE_F32) {                                           \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f32_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f32_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f32_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f32_aau(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
-    } \
-    else { \
-        switch (octx->op) { \
+            default: break;                                              \
+        }                                                                \
+    }                                                                    \
+    else {                                                               \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f16_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f16_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f16_aau(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f16_aau(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
+            default: break;                                              \
+        }                                                                \
     }

 // Macro for vector op switch (All Unaligned - generic loop used in element repeat)
-#define COMPUTE_VECTOR_OP_UUU(DST, SRC0, SRC1, TYPE, N) \
-    if(TYPE == HTP_TYPE_F32) { \
-        switch (octx->op) { \
+#define COMPUTE_VECTOR_OP_UUU(DST, SRC0, SRC1, TYPE, N)                  \
+    if(TYPE == HTP_TYPE_F32) {                                           \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f32_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f32_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f32_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f32_uuu(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
-    } \
-    else { \
-        switch (octx->op) { \
+            default: break;                                              \
+        }                                                                \
+    }                                                                    \
+    else {                                                               \
+        switch (octx->op) {                                              \
             case HTP_OP_ADD: hvx_add_f16_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_SUB: hvx_sub_f16_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_MUL: hvx_mul_f16_uuu(DST, SRC0, SRC1, N); break; \
             case HTP_OP_DIV: hvx_div_f16_uuu(DST, SRC0, SRC1, N); break; \
-            default: break; \
-        } \
+            default: break;                                              \
+        }                                                                \
     }

 // 1. Scalar src1 (ne10 == 1)
@@ -184,9 +187,8 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {

     const uint32_t src0_type = octx->src[0]->type;
     const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row = bctx->nrows_per_thread * ith;
-    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     FARF(HIGH, "binary-scalar: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -222,6 +224,8 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
     }

     // Main loop
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);

@@ -242,12 +246,14 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
         uint8_t * src1_ptr = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
         uint32_t s1_stride = (ne11 == 1) ? 0 : nb11;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
             uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
             COMPUTE_SCALAR_OP(r_dst, r_src0, src1_ptr, src0_type, ne00);
             src1_ptr += s1_stride;
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
         dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -266,6 +272,7 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -277,9 +284,8 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi

     const uint32_t src0_type = octx->src[0]->type;
     const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row = bctx->nrows_per_thread * ith;
-    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     FARF(HIGH, "binary-same-shape: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -323,18 +329,22 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
         uint8_t * d_spad  = (uint8_t *) dma_queue_pop(q).src;
         uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
         uint8_t * s1_spad = (uint8_t *) dma_queue_pop(q).dst;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
             uint8_t * r_src1 = s1_spad + r * bctx->src1_row_size_aligned;
             uint8_t * r_dst  = d_spad  + r * bctx->dst_row_size_aligned;
             COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint32_t i03, i02, i01, rem;
         i03 = fastdiv(ir, &bctx->src0_dim12_div);
@@ -366,6 +376,7 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -377,9 +388,8 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,

     const uint32_t src0_type  = octx->src[0]->type;
     const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row  = bctx->nrows_per_thread * ith;
-    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     FARF(HIGH, "binary-row-bcast: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -416,17 +426,21 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
         uint8_t * d_spad  = (uint8_t *) dma_queue_pop(q).src;
         uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
             uint8_t * r_src1 = (uint8_t *)s1_ptr; // Constant
             uint8_t * r_dst  = d_spad + r * bctx->dst_row_size_aligned;
             COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
         uint32_t rem = ir - i03 * (ne02 * ne01);
@@ -447,6 +461,7 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -458,9 +473,8 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *

     const uint32_t src0_type = octx->src[0]->type;
     const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row  = bctx->nrows_per_thread * ith;
-    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     FARF(HIGH, "binary-complex: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -493,6 +507,8 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
         uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -503,6 +519,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint32_t r_i01 = i01 + r;
             uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
@@ -516,6 +533,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
             // Read src1 from DDR (unaligned)
             COMPUTE_VECTOR_OP_AAU(r_dst, r_src0, r_src1, src0_type, ne00);
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
         dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -532,6 +550,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -544,9 +563,8 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
     const uint32_t src0_type = octx->src[0]->type;
     const uint32_t elem_size_bytes = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
     const uint32_t row_size_bytes = ne00 * elem_size_bytes;;
-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row  = bctx->nrows_per_thread * ith;
-    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row  = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row    = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
@@ -579,6 +597,8 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
         uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -589,6 +609,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint32_t r_i01 = i01 + r;
             uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
@@ -606,6 +627,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
                 COMPUTE_VECTOR_OP_UUU(r_dst + c * elem_size_bytes, r_src0 + c * elem_size_bytes, r_src1_row, src0_type, len);
             }
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
         dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -622,6 +644,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -650,9 +673,8 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t nb2 = dst->nb[2];
     const uint32_t nb3 = dst->nb[3];

-    const uint32_t total_rows = ne01 * ne02 * ne03;
-    const uint32_t start_row = bctx->nrows_per_thread * ith;
-    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, total_rows);
+    const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
+    const uint32_t end_row   = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
     if (start_row >= end_row) return;

     uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
@@ -683,6 +705,8 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = start_row; ir < end_row; ) {
         uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
         uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -693,6 +717,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
         uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
         uint32_t i01 = rem - i02 * ne01;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         for (uint32_t r = 0; r < current_block_size; r++) {
             uint32_t r_i01 = i01 + r; // linear within block since we split at ne01

@@ -704,6 +729,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {

             hvx_add_f32_aau(r_dst, r_src0, r_src1, ne00);
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
         dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, ne00 * sizeof(float), current_block_size);
@@ -720,6 +746,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
         }
         ir += current_block_size;
     }
+
     dma_queue_flush(q);
 }

@@ -729,15 +756,31 @@ static int execute_op_binary(struct htp_ops_context * octx) {
     const struct htp_tensor * dst  = octx->dst;

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads  = MIN(octx->n_threads, src0_nrows);

-    // Use packed row sizes for VTCM allocation
+    // Use packed row sizes for VTCM allocation and alignment
     const uint32_t src0_type = octx->src[0]->type;
     const size_t elem_size = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
     const size_t src0_row_size = src0->ne[0] * elem_size;
     const size_t src1_row_size = src1->ne[0] * elem_size;
     const size_t dst_row_size  = dst->ne[0]  * elem_size;

+    uint32_t row_start = 0;
+    uint32_t nrows     = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     size_t src0_row_size_aligned = hex_round_up(src0_row_size, VLEN);
     size_t src1_row_size_aligned = hex_round_up(src1_row_size, VLEN);
     size_t dst_row_size_aligned  = hex_round_up(dst_row_size,  VLEN);
@@ -815,7 +858,9 @@ static int execute_op_binary(struct htp_ops_context * octx) {

     struct htp_binary_context bctx;
     bctx.octx                  = octx;
-    bctx.nrows_per_thread      = (src0_nrows + n_threads - 1) / n_threads;
+    bctx.nrows_per_thread      = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+    bctx.total_rows            = nrows;
+    bctx.row_start             = row_start;
     bctx.block_max             = rows_per_buffer;
     bctx.src0_row_size_aligned = src0_row_size_aligned;
     bctx.src1_row_size_aligned = src1_row_size_aligned;
@@ -850,7 +895,7 @@ static int execute_op_binary(struct htp_ops_context * octx) {
         dma_queue_pop(q);
     }

-    worker_pool_run_func(octx->ctx->worker_pool, worker_func, &bctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, worker_func, &bctx, n_threads);

     return HTP_STATUS_OK;
 }
@@ -870,4 +915,3 @@ int op_binary(struct htp_ops_context * octx) {

     return HTP_STATUS_NO_SUPPORT;
 }
-
diff --git a/ggml/src/ggml-hexagon/htp/concat-ops.c b/ggml/src/ggml-hexagon/htp/concat-ops.c
index 51d39e8d9..966e867b3 100644
--- a/ggml/src/ggml-hexagon/htp/concat-ops.c
+++ b/ggml/src/ggml-hexagon/htp/concat-ops.c
@@ -1,5 +1,8 @@
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"
 #include "hexagon_types.h"
 #include "hexagon_protos.h"
 #include "hvx_hexagon_protos.h"
@@ -13,6 +16,10 @@ struct htp_concat_context {
     struct htp_ops_context * octx;
     uint32_t dim;
     uint32_t nrows_per_thread;
+    uint32_t row_start;
+    uint32_t nrows;
+    uint32_t elem_start;
+    uint32_t nelems;
     struct fastdiv_values div_ne0;
     struct fastdiv_values div_ne1;
     struct fastdiv_values div_ne2;
@@ -28,10 +35,10 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *

     const uint32_t src0_ne0 = src0->ne[0];
     const uint32_t src1_ne0 = src1->ne[0];
-    const uint32_t ne1      = dst->ne[1];

-    const uint32_t start_i = ith * cctx->nrows_per_thread;
-    const uint32_t end_i   = (start_i + cctx->nrows_per_thread < ne1) ? (start_i + cctx->nrows_per_thread) : ne1;
+    const uint32_t row_end = cctx->row_start + cctx->nrows;
+    const uint32_t start_i = cctx->row_start + ith * cctx->nrows_per_thread;
+    const uint32_t end_i   = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
     if (start_i >= end_i) return;

     dma_queue * q = octx->ctx->dma[ith];
@@ -51,6 +58,8 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
     const uint32_t spad0_row_bytes = hex_round_up((src0_ne0 + src1_ne0_padded) * sizeof(float), VLEN);
     uint32_t mu = src1_ne0_padded * spad1_stride;

+    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;

@@ -66,6 +75,7 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *

         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++) {
@@ -75,6 +85,7 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
                 hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
             }
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);

         dma_queue_pop(q); // src0

@@ -95,10 +106,10 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *

     const uint32_t src0_ne0 = src0->ne[0];
     const uint32_t src1_ne0 = src1->ne[0];
-    const uint32_t ne1      = dst->ne[1];

-    const uint32_t start_i = ith * cctx->nrows_per_thread;
-    const uint32_t end_i   = (start_i + cctx->nrows_per_thread < ne1) ? (start_i + cctx->nrows_per_thread) : ne1;
+    const uint32_t row_end = cctx->row_start + cctx->nrows;
+    const uint32_t start_i = cctx->row_start + ith * cctx->nrows_per_thread;
+    const uint32_t end_i   = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
     if (start_i >= end_i) return;

     dma_queue * q = octx->ctx->dma[ith];
@@ -118,6 +129,8 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
     const uint32_t spad0_row_bytes = hex_round_up((src0_ne0 + src1_ne0_padded) * sizeof(__fp16), VLEN);
     uint32_t mu = src1_ne0_padded * spad1_stride;

+    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;

@@ -133,6 +146,7 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *

         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++) {
@@ -142,6 +156,7 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
                 hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
             }
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);

         dma_queue_pop(q); // src0

@@ -164,11 +179,14 @@ static void concat_generic(unsigned int nth, unsigned int ith, void * data) {
     const uint32_t type_size = (dst->type == HTP_TYPE_F32 || dst->type == HTP_TYPE_I32) ? 4 : 2;

     const uint32_t ne[4] = {dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]};
-    const uint32_t total_elements = ne[0] * ne[1] * ne[2] * ne[3];
-    const uint32_t chunk_size = (total_elements + nth - 1) / nth;

-    const uint32_t start_idx = MIN(ith * chunk_size, total_elements);
-    const uint32_t end_idx   = MIN(start_idx + chunk_size, total_elements);
+    // Per-device element range aligned to prevent false sharing
+    const uint32_t elem_start = cctx->elem_start;
+    const uint32_t nelems     = cctx->nelems;
+    const uint32_t chunk_size = (nelems + nth - 1) / nth;
+
+    const uint32_t start_idx = MIN(elem_start + ith * chunk_size, elem_start + nelems);
+    const uint32_t end_idx   = MIN(start_idx + chunk_size, elem_start + nelems);

     // Naive scalar element-wise copy
     for (uint32_t idx = start_idx; idx < end_idx; idx++) {
@@ -236,13 +254,28 @@ 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) {
-        n_threads = MIN(dst->ne[1], n_threads);
-        if (n_threads < 1) {
-            n_threads = 1;
+        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;
+        uint32_t nrows     = total_rows;
+        if (octx->ctx->mdev.count > 1) {
+            uint32_t rows_per_chunk = 0;
+            htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_data_row_size, &rows_per_chunk);
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            row_start = range.start;
+            nrows     = range.count;
         }
+
+        if (nrows == 0) {
+            return HTP_STATUS_OK;
+        }
+
+        cctx.row_start = row_start;
+        cctx.nrows     = nrows;
+
         uint32_t block_i = (type_size == 4) ? 32 : 64;

-        cctx.nrows_per_thread = hmx_ceil_div(dst->ne[1], n_threads);
+        cctx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);

         // Allocate VTCM
         uint32_t spad1_stride = block_i * type_size;
@@ -270,8 +303,26 @@ int op_concat(struct htp_ops_context * octx) {
         } else {
             worker_func = concat_2d_f16_transposed;
         }
+    } else {
+        const uint32_t total_elements = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
+        uint32_t elem_start = 0;
+        uint32_t nelems     = total_elements;
+        if (octx->ctx->mdev.count > 1) {
+            const uint32_t elems_per_chunk = HEX_L2_LINE_SIZE / type_size;
+            const bool can_split = htp_tensor_mdev_data_aligned(dst) && htp_tensor_is_contiguous(dst, type_size) && !htp_tensor_is_permuted(dst);
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_elements, can_split ? elems_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            elem_start = range.start;
+            nelems     = range.count;
+        }
+
+        if (nelems == 0) {
+            return HTP_STATUS_OK;
+        }
+
+        cctx.elem_start = elem_start;
+        cctx.nelems     = nelems;
     }

-    worker_pool_run_func(octx->ctx->worker_pool, worker_func, &cctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, worker_func, &cctx, n_threads);
     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/cpy-ops.c b/ggml/src/ggml-hexagon/htp/cpy-ops.c
index b151b757f..7f01a8c1e 100644
--- a/ggml/src/ggml-hexagon/htp/cpy-ops.c
+++ b/ggml/src/ggml-hexagon/htp/cpy-ops.c
@@ -16,6 +16,7 @@
 #include "htp-ops.h"
 #include "hvx-utils.h"
 #include "htp-tensor.h"
+#include "htp-fence.h"

 struct htp_copy_context {
     struct htp_ops_context * octx;
@@ -29,7 +30,23 @@ struct htp_copy_context {
     uint32_t          src0_blocks_per_row;
     uint32_t          dst_blocks_per_row;

+    uint32_t          elem_start;
+    uint32_t          nelem;
+    uint32_t          elem_per_thread;
+
     uint32_t          src0_nrows_per_thread;
+    uint32_t          row_start;
+    uint32_t          nrows;
+
+    struct fastdiv_values div_ne01;
+    struct fastdiv_values div_ne02_ne01;
+
+    struct fastdiv_values div_ne0;
+    struct fastdiv_values div_ne1_ne0;
+    struct fastdiv_values div_ne2_ne1_ne0;
+    struct fastdiv_values div_ne00;
+    struct fastdiv_values div_ne01_ne00;
+    struct fastdiv_values div_ne02_ne01_ne00;
 };

 #define cpy_preamble                              \
@@ -54,131 +71,113 @@ struct htp_copy_context {
     const uint32_t  nb0 = dst->nb[0];             \
     const uint32_t  nb1 = dst->nb[1];             \
     const uint32_t  nb2 = dst->nb[2];             \
-    const uint32_t  nb3 = dst->nb[3];             \
-                                                  \
-    const uint32_t   nr = ne01;
-
-#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                                       \
-static void cpy_thread_##NAME##_sameshape(unsigned int nth, unsigned int ith, void * data) {                   \
-    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                           \
-    struct htp_ops_context * octx = ct->octx;                                                                  \
-    cpy_preamble;                                                                                              \
-    const uint32_t dr  = ct->src0_nrows_per_thread;                                                            \
-    const uint32_t ir0 = dr * ith;                                                                             \
-    const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;                                                    \
-    if (ir0 >= nr) return;                                                                                     \
-    for (uint32_t i03 = 0; i03 < ne03; i03++) {                                                                \
-        for (uint32_t i02 = 0; i02 < ne02; i02++) {                                                            \
-            _Pragma("unroll(4)")                                                                               \
-            for (uint32_t i01 = ir0; i01 < ir1; i01++) {                                                       \
-                uint8_t* dst_ptr  = (uint8_t*) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;                     \
-                uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;                    \
-                hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2);                                              \
-                hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE);                                               \
-            }                                                                                                  \
-        }                                                                                                      \
-    }                                                                                                          \
+    const uint32_t  nb3 = dst->nb[3];
+
+#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                       \
+static void cpy_thread_##NAME##_sameshape(unsigned int nth, unsigned int ith, void * data) {   \
+    struct htp_copy_context * ct = (struct htp_copy_context *) data;                           \
+    struct htp_ops_context * octx = ct->octx;                                                  \
+    cpy_preamble;                                                                              \
+    const uint32_t dr  = ct->src0_nrows_per_thread;                                            \
+    const uint32_t ir0 = ct->row_start + dr * ith;                                             \
+    const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);                             \
+    if (ir0 >= ir1) return;                                                                    \
+    const bool contiguous = (nb01 == ne00 * ELEM_SIZE) && (nb1 == nb01) &&                     \
+                            (nb02 == ne01 * nb01)      && (nb2 == nb02) &&                     \
+                            (nb03 == ne02 * nb02)      && (nb3 == nb03);                       \
+    const uint32_t ne02_ne01 = ne02 * ne01;                                                    \
+    uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);                                           \
+    uint32_t rem = ir0 - i03 * ne02_ne01;                                                      \
+    uint32_t i02 = fastdiv(rem, &ct->div_ne01);                                                \
+    uint32_t i01 = rem - i02 * ne01;                                                           \
+    uint8_t * dst_ptr  = (uint8_t *) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;               \
+    uint8_t * src0_ptr = (uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;              \
+    if (contiguous) {                                                                          \
+        hvx_copy_uu(dst_ptr, src0_ptr, (ir1 - ir0) * ne00, ELEM_SIZE);                         \
+        return;                                                                                \
+    }                                                                                          \
+    for (uint32_t r = ir0; r < ir1; r++) {                                                     \
+        hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2);                                      \
+        hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE);                                       \
+        dst_ptr  += nb1;                                                                       \
+        src0_ptr += nb01;                                                                      \
+        if (++i01 == ne01) {                                                                   \
+            i01 = 0;                                                                           \
+            if (++i02 == ne02) {                                                               \
+                i02 = 0;                                                                       \
+                i03++;                                                                         \
+            }                                                                                  \
+            dst_ptr  = (uint8_t *) dst->data  + i02*nb2  + i03*nb3;                            \
+            src0_ptr = (uint8_t *) src0->data + i02*nb02 + i03*nb03;                           \
+        }                                                                                      \
+    }                                                                                          \
 }

 DEFINE_CPY_SAMESHAPE(f32,  float, 4)
 DEFINE_CPY_SAMESHAPE(f16, __fp16, 2)

-#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                                         \
-static void cpy_thread_##NAME##_reshape(unsigned int nth, unsigned int ith, void * data) {                     \
-    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                           \
-    struct htp_ops_context * octx = ct->octx;                                                                  \
-    cpy_preamble;                                                                                              \
-    const uint32_t dr  = ct->src0_nrows_per_thread;                                                            \
-    const uint32_t ir0 = dr * ith;                                                                             \
-    const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;                                                    \
-    if (ir0 >= nr) return;                                                                                     \
-    const bool src0_contig = (nb00 == ELEM_SIZE)   &&                                                          \
-                             (nb01 == ne00 * nb00) &&                                                          \
-                             (nb02 == ne01 * nb01) &&                                                          \
-                             (nb03 == ne02 * nb02);                                                            \
-    const bool dst_contig  = (nb0  == ELEM_SIZE)   &&                                                          \
-                             (nb1  == ne0  * nb0)  &&                                                          \
-                             (nb2  == ne1  * nb1)  &&                                                          \
-                             (nb3  == ne2  * nb2);                                                             \
-    if (src0_contig && dst_contig) {                                                                           \
-        for (int64_t i03 = 0; i03 < ne03; i03++) {                                                             \
-            for (int64_t i02 = 0; i02 < ne02; i02++) {                                                         \
-                uint8_t * src_ptr = (uint8_t *) src0->data + i03*nb03 + i02*nb02 + ir0*nb01;                   \
-                uint32_t  flat    = ((i03*ne02 + i02)*ne01 + ir0) * ne00;                                      \
-                uint8_t * dst_ptr = (uint8_t *) dst->data  + flat * ELEM_SIZE;                                 \
-                hvx_copy_uu(dst_ptr, src_ptr, (ir1 - ir0) * ne00, ELEM_SIZE);                                  \
-            }                                                                                                  \
-        }                                                                                                      \
-        return;                                                                                                \
-    }                                                                                                          \
-    const bool reshape_flat_fast = (ne03 == 1 && ne2 == 1 && ne3 == 1) &&                                      \
-                                   (ne0 == ne00 * ne01) && (ne1 == ne02) &&                                    \
-                                   (nb00 == ELEM_SIZE) && (nb0 == ELEM_SIZE);                                  \
-    if (reshape_flat_fast) {                                                                                   \
-        for (uint32_t i02 = 0; i02 < ne02; i02++) {                                                            \
-            for (uint32_t i01 = ir0; i01 < ir1; i01++) {                                                       \
-                uint8_t * src0_ptr = (uint8_t *) src0->data + i01 * nb01 + i02 * nb02;                         \
-                uint8_t * dst_ptr  = (uint8_t *) dst->data  + i01 * ne00 * ELEM_SIZE + i02 * nb1;              \
-                hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE);                                               \
-            }                                                                                                  \
-        }                                                                                                      \
-        return;                                                                                                \
-    }                                                                                                          \
-    int64_t k10 = 0;                                                                                           \
-    int64_t i11 = 0;                                                                                           \
-    int64_t i12 = 0;                                                                                           \
-    int64_t i13 = 0;                                                                                           \
-    const int64_t nk00 = ct->src0_blocks_per_row;                                                              \
-    const int64_t nk0  = ct->dst_blocks_per_row;                                                               \
-    for (int64_t i03 = 0; i03 < ne03; i03++) {                                                                 \
-        for (int64_t i02 = 0; i02 < ne02; i02++) {                                                             \
-            k10 += nk00 * ir0;                                                                                 \
-            while (k10 >= nk0) {                                                                               \
-                k10 -= nk0;                                                                                    \
-                if (++i11 == ne1) {                                                                            \
-                    i11 = 0;                                                                                   \
-                    if (++i12 == ne2) {                                                                        \
-                        i12 = 0;                                                                               \
-                        if (++i13 == ne3) {                                                                    \
-                            i13 = 0;                                                                           \
-                        }                                                                                      \
-                    }                                                                                          \
-                }                                                                                              \
-            }                                                                                                  \
-            for (int64_t i01 = ir0; i01 < ir1; i01++) {                                                        \
-                for (int64_t k00 = 0; k00 < nk00; k00++) {                                                     \
-                    const char * src0_ptr = ((char *) src0->data + k00*nb00 + i01*nb01 + i02*nb02 + i03*nb03); \
-                          char * dst_ptr  = ((char *)  dst->data + k10*nb0  + i11*nb1  + i12*nb2  + i13*nb3);  \
-                    memcpy(dst_ptr, src0_ptr, ELEM_SIZE);                                                      \
-                    if (++k10 == nk0) {                                                                        \
-                        k10 = 0;                                                                               \
-                        if (++i11 == ne1) {                                                                    \
-                            i11 = 0;                                                                           \
-                            if (++i12 == ne2) {                                                                \
-                                i12 = 0;                                                                       \
-                                if (++i13 == ne3) {                                                            \
-                                    i13 = 0;                                                                   \
-                                }                                                                              \
-                            }                                                                                  \
-                        }                                                                                      \
-                    }                                                                                          \
-                }                                                                                              \
-            }                                                                                                  \
-            k10 += nk00 * (ne01 - ir1);                                                                        \
-            while (k10 >= nk0) {                                                                               \
-                k10 -= nk0;                                                                                    \
-                if (++i11 == ne1) {                                                                            \
-                    i11 = 0;                                                                                   \
-                    if (++i12 == ne2) {                                                                        \
-                        i12 = 0;                                                                               \
-                        if (++i13 == ne3) {                                                                    \
-                            i13 = 0;                                                                           \
-                        }                                                                                      \
-                    }                                                                                          \
-                }                                                                                              \
-            }                                                                                                  \
-        }                                                                                                      \
-    }                                                                                                          \
+#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE)                                               \
+static void cpy_thread_##NAME##_reshape(unsigned int nth, unsigned int ith, void * data) {           \
+    struct htp_copy_context * ct = (struct htp_copy_context *) data;                                 \
+    struct htp_ops_context * octx = ct->octx;                                                        \
+    cpy_preamble;                                                                                    \
+    const uint32_t th_nelem = ct->elem_per_thread;                                                   \
+    const uint32_t th_start = ct->elem_start + ith * th_nelem;                                       \
+    const uint32_t th_end   = MIN(th_start + th_nelem, ct->elem_start + ct->nelem);                  \
+    if (th_start >= th_end) return;                                                                  \
+                                                                                                     \
+    const uint32_t ne01_ne00      = ne01 * ne00;                                                     \
+    const uint32_t ne02_ne01_ne00 = ne02 * ne01_ne00;                                                \
+    const uint32_t ne1_ne0        = ne1 * ne0;                                                       \
+    const uint32_t ne2_ne1_ne0    = ne2 * ne1_ne0;                                                   \
+                                                                                                     \
+    uint32_t e = th_start;                                                                           \
+    uint32_t i13 = fastdiv(e, &ct->div_ne2_ne1_ne0);                                                 \
+    uint32_t rem = e - i13 * ne2_ne1_ne0;                                                            \
+    uint32_t i12 = fastdiv(rem, &ct->div_ne1_ne0);                                                   \
+    uint32_t rem2 = rem - i12 * ne1_ne0;                                                             \
+    uint32_t i11 = fastdiv(rem2, &ct->div_ne0);                                                      \
+    uint32_t i10 = rem2 - i11 * ne0;                                                                 \
+                                                                                                     \
+    uint32_t i03 = fastdiv(e, &ct->div_ne02_ne01_ne00);                                              \
+    uint32_t rem_s = e - i03 * ne02_ne01_ne00;                                                       \
+    uint32_t i02 = fastdiv(rem_s, &ct->div_ne01_ne00);                                               \
+    uint32_t rem2_s = rem_s - i02 * ne01_ne00;                                                       \
+    uint32_t i01 = fastdiv(rem2_s, &ct->div_ne00);                                                   \
+    uint32_t i00 = rem2_s - i01 * ne00;                                                              \
+                                                                                                     \
+    char * dst_ptr        = (char *)       dst->data  + i10*nb0  + i11*nb1  + i12*nb2  + i13*nb3;    \
+    const char * src0_ptr = (const char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03;   \
+                                                                                                     \
+    for (; e < th_end; e++) {                                                                        \
+        *((ELEM_TYPE *) dst_ptr) = *((const ELEM_TYPE *) src0_ptr);                                  \
+                                                                                                     \
+        dst_ptr += nb0;                                                                              \
+        if (++i10 == ne0) {                                                                          \
+            i10 = 0;                                                                                 \
+            if (++i11 == ne1) {                                                                      \
+                i11 = 0;                                                                             \
+                if (++i12 == ne2) {                                                                  \
+                    i12 = 0;                                                                         \
+                    i13++;                                                                           \
+                }                                                                                    \
+            }                                                                                        \
+            dst_ptr = (char *) dst->data + i11*nb1 + i12*nb2 + i13*nb3;                              \
+        }                                                                                            \
+                                                                                                     \
+        src0_ptr += nb00;                                                                            \
+        if (++i00 == ne00) {                                                                         \
+            i00 = 0;                                                                                 \
+            if (++i01 == ne01) {                                                                     \
+                i01 = 0;                                                                             \
+                if (++i02 == ne02) {                                                                 \
+                    i02 = 0;                                                                         \
+                    i03++;                                                                           \
+                }                                                                                    \
+            }                                                                                        \
+            src0_ptr = (const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;                   \
+        }                                                                                            \
+    }                                                                                                \
 }

 DEFINE_CPY_RESHAPE(f32,  float, 4)
@@ -189,22 +188,33 @@ static void cpy_thread_f16_f32_sameshape(unsigned int nth, unsigned int ith, voi
     struct htp_ops_context * octx = ct->octx;
     cpy_preamble;

-    // parallelize by src0 rows
     const uint32_t dr  = ct->src0_nrows_per_thread;
-    const uint32_t ir0 = dr * ith;
-    const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;
-    if (ir0 >= nr) return;
-
-    // copy by rows
-    for (uint32_t i03 = 0; i03 < ne03; i03++) {
-        for (uint32_t i02 = 0; i02 < ne02; i02++) {
-            #pragma unroll(2)
-            for (uint32_t i01 = ir0; i01 < ir1; i01++) {
-                uint8_t* dst_ptr  = (uint8_t*) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;
-                uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
-                hex_l2fetch(src0_ptr, ne00 * sizeof(float), nb01, 2);
-                hvx_copy_f16_f32_uu(dst_ptr, src0_ptr, ne00);
+    const uint32_t ir0 = ct->row_start + dr * ith;
+    const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
+    if (ir0 >= ir1) return;
+
+    const uint32_t ne02_ne01 = ne02 * ne01;
+    uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
+    uint32_t rem = ir0 - i03 * ne02_ne01;
+    uint32_t i02 = fastdiv(rem, &ct->div_ne01);
+    uint32_t i01 = rem - i02 * ne01;
+
+    uint8_t* dst_ptr  = (uint8_t*) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;
+    uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
+
+    for (uint32_t r = ir0; r < ir1; r++) {
+        hex_l2fetch(src0_ptr, ne00 * sizeof(float), nb01, 2);
+        hvx_copy_f16_f32_uu(dst_ptr, src0_ptr, ne00);
+        dst_ptr  += nb1;
+        src0_ptr += nb01;
+        if (++i01 == ne01) {
+            i01 = 0;
+            if (++i02 == ne02) {
+                i02 = 0;
+                i03++;
             }
+            dst_ptr  = (uint8_t*) dst->data  + i02*nb2  + i03*nb3;
+            src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
         }
     }
 }
@@ -214,22 +224,33 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi
     struct htp_ops_context * octx = ct->octx;
     cpy_preamble;

-    // parallelize by src0 rows
     const uint32_t dr  = ct->src0_nrows_per_thread;
-    const uint32_t ir0 = dr * ith;
-    const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;
-    if (ir0 >= nr) return;
-
-    // copy by rows
-    for (uint32_t i03 = 0; i03 < ne03; i03++) {
-        for (uint32_t i02 = 0; i02 < ne02; i02++) {
-            #pragma unroll(2)
-            for (uint32_t i01 = ir0; i01 < ir1; i01++) {
-                uint8_t* dst_ptr  = (uint8_t*) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;
-                uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
-                hex_l2fetch(src0_ptr, ne00 * sizeof(__fp16), nb01, 2);
-                hvx_copy_f32_f16_uu(dst_ptr, src0_ptr, ne00);
+    const uint32_t ir0 = ct->row_start + dr * ith;
+    const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
+    if (ir0 >= ir1) return;
+
+    const uint32_t ne02_ne01 = ne02 * ne01;
+    uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
+    uint32_t rem = ir0 - i03 * ne02_ne01;
+    uint32_t i02 = fastdiv(rem, &ct->div_ne01);
+    uint32_t i01 = rem - i02 * ne01;
+
+    uint8_t* dst_ptr  = (uint8_t*) dst->data  + i01*nb1  + i02*nb2  + i03*nb3;
+    uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
+
+    for (uint32_t r = ir0; r < ir1; r++) {
+        hex_l2fetch(src0_ptr, ne00 * sizeof(__fp16), nb01, 2);
+        hvx_copy_f32_f16_uu(dst_ptr, src0_ptr, ne00);
+        dst_ptr  += nb1;
+        src0_ptr += nb01;
+        if (++i01 == ne01) {
+            i01 = 0;
+            if (++i02 == ne02) {
+                i02 = 0;
+                i03++;
             }
+            dst_ptr  = (uint8_t*) dst->data  + i02*nb2  + i03*nb3;
+            src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
         }
     }
 }
@@ -250,15 +271,19 @@ static inline void cpy_dma_sametype_sameshape(
     dma_queue * q = octx->ctx->dma[0];

     if (contiguous_outer) {
-        dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03);
-        dma_queue_pop(q);
+        if (!dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03)) {
+            dma_queue_flush(q);
+            dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03);
+        }
+        dma_queue_flush(q);
         return;
     }

     for (uint32_t i03 = 0; i03 < ne03; i03++) {
         for (uint32_t i02 = 0; i02 < ne02; i02++) {
-            uint8_t* dst_ptr  = (uint8_t*) dst->data  + i02*nb2  + i03*nb3;
-            uint8_t* src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
+            uint8_t * dst_ptr  = (uint8_t *) dst->data  + i02 * nb2  + i03 * nb3;
+            uint8_t * src0_ptr = (uint8_t *) src0->data + i02 * nb02 + i03 * nb03;
+
             if (!dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01)) {
                 dma_queue_flush(q);
                 dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01);
@@ -269,10 +294,9 @@ static inline void cpy_dma_sametype_sameshape(
     dma_queue_flush(q);
 }

-int op_cpy(struct htp_ops_context * octx) {
+static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
     cpy_preamble;
-
-    const uint32_t n_threads = MIN(nr, octx->n_threads);
+    *use_dma = false;

     struct htp_copy_context ct;
     ct.octx = octx;
@@ -296,59 +320,117 @@ int op_cpy(struct htp_ops_context * octx) {
     }

     const bool sametype   = (src0->type == dst->type);
-    const bool transposed = (nb00 > nb01) || (nb0 > nb1);
+    const bool transposed = (nb00 > nb01) || (nb0 > nb1) ||
+                            (nb00 != ct.src0_type_size) || (nb0 != ct.dst_type_size) ||
+                            (nb01 < ne00 * ct.src0_type_size) || (nb1 < ne0 * ct.dst_type_size);
     const bool sameshape  = !transposed && (ne00 == ne0 && ne01 == ne1 && ne02 == ne2 && ne03 == ne3);

-    ct.src0_nrows_per_thread = (nr + n_threads - 1) / n_threads;
+    const uint32_t n_threads = octx->n_threads;

-    worker_callback_t copy_fun = NULL;
-    bool use_dma = false;
+    const bool dst_is_contiguous = htp_tensor_is_contiguous(dst, ct.dst_type_size);

-    if (sametype && sameshape) {
-        use_dma = true;
-    } else if (sameshape) {
-        /**/ if (dst->type == HTP_TYPE_F16 && src0->type == HTP_TYPE_F32)
-            copy_fun = cpy_thread_f16_f32_sameshape;
-        else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_F16)
-            copy_fun = cpy_thread_f32_f16_sameshape;
-        else
-            return HTP_STATUS_NO_SUPPORT;
-    } else if (sametype) {
-        if (src0->type == HTP_TYPE_F32) {
-            copy_fun = cpy_thread_f32_reshape;
+    if (sameshape) {
+        const uint32_t total_rows = ne01 * ne02 * ne03;
+        const uint32_t row_size   = ne00 * ct.dst_type_size;
+
+        ct.div_ne01      = init_fastdiv_values(ne01);
+        ct.div_ne02_ne01 = init_fastdiv_values(ne02 * ne01);
+
+        uint32_t row_start = 0;
+        uint32_t nrows     = total_rows;
+
+        if (octx->ctx->mdev.count > 1) {
+            const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
+            const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, can_split ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            row_start = range.start;
+            nrows     = range.count;
+        }
+
+        if (nrows == 0) {
+            return HTP_STATUS_OK;
+        }
+
+        ct.row_start = row_start;
+        ct.nrows     = nrows;
+        ct.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+
+        if (sametype && octx->ctx->mdev.count <= 1) {
+            *use_dma = true;
+            cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
         } else {
-            copy_fun = cpy_thread_f16_reshape;
+            work_queue_func_t copy_fun = NULL;
+            if (sametype) {
+                copy_fun = (src0->type == HTP_TYPE_F32) ? cpy_thread_f32_sameshape : cpy_thread_f16_sameshape;
+            } else if (dst->type == HTP_TYPE_F16 && src0->type == HTP_TYPE_F32) {
+                copy_fun = cpy_thread_f16_f32_sameshape;
+            } else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_F16) {
+                copy_fun = cpy_thread_f32_f16_sameshape;
+            } else {
+                return HTP_STATUS_NO_SUPPORT;
+            }
+            work_queue_run(octx->ctx->work_queue, copy_fun, &ct, n_threads);
+        }
+    } else if (sametype) {
+        const uint32_t total_elems = ne0 * ne1 * ne2 * ne3;
+        const uint32_t elems_per_line = (ct.dst_type_size == 4) ? 32 : 64;
+
+        ct.div_ne0            = init_fastdiv_values(ne0);
+        ct.div_ne1_ne0        = init_fastdiv_values(ne1 * ne0);
+        ct.div_ne2_ne1_ne0    = init_fastdiv_values(ne2 * ne1 * ne0);
+        ct.div_ne00           = init_fastdiv_values(ne00);
+        ct.div_ne01_ne00      = init_fastdiv_values(ne01 * ne00);
+        ct.div_ne02_ne01_ne00 = init_fastdiv_values(ne02 * ne01 * ne00);
+
+        uint32_t elem_start = 0;
+        uint32_t nelem      = total_elems;
+
+        if (octx->ctx->mdev.count > 1) {
+            const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_elems, can_split ? elems_per_line : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            elem_start = range.start;
+            nelem      = range.count;
+        }
+
+        if (nelem == 0) {
+            return HTP_STATUS_OK;
         }
+
+        ct.elem_start      = elem_start;
+        ct.nelem           = nelem;
+        ct.elem_per_thread = fastdiv(nelem + n_threads - 1, &octx->n_threads_div);
+
+        work_queue_func_t copy_fun = (src0->type == HTP_TYPE_F32) ? cpy_thread_f32_reshape : cpy_thread_f16_reshape;
+        work_queue_run(octx->ctx->work_queue, copy_fun, &ct, n_threads);
     } else {
         return HTP_STATUS_NO_SUPPORT;
     }

-    FARF(HIGH, "cpy-%s-%s: (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_dma=%d n_threads %u\n",
-         src0->type == HTP_TYPE_F32 ? "f32" : "f16", dst->type == HTP_TYPE_F32 ? "f32" : "f16",
-         ne00, ne01, ne02, ne03, ne0, ne1, ne2, ne3, use_dma, n_threads);
+    return HTP_STATUS_OK;
+}

-    if (use_dma) {
-        cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
-    } else {
-        worker_pool_run_func(octx->ctx->worker_pool, copy_fun, &ct, n_threads);
-    }
+int op_cpy(struct htp_ops_context * octx) {
+    bool use_dma = false;
+    int status = exec_cpy(octx, &use_dma);
+
+    htp_ops_context_set_status(octx, status);

-    const struct htp_tensor *sync = octx->src[1];
-    if (sync && (sync->flags & HTP_TENSOR_FENCE)) {
+    if (octx->op == HTP_OP_CPY_FENCE) {
         if (!use_dma) {
-            // htp_tensor_flush_all(octx->ctx, octx->dsts, 1);
-            qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE);
+            htp_flush_dirty_ranges(octx->ctx);
         }

-        atomic_uint * sync_fence = (atomic_uint *) sync->data;
-        const uint32_t seq = (uint32_t) octx->op_params[0];
+        htp_mdev_group_barrier(octx);

-        atomic_store(&sync_fence[0], seq);
-        asm volatile ("syncht" : : : "memory");
-        Q6_dccleaninva_A((void *) sync_fence);
+        if (octx->ctx->mdev.idx == 0) {
+            const struct htp_tensor * sync = octx->src[1];
+            const uint32_t seq = (uint32_t) octx->op_params[0];
+            atomic_uint * sync_fence = (atomic_uint *) (uintptr_t) sync->data;
+            htp_fence_write(sync_fence, seq, octx->status);

-        FARF(HIGH, "ggml-hex: sync-release : fence %p seq %u\n", sync_fence, seq);
+            FARF(HIGH, "ggml-hex: sync-release : fence %p seq 0x%x status %d\n", sync_fence, seq, octx->status);
+        }
     }

-    return HTP_STATUS_OK;
+    return octx->status;
 }
diff --git a/ggml/src/ggml-hexagon/htp/cumsum-ops.c b/ggml/src/ggml-hexagon/htp/cumsum-ops.c
index 2d45c39f2..971fa3bcc 100644
--- a/ggml/src/ggml-hexagon/htp/cumsum-ops.c
+++ b/ggml/src/ggml-hexagon/htp/cumsum-ops.c
@@ -7,6 +7,8 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
@@ -17,25 +19,25 @@
 #define htp_cumsum_tensors_preamble                         \
     const struct htp_tensor * restrict src0 = octx->src[0]; \
     const struct htp_tensor * restrict dst  = octx->dst;    \
-                                                     \
-    const uint32_t ne00 = src0->ne[0];               \
-    const uint32_t ne01 = src0->ne[1];               \
-    const uint32_t ne02 = src0->ne[2];               \
-    const uint32_t ne03 = src0->ne[3];               \
-                                                     \
-    const uint32_t ne0 = dst->ne[0];                 \
-    const uint32_t ne1 = dst->ne[1];                 \
-    const uint32_t ne2 = dst->ne[2];                 \
-    const uint32_t ne3 = dst->ne[3];                 \
-                                                     \
-    const uint32_t nb00 = src0->nb[0];               \
-    const uint32_t nb01 = src0->nb[1];               \
-    const uint32_t nb02 = src0->nb[2];               \
-    const uint32_t nb03 = src0->nb[3];               \
-                                                     \
-    const uint32_t nb0 = dst->nb[0];                 \
-    const uint32_t nb1 = dst->nb[1];                 \
-    const uint32_t nb2 = dst->nb[2];                 \
+                                                            \
+    const uint32_t ne00 = src0->ne[0];                      \
+    const uint32_t ne01 = src0->ne[1];                      \
+    const uint32_t ne02 = src0->ne[2];                      \
+    const uint32_t ne03 = src0->ne[3];                      \
+                                                            \
+    const uint32_t ne0 = dst->ne[0];                        \
+    const uint32_t ne1 = dst->ne[1];                        \
+    const uint32_t ne2 = dst->ne[2];                        \
+    const uint32_t ne3 = dst->ne[3];                        \
+                                                            \
+    const uint32_t nb00 = src0->nb[0];                      \
+    const uint32_t nb01 = src0->nb[1];                      \
+    const uint32_t nb02 = src0->nb[2];                      \
+    const uint32_t nb03 = src0->nb[3];                      \
+                                                            \
+    const uint32_t nb0 = dst->nb[0];                        \
+    const uint32_t nb1 = dst->nb[1];                        \
+    const uint32_t nb2 = dst->nb[2];                        \
     const uint32_t nb3 = dst->nb[3];

 struct htp_cumsum_context {
@@ -46,6 +48,7 @@ struct htp_cumsum_context {
     size_t          dst_row_size_aligned;
     uint32_t        rows_per_thread;
     uint32_t        total_rows;
+    uint32_t        row_start;
 };

 #define htp_cumsum_preamble                                                \
@@ -116,11 +119,8 @@ static inline void hvx_cumsum_row_f32(const float * restrict src, float * restri
 static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * data) {
     htp_cumsum_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
-    const uint32_t ir0 = cctx->rows_per_thread * ith;
-    const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->total_rows);
+    const uint32_t ir0 = cctx->row_start + cctx->rows_per_thread * ith;
+    const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->row_start + cctx->total_rows);

     if (ir0 >= ir1) {
         return;
@@ -149,11 +149,15 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
                                    src_row_size_aligned, src_row_size, 1);
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = ir0; ir < ir1; ir++) {
         float * dst_spad_row = (float *) dma_queue_pop(dma_queue).src;
         float * src_spad_row = (float *) dma_queue_pop(dma_queue).dst;

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         hvx_cumsum_row_f32(src_spad_row, dst_spad_row, ne00);
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         dma_queue_push_vtcm_to_ddr(dma_queue,
                                    dma_make_ptr(dst_data + (ir * dst_row_size), (uint8_t *) dst_spad_row),
@@ -168,12 +172,10 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
     }

     dma_queue_flush(dma_queue);
-    t2 = HAP_perf_get_qtimer_count();

-    FARF(HIGH, "cumsum-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
+    FARF(HIGH, "cumsum-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
-         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
 }

 // ---------------------------------------------------------------------------
@@ -183,14 +185,14 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
 static void cumsum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
     htp_cumsum_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     const uint8_t * src_data = (const uint8_t *) src0->data;
     uint8_t *       dst_data = (uint8_t *) dst->data;

-    const uint32_t ir0 = cctx->rows_per_thread * ith;
-    const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->total_rows);
+    const uint32_t ir0 = cctx->row_start + cctx->rows_per_thread * ith;
+    const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->row_start + cctx->total_rows);
+
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);

     for (uint32_t ir = ir0; ir < ir1; ir++) {
         const float * restrict src_row = (const float *) (src_data + ir * cctx->src_row_size);
@@ -198,12 +200,11 @@ static void cumsum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
         hvx_cumsum_row_f32(src_row, dst_row, ne00);
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);

-    FARF(HIGH, "cumsum-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
+    FARF(HIGH, "cumsum-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
-         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
 }

 int op_cumsum_f32(struct htp_ops_context * octx) {
@@ -214,8 +215,25 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

-    const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads  = MIN(octx->n_threads, total_rows);
+    const uint32_t total_rows      = src0->ne[1] * src0->ne[2] * src0->ne[3];
+    const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = total_rows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;

     const size_t src_row_size         = src0->nb[1];
     const size_t dst_row_size         = dst->nb[1];
@@ -240,14 +258,15 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
         .dst_row_size         = dst_row_size,
         .src_row_size_aligned = src_row_size_aligned,
         .dst_row_size_aligned = dst_row_size_aligned,
-        .rows_per_thread      = (total_rows + n_threads - 1) / n_threads,
-        .total_rows           = total_rows,
+        .rows_per_thread      = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .total_rows           = nrows,
+        .row_start            = row_start,
     };

     if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
-        worker_pool_run_func(octx->ctx->worker_pool, cumsum_thread_f32, &cctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, cumsum_thread_f32, &cctx, n_threads);
     } else {
-        worker_pool_run_func(octx->ctx->worker_pool, cumsum_thread_f32_dma, &cctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, cumsum_thread_f32_dma, &cctx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/diag-ops.c b/ggml/src/ggml-hexagon/htp/diag-ops.c
index 9b3194d90..a69fd89d3 100644
--- a/ggml/src/ggml-hexagon/htp/diag-ops.c
+++ b/ggml/src/ggml-hexagon/htp/diag-ops.c
@@ -5,8 +5,11 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"
 #include "hvx-types.h"
 #include "hex-utils.h"
 #include "hvx-copy.h"
@@ -15,17 +18,17 @@
 #define htp_diag_tensors_preamble                           \
     const struct htp_tensor * restrict src0 = octx->src[0]; \
     const struct htp_tensor * restrict dst  = octx->dst;    \
-                                                     \
-    const uint32_t ne02 = src0->ne[2];               \
-                                                     \
-    const uint32_t ne0 = dst->ne[0];                 \
-    const uint32_t ne1 = dst->ne[1];                 \
-                                                     \
-    const uint32_t nb02 = src0->nb[2];               \
-    const uint32_t nb03 = src0->nb[3];               \
-                                                     \
-    const uint32_t nb1 = dst->nb[1];                 \
-    const uint32_t nb2 = dst->nb[2];                 \
+                                                            \
+    const uint32_t ne02 = src0->ne[2];                      \
+                                                            \
+    const uint32_t ne0 = dst->ne[0];                        \
+    const uint32_t ne1 = dst->ne[1];                        \
+                                                            \
+    const uint32_t nb02 = src0->nb[2];                      \
+    const uint32_t nb03 = src0->nb[3];                      \
+                                                            \
+    const uint32_t nb1 = dst->nb[1];                        \
+    const uint32_t nb2 = dst->nb[2];                        \
     const uint32_t nb3 = dst->nb[3];

 struct htp_diag_context {
@@ -36,6 +39,7 @@ struct htp_diag_context {
     size_t          dst_row_size_aligned;
     uint32_t        batches_per_thread;
     uint32_t        total_batches;
+    uint32_t        batch_start;
 };

 #define htp_diag_preamble                                              \
@@ -57,11 +61,8 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
     htp_diag_preamble;
     dma_queue * dma_queue = octx->ctx->dma[ith];

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
-    const uint32_t ib0 = dctx->batches_per_thread * ith;
-    const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->total_batches);
+    const uint32_t ib0 = dctx->batch_start + dctx->batches_per_thread * ith;
+    const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->batch_start + dctx->total_batches);

     if (ib0 >= ib1) {
         return;
@@ -79,6 +80,8 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
     uint8_t * src_spad = octx->src0_spad.data + (ith * src_batch_size_aligned);
     uint8_t * dst_spad = octx->dst_spad.data  + (ith * dst_row_size_aligned);

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ib = ib0; ib < ib1; ib++) {
         const uint32_t i3 = ib / ne02;
         const uint32_t i2 = ib % ne02;
@@ -96,7 +99,9 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)

         for (uint32_t i1 = 0; i1 < ne1; i1++) {
             // Compute row in VTCM
+            htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (ib * ne1 + i1));
             hvx_diag_row_f32(src_spad_f32, dst_spad_f32, i1, ne0);
+            htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (ib * ne1 + i1));

             // Write completed row back to DDR
             uint8_t * dst_row = dst_data + i3 * nb3 + i2 * nb2 + i1 * nb1;
@@ -107,12 +112,9 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
-
-    FARF(HIGH, "diag-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
+    FARF(HIGH, "diag-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ib0, ib1,
-         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
 }

 // ---------------------------------------------------------------------------
@@ -122,14 +124,14 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
 static void diag_thread_f32(unsigned int nth, unsigned int ith, void * data) {
     htp_diag_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     const uint8_t * src_data = (const uint8_t *) src0->data;
     uint8_t *       dst_data = (uint8_t *) dst->data;

-    const uint32_t ib0 = dctx->batches_per_thread * ith;
-    const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->total_batches);
+    const uint32_t ib0 = dctx->batch_start + dctx->batches_per_thread * ith;
+    const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->batch_start + dctx->total_batches);
+
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ib0);

     for (uint32_t ib = ib0; ib < ib1; ib++) {
         const uint32_t i3 = ib / ne02;
@@ -143,12 +145,11 @@ static void diag_thread_f32(unsigned int nth, unsigned int ith, void * data) {
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ib0);

-    FARF(HIGH, "diag-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
+    FARF(HIGH, "diag-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ib0, ib1,
-         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
 }

 int op_diag_f32(struct htp_ops_context * octx) {
@@ -160,7 +161,36 @@ int op_diag_f32(struct htp_ops_context * octx) {
     }

     const uint32_t total_batches = src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads     = MIN(octx->n_threads, total_batches);
+    const size_t dst_batch_size  = dst->ne[1] * dst->nb[1];
+
+    uint32_t batch_start = 0;
+    uint32_t nbatches    = total_batches;
+
+    if (octx->ctx->mdev.count > 1) {
+        bool can_split = htp_tensor_mdev_data_aligned(dst) && (dst->ne[0] == 1 || dst->nb[0] == sizeof(float)) && !htp_tensor_is_permuted(dst);
+        uint32_t batches_per_chunk = 1;
+        if (can_split) {
+            if (dst->ne[2] > 1 && (dst->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 &&
+                (dst->ne[3] <= 1 || (dst->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0)) {
+                batches_per_chunk = 1;
+            } else if (dst->nb[2] == dst_batch_size &&
+                       (dst->ne[3] <= 1 || dst->nb[3] == dst->nb[2] * dst->ne[2])) {
+                batches_per_chunk = (dst_batch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(dst_batch_size, HEX_L2_LINE_SIZE)) : 1;
+            } else {
+                can_split = false;
+            }
+        }
+
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_batches, can_split ? batches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        batch_start = range.start;
+        nbatches    = range.count;
+    }
+
+    if (nbatches == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads     = octx->n_threads;

     const size_t src_batch_size         = src0->ne[0] * sizeof(float);
     const size_t dst_row_size           = dst->ne[0] * sizeof(float);
@@ -185,14 +215,15 @@ int op_diag_f32(struct htp_ops_context * octx) {
         .dst_row_size           = dst_row_size,
         .src_batch_size_aligned = src_batch_size_aligned,
         .dst_row_size_aligned   = dst_row_size_aligned,
-        .batches_per_thread     = (total_batches + n_threads - 1) / n_threads,
-        .total_batches          = total_batches,
+        .batches_per_thread     = fastdiv(nbatches + n_threads - 1, &octx->n_threads_div),
+        .total_batches          = nbatches,
+        .batch_start            = batch_start,
     };

     if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
-        worker_pool_run_func(octx->ctx->worker_pool, diag_thread_f32, &dctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, diag_thread_f32, &dctx, n_threads);
     } else {
-        worker_pool_run_func(octx->ctx->worker_pool, diag_thread_f32_dma, &dctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, diag_thread_f32_dma, &dctx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/fill-ops.c b/ggml/src/ggml-hexagon/htp/fill-ops.c
index 3ccfbe74e..1f6eaafad 100644
--- a/ggml/src/ggml-hexagon/htp/fill-ops.c
+++ b/ggml/src/ggml-hexagon/htp/fill-ops.c
@@ -3,10 +3,11 @@
 #pragma clang diagnostic ignored "-Wunused-but-set-variable"

 #include <HAP_farf.h>
-#include <HAP_perf.h>
-
 #include <string.h>

+#include "hex-common.h"
+#include "hex-profile.h"
+
 #include "hvx-copy.h"
 #include "hvx-utils.h"

@@ -14,28 +15,30 @@
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"

 // ggml op_params layout for FILL:
 //   op_params[0] (as float) - the scalar fill value

-#define fill_preamble \
+#define fill_preamble                          \
     const struct htp_tensor * dst = octx->dst; \
-    \
-    const uint32_t ne0 = dst->ne[0]; \
-    const uint32_t ne1 = dst->ne[1]; \
-    const uint32_t ne2 = dst->ne[2]; \
-    const uint32_t ne3 = dst->ne[3]; \
-    \
-    const uint32_t nb1 = dst->nb[1]; \
-    const uint32_t nb2 = dst->nb[2]; \
-    const uint32_t nb3 = dst->nb[3]; \
-    \
+                                               \
+    const uint32_t ne0 = dst->ne[0];           \
+    const uint32_t ne1 = dst->ne[1];           \
+    const uint32_t ne2 = dst->ne[2];           \
+    const uint32_t ne3 = dst->ne[3];           \
+                                               \
+    const uint32_t nb1 = dst->nb[1];           \
+    const uint32_t nb2 = dst->nb[2];           \
+    const uint32_t nb3 = dst->nb[3];           \
+                                               \
     const uint32_t nr = ne1 * ne2 * ne3;

 struct htp_fill_context {
     struct htp_ops_context * octx;
     uint32_t nrows_per_thread;
     uint32_t total_rows;  // ne1 * ne2 * ne3
+    uint32_t row_start;
     bool     opt_path;
     HVX_Vector splat_vec;
     uint32_t   elem_size;
@@ -47,10 +50,15 @@ static void fill_thread(unsigned int nth, unsigned int ith, void * data) {
     fill_preamble;

     // Parallelise over the flat row index spanning ne1*ne2*ne3
-    const uint32_t ir0 = fctx->nrows_per_thread * ith;
-    const uint32_t ir1 = MIN(ir0 + fctx->nrows_per_thread, fctx->total_rows);
+    const uint32_t ir0 = fctx->row_start + fctx->nrows_per_thread * ith;
+    const uint32_t ir1 = MIN(ir0 + fctx->nrows_per_thread, fctx->row_start + fctx->total_rows);

-    uint64_t t1 = HAP_perf_get_qtimer_count();
+    if (ir0 >= ir1) {
+        return;
+    }
+
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);

     if (fctx->opt_path) {
         // Opt path: tensor is fully contiguous, treat as flat array
@@ -69,9 +77,8 @@ static void fill_thread(unsigned int nth, unsigned int ith, void * data) {
         }
     }

-    uint64_t t2 = HAP_perf_get_qtimer_count();
-    FARF(HIGH, "fill %u/%u: rows %u:%u usec %u\n",
-         ith, nth, ir0, ir1, (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir1);
+    FARF(HIGH, "fill %u/%u: rows %u:%u\n", ith, nth, ir0, ir1);
 }

 int op_fill(struct htp_ops_context * octx) {
@@ -85,8 +92,23 @@ int op_fill(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

+    uint32_t row_start = 0;
+    uint32_t nrows     = nr;
+
+    if (octx->ctx->mdev.count > 1) {
+        const uint32_t row_size = nb1;
+        const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(nr, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
     // nr = ne1*ne2*ne3 (flat row count across all outer dims); parallelise over it.
-    const uint32_t n_threads = MIN(nr, octx->n_threads);
+    const uint32_t n_threads = octx->n_threads;

     // Optimize if fully contiguous: skip stride arithmetic, treat as flat array
     const bool opt_path = (nb2 == nb1 * ne1) && (nb3 == nb2 * ne2);
@@ -99,8 +121,9 @@ int op_fill(struct htp_ops_context * octx) {

     struct htp_fill_context fctx = {
         .octx             = octx,
-        .nrows_per_thread = (nr + n_threads - 1) / n_threads,
-        .total_rows       = nr,
+        .nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .total_rows       = nrows,
+        .row_start        = row_start,
         .opt_path         = opt_path,
     };

@@ -117,7 +140,7 @@ int op_fill(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    worker_pool_run_func(octx->ctx->worker_pool, fill_thread, &fctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, fill_thread, &fctx, n_threads);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
index c76b4d3a3..8a1caba22 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c
@@ -5,7 +5,6 @@
 #include <assert.h>
 #include <HAP_compute_res.h>
 #include <HAP_farf.h>
-#include <HAP_perf.h>
 #include <math.h>
 #include <stdbool.h>
 #include <stdatomic.h>
@@ -75,6 +74,7 @@ struct htp_fa_context {

     uint32_t qrows;
     uint32_t qrows_per_thread;
+    uint32_t qrow_start;

     bool is_q_fp32;

@@ -89,8 +89,6 @@ struct htp_fa_context {

     const struct htp_tensor * k;
     const struct htp_tensor * v;
-
-    uint64_t t_start;
 };

 struct hmx_fa_context {
@@ -206,10 +204,9 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
     const uint32_t nb3 = dst->nb[3];

     // total rows in q
-    const uint32_t nr = factx->qrows;
-    const uint32_t dr = factx->qrows_per_thread;
-    const uint32_t ir0 = dr * ith;
-    const uint32_t ir1 = MIN(ir0 + dr, nr);
+    const uint32_t dr  = factx->qrows_per_thread;
+    const uint32_t ir0 = factx->qrow_start + dr * ith;
+    const uint32_t ir1 = MIN(ir0 + dr, factx->qrow_start + factx->qrows);

     if (ir0 >= ir1) return;

@@ -1888,6 +1885,24 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     const uint32_t n_threads = factx.n_threads;
     const uint32_t G = factx.G;

+    // Multi-device: split Q blocks across devices
+    const uint32_t n_q_blocks = (neq1 + Br - 1) / Br;
+    uint32_t q_start_min = 0;
+    uint32_t q_start_max = neq1;
+
+    if (octx->ctx->mdev.count > 1) {
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, htp_tensor_mdev_data_aligned(dst) ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        const uint32_t block_start = range.start;
+        const uint32_t block_end   = range.start + range.count;
+
+        if (block_start >= block_end) {
+            return HTP_STATUS_OK;
+        }
+
+        q_start_min = block_start * Br;
+        q_start_max = MIN(block_end * Br, neq1);
+    }
+
     // ======== VTCM allocation (GQA-aware) ========
     // K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used
     // throughout the KV loop below.
@@ -1977,7 +1992,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
     // ======== Main loop ========
     for (uint32_t ib3 = 0; ib3 < neq3; ++ib3) {
         const uint32_t im3 = mask ? fastmodulo(ib3, mask->ne[3], &factx.src3_div3) : 0;
-        for (uint32_t q_start = 0; q_start < neq1; q_start += Br) {
+        for (uint32_t q_start = q_start_min; q_start < q_start_max; q_start += Br) {
             const uint32_t n_rows_q    = hex_smin(Br, neq1 - q_start);
             const size_t   n_rows_g    = n_rows_q * G;
             const size_t   g_br_actual = hex_align_up(n_rows_g, HMX_FP16_TILE_N_ROWS);
@@ -1991,8 +2006,9 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {

                 // 1. Push Q and KV DMAs for the very first iteration.
                 // Subsequent iterations are enqueued early at the end of the previous iteration.
-                if (ib3 == 0 && q_start == 0 && kv_head == 0) {
-                    const uint8_t * q_ptr = (const uint8_t *) q->data;
+                if (ib3 == 0 && q_start == q_start_min && kv_head == 0) {
+                    const uint8_t * q_ptr = (const uint8_t *) q->data + q_start * q->nb[1] +
+                                            (kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3];
                     const size_t q_row_bytes = q_transposed ? n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed;
                     const size_t n_rows      = q_transposed ? factx.G : n_rows_q;
                     dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, q_ptr), q_row_bytes, hex_smax(q_src_stride, q_row_bytes), q_row_bytes, n_rows);
@@ -2311,8 +2327,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
                 if (next_kv_head >= n_kv_heads) {
                     next_kv_head = 0;
                     next_q_start = q_start + Br;
-                    if (next_q_start >= neq1) {
-                        next_q_start = 0;
+                    if (next_q_start >= q_start_max) {
+                        next_q_start = q_start_min;
                         next_ib3     = ib3 + 1;
                     }
                 }
@@ -2398,6 +2414,10 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
     if (kparams->kernel_type == HTP_FA_KERNEL_HMX) {
         return hmx_flash_attn_ext(octx);
     }
@@ -2407,8 +2427,6 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     factx.k = k;
     factx.v = v;

-    factx.t_start = HAP_perf_get_qtimer_count();
-
     factx.src0_div21 = kparams->u.hvx.src0_div21;
     factx.src0_div1  = kparams->u.hvx.src0_div1;

@@ -2451,8 +2469,30 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
     }

     // total rows in q
-    factx.qrows = kparams->qrows;
-    factx.qrows_per_thread = kparams->qrows_per_thread;
+    const uint32_t neq1 = q->ne[1];
+    const uint32_t neq2 = q->ne[2];
+    const uint32_t neq3 = q->ne[3];
+    const uint32_t total_qrows = neq1 * neq2 * neq3;
+
+    uint32_t qrow_start = 0;
+    uint32_t qrows      = total_qrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        qrow_start = range.start;
+        qrows      = range.count;
+    }
+
+    if (qrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
+    factx.qrows            = qrows;
+    factx.qrow_start       = qrow_start;
+    factx.qrows_per_thread = fastdiv(qrows + n_threads - 1, &octx->n_threads_div);

     size_t size_vkq_acc = hex_round_up(v->ne[0] * sizeof(float), 128); // VKQ32

@@ -2461,18 +2501,18 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {

     uint8_t * vtcm_cur = octx->ctx->vtcm_base;

-    factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * octx->n_threads);
-    factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * octx->n_threads);
-    factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * octx->n_threads);
-    factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * octx->n_threads);
-    factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * octx->n_threads);
+    factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * n_threads);
+    factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * n_threads);
+    factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * n_threads);
+    factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * n_threads);
+    factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * n_threads);

     if ((size_t) (vtcm_cur - octx->ctx->vtcm_base) > octx->ctx->vtcm_size) {
         return HTP_STATUS_VTCM_TOO_SMALL;
     }

     if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, octx->n_threads);
+        work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
index c4d190631..027845411 100644
--- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
+++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h
@@ -51,6 +51,7 @@ struct htp_fa_kernel_params {

     uint32_t qrows;
     uint32_t qrows_per_thread;
+    uint32_t qrow_start;
     float    m0;
     float    m1;
     uint32_t n_head_log2;
diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
index 966552152..0b6529571 100644
--- a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
+++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
@@ -4,10 +4,13 @@

 #include "hvx-utils.h"
 #include "hex-fastdiv.h"
+#include "hex-common.h"
+#include "hex-profile.h"

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "htp-ctx.h"
+#include "htp-tensor.h"

 #ifndef MIN
 #define MIN(a, b) ((a) < (b) ? (a) : (b))
@@ -22,6 +25,8 @@ struct htp_gdn_context {
     size_t   state_bytes;
     uint8_t * vtcm_base;
     size_t   vtcm_per_thread;
+    uint32_t row_start;
+    uint32_t nrows;
 };

 static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const float * restrict mul, const float * restrict dot, uint32_t n) {
@@ -586,8 +591,9 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
     const uint32_t n_seqs   = v->ne[3];
     const uint32_t K        = octx->op_params[0];

-    const uint32_t total_rows = H * n_seqs;
-    if (ith >= total_rows) {
+    const uint32_t row_end = gctx->row_start + gctx->nrows;
+
+    if (ith >= gctx->nrows) {
         return;
     }

@@ -621,11 +627,11 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
     const uint64_t state_seq_stride = state->nb[3] / sizeof(float);
     const uint64_t state_size_per_snap = (uint64_t) S_v * S_v * H * n_seqs;

-    uint32_t ir_prefetch = ith;
+    uint32_t ir_prefetch = gctx->row_start + ith;
     int spad_idx = 0;

     // Prefetch preamble (up to 2 steps)
-    for (int k = 0; k < 2 && ir_prefetch < total_rows; k++) {
+    for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
         const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
         const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
         const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -646,8 +652,11 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));
+
     int curr_spad_idx = 0;
-    for (uint32_t ir = ith; ir < total_rows; ir += nth) {
+    for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
         dma_queue_pop(dma);
         dma_queue_pop(dma);

@@ -812,7 +821,7 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
                        S_v * sizeof(float), S_v);

         // Prefetch next block (if any)
-        if (ir_prefetch < total_rows) {
+        if (ir_prefetch < row_end) {
             const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
             const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
             const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -828,6 +837,7 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
         curr_spad_idx ^= 1;
     }
     dma_queue_flush(dma);
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
 }


@@ -847,8 +857,9 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
     const uint32_t H        = v->ne[1];
     const uint32_t n_seqs   = v->ne[3];

-    const uint32_t total_rows = H * n_seqs;
-    if (ith >= total_rows) {
+    const uint32_t row_end = gctx->row_start + gctx->nrows;
+
+    if (ith >= gctx->nrows) {
         return;
     }

@@ -881,11 +892,11 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo

     const uint64_t state_seq_stride = state->nb[3] / sizeof(float);

-    uint32_t ir_prefetch = ith;
+    uint32_t ir_prefetch = gctx->row_start + ith;
     int spad_idx = 0;

     // Prefetch preamble (up to 2 steps)
-    for (int k = 0; k < 2 && ir_prefetch < total_rows; k++) {
+    for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
         const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
         const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
         const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -906,8 +917,11 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
         spad_idx ^= 1;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));
+
     int curr_spad_idx = 0;
-    for (uint32_t ir = ith; ir < total_rows; ir += nth) {
+    for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
         dma_queue_pop(dma);
         dma_queue_pop(dma);

@@ -1057,7 +1071,7 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
                        S_v * sizeof(float), S_v);

         // Prefetch next block (if any)
-        if (ir_prefetch < total_rows) {
+        if (ir_prefetch < row_end) {
             const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
             const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
             const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -1073,6 +1087,7 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
         curr_spad_idx ^= 1;
     }
     dma_queue_flush(dma);
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
 }


@@ -1085,10 +1100,6 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
     const struct htp_tensor * state = octx->src[5];
     const struct htp_tensor * dst   = octx->dst;

-    if (!q || !k || !v || !g || !beta || !state || !dst) {
-        return HTP_STATUS_INVAL_PARAMS;
-    }
-
     if (q->type != HTP_TYPE_F32 || k->type != HTP_TYPE_F32 || v->type != HTP_TYPE_F32 ||
         g->type != HTP_TYPE_F32 || beta->type != HTP_TYPE_F32 || state->type != HTP_TYPE_F32 ||
         dst->type != HTP_TYPE_F32) {
@@ -1124,16 +1135,37 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

+    const uint32_t total_rows = H * n_seqs;
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = total_rows;
+
+    if (octx->ctx->mdev.count > 1) {
+        const uint32_t head_bytes = S_v * sizeof(float);
+        const uint32_t rows_per_chunk = (head_bytes > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(head_bytes, HEX_L2_LINE_SIZE)) : 1;
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0,
+                                                                             octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     struct htp_gdn_context gctx;
     gctx.octx = octx;
-    gctx.rows_per_thread = (H * n_seqs + octx->n_threads - 1) / octx->n_threads;
+    gctx.row_start       = row_start;
+    gctx.nrows           = nrows;
+    gctx.rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
     gctx.state_bytes = (size_t) S_v * S_v * sizeof(float);

     size_t state_aligned = (size_t) S_v * S_v * sizeof(float);
     state_aligned = (state_aligned + 127) & ~(size_t)127;

-    assert(octx->ctx->vtcm_base != NULL);
-    assert(octx->ctx->vtcm_size >= 2 * state_aligned * octx->n_threads);
+    assert(octx->ctx->vtcm_size >= 2 * state_aligned * n_threads);

     gctx.vtcm_base = octx->ctx->vtcm_base;
     gctx.vtcm_per_thread = 2 * state_aligned;
@@ -1148,9 +1180,9 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
          gctx.vtcm_per_thread * octx->n_threads, octx->n_threads);

     if (n_tokens == 1) {
-        worker_pool_run_func(octx->ctx->worker_pool, gated_delta_net_f32_tg_thread, &gctx, octx->n_threads);
+        work_queue_run(octx->ctx->work_queue, gated_delta_net_f32_tg_thread, &gctx, n_threads);
     } else {
-        worker_pool_run_func(octx->ctx->worker_pool, gated_delta_net_f32_pp_thread, &gctx, octx->n_threads);
+        work_queue_run(octx->ctx->work_queue, gated_delta_net_f32_pp_thread, &gctx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/get-rows-ops.c b/ggml/src/ggml-hexagon/htp/get-rows-ops.c
index a87962d22..d294ba57a 100644
--- a/ggml/src/ggml-hexagon/htp/get-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/get-rows-ops.c
@@ -10,6 +10,7 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
@@ -23,9 +24,12 @@ struct get_rows_context {
     const struct htp_get_rows_kernel_params * kparams;
     struct htp_get_rows_vtcm_layout vtcm_layout;
     uint8_t * vtcm_base;
+    uint32_t task_start;
+    uint32_t tasks;
+    uint32_t tasks_per_thread;
 };

-#define get_rows_preamble \
+#define get_rows_preamble                      \
     const uint32_t ne00 = octx->src[0]->ne[0]; \
     const uint32_t ne01 = octx->src[0]->ne[1]; \
     const uint32_t ne02 = octx->src[0]->ne[2]; \
@@ -61,12 +65,12 @@ static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, vo
     struct htp_ops_context * octx = grctx->octx;                                                                       \
     const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                \
     get_rows_preamble;                                                                                                 \
-    const uint32_t dr  = kparams->tasks_per_thread;                                                                    \
-    const uint32_t ir0 = dr * ith;                                                                                     \
-    if (ir0 >= kparams->total_tasks) {                                                                                 \
+    const uint32_t dr  = grctx->tasks_per_thread;                                                                      \
+    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                 \
+    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                     \
         return;                                                                                                        \
     }                                                                                                                  \
-    const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks);                                                          \
+    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                              \
     const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00);                                 \
     dma_queue * dma_queue = octx->ctx->dma[ith];                                                                       \
     for (uint32_t i = ir0; i < ir1; ++i) {                                                                             \
@@ -101,12 +105,12 @@ static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned
     const struct htp_get_rows_kernel_params * kparams = grctx->kparams;                                                \
     get_rows_preamble;                                                                                                 \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                             \
-    const uint32_t dr  = kparams->tasks_per_thread;                                                                    \
-    const uint32_t ir0 = dr * ith;                                                                                     \
-    if (ir0 >= kparams->total_tasks) {                                                                                 \
+    const uint32_t dr  = grctx->tasks_per_thread;                                                                      \
+    const uint32_t ir0 = grctx->task_start + dr * ith;                                                                 \
+    if (ir0 >= grctx->task_start + grctx->tasks) {                                                                     \
         return;                                                                                                        \
     }                                                                                                                  \
-    const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks);                                                          \
+    const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);                                              \
     const uint32_t chunks_per_row = kparams->chunks_per_row;                                                           \
     const uint32_t chunk_size     = kparams->chunk_size;                                                               \
     dma_queue * dma_queue = octx->ctx->dma[ith];                                                                       \
@@ -225,13 +229,41 @@ int op_get_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

+    const struct htp_tensor * dst = octx->dst;
+    const uint32_t total_tasks    = kparams->total_tasks;
+    const size_t dst_row_size     = htp_tensor_get_row_size(dst->type, dst->ne[0]);
+
+    uint32_t task_start = 0;
+    uint32_t tasks      = total_tasks;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t tasks_per_chunk = 1;
+        htp_tensor_mdev_rows_per_chunk(dst, dst_row_size / dst->ne[0], (uint32_t) dst_row_size, &tasks_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_tasks, tasks_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        task_start = range.start;
+        tasks      = range.count;
+    }
+
+    if (tasks == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     struct get_rows_context grctx;
     grctx.octx = octx;
     grctx.kparams = kparams;
     grctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;
+    grctx.task_start = task_start;
+    grctx.tasks = tasks;
+    grctx.tasks_per_thread = fastdiv(tasks + n_threads - 1, &octx->n_threads_div);

     const uint32_t ne00 = octx->src[0]->ne[0];
-    htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, kparams->n_threads);
+    htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, n_threads);

     const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32);

@@ -247,14 +279,14 @@ int op_get_rows(struct htp_ops_context * octx) {
         }
     }

-    FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu use_dma=%d n_threads %d\n",
+    FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu use-dma %d n-threads %d\n",
          octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
          octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3],
          octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
-         grctx.vtcm_layout.src0_bytes_per_thread * kparams->n_threads,
-         grctx.vtcm_layout.dst_bytes_per_thread  * kparams->n_threads,
-         kparams->use_dma, kparams->n_threads);
+         grctx.vtcm_layout.src0_bytes_per_thread * n_threads,
+         grctx.vtcm_layout.dst_bytes_per_thread  * n_threads,
+         kparams->use_dma, n_threads);

-    work_queue_run(octx->ctx->work_queue, q_func, &grctx, kparams->n_threads);
+    work_queue_run(octx->ctx->work_queue, q_func, &grctx, n_threads);
     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/hex-common.h b/ggml/src/ggml-hexagon/htp/hex-common.h
index 4714486a0..e6a52540d 100644
--- a/ggml/src/ggml-hexagon/htp/hex-common.h
+++ b/ggml/src/ggml-hexagon/htp/hex-common.h
@@ -77,4 +77,13 @@ static inline bool hex_add_overflow(size_t a, size_t b, size_t *out) {
     return false;
 }

+static inline uint32_t hex_gcd_u32(uint32_t a, uint32_t b) {
+    while (b != 0) {
+        uint32_t t = b;
+        b = a % b;
+        a = t;
+    }
+    return a;
+}
+
 #endif // HEX_COMMON_H
diff --git a/ggml/src/ggml-hexagon/htp/hex-utils.h b/ggml/src/ggml-hexagon/htp/hex-utils.h
index 1b3965030..853f1c1b2 100644
--- a/ggml/src/ggml-hexagon/htp/hex-utils.h
+++ b/ggml/src/ggml-hexagon/htp/hex-utils.h
@@ -39,7 +39,6 @@ static inline void hex_l2fetch_block(const void * addr, size_t size) {

 #define HEX_L2_LINE_SIZE           128
 #define HEX_L2_BLOCK_SIZE          (HEX_L2_LINE_SIZE * 4) // flush granularity (lines per loop iteration)
-#define HEX_L2_FLUSH_IL_THRESHOLD  1024                   // inline flush threshold
 #define HEX_L2_FLUSH_WQ_THRESHOLD  (4 * 1024)
 #define HEX_L2_FLUSH_ALL_THRESHOLD (4 * 1024 * 1024)

diff --git a/ggml/src/ggml-hexagon/htp/hmx-utils.h b/ggml/src/ggml-hexagon/htp/hmx-utils.h
index 2a61ca734..ad295cb7d 100644
--- a/ggml/src/ggml-hexagon/htp/hmx-utils.h
+++ b/ggml/src/ggml-hexagon/htp/hmx-utils.h
@@ -27,7 +27,7 @@ static inline void hmx_init_column_scales(void *out_scales, HVX_Vector v_scale)
 // vscatter offsets for fused dequant+transpose: write K-values directly to [K][N] tile.
 // word[i] = i*128 maps K-row-pair i to byte offset i*128.
 // Column offset (n*4) is added at runtime.  Entries 0..15 cover one tile (region 2047);
-// entries 16..31 cover the next adjacent tile (region 4095) — pick region size at the
+// entries 16..31 cover the next adjacent tile (region 4095) - pick region size at the
 // call site to scatter into one tile (masked) or two contiguous tiles (unmasked).
 static const int32_t hmx_transpose_scatter_offsets[32] __attribute__((aligned(VLEN))) = {
     0 * 128,  1 * 128,  2 * 128,  3 * 128,  4 * 128,  5 * 128,  6 * 128,  7 * 128,  8 * 128,  9 * 128,  10 * 128,
@@ -198,16 +198,16 @@ static inline void hmx_interleave_cols_to_tiles(__fp16 * restrict tiles_out,
 }

 // --- HMX inline asm macros for load-store packetization ---
-#define HMX_LOAD_MPY_F16(act, wt, range) \
-    "{\n" \
+#define HMX_LOAD_MPY_F16(act, wt, range)              \
+    "{\n"                                             \
     "    activation.hf = mxmem(" act ", " range ")\n" \
-    "    weight.hf = mxmem(" wt ", " range ")\n" \
+    "    weight.hf = mxmem(" wt ", " range ")\n"      \
     "}\n"

-#define HMX_LOAD_MPY_DEEP_F16(act, wt, range) \
-    "{\n" \
+#define HMX_LOAD_MPY_DEEP_F16(act, wt, range)              \
+    "{\n"                                                  \
     "    activation.hf = mxmem(" act ", " range "):deep\n" \
-    "    weight.hf = mxmem(" wt ", " range ")\n" \
+    "    weight.hf = mxmem(" wt ", " range ")\n"           \
     "}\n"

 #define HMX_STORE_AFTER_F16(out, scale_reg) \
diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h
index c8a909d61..3b60c8bdb 100644
--- a/ggml/src/ggml-hexagon/htp/htp-ctx.h
+++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h
@@ -19,7 +19,7 @@
 #endif
 #define HTP_MAX_MMAPS    16

-#define HTP_MAX_DIRTY_RANGES 16
+#define HTP_MAX_DIRTY_RANGES 32

 // Memory mapping
 struct htp_mmap {
@@ -29,6 +29,11 @@ struct htp_mmap {
     uint32_t reserved;
 };

+struct htp_dirty_range {
+    uint32_t start;
+    uint32_t end;
+};
+
 // Scratchpad state
 struct htp_spad {
     const struct htp_tensor * src;             // original src of the data (for reuse)
@@ -38,6 +43,14 @@ struct htp_spad {
     uint32_t                  size_per_thread; // size per thread
 };

+struct htp_mdev_group {
+    uint16_t              idx;
+    uint16_t              count;
+    struct fastdiv_values count_div;
+    uint8_t *             fence_base;
+    uint32_t              fence_seq;
+};
+
 struct htp_context;

 // Context while processing an Op
@@ -65,8 +78,10 @@ struct htp_ops_context {
     struct htp_spad src3_spad;
     struct htp_spad dst_spad;

-    uint32_t n_threads;
-    uint32_t flags;
+    uint32_t              flags;
+    uint32_t              n_threads;
+    struct fastdiv_values n_threads_div;
+    int                   status;
 };

 // Main context for htp DSP backend
@@ -76,6 +91,7 @@ struct htp_context {
     struct htp_mmap        mmap[HTP_MAX_MMAPS];
     dma_queue_t            dma[HTP_MAX_NTHREADS];
     dma_queue_t            dma_cached[HTP_MAX_NTHREADS];
+    struct htp_thread_trace trace[HTP_MAX_NTHREADS + 1];
     work_queue_t           work_queue;
     hmx_queue_t            hmx_queue;

@@ -88,7 +104,6 @@ struct htp_context {
     bool                   hmx_enabled;
     bool                   etm;
     uint32_t               profiler;
-    struct htp_thread_trace trace[HTP_MAX_NTHREADS + 1];

     uint8_t *              vtcm_base;
     size_t                 vtcm_size;
@@ -97,16 +112,13 @@ struct htp_context {
     atomic_bool            vtcm_needs_release;

     uint64_t               max_vmem;
-    struct htp_dirty_range {
-        uint32_t start;
-        uint32_t end;
-        uint32_t bi;
-    } dirty_ranges[HTP_MAX_DIRTY_RANGES];
+    struct htp_dirty_range dirty_ranges[HTP_MAX_DIRTY_RANGES];

     // Persistent DDR scratchpad for MUL_MAT_ID mappings
     void *                 ddr_spad_base;
     size_t                 ddr_spad_size;

+    struct htp_mdev_group  mdev;
     struct htp_ops_context octx;

     qurt_thread_t          main_thread;
@@ -115,6 +127,27 @@ struct htp_context {
     size_t                 footprint;
 };

+static inline bool htp_ops_context_set_n_threads(struct htp_ops_context * octx, uint32_t n_threads) {
+    if (n_threads == 0 || n_threads > octx->ctx->n_threads) {
+        return false;
+    }
+
+    if (n_threads != octx->n_threads) {
+        octx->n_threads = n_threads;
+        octx->n_threads_div = n_threads == octx->ctx->n_threads
+            ? octx->ctx->n_threads_div
+            : init_fastdiv_values(n_threads);
+    }
+
+    return true;
+}
+
+static inline void htp_ops_context_set_status(struct htp_ops_context * octx, int status) {
+    if (status > HTP_STATUS_OK && octx->status == HTP_STATUS_OK) {
+        octx->status = status;
+    }
+}
+
 int op_matmul(struct htp_ops_context * octx);
 int op_matmul_id(struct htp_ops_context * octx);
 int op_matmul_nx(struct htp_ops_context * octx);
diff --git a/ggml/src/ggml-hexagon/htp/htp-fence.h b/ggml/src/ggml-hexagon/htp/htp-fence.h
new file mode 100644
index 000000000..7450b5de5
--- /dev/null
+++ b/ggml/src/ggml-hexagon/htp/htp-fence.h
@@ -0,0 +1,89 @@
+#ifndef HTP_FENCE_H
+#define HTP_FENCE_H
+
+#include <stdatomic.h>
+#include <stdint.h>
+
+#include <HAP_farf.h>
+
+#include "hex-utils.h"
+#include "htp-ops.h"
+#include "htp-ctx.h"
+
+static inline atomic_uint * htp_mdev_fence_slot(const void * fence_base, uint32_t idx) {
+    return (atomic_uint *) ((const uint8_t *) fence_base + (size_t) idx * HTP_FENCE_SLOT_SIZE);
+}
+
+static inline void htp_fence_write(void * fence_ptr, uint32_t seq, uint32_t status) {
+    atomic_uint * fence = (atomic_uint *) fence_ptr;
+    atomic_store(&fence[1], status);
+    atomic_store(&fence[0], seq);
+    asm volatile ("syncht" : : : "memory");
+    Q6_dccleaninva_A((void *) fence);
+}
+
+static inline void htp_fence_read(const void * fence_ptr, uint32_t * seq, uint32_t * status) {
+    const atomic_uint * fence = (const atomic_uint *) fence_ptr;
+    Q6_dccleaninva_A((void *) fence);
+    asm volatile ("syncht" : : : "memory");
+    *seq = atomic_load(&fence[0]);
+    *status = atomic_load(&fence[1]);
+}
+
+static inline void htp_mdev_group_barrier(struct htp_ops_context * octx) {
+    struct htp_context * ctx = octx->ctx;
+    if (ctx->mdev.count <= 1) {
+        return;
+    }
+
+    const uint32_t seq = ++ctx->mdev.fence_seq;
+
+    struct htp_thread_trace * tr = &ctx->trace[0];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
+
+    const uint32_t mdev_idx   = ctx->mdev.idx;
+    const uint32_t mdev_count = ctx->mdev.count;
+
+    uint8_t * fence_base   = ctx->mdev.fence_base;
+    atomic_uint * my_fence = htp_mdev_fence_slot(fence_base, mdev_idx);
+    htp_fence_write(my_fence, seq, octx->status);
+
+    for (uint32_t d = 0; d < mdev_count; d++) {
+        if (d == mdev_idx) continue;
+        atomic_uint * peer_fence = htp_mdev_fence_slot(fence_base, d);
+        uint64_t spins = 0;
+        while (1) {
+            uint32_t peer_seq;
+            uint32_t peer_status;
+            htp_fence_read(peer_fence, &peer_seq, &peer_status);
+            if ((int32_t)(peer_seq - seq) >= 0) {
+                if (peer_status > HTP_STATUS_OK) {
+                    FARF(ERROR, "ggml-hex: mdev %u peer %u failed with status %u : seq 0x%08x\n",
+                         mdev_idx, d, peer_status, seq);
+                    htp_ops_context_set_status(octx, peer_status);
+                }
+                break;
+            }
+            if (++spins == 10000) {
+                FARF(ALWAYS, "ggml-hex: mdev %u waiting for mdev %u : seq 0x%08x (b %u op %u) my-fence %p peer-fence %p peer-seq 0x%08x (diff %d)\n",
+                     mdev_idx, d, seq, seq >> 12, seq & 0xfff, my_fence, peer_fence, peer_seq, (int32_t)(peer_seq - seq));
+            }
+            if (spins > HTP_FENCE_TIMEOUT) {
+                FARF(ERROR, "ggml-hex: mdev %u timeout waiting for mdev %u : seq 0x%08x (b %u op %u) peer-fence %p peer-seq 0x%08x\n",
+                     mdev_idx, d, seq, seq >> 12, seq & 0xfff, peer_fence, peer_seq);
+                htp_ops_context_set_status(octx, HTP_STATUS_INTERNAL_ERR);
+                break;
+            }
+            hex_pause();
+        }
+    }
+    asm volatile ("syncht" : : : "memory");
+
+    if (octx->status > HTP_STATUS_OK) {
+        htp_fence_write(my_fence, seq, octx->status);
+    }
+
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
+}
+
+#endif // HTP_FENCE_H
diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h
index 12a61b67f..869b19b8c 100644
--- a/ggml/src/ggml-hexagon/htp/htp-ops.h
+++ b/ggml/src/ggml-hexagon/htp/htp-ops.h
@@ -77,6 +77,7 @@ enum htp_op_code {
     HTP_OP_GET_ROWS,
     HTP_OP_SCALE,
     HTP_OP_CPY,
+    HTP_OP_CPY_FENCE,
     HTP_OP_ARGSORT,
     HTP_OP_SQR,
     HTP_OP_SQRT,
@@ -100,6 +101,7 @@ enum htp_op_code {
     HTP_OP_ALLREDUCE,
     HTP_OP_ALLREDUCE_ADD,
     HTP_OP_GLU_SWIGLU_CLAMP,
+    HTP_OP_MDEV_GROUP,

     HTP_OP_INVALID
 };
@@ -114,6 +116,7 @@ enum htp_op_code {
 #define HTP_OP_MAX_TENSORS 8192 // must stay under 64K (uint16)

 #define HTP_FENCE_TIMEOUT  (1000000000ULL)
+#define HTP_FENCE_SLOT_SIZE 128

 #define HTP_OP_MAX_VMEM_DEFAULT (3355443200u)

@@ -214,30 +217,26 @@ struct htp_prof_desc {
 };

 struct htp_opbatch_req {
-    uint32_t id;          // Batch id
+    uint64_t seq;         // Sequence number
     uint32_t n_bufs;      // Number of buffers
     uint32_t n_tensors;   // Number of tensors
     uint32_t n_ops;       // Number of ops
     uint32_t n_traces;    // Number of trace descriptors per thread
-    uint32_t pad;         // unused
-    uint64_t seq;         // Sequence number
     // struct htp_buf_desc  bufs[];    -- dspqueue buf 0
     // struct htp_tensor    tensors[]; -- dspqueue buf 0
     // struct htp_op_desc   ops[];     -- dspqueue buf 0
 };

 struct htp_opbatch_rsp {
-    uint32_t id;         // Batch id
-    uint32_t status;     // HTP_STATUS_...
-    uint32_t n_bufs;     // Number of buffers
-    uint32_t n_tensors;  // Number of tensors
-    uint32_t n_ops;      // Number of op profile descriptors
-    uint32_t n_traces[HTP_MAX_NTHREADS + 1];
-    uint32_t usecs;          // Number of usec
-    uint32_t pad;            // align to 8 bytes
+    uint64_t seq;            // Sequence number
     uint64_t cycles_start;   // Start cycle counter
     uint64_t cycles_stop;    // Stop cycle counter
-    uint64_t seq;            // Sequence number
+    uint32_t status;         // HTP_STATUS_...
+    uint32_t n_bufs;         // Number of buffers
+    uint32_t n_tensors;      // Number of tensors
+    uint32_t n_ops;          // Number of op profile descriptors
+    uint32_t usecs;          // Number of usec
+    uint32_t n_traces[HTP_MAX_NTHREADS + 1];
     // struct htp_prof_desc profs[];  -- dspqueue buf 0
 };

diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.c b/ggml/src/ggml-hexagon/htp/htp-tensor.c
index ae377c922..760ccd831 100644
--- a/ggml/src/ggml-hexagon/htp/htp-tensor.c
+++ b/ggml/src/ggml-hexagon/htp/htp-tensor.c
@@ -20,7 +20,7 @@ struct l2flush_range {

 struct l2flush_multi_task {
     struct htp_thread_trace * trace;
-    struct l2flush_range      ranges[HTP_OP_MAX_INPUTS];
+    struct l2flush_range      ranges[HTP_MAX_DIRTY_RANGES];
     uint32_t                  n_ranges;
     uint32_t                  total_blocks;
     uint32_t                  blocks_per_thread;
@@ -73,6 +73,27 @@ static void l2flush_multi_worker(unsigned int n, unsigned int i, void * data) {
     htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, gb_first);
 }

+static void merge_dirty_ranges(struct htp_context * ctx) {
+    for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) {
+        struct htp_dirty_range * r = &ctx->dirty_ranges[i];
+        if (!r->start) continue;
+
+        for (uint32_t j = 0; j < HTP_MAX_DIRTY_RANGES;) {
+            struct htp_dirty_range * s = &ctx->dirty_ranges[j];
+            if (i == j || !s->start || r->end < s->start || s->end < r->start) {
+                j++;
+                continue;
+            }
+
+            r->start = MIN(r->start, s->start);
+            r->end   = MAX(r->end, s->end);
+            s->start = 0;
+            s->end   = 0;
+            j = 0;
+        }
+    }
+}
+
 void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) {
     const struct htp_tensor * pending[HTP_OP_MAX_OUTPUTS];
     uint32_t n_pending = 0;
@@ -83,11 +104,6 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
             continue;
         }

-        if (t->size <= HEX_L2_FLUSH_IL_THRESHOLD) {
-            hex_l2flush((void *) (uintptr_t) t->data, t->size);
-            continue;
-        }
-
         uint32_t t_start = t->data;
         uint32_t t_end   = t_start + t->size;

@@ -110,6 +126,8 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
         }
     }

+    merge_dirty_ranges(ctx);
+
     if (n_pending == 0) {
         return;
     }
@@ -132,8 +150,8 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
             struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
             r->start = pending[i]->data;
             r->end   = pending[i]->data + pending[i]->size;
-            r->bi    = pending[i]->bi;
         }
+        merge_dirty_ranges(ctx);
         return;
     }

@@ -151,12 +169,12 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
             struct htp_dirty_range * r = &ctx->dirty_ranges[i];
             r->start = pending[i]->data;
             r->end   = pending[i]->data + pending[i]->size;
-            r->bi    = pending[i]->bi;
         }
+        merge_dirty_ranges(ctx);
         return;
     }

-    if (total_evict_size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1 && n_evict <= HTP_OP_MAX_INPUTS) {
+    if (total_evict_size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1 && n_evict <= HTP_MAX_DIRTY_RANGES) {
         struct l2flush_multi_task task;
         task.trace    = ctx->trace;
         task.n_ranges = n_evict;
@@ -195,7 +213,6 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
         struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
         r->start = pending[i]->data;
         r->end   = pending[i]->data + pending[i]->size;
-        r->bi    = pending[i]->bi;
     }

     for (uint32_t i = 0; i < n_empty; i++) {
@@ -203,8 +220,9 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
         struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
         r->start = pending[n_evict + i]->data;
         r->end   = pending[n_evict + i]->data + pending[n_evict + i]->size;
-        r->bi    = pending[n_evict + i]->bi;
     }
+
+    merge_dirty_ranges(ctx);
 }

 static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor * t) {
@@ -242,17 +260,50 @@ static inline bool is_tensor_dirty(struct htp_context * ctx, const struct htp_te
     return false;
 }

-void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) {
-    const struct htp_tensor * dirty_tensors[HTP_OP_MAX_INPUTS];
-    uint32_t n_dirty = 0;
+static void flush_dirty_ranges(struct htp_context * ctx, const struct htp_dirty_range * ranges, uint32_t n_ranges, uint64_t total_dirty) {
+    if (total_dirty >= HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) {
+        struct l2flush_multi_task task;
+        task.trace    = ctx->trace;
+        task.n_ranges = n_ranges;
+
+        uint32_t block_acc = 0;
+        for (uint32_t i = 0; i < n_ranges; i++) {
+            const struct htp_dirty_range * r = &ranges[i];
+            struct l2flush_range * rg = &task.ranges[i];
+            rg->start = hex_align_down((size_t) r->start, HEX_L2_LINE_SIZE);
+            rg->end   = hex_align_up((size_t) r->end, HEX_L2_LINE_SIZE);
+            rg->block_first = block_acc;
+            rg->n_blocks = (rg->end - rg->start + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE;
+            block_acc += rg->n_blocks;
+        }
+
+        task.total_blocks      = block_acc;
+        task.blocks_per_thread = fastdiv(block_acc + ctx->n_threads - 1, &ctx->n_threads_div);
+
+        work_queue_run(ctx->work_queue, l2flush_multi_worker, &task, ctx->n_threads);
+    } else {
+        struct htp_thread_trace * tr = &ctx->trace[0];
+        htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, 0);
+        for (uint32_t i = 0; i < n_ranges; i++) {
+            const struct htp_dirty_range * r = &ranges[i];
+            hex_l2flush((void *) (uintptr_t) r->start, r->end - r->start);
+        }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, 0);
+    }
+}
+
+void htp_flush_dirty_ranges(struct htp_context * ctx) {
+    struct htp_dirty_range ranges[HTP_MAX_DIRTY_RANGES];
+    uint32_t n_ranges = 0;
     uint64_t total_dirty = 0;

-    for (uint32_t i = 0; i < n; i++) {
-        const struct htp_tensor * t = tensors[i];
-        if (t && !(t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE)) && is_tensor_dirty(ctx, t)) {
-            dirty_tensors[n_dirty++] = t;
-            total_dirty += t->size;
+    for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) {
+        const struct htp_dirty_range * r = &ctx->dirty_ranges[i];
+        if (!r->start) {
+            continue;
         }
+        ranges[n_ranges++] = *r;
+        total_dirty += r->end - r->start;
     }

     if (total_dirty == 0) {
@@ -264,37 +315,37 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co
         return;
     }

-    if (total_dirty >= HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) {
-        struct l2flush_multi_task task;
-        task.trace    = ctx->trace;
-        task.n_ranges = 0;
+    flush_dirty_ranges(ctx, ranges, n_ranges, total_dirty);
+    memset(ctx->dirty_ranges, 0, sizeof(ctx->dirty_ranges));
+}

-        uint32_t block_acc = 0;
-        for (uint32_t i = 0; i < n_dirty; i++) {
-            const struct htp_tensor * t = dirty_tensors[i];
-            make_tensor_clean(ctx, t);
+void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) {
+    const struct htp_tensor * dirty_tensors[HTP_OP_MAX_INPUTS];
+    struct htp_dirty_range ranges[HTP_OP_MAX_INPUTS];
+    uint32_t n_dirty = 0;
+    uint64_t total_dirty = 0;

-            struct l2flush_range * rg = &task.ranges[task.n_ranges++];
-            rg->start = hex_align_down((size_t) t->data, HEX_L2_LINE_SIZE);
-            rg->end   = hex_align_up((size_t) t->data + t->size, HEX_L2_LINE_SIZE);
-            rg->block_first = block_acc;
-            rg->n_blocks = (rg->end - rg->start + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE;
-            block_acc += rg->n_blocks;
+    for (uint32_t i = 0; i < n; i++) {
+        const struct htp_tensor * t = tensors[i];
+        if (t && is_tensor_dirty(ctx, t)) {
+            dirty_tensors[n_dirty++] = t;
+            ranges[n_dirty - 1].start = t->data;
+            ranges[n_dirty - 1].end   = t->data + t->size;
+            total_dirty += t->size;
         }
+    }

-        task.total_blocks      = block_acc;
-        task.blocks_per_thread = fastdiv(block_acc + ctx->n_threads - 1, &ctx->n_threads_div);
+    if (total_dirty == 0) {
+        return;
+    }

-        work_queue_run(ctx->work_queue, l2flush_multi_worker, &task, ctx->n_threads);
+    if (total_dirty > HEX_L2_FLUSH_ALL_THRESHOLD) {
+        flush_all_dcache(ctx);
         return;
     }

-    struct htp_thread_trace * tr = &ctx->trace[0];
+    flush_dirty_ranges(ctx, ranges, n_dirty, total_dirty);
     for (uint32_t i = 0; i < n_dirty; i++) {
-        const struct htp_tensor * t = dirty_tensors[i];
-        htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, t->ti);
-        hex_l2flush((void *) (uintptr_t) t->data, t->size);
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, t->ti);
-        make_tensor_clean(ctx, t);
+        make_tensor_clean(ctx, dirty_tensors[i]);
     }
 }
diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.h b/ggml/src/ggml-hexagon/htp/htp-tensor.h
index c9cadbae3..3afff6917 100644
--- a/ggml/src/ggml-hexagon/htp/htp-tensor.h
+++ b/ggml/src/ggml-hexagon/htp/htp-tensor.h
@@ -2,8 +2,20 @@
 #define HTP_TENSOR_H

 #include <stdint.h>
+#include <stdbool.h>
 #include "htp-ops.h"
 #include "hex-bitmap.h"
+#include "hex-common.h"
+#include "hex-fastdiv.h"
+
+enum {
+    HTP_TENSOR_MDEV_LINE_SIZE = 128,
+};
+
+struct htp_tensor_mdev_range {
+    uint32_t start;
+    uint32_t count;
+};

 static inline void * htp_tensor_data(const struct htp_tensor * t) {
     return (void *) (uintptr_t) t->data;
@@ -13,6 +25,102 @@ static inline uint32_t * htp_tensor_flags(const struct htp_tensor * t) {
     return (uint32_t *) &t->flags;
 }

+static inline bool htp_tensor_is_contiguous(const struct htp_tensor * t, uint32_t type_size) {
+    uint32_t next_nb = type_size;
+    if (t->ne[0] != 1 && t->nb[0] != next_nb) {
+        return false;
+    }
+    next_nb *= t->ne[0];
+    for (int i = 1; i < HTP_OP_MAX_DIMS; i++) {
+        if (t->ne[i] != 1 && t->nb[i] != next_nb) {
+            return false;
+        }
+        next_nb *= t->ne[i];
+    }
+    return true;
+}
+
+static inline bool htp_tensor_is_permuted(const struct htp_tensor * t) {
+    return t->nb[0] > t->nb[1] || t->nb[1] > t->nb[2] || t->nb[2] > t->nb[3];
+}
+
+static inline bool htp_tensor_mdev_data_aligned(const struct htp_tensor * t) {
+    return ((uintptr_t) t->data & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0;
+}
+
+static inline bool htp_tensor_can_row_partition(const struct htp_tensor * t, uint32_t elem_size) {
+    if (!htp_tensor_mdev_data_aligned(t)) {
+        return false;
+    }
+    if (t->ne[0] != 1 && t->nb[0] != elem_size) {
+        return false;
+    }
+    if (htp_tensor_is_permuted(t)) {
+        return false;
+    }
+    if (t->ne[1] > 1 && (t->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
+    if (t->ne[2] > 1 && (t->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
+    if (t->ne[3] > 1 && (t->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
+    return true;
+}
+
+static inline bool htp_tensor_mdev_rows_per_chunk(const struct htp_tensor * t, uint32_t elem_size, uint32_t row_size, uint32_t * rows_per_chunk) {
+    *rows_per_chunk = 0;
+
+    if (!htp_tensor_mdev_data_aligned(t)) {
+        return false;
+    }
+    if (t->ne[0] != 1 && t->nb[0] != elem_size) {
+        return false;
+    }
+    if (htp_tensor_is_permuted(t)) {
+        return false;
+    }
+    if (t->ne[1] > 1 && (t->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 &&
+        (t->ne[2] <= 1 || (t->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0) &&
+        (t->ne[3] <= 1 || (t->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0)) {
+        *rows_per_chunk = 1;
+        return true;
+    }
+    if (t->nb[1] == row_size &&
+        (t->ne[2] <= 1 || t->nb[2] == t->nb[1] * t->ne[1]) &&
+        (t->ne[3] <= 1 || t->nb[3] == t->nb[2] * t->ne[2])) {
+        *rows_per_chunk = (row_size > 0) ? (HTP_TENSOR_MDEV_LINE_SIZE / hex_gcd_u32(row_size, HTP_TENSOR_MDEV_LINE_SIZE)) : 1;
+        return true;
+    }
+    return false;
+}
+
+static inline struct htp_tensor_mdev_range htp_tensor_mdev_partition(uint32_t total_units, uint32_t units_per_chunk, uint32_t mdev_idx, uint32_t mdev_count, const struct fastdiv_values * mdev_count_div) {
+    struct htp_tensor_mdev_range range = { 0, total_units };
+
+    if (mdev_count <= 1) {
+        return range;
+    }
+
+    if (units_per_chunk == 0) {
+        range.start = (mdev_idx == 0) ? 0 : total_units;
+        range.count = (mdev_idx == 0) ? total_units : 0;
+        return range;
+    }
+
+    const uint32_t total_chunks = total_units / units_per_chunk;
+    if (total_chunks < mdev_count) {
+        range.start = (mdev_idx == 0) ? 0 : total_units;
+        range.count = (mdev_idx == 0) ? total_units : 0;
+        return range;
+    }
+
+    const uint32_t chunks_per_mdev = fastdiv(total_chunks + mdev_count - 1, mdev_count_div);
+    range.start = MIN(mdev_idx * chunks_per_mdev * units_per_chunk, total_units);
+    if (mdev_idx == mdev_count - 1) {
+        range.count = total_units - range.start;
+    } else {
+        range.count = MIN(chunks_per_mdev * units_per_chunk, total_units - range.start);
+    }
+    return range;
+}
+
 static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) {
     switch (type) {
         case HTP_TYPE_F32:  return ne00 * 4;
@@ -23,6 +131,7 @@ static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) {
 }

 struct htp_context;
+void htp_flush_dirty_ranges(struct htp_context * ctx);
 void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n);
 void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n);

diff --git a/ggml/src/ggml-hexagon/htp/hvx-arith.h b/ggml/src/ggml-hexagon/htp/hvx-arith.h
index fe5477c1b..6cbead74c 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-arith.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-arith.h
@@ -16,25 +16,25 @@
 #define UNUSED(x) (void)(x)

 #define hvx_arith_loop_body(dst_type, src0_type, src1_type, elem_size, vec_store, vec_op) \
-    do {                                                                       \
-        dst_type * vdst  = (dst_type *) dst;                                   \
-        src0_type * vsrc0 = (src0_type *) src0;                                \
-        src1_type * vsrc1 = (src1_type *) src1;                                \
-                                                                               \
-        const uint32_t epv  = 128 / (elem_size);                               \
-        const uint32_t nvec = n / epv;                                         \
-        const uint32_t nloe = n % epv;                                         \
-                                                                               \
-        uint32_t i = 0;                                                        \
-                                                                               \
-        _Pragma("unroll(4)")                                                   \
-        for (; i < nvec; i++) {                                                \
-            vdst[i] = vec_op(vsrc0[i], vsrc1[i]);                              \
-        }                                                                      \
-        if (nloe) {                                                            \
-            HVX_Vector v = vec_op(vsrc0[i], vsrc1[i]);                         \
-            vec_store((void *) &vdst[i], nloe * (elem_size), v);               \
-        }                                                                      \
+    do {                                                                                  \
+        dst_type * vdst  = (dst_type *) dst;                                              \
+        src0_type * vsrc0 = (src0_type *) src0;                                           \
+        src1_type * vsrc1 = (src1_type *) src1;                                           \
+                                                                                          \
+        const uint32_t epv  = 128 / (elem_size);                                          \
+        const uint32_t nvec = n / epv;                                                    \
+        const uint32_t nloe = n % epv;                                                    \
+                                                                                          \
+        uint32_t i = 0;                                                                   \
+                                                                                          \
+        _Pragma("unroll(4)")                                                              \
+        for (; i < nvec; i++) {                                                           \
+            vdst[i] = vec_op(vsrc0[i], vsrc1[i]);                                         \
+        }                                                                                 \
+        if (nloe) {                                                                       \
+            HVX_Vector v = vec_op(vsrc0[i], vsrc1[i]);                                    \
+            vec_store((void *) &vdst[i], nloe * (elem_size), v);                          \
+        }                                                                                 \
     } while(0)

 #if __HVX_ARCH__ < 79
@@ -56,43 +56,43 @@
 #define HVX_OP_MUL_F16(a, b) hvx_vec_mul_f16_f16(a, b)

 // Generic macro to define alignment permutations for an op
-#define DEFINE_HVX_BINARY_OP_VARIANTS(OP_NAME, OP_MACRO, ELEM_TYPE) \
-static inline void OP_NAME##_aaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src0 % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
-static inline void OP_NAME##_aau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src0 % 128 == 0); \
-    hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
-static inline void OP_NAME##_aua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
-static inline void OP_NAME##_auu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
-static inline void OP_NAME##_uaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) src0 % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
-static inline void OP_NAME##_uau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) src0 % 128 == 0); \
-    hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
-static inline void OP_NAME##_uua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
-    assert((uintptr_t) src1 % 128 == 0); \
-    hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
-static inline void OP_NAME##_uuu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
+#define DEFINE_HVX_BINARY_OP_VARIANTS(OP_NAME, OP_MACRO, ELEM_TYPE)                                           \
+static inline void OP_NAME##_aaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) dst % 128 == 0);                                                                       \
+    assert((uintptr_t) src0 % 128 == 0);                                                                      \
+    assert((uintptr_t) src1 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);    \
+}                                                                                                             \
+static inline void OP_NAME##_aau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) dst % 128 == 0);                                                                       \
+    assert((uintptr_t) src0 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);   \
+}                                                                                                             \
+static inline void OP_NAME##_aua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) dst % 128 == 0);                                                                       \
+    assert((uintptr_t) src1 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);   \
+}                                                                                                             \
+static inline void OP_NAME##_auu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) dst % 128 == 0);                                                                       \
+    hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);  \
+}                                                                                                             \
+static inline void OP_NAME##_uaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) src0 % 128 == 0);                                                                      \
+    assert((uintptr_t) src1 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO);   \
+}                                                                                                             \
+static inline void OP_NAME##_uau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) src0 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO);  \
+}                                                                                                             \
+static inline void OP_NAME##_uua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
+    assert((uintptr_t) src1 % 128 == 0);                                                                      \
+    hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO);  \
+}                                                                                                             \
+static inline void OP_NAME##_uuu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) {     \
     hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
+}                                                                                                             \

 DEFINE_HVX_BINARY_OP_VARIANTS(hvx_add_f32, HVX_OP_ADD_F32, float)
 DEFINE_HVX_BINARY_OP_VARIANTS(hvx_sub_f32, HVX_OP_SUB_F32, float)
@@ -103,25 +103,25 @@ DEFINE_HVX_BINARY_OP_VARIANTS(hvx_sub_f16, HVX_OP_SUB_F16, _Float16)
 DEFINE_HVX_BINARY_OP_VARIANTS(hvx_mul_f16, HVX_OP_MUL_F16, _Float16)

 // Dispatcher logic
-#define HVX_BINARY_DISPATCHER(OP_NAME) \
+#define HVX_BINARY_DISPATCHER(OP_NAME)                                                                                                       \
 static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, const uint32_t num_elems) { \
-    if (hex_is_aligned((void *) dst, 128)) { \
-        if (hex_is_aligned((void *) src0, 128)) { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_aau(dst, src0, src1, num_elems); \
-        } else { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_auu(dst, src0, src1, num_elems); \
-        } \
-    } else { \
-        if (hex_is_aligned((void *) src0, 128)) { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_uau(dst, src0, src1, num_elems); \
-        } else { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_uuu(dst, src0, src1, num_elems); \
-        } \
-    } \
+    if (hex_is_aligned((void *) dst, 128)) {                                                                                                 \
+        if (hex_is_aligned((void *) src0, 128)) {                                                                                            \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_aau(dst, src0, src1, num_elems);                                               \
+        } else {                                                                                                                             \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_auu(dst, src0, src1, num_elems);                                               \
+        }                                                                                                                                    \
+    } else {                                                                                                                                 \
+        if (hex_is_aligned((void *) src0, 128)) {                                                                                            \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_uau(dst, src0, src1, num_elems);                                               \
+        } else {                                                                                                                             \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_uuu(dst, src0, src1, num_elems);                                               \
+        }                                                                                                                                    \
+    }                                                                                                                                        \
 }

 HVX_BINARY_DISPATCHER(hvx_add_f32)
@@ -166,44 +166,44 @@ static inline void hvx_mul_mul_f32_aa(uint8_t * restrict dst, const uint8_t * re

 // Scalar Operations

-#define hvx_scalar_loop_body(dst_type, src_type, elem_size, vec_store, scalar_op_macro)   \
-    do {                                                                       \
-        dst_type * restrict vdst = (dst_type *) dst;                           \
-        src_type * restrict vsrc = (src_type *) src;                           \
-                                                                               \
-        const uint32_t epv  = 128 / (elem_size);                               \
-        const uint32_t nvec = n / epv;                                         \
-        const uint32_t nloe = n % epv;                                         \
-                                                                               \
-        uint32_t i = 0;                                                        \
-                                                                               \
-        _Pragma("unroll(4)")                                                   \
-        for (; i < nvec; i++) {                                                \
-            HVX_Vector v = vsrc[i];                                            \
-            vdst[i] = scalar_op_macro(v);                                      \
-        }                                                                      \
-        if (nloe) {                                                            \
-            HVX_Vector v = vsrc[i];                                            \
-            v = scalar_op_macro(v);                                            \
-            vec_store((void *) &vdst[i], nloe * (elem_size), v);               \
-        }                                                                      \
+#define hvx_scalar_loop_body(dst_type, src_type, elem_size, vec_store, scalar_op_macro) \
+    do {                                                                                \
+        dst_type * restrict vdst = (dst_type *) dst;                                    \
+        src_type * restrict vsrc = (src_type *) src;                                    \
+                                                                                        \
+        const uint32_t epv  = 128 / (elem_size);                                        \
+        const uint32_t nvec = n / epv;                                                  \
+        const uint32_t nloe = n % epv;                                                  \
+                                                                                        \
+        uint32_t i = 0;                                                                 \
+                                                                                        \
+        _Pragma("unroll(4)")                                                            \
+        for (; i < nvec; i++) {                                                         \
+            HVX_Vector v = vsrc[i];                                                     \
+            vdst[i] = scalar_op_macro(v);                                               \
+        }                                                                               \
+        if (nloe) {                                                                     \
+            HVX_Vector v = vsrc[i];                                                     \
+            v = scalar_op_macro(v);                                                     \
+            vec_store((void *) &vdst[i], nloe * (elem_size), v);                        \
+        }                                                                               \
     } while(0)

-#define HVX_OP_ADD_SCALAR_F32(v) \
-    ({ \
+#define HVX_OP_ADD_SCALAR_F32(v)                                   \
+    ({                                                             \
         const HVX_VectorPred pred_inf = Q6_Q_vcmp_eq_VwVw(inf, v); \
-        HVX_Vector out = HVX_OP_ADD_F32(v, val_vec); \
-        Q6_V_vmux_QVV(pred_inf, inf, out); \
+        HVX_Vector out = HVX_OP_ADD_F32(v, val_vec);               \
+        Q6_V_vmux_QVV(pred_inf, inf, out);                         \
     })

 #define HVX_OP_MUL_SCALAR_F32(v) HVX_OP_MUL_F32(v, val_vec)
 #define HVX_OP_SUB_SCALAR_F32(v) HVX_OP_SUB_F32(v, val_vec)

-#define HVX_OP_ADD_SCALAR_F16(v) \
-    ({ \
+#define HVX_OP_ADD_SCALAR_F16(v)                                   \
+    ({                                                             \
         const HVX_VectorPred pred_inf = Q6_Q_vcmp_eq_VhVh(inf, v); \
-        HVX_Vector out = HVX_OP_ADD_F16(v, val_vec); \
-        Q6_V_vmux_QVV(pred_inf, inf, out); \
+        HVX_Vector out = HVX_OP_ADD_F16(v, val_vec);               \
+        Q6_V_vmux_QVV(pred_inf, inf, out);                         \
     })

 #define HVX_OP_MUL_SCALAR_F16(v) HVX_OP_MUL_F16(v, val_vec)
@@ -212,31 +212,31 @@ static inline void hvx_mul_mul_f32_aa(uint8_t * restrict dst, const uint8_t * re
 // Scalar Variants

 // Generic macro to define alignment permutations for an op
-#define DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(OP_NAME, OP_MACRO, SPLAT_MACRO, ELEM_TYPE) \
+#define DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(OP_NAME, OP_MACRO, SPLAT_MACRO, ELEM_TYPE)                                  \
 static inline void OP_NAME##_aa(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
-    const HVX_Vector val_vec = SPLAT_MACRO(val); \
-    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src % 128 == 0); \
-    hvx_scalar_loop_body(HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
+    const HVX_Vector val_vec = SPLAT_MACRO(val);                                                                         \
+    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf);                                                \
+    assert((uintptr_t) dst % 128 == 0);                                                                                  \
+    assert((uintptr_t) src % 128 == 0);                                                                                  \
+    hvx_scalar_loop_body(HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);                          \
+}                                                                                                                        \
 static inline void OP_NAME##_au(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
-    const HVX_Vector val_vec = SPLAT_MACRO(val); \
-    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
-    assert((uintptr_t) dst % 128 == 0); \
-    hvx_scalar_loop_body(HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
-} \
+    const HVX_Vector val_vec = SPLAT_MACRO(val);                                                                         \
+    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf);                                                \
+    assert((uintptr_t) dst % 128 == 0);                                                                                  \
+    hvx_scalar_loop_body(HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO);                         \
+}                                                                                                                        \
 static inline void OP_NAME##_ua(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
-    const HVX_Vector val_vec = SPLAT_MACRO(val); \
-    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
-    assert((uintptr_t) src % 128 == 0); \
-    hvx_scalar_loop_body(HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
+    const HVX_Vector val_vec = SPLAT_MACRO(val);                                                                         \
+    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf);                                                \
+    assert((uintptr_t) src % 128 == 0);                                                                                  \
+    hvx_scalar_loop_body(HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO);                         \
+}                                                                                                                        \
 static inline void OP_NAME##_uu(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
-    const HVX_Vector val_vec = SPLAT_MACRO(val); \
-    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
-    hvx_scalar_loop_body(HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
-} \
+    const HVX_Vector val_vec = SPLAT_MACRO(val);                                                                         \
+    const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf);                                                \
+    hvx_scalar_loop_body(HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO);                        \
+}                                                                                                                        \

 DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_add_scalar_f32, HVX_OP_ADD_SCALAR_F32, hvx_vec_splat_f32, float)
 DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_sub_scalar_f32, HVX_OP_SUB_SCALAR_F32, hvx_vec_splat_f32, float)
@@ -247,17 +247,17 @@ DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_sub_scalar_f16, HVX_OP_SUB_SCALAR_F16,
 DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_mul_scalar_f16, HVX_OP_MUL_SCALAR_F16, hvx_vec_splat_f16, _Float16)

 // Dispatcher logic
-#define HVX_BINARY_SCALAR_DISPATCHER(OP_NAME, ELEM_TYPE) \
+#define HVX_BINARY_SCALAR_DISPATCHER(OP_NAME, ELEM_TYPE)                                                                          \
 static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, const uint32_t num_elems) { \
-    if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
-        OP_NAME##_aa(dst, src, val, num_elems); \
-    } else if (hex_is_aligned((void *) dst, 128)) { \
-        OP_NAME##_au(dst, src, val, num_elems); \
-    } else if (hex_is_aligned((void *) src, 128)) { \
-        OP_NAME##_ua(dst, src, val, num_elems); \
-    } else { \
-        OP_NAME##_uu(dst, src, val, num_elems); \
-    } \
+    if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) {                                                 \
+        OP_NAME##_aa(dst, src, val, num_elems);                                                                                   \
+    } else if (hex_is_aligned((void *) dst, 128)) {                                                                               \
+        OP_NAME##_au(dst, src, val, num_elems);                                                                                   \
+    } else if (hex_is_aligned((void *) src, 128)) {                                                                               \
+        OP_NAME##_ua(dst, src, val, num_elems);                                                                                   \
+    } else {                                                                                                                      \
+        OP_NAME##_uu(dst, src, val, num_elems);                                                                                   \
+    }                                                                                                                             \
 }

 HVX_BINARY_SCALAR_DISPATCHER(hvx_add_scalar_f32, float)
@@ -350,12 +350,12 @@ static inline void hvx_max_scalar_f32(uint8_t * restrict dst, const uint8_t * re

 // CLAMP Scalar variants

-#define HVX_OP_CLAMP_SCALAR(v) \
-    ({ \
+#define HVX_OP_CLAMP_SCALAR(v)                                           \
+    ({                                                                   \
         HVX_VectorPred pred_cap_right = Q6_Q_vcmp_gt_VsfVsf(v, max_vec); \
         HVX_VectorPred pred_cap_left  = Q6_Q_vcmp_gt_VsfVsf(min_vec, v); \
-        HVX_Vector tmp = Q6_V_vmux_QVV(pred_cap_right, max_vec, v); \
-        Q6_V_vmux_QVV(pred_cap_left, min_vec, tmp); \
+        HVX_Vector tmp = Q6_V_vmux_QVV(pred_cap_right, max_vec, v);      \
+        Q6_V_vmux_QVV(pred_cap_left, min_vec, tmp);                      \
     })

 static inline void hvx_clamp_scalar_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const float min, const float max, uint32_t n) {
diff --git a/ggml/src/ggml-hexagon/htp/hvx-div.h b/ggml/src/ggml-hexagon/htp/hvx-div.h
index 53ee304e7..bb7ab0519 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-div.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-div.h
@@ -219,64 +219,64 @@ static inline HVX_Vector hvx_vec_hybrid_div_f16(HVX_Vector vec1, HVX_Vector vec2
     } while(0)

 // Generic macro to define alignment permutations for an op
-#define DEFINE_HVX_DIV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
+#define DEFINE_HVX_DIV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY)                                                                            \
 static inline void OP_NAME##_aaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src0 % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_Vector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                                                              \
+    assert((uintptr_t) src0 % 128 == 0);                                                                                             \
+    assert((uintptr_t) src1 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_Vector, hvx_vec_store_a);                                                               \
+}                                                                                                                                    \
 static inline void OP_NAME##_aau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src0 % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_UVector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                                                              \
+    assert((uintptr_t) src0 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_UVector, hvx_vec_store_a);                                                              \
+}                                                                                                                                    \
 static inline void OP_NAME##_aua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_Vector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                                                              \
+    assert((uintptr_t) src1 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_Vector, hvx_vec_store_a);                                                              \
+}                                                                                                                                    \
 static inline void OP_NAME##_auu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_UVector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                                                              \
+    OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_UVector, hvx_vec_store_a);                                                             \
+}                                                                                                                                    \
 static inline void OP_NAME##_uaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) src0 % 128 == 0); \
-    assert((uintptr_t) src1 % 128 == 0); \
-    OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_Vector, hvx_vec_store_u); \
-} \
+    assert((uintptr_t) src0 % 128 == 0);                                                                                             \
+    assert((uintptr_t) src1 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_Vector, hvx_vec_store_u);                                                              \
+}                                                                                                                                    \
 static inline void OP_NAME##_uau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) src0 % 128 == 0); \
-    OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_UVector, hvx_vec_store_u); \
-} \
+    assert((uintptr_t) src0 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_UVector, hvx_vec_store_u);                                                             \
+}                                                                                                                                    \
 static inline void OP_NAME##_uua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    assert((uintptr_t) src1 % 128 == 0); \
-    OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_Vector, hvx_vec_store_u); \
-} \
+    assert((uintptr_t) src1 % 128 == 0);                                                                                             \
+    OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_Vector, hvx_vec_store_u);                                                             \
+}                                                                                                                                    \
 static inline void OP_NAME##_uuu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
-    OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_UVector, hvx_vec_store_u); \
-} \
+    OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_UVector, hvx_vec_store_u);                                                            \
+}                                                                                                                                    \

 // Dispatcher logic
-#define HVX_DIV_DISPATCHER(OP_NAME) \
+#define HVX_DIV_DISPATCHER(OP_NAME)                                                                                                          \
 static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, const uint32_t num_elems) { \
-    if (hex_is_aligned((void *) dst, 128)) { \
-        if (hex_is_aligned((void *) src0, 128)) { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_aau(dst, src0, src1, num_elems); \
-        } else { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_auu(dst, src0, src1, num_elems); \
-        } \
-    } else { \
-        if (hex_is_aligned((void *) src0, 128)) { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_uau(dst, src0, src1, num_elems); \
-        } else { \
-            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
-            else                                    OP_NAME##_uuu(dst, src0, src1, num_elems); \
-        } \
-    } \
+    if (hex_is_aligned((void *) dst, 128)) {                                                                                                 \
+        if (hex_is_aligned((void *) src0, 128)) {                                                                                            \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_aau(dst, src0, src1, num_elems);                                               \
+        } else {                                                                                                                             \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_auu(dst, src0, src1, num_elems);                                               \
+        }                                                                                                                                    \
+    } else {                                                                                                                                 \
+        if (hex_is_aligned((void *) src0, 128)) {                                                                                            \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_uau(dst, src0, src1, num_elems);                                               \
+        } else {                                                                                                                             \
+            if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems);                                               \
+            else                                    OP_NAME##_uuu(dst, src0, src1, num_elems);                                               \
+        }                                                                                                                                    \
+    }                                                                                                                                        \
 }

 DEFINE_HVX_DIV_OP_VARIANTS(hvx_div_f32, hvx_div_f32_loop_body)
diff --git a/ggml/src/ggml-hexagon/htp/hvx-inverse.h b/ggml/src/ggml-hexagon/htp/hvx-inverse.h
index f2054f45b..256a8843b 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-inverse.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-inverse.h
@@ -169,36 +169,36 @@ static inline HVX_Vector hvx_vec_inverse_f16_guard(HVX_Vector v_sf, HVX_Vector n
     } while(0)

 // Generic macro to define alignment permutations for an op
-#define DEFINE_HVX_INV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
+#define DEFINE_HVX_INV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY)                                           \
 static inline void OP_NAME##_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    assert((uintptr_t) src % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_Vector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                             \
+    assert((uintptr_t) src % 128 == 0);                                                             \
+    OP_LOOP_BODY(HVX_Vector, HVX_Vector, hvx_vec_store_a);                                          \
+}                                                                                                   \
 static inline void OP_NAME##_au(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
-    assert((uintptr_t) dst % 128 == 0); \
-    OP_LOOP_BODY(HVX_Vector, HVX_UVector, hvx_vec_store_a); \
-} \
+    assert((uintptr_t) dst % 128 == 0);                                                             \
+    OP_LOOP_BODY(HVX_Vector, HVX_UVector, hvx_vec_store_a);                                         \
+}                                                                                                   \
 static inline void OP_NAME##_ua(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
-    assert((uintptr_t) src % 128 == 0); \
-    OP_LOOP_BODY(HVX_UVector, HVX_Vector, hvx_vec_store_u); \
-} \
+    assert((uintptr_t) src % 128 == 0);                                                             \
+    OP_LOOP_BODY(HVX_UVector, HVX_Vector, hvx_vec_store_u);                                         \
+}                                                                                                   \
 static inline void OP_NAME##_uu(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
-    OP_LOOP_BODY(HVX_UVector, HVX_UVector, hvx_vec_store_u); \
-} \
+    OP_LOOP_BODY(HVX_UVector, HVX_UVector, hvx_vec_store_u);                                        \
+}                                                                                                   \

 // Dispatcher logic
-#define HVX_INV_DISPATCHER(OP_NAME) \
+#define HVX_INV_DISPATCHER(OP_NAME)                                                                          \
 static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src, const uint32_t num_elems) { \
-    if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
-        OP_NAME##_aa(dst, src, num_elems); \
-    } else if (hex_is_aligned((void *) dst, 128)) { \
-        OP_NAME##_au(dst, src, num_elems); \
-    } else if (hex_is_aligned((void *) src, 128)) { \
-        OP_NAME##_ua(dst, src, num_elems); \
-    } else { \
-        OP_NAME##_uu(dst, src, num_elems); \
-    } \
+    if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) {                            \
+        OP_NAME##_aa(dst, src, num_elems);                                                                   \
+    } else if (hex_is_aligned((void *) dst, 128)) {                                                          \
+        OP_NAME##_au(dst, src, num_elems);                                                                   \
+    } else if (hex_is_aligned((void *) src, 128)) {                                                          \
+        OP_NAME##_ua(dst, src, num_elems);                                                                   \
+    } else {                                                                                                 \
+        OP_NAME##_uu(dst, src, num_elems);                                                                   \
+    }                                                                                                        \
 }

 DEFINE_HVX_INV_OP_VARIANTS(hvx_inverse_f32, hvx_inverse_f32_loop_body)
diff --git a/ggml/src/ggml-hexagon/htp/hvx-scale.h b/ggml/src/ggml-hexagon/htp/hvx-scale.h
index 9b1a28f52..5d0650307 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-scale.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-scale.h
@@ -68,30 +68,30 @@ static inline void hvx_scale_f32(uint8_t * restrict dst, const uint8_t * restric
     }
 }

-#define hvx_scale_offset_f32_loop_body(dst_type, src_type, vec_store)                \
-    do {                                                                             \
-        dst_type * restrict vdst = (dst_type *) dst;                                 \
-        src_type * restrict vsrc = (src_type *) src;                                 \
-                                                                                     \
-        HVX_Vector vs = hvx_vec_splat_f32(scale);                                    \
-        HVX_Vector vo = hvx_vec_splat_f32(offset);                                   \
-                                                                                     \
-        const uint32_t elem_size = sizeof(float);                                    \
-        const uint32_t epv = 128 / elem_size;                                        \
-        const uint32_t nvec = n / epv;                                               \
-        const uint32_t nloe = n % epv;                                               \
-                                                                                     \
-        uint32_t i = 0;                                                              \
-                                                                                     \
-        _Pragma("unroll(4)")                                                         \
-        for (; i < nvec; ++i) {                                                      \
+#define hvx_scale_offset_f32_loop_body(dst_type, src_type, vec_store)                     \
+    do {                                                                                  \
+        dst_type * restrict vdst = (dst_type *) dst;                                      \
+        src_type * restrict vsrc = (src_type *) src;                                      \
+                                                                                          \
+        HVX_Vector vs = hvx_vec_splat_f32(scale);                                         \
+        HVX_Vector vo = hvx_vec_splat_f32(offset);                                        \
+                                                                                          \
+        const uint32_t elem_size = sizeof(float);                                         \
+        const uint32_t epv = 128 / elem_size;                                             \
+        const uint32_t nvec = n / epv;                                                    \
+        const uint32_t nloe = n % epv;                                                    \
+                                                                                          \
+        uint32_t i = 0;                                                                   \
+                                                                                          \
+        _Pragma("unroll(4)")                                                              \
+        for (; i < nvec; ++i) {                                                           \
             HVX_Vector v = Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(vsrc[i], vs), vo); \
-            vdst[i] = Q6_Vsf_equals_Vqf32(v);                                        \
-        }                                                                            \
-        if (nloe) {                                                                  \
+            vdst[i] = Q6_Vsf_equals_Vqf32(v);                                             \
+        }                                                                                 \
+        if (nloe) {                                                                       \
             HVX_Vector v = Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(vsrc[i], vs), vo); \
-            vec_store((void *) &vdst[i], nloe * elem_size, Q6_Vsf_equals_Vqf32(v));  \
-        }                                                                            \
+            vec_store((void *) &vdst[i], nloe * elem_size, Q6_Vsf_equals_Vqf32(v));       \
+        }                                                                                 \
     } while(0)

 static inline void hvx_scale_offset_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) {
diff --git a/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h b/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h
index dd66dd84c..552017309 100644
--- a/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h
+++ b/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h
@@ -68,50 +68,50 @@ static inline HVX_Vector hvx_vec_tanh_f32(HVX_Vector x) {
     return Q6_Vsf_equals_Vqf32(res);
 }

-#define hvx_sigmoid_loop_body(dst_type, src_type, vec_store)    \
-    do {                                                        \
-        dst_type * restrict vdst = (dst_type *) dst;            \
-        src_type * restrict vsrc = (src_type *) src;            \
-                                                                \
-        const HVX_Vector one     = hvx_vec_splat_f32(1.f);      \
-        const HVX_Vector max_exp = hvx_vec_splat_f32(87.f);     \
-        const HVX_Vector min_exp = hvx_vec_splat_f32(-87.f);    \
-                                                                \
-        const uint32_t epv  = 128 / sizeof(float);              \
-        const uint32_t nvec = n / epv;                          \
-        const uint32_t nloe = n % epv;                          \
-                                                                \
-        uint32_t i = 0;                                         \
-                                                                \
-        _Pragma("unroll(4)")                                    \
-        for (; i < nvec; i++) {                                 \
-             vdst[i] = hvx_vec_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp); \
-        }                                                       \
-        if (nloe) {                                             \
+#define hvx_sigmoid_loop_body(dst_type, src_type, vec_store)                                  \
+    do {                                                                                      \
+        dst_type * restrict vdst = (dst_type *) dst;                                          \
+        src_type * restrict vsrc = (src_type *) src;                                          \
+                                                                                              \
+        const HVX_Vector one     = hvx_vec_splat_f32(1.f);                                    \
+        const HVX_Vector max_exp = hvx_vec_splat_f32(87.f);                                   \
+        const HVX_Vector min_exp = hvx_vec_splat_f32(-87.f);                                  \
+                                                                                              \
+        const uint32_t epv  = 128 / sizeof(float);                                            \
+        const uint32_t nvec = n / epv;                                                        \
+        const uint32_t nloe = n % epv;                                                        \
+                                                                                              \
+        uint32_t i = 0;                                                                       \
+                                                                                              \
+        _Pragma("unroll(4)")                                                                  \
+        for (; i < nvec; i++) {                                                               \
+             vdst[i] = hvx_vec_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp);        \
+        }                                                                                     \
+        if (nloe) {                                                                           \
              HVX_Vector tmp = hvx_vec_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp); \
-             vec_store((void *) &vdst[i], nloe * sizeof(float), tmp); \
-        }                                                       \
+             vec_store((void *) &vdst[i], nloe * sizeof(float), tmp);                         \
+        }                                                                                     \
     } while(0)

-#define hvx_tanh_loop_body(dst_type, src_type, vec_store)       \
-    do {                                                        \
-        dst_type * restrict vdst = (dst_type *) dst;            \
-        src_type * restrict vsrc = (src_type *) src;            \
-                                                                \
-        const uint32_t epv  = 128 / sizeof(float);              \
-        const uint32_t nvec = n / epv;                          \
-        const uint32_t nloe = n % epv;                          \
-                                                                \
-        uint32_t i = 0;                                         \
-                                                                \
-        _Pragma("unroll(4)")                                    \
-        for (; i < nvec; i++) {                                 \
-             vdst[i] = hvx_vec_tanh_f32(vsrc[i]);               \
-        }                                                       \
-        if (nloe) {                                             \
-             HVX_Vector tmp = hvx_vec_tanh_f32(vsrc[i]);        \
+#define hvx_tanh_loop_body(dst_type, src_type, vec_store)             \
+    do {                                                              \
+        dst_type * restrict vdst = (dst_type *) dst;                  \
+        src_type * restrict vsrc = (src_type *) src;                  \
+                                                                      \
+        const uint32_t epv  = 128 / sizeof(float);                    \
+        const uint32_t nvec = n / epv;                                \
+        const uint32_t nloe = n % epv;                                \
+                                                                      \
+        uint32_t i = 0;                                               \
+                                                                      \
+        _Pragma("unroll(4)")                                          \
+        for (; i < nvec; i++) {                                       \
+             vdst[i] = hvx_vec_tanh_f32(vsrc[i]);                     \
+        }                                                             \
+        if (nloe) {                                                   \
+             HVX_Vector tmp = hvx_vec_tanh_f32(vsrc[i]);              \
              vec_store((void *) &vdst[i], nloe * sizeof(float), tmp); \
-        }                                                       \
+        }                                                             \
     } while(0)

 static inline void hvx_sigmoid_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
diff --git a/ggml/src/ggml-hexagon/htp/im2col-ops.c b/ggml/src/ggml-hexagon/htp/im2col-ops.c
index 35fc103df..52bbc37d1 100644
--- a/ggml/src/ggml-hexagon/htp/im2col-ops.c
+++ b/ggml/src/ggml-hexagon/htp/im2col-ops.c
@@ -3,11 +3,12 @@
 #pragma clang diagnostic ignored "-Wunused-but-set-variable"

 #include <HAP_farf.h>
-#include <HAP_perf.h>
 #include <hexagon_protos.h>
 #include <hexagon_types.h>
 #include <string.h>

+#include "hex-common.h"
+
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "htp-ctx.h"
@@ -16,14 +17,19 @@
 #include "hex-dma.h"
 #include "hex-profile.h"
 #include "htp-vtcm.h"
+#include "htp-tensor.h"

 struct htp_im2col_context {
     struct htp_ops_context * octx;
+    uint32_t                 patch_base;           // first patch index assigned to this dev
+    uint32_t                 npatches;             // number of patches assigned to this dev
     uint32_t                 npatches_per_thread;  // patches = N*OH*OW (pure-DDR kernel)

-    uint32_t pe_rows_per_thread;                   // N*OH rows per worker
-    uint32_t pe_src_row_bytes;                     // one output row's source: IC*KH*IW*4, rounded 256
-    uint32_t pe_dst_row_bytes;                     // one output row's dst: OW*patch_stride*2, rounded 256
+    uint32_t pe_row_base;              // first N*OH row index assigned to this dev (DMA path)
+    uint32_t pe_nrows;                 // number of N*OH rows assigned to this dev (DMA path)
+    uint32_t pe_rows_per_thread;       // N*OH rows per worker
+    uint32_t pe_src_row_bytes;         // one output row's source: IC*KH*IW*4, rounded 256
+    uint32_t pe_dst_row_bytes;         // one output row's dst: OW*patch_stride*2, rounded 256

     // Patch-embed DMA path VTCM ping-pong.
     uint8_t * pe_vtcm_src;                         // base of the 2x src buffers region
@@ -58,33 +64,27 @@ static inline void htp_im2col_vtcm_layout_build(struct htp_im2col_vtcm_layout *
         struct htp_im2col_context * ictx        = (struct htp_im2col_context *) data;                     \
         struct htp_ops_context *    octx        = ictx->octx;                                             \
         struct htp_thread_trace * restrict tr   = &octx->ctx->trace[ith];                                 \
+        const struct htp_tensor * restrict src0 = octx->src[0];                                           \
         const struct htp_tensor * restrict src1 = octx->src[1];                                           \
         const struct htp_tensor * restrict dst  = octx->dst;                                              \
-        const int32_t  s0                       = octx->op_params[0];                                     \
-        const int32_t  s1                       = octx->op_params[1];                                     \
-        const int32_t  p0                       = octx->op_params[2];                                     \
-        const int32_t  p1                       = octx->op_params[3];                                     \
-        const int32_t  d0                       = octx->op_params[4];                                     \
-        const int32_t  d1                       = octx->op_params[5];                                     \
-        const uint32_t N                        = src1->ne[3];                                            \
-        const uint32_t IC                       = src1->ne[2];                                            \
-        const uint32_t IH                       = src1->ne[1];                                            \
-        const uint32_t IW                       = src1->ne[0];                                            \
-        const uint32_t KH                       = octx->src[0]->ne[1];                                    \
-        const uint32_t KW                       = octx->src[0]->ne[0];                                    \
+        const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1];                                   \
+        const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3];                                   \
+        const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5];                                   \
+        const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0];             \
+        const uint32_t KH                       = src0->ne[1], KW = src0->ne[0];                          \
         const uint32_t OH                       = dst->ne[2];                                             \
         const uint32_t OW                       = dst->ne[1];                                             \
         const uint32_t patch_stride             = IC * KH * KW;                                           \
         const float * restrict src_data         = (const float *) src1->data;                             \
         DST_CTYPE * restrict dst_data           = (DST_CTYPE *) dst->data;                                \
-        const uint32_t npatches                 = N * OH * OW;                                            \
-        const uint32_t patch_start              = ictx->npatches_per_thread * ith;                        \
-        const uint32_t patch_end                = MIN(patch_start + ictx->npatches_per_thread, npatches); \
-        if (patch_start >= patch_end) {                                                                   \
+        const uint32_t patch_end                = ictx->patch_base + ictx->npatches;                      \
+        const uint32_t patch_start              = ictx->patch_base + ictx->npatches_per_thread * ith;     \
+        const uint32_t patch_stop               = MIN(patch_start + ictx->npatches_per_thread, patch_end);\
+        if (patch_start >= patch_stop) {                                                                  \
             return;                                                                                       \
         }                                                                                                 \
         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start);                                   \
-        for (uint32_t p = patch_start; p < patch_end; p++) {                                              \
+        for (uint32_t p = patch_start; p < patch_stop; p++) {                                             \
             const uint32_t iow             = p % OW;                                                      \
             const uint32_t ioh             = (p / OW) % OH;                                               \
             const uint32_t in              = p / (OW * OH);                                               \
@@ -154,10 +154,10 @@ IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx
         uint8_t *      dst_base         = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread;                    \
         float *        srcb             = (float *) src_base;                                                        \
         DST_CTYPE *    dstb             = (DST_CTYPE *) dst_base;                                                    \
-        const uint32_t nrows            = N * OH;                                                                    \
+        const uint32_t row_end_max      = ictx->pe_row_base + ictx->pe_nrows;                                        \
         const uint32_t per_thread       = ictx->pe_rows_per_thread;                                                  \
-        const uint32_t row_start        = per_thread * ith;                                                          \
-        const uint32_t row_end          = MIN(row_start + per_thread, nrows);                                        \
+        const uint32_t row_start        = ictx->pe_row_base + per_thread * ith;                                      \
+        const uint32_t row_end          = MIN(row_start + per_thread, row_end_max);                                  \
         if (row_start >= row_end)                                                                                    \
             return;                                                                                                  \
         for (uint32_t r = row_start; r < row_end; r++) {                                                             \
@@ -266,26 +266,55 @@ int op_im2col(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    const uint32_t N         = src1->ne[3];
-    const uint32_t OH        = dst->ne[2];
-    const uint32_t OW        = dst->ne[1];
-    const uint32_t npatches  = N * OH * OW;
-    const uint32_t n_threads = MIN(octx->n_threads, npatches);
+    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t N             = src1->ne[3];
+    const uint32_t OH            = dst->ne[2];
+    const uint32_t OW            = dst->ne[1];
+    const uint32_t total_patches = N * OH * OW;
+    const uint32_t total_rows    = N * OH;
+
+    uint32_t patch_base = 0;
+    uint32_t npatches   = total_patches;
+    if (octx->ctx->mdev.count > 1) {
+        const uint32_t patch_size = dst->nb[1];
+        const uint32_t patches_per_chunk = (patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1;
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        patch_base = range.start;
+        npatches   = range.count;
+    }
+
+    uint32_t row_base = 0;
+    uint32_t nrows    = total_rows;
+    if (octx->ctx->mdev.count > 1) {
+        const uint32_t row_size = dst->nb[2];
+        const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_base = range.start;
+        nrows    = range.count;
+    }

-    if ((octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) || n_threads == 0) {
+    if (npatches == 0 && nrows == 0) {
         return HTP_STATUS_OK;
     }

+    const uint32_t n_threads = MIN(octx->n_threads, MAX(npatches, 1));
+
     struct htp_im2col_context ictx = { 0 };
-    ictx.octx                      = octx;
-    ictx.npatches_per_thread       = (npatches + n_threads - 1) / n_threads;
+    ictx.octx                = octx;
+    ictx.patch_base          = patch_base;
+    ictx.npatches            = npatches;
+    ictx.npatches_per_thread = (npatches + n_threads - 1) / n_threads;

     // Clean non-overlapping patch-embed -> DMA kernel (if it fits VTCM);
     // everything else (padding/dilation/stride edges) -> pure-DDR kernel.
-    if (im2col_use_patchembed_dma(octx)) {
-        const uint32_t nrows = N * OH;
-        const uint32_t pth   = MIN(octx->n_threads, nrows);
+    if (im2col_use_patchembed_dma(octx) && nrows > 0) {
+        const uint32_t pth = MIN(octx->n_threads, nrows);
         if (pth > 0 && im2col_patchembed_dma_fits(octx, &ictx, pth)) {
+            ictx.pe_row_base        = row_base;
+            ictx.pe_nrows           = nrows;
             ictx.pe_rows_per_thread = (nrows + pth - 1) / pth;
             if (dst->type == HTP_TYPE_F16) {
                 work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_thread, &ictx, pth);
@@ -297,6 +326,10 @@ int op_im2col(struct htp_ops_context * octx) {
         // else: doesn't fit -> fall through to the pure-DDR kernel below.
     }

+    if (npatches == 0) {
+        return HTP_STATUS_OK;
+    }
+
     if (dst->type == HTP_TYPE_F16) {
         work_queue_run(octx->ctx->work_queue, im2col_patchembed_thread, &ictx, n_threads);
     } else {
diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c
index be54d4fe9..1d291e16b 100644
--- a/ggml/src/ggml-hexagon/htp/main.c
+++ b/ggml/src/ggml-hexagon/htp/main.c
@@ -34,6 +34,7 @@
 #include "work-queue.h"
 #include "hex-profile.h"
 #include "allreduce-ops.h"
+#include "htp-fence.h"

 #define HMX_QUEUE_CAPACITY     16
 #define HMX_QUEUE_STACK_SIZE   16384
@@ -710,22 +711,43 @@ static inline void profile_stop(uint32_t mode, struct profile_data * d) {
 static int op_fence(struct htp_ops_context * octx) {
     struct htp_context *ctx = octx->ctx;
     struct htp_thread_trace * tr = &ctx->trace[0];
-    const uint32_t seq = (uint32_t) octx->op_params[0];
+    const uint32_t seq  = (uint32_t) octx->op_params[0];
+    const uint32_t mode = (uint32_t) octx->op_params[1];

     htp_trace_event_start(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);

     const struct htp_tensor * sync = octx->src[0];
-    atomic_uint * sync_fence = (atomic_uint *) sync->data;
+    atomic_uint * sync_fence = (atomic_uint *) (uintptr_t) sync->data;
+
+    if (mode == 1) {
+        htp_flush_dirty_ranges(ctx);
+
+        htp_mdev_group_barrier(octx);
+
+        if (ctx->mdev.idx == 0) {
+            htp_fence_write(sync_fence, seq, octx->status);
+        }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
+        FARF(HIGH, "ggml-hex: sync-signal : fence %p seq 0x%x status %d\n", sync_fence, seq, octx->status);
+        return octx->status;
+    }
+
+    int status = HTP_STATUS_OK;
     uint64_t spins = 0;
     while (1) {
-        Q6_dccleaninva_A((void *) sync_fence);
-        asm volatile ("syncht" : : : "memory");
-        uint32_t val = atomic_load(&sync_fence[0]);
-        if ((int32_t)(val - seq) >= 0) {
+        uint32_t sync_seq;
+        uint32_t sync_status;
+        htp_fence_read(sync_fence, &sync_seq, &sync_status);
+        if ((int32_t)(sync_seq - seq) >= 0) {
+            if (sync_status > HTP_STATUS_OK) {
+                FARF(ERROR, "ggml-hex: sync-wait peer failed with status %u : fence %p seq 0x%x\n", sync_status, sync_fence, seq);
+                status = sync_status;
+            }
             break;
         }
         if (++spins > HTP_FENCE_TIMEOUT) {
-            FARF(ERROR, "ggml-hex: sync-wait TIMEOUT : fence %p spins %llu seq %u\n", sync_fence, spins, seq);
+            FARF(ERROR, "ggml-hex: sync-wait TIMEOUT : fence %p spins %llu seq 0x%x\n", sync_fence, spins, seq);
+            status = HTP_STATUS_INTERNAL_ERR;
             break;
         }
         hex_pause();
@@ -733,12 +755,27 @@ static int op_fence(struct htp_ops_context * octx) {

     htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);

-    FARF(HIGH, "ggml-hex: sync-done : fence %p spins %llu seq %u\n", sync_fence, spins, seq);
+    FARF(HIGH, "ggml-hex: sync-done : fence %p spins %llu seq 0x%x\n", sync_fence, spins, seq);
+    return status;
+}
+
+static int op_mdev_group(struct htp_ops_context * octx) {
+    struct htp_context * ctx = octx->ctx;
+    const struct htp_tensor * sync = octx->src[0];
+    ctx->mdev.idx   = (uint16_t) octx->op_params[0];
+    ctx->mdev.count = (uint16_t) sync->ne[1];
+    if (ctx->mdev.count > 1) {
+        ctx->mdev.count_div = init_fastdiv_values(ctx->mdev.count);
+        ctx->mdev.fence_base = (uint8_t *) sync->data;
+    }
     return HTP_STATUS_OK;
 }

 static int execute_op(struct htp_ops_context * octx) {
     switch (octx->op) {
+        case HTP_OP_MDEV_GROUP:
+            return op_mdev_group(octx);
+
         case HTP_OP_FENCE:
             return op_fence(octx);

@@ -812,6 +849,7 @@ static int execute_op(struct htp_ops_context * octx) {
             return op_sum_rows(octx);

         case HTP_OP_CPY:
+        case HTP_OP_CPY_FENCE:
             return op_cpy(octx);

         case HTP_OP_REPEAT:
@@ -855,7 +893,7 @@ static int execute_op(struct htp_ops_context * octx) {
     }

     FARF(ERROR, "Unknown Op %u", octx->op);
-    return -1;
+    return HTP_STATUS_NO_SUPPORT;
 }

 static inline bool reuse_buf(struct htp_context *ctx, uint32_t *m_reuse, struct htp_buf_desc *b) {
@@ -984,11 +1022,19 @@ static void prep_tensors(struct htp_context *ctx, struct htp_buf_desc *bufs, str
     }
 }

-static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, uint32_t idx, struct htp_op_desc * op) {
-    memcpy(octx->op_params, op->params, sizeof(octx->op_params));
+static void mdev_group_init(struct htp_context * ctx, const struct htp_opbatch_req * req) {
+    memset(&ctx->mdev, 0, sizeof(ctx->mdev));
+    ctx->mdev.fence_seq = (uint32_t)((req->seq & 0xfffff) << 12);
+}
+
+static int proc_op_req(struct htp_ops_context * octx, struct htp_buf_desc * bufs, uint32_t n_bufs,
+                       struct htp_tensor * tens, uint32_t idx, struct htp_op_desc * op) {
+    memcpy(octx->op_params,     op->params, sizeof(octx->op_params));
     memcpy(octx->kernel_params, op->kernel_params, sizeof(octx->kernel_params));
-    octx->flags = op->flags;
-    octx->op    = op->opcode;
+    octx->flags         = op->flags;
+    octx->op            = op->opcode;
+    octx->n_threads     = octx->ctx->n_threads;
+    octx->n_threads_div = octx->ctx->n_threads_div;

     FARF(HIGH, "proc-op #%u: opcode %u flags 0x%x", idx, octx->op, octx->flags);

@@ -1027,9 +1073,13 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u
             dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
     }

+    htp_tensor_dirty_all(octx->ctx, octx->dsts, HTP_OP_MAX_OUTPUTS);
+
+    htp_mdev_group_barrier(octx);
+
     int status = execute_op(octx);

-    htp_tensor_dirty_all(octx->ctx, octx->dsts, HTP_OP_MAX_OUTPUTS);
+    htp_ops_context_set_status(octx, status);

     octx->src0_spad.src = NULL;
     octx->src1_spad.src = NULL;
@@ -1037,7 +1087,7 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u
     octx->src3_spad.src = NULL;
     octx->dst_spad.src  = NULL;

-    return status;
+    return octx->status;
 }

 static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_req * req, const struct dspqueue_buffer * dbuf) {
@@ -1059,7 +1109,7 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
         return;
     }

-    FARF(HIGH, "processing opbatch #%u: n-bufs %u n-tensors %u n-ops %u n-traces %u : m-size %u b-size %u t-size %u o-size %u", req->id,
+    FARF(HIGH, "processing opbatch #%llu: n-bufs %u n-tensors %u n-ops %u n-traces %u : m-size %u b-size %u t-size %u o-size %u", (unsigned long long) req->seq,
             n_bufs, n_tens, n_ops, req->n_traces, dbuf->size, b_size, t_size, o_size);

     // Setup descriptor pointers
@@ -1096,8 +1146,11 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r

     struct htp_ops_context *octx = &ctx->octx;
     memset(octx, 0, sizeof(*octx));
-    octx->n_threads = ctx->n_threads;
-    octx->ctx       = ctx;
+    octx->n_threads     = ctx->n_threads;
+    octx->n_threads_div = ctx->n_threads_div;
+    octx->ctx           = ctx;
+
+    mdev_group_init(ctx, req);

     work_queue_wakeup(ctx->work_queue);
     if (ctx->hmx_queue) {
@@ -1105,15 +1158,18 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
     }

     int op_status = HTP_STATUS_OK;
-    for (uint32_t i = 0; i < n_ops && op_status == HTP_STATUS_OK; i++) {
+    octx->status  = HTP_STATUS_OK;
+    for (uint32_t i = 0; i < n_ops; i++) {
         struct profile_data prof;

         profile_start(ctx->profiler, &prof);

-        op_status = proc_op_req(octx, tens, i, &ops[i]);
+        op_status = proc_op_req(octx, bufs, n_bufs, tens, i, &ops[i]);

         profile_stop(ctx->profiler, &prof);

+        htp_ops_context_set_status(octx, op_status);
+
         if (ctx->profiler) {
             pds[i].opcode = ops[i].opcode;
             pds[i].usecs  = prof.usecs;
@@ -1136,19 +1192,20 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
     qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE);
     htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0);

+    htp_mdev_group_barrier(octx);
+
     profile_stop(HTP_PROF_BASIC, &batch_prof);

     struct htp_opbatch_rsp rsp;
     memset(&rsp, 0, sizeof(rsp));
-    rsp.id           = req->id;
-    rsp.status       = op_status;
+    rsp.seq          = req->seq;
+    rsp.status       = octx->status;
     rsp.n_bufs       = n_bufs;
     rsp.n_tensors    = n_tens;
     rsp.n_ops        = n_ops;
     rsp.usecs        = batch_prof.usecs;
     rsp.cycles_start = batch_prof.cycles_start;
     rsp.cycles_stop  = batch_prof.cycles_stop;
-    rsp.seq          = req->seq;

     if (ctx->profiler == HTP_PROF_TRACE) {
         for (int t = 0; t <= HTP_MAX_NTHREADS; t++) {
diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c
index 2a87dd19e..1b597dcd9 100644
--- a/ggml/src/ggml-hexagon/htp/matmul-ops.c
+++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c
@@ -21,6 +21,7 @@
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"
 #include "matmul-ops.h"
 #include "htp-vtcm.h"

@@ -89,6 +90,8 @@ struct htp_mm_context {

     // Precomputed values
     uint32_t src0_nrows_per_thread;
+    uint32_t src0_row_start;
+    uint32_t src0_row_end;
     uint32_t src0_row_size_padded;
     uint32_t src1_nrows;

@@ -135,6 +138,23 @@ struct htp_mm_context {
     uint32_t vtcm_dst_size_per_thread;
 };

+static int htp_mm_init_context(
+    struct htp_ops_context * octx,
+    const struct htp_mm_kernel_params * kparams
+) {
+    if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    if (kparams->n_hmx) {
+        if (kparams->n_act_threads <= 0 || kparams->n_act_threads > (int32_t) octx->n_threads) {
+            return HTP_STATUS_INVAL_PARAMS;
+        }
+    }
+
+    return HTP_STATUS_OK;
+}
+
 // vdelta control to expand first 32 e8m0 values into 32 uint32 elements
 static const uint8_t __attribute__((aligned(128))) expand_x32_e8m0[128] = {
     0x00, 0x00, 0x00, 0x00, 0x01, 0x04, 0x00, 0x00, 0x02, 0x00, 0x08, 0x08, 0x01, 0x02, 0x00, 0x04, 0x04, 0x00, 0x00,
@@ -238,22 +258,24 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
     // This is the size of the rest of the dimensions of the result
     const uint32_t nr1 = ne1 * ne2 * ne3;

+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;
+
     // distribute the thread work across the inner or outer loop based on which one is larger
     uint32_t dr0, dr1, ith0, ith1;
     if (nr0 > nr1) {
-        dr0  = fastdiv(nr0 + nth - 1, &octx->ctx->n_threads_div);
+        dr0  = fastdiv(src0_nrows + nth - 1, &octx->n_threads_div);
         dr1  = nr1;
         ith0 = ith;
         ith1 = 0;
     } else {
-        dr0  = nr0;
-        dr1  = fastdiv(nr1 + nth - 1, &octx->ctx->n_threads_div);
+        dr0  = src0_nrows;
+        dr1  = fastdiv(nr1 + nth - 1, &octx->n_threads_div);
         ith0 = 0;
         ith1 = ith;
     }

-    const uint32_t ir0_start = dr0 * ith0;
-    const uint32_t ir0_end   = MIN(ir0_start + dr0, nr0);
+    const uint32_t ir0_start = mmctx->src0_row_start + dr0 * ith0;
+    const uint32_t ir0_end   = MIN(ir0_start + dr0, mmctx->src0_row_end);

     const uint32_t ir1_start = dr1 * ith1;
     const uint32_t ir1_end   = MIN(ir1_start + dr1, nr1);
@@ -312,11 +334,11 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
 static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                        \
     htp_matmul_preamble;                                                                                                          \
                                                                                                                                   \
-    const uint32_t src0_nrows = ne01 * ne02 * ne03;                                                                               \
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                      \
     const uint32_t src1_nrows = ne11 * ne12 * ne13;                                                                               \
                                                                                                                                   \
-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;                                                                 \
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);                                     \
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                         \
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                            \
                                                                                                                                   \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                        \
                                                                                                                                   \
@@ -414,10 +436,10 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
 static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) {                                        \
     htp_matmul_preamble;                                                                                                          \
                                                                                                                                   \
-    const uint32_t src0_nrows = ne01;                                                                                             \
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;                                                      \
                                                                                                                                   \
-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;                                                                 \
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);                                     \
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;                                         \
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);                            \
                                                                                                                                   \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                                        \
                                                                                                                                   \
@@ -549,12 +571,22 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
         uint32_t n_k_tiles_w = ne00 / 32;                                                                                         \
         uint32_t tile_row_stride = n_k_tiles_w * tile_size;                                                                       \
                                                                                                                                   \
-        const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3];                                                           \
-        uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);                                \
+        uint32_t src0_start_row = 0;                                                                                              \
+        uint32_t src0_end_row   = ne01;                                                                                           \
+        if (octx->ctx->mdev.count > 1) {                                                                                          \
+            const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));                                              \
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0,                        \
+                                                         octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); \
+            src0_start_row = range.start;                                                                                         \
+            src0_end_row   = range.start + range.count;                                                                           \
+        }                                                                                                                         \
+                                                                                                                                  \
+        const uint32_t nrows = src0_end_row - src0_start_row;                                                                     \
+        uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);                                          \
         src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32);                                                          \
                                                                                                                                   \
-        const uint32_t start_row = src0_nrows_per_thread * ith;                                                                   \
-        const uint32_t end_row   = MIN(start_row + src0_nrows_per_thread, src0_nrows);                                            \
+        const uint32_t start_row = src0_start_row + src0_nrows_per_thread * ith;                                                  \
+        const uint32_t end_row   = MIN(start_row + src0_nrows_per_thread, src0_end_row);                                          \
         if (start_row >= end_row) continue;                                                                                       \
                                                                                                                                   \
         uint32_t ct_start = start_row / 32;                                                                                       \
@@ -735,11 +767,11 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
     assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
     const uint32_t prefetch_mask = n_prefetch - 1;

-    const uint32_t src0_nrows = ne01 * ne02 * ne03;  // src0 rows
-    const uint32_t src1_nrows = ne11 * ne12 * ne13;  // src1 rows
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;  // src0 rows
+    const uint32_t src1_nrows = ne11 * ne12 * ne13;                          // src1 rows

-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
     const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];
@@ -781,7 +813,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
         const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;

         htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
-        // Process src1 columns in pairs (2×2 tiling)
+        // Process src1 columns in pairs (2x2 tiling)
         uint32_t ir1 = 0;
         for (; ir1 + 1 < src1_nrows; ir1 += 2) {
             const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);
@@ -791,7 +823,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
             mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
         }

-        // Handle remaining src1 rows (fallback to 2×1)
+        // Handle remaining src1 rows (fallback to 2x1)
         for (; ir1 < src1_nrows; ++ir1) {
             const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
             float * restrict dst_row          = (float *) (dst->data + (ir1 * dst_row_size));
@@ -833,10 +865,10 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
 static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
     htp_matmul_preamble;

-    const uint32_t src0_nrows = ne01;
+    const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;

-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);

     struct htp_thread_trace * tr = &octx->ctx->trace[ith];

@@ -943,13 +975,10 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {

     const struct htp_tensor * restrict ids = octx->src[2];

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
-    const uint32_t src0_nrows      = ne01;  // src0 rows per expert
+    const uint32_t src0_nrows      = mmctx->src0_row_end - mmctx->src0_row_start;  // src0 rows per expert
     const uint32_t src1_nrows      = ne11;
-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);

     hvx_mm_run_quant_task(mmctx, ith);

@@ -1036,9 +1065,9 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {

     const struct htp_tensor * restrict ids = octx->src[2];

-    const uint32_t src0_nrows      = ne01;  // src0 rows per expert
-    const uint32_t src0_start_row  = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_nrows      = mmctx->src0_row_end - mmctx->src0_row_start;  // src0 rows per expert
+    const uint32_t src0_start_row  = mmctx->src0_row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);

     hvx_mm_run_quant_task(mmctx, ith);

@@ -1143,12 +1172,22 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const struct htp_tensor * restrict dst   = octx->dsts[p];
             if (!src_w || !dst) continue;

-            const uint32_t src0_nrows = src_w->ne[1];
-            uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
+            const uint32_t ne01 = src_w->ne[1];
+            uint32_t start_row = 0;
+            uint32_t end_row   = ne01;
+            if (octx->ctx->mdev.count > 1) {
+                const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+                const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+                start_row = range.start;
+                end_row   = range.start + range.count;
+            }
+
+            const uint32_t nrows = end_row - start_row;
+            uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
             src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32);

-            const uint32_t src0_start_row = src0_nrows_per_thread * ith;
-            const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+            const uint32_t src0_start_row = start_row + src0_nrows_per_thread * ith;
+            const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, end_row);
             if (src0_start_row >= src0_end_row) continue;

             const uint8_t * restrict src0_row = (const uint8_t *) src_w->data + eid * src_w->nb[2];
@@ -1227,12 +1266,22 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
             const struct htp_tensor * restrict dst   = octx->dsts[p];
             if (!src_w || !dst) continue;

-            const uint32_t src0_nrows = src_w->ne[1];
-            uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
+            const uint32_t ne01 = src_w->ne[1];
+            uint32_t start_row = 0;
+            uint32_t end_row   = ne01;
+            if (octx->ctx->mdev.count > 1) {
+                const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+                const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+                start_row = range.start;
+                end_row   = range.start + range.count;
+            }
+
+            const uint32_t nrows = end_row - start_row;
+            uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
             src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32);

-            const uint32_t src0_start_row = src0_nrows_per_thread * ith;
-            const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+            const uint32_t src0_start_row = start_row + src0_nrows_per_thread * ith;
+            const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, end_row);
             if (src0_start_row >= src0_end_row) continue;

             const uint8_t * src0_row = (const uint8_t *) src_w->data + cur_a * src_w->nb[2];
@@ -1323,15 +1372,33 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {

     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;

-    const uint32_t src0_nrows = ne01 * ne02 * ne03;
+    const uint32_t src0_nrows = ne01;
     const uint32_t src1_nrows = ne11 * ne12 * ne13;

+    uint32_t src0_row_start = 0;
+    uint32_t src0_row_end   = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        src0_row_start = range.start;
+        src0_row_end   = range.start + range.count;
+    }
+
+    if (src0_row_start >= src0_row_end) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t nrows = src0_row_end - src0_row_start;
+    mmctx->src0_row_start = src0_row_start;
+    mmctx->src0_row_end   = src0_row_end;
+
     bool is_repacked = (src0->type == HTP_TYPE_Q4_0 || src0->type == HTP_TYPE_Q4_1 ||
                         src0->type == HTP_TYPE_Q8_0 || src0->type == HTP_TYPE_IQ4_NL ||
                         src0->type == HTP_TYPE_MXFP4);

     // Compute src0_nrows_per_thread
-    mmctx->src0_nrows_per_thread  = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
+    mmctx->src0_nrows_per_thread  = fastdiv(nrows + octx->n_threads - 1, &octx->n_threads_div);
     if (is_repacked) {
         mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
     } else {
@@ -1503,13 +1570,13 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
         kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
         mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
     } else {
-        mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->ctx->n_threads_div);
+        mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->n_threads_div);
     }

-    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
-    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
+    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

-    size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
+    const size_t vtcm_size = L.total_bytes;

     FARF(HIGH, "matmul-%s : src0-vtcm-size %zu src1-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
          L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
@@ -1583,13 +1650,21 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {

         const uint32_t ne00 = src_w->ne[0];
         const uint32_t ne01 = src_w->ne[1];
-        const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3];
+        uint32_t start_row = 0;
+        uint32_t end_row   = ne01;
+        if (octx->ctx->mdev.count > 1) {
+            const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            start_row = range.start;
+            end_row   = range.start + range.count;
+        }

-        uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
+        const uint32_t nrows = end_row - start_row;
+        uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
         src0_nrows_per_thread += (src0_nrows_per_thread & 1);

-        const uint32_t src0_start_row  = src0_nrows_per_thread * ith;
-        const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+        const uint32_t src0_start_row  = start_row + src0_nrows_per_thread * ith;
+        const uint32_t src0_end_row    = MIN(src0_start_row + src0_nrows_per_thread, end_row);
         const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);
         if (src0_start_row >= src0_end_row) continue;

@@ -2638,10 +2713,6 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
     const struct htp_tensor * restrict src0 = octx->src[0];
     const struct htp_tensor * restrict act  = octx->src[n_weights];

-    if (!src0 || !act) {
-        return HTP_STATUS_INVAL_PARAMS;
-    }
-
     const int weight_type = (int) src0->type;
     const int k           = (int) act->ne[0];
     const int k_valid     = (int) act->ne[0];
@@ -2714,16 +2785,31 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k

     hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00));  // scale: 1.0, bias: 0.0 in FP16

-    FARF(HIGH, "hmx-mm-nx-2d: n_weights %u m %d k %d wtype %d mc %d nc %d vtcm %zu/%zu",
-         n_weights, m, k, weight_type, m_chunk_n_rows, n_chunk_n_cols, L.total_bytes, vtcm_budget);
+    int m_start = 0;
+    int m_rows  = m;
+    if (octx->ctx->mdev.count > 1) {
+        const bool can_split = htp_tensor_can_row_partition(octx->dsts[0], sizeof(float));
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        m_start = (int) range.start;
+        m_rows  = (int) range.count;
+    }
+
+    if (m_rows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    FARF(HIGH, "hmx-mm-nx-2d: n_weights %u m %d (%d..%d) k %d wtype %d mc %d nc %d vtcm %zu/%zu",
+         n_weights, m, m_start, m_start + m_rows, k, weight_type, m_chunk_n_rows, n_chunk_n_cols, L.total_bytes, vtcm_budget);

     htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);

+    const size_t mr_end = (size_t)(m_start + m_rows);
+
     if (pipeline) {
         hmx_matmul_job_t job_slots[2];

-        for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) {
-            const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows);
+        for (size_t mr = (size_t) m_start; mr < mr_end; mr += m_chunk_n_rows) {
+            const size_t n_rows = hex_smin(mr_end - mr, m_chunk_n_rows);

             void *vtcm_weight_bufs[2] = { vtcm_scratch0, vtcm_scratch1 };
             void *vtcm_output_bufs[2] = { vtcm_output,   vtcm_scratch2 };
@@ -2822,8 +2908,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
         }
     } else {
         hmx_matmul_job_t job;
-        for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) {
-            const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows);
+        for (size_t mr = (size_t) m_start; mr < mr_end; mr += m_chunk_n_rows) {
+            const size_t n_rows = hex_smin(mr_end - mr, m_chunk_n_rows);

             struct activation_transfer_params act_params = {
                 .ctx = ctx,
@@ -3095,7 +3181,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
                             int chunk_dst_cols = params->n - (int)nc;
                             if (chunk_dst_cols > 0) {
                                 transfer_output_chunk_threaded(ctx, output, src2_chunk, vtcm_output, (int) n_rows, (int) n_cols,
-                                                               params->dst_stride, params->src2_stride, chunk_dst_cols, ctx->n_threads);
+                                                               params->dst_stride, params->src2_stride, chunk_dst_cols, n_threads);
                             }
                         }
                     }
@@ -3216,7 +3302,10 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
                                          int weight_type,
                                          const struct mmid_row_mapping *matrix_rows,
                                          int cur_a,
-                                         int mapping_stride) {
+                                         int mapping_stride,
+                                         int m_start,
+                                         int m_end,
+                                         int n_threads) {
     struct htp_thread_trace * tr = &ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

@@ -3247,7 +3336,6 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
     const int n_k_tiles = k / HTP_MM_HMX_TILE_N_COLS;
     const struct fastdiv_values n_k_tiles_div = init_fastdiv_values(n_k_tiles);

-    const int n_threads = ctx->n_threads;
     const bool is_quant   = (weight_type != HTP_TYPE_F16 && weight_type != HTP_TYPE_F32);

     const size_t vec_dot_size = k * sizeof(__fp16);
@@ -3303,8 +3391,8 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,

     hmx_matmul_job_t job;

-    for (size_t mr = 0; mr < (size_t) m_padded; mr += m_chunk_n_rows) {
-        const size_t n_rows = hex_smin(m_padded - mr, m_chunk_n_rows);
+    for (size_t mr = (size_t) m_start; mr < (size_t) m_end; mr += m_chunk_n_rows) {
+        const size_t n_rows = hex_smin((size_t) m_end - mr, m_chunk_n_rows);
         const size_t n_row_tiles = hmx_ceil_div(n_rows, HTP_MM_HMX_TILE_N_ROWS);

         transfer_activation_chunk_gathered_threaded(
@@ -3368,31 +3456,48 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
     const int act_stride = (int)(src1->nb[1] / sizeof(float));
     const int wgt_stride = (int)(src0->nb[1] / sizeof(__fp16));

+    int m_start = 0;
+    int m_rows  = m_total;
+    if (octx->ctx->mdev.count > 1) {
+        const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_total, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        m_start = (int) range.start;
+        m_rows  = (int) range.count;
+    }
+
+    if (m_rows == 0) {
+        return HTP_STATUS_OK;
+    }
+
     const float * src2_ptr = NULL;
     uint32_t src2_stride = 0;
     size_t src2_nb2 = 0;
     size_t src2_nb3 = 0;
     if (src2) {
-        src2_ptr = (const float *) src2->data;
         src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
+        src2_ptr = (const float *) src2->data + m_start * src2_stride;
         src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
         src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
     }

+    const int dst_stride = (int)(dst->nb[1] / sizeof(float));
+    float       * dst_ptr = (float *)       dst->data  + m_start * dst_stride;
+    const float * act_ptr = (const float *) src1->data + m_start * act_stride;
+
     int ret = -1;
-    const int n_threads = MIN(kparams->n_threads, (int) octx->n_threads);
+    const int n_threads = kparams->n_threads;
     if (kparams->kernel_type == HTP_MM_KERNEL_HMX_F16_BATCHED) {
         hmx_mm_f16_f32_batched_params_t batch_params = {
-            .dst             = (float *) dst->data,
+            .dst             = dst_ptr,
             .src2            = src2_ptr,
-            .activation      = (float *) src1->data,
+            .activation      = act_ptr,
             .weight          = (const __fp16 *) src0->data,
-            .m               = m_total,
+            .m               = m_rows,
             .k               = k,
             .n               = n,
             .act_stride      = act_stride,
             .weight_stride   = wgt_stride,
-            .dst_stride      = (int) (dst->nb[1] / sizeof(float)),
+            .dst_stride      = dst_stride,
             .src2_stride     = src2_stride,
             .ne02            = ne02,
             .ne03            = ne03,
@@ -3420,9 +3525,9 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
                                      kparams->vtcm_size);
     } else {
         ret = hmx_mm_2d_f32(
-            octx->ctx, (float*) dst->data, src2_ptr, (float*) src1->data, (const uint8_t *) src0->data,
-            m_total, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
-            (int)(dst->nb[1] / sizeof(float)), src2_stride, (int)dst->ne[0],
+            octx->ctx, dst_ptr, src2_ptr, act_ptr, (const uint8_t *) src0->data,
+            m_rows, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
+            dst_stride, src2_stride, (int)dst->ne[0],
             kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
             kparams->n_act_threads,
             &kparams->div_n_act_threads,
@@ -3441,6 +3546,11 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
 int op_matmul(struct htp_ops_context * octx) {
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;

+    const int status = htp_mm_init_context(octx, kparams);
+    if (status != HTP_STATUS_OK) {
+        return status;
+    }
+
     if (kparams->n_hmx) {
         return hmx_mm_op_matmul(octx, kparams);
     }
@@ -3463,6 +3573,16 @@ static int hmx_mm_op_matmul_id(
         const int32_t cne1 = matrix_row_counts[cur_a];
         if (cne1 == 0) continue;

+        const int m_padded = hex_align_up(cne1, 32);
+        int m_start = 0, m_end = m_padded;
+        if (octx->ctx->mdev.count > 1) {
+            const bool can_split = htp_tensor_mdev_data_aligned(dst) && (uint32_t) cne1 >= octx->ctx->mdev.count;
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            m_start = (int) range.start;
+            m_end   = (int) (range.start + range.count);
+        }
+        if (m_start >= m_end) continue;
+
         int ret = hmx_mm_id_2d_f32(octx->ctx, (float*) dst->data, (float*) src1->data,
                                    (const uint8_t *) src0->data + cur_a * nb02,
                                    cne1, ne00, ne01,
@@ -3471,7 +3591,8 @@ static int hmx_mm_op_matmul_id(
                                    nb11, nb12,
                                    nb1, nb2,
                                    (int) src0->nb[1], (int) src0->type,
-                                   matrix_rows, cur_a, mmctx->mapping_stride);
+                                   matrix_rows, cur_a, mmctx->mapping_stride,
+                                   m_start, m_end, (int) octx->n_threads);
         if (ret != 0) {
             FARF(ERROR, "HMX matmul failed for expert %u, error %d\n", cur_a, ret);
             return HTP_STATUS_NO_SUPPORT;
@@ -3524,7 +3645,7 @@ static int hvx_mm_matmul_id(
     htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads,
                                  0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);

-    size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
+    const size_t vtcm_size = L.total_bytes;

     FARF(HIGH, "matmul-id-%s : src0-spad-size %zu src1-spad-size %zu src2-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
          L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
@@ -3554,10 +3675,10 @@ static int hvx_mm_matmul_id(
     mmctx->vtcm_src0_stride = src0_row_size_padded;
     mmctx->vtcm_src1_stride = src1_row_size;

-    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
     mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
     mmctx->vtcm_src2_size_per_thread = 0;
-    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->quant_task_func = quant_task_func;
@@ -3587,6 +3708,20 @@ static int hmx_mm_op_matmul_id_nx(
         const int32_t cne1 = matrix_row_counts[cur_a];
         if (cne1 == 0) continue;

+        const int m_padded = hex_align_up(cne1, 32);
+        int m_start = 0, m_end = m_padded;
+        if (octx->ctx->mdev.count > 1) {
+            bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
+            for (uint32_t p = 0; p < n_weights && can_split; ++p) {
+                const struct htp_tensor * restrict dst = octx->dsts[p];
+                can_split = !dst || htp_tensor_mdev_data_aligned(dst);
+            }
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            m_start = (int) range.start;
+            m_end   = (int) (range.start + range.count);
+        }
+        if (m_start >= m_end) continue;
+
         for (uint32_t p = 0; p < n_weights; ++p) {
             const struct htp_tensor * restrict src_w = octx->src[p];
             const struct htp_tensor * restrict dst   = octx->dsts[p];
@@ -3600,7 +3735,8 @@ static int hmx_mm_op_matmul_id_nx(
                                        act->nb[1], act->nb[2],
                                        dst->nb[1], dst->nb[2],
                                        (int) src_w->nb[1], (int) src_w->type,
-                                       matrix_rows, cur_a, mmctx->mapping_stride);
+                                       matrix_rows, cur_a, mmctx->mapping_stride,
+                                       m_start, m_end, (int) octx->n_threads);
             if (ret != 0) {
                 FARF(ERROR, "HMX matmul ID NX failed for expert %u weight %u, error %d\n", cur_a, p, ret);
                 return HTP_STATUS_NO_SUPPORT;
@@ -3656,7 +3792,7 @@ static int hvx_mm_matmul_id_nx(
     htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
                                  0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);

-    size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
+    const size_t vtcm_size = L.total_bytes;

     if (octx->ctx->vtcm_size < vtcm_size) {
         FARF(ERROR, "matmul-id-nx: current VTCM reservation %zu is too small, needed %zu\n",
@@ -3678,9 +3814,9 @@ static int hvx_mm_matmul_id_nx(
     mmctx->vtcm_src0_stride = 0;
     mmctx->vtcm_src1_stride = src1_row_size;

-    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
     mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
-    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->quant_task_func         = quant_task_func;
@@ -3769,16 +3905,21 @@ static inline void scan_expert_ids(
 int op_matmul_id(struct htp_ops_context * octx) {
     htp_matmul_tensors_preamble;

+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+    struct htp_mm_context mmctx_struct = {0};
+    struct htp_mm_context * mmctx = &mmctx_struct;
+
+    const int status = htp_mm_init_context(octx, kparams);
+    if (status != HTP_STATUS_OK) {
+        return status;
+    }
+
     struct htp_thread_trace * tr = &octx->ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

-    struct htp_mm_context mmctx_struct = {0};
-    struct htp_mm_context * mmctx = &mmctx_struct;
     mmctx->octx = octx;
     mmctx->act = src1;

-    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
-
     const struct htp_tensor * restrict ids = octx->src[2];

     const size_t src0_row_size = nb01;
@@ -3789,9 +3930,6 @@ int op_matmul_id(struct htp_ops_context * octx) {
     const uint32_t src0_nrows = ne01;  // per expert
     const uint32_t src1_nrows = ne11 * ne12 * ne13;

-    mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
-    mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
-
     // row groups
     const int n_ids = ids->ne[0];  // n_expert_used
     const int n_as  = ne02;        // n_expert
@@ -3843,6 +3981,29 @@ int op_matmul_id(struct htp_ops_context * octx) {
     if (kparams->n_hmx) {
         s = hmx_mm_op_matmul_id(octx, mmctx);
     } else {
+        uint32_t src0_row_start = 0;
+        uint32_t src0_row_end   = src0_nrows;
+        if (octx->ctx->mdev.count > 1) {
+            const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            src0_row_start = range.start;
+            src0_row_end   = range.start + range.count;
+        }
+
+        if (src0_row_start >= src0_row_end) {
+            if (mapping_buf != octx->ctx->ddr_spad_base) {
+                free(mapping_buf);
+            }
+            return HTP_STATUS_OK;
+        }
+
+        const uint32_t nrows = src0_row_end - src0_row_start;
+        mmctx->src0_row_start = src0_row_start;
+        mmctx->src0_row_end   = src0_row_end;
+
+        mmctx->src0_nrows_per_thread = fastdiv(nrows + octx->n_threads - 1, &octx->n_threads_div);
+        mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
+
         if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) {
             s = hvx_mm_matmul_id(octx, mmctx, src1_nrows > 1 ? hvx_mm_id : hvx_mv_id);
         } else {
@@ -3858,29 +4019,31 @@ int op_matmul_id(struct htp_ops_context * octx) {
 }

 int op_matmul_id_nx(struct htp_ops_context * octx) {
+    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+    struct htp_mm_context mmctx_struct = {0};
+    struct htp_mm_context * mmctx = &mmctx_struct;
+
+    const int status = htp_mm_init_context(octx, kparams);
+    if (status != HTP_STATUS_OK) {
+        return status;
+    }
+
     struct htp_thread_trace * tr = &octx->ctx->trace[0];
     htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);

-    const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+    mmctx->octx = octx;
     const uint32_t n_weights = kparams->n_weights;
     const struct htp_tensor * restrict src0 = octx->src[0];
     const struct htp_tensor * restrict act  = octx->src[n_weights];
     const struct htp_tensor * restrict ids  = octx->src[n_weights + 1];

-    struct htp_mm_context mmctx_struct = {0};
-    struct htp_mm_context * mmctx = &mmctx_struct;
-    mmctx->octx = octx;
     mmctx->act = act;

     const size_t src0_row_size = src0->nb[1];
     const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);

-    const uint32_t src0_nrows = src0->ne[1];
     const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];

-    mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
-    mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
-
     const int n_ids = ids->ne[0];
     const int n_as  = src0->ne[2];

@@ -3946,6 +4109,12 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
 }
 int op_matmul_nx(struct htp_ops_context * octx) {
     const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
+
+    const int status = htp_mm_init_context(octx, kparams);
+    if (status != HTP_STATUS_OK) {
+        return status;
+    }
+
     if (kparams->n_hmx) {
         return hmx_mm_nx_2d_f32(octx, kparams);
     }
@@ -4012,7 +4181,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
                                  0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true);

-    size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
+    const size_t vtcm_size = L.total_bytes;

     if (octx->ctx->vtcm_size < vtcm_size) {
         FARF(ERROR, "matmul-nx: current VTCM reservation %zu is too small, needed %zu\n",
@@ -4034,9 +4203,9 @@ int op_matmul_nx(struct htp_ops_context * octx) {
     mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded;
     mmctx->vtcm_src1_stride = src1_row_size;

-    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
     mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
-    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
+    mmctx->vtcm_dst_size_per_thread  = fastdiv(L.dst_bytes, &octx->n_threads_div);

     mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
     mmctx->quant_task_func = quant_task_func;
diff --git a/ggml/src/ggml-hexagon/htp/pad-ops.c b/ggml/src/ggml-hexagon/htp/pad-ops.c
index aaa72b315..0222f24dc 100644
--- a/ggml/src/ggml-hexagon/htp/pad-ops.c
+++ b/ggml/src/ggml-hexagon/htp/pad-ops.c
@@ -12,8 +12,11 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"

 /* Circular wrap: maps any integer x into [0, n) */
 static inline uint32_t wrap_around(int32_t x, uint32_t n) {
@@ -68,6 +71,7 @@ struct htp_pad_context {

     uint32_t nrows_per_thread;
     uint32_t total_dst_rows;
+    uint32_t row_start;

     size_t   type_size;

@@ -78,39 +82,39 @@ struct htp_pad_context {
     size_t   dst_row_size_aligned;
 };

-#define htp_pad_preamble                            \
-    const struct htp_tensor * src = octx->src[0];   \
-    const struct htp_tensor * dst = octx->dst;      \
-                                                    \
-    const uint32_t ne00 = src->ne[0];               \
-    const uint32_t nb00 = src->nb[0];               \
-                                                    \
-    const uint32_t ne0 = dst->ne[0];                \
-    const uint32_t ne1 = dst->ne[1];                \
-    const uint32_t ne2 = dst->ne[2];                \
-    const uint32_t ne3 = dst->ne[3];                \
-                                                    \
-    const uint32_t nb1 = dst->nb[1];                \
-    const uint32_t nb2 = dst->nb[2];                \
-    const uint32_t nb3 = dst->nb[3];                \
-                                                    \
-    const int32_t lp0 = pctx->lp0, rp0 = pctx->rp0; \
-    const int32_t lp1 = pctx->lp1, rp1 = pctx->rp1; \
-    const int32_t lp2 = pctx->lp2, rp2 = pctx->rp2; \
-    const int32_t lp3 = pctx->lp3, rp3 = pctx->rp3; \
-                                                    \
-    const size_t type_size = pctx->type_size;       \
-                                                    \
-    const uint32_t row_start = pctx->nrows_per_thread * ith;                                 \
-    const uint32_t row_end   = MIN(row_start + pctx->nrows_per_thread, pctx->total_dst_rows);
-
-
-#define htp_pad_dma_preamble                                        \
-    const size_t src_row_size         = pctx->src_row_size;         \
-    const size_t src_row_size_aligned = pctx->src_row_size_aligned; \
-    const size_t dst_row_size         = pctx->dst_row_size;         \
-    const size_t dst_row_size_aligned = pctx->dst_row_size_aligned; \
-                                                                    \
+#define htp_pad_preamble                                                             \
+    const struct htp_tensor * src = octx->src[0];                                    \
+    const struct htp_tensor * dst = octx->dst;                                       \
+                                                                                     \
+    const uint32_t ne00 = src->ne[0];                                                \
+    const uint32_t nb00 = src->nb[0];                                                \
+                                                                                     \
+    const uint32_t ne0 = dst->ne[0];                                                 \
+    const uint32_t ne1 = dst->ne[1];                                                 \
+    const uint32_t ne2 = dst->ne[2];                                                 \
+    const uint32_t ne3 = dst->ne[3];                                                 \
+                                                                                     \
+    const uint32_t nb1 = dst->nb[1];                                                 \
+    const uint32_t nb2 = dst->nb[2];                                                 \
+    const uint32_t nb3 = dst->nb[3];                                                 \
+                                                                                     \
+    const int32_t lp0 = pctx->lp0, rp0 = pctx->rp0;                                  \
+    const int32_t lp1 = pctx->lp1, rp1 = pctx->rp1;                                  \
+    const int32_t lp2 = pctx->lp2, rp2 = pctx->rp2;                                  \
+    const int32_t lp3 = pctx->lp3, rp3 = pctx->rp3;                                  \
+                                                                                     \
+    const size_t type_size = pctx->type_size;                                        \
+                                                                                     \
+    const uint32_t row_start = pctx->row_start + pctx->nrows_per_thread * ith;       \
+    const uint32_t row_end   = MIN(row_start + pctx->nrows_per_thread, pctx->row_start + pctx->total_dst_rows);
+
+
+#define htp_pad_dma_preamble                                                                \
+    const size_t src_row_size         = pctx->src_row_size;                                 \
+    const size_t src_row_size_aligned = pctx->src_row_size_aligned;                         \
+    const size_t dst_row_size         = pctx->dst_row_size;                                 \
+    const size_t dst_row_size_aligned = pctx->dst_row_size_aligned;                         \
+                                                                                            \
     uint8_t * src_spad_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread; \
     uint8_t * dst_spad_base = octx->dst_spad.data  + ith * octx->dst_spad.size_per_thread;  \
                                                                                             \
@@ -125,8 +129,8 @@ static void pad_job_per_thread_hvx(unsigned int nth, unsigned int ith, void * da
     struct htp_ops_context * octx = pctx->octx;
     htp_pad_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

     for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
         uint32_t i1, i2, i3;
@@ -165,18 +169,17 @@ static void pad_job_per_thread_hvx(unsigned int nth, unsigned int ith, void * da
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

-    FARF(HIGH, "pad-hvx %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
+    FARF(HIGH, "pad-hvx %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
          src->ne[0], src->ne[1], src->ne[2], src->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         row_start, row_end,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         row_start, row_end);
 }

 // ---------------------------------------------------------------------------
-// HVX + DMA PAD kernel — aligned, double-buffered
+// HVX + DMA PAD kernel - aligned, double-buffered
 // ---------------------------------------------------------------------------

 static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void * data) {
@@ -185,9 +188,6 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
     htp_pad_preamble;
     htp_pad_dma_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     // -----------------------------------------------------------------------
     // Priming phase: push 2 pairs of (dummy_dst_DMA, src_DMA) to seed the
     // double-buffer pipeline before the main loop begins.
@@ -222,6 +222,8 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
     // Main loop: pop completed DMAs, compute in VTCM with aligned HVX ops,
     // push dst DMA and prefetch src for the next+1 row.
     // -----------------------------------------------------------------------
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = row_start; ir < row_end; ir++) {
         uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
         uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
@@ -236,6 +238,7 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
                                              lp2, rp2, ne2,
                                              lp3, rp3, ne3);

+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         if (!interior) {
             hvx_splat_f32_a(dst_spad_cur, 0.0f, ne0);
         } else {
@@ -249,6 +252,7 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
                 hvx_copy_f32_ua(dst_interior, src_spad_cur, ne00);
             }
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         dma_queue_push_vtcm_to_ddr(dma,
             dma_make_ptr(dst_ptr, dst_spad_cur),
@@ -274,14 +278,11 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void

     dma_queue_flush(dma);

-    t2 = HAP_perf_get_qtimer_count();
-
-    FARF(HIGH, "pad-hvx-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
+    FARF(HIGH, "pad-hvx-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
          src->ne[0], src->ne[1], src->ne[2], src->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         row_start, row_end,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         row_start, row_end);
 }

 // ---------------------------------------------------------------------------
@@ -293,8 +294,8 @@ static void pad_job_per_thread_hvx_circular(unsigned int nth, unsigned int ith,
     struct htp_ops_context * octx = pctx->octx;
     htp_pad_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

     for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
         uint32_t i1, i2, i3;
@@ -344,18 +345,17 @@ static void pad_job_per_thread_hvx_circular(unsigned int nth, unsigned int ith,
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

-    FARF(HIGH, "pad-hvx-circ %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
+    FARF(HIGH, "pad-hvx-circ %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
          src->ne[0], src->ne[1], src->ne[2], src->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         row_start, row_end,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         row_start, row_end);
 }

 // ---------------------------------------------------------------------------
-// HVX + DMA circular PAD kernel — aligned, double-buffered
+// HVX + DMA circular PAD kernel - aligned, double-buffered
 // ---------------------------------------------------------------------------

 static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int ith, void * data) {
@@ -364,9 +364,6 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
     htp_pad_preamble;
     htp_pad_dma_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     // -----------------------------------------------------------------------
     // Priming phase: push 2 pairs of (dummy_dst_DMA, src_DMA) to seed the
     // double-buffer pipeline.  Every row is a real src DMA (no null DMAs).
@@ -390,6 +387,8 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
     // Main loop: pop completed DMAs, assemble circular row in VTCM with
     // aligned HVX ops, push dst DMA and prefetch src for the next+1 row.
     // -----------------------------------------------------------------------
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+
     for (uint32_t ir = row_start; ir < row_end; ir++) {
         uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
         uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
@@ -398,7 +397,7 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
         pad_decompose_row(ir, ne1, ne2, &i1, &i2, &i3);
         uint8_t * dst_ptr = (uint8_t *) dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;

-
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
         if (lp0 > 0) {
             uint8_t * dst_left       = dst_spad_cur;
             const uint8_t * src_left = src_spad_cur + (size_t)(ne00 - (uint32_t)lp0) * type_size;
@@ -430,6 +429,7 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
                 }
             }
         }
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);

         dma_queue_push_vtcm_to_ddr(dma,
             dma_make_ptr(dst_ptr, dst_spad_cur),
@@ -448,14 +448,11 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i

     dma_queue_flush(dma);

-    t2 = HAP_perf_get_qtimer_count();
-
-    FARF(HIGH, "pad-hvx-circ-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
+    FARF(HIGH, "pad-hvx-circ-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth,
          src->ne[0], src->ne[1], src->ne[2], src->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         row_start, row_end,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         row_start, row_end);
 }

 int op_pad(struct htp_ops_context * octx) {
@@ -489,19 +486,33 @@ int op_pad(struct htp_ops_context * octx) {
     const uint32_t ne00 = src0->ne[0];

     const uint32_t total_dst_rows = dst->ne[1] * dst->ne[2] * dst->ne[3];
-    const uint32_t n_threads = MIN(octx->n_threads, total_dst_rows > 0 ? total_dst_rows : 1);
+    const size_t dst_row_size     = (size_t)ne0 * type_size;
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = total_dst_rows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_dst_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;

     const size_t src_row_size         = (size_t)ne00 * type_size;
-    const size_t dst_row_size         = (size_t)ne0  * type_size;
     const size_t src_row_size_aligned = hex_round_up(src_row_size, VLEN);
     const size_t dst_row_size_aligned = hex_round_up(dst_row_size, VLEN);

     // Total VTCM needed: 2 buffers (ping+pong) for src and dst, per thread
     const size_t vtcm_needed = (size_t)n_threads * 2 * (src_row_size_aligned + dst_row_size_aligned);

-    const int use_dma = (src0->nb[0] == (uint32_t)type_size) &&
-                        (ne00 >= 512) &&
-                        (octx->ctx->vtcm_base != NULL) &&
+    const int use_dma = (src0->nb[0] == (uint32_t)type_size) && (ne00 >= 512) &&
                         (octx->ctx->vtcm_size >= vtcm_needed);

     if (use_dma) {
@@ -521,8 +532,9 @@ int op_pad(struct htp_ops_context * octx) {
         .lp1 = lp1, .rp1 = rp1,
         .lp2 = lp2, .rp2 = rp2,
         .lp3 = lp3, .rp3 = rp3,
-        .nrows_per_thread = (total_dst_rows + n_threads - 1) / n_threads,
-        .total_dst_rows   = total_dst_rows,
+        .nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .total_dst_rows   = nrows,
+        .row_start        = row_start,
         .type_size        = type_size,
         .src_row_size         = src_row_size,
         .src_row_size_aligned = src_row_size_aligned,
@@ -537,11 +549,10 @@ int op_pad(struct htp_ops_context * octx) {
          dst->ne[0],  dst->ne[1],  dst->ne[2],  dst->ne[3],
          lp0, rp0, lp1, rp1, lp2, rp2, lp3, rp3);

-    if      (circular && use_dma) { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_circular_dma, &pctx, n_threads); }
-    else if (circular)            { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_circular,     &pctx, n_threads); }
-    else if (use_dma)             { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_dma,          &pctx, n_threads); }
-    else                          { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx,              &pctx, n_threads); }
+    if      (circular && use_dma) { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_circular_dma, &pctx, n_threads); }
+    else if (circular)            { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_circular,     &pctx, n_threads); }
+    else if (use_dma)             { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_dma,          &pctx, n_threads); }
+    else                          { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx,              &pctx, n_threads); }

     return HTP_STATUS_OK;
 }
-
diff --git a/ggml/src/ggml-hexagon/htp/repeat-ops.c b/ggml/src/ggml-hexagon/htp/repeat-ops.c
index a6f2f0ed5..530279d65 100644
--- a/ggml/src/ggml-hexagon/htp/repeat-ops.c
+++ b/ggml/src/ggml-hexagon/htp/repeat-ops.c
@@ -12,8 +12,10 @@
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "htp-ctx.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "htp-tensor.h"

 struct htp_repeat_context {
     struct htp_ops_context * octx;
@@ -25,6 +27,7 @@ struct htp_repeat_context {

     uint32_t nrows_per_thread;
     uint32_t total_dst_rows;  // ne1 * ne2 * ne3
+    uint32_t row_start;

     size_t   type_size;
 };
@@ -62,11 +65,11 @@ static void repeat_job_per_thread(unsigned int nth, unsigned int ith, void * dat

     const size_t row_bytes = ne00 * rctx->type_size;

-    const uint32_t row_start = rctx->nrows_per_thread * ith;
-    const uint32_t row_end   = MIN(row_start + rctx->nrows_per_thread, rctx->total_dst_rows);
+    const uint32_t row_start = rctx->row_start + rctx->nrows_per_thread * ith;
+    const uint32_t row_end   = MIN(row_start + rctx->nrows_per_thread, rctx->row_start + rctx->total_dst_rows);

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

     for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
         // Decompose flat dst row index into (i1, i2, i3)
@@ -89,12 +92,12 @@ static void repeat_job_per_thread(unsigned int nth, unsigned int ith, void * dat
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);

-    FARF(HIGH, "repeat %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
+    FARF(HIGH, "repeat %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
          ith, nth, src->ne[0], src->ne[1], src->ne[2], src->ne[3],
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
-         row_start, row_end, (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         row_start, row_end);
 }

 int op_repeat(struct htp_ops_context * octx) {
@@ -119,21 +122,39 @@ int op_repeat(struct htp_ops_context * octx) {
             return HTP_STATUS_NO_SUPPORT;
     }

-    const uint32_t total_dst_rows = dst->ne[1] * dst->ne[2] * dst->ne[3];
-    const uint32_t n_threads = MIN(octx->n_threads, total_dst_rows);
-
     if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
         return HTP_STATUS_OK;
     }

+    const uint32_t total_dst_rows  = dst->ne[1] * dst->ne[2] * dst->ne[3];
+    const size_t dst_row_size = dst->ne[0] * type_size;
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = total_dst_rows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_dst_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     struct htp_repeat_context rctx = {
         .octx             = octx,
         .nr0              = dst->ne[0] / src0->ne[0],
         .nr1              = dst->ne[1] / src0->ne[1],
         .nr2              = dst->ne[2] / src0->ne[2],
         .nr3              = dst->ne[3] / src0->ne[3],
-        .nrows_per_thread = (total_dst_rows + n_threads - 1) / n_threads,
-        .total_dst_rows   = total_dst_rows,
+        .nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+        .total_dst_rows   = nrows,
+        .row_start        = row_start,
         .type_size        = type_size,
     };

@@ -142,7 +163,7 @@ int op_repeat(struct htp_ops_context * octx) {
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
          rctx.nr0, rctx.nr1, rctx.nr2, rctx.nr3);

-    worker_pool_run_func(octx->ctx->worker_pool, repeat_job_per_thread, &rctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, repeat_job_per_thread, &rctx, n_threads);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/rope-ops.c b/ggml/src/ggml-hexagon/htp/rope-ops.c
index 0a4b31ccb..c36976ed0 100644
--- a/ggml/src/ggml-hexagon/htp/rope-ops.c
+++ b/ggml/src/ggml-hexagon/htp/rope-ops.c
@@ -80,6 +80,8 @@ struct htp_rope_context {
     size_t dst_row_stride;
     size_t src0_row_size_aligned;
     uint32_t src0_nrows;
+    uint32_t row_start;
+    uint32_t nrows;

     struct fastdiv_values div_ne2_ne1;
     struct fastdiv_values div_ne1;
@@ -539,11 +541,11 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {

     htp_rope_preamble;

-    const uint32_t src0_nrows = rctx->src0_nrows;
+    const uint32_t src0_nrows = rctx->nrows;
     const uint32_t src0_nrows_per_thread = rctx->src0_nrows_per_thread;

-    const uint32_t src0_start_row = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_start_row = rctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, rctx->row_start + src0_nrows);

     // no work for this thread
     if (src0_start_row >= src0_end_row) {
@@ -706,9 +708,32 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
     }

     const struct htp_rope_kernel_params * kparams = (const struct htp_rope_kernel_params *) octx->kernel_params;
-    assert(kparams->n_threads > 0);
+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
     assert(octx->ctx->vtcm_size >= kparams->vtcm_size);

+    const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
+    const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = total_rows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
+            total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     const uint32_t ne0 = dst->ne[0];
     const size_t src0_row_size   = src0->ne[0] * sizeof(float);
     const size_t src0_row_stride = src0->nb[1];
@@ -752,15 +777,17 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
     rctx.dst_row_stride        = dst_row_stride;
     rctx.src0_row_size_aligned = kparams->src0_row_size_aligned;

-    rctx.src0_nrows            = kparams->src0_nrows;
-    rctx.src0_nrows_per_thread = kparams->src0_nrows_per_thread;
+    rctx.src0_nrows            = nrows;
+    rctx.nrows                 = nrows;
+    rctx.row_start             = row_start;
+    rctx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
     rctx.div_ne2_ne1           = kparams->div_ne2_ne1;
     rctx.div_ne1               = kparams->div_ne1;

     FARF(HIGH, "rope-f32 n-rows %u n-dims %d ne0 %u ext-factor %.6f theta-scale %.6f attn-factor %.6f\n", rctx.src0_nrows, rctx.n_dims, ne0,
          rctx.ext_factor, rctx.theta_scale, rctx.attn_factor);

-    work_queue_run(octx->ctx->work_queue, rope_job_f32, &rctx, kparams->n_threads);
+    work_queue_run(octx->ctx->work_queue, rope_job_f32, &rctx, n_threads);

     return err;
 }
diff --git a/ggml/src/ggml-hexagon/htp/set-rows-ops.c b/ggml/src/ggml-hexagon/htp/set-rows-ops.c
index 340a497f7..fbd5162a7 100644
--- a/ggml/src/ggml-hexagon/htp/set-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/set-rows-ops.c
@@ -18,6 +18,7 @@
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"

+#include "hex-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
 #include "htp-tensor.h"
@@ -58,6 +59,9 @@ struct set_rows_context {
     const struct htp_set_rows_kernel_params * kparams;
     struct htp_set_rows_vtcm_layout vtcm_layout;
     uint8_t * vtcm_base;
+    uint32_t task_start;
+    uint32_t tasks;
+    uint32_t tasks_per_thread;
 };

 #define SET_ROWS_THREAD_DMA_FN(TYPE_NAME, IDX_TYPE, COMPUTE_EXPR)                                                \
@@ -67,12 +71,12 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
     const struct htp_set_rows_kernel_params * kparams = srctx->kparams;                                          \
     set_rows_preamble;                                                                                           \
     struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                       \
-    const uint32_t dr  = kparams->tasks_per_thread;                                                              \
-    const uint32_t ir0 = dr * ith;                                                                               \
-    if (ir0 >= kparams->total_tasks) {                                                                           \
+    const uint32_t dr  = srctx->tasks_per_thread;                                                                \
+    const uint32_t ir0 = srctx->task_start + dr * ith;                                                           \
+    if (ir0 >= srctx->task_start + srctx->tasks) {                                                               \
         return;                                                                                                  \
     }                                                                                                            \
-    const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks);                                                    \
+    const uint32_t ir1 = MIN(ir0 + dr, srctx->task_start + srctx->tasks);                                        \
     dma_queue * dma_queue = octx->ctx->dma[ith];                                                                 \
     const struct htp_set_rows_vtcm_layout * vtcm_layout = &srctx->vtcm_layout;                                   \
     uint8_t * vtcm_src0 = srctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;   \
@@ -192,18 +196,44 @@ int op_set_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_NO_SUPPORT;
     }

-    if (octx->src[1]->type != HTP_TYPE_I32 && octx->src[1]->type != HTP_TYPE_I64) {
-        return HTP_STATUS_NO_SUPPORT;
+    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
+        return HTP_STATUS_OK;
+    }
+
+    const struct htp_tensor * dst = octx->dst;
+    const uint32_t total_tasks    = kparams->total_tasks;
+
+    uint32_t task_start = 0;
+    uint32_t tasks      = total_tasks;
+
+    if (octx->ctx->mdev.count > 1) {
+        const bool can_split = htp_tensor_mdev_data_aligned(dst) && (dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 && !htp_tensor_is_permuted(dst);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_tasks, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        task_start = range.start;
+        tasks      = range.count;
     }

+    if (tasks == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     // l2fetch the src1 (indices) tensor in the main thread
     hex_l2fetch_block((const void *)octx->src[1]->data, octx->src[1]->ne[3] * octx->src[1]->nb[3]);

     struct set_rows_context srctx;
     srctx.octx = octx;
     srctx.kparams = kparams;
+    srctx.task_start = task_start;
+    srctx.tasks = tasks;
+    srctx.tasks_per_thread = fastdiv(tasks + n_threads - 1, &octx->n_threads_div);

-    htp_set_rows_vtcm_layout_build(&srctx.vtcm_layout, octx->dst->type, ne00, kparams->n_threads);
+    htp_set_rows_vtcm_layout_build(&srctx.vtcm_layout, octx->dst->type, ne00, n_threads);
     srctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;

     work_queue_func_t q_func = NULL;
@@ -216,15 +246,15 @@ int op_set_rows(struct htp_ops_context * octx) {
         default:            return HTP_STATUS_NO_SUPPORT;
     }

-    FARF(HIGH, "set-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu n_threads %d\n",
+    FARF(HIGH, "set-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu n-threads %d\n",
          octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
          octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3],
          octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
-         srctx.vtcm_layout.src0_bytes_per_thread * kparams->n_threads,
-         srctx.vtcm_layout.dst_bytes_per_thread  * kparams->n_threads,
-         kparams->n_threads);
+         srctx.vtcm_layout.src0_bytes_per_thread * n_threads,
+         srctx.vtcm_layout.dst_bytes_per_thread  * n_threads,
+         n_threads);

-    work_queue_run(octx->ctx->work_queue, q_func, &srctx, kparams->n_threads);
+    work_queue_run(octx->ctx->work_queue, q_func, &srctx, n_threads);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/softmax-ops.c b/ggml/src/ggml-hexagon/htp/softmax-ops.c
index d78bcc0eb..2497ec763 100644
--- a/ggml/src/ggml-hexagon/htp/softmax-ops.c
+++ b/ggml/src/ggml-hexagon/htp/softmax-ops.c
@@ -14,9 +14,11 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "htp-tensor.h"

 #define htp_softmax_preamble3                     \
     const uint32_t ne00 = src0->ne[0];            \
@@ -69,6 +71,8 @@ struct htp_softmax_context {
     struct fastdiv_values fastdiv_ne13; // For mask broadcasting

     uint32_t src0_nrows_per_thread;
+    uint32_t row_start;
+    uint32_t nrows;
 };

 static void apply_mask(float * restrict wp0,
@@ -223,19 +227,17 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {

     htp_softmax_preamble3;

-    const uint32_t src0_nrows            = ne01 * ne02 * ne03;  // src0 rows
+    const uint32_t src0_nrows            = smctx->nrows;
     const uint32_t src0_nrows_per_thread = smctx->src0_nrows_per_thread;

-    const uint32_t src0_start_row = src0_nrows_per_thread * ith;
-    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
+    const uint32_t src0_start_row = smctx->row_start + src0_nrows_per_thread * ith;
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, smctx->row_start + src0_nrows);

     // no work for this thread
     if (src0_start_row >= src0_end_row) {
         return;
     }

-    uint64_t qt = HAP_perf_get_qtimer_count();
-
     int is_aligned = 1;
     int opt_path   = 0;

@@ -262,6 +264,9 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
     uint32_t prev_i2 = (uint32_t)-1;
     float slope = 1.0f;

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
+
     for (uint32_t r = src0_start_row; r < src0_end_row; ++r) {
         uint32_t i1 = fastmodulo(r, ne01, &smctx->fastdiv_ne01);
         uint32_t r_div_ne01 = fastdiv(r, &smctx->fastdiv_ne01);
@@ -323,10 +328,11 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
         }
     }

-    qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt);
-    FARF(HIGH, "softmax-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u : opt %u f16 %u usec %u\n", ith, nth,
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
+
+    FARF(HIGH, "softmax-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u : opt %u f16 %u\n", ith, nth,
          ne00, ne01, ne02, ne03, src0_start_row, src0_end_row, ne10, ne11, ne12, ne13,
-         ne0, ne1, ne2, ne3, opt_path, smctx->use_f16, (unsigned) qt);
+         ne0, ne1, ne2, ne3, opt_path, smctx->use_f16);
 }

 static int execute_op_softmax_f32(struct htp_ops_context * octx) {
@@ -342,13 +348,32 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
     init_softmax_ctx(&smctx, octx);

     const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads  = MIN(octx->n_threads, src0_nrows);
+    const size_t elem_size = sizeof(float);
+    const size_t dst_row_size = dst->nb[1];
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }

-    smctx.src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
+    smctx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+    smctx.row_start             = row_start;
+    smctx.nrows                 = nrows;

     const size_t src0_row_size = src0->nb[1];
     const size_t src1_row_size = src0_row_size;
-    const size_t dst_row_size  = dst->nb[1];

     // VTCM scratchpads for all tensors
     // 4 rows per thread, padded to HVX vector size
@@ -383,9 +408,7 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
     octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size; octx->src1_spad.src = NULL;
     octx->dst_spad.data  = octx->src1_spad.data + octx->src1_spad.size; octx->dst_spad.src  = NULL;

-    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) return err;
-
-    worker_pool_run_func(octx->ctx->worker_pool, softmax_job_f32, &smctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, softmax_job_f32, &smctx, n_threads);

     return err;
 }
diff --git a/ggml/src/ggml-hexagon/htp/solve-tri-ops.c b/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
index ae8e1a504..847a78712 100644
--- a/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
+++ b/ggml/src/ggml-hexagon/htp/solve-tri-ops.c
@@ -1,13 +1,16 @@
 #pragma clang diagnostic ignored "-Wunused-but-set-variable"

 #include <HAP_farf.h>
-#include <HAP_perf.h>
 #include <string.h>

+#include "hex-common.h"
+#include "hex-profile.h"
+
 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
+#include "htp-tensor.h"
 #include "hvx-types.h"
 #include "hvx-utils.h"

@@ -15,6 +18,7 @@ struct htp_solve_tri_context {
     struct htp_ops_context * octx;
     uint32_t                 jobs_per_thread;
     uint32_t                 total_jobs;
+    uint32_t                 job_start;
     uint32_t                 k_chunks;
     uint32_t                 col_block;
 };
@@ -89,11 +93,11 @@ static void solve_tri_batch_thread_f32(unsigned int nth, unsigned int ith, void
     const uint32_t col_block = VLEN_FP32;
     const uint32_t k_full    = (k / col_block) * col_block;

-    const uint32_t start_batch = sctx->jobs_per_thread * ith;
-    const uint32_t end_batch   = MIN(start_batch + sctx->jobs_per_thread, sctx->total_jobs);
+    const uint32_t start_batch = sctx->job_start + sctx->jobs_per_thread * ith;
+    const uint32_t end_batch   = MIN(start_batch + sctx->jobs_per_thread, sctx->job_start + sctx->total_jobs);

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_batch);

     for (uint32_t batch = start_batch; batch < end_batch; ++batch) {
         const uint32_t i03 = batch / ne02;
@@ -127,11 +131,10 @@ static void solve_tri_batch_thread_f32(unsigned int nth, unsigned int ith, void
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) end_batch);

-    FARF(HIGH, "solve-tri-batch %d/%d: A=(%ux%u) B=(%ux%u) batch %u:%u usec %u\n",
-         ith, nth, n, n, k, n, start_batch, end_batch,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+    FARF(HIGH, "solve-tri-batch %d/%d: A=(%ux%u) B=(%ux%u) batch %u:%u\n",
+         ith, nth, n, n, k, n, start_batch, end_batch);
 }

 // Chunk-level thread: each job is one (batch, col_chunk) pair.
@@ -148,11 +151,11 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void

     const uint32_t ne02 = src0->ne[2];

-    const uint32_t start_job = sctx->jobs_per_thread * ith;
-    const uint32_t end_job   = MIN(start_job + sctx->jobs_per_thread, sctx->total_jobs);
+    const uint32_t start_job = sctx->job_start + sctx->jobs_per_thread * ith;
+    const uint32_t end_job   = MIN(start_job + sctx->jobs_per_thread, sctx->job_start + sctx->total_jobs);

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_job);

     for (uint32_t job = start_job; job < end_job; ++job) {
         const uint32_t batch = job / sctx->k_chunks;
@@ -161,16 +164,14 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
         const uint32_t i03 = batch / ne02;
         const uint32_t i02 = batch - i03 * ne02;

-        const uint32_t col0 = chunk * sctx->col_block;
-        const uint32_t coln = MIN(sctx->col_block, k - col0);
-
         const float * A_batch =
             (const float *) ((const uint8_t *) (uintptr_t) src0->data + i02 * src0->nb[2] + i03 * src0->nb[3]);
         const float * B_batch =
             (const float *) ((const uint8_t *) (uintptr_t) src1->data + i02 * src1->nb[2] + i03 * src1->nb[3]);
         float * X_batch = (float *) ((uint8_t *) (uintptr_t) dst->data + i02 * dst->nb[2] + i03 * dst->nb[3]);

-        const bool use_hvx = (coln >= 8);
+        const uint32_t col0 = chunk * sctx->col_block;
+        const uint32_t coln = MIN(sctx->col_block, k - col0);

         for (uint32_t row = 0; row < n; ++row) {
             const float diag     = A_batch[row * n + row];
@@ -179,7 +180,7 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
             const float * A_row = A_batch + row * n;
             const float * B_row = B_batch + row * k;

-            if (use_hvx) {
+            if (coln >= 8) {
                 solve_tri_row_hvx(A_row, B_row, X_batch, row, k, col0, coln, inv_diag);
             } else {
                 solve_tri_row_scalar(A_row, B_row, X_batch, row, k, col0, coln, inv_diag);
@@ -187,11 +188,10 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) end_job);

-    FARF(HIGH, "solve-tri-chunk %d/%d: A=(%ux%u) B=(%ux%u) job %u:%u usec %u\n",
-         ith, nth, n, n, k, n, start_job, end_job,
-         (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+    FARF(HIGH, "solve-tri-chunk %d/%d: A=(%ux%u) B=(%ux%u) jobs %u:%u\n",
+         ith, nth, n, n, k, n, start_job, end_job);
 }

 int op_solve_tri(struct htp_ops_context * octx) {
@@ -235,32 +235,64 @@ int op_solve_tri(struct htp_ops_context * octx) {
          dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], batched);

     if (batched) {
+        uint32_t job_start = 0;
+        uint32_t njobs     = total_batches;
+
+        if (octx->ctx->mdev.count > 1) {
+            const uint32_t batch_size = dst->nb[2];
+            const uint32_t batches_per_chunk = (batch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(batch_size, HEX_L2_LINE_SIZE)) : 1;
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_batches, htp_tensor_mdev_data_aligned(dst) ? batches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            job_start = range.start;
+            njobs     = range.count;
+        }
+
+        if (njobs == 0) {
+            return HTP_STATUS_OK;
+        }
+
         // Batch-level parallelism
-        const uint32_t n_threads = MIN((uint32_t) octx->n_threads, total_batches);
+        const uint32_t n_threads = octx->n_threads;

         struct htp_solve_tri_context sctx = {
             .octx            = octx,
-            .jobs_per_thread = (total_batches + n_threads - 1) / n_threads,
-            .total_jobs      = total_batches,
+            .jobs_per_thread = fastdiv(njobs + n_threads - 1, &octx->n_threads_div),
+            .total_jobs      = njobs,
+            .job_start       = job_start,
             .k_chunks        = k_chunks,
             .col_block       = col_block,
         };

-        worker_pool_run_func(octx->ctx->worker_pool, solve_tri_batch_thread_f32, &sctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, solve_tri_batch_thread_f32, &sctx, n_threads);
     } else {
         // Chunk-level parallelism
         const uint32_t total_jobs = total_batches * k_chunks;
-        const uint32_t n_threads  = MIN((uint32_t) octx->n_threads, MAX(total_jobs, 1));
+
+        uint32_t job_start = 0;
+        uint32_t njobs     = total_jobs;
+
+        if (octx->ctx->mdev.count > 1) {
+            const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
+            const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_jobs, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+            job_start = range.start;
+            njobs     = range.count;
+        }
+
+        if (njobs == 0) {
+            return HTP_STATUS_OK;
+        }
+
+        const uint32_t n_threads = octx->n_threads;

         struct htp_solve_tri_context sctx = {
             .octx            = octx,
-            .jobs_per_thread = (total_jobs + n_threads - 1) / n_threads,
-            .total_jobs      = total_jobs,
+            .jobs_per_thread = fastdiv(njobs + n_threads - 1, &octx->n_threads_div),
+            .total_jobs      = njobs,
+            .job_start       = job_start,
             .k_chunks        = k_chunks,
             .col_block       = col_block,
         };

-        worker_pool_run_func(octx->ctx->worker_pool, solve_tri_chunk_thread_f32, &sctx, n_threads);
+        work_queue_run(octx->ctx->work_queue, solve_tri_chunk_thread_f32, &sctx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/ssm-conv.c b/ggml/src/ggml-hexagon/htp/ssm-conv.c
index a48bc9ed8..bef142536 100644
--- a/ggml/src/ggml-hexagon/htp/ssm-conv.c
+++ b/ggml/src/ggml-hexagon/htp/ssm-conv.c
@@ -4,7 +4,6 @@

 #include <HAP_farf.h>
 #include <HAP_mem.h>
-#include <HAP_perf.h>
 #include <HAP_ps.h>
 #include <hexagon_protos.h>
 #include <hexagon_types.h>
@@ -16,8 +15,9 @@
 #include "ggml-common.h"
 #include "htp-ctx.h"
 #include "hex-dma.h"
+#include "hex-profile.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "htp-tensor.h"
 #include "hvx-utils.h"

 #define htp_ssm_conv_tensors_preamble                           \
@@ -63,6 +63,8 @@ struct htp_ssm_conv_context {
     uint32_t nrows_per_thread;
     uint32_t d_inner_tile;
     uint64_t t_start;
+    uint32_t row_start;
+    uint32_t nrows;
 };

 #define htp_ssm_conv_preamble                                                   \
@@ -75,9 +77,6 @@ struct htp_ssm_conv_context {
 static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *data) {
     htp_ssm_conv_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     const uint32_t d_conv  = src1->ne[0];
     const uint32_t d_inner = src0->ne[1];
     const uint32_t n_t     = dst->ne[1];
@@ -95,14 +94,17 @@ static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *da

     // Calculate row range for this thread
     const uint32_t d_inner_per_thread = scctx->nrows_per_thread;
-    const uint32_t d_inner_start = d_inner_per_thread * ith;
-    const uint32_t d_inner_end   = MIN(d_inner_start + d_inner_per_thread, d_inner);
+    const uint32_t d_inner_start = scctx->row_start + d_inner_per_thread * ith;
+    const uint32_t d_inner_end   = MIN(d_inner_start + d_inner_per_thread, scctx->row_start + scctx->nrows);

     // No work for this thread
     if (d_inner_start >= d_inner_end) {
         return;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_start);
+
     for (uint32_t i3 = 0; i3 < n_s; ++i3) {
         for (uint32_t i2 = 0; i2 < n_t; ++i2) {
             for (uint32_t i1 = d_inner_start; i1 < d_inner_end; ++i1) {
@@ -121,12 +123,12 @@ static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *da
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_end);

-    FARF(HIGH, "ssm-conv-f32 %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n",
+    FARF(HIGH, "ssm-conv-f32 %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
          ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], d_inner_start, d_inner_end,
          src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
-         dst->ne[2], dst->ne[3], (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[2], dst->ne[3]);
 }


@@ -257,9 +259,6 @@ static inline void transpose_src0_block(const float * src0_block,
 static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void *data) {
     htp_ssm_conv_preamble;

-    uint64_t t1, t2;
-    t1 = HAP_perf_get_qtimer_count();
-
     const uint32_t d_conv  = src1->ne[0];
     const uint32_t d_inner = src0->ne[1];
     const uint32_t n_t     = dst->ne[1];
@@ -273,13 +272,16 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
     const uint32_t dst_stride_seq    = dst->nb[2]  / sizeof(float);

     const uint32_t dr  = scctx->nrows_per_thread;
-    const uint32_t ir0 = dr * ith;
-    const uint32_t ir1 = MIN(ir0 + dr, d_inner);
+    const uint32_t ir0 = scctx->row_start + dr * ith;
+    const uint32_t ir1 = MIN(ir0 + dr, scctx->row_start + scctx->nrows);

     if (ir0 >= ir1) {
         return;
     }

+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
+
     const uint32_t d_inner_per_thread = ir1 - ir0;
     const uint32_t d_inner_stride     = scctx->nrows_per_thread;
     const uint32_t d_inner_tile       = scctx->d_inner_tile;
@@ -319,97 +321,118 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
                         HVX_Vector w = *(const HVX_Vector *) (src1_T + j * d_inner_stride + tile_off + cb);
                         acc          = Q6_Vqf32_vadd_Vqf32Vqf32(acc, Q6_Vqf32_vmpy_VsfVsf(x, w));
                     }
-                    HVX_Vector res = Q6_Vsf_equals_Vqf32(acc);

-                    float * dst_ptr = dst_data + i3 * dst_stride_seq + t * dst_stride_token + (ir0 + tile_off + cb);
+                    HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
+
+                    float * dst_ptr = dst_data + (ir0 + tile_off + cb) + t * dst_stride_token + i3 * dst_stride_seq;
                     if (cb_n == C_TILE) {
-                        *(HVX_UVector *) dst_ptr = res;
+                        *(HVX_UVector *) dst_ptr = y;
                     } else {
-                        hvx_vec_store_u(dst_ptr, cb_n * sizeof(float), res);
+                        hvx_vec_store_u(dst_ptr, cb_n * sizeof(float), y);
                     }
                 }
             }
         }
     }

-    t2 = HAP_perf_get_qtimer_count();
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir1);

-    FARF(HIGH, "ssm-conv-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) tile=%u * %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n",
-         ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1, d_inner_tile,
+    FARF(HIGH, "ssm-conv-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
+         ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
          src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
-         dst->ne[2], dst->ne[3], (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
+         dst->ne[2], dst->ne[3]);
 }

 int op_ssm_conv_f32(struct htp_ops_context * octx) {
-    htp_ssm_conv_tensors_preamble;
+    const struct htp_tensor * src0 = octx->src[0];
+    const struct htp_tensor * src1 = octx->src[1];
+    const struct htp_tensor * dst  = octx->dst;

     if (src0->type != HTP_TYPE_F32 || src1->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_F32) {
-        FARF(ERROR, "ssm_conv: only (F32 x F32 -> F32) OPs supported");
         return HTP_STATUS_NO_SUPPORT;
     }

-    struct htp_ssm_conv_context scctx = { 0 };
-    scctx.octx = octx;
-
     const uint32_t d_conv  = src1->ne[0];
     const uint32_t d_inner = src0->ne[1];
     const uint32_t n_t     = dst->ne[1];  // tokens per sequence
     const uint32_t n_s     = dst->ne[2];  // number of sequences in the batch

-    const uint32_t n_threads = MIN(octx->n_threads, d_inner);
+    if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
+        return HTP_STATUS_OK;
+    }

-    if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
-        uint32_t use_hvx = 0;
-        if (d_inner >= VLEN_FP32 && n_t >= VLEN_FP32) {
-            use_hvx = 1;
-        }
+    uint32_t row_start = 0;
+    uint32_t nrows     = d_inner;
+
+    if (octx->ctx->mdev.count > 1) {
+        const uint32_t elems_per_chunk = VLEN_FP32;
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(d_inner, htp_tensor_mdev_data_aligned(dst) ? elems_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
+    struct htp_ssm_conv_context scctx = { 0 };
+    scctx.octx      = octx;
+    scctx.row_start = row_start;
+    scctx.nrows     = nrows;
+
+    uint32_t use_hvx = 0;
+    if (nrows >= VLEN_FP32 && n_t >= VLEN_FP32) {
+        use_hvx = 1;
+    }
+
+    const uint32_t raw_rpt = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
+    scctx.nrows_per_thread = hex_round_up(raw_rpt, VLEN_FP32);

-        scctx.nrows_per_thread = hex_round_up((d_inner + n_threads - 1) / n_threads, VLEN_FP32);
+    const uint32_t d_inner_per_thread = scctx.nrows_per_thread;
+    const uint32_t ncs                = src0->ne[0];

-        const uint32_t d_inner_per_thread = scctx.nrows_per_thread;
-        const uint32_t ncs                = src0->ne[0];
+    const uint32_t src1_T_size = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 256);
+    const uint32_t src0_T_max = HTP_SSM_CONV_VTCM_BUDGET > src1_T_size ? HTP_SSM_CONV_VTCM_BUDGET - src1_T_size : 0;

-        const uint32_t src1_T_size = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 256);
-        const uint32_t src0_T_max = HTP_SSM_CONV_VTCM_BUDGET > src1_T_size ? HTP_SSM_CONV_VTCM_BUDGET - src1_T_size : 0;
+    uint32_t d_inner_tile = (src0_T_max / sizeof(float)) / ncs;
+    d_inner_tile -= (d_inner_tile % VLEN_FP32);
+    if (d_inner_tile == 0) {
+        FARF(HIGH, "ssm_conv-f32: inner tile rounds to 0 (ncs=%u), falling back to scalar\n", ncs);
+        use_hvx = 0;
+    } else {
+        scctx.d_inner_tile = d_inner_tile;

-        uint32_t d_inner_tile = (src0_T_max / sizeof(float)) / ncs;
-        d_inner_tile -= (d_inner_tile % VLEN_FP32);
-        if (d_inner_tile == 0) {
-            FARF(HIGH, "ssm_conv-f32: inner tile rounds to 0 (ncs=%u), falling back to scalar\n", ncs);
+        octx->src0_spad.size_per_thread = hex_round_up(d_inner_tile * ncs * sizeof(float), 256);
+        octx->src1_spad.size_per_thread = src1_T_size;
+        octx->dst_spad.size_per_thread  = 0;
+
+        octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
+        octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
+        octx->dst_spad.size  = 0;
+
+        octx->src0_spad.data = octx->ctx->vtcm_base;
+        octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
+        octx->src0_spad.src  = NULL;
+        octx->src1_spad.src  = NULL;
+
+        const size_t total_spad = octx->src0_spad.size + octx->src1_spad.size;
+        if (total_spad > octx->ctx->vtcm_size) {
+            FARF(HIGH, "ssm_conv-f32: scratchpad %zu exceeds VTCM %zu, falling back to scalar\n",
+                 total_spad, octx->ctx->vtcm_size);
             use_hvx = 0;
-        } else {
-            scctx.d_inner_tile = d_inner_tile;
-
-            octx->src0_spad.size_per_thread = hex_round_up(d_inner_tile * ncs * sizeof(float), 256);
-            octx->src1_spad.size_per_thread = src1_T_size;
-            octx->dst_spad.size_per_thread  = 0;
-
-            octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
-            octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
-            octx->dst_spad.size  = 0;
-
-            octx->src0_spad.data = octx->ctx->vtcm_base;
-            octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
-            octx->src0_spad.src  = NULL;
-            octx->src1_spad.src  = NULL;
-
-            const size_t total_spad = octx->src0_spad.size + octx->src1_spad.size;
-            if (total_spad > octx->ctx->vtcm_size) {
-                FARF(HIGH, "ssm_conv-f32: scratchpad %zu exceeds VTCM %zu, falling back to scalar\n",
-                     total_spad, octx->ctx->vtcm_size);
-                use_hvx = 0;
-            }
         }
+    }

-        FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_hvx %d\n", src0->ne[0],
-             src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
-             dst->ne[1], dst->ne[2], dst->ne[3], use_hvx);
+    FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_hvx %d\n", src0->ne[0],
+         src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
+         dst->ne[1], dst->ne[2], dst->ne[3], use_hvx);

-        if (use_hvx) {
-            worker_pool_run_func(octx->ctx->worker_pool, ssm_conv_thread_f32_f32_hvx, &scctx, n_threads);
-        } else {
-            worker_pool_run_func(octx->ctx->worker_pool, ssm_conv_thread_f32_f32, &scctx, n_threads);
-        }
+    if (use_hvx) {
+        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32_hvx, &scctx, n_threads);
+    } else {
+        work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32, &scctx, n_threads);
     }

     return HTP_STATUS_OK;
diff --git a/ggml/src/ggml-hexagon/htp/sum-rows-ops.c b/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
index 874c41ab2..faf716b4b 100644
--- a/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
+++ b/ggml/src/ggml-hexagon/htp/sum-rows-ops.c
@@ -13,35 +13,38 @@

 #define GGML_COMMON_DECL_C
 #include "ggml-common.h"
+#include "hex-common.h"
+#include "hex-profile.h"
 #include "htp-ctx.h"
 #include "htp-ops.h"
-#include "htp-ops.h"
+#include "htp-tensor.h"

 #define sum_rows_preamble                         \
     const struct htp_tensor *src0 = octx->src[0]; \
     const struct htp_tensor *dst  = octx->dst;    \
                                                   \
-    const uint32_t ne00 = src0->ne[0];     \
-    const uint32_t ne01 = src0->ne[1];     \
-    const uint32_t ne02 = src0->ne[2];     \
-    const uint32_t ne03 = src0->ne[3];     \
-                                           \
-    const uint32_t nb00 = src0->nb[0];     \
-    const uint32_t nb01 = src0->nb[1];     \
-    const uint32_t nb02 = src0->nb[2];     \
-    const uint32_t nb03 = src0->nb[3];     \
-                                           \
-    const uint32_t  ne0 = dst->ne[0];      \
-    const uint32_t  ne1 = dst->ne[1];      \
-    const uint32_t  ne2 = dst->ne[2];      \
-    const uint32_t  ne3 = dst->ne[3];      \
-                                           \
-    const uint32_t  nb0 = dst->nb[0];      \
-    const uint32_t  nb1 = dst->nb[1];      \
-    const uint32_t  nb2 = dst->nb[2];      \
-    const uint32_t  nb3 = dst->nb[3];      \
+    const uint32_t ne00 = src0->ne[0];            \
+    const uint32_t ne01 = src0->ne[1];            \
+    const uint32_t ne02 = src0->ne[2];            \
+    const uint32_t ne03 = src0->ne[3];            \
+                                                  \
+    const uint32_t nb00 = src0->nb[0];            \
+    const uint32_t nb01 = src0->nb[1];            \
+    const uint32_t nb02 = src0->nb[2];            \
+    const uint32_t nb03 = src0->nb[3];            \
+                                                  \
+    const uint32_t  ne0 = dst->ne[0];             \
+    const uint32_t  ne1 = dst->ne[1];             \
+    const uint32_t  ne2 = dst->ne[2];             \
+    const uint32_t  ne3 = dst->ne[3];             \
+                                                  \
+    const uint32_t  nb0 = dst->nb[0];             \
+    const uint32_t  nb1 = dst->nb[1];             \
+    const uint32_t  nb2 = dst->nb[2];             \
+    const uint32_t  nb3 = dst->nb[3];             \

 struct sum_rows_context {
+    struct htp_ops_context * octx;
     const uint8_t * src_data;
     uint8_t       * dst_data;
     uint32_t        ne00;
@@ -76,6 +79,9 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
     // Calculate actual number of rows for this thread
     const uint32_t n_rows = end_row - start_row;

+    struct htp_thread_trace * tr = &smctx->octx->ctx->trace[ith];
+    htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
+
     for (uint32_t ir = 0; ir < n_rows; ir++) {
         const float * restrict src_local = src_th + (ir * (src_stride / sizeof(float)));

@@ -89,6 +95,8 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
             dst_th[ir] = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
         }
     }
+
+    htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
 }

 int op_sum_rows(struct htp_ops_context * octx) {
@@ -102,9 +110,26 @@ int op_sum_rows(struct htp_ops_context * octx) {
         return HTP_STATUS_OK;
     }

-    const uint32_t src0_nrows = ne01 * ne02 * ne03;
-    const uint32_t n_threads = MIN(octx->n_threads, src0_nrows);
-    const uint32_t rows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
+    const uint32_t src0_nrows      = ne01 * ne02 * ne03;
+    const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
+
+    uint32_t row_start = 0;
+    uint32_t nrows     = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+    const uint32_t rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);

     bool opt_path = false;
     if ((0 == hex_is_aligned((void *) src0->data, VLEN)) && !(nb01 & (VLEN - 1))) {
@@ -112,17 +137,18 @@ int op_sum_rows(struct htp_ops_context * octx) {
     }

     struct sum_rows_context smctx = {
-        .src_data        = (const uint8_t *) src0->data,
-        .dst_data        = (uint8_t *) dst->data,
+        .octx            = octx,
+        .src_data        = (const uint8_t *) src0->data + row_start * nb01,
+        .dst_data        = (uint8_t *) dst->data + row_start * nb1,
         .ne00            = ne00,
         .src_stride      = nb01,
         .dst_stride      = nb1,
         .rows_per_thread = rows_per_thread,
-        .total_rows      = src0_nrows,
+        .total_rows      = nrows,
         .opt_path        = opt_path,
     };

-    worker_pool_run_func(octx->ctx->worker_pool, sum_rows_thread_f32, &smctx, n_threads);
+    work_queue_run(octx->ctx->work_queue, sum_rows_thread_f32, &smctx, n_threads);

     return HTP_STATUS_OK;
 }
diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c
index 7850ab27e..cb82bfa3c 100644
--- a/ggml/src/ggml-hexagon/htp/unary-ops.c
+++ b/ggml/src/ggml-hexagon/htp/unary-ops.c
@@ -46,6 +46,7 @@ struct htp_unary_context {
     uint32_t                  block;
     uint32_t                  src0_nrows;
     uint32_t                  src0_nrows_per_thread;
+    uint32_t                  row_start;
     uint32_t                  nc;
     uint32_t                  col_tile;             // tiled mode
     bool                      broadcast_weight;
@@ -496,7 +497,7 @@ static void tri_f32(const float * restrict src,
         }
         if (boundary > ne0) boundary = ne0;

-        // Full HVX vectors — each starts at a 128-byte aligned offset
+        // Full HVX vectors - each starts at a 128-byte aligned offset
         for (uint32_t i = 0; i < nvec; i++) {
             const uint32_t vec_start = i * VLEN_FP32;
             const uint32_t vec_end   = vec_start + VLEN_FP32;
@@ -563,7 +564,7 @@ static void softplus_f32(const float * restrict src,

         for (uint32_t i = 0; i < ne0; i++) {
             float x = src_f[i];
-            // For x > 20: softplus(x) ≈ x (avoids exp overflow)
+            // For x > 20: softplus(x) ~ x (avoids exp overflow)
             dst_f[i] = (x > 20.0f) ? x : logf(1.0f + expf(x));
         }
     }
@@ -661,8 +662,8 @@ static void unary_task_##SUFFIX##_##NAME(unsigned int nth, unsigned int ith, voi
     const size_t dst_row_size_aligned  = uctx->dst_row_size_aligned;                                                \
                                                                                                                     \
     const uint32_t src0_nrows = uctx->src0_nrows;                                                                   \
-    const uint32_t src0_start_row = src0_nrows_per_thread * ith;                                                    \
-    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);                        \
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;                                  \
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);      \
                                                                                                                     \
     if (src0_start_row >= src0_end_row) {                                                                           \
         return;                                                                                                     \
@@ -833,124 +834,126 @@ DEFINE_UNARY_TASK_IMPL(unary_abs, _Float16, f16, false, false, abs_f16(src0_vtcm
 DEFINE_UNARY_TASK_IMPL(unary_log, _Float16, f16, false, false, log_f16(src0_vtcm, dst_vtcm, block_size, uctx))

 // Apply a pointwise unary op to one column tile that is already in VTCM.
-#define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR)                                                       \
-static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) {                          \
-    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;                                \
-    struct htp_ops_context * octx = uctx->octx;                                                                     \
-    const struct htp_tensor * src = octx->src[0];                                                                   \
-    const struct htp_tensor * dst = octx->dst;                                                                      \
-    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                          \
-                                                                                                                    \
-    htp_unary_preamble;                                                                                             \
-                                                                                                                    \
-    int32_t *      op_params = octx->op_params;                                                                     \
-    const uint32_t col_tile  = uctx->col_tile;                                                                      \
-                                                                                                                    \
-    const uint32_t src0_nrows     = uctx->src0_nrows;                                                               \
-    const uint32_t src0_start_row = uctx->src0_nrows_per_thread * ith;                                              \
-    const uint32_t src0_end_row   = MIN(src0_start_row + uctx->src0_nrows_per_thread, src0_nrows);                  \
-                                                                                                                    \
-    if (src0_start_row >= src0_end_row) {                                                                           \
-        return;                                                                                                     \
-    }                                                                                                               \
-                                                                                                                    \
-    const uint8_t * restrict data_src = uctx->data_src0;                                                            \
-    uint8_t * restrict       data_dst = uctx->data_dst;                                                             \
-                                                                                                                    \
-    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);                           \
-    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);                             \
-                                                                                                                    \
-    const size_t src0_half = uctx->src0_vtcm_half_size;                                                             \
-    const size_t dst_half  = uctx->dst_vtcm_half_size;                                                              \
-                                                                                                                    \
-    dma_queue * dmaq = octx->ctx->dma[ith];                                                                         \
-                                                                                                                    \
-    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;                                             \
-    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;                                             \
-    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;                                            \
-    const struct fastdiv_values * div_tpr   = &uctx->kparams->div_tpr;                                              \
-                                                                                                                    \
-    const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile;                                                 \
-    const int32_t  tri_ttype     = (IS_TRI) ? op_params[0] : 0;                                                     \
-                                                                                                                    \
-    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&                                                       \
-                             (nb03 == (size_t)ne02 * nb02);                                                         \
-    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&                                                       \
-                             (nb3  == (size_t)ne2  * nb2);                                                          \
-                                                                                                                    \
-    const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row;                                   \
-                                                                                                                    \
-    for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) {                          \
-        const uint32_t row  = src0_start_row + t / tiles_per_row;                                                   \
-        const uint32_t col  = (t % tiles_per_row) * col_tile;                                                       \
-        const uint32_t tw   = MIN(col_tile, ne0 - col);                                                             \
-        const size_t   tb   = (size_t) tw * sizeof(float);                                                          \
-        const size_t   soff = (src0_contig ? (row * nb01) :                                                         \
-                               unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +\
-                               (size_t) col * sizeof(float);                                                        \
-                                                                                                                    \
-        dma_queue_push(dmaq, dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0);            \
-        dma_queue_push(dmaq, dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);\
-    }                                                                                                               \
-                                                                                                                    \
-    uint32_t row = src0_start_row;                                                                                  \
-    uint32_t col = 0;                                                                                               \
-    uint32_t tile_in_row = 0;                                                                                       \
-    uint32_t i01 = fastmodulo(row, ne01, div_ne01);                                                                 \
-                                                                                                                    \
-    uint32_t prow = src0_start_row + fastdiv(2, div_tpr);                                                           \
-    uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile;                                               \
-    uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr);                                                  \
-                                                                                                                    \
-    for (uint32_t t = 0; t < total_tiles; t++) {                                                                    \
-        uint8_t * dst_vtcm = (uint8_t *) dma_queue_pop(dmaq).src;                                                   \
-        uint8_t * src_vtcm = (uint8_t *) dma_queue_pop(dmaq).dst;                                                   \
-                                                                                                                    \
-        const uint32_t tw  = MIN(col_tile, ne0 - col);                                                              \
-                                                                                                                    \
-        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                       \
-        CORE_TILE_EXPR;                                                                                             \
-        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                        \
-                                                                                                                    \
-        const size_t doff = (dst_contig ? (row * nb1) :                                                             \
-                             unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) +       \
-                             (size_t) col * sizeof(float);                                                          \
-        const size_t tb   = (size_t) tw * sizeof(float);                                                            \
-        dma_queue_push(dmaq, dma_make_ptr(data_dst + doff, dst_vtcm), tb, tb, tb, 1);                               \
-                                                                                                                    \
-        const uint32_t pt = t + 2;                                                                                  \
-        if (pt < total_tiles) {                                                                                     \
-            const uint32_t ptw  = MIN(col_tile, ne0 - pcol);                                                        \
-            const size_t   ptb  = (size_t) ptw * sizeof(float);                                                     \
-            const size_t   psoff = (src0_contig ? (prow * nb01) :                                                   \
-                                    unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02,   \
-                                                     nb03)) +                                                       \
-                                   (size_t) pcol * sizeof(float);                                                   \
-            dma_queue_push(dmaq, dma_make_ptr(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1);                       \
-        }                                                                                                           \
-                                                                                                                    \
-        tile_in_row++;                                                                                              \
-        col += col_tile;                                                                                            \
-        if (tile_in_row == tiles_per_row) {                                                                         \
-            tile_in_row = 0;                                                                                        \
-            col = 0;                                                                                                \
-            row++;                                                                                                  \
-            i01++;                                                                                                  \
-            if (i01 == ne01) {                                                                                      \
-                i01 = 0;                                                                                            \
-            }                                                                                                       \
-        }                                                                                                           \
-                                                                                                                    \
-        ptile_in_row++;                                                                                             \
-        pcol += col_tile;                                                                                           \
-        if (ptile_in_row == tiles_per_row) {                                                                        \
-            ptile_in_row = 0;                                                                                       \
-            pcol = 0;                                                                                               \
-            prow++;                                                                                                 \
-        }                                                                                                           \
-    }                                                                                                               \
-                                                                                                                    \
-    dma_queue_flush(dmaq);                                                                                          \
+#define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR)                                                         \
+static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) {                            \
+    const struct htp_unary_context * uctx = (const struct htp_unary_context *) data;                                  \
+    struct htp_ops_context * octx = uctx->octx;                                                                       \
+    const struct htp_tensor * src = octx->src[0];                                                                     \
+    const struct htp_tensor * dst = octx->dst;                                                                        \
+    struct htp_thread_trace * tr = &octx->ctx->trace[ith];                                                            \
+                                                                                                                      \
+    htp_unary_preamble;                                                                                               \
+                                                                                                                      \
+    uint32_t     src0_nrows_per_thread = uctx->src0_nrows_per_thread;                                                 \
+                                                                                                                      \
+    int32_t *      op_params = octx->op_params;                                                                       \
+    const uint32_t col_tile  = uctx->col_tile;                                                                        \
+                                                                                                                      \
+    const uint32_t src0_nrows = uctx->src0_nrows;                                                                     \
+    const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith;                                    \
+    const uint32_t src0_end_row   = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows);        \
+                                                                                                                      \
+    if (src0_start_row >= src0_end_row) {                                                                             \
+        return;                                                                                                       \
+    }                                                                                                                 \
+                                                                                                                      \
+    const uint8_t * restrict data_src = uctx->data_src0;                                                              \
+    uint8_t * restrict       data_dst = uctx->data_dst;                                                               \
+                                                                                                                      \
+    uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread);                             \
+    uint8_t * dst_vtcm_data  = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread);                               \
+                                                                                                                      \
+    const size_t src0_half = uctx->src0_vtcm_half_size;                                                               \
+    const size_t dst_half  = uctx->dst_vtcm_half_size;                                                                \
+                                                                                                                      \
+    dma_queue * dmaq = octx->ctx->dma[ith];                                                                           \
+                                                                                                                      \
+    const struct fastdiv_values * div_ne01  = &uctx->kparams->div_ne01;                                               \
+    const struct fastdiv_values * div_ne02  = &uctx->kparams->div_ne02;                                               \
+    const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012;                                              \
+    const struct fastdiv_values * div_tpr   = &uctx->kparams->div_tpr;                                                \
+                                                                                                                      \
+    const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile;                                                   \
+    const int32_t  tri_ttype     = (IS_TRI) ? op_params[0] : 0;                                                       \
+                                                                                                                      \
+    const bool src0_contig = (nb02 == (size_t)ne01 * nb01) &&                                                         \
+                             (nb03 == (size_t)ne02 * nb02);                                                           \
+    const bool dst_contig  = (nb2  == (size_t)ne1  * nb1)  &&                                                         \
+                             (nb3  == (size_t)ne2  * nb2);                                                            \
+                                                                                                                      \
+    const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row;                                     \
+                                                                                                                      \
+    for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) {                            \
+        const uint32_t row  = src0_start_row + t / tiles_per_row;                                                     \
+        const uint32_t col  = (t % tiles_per_row) * col_tile;                                                         \
+        const uint32_t tw   = MIN(col_tile, ne0 - col);                                                               \
+        const size_t   tb   = (size_t) tw * sizeof(float);                                                            \
+        const size_t   soff = (src0_contig ? (row * nb01) :                                                           \
+                               unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +  \
+                               (size_t) col * sizeof(float);                                                          \
+                                                                                                                      \
+        dma_queue_push(dmaq, dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0);              \
+        dma_queue_push(dmaq, dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);  \
+    }                                                                                                                 \
+                                                                                                                      \
+    uint32_t row = src0_start_row;                                                                                    \
+    uint32_t col = 0;                                                                                                 \
+    uint32_t tile_in_row = 0;                                                                                         \
+    uint32_t i01 = fastmodulo(row, ne01, div_ne01);                                                                   \
+                                                                                                                      \
+    uint32_t prow = src0_start_row + fastdiv(2, div_tpr);                                                             \
+    uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile;                                                 \
+    uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr);                                                    \
+                                                                                                                      \
+    for (uint32_t t = 0; t < total_tiles; t++) {                                                                      \
+        uint8_t * dst_vtcm = (uint8_t *) dma_queue_pop(dmaq).src;                                                     \
+        uint8_t * src_vtcm = (uint8_t *) dma_queue_pop(dmaq).dst;                                                     \
+                                                                                                                      \
+        const uint32_t tw  = MIN(col_tile, ne0 - col);                                                                \
+                                                                                                                      \
+        htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                         \
+        CORE_TILE_EXPR;                                                                                               \
+        htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t);                                                          \
+                                                                                                                      \
+        const size_t doff = (dst_contig ? (row * nb1) :                                                               \
+                             unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) +         \
+                             (size_t) col * sizeof(float);                                                            \
+        const size_t tb   = (size_t) tw * sizeof(float);                                                              \
+        dma_queue_push(dmaq, dma_make_ptr(data_dst + doff, dst_vtcm), tb, tb, tb, 1);                                 \
+                                                                                                                      \
+        const uint32_t pt = t + 2;                                                                                    \
+        if (pt < total_tiles) {                                                                                       \
+            const uint32_t ptw  = MIN(col_tile, ne0 - pcol);                                                          \
+            const size_t   ptb  = (size_t) ptw * sizeof(float);                                                       \
+            const size_t   psoff = (src0_contig ? (prow * nb01) :                                                     \
+                                    unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02,     \
+                                                     nb03)) +                                                         \
+                                   (size_t) pcol * sizeof(float);                                                     \
+            dma_queue_push(dmaq, dma_make_ptr(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1);                         \
+        }                                                                                                             \
+                                                                                                                      \
+        tile_in_row++;                                                                                                \
+        col += col_tile;                                                                                              \
+        if (tile_in_row == tiles_per_row) {                                                                           \
+            tile_in_row = 0;                                                                                          \
+            col = 0;                                                                                                  \
+            row++;                                                                                                    \
+            i01++;                                                                                                    \
+            if (i01 == ne01) {                                                                                        \
+                i01 = 0;                                                                                              \
+            }                                                                                                         \
+        }                                                                                                             \
+                                                                                                                      \
+        ptile_in_row++;                                                                                               \
+        pcol += col_tile;                                                                                             \
+        if (ptile_in_row == tiles_per_row) {                                                                          \
+            ptile_in_row = 0;                                                                                         \
+            pcol = 0;                                                                                                 \
+            prow++;                                                                                                   \
+        }                                                                                                             \
+    }                                                                                                                 \
+                                                                                                                      \
+    dma_queue_flush(dmaq);                                                                                            \
 }

 static inline void tile_scale_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw, const int32_t * op_params) {
@@ -1146,14 +1149,32 @@ static int execute_op_unary(struct htp_ops_context * octx) {

     const struct htp_unary_kernel_params * kparams = (const struct htp_unary_kernel_params *) octx->kernel_params;

-    const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
-    const uint32_t n_threads  = kparams->n_threads;
+    if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
+        return HTP_STATUS_INVAL_PARAMS;
+    }

+    const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
     const size_t elem_size = is_f16 ? sizeof(_Float16) : sizeof(float);
-
     const size_t src0_data_row_size = src0->ne[0] * elem_size;
     const size_t dst_data_row_size  = dst->ne[0]  * elem_size;

+    uint32_t row_start = 0;
+    uint32_t nrows     = src0_nrows;
+
+    if (octx->ctx->mdev.count > 1) {
+        uint32_t rows_per_chunk = 0;
+        htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_data_row_size, &rows_per_chunk);
+        const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
+        row_start = range.start;
+        nrows     = range.count;
+    }
+
+    if (nrows == 0) {
+        return HTP_STATUS_OK;
+    }
+
+    const uint32_t n_threads = octx->n_threads;
+
     const size_t src0_row_size_aligned = kparams->src0_row_size_aligned;
     const size_t dst_row_size_aligned  = kparams->dst_row_size_aligned;

@@ -1191,8 +1212,9 @@ static int execute_op_unary(struct htp_ops_context * octx) {
         struct htp_unary_context uctx = {
             .octx                  = octx,
             .kparams               = kparams,
-            .src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads,
-            .src0_nrows            = src0_nrows,
+            .src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
+            .src0_nrows            = nrows,
+            .row_start             = row_start,

             .data_src0             = (const uint8_t *)src0->data,
             .data_src1             = (octx->op == HTP_OP_RMS_NORM_MUL) ? (const uint8_t *)src1->data : NULL,
@@ -1287,7 +1309,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
         }

         if (task_func) {
-            worker_pool_run_func(octx->ctx->worker_pool, task_func, &uctx, n_threads);
+            work_queue_run(octx->ctx->work_queue, task_func, &uctx, n_threads);
         } else {
             FARF(ERROR, "execute_op_unary: task function is NULL for op %d\n", octx->op);
             err = HTP_STATUS_NO_SUPPORT;
diff --git a/scripts/snapdragon/ggml-hexagon-align-macros.py b/scripts/snapdragon/ggml-hexagon-align-macros.py
new file mode 100755
index 000000000..b64db3e65
--- /dev/null
+++ b/scripts/snapdragon/ggml-hexagon-align-macros.py
@@ -0,0 +1,296 @@
+#!/usr/bin/env python3
+"""
+align-macros.py - Inspect and align trailing backslashes in multiline C/C++ macros.
+
+Usage:
+    align-macros.py [paths...]                 # Check and report misaligned macros
+    align-macros.py --diff [paths...]          # Show unified diff of fixes
+    align-macros.py --fix [paths...]           # Fix misaligned macros in-place
+    align-macros.py --fix --mode majority ...  # Align to the dominant column
+    align-macros.py --fix --pad 2 ...          # Align to (max_content_len + pad)
+
+Safety rules:
+    - Macros that are ALREADY aligned are NEVER touched (unless --all is given).
+    - Whitespace after trailing backslashes is flagged and cleaned.
+"""
+
+import argparse
+import difflib
+import logging
+import os
+import re
+import sys
+from collections import Counter
+from typing import List, Optional, Tuple, NamedTuple
+
+logger = logging.getLogger("ggml-hexagon-align-macros")
+
+
+class MacroLine(NamedTuple):
+    line_num: int       # 1-indexed
+    raw: str            # Original line including newline
+    content: str        # Line content before trailing backslash (stripped of trailing whitespace)
+    bs_col: Optional[int]  # 1-indexed column of backslash, or None if last line has no backslash
+    trailing_ws: bool   # True if whitespace existed after the backslash
+
+
+class MacroDef(NamedTuple):
+    name: str
+    filepath: str
+    start_line: int
+    end_line: int
+    lines: List[MacroLine]
+
+
+def parse_macros(filepath: str) -> List[MacroDef]:
+    """Extract all multiline macros from a C/C++ source file."""
+    try:
+        with open(filepath, "r", encoding="utf-8", errors="replace") as f:
+            lines = f.readlines()
+    except Exception as e:
+        logger.error(f"Error reading {filepath}: {e}")
+        return []
+
+    macros: List[MacroDef] = []
+    i = 0
+    n = len(lines)
+
+    while i < n:
+        line = lines[i]
+        m = re.match(r"^\s*#\s*define\s+([A-Za-z_][A-Za-z0-9_]*)", line)
+        if m:
+            macro_name = m.group(1)
+            macro_start = i + 1
+            macro_lines: List[MacroLine] = []
+            cur = i
+
+            while cur < n:
+                l_raw = lines[cur]
+                l_rstrip = l_raw.rstrip("\r\n")
+
+                # Check if line has a trailing backslash
+                # Note: handle possible accidental spaces after backslash
+                match_bs = re.search(r"\\([ \t]*)$", l_rstrip)
+                if match_bs:
+                    has_trailing_ws = len(match_bs.group(1)) > 0
+                    bs_index = match_bs.start()
+                    content = l_rstrip[:bs_index].rstrip()
+                    # 1-indexed column of the backslash
+                    bs_col = bs_index + 1
+                    macro_lines.append(MacroLine(
+                        line_num=cur + 1,
+                        raw=l_raw,
+                        content=content,
+                        bs_col=bs_col,
+                        trailing_ws=has_trailing_ws
+                    ))
+                    cur += 1
+                else:
+                    # Line does not end with backslash
+                    if cur == i:
+                        # Single-line macro, not multiline
+                        break
+                    else:
+                        # Final line of a multiline macro
+                        macro_lines.append(MacroLine(
+                            line_num=cur + 1,
+                            raw=l_raw,
+                            content=l_rstrip.rstrip(),
+                            bs_col=None,
+                            trailing_ws=False
+                        ))
+                        break
+
+            # Only record if it is a multiline macro (has at least one continuation line)
+            continuation_lines = [ml for ml in macro_lines if ml.bs_col is not None]
+            if continuation_lines:
+                macro_end = macro_lines[-1].line_num
+                macros.append(MacroDef(
+                    name=macro_name,
+                    filepath=filepath,
+                    start_line=macro_start,
+                    end_line=macro_end,
+                    lines=macro_lines
+                ))
+            i = cur
+        i += 1
+
+    return macros
+
+
+def is_macro_aligned(macro: MacroDef) -> bool:
+    """A macro is aligned if all continuation lines have backslashes at the same column."""
+    bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
+    if not bs_cols:
+        return True
+    has_trailing_ws = any(ml.trailing_ws for ml in macro.lines)
+    return len(set(bs_cols)) == 1 and not has_trailing_ws
+
+
+def compute_target_column(macro: MacroDef, mode: str, pad: int, target_col: Optional[int]) -> int:
+    """Determine the column where backslashes should be aligned."""
+    max_content_len = max(len(ml.content) for ml in macro.lines)
+    min_needed = max_content_len + pad
+
+    if target_col is not None:
+        return max(target_col, min_needed)
+
+    bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
+    if not bs_cols:
+        return min_needed
+
+    if mode == "min":
+        return min_needed
+    elif mode == "max":
+        return max(max(bs_cols), min_needed)
+    elif mode == "majority":
+        counts = Counter(bs_cols)
+        # Sort by frequency descending, then by column descending
+        majority_col = sorted(counts.items(), key=lambda x: (-x[1], -x[0]))[0][0]
+        return max(majority_col, min_needed)
+    else:
+        return min_needed
+
+
+def realign_macro_lines(macro: MacroDef, target_col: int) -> List[str]:
+    """Format macro lines with backslashes aligned at target_col."""
+    new_lines: List[str] = []
+    for ml in macro.lines:
+        nl = "\r\n" if ml.raw.endswith("\r\n") else "\n"
+        if ml.bs_col is None:
+            # Last line without backslash
+            new_lines.append(ml.raw)
+        else:
+            if not ml.content:
+                spaces = " " * (target_col - 1)
+                new_lines.append(f"{spaces}\\{nl}")
+            else:
+                spaces_needed = max(1, target_col - len(ml.content) - 1)
+                new_lines.append(f"{ml.content}{' ' * spaces_needed}\\{nl}")
+    return new_lines
+
+
+def process_file(filepath: str, args: argparse.Namespace) -> Tuple[int, int, Optional[str]]:
+    macros = parse_macros(filepath)
+    if not macros:
+        return 0, 0, None
+
+    with open(filepath, "r", encoding="utf-8", errors="replace") as f:
+        file_lines = f.readlines()
+
+    misaligned_count = 0
+    modified = False
+    new_file_lines = list(file_lines)
+
+    for macro in macros:
+        aligned = is_macro_aligned(macro)
+        if not aligned or args.all:
+            if not aligned:
+                misaligned_count += 1
+
+            bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
+            max_content = max(len(ml.content) for ml in macro.lines)
+            col_counts = Counter(bs_cols)
+
+            if not args.quiet:
+                logger.info(f"{filepath}:{macro.start_line}-{macro.end_line} [{macro.name}]")
+                logger.info(f"  Max content width: {max_content}, Min needed column (+{args.pad}): {max_content + args.pad}")
+                logger.info(f"  Current backslash columns: {dict(sorted(col_counts.items()))}")
+                trailing_ws_lines = [ml.line_num for ml in macro.lines if ml.trailing_ws]
+                if trailing_ws_lines:
+                    logger.warning(f"  Warning: Trailing whitespace after backslash on line(s): {trailing_ws_lines}")
+
+            target_col = compute_target_column(macro, args.mode, args.pad, args.target_col)
+            if not args.quiet:
+                logger.info(f"  -> Target alignment column: {target_col}")
+
+            realigned = realign_macro_lines(macro, target_col)
+
+            start_idx = macro.start_line - 1
+            end_idx = start_idx + len(macro.lines)
+            if new_file_lines[start_idx:end_idx] != realigned:
+                new_file_lines[start_idx:end_idx] = realigned
+                modified = True
+
+    diff_text = None
+    if modified:
+        diff = difflib.unified_diff(
+            file_lines,
+            new_file_lines,
+            fromfile=f"a/{filepath}",
+            tofile=f"b/{filepath}",
+            lineterm=""
+        )
+        diff_text = "\n".join(diff)
+
+        if args.fix:
+            with open(filepath, "w", encoding="utf-8") as f:
+                f.writelines(new_file_lines)
+            if not args.quiet:
+                logger.info(f"  [FIXED] Updated {filepath}")
+
+    return len(macros), misaligned_count, diff_text
+
+
+def find_source_files(paths: List[str]) -> List[str]:
+    extensions = {".c", ".cpp", ".cc", ".cxx", ".h", ".hpp", ".inl"}
+    result: List[str] = []
+    for p in paths:
+        if os.path.isfile(p):
+            result.append(p)
+        elif os.path.isdir(p):
+            for root, _, files in os.walk(p):
+                for file in sorted(files):
+                    _, ext = os.path.splitext(file)
+                    if ext.lower() in extensions:
+                        result.append(os.path.join(root, file))
+    return sorted(result)
+
+
+def main():
+    logging.basicConfig(level=logging.INFO, format="%(message)s")
+    parser = argparse.ArgumentParser(
+        description="Inspect and align backslashes in multiline C/C++ macros."
+    )
+    parser.add_argument("paths", nargs="*", default=["."], help="Files or directories to scan (default: current dir)")
+    parser.add_argument("--fix", action="store_true", help="Fix misaligned macros in-place")
+    parser.add_argument("--diff", action="store_true", help="Display unified diff of suggested fixes")
+    parser.add_argument("--check", action="store_true", help="Exit with code 1 if misaligned macros exist")
+    parser.add_argument("--mode", choices=["min", "max", "majority"], default="min",
+                        help="Alignment mode: 'min' (max_len + pad), 'max' (max existing col), 'majority' (dominant col)")
+    parser.add_argument("--pad", type=int, default=2, help="Spaces between longest line and backslash (default: 2)")
+    parser.add_argument("--target-col", type=int, default=None, help="Force alignment to an exact column")
+    parser.add_argument("--all", action="store_true", help="Realign all macros even if already aligned (default: only misaligned)")
+    parser.add_argument("-q", "--quiet", action="store_true", help="Only output errors and diffs/summary")
+
+    args = parser.parse_args()
+
+    files = find_source_files(args.paths)
+    if not files:
+        logger.error("No C/C++ source files found.")
+        sys.exit(0)
+
+    total_macros = 0
+    total_misaligned = 0
+    diffs: List[str] = []
+
+    for filepath in files:
+        num_macros, num_misaligned, diff_text = process_file(filepath, args)
+        total_macros += num_macros
+        total_misaligned += num_misaligned
+        if diff_text:
+            diffs.append(diff_text)
+
+    if args.diff and diffs:
+        logger.info("\n--- Proposed Changes ---\n")
+        for d in diffs:
+            logger.info(d)
+
+    logger.info(f"\nSummary: scanned {len(files)} files, {total_macros} multiline macros, {total_misaligned} misaligned.")
+
+    if args.check and total_misaligned > 0:
+        sys.exit(1)
+
+
+if __name__ == "__main__":
+    main()
diff --git a/scripts/snapdragon/run.py b/scripts/snapdragon/run.py
index 81eecd2e0..dc71d4a32 100755
--- a/scripts/snapdragon/run.py
+++ b/scripts/snapdragon/run.py
@@ -14,6 +14,42 @@ import logging
 logger = logging.getLogger("run")


+MANAGED_ENV_NAMES = (
+    "GGML_HEXAGON_DEVICES",
+    "GGML_HEXAGON_VERBOSE",
+    "GGML_HEXAGON_PROFILE",
+    "GGML_HEXAGON_NHVX",
+    "GGML_HEXAGON_NHMX",
+    "GGML_HEXAGON_HOSTBUF",
+    "GGML_HEXAGON_OPBATCH",
+    "GGML_HEXAGON_OPQUEUE",
+    "GGML_HEXAGON_OPPOLL",
+    "GGML_HEXAGON_OPFILTER",
+    "GGML_HEXAGON_OPFUSION",
+    "GGML_HEXAGON_VMEM",
+    "GGML_HEXAGON_MBUF",
+    "GGML_HEXAGON_MM_SELECT",
+    "GGML_HEXAGON_FA_SELECT",
+    "GGML_HEXAGON_AR_SELECT",
+    "GGML_HEXAGON_ETM",
+    "GGML_HEXAGON_ARCH",
+    "GGML_HEXAGON_OPTRACE",
+    "GGML_OPENCL_PLATFORM",
+    "GGML_OPENCL_DEVICE",
+    "GGML_OPENCL_OPFILTER",
+    "GGML_OPENCL_KERNEL_CACHE_DIR",
+    "GGML_OPENCL_KERNEL_CACHE_DEBUG",
+    "GGML_OPENCL_FA_TUNE",
+    "GGML_OPENCL_DISABLE_FUSION",
+    "GGML_OPENCL_ADRENO_XMEM_GEMM",
+    "GGML_OPENCL_ADRENO_USE_LARGE_BUFFER",
+    "GGML_SCHED_DEBUG",
+    "MTMD_BACKEND_DEVICE",
+    "D",
+    "DEVICE",
+)
+
+
 def parse_target(target_str):
     if not target_str:
         return None, None
@@ -38,6 +74,57 @@ def shlex_join(args_list):
     return " ".join(pipes.quote(x) for x in args_list)


+def split_device_list(devices):
+    parts = []
+    curr = []
+    bracket_depth = 0
+
+    for ch in devices:
+        if ch == '[':
+            bracket_depth += 1
+            curr.append(ch)
+        elif ch == ']':
+            if bracket_depth > 0:
+                bracket_depth -= 1
+            curr.append(ch)
+        elif ch == ',' and bracket_depth == 0:
+            part = "".join(curr).strip()
+            if part:
+                parts.append(part)
+            curr = []
+        else:
+            curr.append(ch)
+
+    part = "".join(curr).strip()
+    if part:
+        parts.append(part)
+
+    return parts
+
+
+def device_arg_from_devices(devices):
+    if devices.isdigit():
+        n = int(devices)
+        return ",".join(f"HTP{i}" for i in range(n))
+
+    names = []
+    for part in split_device_list(devices):
+        if "[" in part:
+            part = part.split("[", 1)[0].strip()
+        if part:
+            names.append(part)
+
+    return ",".join(names)
+
+
+def normalize_cmd_device_args(cmd_args):
+    for i, arg in enumerate(cmd_args):
+        if arg == "--device" and i + 1 < len(cmd_args):
+            cmd_args[i + 1] = device_arg_from_devices(cmd_args[i + 1])
+        elif arg.startswith("--device="):
+            cmd_args[i] = "--device=" + device_arg_from_devices(arg.split("=", 1)[1])
+
+
 def main():
     logging.basicConfig(level=logging.INFO, format='%(message)s')
     # Split arguments at '--'
@@ -142,8 +229,6 @@ def main():
     def set_env(env_name, opt_val):
         if opt_val is not None:
             env_vars[env_name] = str(opt_val)
-        elif env_name in os.environ:
-            env_vars[env_name] = os.environ[env_name]

     # Resolve and filter devices (HTP vs OpenCL)
     device_in_cmd = None
@@ -166,7 +251,7 @@ def main():
         hex_devices = devices_val
         cl_device = ""
     else:
-        parts = [p.strip() for p in devices_val.split(",")]
+        parts = split_device_list(devices_val)
         # Any device containing "htp" is Hexagon, rest is OpenCL
         hex_parts = [p for p in parts if "htp" in p.lower()]
         cl_parts = [
@@ -181,15 +266,13 @@ def main():
     # Set Hexagon devices
     if hex_devices:
         env_vars["GGML_HEXAGON_DEVICES"] = hex_devices
-    elif "GGML_HEXAGON_DEVICES" in os.environ:
-        env_vars["GGML_HEXAGON_DEVICES"] = os.environ["GGML_HEXAGON_DEVICES"]
+
+    normalize_cmd_device_args(cmd_args)

     # Set OpenCL device (unless overridden by --cl-device)
     final_cl_device = args.cl_device if args.cl_device is not None else cl_device
     if final_cl_device:
         env_vars["GGML_OPENCL_DEVICE"] = final_cl_device
-    elif "GGML_OPENCL_DEVICE" in os.environ:
-        env_vars["GGML_OPENCL_DEVICE"] = os.environ["GGML_OPENCL_DEVICE"]

     # Map shared & backend-specific parameters with correct overrides

@@ -206,8 +289,6 @@ def main():

     if args.cl_fa_tune or args.profile is not None:
         env_vars["GGML_OPENCL_FA_TUNE"] = "1"
-    elif "GGML_OPENCL_FA_TUNE" in os.environ:
-        env_vars["GGML_OPENCL_FA_TUNE"] = os.environ["GGML_OPENCL_FA_TUNE"]

     # Other Hexagon environment variables
     set_env("GGML_HEXAGON_NHVX", args.hex_nhvx)
@@ -235,18 +316,12 @@ def main():

     if args.cl_disable_fusion:
         env_vars["GGML_OPENCL_DISABLE_FUSION"] = "1"
-    elif "GGML_OPENCL_DISABLE_FUSION" in os.environ:
-        env_vars["GGML_OPENCL_DISABLE_FUSION"] = os.environ["GGML_OPENCL_DISABLE_FUSION"]

     if args.cl_adreno_xmem:
         env_vars["GGML_OPENCL_ADRENO_XMEM_GEMM"] = "1"
-    elif "GGML_OPENCL_ADRENO_XMEM_GEMM" in os.environ:
-        env_vars["GGML_OPENCL_ADRENO_XMEM_GEMM"] = os.environ["GGML_OPENCL_ADRENO_XMEM_GEMM"]

     if args.cl_adreno_large_buffer:
         env_vars["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"] = "1"
-    elif "GGML_OPENCL_ADRENO_USE_LARGE_BUFFER" in os.environ:
-        env_vars["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"] = os.environ["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"]

     if args.sched_debug:
         env_vars["GGML_SCHED_DEBUG"] = "2"
@@ -288,15 +363,7 @@ def main():
         has_b = any(arg == "-b" for arg in cmd_args)
         if not has_b:
             if args.devices:
-                if args.devices.isdigit():
-                    n = int(args.devices)
-                    device_val = ",".join(f"HTP{i}" for i in range(n))
-                else:
-                    device_val = args.devices
-            elif "D" in os.environ:
-                device_val = os.environ["D"]
-            elif "DEVICE" in os.environ:
-                device_val = os.environ["DEVICE"]
+                device_val = device_arg_from_devices(args.devices)
             else:
                 device_val = "HTP0"
             if device_val:
@@ -305,17 +372,10 @@ def main():
         has_device = any(arg.startswith("--device") for arg in cmd_args)
         if not has_device:
             if args.devices:
-                if args.devices.isdigit():
-                    n = int(args.devices)
-                    device_val = ",".join(f"HTP{i}" for i in range(n))
-                else:
-                    device_val = args.devices
-            elif "D" in os.environ:
-                device_val = os.environ["D"]
-            elif "DEVICE" in os.environ:
-                device_val = os.environ["DEVICE"]
+                device_val = device_arg_from_devices(args.devices)
             else:
                 device_val = "HTP0"
+
             if device_val:
                 cmd_args += ["--device", device_val]

@@ -415,6 +475,8 @@ def main():
         else:
             local_env["LD_LIBRARY_PATH"] = lib_dir + os.path.pathsep + local_env.get("LD_LIBRARY_PATH", "")

+        for k in MANAGED_ENV_NAMES:
+            local_env.pop(k, None)
         for k, v in env_vars.items():
             local_env[k] = v