# 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 torchtitan_recipes.tests.suites.models as recipes

from torchtitan_recipes.tests.models.deepseek_v3 import deepseek_v3_debugmodel
from torchtitan_recipes.tests.models.gpt_oss import gpt_oss_debugmodel_flex
from torchtitan_recipes.tests.models.llama3 import llama3_debugmodel

from tests.integration_tests import IntegrationTestDefinition


def build_model_tests_list() -> list[IntegrationTestDefinition]:
    """
    Build the list of model parallelism test configurations.
    This test suite is aimed at testing the model parallelism of torchtitan, and will
    broadly cover all the supported model parallelism patterns on all the supported
    models.
    """
    return [
        IntegrationTestDefinition(
            configs=[recipes.llama3_debugmodel_fsdp2_tp2_cp2],
            test_descr="Llama 3 FSDP+TP+CP",
            test_name="llama3_fsdp+tp+cp",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/llama3.txt"
            ),
            loss_compare_seed_config=llama3_debugmodel,
        ),
        IntegrationTestDefinition(
            configs=[recipes.llama3_debugmodel_region_ac_fsdp2_tp2_cp2],
            test_descr="Llama 3 FSDP+TP+CP+RegionAC",
            test_name="llama3_fsdp+tp+cp+region_ac",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/llama3.txt"
            ),
            loss_compare_seed_config=llama3_debugmodel,
        ),
        IntegrationTestDefinition(
            configs=[recipes.llama3_debugmodel_fsdp2_tp2_pp2],
            test_descr="Llama 3 FSDP+TP+PP",
            test_name="llama3_fsdp+tp+pp",
            ngpu=8,
            golden_numerics_path="tests/assets/losses/real_pg/{gpu_arch}/llama3_pp.txt",
            use_real_pg=True,
        ),
        # Integration Test Cases for DeepSeek V3
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_mtp_fsdp4_ep2],
            test_descr="DeepSeek V3 MTP FSDP+EP",
            test_name="deepseek_v3_mtp_fsdp+ep",
            ngpu=4,
            # The Helion fused RoPE kernels are CUDA-only and tuned for NVIDIA
            # H100/GB200; skip on ROCm where they are unvalidated.
            skip_rocm_test=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_mtp_tp2_cp2],
            test_descr="DeepSeek V3 MTP TP+CP with SP",
            test_name="deepseek_v3_mtp_tp+cp",
            ngpu=4,
            use_real_pg=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_fsdp8_ep8],
            test_descr="DeepSeek V3 FSDP+EP",
            test_name="deepseek_v3_fsdp+ep",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/deepseek_v3.txt"
            ),
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_fsdp2_cp2_pp2_ep4],
            test_descr="DeepSeek V3 FSDP+CP+PP+EP",
            test_name="deepseek_v3_fsdp+cp+pp+ep",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/real_pg/{gpu_arch}/deepseek_v3_cp_pp.txt"
            ),
            loss_compare_seed_config=deepseek_v3_debugmodel,
            use_real_pg=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_hsdp2x2_ep2],
            test_descr="DeepSeek V3 HSDP+EP",
            test_name="deepseek_v3_hsdp+ep",
            ngpu=4,
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v3_debugmodel_fused_mla_swiglu_fsdp4_ep2],
            test_descr="DeepSeek V3 fused MLA+SwiGLU FSDP+EP",
            test_name="deepseek_v3_fused_mla_swiglu_fsdp+ep",
            ngpu=4,
            skip_rocm_test=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.deepseek_v4_debugmodel_fsdp2_tp2_ep2],
            test_descr="DeepSeek V4 FSDP+TP+EP",
            test_name="deepseek_v4_fsdp+tp+ep",
            ngpu=4,
            # Sparse attention / indexer kernels are CUDA-only and unvalidated
            # on ROCm.
            skip_rocm_test=True,
            # Runs on a real PG. Under Fake PG this config's sequence-parallel
            # collectives return activations that alias their inputs, which
            # corrupts a saved-for-backward tensor and blows up grad_norm at
            # step 1. The same config trains cleanly on a real 4-GPU PG
            # (grad_norm ~3.8), so keep it on a real PG until the Fake PG
            # collective aliasing under spmd_types is fixed.
            use_real_pg=True,
        ),
        # Integration Test Cases for Qwen3 dense and MoE model
        IntegrationTestDefinition(
            configs=[recipes.qwen3_debugmodel_moe_param_groups_fsdp2_tp2_cp2_ep8],
            test_descr="Qwen3 MoE FSDP+TP+CP+EP (param groups)",
            test_name="qwen3_moe_fsdp+tp+cp+ep_param_groups",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/qwen3.txt"
            ),
            loss_compare_seed_config=recipes.qwen3_debugmodel_moe_param_groups_seed,
        ),
        IntegrationTestDefinition(
            configs=[
                recipes.qwen3_debugmodel_fsdp2_tp2_cp2_no_sp,
                recipes.qwen3_debugmodel_fsdp2_tp2_cp2,
            ],
            test_descr="Qwen3 FSDP+TP+CP (SP disabled)",
            test_name="qwen3_fsdp+tp+cp_no_sp",
            ngpu=8,
        ),
        # Integration Test Cases for Qwen3.5
        IntegrationTestDefinition(
            configs=[recipes.qwen35_debugmodel_moe_fsdp2_tp2_pp2_ep4],
            test_descr="Qwen3.5 MoE FSDP+TP+EP+PP",
            test_name="qwen3_5_moe_fsdp+tp+ep+pp",
            ngpu=8,
            use_real_pg=True,
            # short_conv's CuTe/CUTLASS kernel (attn_gym) is CUDA-only.
            skip_rocm_test=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.qwen35_debugmodel_moe_fsdp4_tp2_ep4],
            test_descr="Qwen3.5 MoE FSDP+TP+EP",
            test_name="qwen3_5_moe_fsdp+tp+ep",
            ngpu=8,
            # NOTE: This topology is not bitwise deterministic with Real PG on
            # A10G, so this case provides end-to-end coverage without a golden.
            # short_conv's CuTe/CUTLASS kernel (attn_gym) is CUDA-only.
            skip_rocm_test=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.qwen35_debugmodel_varlen_attn_fsdp2_tp2_sac],
            test_descr="Qwen3.5 FSDP+TP+VARLEN_ATTN + selective AC",
            test_name="qwen3_5_fsdp+tp+varlen_attn+per_op_sac",
            ngpu=4,
            skip_rocm_test=True,
            use_real_pg=True,
        ),
        # Integration Test Cases for gpt-oss
        IntegrationTestDefinition(
            configs=[recipes.gpt_oss_debugmodel_fsdp4_tp2_ep4],
            test_descr="GPT-OSS FSDP+TP+EP",
            test_name="gpt_oss_fsdp+tp+ep",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/gpt_oss.txt"
            ),
        ),
        IntegrationTestDefinition(
            configs=[recipes.gpt_oss_debugmodel_flex_fsdp2_cp2_pp2_ep4_sac],
            test_descr="GPT-OSS PP+FSDP+CP+EP+selective AC",
            test_name="gpt_oss_pp+fsdp+cp+ep+sacop",
            ngpu=8,
            golden_numerics_path="tests/assets/losses/real_pg/{gpu_arch}/gpt_oss_pp.txt",
            loss_compare_seed_config=gpt_oss_debugmodel_flex,
            use_real_pg=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.gpt_oss_debugmodel_fsdp4_pp2_ep4_sac],
            test_descr="GPT-OSS PP+FSDP+EP+selective AC with VarlenInnerAttention",
            test_name="gpt_oss_pp+fsdp+ep+sacop",
            ngpu=8,
            use_real_pg=True,
        ),
        # Integration Test Cases for Kimi K2.7
        IntegrationTestDefinition(
            configs=[recipes.kimi_k2_5_debugmodel_muon_fsdp2_pp2_ep2],
            test_descr="Kimi K2.7 DistMuon PP+FSDP+EP",
            test_name="kimi_k2_5_muon_pp+fsdp+ep",
            ngpu=4,
            timeout=600,
            use_real_pg=True,
        ),
        IntegrationTestDefinition(
            configs=[recipes.kimi_k2_5_debugmodel_muon_fsdp8_ep8],
            test_descr="Kimi K2.5 DistMuon FSDP+EP",
            test_name="kimi_k2_5_muon_fsdp+ep",
            ngpu=8,
        ),
        # Integration Test Cases for Muse Glimmer
        IntegrationTestDefinition(
            configs=[recipes.muse_glimmer_debugmodel_fsdp8],
            test_descr="Muse Glimmer text FSDP",
            test_name="muse_glimmer_text_fsdp",
            ngpu=8,
            golden_numerics_path=(
                "tests/assets/losses/{execution_mode}/{gpu_arch}/muse_glimmer.txt"
            ),
        ),
        IntegrationTestDefinition(
            configs=[recipes.muse_glimmer_debugmodel_mm_fsdp2_tp2],
            test_descr="Muse Glimmer multimodal FSDP+TP+SP",
            test_name="muse_glimmer_mm_fsdp+tp+sp",
            ngpu=4,
        ),
        IntegrationTestDefinition(
            configs=[recipes.muse_glimmer_debugmodel_mm_tp2_cp2_pp2],
            test_descr="Muse Glimmer multimodal TP+CP+PP+SP",
            test_name="muse_glimmer_mm_tp+cp+pp+sp",
            ngpu=8,
            use_real_pg=True,
        ),
    ]
