# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from abc import ABC, abstractmethod
from dataclasses import dataclass

import spmd_types as spmd
import torch
import torch_remat as remat
from torch.distributed._functional_collectives import all_to_all_single
from torch.distributed.tensor import DeviceMesh

from torchtitan.distributed.spmd_types import maybe_set_sparse_mesh, spmd_sparse_mesh
from torchtitan.ops.scatter_add import deterministic_scatter_add
from torchtitan.protocols.module import Module


@dataclass(frozen=True, kw_only=True)
class LocalDispatchMetadata:
    """Metadata returned by LocalTokenDispatcher.dispatch() for use in combine()."""

    token_indices_experts_sorted_N: torch.Tensor  # noqa: N815
    topk_scores_experts_sorted_N: torch.Tensor  # noqa: N815


@dataclass(frozen=True, kw_only=True)
class AllToAllDispatchMetadata(LocalDispatchMetadata):
    """Metadata returned by AllToAllTokenDispatcher.dispatch() for use in combine()."""

    input_shape: tuple  # for _unpermute
    permuted_indices: torch.Tensor  # for _unpermute
    input_splits: list[int]
    output_splits: list[int]


class LocalTokenDispatcher(Module):
    """Token dispatcher for EP=1. Handles local token reordering only.

    Dispatchers are parameterless modules so they can own remat region names
    and policy state.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Module.Config):
        num_experts: int
        top_k: int

    def __init__(self, config: Config):
        super().__init__()
        self.num_experts = config.num_experts
        self.top_k = config.top_k

    def init_buffer(self) -> None:
        """Initialize backend communication buffers, if any."""

    def _local_reorder(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """Reorder tokens by expert assignment for local expert computation.

        Groups tokens by expert index via argsort. Routing scores are applied
        to the expert outputs in ``combine``, after expert computation.

        Args:
            x_TD: ``(T, D)`` input tokens
            topk_scores_TK: ``(T, K)`` routing scores
            topk_expert_ids_TK: ``(T, K)`` expert indices

        Returns:
            routed_input_ND: ``(N, D)`` where N = T*K. Tokens in expert-sorted
                order.
            token_indices_experts_sorted_N: ``(N,)`` token-to-original mapping
            topk_scores_experts_sorted_N: ``(N,)`` scores in expert-sorted order
        """
        # Reorder the token indices to match the order of the experts where N = T*K
        token_indices_experts_sorted_N = torch.argsort(
            topk_expert_ids_TK.view(-1), stable=True
        )
        topk_scores_experts_sorted_N = topk_scores_TK.view(-1)[
            token_indices_experts_sorted_N
        ]
        token_indices_experts_sorted_N = token_indices_experts_sorted_N // self.top_k
        routed_input_ND = x_TD[token_indices_experts_sorted_N]

        return (
            routed_input_ND,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
        )

    def dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor, LocalDispatchMetadata]:
        """Reorder tokens by expert assignment for local expert computation.

        Args:
            x_TD: ``(T, D)`` all input tokens
            topk_scores_TK: ``(T, K)`` routing scores
            topk_expert_ids_TK: ``(T, K)`` expert indices per token
            num_local_tokens_per_expert_E: ``(E,)`` token counts per expert

        Returns:
            routed_input_RD: ``[R = sum(num_local_tokens_per_expert_E), input_dim(D)]``.
                Tokens sorted by expert index.
            num_local_tokens_per_expert_E: ``(E,)`` token counts per expert
            metadata: LocalDispatchMetadata for combine()
        """
        if spmd.is_type_checking():
            spmd.mutate_type(
                num_local_tokens_per_expert_E,
                src=spmd.P,
                dst={"dp": spmd.V, "cp": spmd.V, "tp": spmd.V},
            )
        # R = N (no EP all-to-all)
        (
            routed_input_RD,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
        ) = self._local_reorder(x_TD, topk_scores_TK, topk_expert_ids_TK)
        metadata = LocalDispatchMetadata(
            token_indices_experts_sorted_N=token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N=topk_scores_experts_sorted_N,
        )
        return routed_input_RD, num_local_tokens_per_expert_E, metadata

    def combine(
        self,
        routed_output_RD: torch.Tensor,
        metadata: LocalDispatchMetadata,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Score and scatter_add routed expert outputs.

        Args:
            routed_output_RD: ``(R, D)`` expert outputs
            metadata: LocalDispatchMetadata from dispatch()
            x_TD: ``(T, D)`` original input tokens
        Returns:
            out_TD: ``(T, D)`` combined output.
        """
        # The scatter-add reads the expert outputs with bare ops.
        remat.recompute_needs_tensor(routed_output_RD)
        return self._score_and_scatter_add(
            routed_output_RD,
            metadata.topk_scores_experts_sorted_N,
            metadata.token_indices_experts_sorted_N,
            x_TD,
        )

    def _score_and_scatter_add(
        self,
        routed_output_ND: torch.Tensor,
        topk_scores_N: torch.Tensor,
        token_indices_N: torch.Tensor,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Weight expert-sorted outputs by their router scores and sum them per token."""
        # Type promotion computes the product in float32 without materializing a
        # float32 copy of the routed output, which the multiply would save for
        # backward (twice the bytes of the routed output itself).
        routed_output_ND = (
            routed_output_ND * topk_scores_N.reshape(-1, 1).to(torch.float32)
        ).to(routed_output_ND.dtype)
        return deterministic_scatter_add(
            torch.zeros_like(x_TD),
            token_indices_N.reshape(-1, 1).expand(-1, x_TD.shape[-1]),
            routed_output_ND,
        )


class BaseEPTokenDispatcher(LocalTokenDispatcher, ABC):
    """Base class for EP token dispatchers.

    Resolves the EP mesh from the ambient SPMD runtime.
    LocalTokenDispatcher intentionally does not know about SP: local dispatch is
    used when EP is off, and expert activations are replicated for expert TP.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(LocalTokenDispatcher.Config):
        pass

    def __init__(self, config: Config):
        super().__init__(config)

    @property
    def ep_mesh(self) -> DeviceMesh | None:
        """Return the active one-dimensional EP mesh, if EP is enabled."""
        with spmd.no_typecheck():
            # weirdly, accessing ep PG hits DeviceMesh internals that expects type annotations
            mesh = spmd_sparse_mesh()
            return None if mesh is None else mesh["ep"]

    def init_buffer(self) -> None:
        """Initialize backend communication buffers, if any."""

    @abstractmethod
    # pyrefly: ignore [bad-override]
    def dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor, object]:
        """Dispatch physically padded local tokens.

        MoE pads the sequence before routing, so ``x_TD.shape[0]`` is the common
        current token count across ranks. Persistent backends use the lifetime
        ``num_max_tokens_per_rank`` from their config only to preallocate storage.
        """
        raise NotImplementedError

    @abstractmethod
    def combine(
        self,
        routed_output_RD: torch.Tensor,
        metadata: object,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Combine expert outputs."""
        raise NotImplementedError


class AllToAllTokenDispatcher(BaseEPTokenDispatcher):
    """Token dispatcher for EP>1. Handles token reorder + all-to-all dispatch/combine.

    Handles the full token routing lifecycle:
    dispatch (reorder + EP all-to-all) and combine (reverse).

    With EP, each of dispatch and combine is one remat region,
    ``<fqn>.dispatch`` and ``<fqn>.combine``, covering its local reordering,
    collectives, and (for combine) the score-weighted scatter-add.

    The EP mesh is resolved from the ambient SPMD runtime.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(BaseEPTokenDispatcher.Config):
        pass

    def __init__(self, config: Config):
        super().__init__(config)

    def _token_count_exchange(
        self,
        num_local_tokens_per_expert_E: torch.Tensor,
        pg,
        ep_size: int,
    ) -> torch.Tensor:
        """Exchange per-rank expert token counts before the data all-to-all.

        This method is separate from ``dispatch`` so graph passes can annotate
        the count exchange independently from the true token-exchange
        scheduling markers.
        """
        assert self.ep_mesh is not None
        if torch.compiler.is_compiling() or torch.compiler._is_non_strict_tracing():
            return all_to_all_single(
                num_local_tokens_per_expert_E.view(ep_size, -1),
                None,
                None,
                group=self.ep_mesh,
            )

        return spmd.all_to_all(
            num_local_tokens_per_expert_E.view(ep_size, -1),
            pg,
            src=spmd.V,
            dst=spmd.V,
        )

    def _sync_token_count_exchange(
        self,
        num_local_tokens_per_expert_E: torch.Tensor,
        num_global_tokens_per_local_expert_EP_e: torch.Tensor,
        ep_size: int,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """Wait for token counts and copy the per-rank splits to CPU.

        Local input splits can copy to CPU non-blocking; remote output splits
        must be ready before launching the variable-size data all-to-all.
        """
        # Need to wait explicitly because it is used by a triton kernel later
        # which doesn't realize that AsyncCollectiveTensor needs unwrapping
        num_global_tokens_per_local_expert_EP_e = (
            torch.ops._c10d_functional.wait_tensor(
                num_global_tokens_per_local_expert_EP_e
            )
        )
        num_global_tokens_per_local_expert_E = (
            num_global_tokens_per_local_expert_EP_e.reshape(-1)
        )
        input_splits = (
            num_local_tokens_per_expert_E.view(ep_size, -1)
            .sum(dim=1)
            .to(torch.device("cpu"), non_blocking=True)
        )
        # NOTE: this would incur a device-to-host sync
        output_splits = (
            num_global_tokens_per_local_expert_E.view(ep_size, -1)
            .sum(dim=1)
            .to(torch.device("cpu"), non_blocking=False)
        )
        return num_global_tokens_per_local_expert_E, input_splits, output_splits

    def _dispatch_token_exchange(
        self,
        routed_input_ND: torch.Tensor,
        pg,
        output_splits: list[int],
        input_splits: list[int],
    ) -> torch.Tensor:
        """Launch the dispatch all-to-all that moves routed tokens to experts."""
        assert self.ep_mesh is not None
        if torch.compiler.is_compiling() or torch.compiler._is_non_strict_tracing():
            return all_to_all_single(
                routed_input_ND,
                output_splits,
                input_splits,
                self.ep_mesh,
            )

        return spmd.all_to_all(
            routed_input_ND,
            pg,
            src=spmd.V,
            dst=spmd.V,
            output_split_sizes=output_splits,
            input_split_sizes=input_splits,
        )

    def _combine_token_exchange(
        self,
        routed_output_RD: torch.Tensor,
        pg,
        input_splits: list[int],
        output_splits: list[int],
    ) -> torch.Tensor:
        """Launch the combine all-to-all that returns expert outputs to tokens."""
        assert self.ep_mesh is not None
        if torch.compiler.is_compiling() or torch.compiler._is_non_strict_tracing():
            return all_to_all_single(
                routed_output_RD,
                input_splits,
                output_splits,
                self.ep_mesh,
            )

        return spmd.all_to_all(
            routed_output_RD,
            pg,
            src=spmd.V,
            dst=spmd.V,
            output_split_sizes=input_splits,
            input_split_sizes=output_splits,
        )

    def dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[
        torch.Tensor, torch.Tensor, AllToAllDispatchMetadata | LocalDispatchMetadata
    ]:
        """Reorder tokens, then all-to-all dispatch to expert-parallel ranks.

        When ep_mesh is None (EP=1), falls back to local dispatch — no
        all-to-all communication, just local token reordering with padding.

        With SP, x_TD/topk_scores_TK/topk_expert_ids_TK are already the local
        SP shard provided by the surrounding local SPMD region.

        Args:
            x_TD: ``(T, D)`` local token shard
            topk_scores_TK: ``(T, K)`` routing scores
            topk_expert_ids_TK: ``(T, K)`` expert indices
            num_local_tokens_per_expert_E: ``(E,)`` token counts for this local
                token shard

        Returns:
            routed_input_RD: ``[R = sum(num_tokens_per_local_expert_e), input_dim(D)]``.
                Tokens in expert-major order for local experts.
            num_tokens_per_local_expert_e: ``(num_local_experts,)`` token counts
            metadata: dispatch metadata for combine()
        """
        # EP=1: fall back to local dispatch (no all-to-all needed)
        if self.ep_mesh is None:
            return LocalTokenDispatcher.dispatch(
                self,
                x_TD,
                topk_scores_TK,
                topk_expert_ids_TK,
                num_local_tokens_per_expert_E,
            )

        if spmd.is_type_checking():  # sparse mesh reinterpret
            spmd.mutate_type(
                num_local_tokens_per_expert_E,
                src=spmd.P,
                dst={"dp": spmd.V, "cp": spmd.V, "tp": spmd.V},
            )
        (
            routed_input_RD,
            num_global_tokens_per_local_expert_e,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
            permuted_indices,
            input_splits,
            output_splits,
        ) = remat.region(
            self._dispatch,
            self.remat_region_name("dispatch"),
            recompute=self.remat_should_recompute("dispatch"),
        )(
            x_TD,
            topk_scores_TK,
            topk_expert_ids_TK,
            num_local_tokens_per_expert_E,
        )
        # MoE.forward reads the counts with bare ops, and the metadata below
        # reads the splits.
        remat.recompute_needs_tensor(
            num_global_tokens_per_local_expert_e,
            input_splits,
            output_splits,
        )
        # The splits are CPU tensors, so tolist() does not sync.
        input_splits_list = input_splits.tolist()
        output_splits_list = output_splits.tolist()
        metadata = AllToAllDispatchMetadata(
            token_indices_experts_sorted_N=token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N=topk_scores_experts_sorted_N,
            # Pre-permute shape: the all-to-all output, which padding
            # dispatchers (TorchAO) may grow in _permute.
            input_shape=torch.Size((sum(output_splits_list), x_TD.shape[-1])),
            permuted_indices=permuted_indices,
            input_splits=input_splits_list,
            output_splits=output_splits_list,
        )
        return routed_input_RD, num_global_tokens_per_local_expert_e, metadata

    def _dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[torch.Tensor, ...]:
        """Sort, exchange counts and tokens, and permute to expert-major order.

        Runs as the single ``<fqn>.dispatch`` remat region, so a saved dispatch
        replays neither the collectives nor the device-to-host split sync, and
        a recomputed one replays all of them. Regions return only tensors, so
        the splits come back as CPU tensors.
        """
        assert self.ep_mesh is not None
        ep_size = self.ep_mesh.size()
        # (N, D) where N = T*K; the all-to-all below produces (R, D).
        (
            routed_input_ND,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
        ) = self._local_reorder(x_TD, topk_scores_TK, topk_expert_ids_TK)

        with maybe_set_sparse_mesh():
            pg = "ep"
            if spmd.is_type_checking():
                num_local_tokens_per_expert_E = spmd.reinterpret_mesh(
                    num_local_tokens_per_expert_E, spmd.current_mesh()
                )
                routed_input_ND = spmd.reinterpret_mesh(
                    routed_input_ND, spmd.current_mesh()
                )

            # generate the input splits and output splits for all-to-all
            with torch.no_grad():
                num_global_tokens_per_local_expert_EP_e = self._token_count_exchange(
                    num_local_tokens_per_expert_E, pg, ep_size
                )
                (
                    num_global_tokens_per_local_expert_E,
                    input_splits,
                    output_splits,
                ) = self._sync_token_count_exchange(
                    num_local_tokens_per_expert_E,
                    num_global_tokens_per_local_expert_EP_e,
                    ep_size,
                )

            routed_input_RD = self._dispatch_token_exchange(
                routed_input_ND,
                pg,
                output_splits.tolist(),
                input_splits.tolist(),
            )
            # Reorder from rank-major to expert-major via _permute.
            #
            # num_global_tokens_per_local_expert_E layout after all-to-all
            # (e = local experts, EP = EP ranks):
            #   (e0,r0), (e1,r0), ..., (e0,r1), (e1,r1), ...  (rank-major)
            # _permute reshuffles to:
            #   (e0,r0), (e0,r1), ..., (e1,r0), (e1,r1), ...  (expert-major)
            # TODO: Consider using num_global_tokens_per_local_expert_e as the
            # expert_bias_e update buffer, then all-gather on EP ranks. This
            # is blocked by clarification on HybridEP token dropping.
            (
                routed_input_RD,
                permuted_indices,
                num_global_tokens_per_local_expert_e,
            ) = self._permute(routed_input_RD, num_global_tokens_per_local_expert_E)

        return (
            routed_input_RD,
            num_global_tokens_per_local_expert_e,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
            permuted_indices,
            input_splits,
            output_splits,
        )

    def _permute(
        self,
        routed_input_RD,
        num_global_tokens_per_local_expert_E,
    ):
        """Reorder tokens from rank-major to expert-major layout.

        Input layout:  (e0,r0), (e1,r0), ..., (e0,r1), (e1,r1), ...  (rank-major)
        Output layout: (e0,r0), (e0,r1), ..., (e1,r0), (e1,r1), ...  (expert-major)

        Collapses token count matrix ``t_mat`` from ``(EP, e)`` to
        ``num_global_tokens_per_local_expert_e`` ``(e,)`` by summing across ranks.
        """
        # pyrefly: ignore [missing-attribute]
        ep_size = self.ep_mesh.size()
        e = num_global_tokens_per_local_expert_E.shape[0] // ep_size
        device = num_global_tokens_per_local_expert_E.device
        total = routed_input_RD.shape[0]

        # (EP, e) matrix of token counts per (rank, local_expert)
        t_mat = num_global_tokens_per_local_expert_E.view(ep_size, e)

        # Where each (r, e) segment starts in the input (rank-major order)
        input_starts = (
            num_global_tokens_per_local_expert_E.cumsum(0)
            - num_global_tokens_per_local_expert_E
        ).view(ep_size, e)

        # Transpose to expert-major (e, EP) and flatten
        segment_lens = t_mat.t().reshape(-1)
        input_starts = input_starts.t().reshape(-1)

        # For each output position, find its input position:
        #   output[p] = input[input_starts[seg] + (p - output_starts[seg])]
        seg_ids = torch.arange(segment_lens.shape[0], device=device).repeat_interleave(
            segment_lens, output_size=total
        )
        output_starts = segment_lens.cumsum(0) - segment_lens
        # seg_ids.shape[0] == segment_lens.sum(), but reuses the unbacked symint
        # already created by repeat_interleave above.
        permuted_indices = (
            input_starts[seg_ids]
            + torch.arange(seg_ids.shape[0], device=device)
            - output_starts[seg_ids]
        )

        num_global_tokens_per_local_expert_e = t_mat.sum(0)
        return (
            routed_input_RD[permuted_indices, :],
            permuted_indices,
            num_global_tokens_per_local_expert_e,
        )

    def _unpermute(self, routed_output_RD, input_shape, permuted_indices):
        """Reverse expert-major reordering."""
        out_unpermuted_RD = routed_output_RD.new_empty(input_shape)
        out_unpermuted_RD[permuted_indices, :] = routed_output_RD
        return out_unpermuted_RD

    # pyrefly: ignore [bad-override]
    def combine(
        self,
        routed_output_RD: torch.Tensor,
        metadata: AllToAllDispatchMetadata,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Reverse the dispatch: unpermute + all-to-all + score + scatter_add.

        Args:
            routed_output_RD: ``(R, D)`` expert outputs in expert-major order
            metadata: AllToAllDispatchMetadata from dispatch()
            x_TD: ``(T, D)`` original input tokens

        Returns:
            out_TD: Combined local output ``(T, D)``.
        """
        # EP=1: fall back to local combine (no all-to-all needed)
        if self.ep_mesh is None:
            return LocalTokenDispatcher.combine(
                self,
                routed_output_RD,
                metadata,
                x_TD,
            )

        out_TD = remat.region(
            self._combine,
            self.remat_region_name("combine"),
            recompute=self.remat_should_recompute("combine"),
        )(
            routed_output_RD,
            metadata.topk_scores_experts_sorted_N,
            metadata.token_indices_experts_sorted_N,
            metadata.permuted_indices,
            metadata.input_shape,
            metadata.input_splits,
            metadata.output_splits,
            x_TD,
        )
        return out_TD

    def _combine(
        self,
        routed_output_RD: torch.Tensor,
        topk_scores_experts_sorted_N: torch.Tensor,
        token_indices_experts_sorted_N: torch.Tensor,
        permuted_indices: torch.Tensor,
        input_shape: torch.Size,
        input_splits: list[int],
        output_splits: list[int],
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Unpermute, all-to-all back to token ranks, then score and scatter-add.

        Runs as the single ``<fqn>.combine`` remat region.
        """
        with maybe_set_sparse_mesh():
            routed_output_RD = self._unpermute(
                routed_output_RD, input_shape, permuted_indices
            )
            routed_output_RD = self._combine_token_exchange(
                routed_output_RD, "ep", input_splits, output_splits
            )
        if spmd.is_type_checking():  # dense mesh reinterpret
            routed_output_RD = spmd.reinterpret_mesh(
                routed_output_RD, spmd.current_mesh()
            )
        return self._score_and_scatter_add(
            routed_output_RD,
            topk_scores_experts_sorted_N,
            token_indices_experts_sorted_N,
            x_TD,
        )


class TorchAOTokenDispatcher(AllToAllTokenDispatcher):
    """Token dispatcher with token group padding for quantized grouped GEMMs.

    Uses torchao's ``permute_and_pad`` instead of the standard ``_permute`` to
    reorder tokens into expert-major order and pad each expert's token group to
    a multiple of ``pad_multiple``. This alignment is required by quantized
    grouped GEMM kernels (e.g. 32 for MXFP8).

    Works with EP enabled (all-to-all dispatch + padded permute) and with
    EP=1 (``ep_mesh is None``), where it skips the all-to-all and only applies
    the local padded permute. The padding is what the quantized grouped GEMM
    needs; the all-to-all is orthogonal, so EP=1 is supported for single-GPU
    debugging / numerics by running ``permute_and_pad`` with ``ep_degree=1``.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(AllToAllTokenDispatcher.Config):
        pad_multiple: int

    def __init__(self, config: Config):
        super().__init__(config)
        self.pad_multiple = config.pad_multiple

    def dispatch(
        self,
        x_TD,
        topk_scores_TK,
        topk_expert_ids_TK,
        num_local_tokens_per_expert_E,
    ):
        if self.ep_mesh is not None:
            return super().dispatch(
                x_TD,
                topk_scores_TK,
                topk_expert_ids_TK,
                num_local_tokens_per_expert_E,
            )

        # EP=1: no all-to-all. Locally reorder tokens to expert-sorted order,
        # then apply the padded permute so the quantized grouped GEMM sees
        # token groups aligned to pad_multiple. _permute reads ep_size=1 when
        # ep_mesh is None, so num_local_tokens_per_expert_E is already the
        # full per-expert count.
        (
            routed_input_ND,
            token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N,
        ) = self._local_reorder(x_TD, topk_scores_TK, topk_expert_ids_TK)

        input_shape = routed_input_ND.shape
        (
            routed_input_RD,
            permuted_indices,
            num_tokens_per_local_expert_padded_e,
        ) = self._permute(routed_input_ND, num_local_tokens_per_expert_E)

        metadata = AllToAllDispatchMetadata(
            token_indices_experts_sorted_N=token_indices_experts_sorted_N,
            topk_scores_experts_sorted_N=topk_scores_experts_sorted_N,
            input_shape=input_shape,
            permuted_indices=permuted_indices,
            # Unused in the EP=1 combine path (no all-to-all to reverse).
            input_splits=[],
            output_splits=[],
        )
        return routed_input_RD, num_tokens_per_local_expert_padded_e, metadata

    def combine(
        self,
        routed_output_RD,
        metadata,
        x_TD,
    ):
        if self.ep_mesh is not None:
            return super().combine(
                routed_output_RD,
                metadata,
                x_TD,
            )

        # EP=1: strip the padding (via _unpermute) to recover expert-sorted
        # order, then apply the local score + scatter_add used by the EP=1
        # path. Mirrors LocalTokenDispatcher.combine, plus the unpad.
        assert isinstance(metadata, AllToAllDispatchMetadata)
        # The unpermute reads the expert outputs with bare ops.
        remat.recompute_needs_tensor(routed_output_RD)
        routed_output_RD = self._unpermute(
            routed_output_RD, metadata.input_shape, metadata.permuted_indices
        )

        return self._score_and_scatter_add(
            routed_output_RD,
            metadata.topk_scores_experts_sorted_N,
            metadata.token_indices_experts_sorted_N,
            x_TD,
        )

    def _permute(
        self,
        routed_input_RD,
        num_global_tokens_per_local_expert_E,
    ):
        # MXFP8 requires groups to be permuted to expert major order AND
        # padded to nearest multiple of 16.
        # It also does padding to make sure the number of tokens each expert
        # gets locally is a multiple of `self.pad_multiple`.
        # Note that this will create side effects when wrapping the for-loop
        # implementation of routed experts, as it does not need padding.
        from torchao.prototype.moe_training.ep.permute import permute_and_pad

        # ep_size=1 when EP is disabled: permute_and_pad then only pads token
        # groups (rank-major == expert-major for a single rank).
        ep_size = self.ep_mesh.size() if self.ep_mesh is not None else 1
        e = num_global_tokens_per_local_expert_E.shape[0] // ep_size

        (
            _padded_input_shape,
            routed_input_RD,
            permuted_indices,
            num_global_tokens_per_local_expert_padded_e,
            _group_offsets,
        ) = permute_and_pad(
            routed_input_RD,
            num_global_tokens_per_local_expert_E,
            ep_size,
            e,
            self.pad_multiple,
        )
        return (
            routed_input_RD,
            permuted_indices,
            num_global_tokens_per_local_expert_padded_e,
        )

    def _unpermute(self, routed_output_RD, input_shape, permuted_indices):
        # permute_and_pad appends a zero sentinel row that padding rows index;
        # scatter into it too, then strip it.
        num_rows, *feature_shape = input_shape
        out_unpermuted_RD = routed_output_RD.new_empty((num_rows + 1, *feature_shape))
        out_unpermuted_RD[permuted_indices, :] = routed_output_RD
        return out_unpermuted_RD[:-1]


@dataclass(frozen=True, kw_only=True)
class EPDispatchMetadata:
    """Metadata for DeepEP and HybridEP token dispatch."""

    state: object  # Backend-specific dispatch state.


class DeepEPTokenDispatcher(BaseEPTokenDispatcher):
    """Token dispatcher using DeepEP v2's unified ``ElasticBuffer`` dispatch/combine.

    DeepEP v2 (>= 2.0.0) collapses the v1 high-throughput (HT) and low-latency (LL)
    paths into a single ``buffer.dispatch``/``combine``. Compact dispatch is gathered
    from its deduplicated output into expert-major order; expand dispatch already returns
    the static expert-major layout. Combine is synchronized before returning its result.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(BaseEPTokenDispatcher.Config):
        # Select the dispatch layout. False (default, also forced under autograd): compact,
        # host-synced, backward-able path for training. True: static, no-host-sync expand
        # layout so the MoE forward is CUDA-graph-capturable -- inference only (covers BOTH
        # prefill and decode, since both run under no_grad), no backward. The deepep
        # primitives gate on grad context, so a True spec falls back to compact in training.
        cuda_graph_compatible: bool = False
        # Hard per-rank input-token bound used to preallocate the communication buffer.
        # Runtime configuration must fill it before dispatcher construction.
        num_max_tokens_per_rank: int | None = None
        # Model hidden dim, threaded by the builder for eager buffer initialization.
        hidden_dim: int | None = None

    def __init__(self, config: Config):
        super().__init__(config)
        if config.num_max_tokens_per_rank is None:
            raise ValueError(
                "DeepEP requires num_max_tokens_per_rank for buffer initialization."
            )
        if config.num_max_tokens_per_rank <= 0:
            raise ValueError(
                "DeepEP num_max_tokens_per_rank must be positive, got "
                f"{config.num_max_tokens_per_rank}."
            )
        self.num_max_tokens_per_rank = config.num_max_tokens_per_rank
        self.hidden_dim = config.hidden_dim
        self.cuda_graph_compatible = config.cuda_graph_compatible

        # Import to register custom ops so SAC saves communication outputs
        # instead of recomputing them. This must happen before apply_ac.
        from torchtitan.distributed.deepep import deepep  # noqa: F401

    def init_buffer(self) -> None:
        """Eagerly create the DeepEP buffer."""
        assert self.ep_mesh is not None
        assert self.hidden_dim is not None

        from torchtitan.distributed.deepep.deepep import get_buffer

        get_buffer(
            self.ep_mesh.get_group(),
            hidden=self.hidden_dim,
            num_max_tokens_per_rank=self.num_max_tokens_per_rank,
            num_topk=self.top_k,
        )

    def dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor, EPDispatchMetadata]:
        """Dispatch through the preallocated ElasticBuffer."""
        # Ignore input num_local_tokens_per_expert_E. DeepEP returns the number
        # of global routed tokens for every local expert using other inputs.
        del num_local_tokens_per_expert_E
        assert self.ep_mesh is not None, (
            "ep_mesh must be set before dispatch. "
            "ExpertParallel._partition_fn() should set it."
        )
        ep_group = self.ep_mesh.get_group()
        num_local_experts = self.num_experts // ep_group.size()

        from torchtitan.distributed.deepep.deepep import dispatch_tokens

        hidden_states_RD, num_global_tokens_per_local_expert_e, state = dispatch_tokens(
            x_TD,
            topk_expert_ids_TK,
            topk_scores_TK,
            num_local_experts,
            self.num_experts,
            num_tokens_per_rank=x_TD.shape[0],
            remat_region_name=self.remat_region_name("ep_communication.dispatch"),
            cuda_graph_compatible=self.cuda_graph_compatible,
        )

        metadata = EPDispatchMetadata(state=state)
        return hidden_states_RD, num_global_tokens_per_local_expert_e, metadata

    # pyrefly: ignore [bad-override]
    def combine(
        self,
        routed_output_RD: torch.Tensor,
        metadata: EPDispatchMetadata,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Combine tokens via DeepEP and wait for completion."""
        del x_TD
        from torchtitan.distributed.deepep.deepep import combine_tokens, sync_combine

        # combine_tokens applies routing scores with bare ops.
        remat.recompute_needs_tensor(routed_output_RD)

        combined_TD = combine_tokens(
            routed_output_RD,
            metadata.state,  # pyrefly: ignore [bad-argument-type]
            remat_region_name=self.remat_region_name("ep_communication.combine"),
        )
        sync_combine()
        return combined_TD


class HybridEPTokenDispatcher(BaseEPTokenDispatcher):
    """Token dispatcher using HybridEP for efficient token dispatch/combine.

    Uses HybridEP library kernels (GB200/NVLink72) instead of standard
    all-to-all collectives.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(BaseEPTokenDispatcher.Config):
        """Config for HybridEP token dispatcher.

        Args:
            non_blocking_capacity_factor: Enable non-blocking HybridEP dispatch
                with a given capacity factor.

                Setting this to a float in (0, 1] enables CPU-free non-blocking
                dispatch and controls num_permuted_tokens — the fused-permute
                output capacity, estimated as:
                num_max_tokens_per_rank * ep_size *
                min(num_local_experts, top_k) * capacity_factor, aligned for
                MXFP8. Tokens whose permuted offset exceeds this limit are
                silently dropped (overflow_flag is set on GPU).

                - None = blocking mode (default).  HybridEP calls
                  cudaStreamSynchronize after dispatch, copies
                  tokens_per_expert to pinned CPU memory, and computes the
                  exact num_permuted_tokens on the host.  No token dropping.
                - 1.0 = non-blocking, worst-case sizing: every token can reach
                  every local expert, no drops, highest memory.
                - < 1.0 = non-blocking, reduced memory; controls the
                  fused-permute output tensor size (num_permuted_tokens).
                  Safe in practice when forced load balancing (e.g. aux-loss /
                  round-robin) keeps distribution roughly uniform.

                This factor does not affect the all-to-all communication
                buffer, which is initialized separately from
                ``num_max_tokens_per_rank``.
        """

        non_blocking_capacity_factor: float | None = None
        pad_multiple: int | None = None
        hidden_dim: int | None = None
        num_max_tokens_per_rank: int | None = None

    def __init__(self, config: Config):
        super().__init__(config)
        self.non_blocking_capacity_factor = config.non_blocking_capacity_factor
        self.pad_multiple = config.pad_multiple
        self.hidden_dim = config.hidden_dim
        self.num_max_tokens_per_rank = config.num_max_tokens_per_rank

        # Import to register custom ops so SAC saves communication outputs
        # instead of recomputing them. This must happen before apply_ac.
        from torchtitan.distributed.deepep import hybridep  # noqa: F401

    def init_buffer(self) -> None:
        """Eagerly create the HybridEP buffer."""
        assert self.ep_mesh is not None
        assert self.hidden_dim is not None

        if self.num_max_tokens_per_rank is None:
            raise ValueError(
                "HybridEP requires num_max_tokens_per_rank for buffer initialization."
            )

        from torchtitan.distributed.deepep.hybridep import get_buffer

        get_buffer(
            group=self.ep_mesh.get_group(),
            hidden_dim=self.hidden_dim,
            num_max_tokens_per_rank=self.num_max_tokens_per_rank,
            num_local_experts=self.num_experts // self.ep_mesh.size(),
        )

    def dispatch(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor, EPDispatchMetadata]:
        """Dispatch through the preallocated HybridEP buffer."""
        # Ignore input num_local_tokens_per_expert_E. HybridEP returns the
        # number of global routed tokens for every local expert using other inputs.
        del num_local_tokens_per_expert_E
        assert self.ep_mesh is not None, (
            "ep_mesh must be set before dispatch. "
            "ExpertParallel._partition_fn() should set it."
        )
        ep_group = self.ep_mesh.get_group()
        num_local_experts = self.num_experts // ep_group.size()

        from torchtitan.distributed.deepep.hybridep import dispatch_tokens

        hidden_states_RD, num_global_tokens_per_local_expert_e, state = dispatch_tokens(
            x_TD,
            topk_expert_ids_TK,
            topk_scores_TK,
            num_local_experts,
            self.num_experts,
            ep_group,
            non_blocking_expert_capacity_factor=self.non_blocking_capacity_factor,
            pad_multiple=self.pad_multiple,
        )

        metadata = EPDispatchMetadata(state=state)
        return hidden_states_RD, num_global_tokens_per_local_expert_e, metadata

    # pyrefly: ignore [bad-override]
    def combine(
        self,
        routed_output_RD: torch.Tensor,
        metadata: EPDispatchMetadata,
        x_TD: torch.Tensor,
    ) -> torch.Tensor:
        """Combine tokens via HybridEP."""
        del x_TD

        from torchtitan.distributed.deepep import hybridep

        # combine_tokens applies routing scores and combines with bare ops.
        remat.recompute_needs_tensor(routed_output_RD)
        combined_TD = hybridep.combine_tokens(
            routed_output_RD,
            metadata.state,  # pyrefly: ignore [bad-argument-type]
            pad_multiple=self.pad_multiple,
        )
        return combined_TD
