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

from typing import TYPE_CHECKING

import spmd_types as spmd

from torchtitan.models.common.attention import GQAttention

from torchtitan.models.common.decoder_sharding import (
    attention_activation_placement,
    dense_activation_placement,
    dense_param_placement,
    dense_sequence_parallel_placement,
    norm_config,
    set_decoder_sharding_config,
    set_dense_ffn_sharding,
    set_gqa_attention_sharding,
    set_gqa_inner_attention_local_spmd,
)
from torchtitan.models.common.moe_sharding import set_moe_sharding_config
from torchtitan.protocols.sharding import ShardingConfig

if TYPE_CHECKING:
    from torchtitan.models.qwen3.model import Qwen3Model, Qwen3TransformerBlock


def set_qwen3_sharding_config(
    config: "Qwen3Model.Config",
    *,
    enable_sp: bool,
    enable_ep: bool,
) -> None:
    """Fill ``sharding_config`` on all Qwen3 sub-configs.

    Dense sub-configs (attention, norms, dense FFN) are populated
    unconditionally — ``Module.parallelize`` filters disabled axes
    at runtime.

    MoE sub-configs (router, shared experts, routed experts) are
    populated unconditionally — ``resolve_mesh`` filters disabled
    axes at runtime.
    """

    set_decoder_sharding_config(config, enable_sp=enable_sp)
    for layer_cfg in config.layers:
        _set_qwen3_layer_sharding(layer_cfg, enable_sp=enable_sp, enable_ep=enable_ep)


def _set_qwen3_layer_sharding(
    layer_cfg: "Qwen3TransformerBlock.Config",
    *,
    enable_sp: bool,
    enable_ep: bool,
) -> None:
    """Set sharding on one Qwen3 transformer layer.

    Attention and norms are sharded on all blocks (MoE and non-MoE).
    Dense FFN is only sharded on non-MoE blocks; MoE FFN is routed
    through ``set_moe_sharding_config``.
    """
    attention = layer_cfg.attention
    assert isinstance(attention, GQAttention.Config)

    norm = norm_config(enable_sp=enable_sp)
    layer_cfg.attention_norm.sharding_config = norm
    layer_cfg.ffn_norm.sharding_config = norm

    set_gqa_attention_sharding(attention, enable_sp=enable_sp)
    set_gqa_inner_attention_local_spmd(attention.inner_attention)

    # QK norms: shard on head dim (dim=1), independent of SP.
    if attention.qk_norm is not None:
        head_layout = attention_activation_placement()
        attention.qk_norm.sharding_config = ShardingConfig(
            state_shardings={"weight": dense_param_placement(tp=spmd.R)},
            in_src_shardings={"input": head_layout},
            in_dst_shardings={"input": head_layout},
            out_src_shardings=head_layout,
            out_dst_shardings=head_layout,
        )

    # Dense FFN (non-MoE layers only)
    if layer_cfg.feed_forward is not None:
        attn_x_layout = (
            dense_sequence_parallel_placement()
            if enable_sp
            else dense_activation_placement(tp=spmd.I, cp=spmd.S(0))
        )
        set_dense_ffn_sharding(
            layer_cfg.feed_forward,
            attn_x_layout=attn_x_layout,
            enable_sp=enable_sp,
        )

    # MoE FFN (MoE-enabled layers only).
    if layer_cfg.moe is not None:
        set_moe_sharding_config(
            layer_cfg.moe,
            enable_ep=enable_ep,
            enable_sp=enable_sp,
        )
