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

"""The fused activation and DeepEP overrides compose independently."""

import unittest
from functools import partial

from torch.nn import init

from torchtitan.config.override import _REGISTRY, apply_overrides, OverrideConfig
from torchtitan.models.common.activation import Sigmoid
from torchtitan.models.common.config_utils import (
    make_moe_config,
    make_routed_experts_config,
    make_router_config,
)
from torchtitan.models.common.linear import GroupedLinear
from torchtitan.models.common.token_dispatcher import (
    AllToAllTokenDispatcher,
    DeepEPTokenDispatcher,
)
from torchtitan_recipes.overrides.fused_swiglu import fused_swiglu, FusedSwiGLU
from torchtitan_recipes.overrides.moe_token_dispatcher import deepep_override

_DIM = 16
_HIDDEN = 32
_E = 4

_FUSED_SWIGLU = "torchtitan_recipes.overrides.fused_swiglu.fused_swiglu"
_DEEPEP_OVERRIDE = (
    "torchtitan_recipes.overrides.moe_token_dispatcher.deepep_override",
    {"cuda_graph_compatible": True},
)

# The @override decorators register once, at the imports above. Capture the
# entries this test needs so a sibling test that calls clear_overrides() (e.g.
# test_override.py) can't leave the registry empty; setUp restores them.
_OVERRIDES = {
    key: _REGISTRY[key]
    for key in (
        "torchtitan_recipes.overrides.fused_swiglu.fused_swiglu",
        "torchtitan_recipes.overrides.moe_token_dispatcher.deepep_override",
    )
    if key in _REGISTRY
}


def _moe_config(dispatcher: type = AllToAllTokenDispatcher):
    param_init = {
        "w1_EFD": partial(init.trunc_normal_, std=0.02),
        "w2_EDF": partial(init.trunc_normal_, std=0.02),
        "w3_EFD": partial(init.trunc_normal_, std=0.02),
    }
    routed_experts = make_routed_experts_config(
        dim=_DIM,
        hidden_dim=_HIDDEN,
        num_experts=_E,
        top_k=1,
        param_init=param_init,
    )
    if dispatcher is DeepEPTokenDispatcher:
        routed_experts.token_dispatcher = DeepEPTokenDispatcher.Config(
            num_experts=_E,
            top_k=1,
            hidden_dim=_DIM,
            num_max_tokens_per_rank=16,
        )
    router = make_router_config(
        dim=_DIM,
        num_experts=_E,
        gate_param_init={"weight": partial(init.trunc_normal_, std=0.02)},
        score_func=Sigmoid.Config(),
        top_k=1,
    )
    return make_moe_config(num_experts=_E, router=router, routed_experts=routed_experts)


class TestInferenceMoEOverrides(unittest.TestCase):
    def setUp(self):
        # Restore the overrides if a previously run test cleared the registry.
        for name, ov in _OVERRIDES.items():
            _REGISTRY.setdefault(name, ov)

    def test_grouped_linears_and_dispatcher_are_siblings(self):
        """Expert projections and the dispatcher are first-class siblings."""
        cfg = _moe_config(DeepEPTokenDispatcher)
        self.assertIsInstance(cfg.routed_experts.w13, GroupedLinear.Config)
        self.assertIsInstance(cfg.routed_experts.w2, GroupedLinear.Config)
        self.assertIsInstance(
            cfg.routed_experts.token_dispatcher, DeepEPTokenDispatcher.Config
        )

    def test_deepep_both_overrides_apply_without_conflict(self):
        cfg = _moe_config(DeepEPTokenDispatcher)

        replacements = apply_overrides(
            OverrideConfig(imports=[_FUSED_SWIGLU, _DEEPEP_OVERRIDE]),
            cfg,
        )

        self.assertEqual(len(replacements), 2)
        self.assertIsInstance(cfg.routed_experts.w13, GroupedLinear.Config)
        self.assertIsInstance(cfg.routed_experts.activation_fn, FusedSwiGLU.Config)
        self.assertIsInstance(
            cfg.routed_experts.token_dispatcher, DeepEPTokenDispatcher.Config
        )
        self.assertTrue(cfg.routed_experts.token_dispatcher.cuda_graph_compatible)

    def test_non_deepep_dispatcher_flip_is_noop(self):
        cfg = _moe_config()

        # deepep_override targets DeepEP only; on a standard dispatcher just fusion applies.
        replacements = apply_overrides(
            OverrideConfig(imports=[_FUSED_SWIGLU, _DEEPEP_OVERRIDE]),
            cfg,
        )

        self.assertEqual(len(replacements), 1)
        self.assertIsInstance(cfg.routed_experts.w13, GroupedLinear.Config)
        self.assertIsInstance(cfg.routed_experts.activation_fn, FusedSwiGLU.Config)

    def test_composition_is_order_independent(self):
        # Disjoint sibling nodes -> either application order yields the same result.
        def summarize(ge):
            return (
                type(ge.activation_fn).__qualname__,
                type(ge.token_dispatcher).__qualname__,
                ge.token_dispatcher.cuda_graph_compatible,
            )

        a = _moe_config(DeepEPTokenDispatcher).routed_experts
        a.activation_fn = fused_swiglu(a.activation_fn)
        a.token_dispatcher = deepep_override(
            a.token_dispatcher, cuda_graph_compatible=True
        )

        b = _moe_config(DeepEPTokenDispatcher).routed_experts
        b.token_dispatcher = deepep_override(
            b.token_dispatcher, cuda_graph_compatible=True
        )
        b.activation_fn = fused_swiglu(b.activation_fn)

        self.assertEqual(summarize(a), summarize(b))
        self.assertIsInstance(a.activation_fn, FusedSwiGLU.Config)
        self.assertTrue(a.token_dispatcher.cuda_graph_compatible)

    def test_trainer_uses_only_experts_fusion(self):
        cfg = _moe_config(DeepEPTokenDispatcher)

        # Trainer imports only fused_swiglu: activation fused, dispatcher unchanged.
        apply_overrides(OverrideConfig(imports=[_FUSED_SWIGLU]), cfg)

        self.assertIsInstance(cfg.routed_experts.w13, GroupedLinear.Config)
        self.assertIsInstance(cfg.routed_experts.activation_fn, FusedSwiGLU.Config)
        self.assertFalse(cfg.routed_experts.token_dispatcher.cuda_graph_compatible)


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