# 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 __future__ import annotations

from collections.abc import Iterator
from dataclasses import dataclass, field
from typing import cast, Protocol

import spmd_types as spmd

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch_remat as remat
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl
from torch.optim import Optimizer

from torchtitan.distributed import ParallelismContext
from torchtitan.distributed.parallelism_context import MeshAxisName
from torchtitan.distributed.spmd_types import (
    maybe_set_sparse_mesh,
    spmd_dense_sp_enabled,
    spmd_local_context,
    spmd_mesh_group,
    spmd_mesh_size,
    spmd_sparse_mesh,
)
from torchtitan.models.common.activation import (
    BinaryActivationFn,
    Sigmoid,
    SwiGLU,
    UnaryActivationFn,
)
from torchtitan.models.common.aux_loss import AuxLoss
from torchtitan.models.common.feed_forward import FeedForward
from torchtitan.models.common.hi_mid_lo_linear import HiMidLoLinear
from torchtitan.models.common.linear import GroupedLinear
from torchtitan.protocols.module import Module

from .token_dispatcher import LocalTokenDispatcher

# Shape suffix legend
# (https://medium.com/@NoamShazeer/shape-suffixes-good-coding-style-f836e72e24fd):
#   T = num tokens, D = model dimension,
#   F = hidden (FFN intermediate) dimension, E = num experts,
#   B = histogram bins,
#   e = num local experts (E / EP, used in token dispatcher for
#       per-local-expert token counts after EP dispatch /_permute),
#   K = top-k, N = routed tokens (T*K),
#   R = routed tokens assigned to local experts,
#   O = expert output features, I = expert input features
#       (roles, not model dims: the _grouped_mm seam takes the expert
#        weight in its stored (E, O, I) orientation, which is (E, F, D)
#        for the up/gate projections and (E, D, F) for the down one)


class RoutedExperts(Module):
    """Local SPMD region with first-class grouped expert projections."""

    @dataclass(kw_only=True, slots=True)
    class Config(Module.Config):
        w13: GroupedLinear.Config
        w2: GroupedLinear.Config
        token_dispatcher: LocalTokenDispatcher.Config
        activation_fn: BinaryActivationFn.Config = field(default_factory=SwiGLU.Config)
        output_postprocess: Module.Config | None = None

        def __post_init__(self) -> None:
            if self.w13.group_size != self.w2.group_size:
                raise ValueError("w13 and w2 must contain the same number of experts")
            if self.token_dispatcher.num_experts != self.w13.group_size:
                raise ValueError(
                    "token dispatcher and grouped linears must contain the same "
                    "number of experts"
                )
            if self.w13.in_features != self.w2.out_features:
                raise ValueError("w13 input and w2 output dimensions must match")
            if self.w13.num_linears != 2:
                raise ValueError("w13 output must contain gate and up projections")
            if self.w13.out_features != self.w2.in_features:
                raise ValueError("w13 output and w2 input dimensions must match")
            if self.w2.num_linears != 1:
                raise ValueError("w2 must contain one down projection")

    def __init__(self, config: Config):
        super().__init__()
        self.w13 = config.w13.build()
        self.w2 = config.w2.build()
        self.activation_fn = config.activation_fn.build()
        self.output_postprocess = (
            config.output_postprocess.build()
            if config.output_postprocess is not None
            else None
        )
        self.token_dispatcher = config.token_dispatcher.build()

    def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None:
        del buffer_device
        self.token_dispatcher.init_buffer()

    def forward(
        self,
        x_TD: torch.Tensor,
        topk_scores_TK: torch.Tensor,
        topk_expert_ids_TK: torch.Tensor,
        num_local_tokens_per_expert_E: torch.Tensor,
    ) -> torch.Tensor:
        """Dispatch tokens to experts, compute, combine, and scatter_add.

        When parallelized, ``local_spmd`` (from ``sharding_config``) establishes
        the local SPMD types for the forward body.
        """
        (
            routed_input_RD,
            num_global_tokens_per_local_expert_e,
            metadata,
        ) = self.token_dispatcher.dispatch(
            x_TD,
            topk_scores_TK,
            topk_expert_ids_TK,
            num_local_tokens_per_expert_E,
        )
        offsets_E = torch.cumsum(
            num_global_tokens_per_local_expert_e,
            dim=0,
            dtype=torch.int32,
        )

        with maybe_set_sparse_mesh():
            # w13 and w2 declare their own remat regions (<fqn>.grouped_mm).
            # The bf16 cast reads the dispatched tokens with bare ops.
            remat.recompute_needs_tensor(routed_input_RD)
            gate_up_R2F = self.w13(routed_input_RD.bfloat16(), offsets_E)
            remat.recompute_needs_tensor(gate_up_R2F)
            hidden_RF = self.activation_fn(gate_up_R2F, offsets=offsets_E)
            routed_output_RD = self.w2(hidden_RF, offsets_E)
            # A real dtype cast and the output postprocess read the w2 output with
            # bare ops, so pin it only then. In the common bf16 case without a
            # postprocess, type_as is a no-op and the w2 output goes straight to
            # the combine region, so an unconditional pin would keep it alive
            # even when w2 and the combine are both saved.
            if (
                routed_output_RD.dtype != routed_input_RD.dtype
                or self.output_postprocess is not None
            ):
                remat.recompute_needs_tensor(routed_output_RD)
            routed_output_RD = routed_output_RD.type_as(routed_input_RD)
            if self.output_postprocess is not None:
                routed_output_RD = self.output_postprocess(routed_output_RD)
        out_TD = self.token_dispatcher.combine(
            routed_output_RD,
            metadata,
            x_TD,
        )
        return out_TD


class TokenChoiceTopKRouter(Module):
    """This class implements token-choice routing. In token-choice top-K routing, each token is
    routed to top K experts based on the router scores.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Module.Config):
        num_experts: int
        gate: HiMidLoLinear.Config
        score_func: UnaryActivationFn.Config
        top_k: int = 1
        route_norm: bool = False
        route_norm_epsilon: float = 1e-20
        route_scale: float = 1.0
        aux_loss: AuxLoss.Config | None = None

    def __init__(self, config: Config):
        super().__init__()
        self.gate = config.gate.build()
        self.num_experts = config.num_experts
        self.top_k = config.top_k
        self.score_func = config.score_func.build()
        self.route_norm = config.route_norm
        self.route_norm_epsilon = config.route_norm_epsilon
        self.route_scale = config.route_scale
        self.aux_loss = config.aux_loss.build() if config.aux_loss is not None else None
        # tokens_per_expert_E will be used to track expert usage and to update the expert bias for load balancing
        self.register_buffer(
            "tokens_per_expert_E",
            torch.zeros(config.num_experts, dtype=torch.float32),
            persistent=False,
        )

    def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None:
        if buffer_device is None:
            # After ``to_empty()``, the existing buffer records the target device.
            buffer_device = self.tokens_per_expert_E.device
        self.tokens_per_expert_E = torch.zeros(
            self.num_experts,
            dtype=torch.float32,
            device=buffer_device,
        )

    def _select_experts(
        self,
        scores_TE: torch.Tensor,
        expert_bias_E: torch.Tensor | None = None,
        **router_kwargs,
    ) -> torch.Tensor:
        scores_for_choice_TE = (
            scores_TE if expert_bias_E is None else scores_TE + expert_bias_E
        )
        return torch.topk(
            scores_for_choice_TE, k=self.top_k, dim=-1, sorted=False
        ).indices

    def forward(
        self,
        x_TD: torch.Tensor,
        expert_bias_E: torch.Tensor | None = None,
        *,
        padding_mask_T: torch.Tensor | None = None,
        aux_loss_denominator: torch.Tensor | None = None,
        **router_kwargs,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Args:
            x_TD: Input ``(T, D)``.
            expert_bias_E: Optional load-balancing bias ``(E,)``.
            padding_mask_T: Boolean ``(T,)`` mask that is true for padding.

        Returns:
            topk_scores_TK: Routing scores ``(T, K)``.
            topk_expert_ids_TK: Expert indices ``(T, K)``.
            routing_map_TE: One-hot boolean routing map ``(T, E)``.
        """
        # HiMidLoLinear returns FP32, so configured scoring runs in FP32.
        gate_TE = self.gate(x_TD)
        # The scoring function reads the router gate projection output with bare ops.
        remat.recompute_needs_tensor(gate_TE)
        scores_TE = self.score_func(gate_TE)

        if padding_mask_T is not None:
            if padding_mask_T.dtype != torch.bool:
                raise ValueError(
                    "padding_mask_T must have dtype bool, "
                    f"got {padding_mask_T.dtype}."
                )
            if padding_mask_T.shape != scores_TE.shape[:-1]:
                raise ValueError(
                    "padding_mask_T must have shape matching the routing-map "
                    f"token axis, got {tuple(padding_mask_T.shape)} for scores "
                    f"{tuple(scores_TE.shape)}."
                )

        topk_expert_ids_TK = remat.region(
            self._select_experts,
            "routing_decision",
            recompute=False,
        )(
            scores_TE,
            expert_bias_E,
            padding_mask_T=padding_mask_T,
            **router_kwargs,
        )
        remat.recompute_needs_tensor(topk_expert_ids_TK)
        # The expert bias is only used for routing. The gating value is
        # still derived from the original scores.
        topk_scores_TK = scores_TE.gather(dim=-1, index=topk_expert_ids_TK)

        if self.route_norm:
            denominator_T1 = (
                topk_scores_TK.sum(dim=-1, keepdim=True) + self.route_norm_epsilon
            )
            topk_scores_TK = topk_scores_TK / denominator_T1
        topk_scores_TK = topk_scores_TK * self.route_scale

        # Build a one-hot boolean routing map (T, E) marking the experts each
        # token is routed to. Under TP/SP the router outputs are sharded on the
        # token dimension; scatter_ writes along the replicated expert
        # dimension and therefore needs no redistribution.
        routing_map_TE = torch.zeros_like(scores_TE, dtype=torch.bool).scatter_(
            -1,
            topk_expert_ids_TK,
            True,
        )
        # Keep the full routing map for dispatch, and build the masked view once
        # for all load-balancing statistics. The auxiliary-loss gradient is
        # injected into topk_scores_TK on backward; see ``AuxLoss.inject``.
        if self.training:
            masked_routing_map_TE = (
                routing_map_TE
                if padding_mask_T is None
                else routing_map_TE & ~padding_mask_T.unsqueeze(-1)
            )
            if not remat.is_recomputing():
                with torch.no_grad():
                    self.tokens_per_expert_E.add_(masked_routing_map_TE.sum(dim=0))
            if self.aux_loss is not None:
                if aux_loss_denominator is None:
                    raise ValueError("An auxiliary-loss denominator is required.")
                topk_scores_TK = self.aux_loss(
                    scores_TE,
                    masked_routing_map_TE,
                    carrier=topk_scores_TK,
                    padding_mask_T=padding_mask_T,
                    denominator=aux_loss_denominator,
                )
        return (
            topk_scores_TK,
            topk_expert_ids_TK,
            routing_map_TE,
        )


class RoundRobinTokenChoiceTopKRouter(TokenChoiceTopKRouter):
    """Route token-expert assignments round-robin with exact balance."""

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

    def _select_experts(
        self,
        scores_TE: torch.Tensor,
        expert_bias_E: torch.Tensor | None = None,
        **router_kwargs,
    ) -> torch.Tensor:
        del expert_bias_E, router_kwargs
        num_tokens = scores_TE.shape[0]
        return (
            torch.arange(
                num_tokens * self.top_k,
                device=scores_TE.device,
                dtype=torch.int64,
            ).reshape(num_tokens, self.top_k)
            % self.num_experts
        )


class QuantileBalancedTopKRouter(TokenChoiceTopKRouter):
    """Top-k router that uses a biased Top-(k+1) cutoff during training."""

    @dataclass(kw_only=True, slots=True)
    class Config(TokenChoiceTopKRouter.Config):
        num_bins: int

    def __init__(self, config: Config):
        super().__init__(config)
        if not isinstance(self.score_func, Sigmoid):
            raise ValueError("Quantile balancing requires sigmoid router scores.")
        self.quantile_balancer = QuantileBalancer.Config(
            num_experts=self.num_experts,
            top_k=self.top_k,
            num_bins=config.num_bins,
        ).build()

    def _select_experts(
        self,
        scores_TE: torch.Tensor,
        expert_bias_E: torch.Tensor | None = None,
        padding_mask_T: torch.Tensor | None = None,
        **router_kwargs,
    ) -> torch.Tensor:
        if expert_bias_E is None:
            raise ValueError("Quantile balancing requires an expert bias.")
        biased_scores_TE = scores_TE + expert_bias_E
        topk_plus_one_scores, topk_plus_one_expert_ids = torch.topk(
            biased_scores_TE,
            k=self.top_k + 1,
            dim=-1,
            sorted=True,
        )
        if self.training and not remat.is_recomputing():
            self.quantile_balancer.observe(
                scores_TE,
                topk_plus_one_scores[:, self.top_k :],
                expert_bias_E,
                padding_mask_T,
            )
        return topk_plus_one_expert_ids[:, : self.top_k].contiguous()


class QuantileBalancer(Module):
    """Accumulate and recover histogram-based quantile bias updates.

    For scores bounded between zero and one, required expert biases lie between
    the current minimum bias minus one and maximum bias plus one. Each training
    micro-batch is accumulated into uniform bins over that interval.
    """

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

    def __init__(self, config: Config):
        super().__init__()
        if not 0 < config.top_k < config.num_experts:
            raise ValueError("top_k must be between zero and num_experts.")
        self.num_experts = config.num_experts
        self.top_k = config.top_k
        self.num_bins = config.num_bins
        self.register_buffer(
            "required_bias_histogram_EB",
            torch.zeros(config.num_experts, config.num_bins, dtype=torch.int32),
            persistent=False,
        )

    def observe(
        self,
        scores_TE: torch.Tensor,
        cutoff_T1: torch.Tensor,
        expert_bias_E: torch.Tensor,
        padding_mask_T: torch.Tensor | None,
    ) -> None:
        """Accumulate required-bias histograms for one local micro-batch."""
        if not self.training:
            return

        with spmd.no_typecheck(), torch.no_grad():
            lower_bound = expert_bias_E.min() - 1.0
            bin_width = (
                expert_bias_E.max() - expert_bias_E.min() + 2.0
            ) / self.num_bins
            required_bias_TE = cutoff_T1 - scores_TE
            bin_indices_TE = torch.floor(
                (required_bias_TE - lower_bound) / bin_width
            ).to(torch.int64)
            bin_indices_ET = bin_indices_TE.clamp_(0, self.num_bins - 1).transpose(0, 1)
            histogram_updates_ET = torch.ones_like(
                bin_indices_ET,
                dtype=self.required_bias_histogram_EB.dtype,
            )
            if padding_mask_T is not None:
                histogram_updates_ET = histogram_updates_ET * (
                    ~padding_mask_T
                ).unsqueeze(0)
            self.required_bias_histogram_EB.scatter_add_(
                1,
                bin_indices_ET,
                histogram_updates_ET,
            )

    def estimate_expert_bias(
        self,
        histogram_EB: torch.Tensor,
        expert_bias_E: torch.Tensor,
    ) -> torch.Tensor:
        """Estimate the next mean-centered expert bias from the histogram."""
        counts_E = histogram_EB.sum(dim=-1, dtype=torch.int64)
        target_count_E = counts_E.float() * (self.top_k / self.num_experts)
        cumulative_counts_EB = histogram_EB.cumsum(dim=-1, dtype=torch.int64)
        target_rank_E = target_count_E.ceil().to(torch.int64)
        target_bin_E = (cumulative_counts_EB < target_rank_E.unsqueeze(-1)).sum(dim=-1)

        target_bin_E1 = target_bin_E.unsqueeze(-1)
        counts_in_bin_E = histogram_EB.gather(-1, target_bin_E1).squeeze(-1)
        counts_before_E = (
            cumulative_counts_EB.gather(-1, target_bin_E1).squeeze(-1) - counts_in_bin_E
        )
        fraction_E = (
            target_count_E - counts_before_E.float()
        ) / counts_in_bin_E.float()

        bin_width = (expert_bias_E.max() - expert_bias_E.min() + 2.0) / self.num_bins
        quantile_position_E = target_bin_E.float() + fraction_E
        return (quantile_position_E - quantile_position_E.mean()) * bin_width

    def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None:
        if buffer_device is None:
            buffer_device = self.required_bias_histogram_EB.device
        with torch.device(buffer_device):
            self.required_bias_histogram_EB = torch.zeros(
                self.num_experts,
                self.num_bins,
                dtype=torch.int32,
            )


class MicrobatchWiseLoadBalanceLoss(AuxLoss):
    """Per-forward MoE load-balance gradient (DeepSeek-V3 Sec 2.1.2 Eqs 17-20).

    The balancing unit is one forward's folded token stream (a DP-local
    microbatch).  Global (corpus-level) balance is left to the
    auxiliary-loss-free bias path (``expert_bias_E``); this loss only
    discourages extreme load imbalance within individual forwards (samples),
    per the DeepSeek-V3 design (Sec 2.1.2, "Complementary Sequence-Wise
    Auxiliary Loss").

    With ``E`` experts, top-``K`` selection and ``T`` valid tokens per forward:

    Eq. 18: ``f_i = (E / (K T)) * sum_t 1[token t routes to expert i]``
    Eq. 19: ``p_i = (1 / T) * sum_t s'_t,i``,
            where ``s'_t,i = s_t,i / sum_j s_t,j`` is the per-token
            normalized score.
    Eq. 17: ``L_bal = sum_i f_i * p_i``

    The returned value is ``T * L_bal`` (token-mode): Eqs 17-20 define a
    per-token-normalized value, while ``AuxLoss`` scales every auxiliary
    loss by the reciprocal of the step's global routing-token count, so the
    sum-type form keeps the injected weight at ``coeff * L_bal``.

    The counts (Eq. 18) and normalized-score sums (Eq. 19) are sums over the
    folded token dim, hence Partial over the mesh axes that shard it (CP
    always, plus TP under EP).  They are all-reduced to Invariant before
    the formula, so every rank computes the same per-forward loss.  The
    one-hot counts are non-differentiable: the gradient reaches the router
    only through the normalized-score sums and the top-k score carrier.
    ``T`` is the forward's token count: the code evaluates Eq. 18 in the
    T-free form ``f_i = E * counts_i / sum_j counts_j``, which equals
    ``(E / (K T)) * counts_i`` because each token contributes K entries, so
    ``sum_j counts_j = K T``.  That needs no shape or mesh-degree assumption
    and follows any masking the router applies to the routing map.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(AuxLoss.Config):
        """Same fields as ``AuxLoss.Config``; this loss adds no knobs.

        A distinct Config is required even without new fields: ``Config.build()``
        constructs the class that owns the config (``__init_subclass__`` sets
        ``_owner``), so a router configured with ``AuxLoss.Config`` would
        build a plain ``AuxLoss``, which has no ``forward``.
        """

    def _reduce_token_partials(
        self, partial_E: torch.Tensor, axes: tuple[str, ...]
    ) -> torch.Tensor:
        """Partial -> Invariant all-reduce over the token-partition axes.

        Axes are passed by name, so spmd_types resolves them against the
        ambient mesh and no DeviceMesh escapes into model code; an inactive
        axis is skipped rather than run as a no-op collective.  ``P -> I`` is
        an all-reduce in forward with an identity backward: the reduced sums,
        and hence the loss and its gradient, are identical on every rank of
        the reduction group.
        """
        for axis in axes:
            if spmd_mesh_size(axis) == 1:
                # No mesh context or a size-1 axis: nothing shards the tokens.
                continue
            partial_E = spmd.redistribute(
                partial_E,
                axis,
                src=spmd.Partial,
                dst=spmd.Invariant,
                backward_options={"op_dtype": partial_E.dtype},
            )
        return partial_E

    def forward(
        self,
        scores_TE: torch.Tensor,
        routing_map_TE: torch.Tensor,
        *,
        carrier: torch.Tensor,
        padding_mask_T: torch.Tensor | None = None,
        denominator: torch.Tensor,
    ) -> torch.Tensor:
        """Compute the per-forward balance loss and inject its gradient.

        Args:
            scores_TE: Router scores ``(T, E)`` for the forward's tokens.
            routing_map_TE: One-hot routing map ``(T, E)`` for the same tokens,
                as counted by the router.
            carrier: Tensor whose backward path carries the injected
                gradient (the router's top-k scores).
            padding_mask_T: Boolean ``(T,)`` mask that is true for padding.

        Returns:
            ``carrier`` unchanged (identity forward).
        """
        # Mark DP local for the counts arithmetic: each DP rank owns an
        # independent token stream, so DP must not be reduced; only the
        # global axes that shard the stream (CP, TP under EP) are.
        with spmd_local_context("dp"):
            E = scores_TE.size(-1)
            # Axes that shard the router output's token dim: CP in every
            # layout, TP only under EP, which distributes tokens over TP (the
            # gate computes and emits dense_sequence_parallel_placement
            # whenever EP is on, and tokens_per_expert_E is TP-Partial for the
            # same reason).
            axes = ("cp", "tp")

            # Eq. 18: per-expert routing frequency counts_i over the forward's
            # tokens, then f_i = E * counts_i / sum_j counts_j (so
            # sum_i f_i = E).  The latter is the (E / (K T)) form with
            # T = sum_j counts_j / K, so it needs no token count, shape or mesh
            # degree and follows any masking the router applies to the map.
            # The map is cast to float before the reduction: casting a Partial
            # tensor is non-linear and rejected by spmd_types.
            counts_E = self._reduce_token_partials(
                routing_map_TE.to(scores_TE.dtype).sum(dim=0), axes
            )
            f_E = F.normalize(counts_E, p=1, dim=0) * E

            # Eq. 19: p_i = (1/T) sum_t s'_t,i, the per-token L1-normalized
            # scores.  F.normalize's eps clamp only guards an all-zero score
            # row: the scores are non-negative, so the norm is a plain sum.
            probs_TE = F.normalize(scores_TE, p=1, dim=-1)
            if padding_mask_T is not None:
                probs_TE = probs_TE * ~padding_mask_T.unsqueeze(-1)
            p_E = self._reduce_token_partials(probs_TE.sum(dim=0), axes)

            # Eq. 17: L_bal = sum_i f_i * p_i
            loss = (f_E * p_E).sum()
            return self.inject(loss, carrier=carrier, denominator=denominator)


class MoE(Module):
    """Mixture of Experts layer.

    The forward pass proceeds as:
    1. Router computes expert assignments.
    2. RoutedExperts.forward() enters a local SPMD region, then handles:
       a. dispatch (TokenDispatcher) — reorder tokens by expert assignment.
          With EP, also performs all-to-all communication to send tokens
          to expert-owning ranks.
       b. expert computation (W13, activation, and W2 on local tensors)
       c. combine (TokenDispatcher) — reverse the dispatch reordering.
          - LocalTokenDispatcher (no EP): scatter_add only.
          - AllToAll: all-to-all communication, then scatter_add.
          - DeepEP: combine_tokens followed by backend synchronization.
          - HybridEP: synchronous combine_tokens.
    3. Shared experts compute their output.
    4. Routed and shared expert outputs are summed.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Module.Config):
        num_experts: int = 8
        routed_experts: RoutedExperts.Config
        router: TokenChoiceTopKRouter.Config
        load_balance_coeff: float | None = 1e-3
        shared_experts: FeedForward.Config | None = None

        def __post_init__(self) -> None:
            expert_counts = {
                "moe": self.num_experts,
                "router": self.router.num_experts,
                "routed_experts": self.routed_experts.w13.group_size,
            }
            if len(set(expert_counts.values())) != 1:
                raise ValueError(
                    "MoE expert counts must match: "
                    + ", ".join(
                        f"{owner}={count}" for owner, count in expert_counts.items()
                    )
                )

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

        self.num_experts = config.num_experts
        self.routed_experts = config.routed_experts.build()
        self.router = config.router.build()
        self.shared_experts = (
            config.shared_experts.build() if config.shared_experts is not None else None
        )

        # define fields for auxiliary-loss-free load balancing (https://arxiv.org/abs/2408.15664)
        # NOTE: router.tokens_per_expert_E is accumulated in the router forward pass.
        #       expert_bias_E is updated outside the model in an optimizer step pre hook
        #       to work with gradient accumulation.
        self.load_balance_coeff = config.load_balance_coeff
        if self.load_balance_coeff is not None:
            assert self.load_balance_coeff > 0.0
            self.register_buffer(
                "expert_bias_E",
                torch.zeros(self.num_experts, dtype=torch.float32),
                persistent=True,
            )
        else:
            self.expert_bias_E = None

    def forward(
        self,
        x_TD: torch.Tensor,
        *,
        padding_mask_T: torch.Tensor | None = None,
        aux_loss_denominator: torch.Tensor | None = None,
        **router_kwargs,
    ) -> torch.Tensor:
        """
        Args:
            x_TD: Input ``(T, D)``.
            padding_mask_T: Boolean ``(T,)`` mask that is true for padding.

        Returns:
            Output ``(T, D)``.

        The MoE wrapper owns the TP transitions shared across its router,
        routed-expert, and shared-expert branches. Routed expert computation
        runs in a local SPMD region. When EP internally sequence-shards tokens
        across TP, the caller must provide a TP-divisible token count.
        """
        (
            routed_x_TD,
            routed_padding_mask_T,
        ) = self._maybe_shard_routed_branch_inputs_across_tp(x_TD, padding_mask_T)

        # topk scores and expert IDs have shape (T, K); the routing map (T, E)
        # marks the experts each token is routed to (built inside the router).
        (topk_scores_TK, topk_expert_ids_TK, routing_map_TE,) = self.router(
            routed_x_TD,
            self.expert_bias_E,
            padding_mask_T=routed_padding_mask_T,
            aux_loss_denominator=aux_loss_denominator,
            **router_kwargs,
        )
        num_local_tokens_per_expert_E = routing_map_TE.sum(dim=0)

        out_TD = self.routed_experts(
            routed_x_TD,
            topk_scores_TK,
            topk_expert_ids_TK,
            num_local_tokens_per_expert_E,
        )
        out_TD = self._maybe_zero_fill_routed_output_to_tp_partial(out_TD)
        if self.shared_experts is not None:
            shared_TD = self.shared_experts(x_TD)
            # Trailing add, always saved: it saves nothing for backward, so replay skips
            # it and its inputs need no persisting, matching checkpoint early stop.
            out_TD = remat.region(
                torch.add, self.remat_region_name("shared_add"), recompute=False
            )(out_TD, shared_TD)
        return self._maybe_all_reduce_moe_output_across_tp(out_TD)

    def _maybe_shard_routed_branch_inputs_across_tp(
        self,
        x_TD: torch.Tensor,
        padding_mask_T: torch.Tensor | None,
    ) -> tuple[torch.Tensor, torch.Tensor | None]:
        """Prepare the router and routed-expert inputs for TP token sharding.

        With EP, the routed branch consumes ``Shard(0)`` tokens across TP. Dense
        SP already provides that layout for ``x_TD``; otherwise this boundary
        explicitly shards it. The padding mask enters the MoE replicated in
        either case and is explicitly sharded here to follow the routed tokens.
        """
        if spmd_sparse_mesh() is None:
            assert (
                spmd_mesh_group(MeshAxisName.TP) is None
            ), "MoE requires expert parallelism when tensor parallelism is enabled on dense modules"
            return x_TD, padding_mask_T

        tp_group = spmd_mesh_group(MeshAxisName.TP)
        if tp_group is None:
            return x_TD, padding_mask_T

        if not spmd_dense_sp_enabled():
            x_TD = spmd.redistribute(
                x_TD,
                tp_group,
                src=spmd.I,
                dst=spmd.S(0),
                backward_options={"op_dtype": x_TD.dtype},
            )
        # The padding mask is replicated even when SP has already sharded x_TD.
        if padding_mask_T is not None:
            padding_mask_T = spmd.redistribute(
                padding_mask_T,
                tp_group,
                src=spmd.R,
                dst=spmd.S(0),
                backward_options={"op_dtype": padding_mask_T.dtype},
            )
        return x_TD, padding_mask_T

    def _maybe_zero_fill_routed_output_to_tp_partial(
        self, routed_output_TD: torch.Tensor
    ) -> torch.Tensor:
        """Convert the routed token shard to a TP partial without communication.

        With EP enabled and dense SP disabled, each TP rank computes only its
        ``Shard(0)`` of routed tokens, while the shared expert produces a
        ``Partial`` contribution over all tokens. Zero-filling positions owned
        by other TP ranks converts the routed output to the same ``Partial``
        representation. The two branches can then be added locally and kept
        partial until one final all-reduce materializes the complete MoE output.
        """
        if spmd_sparse_mesh() is None or spmd_dense_sp_enabled():
            return routed_output_TD

        tp_group = spmd_mesh_group(MeshAxisName.TP)
        if tp_group is None:
            return routed_output_TD
        # The zero-fill reads the routed output with bare ops.
        remat.recompute_needs_tensor(routed_output_TD)
        return spmd.redistribute(
            routed_output_TD,
            tp_group,
            src=spmd.S(0),
            dst=spmd.P,
            backward_options={"op_dtype": routed_output_TD.dtype},
        )

    def _maybe_all_reduce_moe_output_across_tp(
        self, out_TD: torch.Tensor
    ) -> torch.Tensor:
        """Reduce the combined partial output when dense SP is disabled."""
        if spmd_dense_sp_enabled():
            return out_TD
        tp_group = spmd_mesh_group(MeshAxisName.TP)
        if tp_group is None:
            return out_TD
        # This reduction needs a standalone region: routed token combine and
        # shared/routed branch addition separate it from either w2 computation.
        out_TD = remat.region(
            spmd.redistribute,
            self.remat_region_name("tp_output_reduction"),
            recompute=self.remat_should_recompute("tp_output_reduction"),
        )(
            out_TD,
            tp_group,
            src=spmd.P,
            dst=spmd.I,
            backward_options={"op_dtype": out_TD.dtype},
        )
        return out_TD

    def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None:
        if buffer_device is None:
            # After ``to_empty()``, the existing buffer records the target device.
            # Reinitialize MoE buffers there when no explicit buffer device is passed.
            buffer_device = self.router.tokens_per_expert_E.device

        with torch.device(buffer_device):
            if self.load_balance_coeff is not None:
                self.expert_bias_E = torch.zeros(
                    self.router.num_experts,
                    dtype=torch.float32,
                )


class _MoERouterLike(Protocol):
    tokens_per_expert_E: torch.Tensor  # noqa: N815


class _MoELike(Protocol):
    load_balance_coeff: float | None
    expert_bias_E: torch.Tensor  # noqa: N815
    router: _MoERouterLike


def register_moe_load_balancing_hook(
    optimizers: Optimizer,
    model_parts: list[nn.Module],
    parallelism_context: ParallelismContext,
) -> None:
    """Register an optimizer step pre-hook for MoE auxiliary-loss-free load balancing.

    This function checks if MoE load balancing is enabled and, if so, registers
    a hook that updates expert biases before each optimizer step.

    Args:
        optimizers: The optimizers container to register the hook on.
        model_parts: List of model parts that may contain MoE layers.
        parallelism_context: Parallel dimensions for distributed communication.
    """

    def _iter_moe_layers(
        model_parts: list[nn.Module],
    ) -> Iterator[tuple[nn.Module, _MoELike]]:
        for model_part in model_parts:
            # MTP decoder blocks live in ``mtp_layers`` after FSDP wrapping;
            # they are only temporarily inserted into ``layers`` while FSDP
            # is applied. Keep both containers in the runtime hook so MTP
            # expert bias and usage counters receive the same update as the
            # main decoder layers.
            layer_containers = [model_part.get_submodule("layers")]
            mtp_layers = getattr(model_part, "mtp_layers", None)
            if mtp_layers is not None:
                layer_containers.append(mtp_layers)
            for layers in layer_containers:
                assert isinstance(layers, (nn.ModuleDict, nn.ModuleList))
                for transformer_block in layers.children():
                    if getattr(transformer_block, "moe_enabled", False):
                        yield transformer_block, cast(_MoELike, transformer_block.moe)

    def _should_register_moe_balancing_hook(model_parts: list[nn.Module]) -> bool:
        moe_layers = list(_iter_moe_layers(model_parts))
        if not moe_layers:
            return False

        load_balance_enabled = moe_layers[0][1].load_balance_coeff is not None
        for _transformer_block, moe in moe_layers[1:]:
            if (moe.load_balance_coeff is not None) != load_balance_enabled:
                raise ValueError(
                    "MoE load_balance_coeff must be configured consistently "
                    "across all MoE layers. Either set it for every MoE layer "
                    "or leave it unset for all MoE layers."
                )
        return load_balance_enabled

    # for MoE auxiliary-loss-free load balancing
    def _is_recomputation_enabled(module):
        return getattr(module, "checkpoint_impl", None) is CheckpointImpl.NO_REENTRANT

    def _update_expert_bias(
        model_parts: list[nn.Module],
        parallelism_context: ParallelismContext,
    ):
        loss_mesh = parallelism_context.get_optional_mesh("loss")
        # TODO: Currently this sync is blocking (thus exposed) and happens on the
        # default compute stream. Need to assess if this is OK performance-wise.
        tokens_per_expert_E_list = []
        for transformer_block, moe in _iter_moe_layers(model_parts):
            tokens_per_expert_E = moe.router.tokens_per_expert_E
            if _is_recomputation_enabled(transformer_block):
                # TODO: This is a hack, we assume with full AC, the tokens_per_expert_E is counted twice.
                # This does not affect to expert choice, but affects the experts usage metrics.
                # We divide by 2 to correct for this double-counting due to recomputation
                # TODO: new API to help determine if AC is enabled https://github.com/pytorch/pytorch/pull/160888
                tokens_per_expert_E = tokens_per_expert_E // 2
            tokens_per_expert_E_list.append(tokens_per_expert_E)

        if not tokens_per_expert_E_list:
            return

        tokens_per_expert_E_by_layer = torch.vstack(tokens_per_expert_E_list)

        if parallelism_context.ep_enabled and parallelism_context.tp > 1:
            torch.distributed.all_reduce(
                tokens_per_expert_E_by_layer,
                group=parallelism_context.get_dense_tp_mesh().get_group(),
            )
        if loss_mesh is not None:
            torch.distributed.all_reduce(
                tokens_per_expert_E_by_layer,
                group=loss_mesh.get_group(),
                op=torch.distributed.ReduceOp.SUM,
            )
        moe_layer_idx = 0
        with torch.no_grad():
            for _transformer_block, moe in _iter_moe_layers(model_parts):
                load_balance_coeff = moe.load_balance_coeff
                assert load_balance_coeff is not None

                tokens_per_expert_E = tokens_per_expert_E_by_layer[
                    moe_layer_idx
                ].float()
                moe_layer_idx += 1

                # update the expert bias
                # this is not exactly the same as https://arxiv.org/pdf/2408.15664 proposed
                expert_bias_delta_E = load_balance_coeff * torch.sign(
                    tokens_per_expert_E.mean() - tokens_per_expert_E
                )
                expert_bias_delta_E = expert_bias_delta_E - expert_bias_delta_E.mean()
                moe.expert_bias_E.add_(expert_bias_delta_E)
                moe.router.tokens_per_expert_E.zero_()

    if _should_register_moe_balancing_hook(model_parts):
        optimizers.register_step_pre_hook(
            lambda *args, **kwargs: _update_expert_bias(
                model_parts, parallelism_context=parallelism_context
            )
        )


def register_moe_quantile_balancing_hook(
    optimizers: Optimizer,
    model_parts: list[nn.Module],
    parallelism_context: ParallelismContext,
) -> None:
    """Update quantile-balanced expert biases before each optimizer step."""
    moe_layers: list[tuple[MoE, QuantileBalancedTopKRouter]] = []
    for model_part in model_parts:
        for module in model_part.modules():
            if isinstance(module, MoE) and isinstance(
                module.router, QuantileBalancedTopKRouter
            ):
                moe_layers.append((module, module.router))

    if not moe_layers:
        return

    @torch.no_grad()
    def _update_expert_bias() -> None:
        reduction_groups = []
        # With EP, the router is token-sharded on the dense TP axis even when
        # model-wide sequence parallelism is disabled.
        if parallelism_context.ep_enabled and parallelism_context.tp > 1:
            reduction_groups.append(parallelism_context.get_dense_tp_mesh().get_group())
        loss_mesh = parallelism_context.get_optional_mesh("loss")
        if loss_mesh is not None:
            reduction_groups.append(loss_mesh.get_group())

        histograms = [
            router.quantile_balancer.required_bias_histogram_EB
            for _moe, router in moe_layers
        ]
        if reduction_groups:
            reduced_histograms_LEB = torch.stack(histograms)
            for group in reduction_groups:
                torch.distributed.all_reduce(
                    reduced_histograms_LEB,
                    group=group,
                    op=torch.distributed.ReduceOp.SUM,
                )
            histograms = list(reduced_histograms_LEB.unbind())

        for histogram_EB, (moe, router) in zip(
            histograms,
            moe_layers,
            strict=True,
        ):
            expert_bias_E = moe.expert_bias_E
            assert expert_bias_E is not None
            quantile_balancer = router.quantile_balancer
            next_expert_bias_E = quantile_balancer.estimate_expert_bias(
                histogram_EB,
                expert_bias_E,
            )
            expert_bias_E.copy_(next_expert_bias_E)
            quantile_balancer.required_bias_histogram_EB.zero_()
            router.tokens_per_expert_E.zero_()

    optimizers.register_step_pre_hook(lambda *args, **kwargs: _update_expert_bias())
