# 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 dataclasses import dataclass
from typing import Any, cast

import spmd_types as spmd
import torch
from spmd_types import SpmdType

from torchtitan.config import TORCH_DTYPE_MAP, TrainingConfig
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed.parallelism_context import ParallelismContext
from torchtitan.distributed.spmd_types import annotate_input_spmd_types
from torchtitan.models.common.attention import (
    AttentionMetadata,
    AttentionMetadataMap,
    BaseAttention,
    InnerAttention,
)
from torchtitan.models.common.aux_loss import AuxLoss
from torchtitan.models.common.decoder_sharding import decoder_input_sharding
from torchtitan.models.common.embedding import Embedding
from torchtitan.models.common.feed_forward import FeedForward
from torchtitan.models.common.linear import Linear
from torchtitan.models.common.moe import MoE
from torchtitan.models.common.nn_modules import RMSNorm
from torchtitan.protocols.model import BaseModel
from torchtitan.protocols.module import Module, ModuleDict

__all__ = ["Decoder", "TransformerBlock"]


# TODO: we can unify the TransformerBlock impl across all models when
# there is no special logic for each model, including
# ffn vs. moe naming and creation, etc.
class TransformerBlock(Module):
    """Base class for all language model transformer blocks.

    All language model TransformerBlocks share:
    - Attention module (from ``attention.build()``)
    - FFN or MoE (from ``feed_forward.build()`` / ``moe.build()``)
    - Two RMSNorms (``attention_norm``, ``ffn_norm``)
    - Forward: ``x + attn(norm(x), ...); x + ffn(norm(x))``

    Forward accepts ``aux_loss_denominator``. Dense blocks ignore it; MoE
    blocks pass it to routers configured with an auxiliary loss.

    Children implement ``__init__`` and ``forward``.
    """

    attention: BaseAttention

    @dataclass(kw_only=True, slots=True)
    class Config(Module.Config):
        attention: BaseAttention.Config  # required, no default
        feed_forward: FeedForward.Config | None = None
        moe: MoE.Config | None = None
        attention_norm: RMSNorm.Config
        ffn_norm: RMSNorm.Config


class Decoder(BaseModel):
    """Base class for autoregressive decoder-only language models.

    Provides shared ``__init__``, ``forward``, ``init_states``, and
    ``_get_attention_metadata`` (flex/varlen dispatch) used by most models.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(BaseModel.Config):
        max_context_length: int
        dim: int
        vocab_size: int
        lm_head: Linear.Config
        tok_embeddings: Embedding.Config
        norm: RMSNorm.Config
        # TODO(fegin): revisit
        # https://github.com/pytorch/torchtitan/pull/2785#discussion_r3033849265
        # and fix the typing here
        layers: list  # list[TransformerBlock.Config] or subclass configs
        # Tie ``tok_embeddings`` and ``lm_head`` to share one weight. Models
        # that support it set this True in their config factories; the tying
        # itself is handled by ``Decoder.__init__`` / ``Decoder.init_states``.
        enable_weight_tying: bool = False

        @property
        def first_base_attention(self) -> BaseAttention.Config | None:
            """First ``BaseAttention`` config, else ``None``.

            Hybrid models do not carry a ``BaseAttention`` config on every
            layer, so callers needing its metadata look up the first such layer
            rather than assuming ``layers[0]``.
            """
            return next(
                (
                    layer.attention
                    for layer in self.layers
                    if layer.attention is not None
                ),
                None,
            )

        @property
        def base_attention_backends(self) -> tuple[Module.Config, ...]:
            """Inner backend configs for all ``BaseAttention`` layers."""
            return tuple(
                layer.attention.inner_attention
                for layer in self.layers
                if layer.attention is not None
            )

        @property
        def first_base_attention_backend(self) -> Module.Config | None:
            """Inner backend config of the first ``BaseAttention`` layer."""
            return next(iter(self.base_attention_backends), None)

        @property
        def first_feed_forward(self) -> FeedForward.Config | None:
            """First dense feed-forward config, else None."""
            return next(
                (
                    layer.feed_forward
                    for layer in self.layers
                    if layer.feed_forward is not None
                ),
                None,
            )

        @property
        def first_moe(self) -> MoE.Config | None:
            """First mixture-of-experts config, else None."""
            return next(
                (layer.moe for layer in self.layers if layer.moe is not None),
                None,
            )

    # Set by the trainer when ChunkedLossWrapper is used, so lm_head is applied
    # per-chunk inside the loss function instead of in forward().
    # TODO(#ISSUE): Remove after fixing PP backward to skip non-tensor
    # inputs (bool kwargs cause 'has no attribute requires_grad' errors).
    _skip_lm_head: bool = False

    def _apply_fsdp(
        self,
        *,
        parallelism_context: ParallelismContext,
        training: TrainingConfig,
        parallelism: ParallelismConfig,
    ) -> None:
        from torchtitan.distributed.fsdp import (
            apply_fsdp_to_decoder,
            resolve_fsdp_mesh,
            resolve_sparse_fsdp_mesh,
        )

        dp_mesh, dp_mesh_dims = resolve_fsdp_mesh(parallelism_context)
        edp_mesh, edp_mesh_dims = resolve_sparse_fsdp_mesh(parallelism_context)
        apply_fsdp_to_decoder(
            self,
            dp_mesh,
            param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param],
            reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce],
            pp_enabled=parallelism_context.pp_enabled,
            cpu_offload=training.enable_cpu_offload,
            reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward,
            ep_degree=parallelism_context.ep,
            edp_mesh=edp_mesh,
            dp_mesh_dims=dp_mesh_dims,
            edp_mesh_dims=edp_mesh_dims,
            symm_mem_scope=parallelism.fsdp_symm_mem_scope,
        )

    def __init__(self, config: Config):
        from torchtitan.distributed.spmd_types import spmd_mesh_size, spmd_sparse_mesh

        tp = spmd_mesh_size("tp")
        attention = config.first_base_attention
        if tp > 1 and attention is not None:
            num_heads = attention.n_heads
            num_kv_heads = getattr(attention, "n_kv_heads", None) or num_heads
            if num_heads % tp != 0:
                raise ValueError(
                    f"tensor parallel degree ({tp}) must divide "
                    f"n_heads ({num_heads})."
                )
            # Fused QKV projections shard whole KV-head groups. Attention with
            # separate K/V projections may instead shard each head's features.
            if hasattr(attention, "qkv_linear") and num_kv_heads % tp != 0:
                raise ValueError(
                    f"tensor parallel degree ({tp}) must divide "
                    f"n_kv_heads ({num_kv_heads})."
                )

        sparse_mesh = spmd_sparse_mesh()
        ep = sparse_mesh["ep"].size() if sparse_mesh is not None else 1
        moe_configs = list(config.traverse(MoE.Config))
        if moe_configs and ep < tp:
            raise ValueError(
                f"MoE models require expert parallel degree ({ep}) to be "
                f"greater than or equal to tensor parallel degree ({tp})."
            )
        for moe_fqn, moe, _, _ in moe_configs:
            if moe.num_experts % ep != 0:
                raise ValueError(
                    f"{moe_fqn}.num_experts ({moe.num_experts}) must be "
                    f"divisible by expert parallel degree ({ep})."
                )

        super().__init__()
        self.config = config

        self.tok_embeddings = config.tok_embeddings.build()

        self.layers = ModuleDict()
        for i, layer_config in enumerate(config.layers):
            self.layers[str(i)] = layer_config.build()

        self.norm = config.norm.build()
        self.lm_head = config.lm_head.build()

        self.enable_weight_tying = config.enable_weight_tying
        if self.enable_weight_tying:
            self.tok_embeddings.weight = self.lm_head.weight

    def init_states(
        self,
        *,
        buffer_device: torch.device | None = None,
    ) -> None:
        if self.enable_weight_tying:
            # Re-tie before init: on meta device the ``__init__`` tying may not
            # have taken effect, and ``tok_embeddings.weight`` is skipped by
            # ``skip_param_init``, so re-point it at the initialized lm_head
            # weight.
            assert self.tok_embeddings is not None and self.lm_head is not None
            self.tok_embeddings.weight = self.lm_head.weight
        super().init_states(buffer_device=buffer_device)

    def forward(
        self,
        tokens: torch.Tensor,
        positions: torch.Tensor | None = None,
        attention_metadata: AttentionMetadataMap | None = None,
        *,
        padding_mask: torch.Tensor | None = None,
        aux_loss_denominators: torch.Tensor | None = None,
    ):
        # positions is listed before attention_metadata so AutoParallel's input_fn,
        # which returns (tokens, positions) and binds them positionally, maps
        # positions to the right parameter (it would otherwise land in the
        # attention_metadata slot and break the maskless SDPA backend).
        # passthrough for nonexistent layers, allows easy configuration of pipeline parallel stages
        h = self.tok_embeddings(tokens) if self.tok_embeddings is not None else tokens

        with spmd.no_typecheck():
            aux_loss_denominator = (
                None if aux_loss_denominators is None else aux_loss_denominators[0]
            )
        for layer in self.layers.values():
            layer_attention_metadata = (
                None
                if attention_metadata is None
                else attention_metadata.get(
                    cast(TransformerBlock, layer).attention.attention_metadata_key
                )
            )
            h = layer(
                h,
                layer_attention_metadata,
                positions,
                padding_mask=padding_mask,
                aux_loss_denominator=aux_loss_denominator,
            )

        h = self.norm(h) if self.norm is not None else h

        # _skip_lm_head is an attribute rather than a forward kwarg because PP backward
        # calls .requires_grad on all stage inputs, which fails on bool kwargs.
        # TODO: fix PP backward upstream to skip non-tensor inputs
        if self._skip_lm_head:
            return h
        output = self.lm_head(h) if self.lm_head is not None else h
        return output

    def preprocess_inputs(
        self,
        input_dict: dict[str, Any],
        *,
        parallelism_context: ParallelismContext,
        parallelism: ParallelismConfig,
        max_num_documents: int | None = None,
        max_context_length: int | None = None,
        **kwargs: Any,
    ) -> tuple[
        torch.Tensor | tuple[torch.Tensor, ...],
        torch.Tensor | tuple[torch.Tensor, ...],
        dict[str, Any],
    ]:
        """Build masks (flex/varlen), CP-shard, SPMD-wrap, and return the batch."""
        del kwargs
        positions = input_dict.get("positions", None)
        padding_mask = input_dict.get("padding_mask", None)
        if positions is not None:
            attention_metadata = self._get_attention_metadata(
                positions=positions,
                padding_mask=padding_mask,
                max_num_documents=max_num_documents,
                max_context_length=max_context_length,
            )
            input_dict["attention_metadata"] = attention_metadata

        input_shardings = decoder_input_sharding()
        if parallelism_context.cp_enabled:
            input_dict = self._cp_shard(
                input_dict,
                input_shardings=input_shardings,
                parallelism_context=parallelism_context,
                parallelism=parallelism,
            )
        input_dict = annotate_input_spmd_types(
            parallelism_context, input_dict, input_shardings
        )

        inputs = input_dict.pop("input")
        labels = input_dict.pop("labels")
        if next(self.config.traverse(AuxLoss.Config), None) is not None:
            input_dict["aux_loss_denominators"] = None
        return inputs, labels, input_dict

    def _cp_shard(
        self,
        input_dict: dict[str, Any],
        input_shardings: dict[str, SpmdType],
        parallelism_context: ParallelismContext,
        parallelism: ParallelismConfig,
    ) -> dict[str, Any]:
        """Prepare attention metadata and shard model inputs for CP."""
        from torchtitan.distributed import context_parallel
        from torchtitan.models.common.attention.cp_attention import (
            canonicalize_cp_inner_attention,
            CPInnerAttention,
        )

        attention_metadata = input_dict.get("attention_metadata")
        load_balancer_config = parallelism.context_parallel_load_balancer
        selected_attention_metadata = None
        if load_balancer_config is not None and attention_metadata is not None:
            first_base_attention = self.config.first_base_attention
            assert first_base_attention is not None
            first_cp_inner_attention = first_base_attention.inner_attention._owner
            assert first_cp_inner_attention is not None and issubclass(
                first_cp_inner_attention, CPInnerAttention
            )
            # One permutation is shared across layers. Prefer full-attention
            # metadata, falling back to the first quadratic inner attention
            # for models containing only sliding-window attention.
            full_cp_inner_attention = canonicalize_cp_inner_attention(
                first_cp_inner_attention
            )
            selected_attention_metadata = attention_metadata.get(
                full_cp_inner_attention
            )
            if selected_attention_metadata is None:
                selected_attention_metadata = attention_metadata.get(
                    first_cp_inner_attention
                )
        load_balancer = (
            # TODO: `attention_metadata` alone cannot determine the load-balancing strategy.
            load_balancer_config.build(
                seq_len=context_parallel.get_cp_input_seq_len(
                    input_dict, input_shardings=input_shardings
                ),
                attention_metadata=selected_attention_metadata,
            )
            if load_balancer_config is not None
            else None
        )
        permutation = (
            load_balancer.generate_permutation() if load_balancer is not None else None
        )
        if "attention_metadata" in input_dict:
            attention_metadata = input_dict["attention_metadata"]
            assert isinstance(attention_metadata, dict)
            for inner_attention, metadata in attention_metadata.items():
                if not issubclass(inner_attention, CPInnerAttention):
                    continue
                attention_metadata[
                    inner_attention
                ] = inner_attention.prepare_cp_metadata(
                    metadata,
                    permutation=permutation,
                )
        return context_parallel.shard_tensors(
            input_dict,
            input_shardings=input_shardings,
            permutation=permutation,
        )

    def _get_attention_metadata(
        self,
        positions: torch.Tensor,
        *,
        padding_mask: torch.Tensor | None = None,
        max_num_documents: int | None = None,
        max_context_length: int | None = None,
    ) -> AttentionMetadataMap:
        attention_metadata: dict[type[InnerAttention], AttentionMetadata] = {}
        for layer_config in self.config.layers:
            for _, config, _, _ in layer_config.traverse(InnerAttention.Config):
                backend = config._owner
                assert backend is not None and issubclass(backend, InnerAttention)
                if backend in attention_metadata:
                    continue
                metadata = config.build_attention_metadata(
                    positions,
                    padding_mask=padding_mask,
                    max_num_documents=max_num_documents,
                    max_context_length=max_context_length,
                )
                if metadata is not None:
                    attention_metadata[backend] = metadata
        return attention_metadata
