# 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 functools
from dataclasses import dataclass
from typing import cast, TypeVar

from torchtitan.models.common.linear import GroupedLinear, Linear
from torchtitan.models.common.token_dispatcher import (
    AllToAllTokenDispatcher,
    HybridEPTokenDispatcher,
    TorchAOTokenDispatcher,
)

_LinearT = TypeVar("_LinearT", bound=Linear)


@functools.cache
def get_quantized_linear(
    quantized_cls: type[_LinearT],
    parent_cls: type[Linear],
) -> type[_LinearT]:
    """Get a cached quantized version of a linear module class."""
    if parent_cls is Linear:
        return quantized_cls

    quantized_config_cls = quantized_cls.Config

    class QuantizedLinear(
        quantized_cls,  # pyrefly: ignore [invalid-inheritance]
        parent_cls,
    ):
        @dataclass(kw_only=True, slots=True)
        class Config(quantized_config_cls):  # type: ignore[misc]
            pass

    quantized_name = quantized_cls.__name__.removesuffix("Linear")
    linear_name = f"{quantized_name}{parent_cls.__name__}"
    QuantizedLinear.__name__ = linear_name
    QuantizedLinear.__qualname__ = linear_name
    QuantizedLinear.__module__ = quantized_cls.__module__
    QuantizedLinear.Config.__qualname__ = f"{linear_name}.Config"
    QuantizedLinear.Config.__module__ = quantized_cls.__module__
    return cast(type[_LinearT], QuantizedLinear)


def swap_token_dispatcher(routed_experts_config, pad_multiple: int) -> None:
    """Swap the routed-experts token dispatcher config to support padded grouped GEMMs.

    Takes the ``RoutedExperts.Config`` (which owns the ``token_dispatcher`` child) and
    swaps its dispatcher in place. Requires a dispatcher that handles padding
    (TorchAOTokenDispatcher or DeepEP hybridep). Raises ValueError if the
    dispatcher doesn't support it.
    """
    dispatcher = routed_experts_config.token_dispatcher
    if isinstance(dispatcher, AllToAllTokenDispatcher.Config) and not isinstance(
        dispatcher, TorchAOTokenDispatcher.Config
    ):
        routed_experts_config.token_dispatcher = TorchAOTokenDispatcher.Config(
            num_experts=dispatcher.num_experts,
            top_k=dispatcher.top_k,
            pad_multiple=pad_multiple,
        )
    elif isinstance(dispatcher, HybridEPTokenDispatcher.Config):
        routed_experts_config.token_dispatcher = HybridEPTokenDispatcher.Config(
            num_experts=dispatcher.num_experts,
            top_k=dispatcher.top_k,
            non_blocking_capacity_factor=dispatcher.non_blocking_capacity_factor,
            pad_multiple=pad_multiple,
            hidden_dim=dispatcher.hidden_dim,
            num_max_tokens_per_rank=dispatcher.num_max_tokens_per_rank,
        )
    else:
        raise ValueError(
            f"MoE quantization requires a token dispatcher that supports "
            f"padding (TorchAOTokenDispatcher or HybridEPTokenDispatcher), "
            f"got {type(dispatcher).__name__}."
        )


def has_quantization(model_config) -> bool:
    """Check if any module in the model config has quantization applied."""
    from .mxfp8 import MXFP8Linear
    from .mxfp8.experts import _mxfp8_grouped_linear_cache
    from .nvfp4 import _nvfp4_grouped_linear_cache, NVFP4Linear

    quant_linear_types: list[type] = []
    if MXFP8Linear is not None:
        quant_linear_types.append(MXFP8Linear.Config)
    if NVFP4Linear is not None:
        quant_linear_types.append(NVFP4Linear.Config)

    has_quant_linear = bool(quant_linear_types) and any(
        isinstance(config, tuple(quant_linear_types))
        for _fqn, config, _parent, _attr in model_config.traverse(Linear.Config)
    )
    quant_grouped_linear_types = tuple(
        cls.Config  # type: ignore[attr-defined]
        for cls in (
            *_mxfp8_grouped_linear_cache.values(),
            *_nvfp4_grouped_linear_cache.values(),
        )
    )
    has_quant_moe = bool(quant_grouped_linear_types) and any(
        isinstance(config, quant_grouped_linear_types)
        for _fqn, config, _parent, _attr in model_config.traverse(GroupedLinear.Config)
    )
    return has_quant_linear or has_quant_moe
