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

"""GPT-OSS model flavors."""

import dataclasses
from collections.abc import Callable
from functools import partial

import torch.nn as nn

from torchtitan.config.transform import (
    ModelConfigConverter,
    validate_converter_compatibility,
)
from torchtitan.models.common import (
    ColumnParallelLinear,
    CosSinRoPE,
    Embedding,
    HiMidLoLinear,
    Linear,
    RMSNorm,
    RoPE,
    RowParallelLinear,
    Softmax,
    TransformerBlock,
)
from torchtitan.models.common.attention import (
    FlexInnerAttention,
    QKVLinear,
    SlidingWindowFlexInnerAttention,
    VarlenInnerAttention,
)
from torchtitan.models.common.config_utils import get_attention_config
from torchtitan.models.common.moe import MoE, RoutedExperts, TokenChoiceTopKRouter
from torchtitan.models.common.param_init import depth_scaled_std
from torchtitan.models.common.token_dispatcher import AllToAllTokenDispatcher
from .model import Attention, GptOssModel, GptOssTransformerBlock
from .moe import GptOssGroupedLinear, GptOssSwiGLU

__all__ = [
    "GptOssModel",
    "MODEL_FLAVORS",
    "build_model_config",
]

_NORM_INIT = {"weight": nn.init.ones_}
# GPT-OSS uses std=0.02 for embeddings (model-specific)
_EMBEDDING_INIT = {"weight": partial(nn.init.normal_, std=0.02)}


def _output_linear_init(dim: int) -> dict[str, Callable]:
    s = dim**-0.5
    return {
        "weight": partial(nn.init.trunc_normal_, std=s, a=-3 * s, b=3 * s),
        "bias": nn.init.zeros_,
    }


def _depth_init(layer_id: int) -> dict[str, Callable]:
    return {
        "weight": partial(nn.init.trunc_normal_, std=depth_scaled_std(0.02, layer_id)),
        "bias": nn.init.zeros_,
    }


def _make_gptoss_attn_config(
    *,
    dim: int,
    layer_id: int,
    attn_backend: str = "varlen",
    n_heads: int = 64,
    n_kv_heads: int = 8,
    head_dim: int = 64,
    sliding_window_size: int | None = None,
    rope: RoPE.Config,
) -> Attention.Config:
    """Build a fully-specified GPT-OSS Attention.Config for a single layer.

    All linear sub-configs have their in_features/out_features set.
    All linear params use depth-scaled init (including wq/wkv/wo).
    Sinks also use depth-scaled init.
    """

    inner_attention = get_attention_config(attn_backend)

    if sliding_window_size is not None and isinstance(
        inner_attention, VarlenInnerAttention.Config
    ):
        inner_attention = dataclasses.replace(
            inner_attention, window_size=(sliding_window_size - 1, 0)
        )
    elif sliding_window_size is not None and isinstance(
        inner_attention, FlexInnerAttention.Config
    ):
        inner_attention = SlidingWindowFlexInnerAttention.Config(
            block_size=inner_attention.block_size,
            kernel_options=inner_attention.kernel_options,
            window_size=sliding_window_size,
        )

    sinks_init = {
        "sinks": partial(nn.init.trunc_normal_, std=depth_scaled_std(0.02, layer_id))
    }

    qkv = QKVLinear.Config(
        head_dim=head_dim,
        n_heads=n_heads,
        n_kv_heads=n_kv_heads,
        wqkv=ColumnParallelLinear.Config(
            in_features=dim,
            out_features=(n_heads + 2 * n_kv_heads) * head_dim,
            bias=True,
            param_init=_depth_init(layer_id),
        ),
    )

    return Attention.Config(
        n_heads=n_heads,
        n_kv_heads=n_kv_heads,
        head_dim=head_dim,
        dim=dim,
        qkv_linear=qkv,
        wo=RowParallelLinear.Config(
            in_features=n_heads * head_dim,
            out_features=dim,
            bias=True,
            param_init=_depth_init(layer_id),
        ),
        sliding_window_size=sliding_window_size,
        inner_attention=inner_attention,
        param_init=sinks_init,
        rope=dataclasses.replace(rope),
    )


def _make_gptoss_experts_config(
    *,
    dim: int,
    hidden_dim: int,
    num_experts: int,
    layer_id: int,
    top_k: int,
) -> RoutedExperts.Config:
    """Build a fully-specified RoutedExperts.Config for a single GPT-OSS layer."""
    # TODO: Audit GPT-OSS gate/up depth scaling separately, including its biases.
    std = depth_scaled_std(0.02, layer_id)
    experts_init = {
        "weight": partial(nn.init.trunc_normal_, std=std),
        "bias": partial(nn.init.trunc_normal_, std=std),
    }
    return RoutedExperts.Config(
        w13=GptOssGroupedLinear.Config(
            group_size=num_experts,
            in_features=dim,
            out_features=hidden_dim,
            num_linears=2,
            param_init=experts_init,
        ),
        w2=GptOssGroupedLinear.Config(
            group_size=num_experts,
            in_features=hidden_dim,
            out_features=dim,
            param_init=experts_init,
        ),
        activation_fn=GptOssSwiGLU.Config(),
        token_dispatcher=AllToAllTokenDispatcher.Config(
            num_experts=num_experts,
            top_k=top_k,
        ),
    )


def _build_gptoss_layers(
    *,
    dim: int,
    n_layers: int,
    hidden_dim: int,
    num_experts: int,
    top_k: int,
    load_balance_coeff: float,
    attn_backend: str = "varlen",
    rope: RoPE.Config,
) -> list[TransformerBlock.Config]:
    """Build per-layer configs for GPT-OSS.

    Even-indexed layers (0, 2, 4, ...) use sliding window attention.
    All dimensional fields are set directly.
    """
    layers = []
    for layer_id in range(n_layers):
        attn_cfg = _make_gptoss_attn_config(
            dim=dim,
            layer_id=layer_id,
            attn_backend=attn_backend,
            sliding_window_size=128 if layer_id % 2 == 0 else None,
            rope=rope,
        )
        routed_experts_cfg = _make_gptoss_experts_config(
            dim=dim,
            hidden_dim=hidden_dim,
            num_experts=num_experts,
            layer_id=layer_id,
            top_k=top_k,
        )
        moe_cfg = MoE.Config(
            num_experts=num_experts,
            load_balance_coeff=load_balance_coeff,
            routed_experts=routed_experts_cfg,
            router=TokenChoiceTopKRouter.Config(
                num_experts=num_experts,
                score_func=Softmax.Config(),
                route_norm=True,
                gate=HiMidLoLinear.Config(
                    in_features=dim,
                    out_features=num_experts,
                    backward_mode="hi_mid_lo",
                    bias=True,
                    param_init=_depth_init(layer_id),
                ),
                top_k=top_k,
            ),
        )
        layer_cfg = GptOssTransformerBlock.Config(
            attention=attn_cfg,
            attention_norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
            ffn_norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
            moe=moe_cfg,
        )
        layers.append(layer_cfg)
    return layers


def _debugmodel(
    attn_backend: str = "varlen",
    *,
    seq_len: int,
) -> GptOssModel.Config:
    dim = 256
    hidden_dim = 2880
    n_layers = 4
    return GptOssModel.Config(
        max_context_length=seq_len,
        vocab_size=2048,
        dim=dim,
        tok_embeddings=Embedding.Config(
            num_embeddings=2048, embedding_dim=dim, param_init=_EMBEDDING_INIT
        ),
        norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
        lm_head=Linear.Config(
            in_features=dim,
            out_features=2048,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_gptoss_layers(
            dim=dim,
            n_layers=n_layers,
            hidden_dim=hidden_dim,
            num_experts=8,
            top_k=4,
            load_balance_coeff=1e-3,
            attn_backend=attn_backend,
            rope=CosSinRoPE.Config(
                dim=64,
                max_context_length=seq_len,
                theta=150000.0,
                scaling="yarn",
                rope_factor=32,
                beta_fast=32.0,
                beta_slow=1.0,
                truncate=False,
                original_seq_len=4096,
            ),
        ),
    )


def _20b(
    attn_backend: str = "varlen",
    *,
    seq_len: int,
) -> GptOssModel.Config:
    dim = 2880
    hidden_dim = 2880
    n_layers = 24
    return GptOssModel.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=201088,
        tok_embeddings=Embedding.Config(
            num_embeddings=201088, embedding_dim=dim, param_init=_EMBEDDING_INIT
        ),
        norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
        lm_head=Linear.Config(
            in_features=dim,
            out_features=201088,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_gptoss_layers(
            dim=dim,
            n_layers=n_layers,
            hidden_dim=hidden_dim,
            num_experts=32,
            top_k=4,
            load_balance_coeff=1e-3,
            attn_backend=attn_backend,
            rope=CosSinRoPE.Config(
                dim=64,
                max_context_length=seq_len,
                theta=150000.0,
                scaling="yarn",
                rope_factor=32,
                beta_fast=32.0,
                beta_slow=1.0,
                truncate=False,
                original_seq_len=4096,
            ),
        ),
    )


def _120b(
    attn_backend: str = "varlen",
    *,
    seq_len: int,
) -> GptOssModel.Config:
    dim = 2880
    hidden_dim = 2880
    n_layers = 36
    return GptOssModel.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=201088,
        tok_embeddings=Embedding.Config(
            num_embeddings=201088, embedding_dim=dim, param_init=_EMBEDDING_INIT
        ),
        norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
        lm_head=Linear.Config(
            in_features=dim,
            out_features=201088,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_gptoss_layers(
            dim=dim,
            n_layers=n_layers,
            hidden_dim=hidden_dim,
            num_experts=128,
            top_k=4,
            load_balance_coeff=1e-3,
            attn_backend=attn_backend,
            rope=CosSinRoPE.Config(
                dim=64,
                max_context_length=seq_len,
                theta=150000.0,
                scaling="yarn",
                rope_factor=32,
                beta_fast=32.0,
                beta_slow=1.0,
                truncate=False,
                original_seq_len=4096,
            ),
        ),
    )


MODEL_FLAVORS = {
    "debugmodel": (_debugmodel, 131072),
    "20b": (_20b, 131072),
    "120b": (_120b, 131072),
}


def build_model_config(
    flavor: str,
    *,
    seq_len: int | None = None,
    attn_backend: str = "varlen",
    converters: list[ModelConfigConverter.Config] | None = None,
) -> GptOssModel.Config:
    get_config, max_context_len = MODEL_FLAVORS[flavor]
    context_len = seq_len or max_context_len
    if context_len > max_context_len:
        raise ValueError(
            f"Requested seq_len {context_len} exceeds max context length "
            f"{max_context_len} for flavor {flavor}"
        )
    config = get_config(
        attn_backend=attn_backend,
        seq_len=context_len,
    )
    if converters is not None:
        validate_converter_compatibility(converters)
        for c in converters:
            config = c.build().convert(config)
    return config
