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

"""
Single-process precompile entry point for graph_trainer.

Uses compile-on-one-rank (CooR) to generate a rank-agnostic compiled
artifact from a single process, which can then be loaded by all ranks
during torchrun training. This avoids the need to run torchrun with N
GPUs just for precompilation.

Usage:
    python -m torchtitan.experiments.graph_trainer.precompile_main \
        --module torchtitan_recipes.tests.graph_trainer.llama3 \
        --config graph_trainer_llama3_debugmodel
"""

import contextlib
import copy
import logging
from typing import Any, cast

import torch
import torch.distributed as dist

from torchtitan.components.loss import ChunkedLossWrapper
from torchtitan.config import apply_overrides, ConfigLoader, TORCH_DTYPE_MAP
from torchtitan.distributed import ParallelismContext
from torchtitan.experiments.graph_trainer.common_utils import (
    maybe_register_blockmask_pytree_node,
)
from torchtitan.experiments.graph_trainer.memory_policy import (
    validate_memory_policy_config,
)
from torchtitan.experiments.graph_trainer.precompile import (
    _FX_TRACE_ARTIFACT_KEY,
    _register_coor_ops,
)
from torchtitan.experiments.graph_trainer.storage import DiskStorageAdapter
from torchtitan.models.common.attention import FlexInnerAttention, VarlenInnerAttention
from torchtitan.models.common.aux_loss import AuxLoss
from torchtitan.models.common.decoder import Decoder
from torchtitan.models.deepseek_v3.mtp import get_mtp_token_counts
from torchtitan.observability.logging import init_logger
from torchtitan.tools import utils


logger = logging.getLogger(__name__)


def _common_setup(config):
    """Common setup for precompile: fake PG, CooR, model build."""
    compile_config = config.compile

    if not compile_config.precompile_artifact_dir:
        raise ValueError(
            "precompile_main requires compile.precompile_artifact_dir in the recipe."
        )

    parallelism = config.parallelism
    dp_replicate = parallelism.data_parallel_replicate_degree
    dp_shard = parallelism.data_parallel_shard_degree
    cp = parallelism.context_parallel_degree
    tp = parallelism.tensor_parallel_degree
    pp = parallelism.pipeline_parallel_degree

    # dp_shard=-1 means "use remaining ranks" which can't be inferred
    # in single-process mode. The compiled graph bakes in tensor shapes
    # that depend on dp_shard, so the exact value must match training.
    if dp_shard < 0:
        raise ValueError(
            "precompile_main requires an explicit "
            "parallelism.data_parallel_shard_degree (not -1) in the recipe. "
            "It must match the value used during torchrun training."
        )
    world_size = dp_replicate * dp_shard * cp * tp * pp

    logger.info(f"Initializing single-process precompile with world_size={world_size}")

    # rank must be 0 because --virtual-local-rank maps every torchrun rank
    # to local rank 0, so the precompiled artifact needs to match that setup.
    # Fake backend produces correct collective output shapes without real
    # communication, letting us trace distributed ops on a single process.
    dist.init_process_group("fake", rank=0, world_size=world_size)

    # CooR must be enabled globally (not just during tracing) so that the
    # parallelization phase (TP, FSDP mesh setup) also uses symbolic
    # coordinates rather than hardcoding rank-specific values.
    import torch.distributed.config as dist_config

    dist_config.compile_on_one_rank = True
    _register_coor_ops()

    # Match the deterministic mode that the training loop will use.
    # The backward graph captures use_deterministic_algorithms() at
    # compile time and asserts it matches at runtime.
    if config.debug.deterministic:
        torch.use_deterministic_algorithms(True)
        torch.backends.cudnn.deterministic = True
        torch.backends.cudnn.benchmark = False

    device = torch.device("cuda:0")
    torch.cuda.set_device(device)

    parallelism_context = ParallelismContext(
        dp_shard=dp_shard,
        dp_replicate=dp_replicate,
        cp=cp,
        tp=tp,
        pp=pp,
        ep=parallelism.expert_parallel_degree,
        world_size=world_size,
        enable_sequence_parallel=parallelism.enable_sequence_parallel,
    )
    parallelism_context.build_mesh()

    # TODO: Factor the model setup below with the training path so precompile
    # and training share a single implementation of build/parallelize/init.
    model_config = copy.deepcopy(config.model)
    model_config.set_sharding_(config.parallelism)
    config.model = model_config
    if config.override.imports:
        apply_overrides(config.override, config)
    model_config = config.model
    logger.info(f"Building {type(model_config).__qualname__} on meta device")
    with (
        parallelism_context.activate_spmd(),
        torch.device("meta"),
        utils.set_default_dtype(TORCH_DTYPE_MAP[config.training.dtype]),
    ):
        model = model_config.build()

    # For aot_fx_trace, apply_compile inside model.parallelize is a no-op
    # (returns model unchanged), so we pass the real compile_config.
    model = model.parallelize(
        parallelism_context=parallelism_context,
        training=config.training,
        parallelism=parallelism,
        compile_config=compile_config,
        ac_config=config.activation_checkpoint,
        dump_folder=config.dump_folder,
    )

    # CooR must be disabled during init_weights because DTensor RNG ops
    # (weight initialization seeding) raise NotImplementedError under
    # compile_on_one_rank=True. Re-enable for the tracing phase after.
    device_type = utils.device_type
    model.to_empty(device=device_type)
    dist_config.compile_on_one_rank = False
    try:
        with torch.no_grad():
            model.init_weights(buffer_device=None)
    finally:
        dist_config.compile_on_one_rank = True
    model.train()

    logger.info("Model parallelized and materialized")

    tokenizer = config.tokenizer.build(tokenizer_path=config.hf_assets_path)

    return (
        model,
        model_config,
        compile_config,
        parallelism_context,
        device,
        tokenizer,
    )


def _prepare_loss_for_precompile(model, loss_fn) -> None:
    """Match Trainer's post-parallelization loss setup for precompile tracing."""
    if not isinstance(loss_fn, ChunkedLossWrapper):
        return

    lm_head = getattr(model, "lm_head", None)
    if lm_head is None:
        raise ValueError("Model must have lm_head for ChunkedLossWrapper precompile")

    loss_fn.set_lm_head(lm_head)
    model._skip_lm_head = True


def _precompile_aot_fx_trace(
    config,
    model,
    model_config,
    compile_config,
    parallelism_context,
    device,
    tokenizer,
):
    """aot_fx_trace mode precompilation: make_fx tracing + Inductor."""
    from torchtitan.experiments.graph_trainer.make_fx_tracer import minimal_fx_tracer
    from torchtitan.experiments.graph_trainer.precompile import (
        compute_config_fingerprint,
        get_spmd_precompile_meshes,
        precompile_fx_trace_save,
    )
    from torchtitan.experiments.graph_trainer.spmd_graph_builder import (
        make_fwd_bwd_step,
    )

    loss_fn = config.loss.build()
    _prepare_loss_for_precompile(model, loss_fn)

    fwd_bwd_fn = make_fwd_bwd_step(model, loss_fn)

    num_tokens = config.training.num_tokens_per_microbatch_per_dp_rank
    vocab_size = model_config.vocab_size

    dummy_inputs = torch.randint(0, vocab_size, (num_tokens,), device=device)
    dummy_labels = torch.randint(0, vocab_size, (num_tokens,), device=device)
    # Match Trainer.train_step, which passes dense loss normalization as a
    # standalone int64 scalar tensor on the training device.
    global_num_tokens = (
        num_tokens
        * parallelism_context.dp_shard
        * parallelism_context.dp_replicate
        * parallelism_context.cp
    )
    dummy_global_loss_token_counts = torch.tensor(
        global_num_tokens, dtype=torch.int64, device=device
    )
    extra_kwargs: dict[str, Any] = {}

    if isinstance(model_config, Decoder.Config) and model_config.layers:
        attn_config = model_config.layers[0].attention
        inner_attention = attn_config.inner_attention

        positions = (
            torch.arange(num_tokens, dtype=torch.int32, device=dummy_inputs.device)
            % config.training.max_context_length
        )
        extra_kwargs["positions"] = positions
        extra_kwargs["padding_mask"] = torch.zeros(
            num_tokens, dtype=torch.bool, device=dummy_inputs.device
        )

        if isinstance(
            inner_attention, (FlexInnerAttention.Config, VarlenInnerAttention.Config)
        ):
            extra_kwargs["attention_metadata"] = cast(
                Decoder, model
            )._get_attention_metadata(
                positions=positions,
            )

        uses_aux_loss = next(model_config.traverse(AuxLoss.Config), None) is not None
        if not uses_aux_loss:
            uses_aux_loss = any(
                getattr(layer, "moe", None) is not None for layer in model_config.layers
            )
        if uses_aux_loss:
            _, routing_token_counts = get_mtp_token_counts(
                target_mask=torch.ones_like(dummy_labels, dtype=torch.bool),
                positions=positions,
                padding_mask=extra_kwargs["padding_mask"],
                num_mtp_layers=config.dataloader.num_mtp_layers,
            )
            num_pp_microbatches = (
                config.parallelism.num_pp_microbatches
                if parallelism_context.pp_enabled
                else 1
            )
            extra_kwargs["aux_loss_denominators"] = routing_token_counts * (
                parallelism_context.dp_replicate
                * parallelism_context.dp_shard
                * num_pp_microbatches
            )

    # TODO: Add CP support by generating a permutation and
    # sharding inputs here.
    # to shard dummy_inputs/dummy_labels/extra_kwargs along the sequence
    # dimension, matching the trainer's preprocess_inputs path.
    if parallelism_context.cp_enabled:
        raise NotImplementedError(
            "CooR precompile does not yet support context parallelism. "
            "Set parallelism.context_parallel_degree=1."
        )

    loss_parallel_ctx = (
        # TODO(bobrenjc93): Migrate graph trainer to the manual loss-parallel
        # custom autograd function and remove this DTensor context manager.
        torch.distributed.tensor.parallel.loss_parallel()
        if parallelism_context.tp_enabled
        else contextlib.nullcontext()
    )

    maybe_register_blockmask_pytree_node()

    logger.info("Tracing fwd+loss+bwd via make_fx...")
    with parallelism_context.activate_spmd(), loss_parallel_ctx:
        traced_result = minimal_fx_tracer(
            fwd_bwd_fn,
            module=model,
            precompile_meshes=get_spmd_precompile_meshes(parallelism_context),
        )(dummy_inputs, dummy_labels, dummy_global_loss_token_counts, extra_kwargs)
    logger.info(
        f"Traced graph has {len(list(traced_result.gm.graph.nodes))} nodes, "
        f"{len(traced_result.state_fqns)} state entries"
    )

    # Apply precompile-time graph passes (cleanup + regional_inductor)
    # so compiled Triton kernels are baked into the serialized artifact.
    # CUDA graph is excluded — it runs at load time on each rank.
    from torchtitan.experiments.graph_trainer.passes import (
        apply_graph_passes,
        compile_time_passes,
    )

    passes = compile_time_passes(
        traced_result, config, parallelism_context=parallelism_context
    )

    traced_result.gm = apply_graph_passes(
        traced_result.gm, traced_result.example_inputs, passes
    )
    logger.info(
        f"Applied {len(passes)} precompile graph passes, "
        f"graph now has {len(list(traced_result.gm.graph.nodes))} nodes"
    )

    storage = DiskStorageAdapter(compile_config.precompile_artifact_dir)
    config_fingerprint = compute_config_fingerprint(
        model, compile_config, parallelism_context
    )

    precompile_fx_trace_save(
        traced_result,
        storage,
        config_fingerprint=config_fingerprint,
    )

    logger.info(
        f"Precompile complete. Artifact saved to "
        f"{compile_config.precompile_artifact_dir}/{_FX_TRACE_ARTIFACT_KEY}.bin"
    )


def main():
    init_logger()
    config = ConfigLoader().load()

    (
        model,
        model_config,
        compile_config,
        parallelism_context,
        device,
        tokenizer,
    ) = _common_setup(config)
    validate_memory_policy_config(compile_config)

    _precompile_aot_fx_trace(
        config,
        model,
        model_config,
        compile_config,
        parallelism_context,
        device,
        tokenizer,
    )

    dist.destroy_process_group()


if __name__ == "__main__":
    main()
