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.