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

import fnmatch
import logging
import time
from collections.abc import Callable, Iterable
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Any, TypeAlias

import torch
import torch.nn as nn
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.tensor import DTensor, Replicate
from torch.fx.traceback import annotate, annotate_fn
from torch.utils._pytree import register_constant, register_pytree_node, tree_map

from torchtitan.config import TORCH_DTYPE_MAP, TrainingConfig
from torchtitan.distributed import ParallelismContext
from torchtitan.experiments.graph_trainer.simple_fsdp import (
    data_parallel,
    MixedPrecisionPolicy,
)
from torchtitan.models.common.attention import ScaledDotProductInnerAttention
from torchtitan.models.common.decoder import Decoder, TransformerBlock


logger = logging.getLogger(__name__)


BOXED_CODEGEN_META = "graph_trainer_boxed_codegen"
LossResult: TypeAlias = torch.Tensor | tuple[torch.Tensor, dict[str, Any]]
AnnotatedLossFn: TypeAlias = Callable[..., LossResult]


def _get_graph_modules(
    gm: torch.fx.GraphModule,
    *,
    recurse: bool,
    apply_to_root: bool,
) -> list[torch.fx.GraphModule]:
    modules = [gm] if apply_to_root else []
    if recurse:
        modules.extend(
            module
            for name, module in gm.named_modules()
            if name and isinstance(module, torch.fx.GraphModule)
        )
    return modules


class GraphTrainerScaledDotProductInnerAttention(ScaledDotProductInnerAttention):
    """Adapt flat graph-trainer attention inputs to the batched SDPA interface."""

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

    def forward(
        self,
        q_THK: torch.Tensor,
        k_THK: torch.Tensor,
        v_THV: torch.Tensor,
        **kwargs,
    ) -> torch.Tensor:
        out_1THV = super().forward(
            q_THK.unsqueeze(0),
            k_THK.unsqueeze(0),
            v_THV.unsqueeze(0),
            **kwargs,
        )
        return out_1THV.squeeze(0)


@contextmanager
def log_timer(label: str):
    start = time.perf_counter()
    yield
    elapsed_s = time.perf_counter() - start
    logger.info("%s took %.3fs", label, elapsed_s)


def build_decoder_config_for_backend(
    config_builder: Callable, attn_backend: str, **builder_kwargs
):
    """Build a Decoder model config for ``attn_backend``, allowing test-only SDPA.

    ``SDPA`` is not a valid production language-model backend — ``get_attention_config``
    rejects it because the dataloaders always emit per-document positions and SDPA
    cannot consume them (it only has a boolean ``is_causal``). The graph_trainer
    tests, however, use SDPA to exercise *backend-agnostic* graph machinery
    (precompile-artifact serialization, custom codegen, context parallel, bitwise
    determinism) without FlexInnerAttention's ``BlockMask``, which is unpicklable (its
    ``mask_mod`` closures are Python code objects), is not a tensor (so it breaks
    pipeline-parallel split-backward, which calls ``.requires_grad`` on every stage
    input), and overflows the fp32 Triton shared-memory limit on large head dims.

    For SDPA we build the flex config (a valid backend) and swap each layer's
    ``inner_attention`` to ``GraphTrainerScaledDotProductInnerAttention.Config()``.
    The adapter adds a singleton batch around the flat graph-trainer inputs and
    delegates to the common batched SDPA implementation. Production code never
    reaches this path: ``get_attention_config`` still rejects ``sdpa``, so no model
    registry can construct an SDPA language model outside these tests.
    """
    if attn_backend != "sdpa":
        return config_builder(attn_backend=attn_backend, **builder_kwargs)

    config = config_builder(attn_backend="flex", **builder_kwargs)
    for layer in config.layers:
        layer.attention.inner_attention = (
            GraphTrainerScaledDotProductInnerAttention.Config()
        )
    return config


def _local_stride(tensor: torch.Tensor) -> tuple[int, ...]:
    return (
        tensor.to_local().stride() if isinstance(tensor, DTensor) else tensor.stride()
    )


def _maybe_materialize_grad_for_param_layout(
    param: torch.Tensor, grad: torch.Tensor
) -> torch.Tensor:
    """Match eager autograd's ``param.grad`` layout contract after graph replay.

    Graph replay assigns ``torch.autograd.grad`` outputs manually, bypassing
    AccumulateGrad's normal stride materialization. Copying through
    ``empty_like(param)`` restores the param's global and DTensor-local layout
    when Inductor returns an equivalent but differently-strided grad.
    """
    if grad.stride() == param.stride() and _local_stride(grad) == _local_stride(param):
        return grad

    materialized_grad = torch.empty_like(param)
    materialized_grad.copy_(grad)
    return materialized_grad


def _is_same_tensor_view(lhs: torch.Tensor, rhs: torch.Tensor) -> bool:
    if lhs is rhs:
        return True
    if isinstance(lhs, DTensor):
        if not isinstance(rhs, DTensor):
            return False
        if lhs.device_mesh != rhs.device_mesh or lhs.placements != rhs.placements:
            return False
        lhs = lhs._local_tensor
        rhs = rhs._local_tensor
    elif isinstance(rhs, DTensor):
        return False
    return (
        lhs.shape == rhs.shape
        and lhs.stride() == rhs.stride()
        and lhs.storage_offset() == rhs.storage_offset()
        # pyrefly: ignore [missing-attribute]
        and torch._C._is_alias_of(lhs, rhs)
    )


def set_graph_module_boxed_codegen(
    gm: torch.fx.GraphModule,
    *,
    boxed: bool,
) -> torch.fx.GraphModule:
    """Set the FX calling convention and record it in graph metadata.

    Args:
        gm: Graph module whose generated Python wrapper should be updated.
        boxed: Whether the wrapper should accept one mutable argument list and
            clear it after placeholder extraction.

    Returns:
        The same graph module after recompilation if the calling convention
        changed.
    """

    if gm.meta.get(BOXED_CODEGEN_META) is boxed:
        return gm
    codegen = torch.fx.graph._BoxedCodeGen() if boxed else torch.fx.graph.CodeGen()
    gm.graph.set_codegen(codegen)
    gm.recompile()
    gm.meta[BOXED_CODEGEN_META] = boxed
    return gm


def ensure_boxed_graph_module(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
    """Use boxed FX codegen for runtime-owned graph modules.

    Args:
        gm: Graph module to box.

    Returns:
        The same graph module with boxed FX codegen enabled.
    """

    return set_graph_module_boxed_codegen(gm, boxed=True)


_MODULE_FQN = "module_fqn"
# Tuple of unique parameter FQNs naming one gradient value. Consumers must
# treat it as a set: tied parameters can associate multiple FQNs with one node.
PARAMETER_GRADIENT_FQNS_META = "parameter_gradient_fqns"
_EP_TOKEN_COUNT_EXCHANGE = "EP_token_count_exchange"
_EP_TOKEN_COUNT_SYNC = "EP_token_count_sync"
_EP_TOKEN_EXCHANGE = "EP_token_exchange"
_EP_TOKEN_EXCHANGE_WAIT = "EP_token_exchange_wait"
_NOT_IN_LAYERS = -1


def compute_parameter_gradients(
    loss: torch.Tensor,
    named_parameters: Iterable[tuple[str, torch.Tensor]],
) -> tuple[torch.Tensor, ...]:
    """Compute and identify parameter gradients in the traced graph.

    Each gradient is passed through an annotated ``aten.alias`` immediately
    after ``torch.autograd.grad`` establishes its parameter correspondence.
    The alias is a view of the same storage, not a copy or device computation,
    and keeps the identity available even when an optimizer consumes the
    gradient inside the graph instead of returning it as a graph output.
    GraphTrainer moves the FQN set onto the underlying gradient node and
    erases the markers before later graph passes and backend compilation.
    """
    named_parameters = tuple(named_parameters)
    gradients = torch.autograd.grad(
        loss, tuple(parameter for _, parameter in named_parameters)
    )
    return tuple(
        annotate_parameter_gradient(gradient, parameter_fqn)
        for (parameter_fqn, _), gradient in zip(
            named_parameters, gradients, strict=True
        )
    )


def annotate_parameter_gradient(
    gradient: torch.Tensor,
    parameter_fqn: str,
) -> torch.Tensor:
    """Attach a parameter identity to one gradient in the traced graph."""
    with annotate({PARAMETER_GRADIENT_FQNS_META: (parameter_fqn,)}):
        return torch.ops.aten.alias.default(gradient)


def compute_annotated_loss(
    loss_fn: AnnotatedLossFn,
    pred: torch.Tensor,
    labels: torch.Tensor,
    loss_kwargs: dict[str, Any] | None = None,
) -> torch.Tensor:
    """Compute the loss tensor with the same FX metadata convention as GraphTrainer."""
    annotated_loss_fn = annotate_fn({_MODULE_FQN: "loss"})(loss_fn)
    result = annotated_loss_fn(pred, labels, **(loss_kwargs or {}))
    if isinstance(result, tuple):
        if len(result) != 2:
            raise ValueError(
                "GraphTrainer loss functions must return a loss tensor or "
                "(loss tensor, metrics)."
            )
        loss, _metrics = result
        return loss
    return result


def accumulate_param_grads_(
    params: Iterable[torch.Tensor],
    grads: Iterable[torch.Tensor | None],
    *,
    clone_grads_to_initialize_param_grad: bool = False,
) -> None:
    """Accumulate explicit graph-produced gradients into live parameters.

    Args:
        params: Parameters that receive the explicit gradients.
        grads: Gradients returned by the traced forward-backward graph.
        clone_grads_to_initialize_param_grad: Whether to initialize empty
            ``param.grad`` fields with clones instead of input aliases.
    """
    for param, grad in zip(params, grads, strict=True):
        if grad is None:
            continue
        grad = _maybe_materialize_grad_for_param_layout(param, grad)
        if param.grad is None:
            param.grad = grad.clone() if clone_grads_to_initialize_param_grad else grad
        # The graph-owned buffer may already be the optimizer-visible gradient.
        elif _is_same_tensor_view(param.grad, grad):
            continue
        else:
            param.grad += grad


def _is_backward_node(node: torch.fx.Node) -> bool:
    return node.meta.get("autograd_backward", False)


def _get_module_fqn(node: torch.fx.Node) -> str:
    return node.meta.get("custom", {}).get(_MODULE_FQN, "")


def _get_layer_id(node: torch.fx.Node) -> int:
    """Extract the layer index from the node's module_fqn metadata.

    Nodes under ``layers.<N>`` return ``N``.
    All other nodes (tok_embeddings, norm, output) return ``_NOT_IN_LAYERS``.
    """
    fqn = _get_module_fqn(node)
    parts = fqn.split(".")
    if parts[0] == "layers" and len(parts) >= 2:
        try:
            return int(parts[1])
        except ValueError:
            pass
    return _NOT_IN_LAYERS


def annotate_module_fqns(model: nn.Module) -> None:
    """Annotate all modules' forward with their fully-qualified names.

    Every named submodule (excluding the root) gets its forward method wrapped
    with ``annotate_fn`` so that FX nodes carry ``module_fqn`` in
    ``node.meta["custom"]``.

    Call once after model construction, before tracing/compilation.
    """
    for fqn, submodule in model.named_modules():
        if fqn:  # skip root module
            submodule.forward = annotate_fn({_MODULE_FQN: fqn})(submodule.forward)


def annotate_graph_trainer_model(model: Decoder) -> None:
    """Attach the annotations consumed by GraphTrainer passes."""
    if any(getattr(layer, "moe", None) is not None for layer in model.config.layers):
        annotate_moe_ep_regions()
    annotate_module_fqns(model)


def matches_module_fqn_pattern(pattern: str, fqn: str) -> bool:
    """Match one module FQN against a component-wise fnmatch pattern."""
    pattern_parts = pattern.split(".")
    fqn_parts = fqn.split(".")
    return len(pattern_parts) == len(fqn_parts) and all(
        fnmatch.fnmatchcase(fqn_part, pattern_part)
        for pattern_part, fqn_part in zip(pattern_parts, fqn_parts)
    )


_MOE_EP_REGIONS_ANNOTATED = False


def annotate_moe_ep_regions() -> None:
    """Annotate MoE EP compute, dispatch, and combine regions for FX passes."""
    global _MOE_EP_REGIONS_ANNOTATED
    if _MOE_EP_REGIONS_ANNOTATED:
        return

    from torchtitan.models.common.moe import MoE
    from torchtitan.models.common.token_dispatcher import (
        AllToAllTokenDispatcher,
        LocalTokenDispatcher,
    )

    LocalTokenDispatcher.dispatch = annotate_fn({"EP": "dispatch"})(
        LocalTokenDispatcher.dispatch
    )
    LocalTokenDispatcher.combine = annotate_fn({"EP": "combine"})(
        LocalTokenDispatcher.combine
    )
    AllToAllTokenDispatcher.dispatch = annotate_fn({"EP": "dispatch"})(
        AllToAllTokenDispatcher.dispatch
    )
    AllToAllTokenDispatcher.combine = annotate_fn({"EP": "combine"})(
        AllToAllTokenDispatcher.combine
    )
    AllToAllTokenDispatcher._token_count_exchange = annotate_fn(
        {_EP_TOKEN_COUNT_EXCHANGE: "dispatch"}
    )(AllToAllTokenDispatcher._token_count_exchange)
    AllToAllTokenDispatcher._sync_token_count_exchange = annotate_fn(
        {_EP_TOKEN_COUNT_SYNC: "dispatch"}
    )(AllToAllTokenDispatcher._sync_token_count_exchange)
    AllToAllTokenDispatcher._dispatch_token_exchange = annotate_fn(
        {_EP_TOKEN_EXCHANGE: "dispatch"}
    )(AllToAllTokenDispatcher._dispatch_token_exchange)
    AllToAllTokenDispatcher._combine_token_exchange = annotate_fn(
        {_EP_TOKEN_EXCHANGE: "combine"}
    )(AllToAllTokenDispatcher._combine_token_exchange)
    MoE.forward = annotate_fn({"EP": "compute"})(MoE.forward)
    _MOE_EP_REGIONS_ANNOTATED = True


def parallelize_inputs(parallelism_context, args, kwargs):
    if not parallelism_context.tp_enabled:
        return args, kwargs

    def to_dtensor(tensor):
        if isinstance(tensor, torch.Tensor):
            return DTensor.from_local(
                tensor, parallelism_context.get_mesh("tp"), [Replicate()]
            )
        return tensor

    dt_args = tree_map(to_dtensor, args)

    # TODO: When using flex_attention, BlockMask would show up in kwargs,
    # and it's unclear how to convert it to DTensor. If I use to_dtensor,
    # it would fail with Dynamo Error: P2011360347
    # dt_kwargs = tree_map(to_dtensor, kwargs)

    dt_kwargs = kwargs

    return dt_args, dt_kwargs


def register_blockmask_pytree_node():
    from torch.nn.attention.flex_attention import BlockMask

    if BlockMask not in torch.utils._pytree.SUPPORTED_NODES:
        register_pytree_node(
            BlockMask,
            BlockMask._flatten,
            BlockMask._unflatten,
            flatten_with_keys_fn=BlockMask._flatten_with_keys,
            serialized_type_name="torch.nn.attention.flex_attention.BlockMask",
        )


def maybe_register_blockmask_pytree_node() -> None:
    """Register flex-attention pytree helpers if they are missing."""
    from torch.nn.attention.flex_attention import _MaskModWrapper, BlockMask

    if BlockMask not in torch.utils._pytree.SUPPORTED_NODES:
        register_blockmask_pytree_node()
    if _MaskModWrapper not in torch.utils._pytree.SUPPORTED_NODES:
        register_constant(_MaskModWrapper)


def end_with_pass(passes: list[Callable], names: list[str]) -> bool:
    return (
        len(passes) > 0
        and (last_pass_name := getattr(passes[-1], "__name__", None))
        and (last_pass_name in names)
    )


def get_default_transformer_block_buckets(
    n_layers: int,
    *,
    chunked_loss_enabled: bool = False,
    moe_layer_ids: frozenset[int] = frozenset(),
    split_moe_expert_buckets: bool = False,
) -> list[list[str] | str]:
    """Get default transformer block buckets for manual bucketing passes.

    Assumes the standard Decoder layout: tok_embeddings, layers.0..N-1,
    norm, and output (e.g., Llama3, DeepSeekV3, Qwen3).
    """
    layer_buckets: list[list[str] | str] = []
    for layer_id in range(n_layers):
        if layer_id in moe_layer_ids and split_moe_expert_buckets:
            layer_buckets.extend(
                [
                    [
                        f"layers.{layer_id}.attention_norm",
                        f"layers.{layer_id}.attention",
                        f"layers.{layer_id}.ffn_norm",
                        f"layers.{layer_id}.moe.router",
                        f"layers.{layer_id}.moe.shared_experts",
                    ],
                    f"layers.{layer_id}.moe.routed_experts",
                ]
            )
        else:
            layer_buckets.append(f"layers.{layer_id}")
    final_bucket = ["norm", "lm_head"]
    if chunked_loss_enabled:
        # Chunked loss moves the lm_head weight use under module_fqn "loss".
        final_bucket.append("loss")

    return [
        "tok_embeddings",
        *layer_buckets,
        final_bucket,
    ]


def get_transformer_block_buckets(model) -> list[list[str] | str]:
    """Get transformer block buckets for manual bucketing passes.

    Works for any model with tok_embeddings, layers (OrderedDict), norm, and output
    attributes (e.g., Llama3, DeepSeekV3).
    """
    # [TODO](ruisizhang123) add EP support for transformer block bucketing
    module_list = [
        model.tok_embeddings,
        [model.norm, model.lm_head],
    ]
    for layer_id, transformer_block in model.layers.items():
        module_list.append(transformer_block)

    def convert_modules_to_fqns(modules, module_to_fqn_mapping):
        """Convert a (possibly nested) list of modules to FQN strings."""
        result = []
        for m in modules:
            if isinstance(m, list):
                if fqn_list := convert_modules_to_fqns(m, module_to_fqn_mapping):
                    result.append(fqn_list)
            else:
                if fqn := module_to_fqn_mapping.get(m):
                    result.append(fqn)
        return result

    module_to_name = {m: n for n, m in model.named_modules()}
    module_fqns = convert_modules_to_fqns(module_list, module_to_name)
    return module_fqns


def get_simple_fsdp_mesh(parallelism_context: ParallelismContext) -> DeviceMesh:
    """Return the flattened DP-shard/CP mesh used by SimpleFSDP."""
    fsdp_mesh = parallelism_context.get_optional_mesh(
        ["dp_shard", "cp"], include_singleton_axes=True
    )
    assert fsdp_mesh is not None
    return fsdp_mesh._flatten("fsdp")


def apply_simple_fsdp(
    model: nn.Module,
    *,
    parallelism_context: ParallelismContext,
    training: TrainingConfig,
) -> nn.Module:
    """Wrap the model (and any MoE experts) with graph_trainer's simple_fsdp.

    For MoE-enabled models, routed W13 and W2 projections are separately
    wrapped on the EDP mesh when expert parallelism is enabled.
    """
    fsdp_mesh = get_simple_fsdp_mesh(parallelism_context)

    if parallelism_context.dp_replicate_enabled:
        if parallelism_context.dp_shard_enabled or parallelism_context.cp_enabled:
            dp_replicate_mesh = parallelism_context.get_optional_mesh(
                "dp_replicate", include_singleton_axes=True
            )
            assert dp_replicate_mesh is not None
            dp_mesh = DeviceMesh._concatenate([dp_replicate_mesh, fsdp_mesh])
            dp_mode = "hybrid_shard"
        else:
            dp_mesh = parallelism_context.get_mesh("dp_replicate")
            dp_mode = "replicate"
    else:
        dp_mesh = fsdp_mesh
        dp_mode = "fully_shard"

    mp_policy = MixedPrecisionPolicy(
        param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param],
        reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce],
    )

    if parallelism_context.ep_enabled and isinstance(model, Decoder):
        edp_mesh_names = (
            ["dp_replicate", "edp_shard"]
            if parallelism_context.dp_replicate_enabled
            else ["edp_shard"]
        )
        edp_mesh = parallelism_context.get_optional_mesh(edp_mesh_names)
        assert edp_mesh is not None

        for _, transformer_block in model.layers.items():
            if not isinstance(transformer_block, TransformerBlock):
                continue
            moe = getattr(transformer_block, "moe", None)
            if moe is None:
                continue
            routed_experts = moe.routed_experts
            experts_shard_dim = 0
            if (
                edp_mesh["edp_shard"].size() * parallelism_context.ep
                > routed_experts.w13.group_size
            ):
                experts_shard_dim = 1

            if experts_shard_dim == 0:
                data_parallel(
                    routed_experts,
                    edp_mesh,
                    dp_mode,
                    mp_policy=mp_policy,
                    shard_dim=0,
                    non_dp_mesh=parallelism_context.get_optional_mesh("ep"),
                    # Match core FSDP: every routed-expert parameter, including
                    # stacked W13, shards the expert dimension.
                    param_shard_placements={},
                )
            else:
                data_parallel(
                    routed_experts.w13,
                    edp_mesh,
                    dp_mode,
                    mp_policy=mp_policy,
                    shard_dim=2,
                    non_dp_mesh=parallelism_context.get_optional_mesh("ep"),
                )
                data_parallel(
                    routed_experts.w2,
                    edp_mesh,
                    dp_mode,
                    mp_policy=mp_policy,
                    shard_dim=1,
                    non_dp_mesh=parallelism_context.get_optional_mesh("ep"),
                )

    model = data_parallel(
        model,
        dp_mesh,
        dp_mode,
        mp_policy=mp_policy,
        non_dp_mesh=parallelism_context.get_optional_mesh("tp"),
    )
    logger.info(
        "Applied Data Parallel (simple_fsdp) (dp mode=%s) to the model", dp_mode
    )
    return model
