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

"""
AutoParallel-based parallelization for DeepSeek V3.

Uses AutoParallelGraph to apply solver-based SPMD sharding on AutoParallel's
local_map DSv3 model (whose ops the solver supports), then lets graph_trainer
trace and compile the placed model through its normal `aot_fx_trace` train-step
pipeline. Requires a 2D sparse mesh (``edp_shard`` + ``ep``).

The torchtitan DSv3 model is replaced with AutoParallel's DeepSeekV3Model
because the solver doesn't support torchtitan's token_dispatcher ops
(aten::div.Tensor_mode). The two models share the same hierarchical config
layout via duck typing.
"""

import logging
import time

import torch
from torch.distributed.fsdp import MixedPrecisionPolicy
from torch.distributed.tensor.placement_types import Shard

from torchtitan.config import TORCH_DTYPE_MAP, TrainingConfig
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed import ParallelismContext
from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig
from torchtitan.distributed.fsdp import get_fsdp_reshard_after_forward_policy
from torchtitan.experiments.graph_trainer.autoparallel_api import AutoParallelGraph
from torchtitan.experiments.graph_trainer.compile import apply_compile
from torchtitan.experiments.graph_trainer.configs import GraphTrainerCompileConfig
from torchtitan.tools.utils import device_type


logger = logging.getLogger(__name__)


def _load_autoparallel_dsv3_dependency():
    """Load the temporary AutoParallel DSv3 integration dependency."""
    try:
        from autoparallel._testing.models.dsv3 import (
            annotate_deepseekv3_for_graph_trainer,
            DeepSeekV3Model,
        )
    except ImportError as exc:
        raise ImportError(
            "AutoParallel graph_trainer DeepSeek V3 currently depends on "
            "autoparallel._testing.models.dsv3. Move that model and annotation "
            "helper into a supported AutoParallel namespace before treating this "
            "route as a stable production dependency."
        ) from exc
    return DeepSeekV3Model, annotate_deepseekv3_for_graph_trainer


def _set_torchtitan_fields(parallel_model):
    if hasattr(parallel_model, "layers") and isinstance(
        parallel_model.layers, torch.nn.ModuleDict
    ):
        for block in parallel_model.layers.values():
            block.moe_enabled = hasattr(block, "moe")


def _preserve_moe_attributes(original_model, parallel_model):
    """Preserve MoE attributes (moe_enabled, load_balance_coeff) from original."""

    def get_moe_modules(model):
        moe_modules = []
        if hasattr(model, "layers"):
            blocks = (
                model.layers.values()
                if isinstance(model.layers, torch.nn.ModuleDict)
                else []
            )
            for block in blocks:
                if hasattr(block, "moe"):
                    moe_modules.append(block.moe)
        return moe_modules

    for orig_moe, par_moe in zip(
        get_moe_modules(original_model), get_moe_modules(parallel_model)
    ):
        if hasattr(orig_moe, "moe_enabled"):
            par_moe.moe_enabled = orig_moe.moe_enabled
        if hasattr(orig_moe, "load_balance_coeff"):
            par_moe.load_balance_coeff = orig_moe.load_balance_coeff


def parallelize_autoparallel_deepseekv3(
    model,
    *,
    parallelism_context: ParallelismContext,
    training: TrainingConfig,
    parallelism: ParallelismConfig,
    compile_config: GraphTrainerCompileConfig,
    ac_config: ActivationCheckpointingConfig,
    dump_folder: str,
):
    """Apply AutoParallelGraph SPMD sharding to DeepSeek V3.

    Returns a sharded model carrying AutoParallel train-step metadata.
    Requires a 2D sparse mesh (``edp_shard`` + ``ep``).
    """
    if parallelism_context.dp_replicate_enabled:
        raise ValueError("AutoParallel DeepSeek V3 does not support DDP yet")
    if parallelism_context.cp_enabled:
        raise ValueError("AutoParallel DeepSeek V3 does not support CP yet")
    if parallelism_context.pp_enabled:
        raise ValueError("AutoParallel DeepSeek V3 does not support PP yet")
    if parallelism_context.tp_enabled:
        raise ValueError("AutoParallel DeepSeek V3 does not support TP yet")

    required_sparse_axes = ("edp_shard", "ep")
    missing_sparse_axes = [
        name
        for name in required_sparse_axes
        if parallelism_context.get_optional_mesh(name) is None
    ]
    if missing_sparse_axes:
        raise ValueError(
            "AutoParallel DeepSeek V3 requires edp_shard and ep axes, but missing "
            f"{missing_sparse_axes}"
        )

    sparse_mesh = parallelism_context.get_mesh(list(required_sparse_axes))
    if sparse_mesh.ndim != 2 or sparse_mesh.mesh_dim_names != required_sparse_axes:
        raise ValueError(
            "AutoParallel DeepSeek V3 requires a 2D sparse mesh with edp_shard and ep "
            f"axes, but got mesh axes {sparse_mesh.mesh_dim_names}"
        )

    param_dtype = TORCH_DTYPE_MAP[training.mixed_precision_param]
    reduce_dtype = TORCH_DTYPE_MAP[training.mixed_precision_reduce]
    mp_policy = MixedPrecisionPolicy(
        param_dtype=param_dtype,
        reduce_dtype=reduce_dtype,
        cast_forward_inputs=False,
    )
    reshard_after_forward = get_fsdp_reshard_after_forward_policy(
        parallelism.fsdp_reshard_after_forward,
        parallelism_context.pp_enabled,
    )
    (
        APDeepSeekV3Model,
        annotate_deepseekv3_for_graph_trainer,
    ) = _load_autoparallel_dsv3_dependency()

    # Use AutoParallel's DSv3 model: torchtitan's token_dispatcher uses
    # aten::div.Tensor_mode which the AP solver doesn't support yet.
    # The AP model accepts torchtitan's config via duck typing (same
    # hierarchical attribute paths).
    with torch.device("meta"):
        ap_model = APDeepSeekV3Model(
            model.config,
            mesh=sparse_mesh,
            compute_dtype=param_dtype,
        )

    def input_fn():
        dp_degree = parallelism_context.dp_replicate * parallelism_context.dp_shard
        num_tokens_per_train_step = training.num_tokens_per_train_step
        if num_tokens_per_train_step < 0:
            num_tokens_per_train_step = (
                training.num_tokens_per_microbatch_per_dp_rank * dp_degree
            )
        tokens = torch.randint(
            0,
            ap_model.model_args.vocab_size,
            (num_tokens_per_train_step,),
            device=torch.device(device_type),
        )
        return tokens

    x_sharding = (Shard(0), Shard(0))

    autop = AutoParallelGraph(
        ap_model,
        input_fn,
        sparse_mesh,
        mp_policy=mp_policy,
        reshard_after_forward=reshard_after_forward,
        dynamic=True,
    )

    annotate_deepseekv3_for_graph_trainer(autop.model)

    with autop:
        autop.add_parameter_memory_constraint(low=None, high=None)
        autop.add_input_constraints([x_sharding])
        autop.add_output_constraints([x_sharding])

        t0 = time.time()
        sharding_placement = autop.optimize_placement()
        t1 = time.time()
        logger.info(f"AutoParallelGraph took {t1 - t0:.2f} seconds")

        # The solved output is logically batch-sharded over edp_shard and ep.
        # (Shard(0), Shard(0)). Those axes are data-parallel factors for
        # loss computation, so graph_trainer can consume each rank's local
        # logits as a plain tensor and pair them with local labels. Only TP
        # vocab sharding needs a DTensor output boundary for loss_parallel().
        parallel_mod = autop.apply_placement_for_fx_module(
            sharding_placement,
            compile_config=compile_config,
        )

    _set_torchtitan_fields(parallel_mod)
    _preserve_moe_attributes(ap_model, parallel_mod)

    model = apply_compile(
        parallel_mod,
        compile_config=compile_config,
        parallelism_context=parallelism_context,
    )
    return model
