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

"""Llama 3 model flavors."""

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 (
    ComplexRoPE,
    compute_ffn_hidden_dim,
    Embedding,
    Linear,
    RMSNorm,
    RoPE,
    TransformerBlock,
)
from torchtitan.models.common.config_utils import (
    get_attention_config,
    make_ffn_config,
    make_gqa_config,
)
from torchtitan.models.common.param_init import depth_scaled_std, skip_param_init

from .model import Llama3Model, Llama3TransformerBlock

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


_LINEAR_INIT = {
    "weight": partial(nn.init.trunc_normal_, std=0.02),
    "bias": nn.init.zeros_,
}
_NORM_INIT = {"weight": nn.init.ones_}
_EMBEDDING_INIT = {"weight": partial(nn.init.normal_, std=1.0)}
_EMBEDDING_SKIP_INIT = {"weight": skip_param_init}


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 _build_llama3_layers(
    *,
    n_layers: int,
    dim: int,
    n_heads: int,
    hidden_dim: int,
    rope: RoPE.Config,
    n_kv_heads: int | None = None,
    attn_backend: str,
) -> list[TransformerBlock.Config]:
    """Build a list of per-layer TransformerBlock configs with depth-scaled inits."""
    inner_attention = get_attention_config(attn_backend)
    layers = []
    for layer_id in range(n_layers):
        layers.append(
            Llama3TransformerBlock.Config(
                attention_norm=RMSNorm.Config(
                    normalized_shape=dim, param_init=_NORM_INIT
                ),
                ffn_norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
                attention=make_gqa_config(
                    dim=dim,
                    n_heads=n_heads,
                    n_kv_heads=n_kv_heads,
                    wqkv_param_init=_LINEAR_INIT,
                    wo_param_init=_depth_init(layer_id),
                    inner_attention=inner_attention,
                    rope=rope,
                ),
                feed_forward=make_ffn_config(
                    dim=dim,
                    hidden_dim=hidden_dim,
                    w13_param_init=_LINEAR_INIT,
                    w2_param_init=_depth_init(layer_id),
                ),
            )
        )
    return layers


def _debugmodel(
    attn_backend: str,
    *,
    seq_len: int,
    n_heads: int = 16,
) -> Llama3Model.Config:
    dim = 256
    n_layers = 6
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=2048,
        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_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            hidden_dim=compute_ffn_hidden_dim(dim, multiple_of=256),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


def _1b(
    attn_backend: str,
    *,
    seq_len: int,
) -> Llama3Model.Config:
    dim = 2048
    n_heads = 32
    n_kv_heads = 8
    n_layers = 16
    vocab_size = 128256
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=vocab_size,
        enable_weight_tying=True,
        tok_embeddings=Embedding.Config(
            num_embeddings=vocab_size,
            embedding_dim=dim,
            param_init=_EMBEDDING_SKIP_INIT,
        ),
        norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
        lm_head=Linear.Config(
            in_features=dim,
            out_features=vocab_size,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            n_kv_heads=n_kv_heads,
            hidden_dim=compute_ffn_hidden_dim(
                dim, multiple_of=1024, ffn_dim_multiplier=1.5
            ),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


def _3b(
    attn_backend: str,
    *,
    seq_len: int,
) -> Llama3Model.Config:
    dim = 3072
    n_heads = 24
    n_kv_heads = 8
    n_layers = 28
    vocab_size = 128256
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=vocab_size,
        enable_weight_tying=True,
        tok_embeddings=Embedding.Config(
            num_embeddings=vocab_size,
            embedding_dim=dim,
            param_init=_EMBEDDING_SKIP_INIT,
        ),
        norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
        lm_head=Linear.Config(
            in_features=dim,
            out_features=vocab_size,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            n_kv_heads=n_kv_heads,
            hidden_dim=compute_ffn_hidden_dim(
                dim, multiple_of=1024, ffn_dim_multiplier=1.0
            ),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


def _8b(
    attn_backend: str,
    *,
    seq_len: int,
) -> Llama3Model.Config:
    dim = 4096
    n_heads = 32
    n_kv_heads = 8
    n_layers = 32
    vocab_size = 128256
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=vocab_size,
        tok_embeddings=Embedding.Config(
            num_embeddings=vocab_size, 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=vocab_size,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            n_kv_heads=n_kv_heads,
            hidden_dim=compute_ffn_hidden_dim(
                dim, multiple_of=1024, ffn_dim_multiplier=1.3
            ),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


def _70b(
    attn_backend: str,
    *,
    seq_len: int,
) -> Llama3Model.Config:
    dim = 8192
    n_heads = 64
    n_kv_heads = 8
    n_layers = 80
    vocab_size = 128256
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=vocab_size,
        tok_embeddings=Embedding.Config(
            num_embeddings=vocab_size, 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=vocab_size,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            n_kv_heads=n_kv_heads,
            hidden_dim=compute_ffn_hidden_dim(
                dim, multiple_of=4096, ffn_dim_multiplier=1.3
            ),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


def _405b(
    attn_backend: str,
    *,
    seq_len: int,
) -> Llama3Model.Config:
    dim = 16384
    n_heads = 128
    n_kv_heads = 8
    n_layers = 126
    vocab_size = 128256
    return Llama3Model.Config(
        max_context_length=seq_len,
        dim=dim,
        vocab_size=vocab_size,
        tok_embeddings=Embedding.Config(
            num_embeddings=vocab_size, 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=vocab_size,
            param_init=_output_linear_init(dim),
        ),
        layers=_build_llama3_layers(
            n_layers=n_layers,
            dim=dim,
            n_heads=n_heads,
            n_kv_heads=n_kv_heads,
            hidden_dim=compute_ffn_hidden_dim(
                dim, multiple_of=4096, ffn_dim_multiplier=1.2
            ),
            rope=ComplexRoPE.Config(
                dim=dim // n_heads,
                max_context_length=seq_len,
                theta=500000,
                scaling="llama",
            ),
            attn_backend=attn_backend,
        ),
    )


MODEL_FLAVORS = {
    "debugmodel": (_debugmodel, 131072),
    # Preserve the debug model's dimensions and QKV GEMM shape, but use
    # 32-wide heads so MXFP8 weight-scale tiles align with head boundaries.
    "debugmodel_mxfp8": (partial(_debugmodel, n_heads=8), 131072),
    "1B": (_1b, 131072),
    "3B": (_3b, 131072),
    "8B": (_8b, 131072),
    "70B": (_70b, 131072),
    "405B": (_405b, 131072),
}


def build_model_config(
    flavor: str,
    *,
    seq_len: int | None = None,
    attn_backend: str = "flex",
    converters: list[ModelConfigConverter.Config] | None = None,
) -> Llama3Model.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
