# 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 importlib
import inspect
import unittest
from dataclasses import dataclass
from unittest import mock

from torchtitan.config.transform import ContextParallelTransform
from torchtitan.distributed.context_parallel import ContextParallelLoadBalancer
from torchtitan.models.common.attention import FlexInnerAttention
from torchtitan.protocols.module import Module


class TestDecoderConfigCpValidation(unittest.TestCase):
    """``Trainer.Config.__post_init__`` applies the CP gates at config time."""

    @staticmethod
    def _config(*, cp: int, varlen: bool = False, cp_kernel: bool = False):
        from torchtitan.models.common.attention.cp_attention import (
            KVAllGatherCPFlexInnerAttention,
        )
        from torchtitan_recipes.tests.models.llama3 import (
            llama3_debugmodel,
            llama3_debugmodel_varlen_attn,
        )

        config = (llama3_debugmodel_varlen_attn if varlen else llama3_debugmodel)(
            seq_len=512
        )
        if cp_kernel:
            # Apply the transform without its final validation.
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            ).transform(config.model)
        config.parallelism.context_parallel_degree = cp
        return config

    def test_allows_cp_kernel(self):
        config = self._config(cp=2, cp_kernel=True)
        config.__post_init__()

    def test_contiguous_cp_requires_divisibility_by_cp_only(self):
        config = self._config(cp=2, cp_kernel=True)
        config.parallelism.context_parallel_load_balancer = None
        config.training.num_tokens_per_microbatch_per_dp_rank = 6
        config.__post_init__()

    def test_ptrr_cp_requires_divisibility_by_cp_only(self):
        from torchtitan.distributed.context_parallel import (
            PTRRFlexAttentionCPLoadBalancer,
        )

        config = self._config(cp=2, cp_kernel=True)
        config.parallelism.context_parallel_load_balancer = (
            PTRRFlexAttentionCPLoadBalancer.Config()
        )
        config.training.num_tokens_per_microbatch_per_dp_rank = 6
        config.__post_init__()

    def test_allows_plain_flex_without_cp(self):
        config = self._config(cp=1)
        config.__post_init__()

    def test_rejects_cp_kernel_without_cp(self):
        config = self._config(cp=1, cp_kernel=True)
        with self.assertRaisesRegex(ValueError, "context parallel degree is 1"):
            config.__post_init__()

    def test_rejects_plain_flex_cp(self):
        config = self._config(cp=2)
        with self.assertRaisesRegex(ValueError, "KVAllGatherCPFlexInnerAttention"):
            config.__post_init__()

    def test_rejects_varlen_cp(self):
        config = self._config(cp=2, varlen=True)
        with self.assertRaisesRegex(ValueError, "CPInnerAttention"):
            config.__post_init__()

    def test_rejects_an_unrecognized_kernel_cp(self):
        class LocalOnlyAttention(Module):
            @dataclass(kw_only=True, slots=True)
            class Config(Module.Config):
                pass

        config = self._config(cp=2)
        for layer in config.model.layers:
            layer.attention.inner_attention = LocalOnlyAttention.Config()
        with self.assertRaisesRegex(ValueError, "CPInnerAttention"):
            config.__post_init__()


class TestUlyssesConfigValidation(unittest.TestCase):
    """Validate Ulysses configuration requirements."""

    @staticmethod
    def _config(
        *,
        cp: int = 2,
        tp: int = 1,
        load_balancer: ContextParallelLoadBalancer.Config | None = None,
        n_heads: int | None = None,
        n_kv_heads: int | None = None,
    ):
        from torchtitan.models.common.attention.cp_attention import (
            UlyssesCPFlexInnerAttention,
        )
        from torchtitan_recipes.tests.models.llama3 import llama3_debugmodel

        config = llama3_debugmodel(seq_len=512)
        attention = config.model.layers[0].attention
        if n_heads is not None:
            attention.n_heads = n_heads
        if n_kv_heads is not None:
            attention.n_kv_heads = n_kv_heads
        ContextParallelTransform(
            inner_attention_map={FlexInnerAttention: UlyssesCPFlexInnerAttention}
        ).transform(config.model)
        config.parallelism.context_parallel_degree = cp
        config.parallelism.tensor_parallel_degree = tp
        config.parallelism.context_parallel_load_balancer = load_balancer
        return config

    def test_allows_none_for_contiguous_sharding(self):
        config = self._config(load_balancer=None)
        config.__post_init__()

    def test_rejects_reordering_load_balancer(self):
        from torchtitan.distributed.context_parallel import HeadTailCPLoadBalancer

        config = self._config(load_balancer=HeadTailCPLoadBalancer.Config())
        with self.assertRaisesRegex(ValueError, "must be None"):
            config.__post_init__()

    def test_rejects_unknown_load_balancer(self):
        class OtherLoadBalancer(ContextParallelLoadBalancer):
            @dataclass(kw_only=True, slots=True)
            class Config(ContextParallelLoadBalancer.Config):
                pass

        config = self._config(load_balancer=OtherLoadBalancer.Config())
        with self.assertRaisesRegex(ValueError, "must be None"):
            config.__post_init__()

    def test_rejects_kv_heads_indivisible_by_cp(self):
        config = self._config(cp=4, tp=1, n_heads=8, n_kv_heads=2)
        with self.assertRaisesRegex(ValueError, r"n_kv_heads \(2\)"):
            config.__post_init__()

    def test_rejects_heads_indivisible_by_tp_times_cp(self):
        config = self._config(cp=8, tp=2, n_heads=8, n_kv_heads=8)
        with self.assertRaisesRegex(ValueError, r"n_heads \(8\)"):
            config.__post_init__()

    def test_allows_heads_divisible_by_tp_times_cp(self):
        config = self._config(cp=4, tp=2, n_heads=8, n_kv_heads=8)
        config.__post_init__()

    def test_allows_different_cp_backends_with_contiguous_sharding(self):
        from dataclasses import fields

        from torchtitan.models.common.attention.cp_attention import (
            KVAllGatherCPFlexInnerAttention,
            UlyssesCPFlexInnerAttention,
        )

        config = self._config(cp=2, tp=1)
        layer = config.model.layers[1]
        existing = layer.attention.inner_attention
        self.assertIsInstance(existing, UlyssesCPFlexInnerAttention.Config)
        layer.attention.inner_attention = KVAllGatherCPFlexInnerAttention.Config(
            **{f.name: getattr(existing, f.name) for f in fields(existing)}
        )
        config.__post_init__()


class TestGptOssUlysses(unittest.TestCase):
    def _parallelize(self, inner_attention):
        from types import SimpleNamespace

        from torchtitan.models.gpt_oss.model import GptOssModel

        with self.assertRaisesRegex(NotImplementedError, "Ulysses CP"):
            GptOssModel.parallelize(
                SimpleNamespace(
                    config=SimpleNamespace(
                        base_attention_backends=(inner_attention.Config(),)
                    )
                ),
                parallelism_context=SimpleNamespace(cp_enabled=True),
                training=None,
                parallelism=None,
                local_compile_regions=[],
                ac_config=None,
                dump_folder="",
            )

    def test_rejects_flex(self):
        from torchtitan.models.common.attention.cp_attention import (
            UlyssesCPFlexInnerAttention,
        )

        self._parallelize(UlyssesCPFlexInnerAttention)

    def test_rejects_varlen(self):
        from torchtitan.models.common.attention.cp_attention import (
            UlyssesCPVarlenInnerAttention,
        )

        self._parallelize(UlyssesCPVarlenInnerAttention)


class TestHeadDivisibility(unittest.TestCase):
    """Validate when CP adds to head sharding."""

    @staticmethod
    def _config(
        *, inner_attention=None, cp: int = 1, tp: int = 1, n_heads: int, n_kv_heads: int
    ):
        from torchtitan_recipes.tests.models.llama3 import llama3_debugmodel

        config = llama3_debugmodel(seq_len=512)
        attention = config.model.layers[0].attention
        attention.n_heads = n_heads
        attention.n_kv_heads = n_kv_heads
        if inner_attention is not None:
            ContextParallelTransform(
                inner_attention_map={FlexInnerAttention: inner_attention}
            ).transform(config.model)
        config.parallelism.context_parallel_degree = cp
        config.parallelism.tensor_parallel_degree = tp
        config.parallelism.context_parallel_load_balancer = None
        return config

    def test_all_gather_cp_keeps_cp_out_of_the_divisor(self):
        from torchtitan.models.common.attention.cp_attention import (
            KVAllGatherCPFlexInnerAttention,
        )

        config = self._config(
            inner_attention=KVAllGatherCPFlexInnerAttention,
            cp=4,
            tp=1,
            n_heads=2,
            n_kv_heads=2,
        )
        config.__post_init__()


class TestShippedCpRecipes(unittest.TestCase):
    """Validate every shipped CP recipe after construction."""

    _MODULES = (
        "torchtitan_recipes.models.muse_glimmer",
        "torchtitan_recipes.tests.suites.models",
        "torchtitan_recipes.tests.suites.features",
        "torchtitan_recipes.tests.suites.h100",
    )

    @classmethod
    def _recipes(cls):
        for name in cls._MODULES:
            module = importlib.import_module(name)
            for fn_name, fn in vars(module).items():
                if fn_name.startswith("_") or not inspect.isfunction(fn):
                    continue
                # Include local functions that take no arguments.
                if fn.__module__ != name or inspect.signature(fn).parameters:
                    continue
                yield f"{name}.{fn_name}", fn

    def test_every_cp_recipe_passes_the_gate(self):
        checked = 0
        for name, fn in self._recipes():
            with mock.patch(
                "torchtitan.config.transform.quantization.has_cuda_capability",
                return_value=True,
            ):
                config = fn()
            if config.parallelism.context_parallel_degree == 1:
                continue
            with self.subTest(recipe=name):
                config.__post_init__()
            checked += 1
        # Ensure recipe discovery found at least one CP recipe.
        self.assertGreater(checked, 0)

    def test_allows_mtp_cp(self):
        from torchtitan.config.transform import apply_transforms
        from torchtitan.models.common.attention.cp_attention import (
            KVAllGatherCPFlexInnerAttention,
        )
        from torchtitan_recipes.tests.models.deepseek_v3 import (
            deepseek_v3_debugmodel_mtp,
        )

        config = deepseek_v3_debugmodel_mtp()
        config.parallelism.context_parallel_degree = 2
        apply_transforms(
            config,
            [
                ContextParallelTransform(
                    inner_attention_map={
                        FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                    }
                )
            ],
        )


if __name__ == "__main__":
    unittest.main()
