# 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.

"""Replace HF MoE blocks with Titan MoE modules.

Two-phase replacement:
  Phase 1 (init time): ``prepare_native_moe_configs`` probes the HF MoE block
      and builds a Titan ``MoE.Config``, stored on each layer as
      ``_native_moe_config``.
  Phase 2 (parallelize time): ``build_and_swap_native_moe`` calls
      ``set_moe_sharding_config`` on each stored config, builds the Titan MoE,
      initializes it, and swaps it into the layer. Actual parallelization
      happens later via ``model._parallelize(parallelism_context)``.
"""

import logging
from collections.abc import Callable
from dataclasses import replace
from functools import partial

import spmd_types as spmd
import torch
import torch.nn as nn

from torchtitan.distributed.parallelism_context import ParallelismContext
from torchtitan.experiments.transformers_modeling_backend.hf_sharding import (
    _hf_activation_placement,
    _hf_sequence_parallel_placement,
)
from torchtitan.models.common import Sigmoid, Softmax
from torchtitan.models.common.config_utils import (
    fused_gate_up_param_init,
    make_moe_config,
    make_routed_experts_config,
    make_router_config,
    make_shared_expert_ffn_config,
)
from torchtitan.models.common.linear import Linear
from torchtitan.models.common.moe import MoE
from torchtitan.models.common.moe_sharding import (
    set_moe_sharding_config,
    set_routed_moe_sharding_config,
)
from torchtitan.models.deepseek_v3.flavors import make_deepseek_v3_router_config


logger = logging.getLogger(__name__)


class _HFBatchedMoE(MoE):
    """Adapt HF's singleton batch to Titan MoE's flat token interface."""

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        return super().forward(hidden_states.squeeze(0)).unsqueeze(0)


# ---------------------------------------------------------------------------
# Phase 1 — config preparation (called from HFTransformerModel.__init__)
# ---------------------------------------------------------------------------


def prepare_native_moe_configs(model: nn.Module, config) -> None:
    """Probe each HF MoE layer and store a Titan ``MoE.Config`` for later build.

    Called during ``HFTransformerModel.__init__`` on meta device.
    Does NOT instantiate the Titan MoE modules yet -- that happens in Phase 2.
    """
    for layer in model.layers.values():
        if not getattr(layer, "moe_enabled", False):
            continue

        is_layer_level = getattr(layer, "_layer_level_moe", False)
        if is_layer_level:
            # Layer-level MoE (Gemma4): router/experts are siblings of dense
            # MLP. Probe from the layer itself, treating the dense MLP as
            # the shared expert.
            moe_params = _probe_layer_level_moe(layer, config)
        else:
            moe_block = _get_moe_block(layer)
            moe_params = _probe_hf_moe_block(moe_block, config)
        moe_config = _build_moe_config(moe_params, config)
        layer._native_moe_config = moe_config

    logger.info("Prepared Titan MoE configs for all MoE layers")


# ---------------------------------------------------------------------------
# Phase 2 — build, init, swap (called from parallelize_hf_transformers)
# ---------------------------------------------------------------------------


def build_and_swap_native_moe(
    model: nn.Module,
    parallelism_context: ParallelismContext,
) -> None:
    """Build Titan MoE modules and swap them into the model.

    For each MoE layer with a stored ``_native_moe_config``:
    1. Set sharding config on the MoE.Config (now that EP/TP is known)
    2. Build the Titan MoE module
    3. Initialize parameters and buffers
    4. Swap into the layer's MoE attribute, set ``layer.moe`` for load-balancing hook

    Args:
        model: The HFTransformerModel with ``_native_moe_config`` stored on
            each MoE-enabled layer (from ``prepare_native_moe_configs``).
        parallelism_context: Parallel dimensions for EP/TP mesh resolution.
    """
    if parallelism_context.ep < parallelism_context.tp:
        raise ValueError(
            f"MoE models require expert_parallel_degree ({parallelism_context.ep}) to be "
            "greater than or equal to tensor_parallel_degree "
            f"({parallelism_context.tp})."
        )
    enable_ep = parallelism_context.ep_enabled
    enable_sp = parallelism_context.tp_enabled

    for layer in model.layers.values():
        moe_config = getattr(layer, "_native_moe_config", None)
        if moe_config is None:
            continue

        shared_experts = moe_config.shared_experts
        if shared_experts is not None and hasattr(shared_experts, "gate"):
            # Avoid loading Qwen3.5's optional model dependencies for other HF
            # architectures handled by this generic experiment.
            from torchtitan.models.qwen3_5.moe import SigmoidGatedFeedForward
            from torchtitan.models.qwen3_5.sharding import (
                set_sigmoid_gated_feed_forward_sharding_config,
            )

            assert isinstance(shared_experts, SigmoidGatedFeedForward.Config)
            set_routed_moe_sharding_config(
                moe_config,
                enable_ep=enable_ep,
            )
            set_sigmoid_gated_feed_forward_sharding_config(
                shared_experts, enable_sp=enable_sp
            )
        else:
            set_moe_sharding_config(
                moe_config,
                enable_ep=enable_ep,
                enable_sp=enable_sp,
            )
        root_sharding = moe_config.sharding_config
        assert root_sharding is not None
        # Only the MoE root sees HF's singleton batch. Its children retain the
        # standard Titan ``(tokens, hidden)`` layouts used inside MoE.forward.
        hf_sp_layout = (
            _hf_sequence_parallel_placement()
            if enable_sp
            else _hf_activation_placement(tp=spmd.I)
        )
        moe_config.sharding_config = replace(
            root_sharding,
            in_src_shardings={"hidden_states": hf_sp_layout},
            out_src_shardings=hf_sp_layout,
        )

        with torch.device("meta"):
            native_moe = moe_config.build()
        # Preserve the Titan MoE state-dict paths while adapting its root
        # forward boundary; an nn.Module wrapper would add another name level.
        native_moe.__class__ = _HFBatchedMoE

        # Materialize meta params to real tensors, then initialize values.
        # This mirrors the trainer's flow: to_empty → init_states.
        native_moe.to_empty(device=torch.device("cpu"))
        native_moe.init_states(buffer_device=torch.device("cpu"))

        # Swap the Titan MoE into the layer's original ``mlp`` attribute.
        moe_attr = _get_moe_attr_name(layer)
        setattr(layer, moe_attr, native_moe)
        object.__setattr__(layer, "moe", native_moe)

        # For layer-level MoE (Gemma4), the original router and experts
        # are layer-level siblings of the dense MLP. After swapping in
        # the Titan MoE at ``layer.mlp``, disable the HF forward's
        # separate MoE path (which references ``self.router`` and
        # ``self.experts``) and delete the original modules to prevent
        # duplicate parameter registration.
        if getattr(layer, "_layer_level_moe", False):
            if hasattr(layer, "enable_moe_block"):
                layer.enable_moe_block = False
            if hasattr(layer, "router"):
                delattr(layer, "router")
            if hasattr(layer, "experts"):
                delattr(layer, "experts")
            # Remove unused norms from the separate MoE path to keep
            # the parameter count clean.
            for norm_name in (
                "pre_feedforward_layernorm_2",
                "post_feedforward_layernorm_1",
                "post_feedforward_layernorm_2",
            ):
                if hasattr(layer, norm_name):
                    delattr(layer, norm_name)

        del layer._native_moe_config

    logger.info("Built and swapped Titan MoE modules into the model")


# ---------------------------------------------------------------------------
# MoE attribute helpers
# ---------------------------------------------------------------------------


def _get_moe_attr_name(layer: nn.Module) -> str:
    """Return the attribute name holding the MoE block on a decoder layer.

    Models use ``mlp``; for layer-level MoE (Gemma4) the Titan MoE
    replaces ``mlp`` too.
    """
    if hasattr(layer, "mlp"):
        return "mlp"
    raise AttributeError(f"Layer {type(layer).__name__} has no 'mlp'")


def _get_moe_block(layer: nn.Module) -> nn.Module:
    """Return the MoE block module from a decoder layer."""
    return getattr(layer, _get_moe_attr_name(layer))


# ---------------------------------------------------------------------------
# HF MoE block probing
# ---------------------------------------------------------------------------


def _probe_hf_moe_block(moe_block: nn.Module, config) -> dict:
    """Extract MoE configuration from an HF MoE block.

    Args:
        moe_block: The HF MoE block (e.g., ``Qwen3MoeSparseMoeBlock``).
        config: The HF model config with MoE-related attributes.

    Returns:
        Dict with all parameters needed to build a Titan ``MoE.Config``.
    """
    gate = getattr(moe_block, "gate", None) or getattr(moe_block, "router", None)
    experts = moe_block.experts

    num_experts = _resolve_num_experts(experts, gate, moe_block, config)
    dim = config.hidden_size

    # Intermediate size: HF MoE models use fused gate_up_proj with the
    # standard (E, 2*I, H) layout, so dim 1 is 2*I.
    if hasattr(experts, "gate_up_proj"):
        moe_intermediate_size = experts.gate_up_proj.shape[1] // 2
    elif hasattr(config, "moe_intermediate_size") and config.moe_intermediate_size:
        moe_intermediate_size = config.moe_intermediate_size
    else:
        moe_intermediate_size = getattr(config, "intermediate_size", dim * 4)

    top_k = _resolve_top_k(moe_block, gate, config)

    # Router scoring function
    score_func = _resolve_score_func(gate, config)

    # Route normalization: some models (Mixtral, Qwen3.5) always normalize
    # but don't have a config flag. Detect by checking if norm_topk_prob is
    # absent (meaning the router hardcodes normalization).
    # Sigmoid-routing models without explicit norm_topk_prob default to
    # False — sigmoid outputs are used directly as scores without normalization.
    if hasattr(config, "norm_topk_prob"):
        route_norm = config.norm_topk_prob
    elif hasattr(gate, "norm_topk_prob"):
        route_norm = gate.norm_topk_prob
    else:
        route_norm = score_func != "sigmoid"

    # Route scaling factor (DeepSeek V2/V3)
    route_scale = getattr(config, "routed_scaling_factor", 1.0)

    # Group-limited routing (DeepSeek V2/V3)
    num_expert_groups = getattr(config, "n_group", None)
    num_limited_groups = getattr(config, "topk_group", None)

    # Load balance coefficient
    load_balance_coeff = getattr(config, "load_balance_coeff", 1e-3)

    # Shared experts
    shared_expert_info = _probe_shared_experts(moe_block, config)

    return {
        "num_experts": num_experts,
        "dim": dim,
        "moe_intermediate_size": moe_intermediate_size,
        "top_k": top_k,
        "score_func": score_func,
        "route_norm": route_norm,
        "route_scale": route_scale,
        "num_expert_groups": num_expert_groups,
        "num_limited_groups": num_limited_groups,
        "load_balance_coeff": load_balance_coeff,
        "shared_expert_info": shared_expert_info,
    }


def _resolve_num_experts(
    experts: nn.Module, gate: nn.Module | None, moe_block: nn.Module, config
) -> int:
    """Infer the total expert count from the HF MoE block or config."""
    for owner in (experts, gate, moe_block):
        if owner is None:
            continue
        for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
            val = getattr(owner, attr, None)
            if val is not None:
                return int(val)
    for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
        val = getattr(config, attr, None)
        if val is not None:
            return int(val)
    if gate is not None and hasattr(gate, "weight"):
        return gate.weight.shape[0]
    raise ValueError("Could not determine number of experts from HF MoE block")


def _resolve_top_k(moe_block: nn.Module, gate: nn.Module | None, config) -> int:
    """Infer top-k routing from the HF MoE block or config."""
    for owner in (moe_block, gate):
        if owner is None:
            continue
        for attr in ("top_k", "num_experts_per_tok", "top_k_experts"):
            val = getattr(owner, attr, None)
            if val is not None:
                return int(val)
    for attr in ("num_experts_per_tok", "top_k_experts"):
        val = getattr(config, attr, None)
        if val is not None:
            return int(val)
    raise ValueError("Could not determine top_k from HF MoE block")


def _resolve_score_func(gate: nn.Module | None, config) -> str:
    """Determine the router scoring function (softmax or sigmoid)."""
    # DeepSeek-V3 / GLM use sigmoid with e_score_correction_bias.
    if gate is not None and "e_score_correction_bias" in getattr(gate, "_buffers", {}):
        return "sigmoid"

    scoring_func = getattr(config, "scoring_func", None)
    if scoring_func is not None:
        if scoring_func in ("softmax", "sigmoid"):
            return scoring_func
        raise ValueError(
            f"Unsupported scoring function '{scoring_func}'. "
            "Titan MoE router supports 'softmax' and 'sigmoid'."
        )

    return "softmax"


def _probe_shared_experts(moe_block: nn.Module, config) -> dict | None:
    """Detect shared expert configuration from the HF MoE block."""
    shared = None
    for name in ("shared_expert", "shared_experts", "shared_mlp"):
        shared = getattr(moe_block, name, None)
        if shared is not None:
            break

    if shared is None:
        return None

    # Determine shared expert intermediate size
    for attr in ("intermediate_size", "hidden_size"):
        if hasattr(shared, attr):
            shared_hidden_dim = getattr(shared, attr)
            break
    else:
        # Try to infer from weight shapes
        gate_proj = getattr(shared, "gate_proj", None)
        if gate_proj is not None and hasattr(gate_proj, "weight"):
            shared_hidden_dim = gate_proj.weight.shape[0]
        else:
            shared_hidden_dim = getattr(config, "shared_expert_intermediate_size", None)
            if shared_hidden_dim is None:
                n_shared = getattr(config, "n_shared_experts", 1)
                shared_hidden_dim = (
                    getattr(config, "moe_intermediate_size", config.hidden_size)
                    * n_shared
                )

    # Check for sigmoid-gated shared expert (Qwen3.5 pattern)
    shared_expert_gate = getattr(moe_block, "shared_expert_gate", None)
    has_sigmoid_gate = shared_expert_gate is not None

    return {
        "hidden_dim": shared_hidden_dim,
        "dim": config.hidden_size,
        "has_sigmoid_gate": has_sigmoid_gate,
    }


def _probe_layer_level_moe(layer: nn.Module, config) -> dict:
    """Probe layer-level MoE (Gemma4) where router/experts are layer siblings.

    The dense MLP is treated as a shared expert: its output is summed with
    the routed expert output in the original HF forward. The Titan MoE
    replaces ``layer.mlp`` and contains all three components (router,
    experts, shared_experts=dense MLP).
    """
    gate = getattr(layer, "gate", None) or getattr(layer, "router", None)
    experts = layer.experts

    num_experts = _resolve_num_experts(experts, gate, layer, config)
    dim = config.hidden_size
    top_k = _resolve_top_k(layer, gate, config)
    score_func = _resolve_score_func(gate, config)

    # Expert intermediate size from fused gate_up_proj
    if hasattr(experts, "gate_up_proj"):
        shape = experts.gate_up_proj.shape
        if shape[2] == dim and shape[1] != dim:
            moe_intermediate_size = shape[1] // 2
        elif shape[1] == dim and shape[2] != dim:
            moe_intermediate_size = shape[2] // 2
        else:
            moe_intermediate_size = getattr(
                config, "moe_intermediate_size", shape[1] // 2
            )
    else:
        moe_intermediate_size = getattr(config, "moe_intermediate_size", dim * 4)

    # Route normalization
    if hasattr(config, "norm_topk_prob"):
        route_norm = config.norm_topk_prob
    elif hasattr(gate, "norm_topk_prob") if gate else False:
        route_norm = gate.norm_topk_prob
    else:
        route_norm = score_func != "sigmoid"

    route_scale = getattr(config, "routed_scaling_factor", 1.0)
    num_expert_groups = getattr(config, "n_group", None)
    num_limited_groups = getattr(config, "topk_group", None)
    load_balance_coeff = getattr(config, "load_balance_coeff", 1e-3)
    # Dense MLP is the shared expert
    mlp = getattr(layer, "mlp", None)
    shared_expert_info = None
    if mlp is not None:
        gate_proj = getattr(mlp, "gate_proj", None)
        if gate_proj is not None and hasattr(gate_proj, "weight"):
            shared_hidden_dim = gate_proj.weight.shape[0]
        else:
            shared_hidden_dim = getattr(config, "intermediate_size", dim * 4)
        shared_expert_info = {
            "hidden_dim": shared_hidden_dim,
            "dim": dim,
            "has_sigmoid_gate": False,
        }

    return {
        "num_experts": num_experts,
        "dim": dim,
        "moe_intermediate_size": moe_intermediate_size,
        "top_k": top_k,
        "score_func": score_func,
        "route_norm": route_norm,
        "route_scale": route_scale,
        "num_expert_groups": num_expert_groups,
        "num_limited_groups": num_limited_groups,
        "load_balance_coeff": load_balance_coeff,
        "shared_expert_info": shared_expert_info,
    }


# ---------------------------------------------------------------------------
# Sigmoid-gated shared expert wrapper (Qwen3.5 pattern)
# ---------------------------------------------------------------------------


# ---------------------------------------------------------------------------
# Config building
# ---------------------------------------------------------------------------

_LINEAR_INIT = {
    "weight": partial(nn.init.trunc_normal_, std=0.02),
    "bias": nn.init.zeros_,
}


def _get_expert_param_init() -> dict[str, Callable]:
    """Return initializers for the three logical expert projections."""
    init_fn = partial(nn.init.trunc_normal_, std=0.02)
    return {"w1_EFD": init_fn, "w2_EDF": init_fn, "w3_EFD": init_fn}


def _build_moe_config(params: dict, config) -> MoE.Config:
    """Build a fully-specified MoE.Config from probed parameters."""
    router_config_factory = (
        make_deepseek_v3_router_config
        if params["num_expert_groups"] is not None
        else make_router_config
    )
    router_kwargs = {}
    if params["num_expert_groups"] is not None:
        router_kwargs = {
            "num_expert_groups": params["num_expert_groups"],
            "num_limited_groups": params["num_limited_groups"],
        }
    router = router_config_factory(
        dim=params["dim"],
        num_experts=params["num_experts"],
        gate_param_init=_LINEAR_INIT,
        top_k=params["top_k"],
        score_func=(
            Sigmoid.Config() if params["score_func"] == "sigmoid" else Softmax.Config()
        ),
        route_norm=params["route_norm"],
        route_scale=params["route_scale"],
        **router_kwargs,
    )

    routed_experts = make_routed_experts_config(
        dim=params["dim"],
        hidden_dim=params["moe_intermediate_size"],
        num_experts=params["num_experts"],
        top_k=params["top_k"],
        param_init=_get_expert_param_init(),
    )

    shared_experts = None
    shared_info = params["shared_expert_info"]
    if shared_info is not None:
        ffn_config = make_shared_expert_ffn_config(
            dim=shared_info["dim"],
            hidden_dim=shared_info["hidden_dim"],
            w13_param_init=_LINEAR_INIT,
            w2_param_init=_LINEAR_INIT,
        )
        if shared_info["has_sigmoid_gate"]:
            # Import only for the Qwen3.5 topology so unrelated HF models do
            # not load Qwen3.5's optional model dependencies.
            from torchtitan.models.qwen3_5.moe import SigmoidGatedFeedForward

            shared_experts = SigmoidGatedFeedForward.Config(
                # The enclosing MoE gathers once because w13 and the sigmoid
                # gate consume the same input.
                w13=Linear.Config(
                    in_features=shared_info["dim"],
                    out_features=shared_info["hidden_dim"],
                    num_linears=2,
                    param_init=fused_gate_up_param_init(_LINEAR_INIT, _LINEAR_INIT),
                ),
                w2=ffn_config.w2,
                activation_fn=ffn_config.activation_fn,
                gate=Linear.Config(
                    in_features=shared_info["dim"],
                    out_features=1,
                    bias=False,
                    param_init=_LINEAR_INIT,
                ),
            )
        else:
            shared_experts = ffn_config

    return make_moe_config(
        num_experts=params["num_experts"],
        router=router,
        routed_experts=routed_experts,
        shared_experts=shared_experts,
        load_balance_coeff=params["load_balance_coeff"],
    )
