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

"""
Separate test runner for CooR precompile integration tests.

Each test has two steps:
1. Run precompile_main.py on a single process to generate a rank-agnostic
   compiled artifact.
2. Run training via graph_trainer/run_train_precompile.sh (which passes
   --virtual-local-rank to torchrun) to load and train with the artifact.

Usage:
    python -m torchtitan.experiments.graph_trainer.tests.run_precompile_tests \
        output_dir --ngpu 8
"""

import argparse

import logging
import os
import subprocess
import tempfile
import time
from collections.abc import Callable
from dataclasses import dataclass

from tests.integration_tests import get_importable_config_module

from torchtitan_recipes.tests.graph_trainer.precompile import (
    deepseek_v3_precompile_fsdp_tp_ep,
    llama3_precompile_fsdp_tp,
)

from torchtitan.experiments.graph_trainer.trainer import GraphTrainer
from torchtitan.observability.logging import init_logger


logger = logging.getLogger(__name__)


@dataclass
class PrecompileTestDefinition:
    config: Callable[[], GraphTrainer.Config]
    artifact_dir: str
    test_descr: str
    test_name: str
    ngpu: int = 8
    disabled: bool = False


def _build_precompile_tests() -> list[PrecompileTestDefinition]:
    fx_trace_precompile_dir = tempfile.mkdtemp(prefix="fx_trace_precompile_")
    dsv3_fx_trace_precompile_dir = tempfile.mkdtemp(prefix="dsv3_fx_trace_precompile_")
    return [
        # Uses the SDPA backend: the default FlexInnerAttention backend bakes a
        # BlockMask into the precompiled artifact, whose mask_mod closures are
        # Python code objects that pickle.dumps cannot serialize ("TypeError:
        # cannot pickle code objects" in precompile_fx_trace_save). SDPA carries
        # no such object, so it exercises the precompile machinery cleanly.
        # TODO: re-test on FlexInnerAttention once BlockMask is excluded/rebuilt at
        # load time (or becomes picklable).
        PrecompileTestDefinition(
            config=llama3_precompile_fsdp_tp,
            artifact_dir=fx_trace_precompile_dir,
            test_descr="aot_fx_trace llama3 precompile FSDP+TP",
            test_name="aot_fx_trace_llama3_precompile_fsdp_tp",
            ngpu=8,
        ),
        # TODO: disabled — precompile sharding propagation fails on aten.view
        # with a data-dependent unbacked symint ("Could not extract specialized
        # integer from u13") for DSv3 MoE. Separate from the empty_strided
        # shadow-node fix; re-enable once the precompile symint issue is fixed.
        PrecompileTestDefinition(
            config=deepseek_v3_precompile_fsdp_tp_ep,
            artifact_dir=dsv3_fx_trace_precompile_dir,
            test_descr="aot_fx_trace deepseek_v3 precompile FSDP+TP+EP",
            test_name="aot_fx_trace_deepseek_v3_precompile_fsdp_tp_ep",
            ngpu=8,
            disabled=True,
        ),
    ]


RUN_TRAIN_SCRIPT = "torchtitan/experiments/graph_trainer/run_train_precompile.sh"


def run_precompile_tests(args):
    test_list = _build_precompile_tests()

    ran_any = False
    for test in test_list:
        if args.test_name != "all" and test.test_name != args.test_name:
            continue
        if test.disabled:
            logger.info(f"Skipping disabled test: {test.test_name}")
            continue
        if args.ngpu < test.ngpu:
            logger.info(
                f"Skipping test {test.test_name} that requires {test.ngpu} gpus,"
                f" because --ngpu arg is {args.ngpu}"
            )
            continue

        ran_any = True
        all_ranks = ",".join(map(str, range(test.ngpu)))
        output_dir_arg = f"--output-dir {args.output_dir}/{test.test_name}"
        env = os.environ.copy()
        env["TORCHTITAN_PRECOMPILE_ARTIFACT_DIR"] = test.artifact_dir
        config_args = (
            f"--module {get_importable_config_module(test.config)} "
            f"--config {test.config.__name__}"
        )
        precompile_command = (
            "python -m torchtitan.experiments.graph_trainer.precompile_main "
            + config_args
        )

        # Step 1: precompile
        logger.info(
            f"===== {time.strftime('%Y-%m-%d %H:%M:%S')} "
            f"Precompile step for {test.test_descr}: "
            f"{precompile_command} ====="
        )
        result = subprocess.run(precompile_command, text=True, shell=True, env=env)
        logger.info(result.stdout)
        if result.returncode != 0:
            raise Exception(
                f"Precompile step failed for: {test.test_descr}, "
                f"command: {precompile_command}"
            )

        # Step 2: training with the precompiled artifact
        cmd = f"NGPU={test.ngpu} LOG_RANK={all_ranks} " f"./{RUN_TRAIN_SCRIPT}"
        cmd = f'TORCH_TRACE="{args.output_dir}/{test.test_name}/compile_trace" ' + cmd
        cmd += " " + output_dir_arg
        cmd += " " + config_args

        logger.info(
            f"===== {time.strftime('%Y-%m-%d %H:%M:%S')} "
            f"Training step for {test.test_descr}: {cmd} ====="
        )
        result = subprocess.run(cmd, text=True, shell=True, env=env)
        logger.info(result.stdout)
        if result.returncode != 0:
            raise Exception(
                f"Training step failed for: {test.test_descr}, command: {cmd}"
            )

    if not ran_any:
        available = [t.test_name for t in test_list]
        logger.warning(
            f"No precompile tests were run for --test_name '{args.test_name}'.\n"
            f"Available test names: {available}"
        )


def main():
    init_logger()
    parser = argparse.ArgumentParser()
    parser.add_argument("output_dir")
    parser.add_argument(
        "--test_name",
        default="all",
        help="Specific test name to run (default: all)",
    )
    parser.add_argument("--ngpu", default=8, type=int)
    args = parser.parse_args()

    if not os.path.exists(args.output_dir):
        os.makedirs(args.output_dir)
    if os.listdir(args.output_dir):
        raise RuntimeError("Please provide an empty output directory.")

    run_precompile_tests(args)


if __name__ == "__main__":
    main()
