integrations.expert_parallel.torch_dispatch
integrations.expert_parallel.torch_dispatch
Token-dispatch expert parallelism over plain all_to_all_single (NCCL or gloo).
The layout matches DeepEP’s contract so every local experts kernel runs unchanged: each
token is sent ONCE to every rank that owns at least one of its top-k experts, carrying a
K-wide recv_topk_idx (local expert ids in [0, E_local) for that rank, -1 for
slots owned elsewhere) and K-wide fp32 recv_topk_weights (0 on -1 slots). The
local kernel weights and sums over K, and combine sums the per-rank partial outputs
back onto the source tokens.
The collectives are registered as torch.ops.axolotl.ep_all_to_all_single and
torch.ops.axolotl.ep_all_to_all_single_equal so selective activation checkpointing
can match (and save) them by name.
Hang safety: the split sizes are derived from the routing, so every rank must see the
same routing for the same step. Recomputation under activation checkpointing must reuse
the forward’s topk result and CPU split copies rather than re-running them.
Classes
| Name | Description |
|---|---|
| TorchEPHandle | State combine needs to return expert outputs to their source tokens. |
TorchEPHandle
integrations.expert_parallel.torch_dispatch.TorchEPHandle(
num_tokens,
send_token_idx=None,
send_splits=None,
recv_splits=None,
group=None,
recv_x=None,
recv_w=None,
)State combine needs to return expert outputs to their source tokens.
Functions
| Name | Description |
|---|---|
| all_to_all_single | Differentiable uneven all-to-all along dim 0. output_splits / input_splits are |
| all_to_all_single_async | Differentiable uneven all-to-all whose wait is deferred to the output’s first use. |
| all_to_all_single_equal | Differentiable equal-split all-to-all along dim 0. |
| build_send_layout | Rank-major, token-ascending send order with per-destination local expert ids. |
| chunk_bounds | chunks contiguous [start, end) token ranges; the last takes the remainder, |
| combine | Return [N_recv,H] expert outputs to their source ranks and sum them per token. |
| compute_send_counts | Return (dest [T,K], is_in_rank [T,P], send_counts [P]) for global topk_idx |
| dispatch | Send each token once to every rank owning one of its experts. |
| dispatch_chunked_forward | dispatch -> local_kernel -> combine over chunks token ranges, pipelined. |
| to_host | Device->host copy of the split counts, dispatched as axolotl::ep_to_host. |
all_to_all_single
integrations.expert_parallel.torch_dispatch.all_to_all_single(
x,
output_splits,
input_splits,
group,
)Differentiable uneven all-to-all along dim 0. output_splits / input_splits are
CPU int64 tensors of length group.size(); backward is the same op with them swapped.
all_to_all_single_async
integrations.expert_parallel.torch_dispatch.all_to_all_single_async(
x,
output_splits,
input_splits,
group,
)Differentiable uneven all-to-all whose wait is deferred to the output’s first use.
Returns an AsyncCollectiveTensor so compute enqueued before that use overlaps the
collective. Dispatches as _c10d_functional::all_to_all_single / wait_tensor.
all_to_all_single_equal
integrations.expert_parallel.torch_dispatch.all_to_all_single_equal(x, group)Differentiable equal-split all-to-all along dim 0.
build_send_layout
integrations.expert_parallel.torch_dispatch.build_send_layout(
topk_idx,
topk_weights,
dest,
is_in_rank,
num_send,
num_local_experts,
)Rank-major, token-ascending send order with per-destination local expert ids.
Returns (send_token_idx [N], send_rank [N], send_idx [N,K], send_w [N,K]). num_send
must equal is_in_rank.sum(); passing it from the host avoids a device sync.
chunk_bounds
integrations.expert_parallel.torch_dispatch.chunk_bounds(num_tokens, chunks)chunks contiguous [start, end) token ranges; the last takes the remainder,
so some are empty when num_tokens < chunks.
combine
integrations.expert_parallel.torch_dispatch.combine(
local_out,
handle,
dtype=None,
)Return [N_recv,H] expert outputs to their source ranks and sum them per token.
compute_send_counts
integrations.expert_parallel.torch_dispatch.compute_send_counts(
topk_idx,
num_ranks,
num_local_experts,
)Return (dest [T,K], is_in_rank [T,P], send_counts [P]) for global topk_idx
(-1 = unrouted). dest is the owning rank per slot, -1 where unrouted.
dispatch
integrations.expert_parallel.torch_dispatch.dispatch(
x,
topk_idx,
topk_weights,
*,
num_local_experts,
group,
)Send each token once to every rank owning one of its experts.
x is [T,H], topk_idx [T,K] int64 global expert ids (-1 = unrouted),
topk_weights [T,K] fp32. Returns (recv_x, recv_idx, recv_w, handle) in the
DeepEP layout. Tokens with no routed expert are not sent; their combined output is zero.
group=None or a size-1 group is the collective-free local path.
dispatch_chunked_forward
integrations.expert_parallel.torch_dispatch.dispatch_chunked_forward(
x,
topk_idx,
topk_weights,
local_kernel,
*,
num_local_experts,
group,
chunks,
)dispatch -> local_kernel -> combine over chunks token ranges, pipelined.
Chunk i+1’s dispatch all-to-alls are issued before chunk i’s local kernel, and
chunk i’s combine is consumed only after chunk i+1’s kernel is enqueued, so the
collectives overlap expert compute in forward. The split counts of every chunk come
from ONE count exchange and ONE device->host copy. Every rank issues every chunk’s
collectives in the same order, including empty ones.
to_host
integrations.expert_parallel.torch_dispatch.to_host(x)Device->host copy of the split counts, dispatched as axolotl::ep_to_host.