# 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 logging
from typing import Any

import torch
import torch.nn as nn
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.fsdp import (
    CPUOffloadPolicy,
    DataParallelMeshDims,
    fully_shard,
    MixedPrecisionPolicy,
)
from torch.distributed.tensor import Shard

from torchtitan.config import TrainingConfig
from torchtitan.config.parallelism import FSDPSymmMemScope, ParallelismConfig
from torchtitan.distributed import ParallelismContext
from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig
from torchtitan.distributed.fsdp import (
    disable_fsdp_gradient_division,
    enable_fsdp_symm_mem,
    get_fsdp_reshard_after_forward_policy,
)
from torchtitan.distributed.local_compile import apply_local_compile


logger = logging.getLogger(__name__)


def _wrap_flex_kernel_cp(model: nn.Module, cp_mesh: DeviceMesh) -> None:
    """All-gather k/v across the CP axis inside each flex kernel forward.

    q/k/v reach the flex kernel seq-sharded on the CP axis (tensor dim 2 of the
    ``(b, heads, seq, dim)`` layout). Attention needs full-length k/v, so we
    all-gather them across the CP mesh with a funcol collective (autograd-aware:
    the backward reduce-scatters gradients back to the local shard). q stays
    sharded, so each rank computes attention for its seq shard against the full
    keys -- the BlockMask is Q-sharded / KV-full to match.

    This is the explicit-collective analogue of Titan's ``flex_cp_allgather``
    path: the kernel runs nested inside the attention module's local SPMD region
    where the CP mesh dim is no longer visible to a declarative redistribute, so
    the gather is done here on local tensors. Called before ``model.parallelize``
    so the wrap is captured inside the local SPMD wrapper.
    """
    import torch.distributed as dist
    from torch.distributed.tensor.experimental._context_parallel._attention import (
        flex_cp_allgather,
    )

    pg_name = dist._get_process_group_name(cp_mesh.get_group())
    seq_dim = 2  # (b, heads, seq, dim)

    for module in model.modules():
        kernel = getattr(module, "_titan_flex_kernel", None)
        if kernel is None:
            continue

        def _make_cp_forward(orig_forward):
            def cp_forward(query, key, value, **kwargs):
                key = key.contiguous()
                value = value.contiguous()
                global_key, global_value = flex_cp_allgather(
                    key, value, seq_dim, pg_name
                )
                return orig_forward(query, global_key, global_value, **kwargs)

            return cp_forward

        kernel.forward = _make_cp_forward(kernel.forward)


# ---------------------------------------------------------------------------
# Main parallelization entry point
# ---------------------------------------------------------------------------


def parallelize_hf_transformers(
    model: nn.Module,
    *,
    parallelism_context: ParallelismContext,
    training: TrainingConfig,
    parallelism: ParallelismConfig,
    local_compile_regions: list[str],
    ac_config: ActivationCheckpointingConfig,
    dump_folder: str,
    **kwargs: Any,
):
    """Apply parallelism to the HF model using the titan Module protocol.

    Flow:
    1. Build and swap Titan MoE modules (sets _sharding_config on MoE tree)
    2. Convert all remaining HF nn.Modules to Module protocol via __class__ swap
    3. Set ShardingConfig on every module based on its role
    4. Single model._parallelize(parallelism_context) call -- shards states, wraps forward
    5. Apply AC and FSDP
    """
    # Bind local implementations early; torch.compile traces on first use.
    apply_local_compile(local_compile_regions)
    # Flex attention supports FSDP, TP, CP, and PP (in any combination). Under CP
    # the flex kernel's local SPMD boundary redistributes
    # k/v from seq-sharded to CP-Replicate (all-gather); see _attach_flex_kernel
    # in hf_sharding.py. The CP-sharded BlockMask is built and sharded on its Q
    # axis upstream (trainer, ptrr balancer). Note: the ptrr balancer requires
    # the number of Q blocks (seq_len / flex BLOCK_SIZE) to be divisible by the
    # CP degree; too-short sequences raise "num_tasks N must be divisible by
    # group_size" from the balancer -- this is a CP+ptrr constraint, independent
    # of PP.

    # 0. Un-tie embedding/lm_head weights for FSDP compatibility.
    # Some models (Gemma4) share the embedding and lm_head weight
    # (tie_word_embeddings=True). FSDP2 cannot handle parameters shared
    # across FSDP groups. Un-tying here creates an independent copy for
    # the lm_head so both can be sharded separately. Since we train from
    # scratch, the un-tied weights will diverge during training (which is
    # expected — the model learns independent embedding and output layers).
    if (
        model.tok_embeddings is not None
        and model.lm_head is not None
        and any(
            p1 is p2
            for p1 in model.tok_embeddings.parameters()
            for p2 in model.lm_head.parameters()
        )
    ):
        model.lm_head.weight = nn.Parameter(
            model.lm_head.weight.clone(),
            requires_grad=model.lm_head.weight.requires_grad,
        )
        logger.info("Un-tied embedding/lm_head weights for FSDP compatibility")

    # 1. Build and swap Titan MoE (sets _sharding_config, does NOT parallelize)
    if any(getattr(b, "moe_enabled", False) for b in model.layers):
        from torchtitan.experiments.transformers_modeling_backend.moe_replacement import (
            build_and_swap_native_moe,
        )

        build_and_swap_native_moe(model, parallelism_context)

    # 2. Convert HF modules to Module protocol.
    # The spmd_types backend uses the Module protocol for state distribution,
    # activation checks, and any TP/EP/CP redistribution.
    from torchtitan.experiments.transformers_modeling_backend.hf_sharding import (
        set_hf_sharding_configs,
    )
    from torchtitan.experiments.transformers_modeling_backend.module_conversion import (
        convert_hf_to_module,
    )

    convert_hf_to_module(model)

    # 3. Set sharding configs on all non-MoE modules
    set_hf_sharding_configs(
        model,
        enable_sp=parallelism_context.tp_enabled,
    )

    # 3b. Under CP, wrap each flex kernel forward to all-gather k/v across
    # the CP axis (on the seq dim). Must run before model._parallelize so the
    # wrap is captured inside the local SPMD region and operates on the local
    # (already TP-head-sharded, CP-seq-sharded) tensors.
    if parallelism_context.cp_enabled:
        _wrap_flex_kernel_cp(model, parallelism_context.get_mesh("cp"))

    # 4. Single parallelize call -- handles TP, EP, MoE, everything
    model._parallelize(parallelism_context)

    if ac_config is not None:
        ac_config.build(dump_folder=dump_folder).apply(model)

    model._apply_fsdp(
        parallelism_context=parallelism_context,
        training=training,
        parallelism=parallelism,
    )

    if training.enable_cpu_offload:
        logger.info("Applied CPU Offloading to the model")

    if parallelism_context.cp_enabled:
        model.set_cp_mesh(parallelism_context.get_mesh("cp"))
        logger.info("Applied Context Parallel to the model")

    return model


# ---------------------------------------------------------------------------
# FSDP
# ---------------------------------------------------------------------------


def apply_fsdp(
    model: nn.Module,
    dp_mesh: DeviceMesh,
    param_dtype: torch.dtype,
    reduce_dtype: torch.dtype,
    pp_enabled: bool,
    cpu_offload: bool = False,
    reshard_after_forward_policy: str = "default",
    ep_degree: int = 1,
    dp_mod_ep_mesh: DeviceMesh | None = None,
    dp_mesh_dims: DataParallelMeshDims | None = None,
    edp_mesh_dims: DataParallelMeshDims | None = None,
    gradient_divide_factor: int | None = None,
    symm_mem_scope: FSDPSymmMemScope = None,
):
    """Apply data parallelism (via FSDP2) to the model.

    When EP is enabled (``ep_degree > 1``), uses flat FSDP with
    ``shard_placement_fn`` to route expert params to ``dp_mod_ep_mesh``
    and other params to ``dp_mesh`` within a single ``fully_shard`` call
    per transformer block — matching Titan's approach and avoiding
    nested FSDP hooks that cause SAC op-count mismatches during recompute.
    """
    mp_policy = MixedPrecisionPolicy(
        param_dtype=param_dtype,
        reduce_dtype=reduce_dtype,
        cast_forward_inputs=False,
    )
    fsdp_config: dict[str, Any] = {"mesh": dp_mesh, "mp_policy": mp_policy}
    if dp_mesh_dims is not None:
        fsdp_config["dp_mesh_dims"] = dp_mesh_dims
    if cpu_offload:
        fsdp_config["offload_policy"] = CPUOffloadPolicy()

    reshard_after_forward = get_fsdp_reshard_after_forward_policy(
        reshard_after_forward_policy, pp_enabled
    )

    # When input/output embeddings are tied (e.g. Qwen3), tok_embeddings and
    # lm_head share one parameter. FSDP2 forbids a parameter being managed by
    # two FSDP groups, so they must be grouped into a single unit.
    tok_emb_weight = getattr(model.tok_embeddings, "weight", None)
    lm_head_weight = getattr(model.lm_head, "weight", None)
    # Detect tying by parameter identity (the exact thing FSDP2 checks); this
    # also skips PP nn.Identity placeholders, which have no `.weight`.
    tie_word_embeddings = (
        tok_emb_weight is not None
        and lm_head_weight is not None
        and tok_emb_weight is lm_head_weight
    )

    if tie_word_embeddings:
        fully_shard(
            [
                m
                for m in (model.tok_embeddings, model.norm, model.lm_head)
                if m is not None
            ],
            **fsdp_config,
            reshard_after_forward=reshard_after_forward_policy == "always",
        )
    elif model.tok_embeddings is not None:
        fully_shard(
            model.tok_embeddings,
            **fsdp_config,
            reshard_after_forward=reshard_after_forward,
        )

    for transformer_block in model.layers:
        if (
            hasattr(transformer_block, "moe_enabled")
            and transformer_block.moe_enabled
            and ep_degree > 1
        ):
            from torch.distributed.fsdp._fully_shard._fsdp_common import (
                FSDPMeshInfo,
                ShardPlacementResult,
            )
            from torch.distributed.fsdp._fully_shard._fsdp_init import _get_mesh_info

            assert dp_mod_ep_mesh is not None
            moe_module = getattr(transformer_block, "mlp", None)
            routed_experts = moe_module.routed_experts
            w13_params = set(routed_experts.w13.parameters())
            w2_params = set(routed_experts.w2.parameters())
            expert_params = w13_params | w2_params
            num_experts = routed_experts.w13.group_size

            expert_sharding_size = dp_mod_ep_mesh["edp_shard"].size() * ep_degree
            shard_expert_dim = expert_sharding_size <= num_experts

            edp_mesh_info = _get_mesh_info(dp_mod_ep_mesh, edp_mesh_dims)
            dp_mesh_info = _get_mesh_info(dp_mesh, dp_mesh_dims)
            assert isinstance(edp_mesh_info, FSDPMeshInfo)
            assert isinstance(dp_mesh_info, FSDPMeshInfo)

            def _shard_placement_fn(
                param: nn.Parameter,
                _expert_params: set = expert_params,
                _w13_params: set = w13_params,
                _shard_expert_dim: bool = shard_expert_dim,
                _edp_mesh_info: FSDPMeshInfo = edp_mesh_info,
                _dp_mesh_info: FSDPMeshInfo = dp_mesh_info,
            ) -> ShardPlacementResult:
                if param in _expert_params:
                    placement = (
                        Shard(0)
                        if _shard_expert_dim
                        else Shard(2 if param in _w13_params else 1)
                    )
                    return ShardPlacementResult(
                        placement=placement, mesh_info=_edp_mesh_info
                    )
                return ShardPlacementResult(placement=Shard(0), mesh_info=_dp_mesh_info)

            fully_shard(
                transformer_block,
                **fsdp_config,
                reshard_after_forward=reshard_after_forward,
                shard_placement_fn=_shard_placement_fn,
            )
        else:
            fully_shard(
                transformer_block,
                **fsdp_config,
                reshard_after_forward=reshard_after_forward,
            )

    # As an optimization, do not reshard_after_forward the last layers by default
    # since FSDP would prefetch them immediately after the forward pass. When
    # weights are tied, norm/lm_head are already grouped with tok_embeddings above.
    if not tie_word_embeddings and model.norm is not None and model.lm_head is not None:
        fully_shard(
            [model.norm, model.lm_head],
            **fsdp_config,
            reshard_after_forward=reshard_after_forward_policy == "always",
        )

    fully_shard(model, **fsdp_config)

    enable_fsdp_symm_mem(model, symm_mem_scope)

    # Disable FSDP's automatic gradient division for all FSDP modules
    disable_fsdp_gradient_division(model)

    # NOTE: set up explicit prefetching when EP is enabled, as D2H syncs
    # in EP could interfere with implicit prefetching in FSDP.
    if ep_degree == 1:
        return

    # forward
    transformer_blocks = list(model.layers.values())
    next_transformer_blocks = transformer_blocks[1:] + [None]

    if model.tok_embeddings is not None and model.layers is not None:
        model.tok_embeddings.set_modules_to_forward_prefetch([transformer_blocks[0]])

    for transformer_block, next_transformer_block in zip(
        transformer_blocks, next_transformer_blocks
    ):
        if next_transformer_block is not None:
            transformer_block.set_modules_to_forward_prefetch([next_transformer_block])
        elif model.norm is not None and model.lm_head is not None:
            transformer_block.set_modules_to_forward_prefetch(
                [model.norm, model.lm_head]
            )

    # backward
    reversed_transformer_blocks = list(reversed(model.layers.values()))
    prev_transformer_blocks = reversed_transformer_blocks[1:] + [None]

    if (
        model.norm is not None
        and model.lm_head is not None
        and model.layers is not None
    ):
        model.lm_head.set_modules_to_backward_prefetch([reversed_transformer_blocks[0]])

    for transformer_block, prev_transformer_block in zip(
        reversed_transformer_blocks, prev_transformer_blocks
    ):
        if prev_transformer_block is not None:
            transformer_block.set_modules_to_backward_prefetch([prev_transformer_block])
        elif model.tok_embeddings is not None:
            transformer_block.set_modules_to_backward_prefetch([model.tok_embeddings])
