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

"""Shared helpers for building model configurations.

These helpers construct fully-specified sub-configs with all dimensional
fields set at config creation time.
"""

import dataclasses
from collections.abc import Callable

import torch
from torch.distributed.tensor import DTensor

from torchtitan.distributed.spmd_types import current_spmd_mesh, spmd_mesh_size
from torchtitan.models.common.activation import UnaryActivationFn
from torchtitan.models.common.attention import (
    FlexInnerAttention,
    GQAttention,
    QKVLinear,
    VarlenInnerAttention,
)
from torchtitan.models.common.decoder import Decoder
from torchtitan.models.common.feed_forward import FeedForward
from torchtitan.models.common.hi_mid_lo_linear import HiMidLoLinear
from torchtitan.models.common.linear import (
    ColumnParallelLinear,
    GroupedLinear,
    RowParallelLinear,
    SharedExpertRowParallelLinear,
)
from torchtitan.models.common.moe import (
    MicrobatchWiseLoadBalanceLoss,
    MoE,
    RoutedExperts,
    TokenChoiceTopKRouter,
)
from torchtitan.models.common.nn_modules import RMSNorm
from torchtitan.models.common.rope import RoPE
from torchtitan.models.common.token_dispatcher import AllToAllTokenDispatcher
from torchtitan.protocols.module import Module


DEFAULT_DEBUG_MODEL_SEQ_LEN = 2048


def _make_fused_linear_init(gate_init: Callable, up_init: Callable) -> Callable:
    """Build an initializer for a stacked gate/up linear weight."""

    def _init(t: torch.Tensor) -> None:
        gate_init(t[0])
        up_init(t[1])

    return _init


def decoder_vocab_size(model_config: Module.Config) -> int:
    """Assert Decoder.Config type so lint is not annoyed."""
    assert isinstance(model_config, Decoder.Config)
    return model_config.vocab_size


def get_attention_config(
    backend: str,
) -> Module.Config:
    """Map backend string to an inner_attention config.

    Language models always use block_causal masking (the dataloaders always
    emit per-document positions), so every backend here is a masked attention
    backend. ``ScaledDotProductInnerAttention`` only supports a boolean ``is_causal``
    flag and cannot consume per-document positions, so it is not a valid
    language-model backend (it remains available for Flux, which builds it
    directly).
    """
    if backend == "flex":
        return FlexInnerAttention.Config()
    elif backend == "flex_flash":
        from torchtitan.tools.utils import has_cuda_capability

        if not has_cuda_capability(9, 0):
            raise ValueError(
                "Flash backend of FlexInnerAttention is only supported on Hopper or Blackwell"
            )
        return FlexInnerAttention.Config(
            block_size=(256, 128), kernel_options={"BACKEND": "FLASH"}
        )
    elif backend == "varlen":
        return VarlenInnerAttention.Config()
    elif backend == "sdpa":
        raise ValueError(
            "sdpa is no longer supported for language models; positions are "
            "always available so use flex, flex_flash, or varlen."
        )
    else:
        raise ValueError(f"Unknown backend: {backend}")


def fused_qkv_param_init(
    base_param_init: dict[str, Callable],
    *,
    n_heads: int,
    n_kv_heads: int,
    head_dim: int,
) -> dict[str, Callable]:
    """Initialize fused ``wqkv`` from logical ``wq``/``wk``/``wv`` draws.

    Q, K, and V are initialized as separate contiguous tensors, then packed into
    the fused ``(n_kv_heads, R, head_dim, dim)`` layout, where
    ``R = heads_per_kv + 2``. This preserves logical initialization order and
    matches the packing used when loading separate checkpoint tensors.

    Parallelism-agnostic RNG: at init ``t`` is the (possibly sharded) matrix --
    e.g. a ``Shard(0)`` DTensor for the colwise wqkv. ``t.new_empty(...)``
    returns ``Replicate`` DTensors, so each ``base_init`` runs on the full tensor
    and draws the same values on every rank (the weights do not depend on the
    TP/FSDP degree). ``cat`` of ``Replicate`` stays ``Replicate``, and the final
    ``copy_`` scatters it into the sharded ``t`` (each rank keeps its shard).
    This "init replicated, then shard" path keeps RNG independent of the
    parallelism.
    """
    heads_per_kv = n_heads // n_kv_heads

    def _make_init(base_init: Callable) -> Callable:
        # ``tail`` is the per-row shape: () for bias, (in_features,) for weight.
        # Building q/k/v with their logical shapes preserves their independent
        # initialization order before packing them into wqkv.
        def _init(t):
            tail = t.shape[1:]
            # If t is a sharded DTensor, new_empty (with the full logical shape)
            # returns Replicate DTensors, so base_init runs replicated and draws
            # the same values on every rank (parallelism-agnostic RNG).
            q = t.new_empty(n_heads * head_dim, *tail)
            k = t.new_empty(n_kv_heads * head_dim, *tail)
            v = t.new_empty(n_kv_heads * head_dim, *tail)
            base_init(q)
            base_init(k)
            base_init(v)
            fused = torch.cat(
                [
                    q.view(n_kv_heads, heads_per_kv, head_dim, *tail),
                    k.view(n_kv_heads, 1, head_dim, *tail),
                    v.view(n_kv_heads, 1, head_dim, *tail),
                ],
                dim=1,
            ).view(-1, *tail)
            with torch.no_grad():
                # fused is Replicate (cat of Replicates). Flatten it to t's native
                # 2D [(n_kv_heads*r_dim*head_dim), *tail] shape and copy_ into t,
                # which scatters each rank's own shard. Reshaping the Replicate
                # source (instead of t.view(n_kv_heads, ...) on the sharded t)
                # avoids the "unflatten unevenly sharded" error when dp_shard*tp
                # does not divide n_kv_heads (e.g. dp_shard=8, n_kv_heads=4); no
                # gather, since fused is already replicated.
                if not isinstance(t, DTensor) and (tp_size := spmd_mesh_size("tp")) > 1:
                    # RL generator init_weights() only needs non-persistent
                    # buffers; weights come from trainer state dict. Until it has
                    # a DTensor static state dict path, copy this TP shard here.
                    # TODO: Remove once RL can init buffers without weight init.
                    mesh = current_spmd_mesh()
                    assert mesh is not None
                    tp_rank = mesh.get_local_rank("tp")
                    fused = fused.chunk(tp_size, dim=0)[tp_rank]
                t.copy_(fused)

        return _init

    out: dict[str, Callable] = {}
    for param in ("weight", "bias"):
        base_init = base_param_init.get(param)
        if base_init is not None:
            out[param] = _make_init(base_init)
    return out


def fused_gate_up_param_init(
    gate_param_init: dict[str, Callable],
    up_param_init: dict[str, Callable],
) -> dict[str, Callable] | None:
    """Initialize the logical gate and up slices of a fused ``w13`` weight."""
    gate_init = gate_param_init.get("weight")
    up_init = up_param_init.get("weight")
    if gate_init is None or up_init is None:
        return None
    return {"weight": _make_fused_linear_init(gate_init, up_init)}


def fused_grouped_gate_up_param_init(
    param_init: dict[str, Callable],
) -> dict[str, Callable]:
    """Build ``w13.weight`` initialization from logical expert projections."""
    missing = {"w1_EFD", "w3_EFD"} - param_init.keys()
    if missing:
        raise ValueError(f"Missing routed-expert initializers: {sorted(missing)}")

    def init(weight_E2FD: torch.Tensor) -> None:
        param_init["w1_EFD"](weight_E2FD[:, 0])
        param_init["w3_EFD"](weight_E2FD[:, 1])

    return {"weight": init}


def make_gqa_config(
    *,
    dim: int,
    n_heads: int,
    wqkv_param_init: dict[str, Callable],
    wo_param_init: dict[str, Callable],
    inner_attention: Module.Config,
    rope: RoPE.Config | None,
    n_kv_heads: int | None = None,
    head_dim: int | None = None,
    qk_norm: RMSNorm.Config | None = None,
) -> GQAttention.Config:
    """Build a fully-specified GQAttention.Config.

    ``rope=None`` builds a NoPE layer (no positional encoding); see
    :class:`GQAttention`.

    The projection types make the standard synchronous TP collectives explicit.
    Without a TP mesh, they execute as ordinary linear modules.
    """
    n_kv = n_kv_heads if n_kv_heads is not None else n_heads
    per_head_dim = head_dim if head_dim is not None else dim // n_heads
    rope = dataclasses.replace(rope) if rope is not None else None

    qkv = QKVLinear.Config(
        head_dim=per_head_dim,
        n_heads=n_heads,
        n_kv_heads=n_kv,
        wqkv=ColumnParallelLinear.Config(
            in_features=dim,
            out_features=(n_heads + 2 * n_kv) * per_head_dim,
            param_init=fused_qkv_param_init(
                wqkv_param_init,
                n_heads=n_heads,
                n_kv_heads=n_kv,
                head_dim=per_head_dim,
            ),
        ),
    )

    return GQAttention.Config(
        n_heads=n_heads,
        n_kv_heads=n_kv_heads,
        head_dim=head_dim,
        dim=dim,
        qkv_linear=qkv,
        wo=RowParallelLinear.Config(
            in_features=n_heads * per_head_dim,
            out_features=dim,
            param_init=wo_param_init,
        ),
        qk_norm=qk_norm,
        inner_attention=inner_attention,
        rope=rope,
    )


def make_ffn_config(
    *,
    dim: int,
    hidden_dim: int,
    w13_param_init: dict[str, Callable],
    w2_param_init: dict[str, Callable],
) -> FeedForward.Config:
    """Build a fully-specified FeedForward.Config.

    ``w1`` and ``w3`` are the gate/up projections and share
    ``w13_param_init``; ``w2`` is the residual output projection.
    """
    return FeedForward.Config(
        w13=ColumnParallelLinear.Config(
            in_features=dim,
            out_features=hidden_dim,
            num_linears=2,
            param_init=fused_gate_up_param_init(w13_param_init, w13_param_init),
        ),
        w2=RowParallelLinear.Config(
            in_features=hidden_dim,
            out_features=dim,
            param_init=w2_param_init,
        ),
    )


def make_shared_expert_ffn_config(
    *,
    dim: int,
    hidden_dim: int,
    w13_param_init: dict[str, Callable],
    w2_param_init: dict[str, Callable],
) -> FeedForward.Config:
    """Build a shared FFN whose output reduction is selected at runtime.

    ``w1`` and ``w3`` are the gate/up projections and share
    ``w13_param_init``; ``w2`` is the residual output projection.
    """
    return FeedForward.Config(
        w13=ColumnParallelLinear.Config(
            in_features=dim,
            out_features=hidden_dim,
            num_linears=2,
            param_init=fused_gate_up_param_init(w13_param_init, w13_param_init),
        ),
        w2=SharedExpertRowParallelLinear.Config(
            in_features=hidden_dim,
            out_features=dim,
            param_init=w2_param_init,
        ),
    )


def make_moe_config(
    *,
    num_experts: int = 8,
    router: TokenChoiceTopKRouter.Config,
    routed_experts: RoutedExperts.Config,
    shared_experts: FeedForward.Config | None = None,
    load_balance_coeff: float | None = 1e-3,
    aux_loss_coeff: float | None = None,
) -> MoE.Config:
    """Build a fully-specified MoE.Config."""
    if aux_loss_coeff is not None:
        router = dataclasses.replace(
            router,
            aux_loss=MicrobatchWiseLoadBalanceLoss.Config(coeff=aux_loss_coeff),
        )
    return MoE.Config(
        num_experts=num_experts,
        load_balance_coeff=load_balance_coeff,
        router=router,
        routed_experts=routed_experts,
        shared_experts=shared_experts,
    )


def make_router_config(
    *,
    dim: int,
    num_experts: int,
    gate_param_init: dict[str, Callable],
    score_func: UnaryActivationFn.Config,
    top_k: int = 1,
    route_norm: bool = False,
    route_norm_epsilon: float = 1e-20,
    route_scale: float = 1.0,
    bias: bool = False,
) -> TokenChoiceTopKRouter.Config:
    """Build a fully-specified TokenChoiceTopKRouter.Config."""
    return TokenChoiceTopKRouter.Config(
        num_experts=num_experts,
        gate=HiMidLoLinear.Config(
            in_features=dim,
            out_features=num_experts,
            backward_mode="hi_mid_lo",
            bias=bias,
            param_init=gate_param_init,
        ),
        top_k=top_k,
        score_func=score_func,
        route_norm=route_norm,
        route_norm_epsilon=route_norm_epsilon,
        route_scale=route_scale,
    )


def make_routed_experts_config(
    *,
    dim: int,
    hidden_dim: int,
    num_experts: int,
    top_k: int,
    param_init: dict[str, Callable],
) -> RoutedExperts.Config:
    """Build routed experts with structured gate/up and down projections."""
    missing = {"w1_EFD", "w2_EDF", "w3_EFD"} - param_init.keys()
    if param_init and missing:
        raise ValueError(f"Missing routed-expert initializers: {sorted(missing)}")

    return RoutedExperts.Config(
        w13=GroupedLinear.Config(
            group_size=num_experts,
            in_features=dim,
            out_features=hidden_dim,
            num_linears=2,
            param_init=(
                fused_grouped_gate_up_param_init(param_init) if param_init else None
            ),
        ),
        w2=GroupedLinear.Config(
            group_size=num_experts,
            in_features=hidden_dim,
            out_features=dim,
            param_init={"weight": param_init["w2_EDF"]} if param_init else None,
        ),
        token_dispatcher=AllToAllTokenDispatcher.Config(
            num_experts=num_experts,
            top_k=top_k,
        ),
    )
