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

"""Reference (SGLang):
https://github.com/sgl-project/sglang/blob/e0c0c0a45cb1bda90392bfa2bba4184f5b0638a0/python/sglang/srt/models/kimi_k25.py
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Any, cast

import spmd_types as spmd
import torch
from torch import nn

from torchtitan.config import TrainingConfig
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig
from torchtitan.distributed.parallelism_context import ParallelismContext
from torchtitan.distributed.spmd_types import (
    annotate_input_spmd_types,
    spmd_local_context,
    spmd_mesh_size,
)
from torchtitan.models.common.attention import (
    AttentionMetadataMap,
    FlexInnerAttention,
    VarlenInnerAttention,
)
from torchtitan.models.common.decoder import TransformerBlock
from torchtitan.models.common.decoder_sharding import decoder_input_sharding
from torchtitan.models.common.multimodal import (
    add_zero_vision_dependency,
    build_dummy_vision_inputs,
    get_vision_positions,
    MultimodalModel,
    scatter_vision_embeds,
)
from torchtitan.models.common.vision_encoder_sharding import multimodal_input_sharding
from torchtitan.models.deepseek_v3.model import (
    DeepSeekV3Model,
    get_deepseek_v3_nparams_and_flops as get_kimi_k2_7_nparams_and_flops,
)
from torchtitan.models.deepseek_v3.mtp import MTPDecoder

from .state_dict_adapter import KimiK25StateDictAdapter
from .vision_encoder import KimiK25VisionEncoder


class KimiK25Model(MultimodalModel, MTPDecoder):
    state_dict_adapter_cls = KimiK25StateDictAdapter
    multimodal_encoder_fqns = ("vision_encoder",)

    @classmethod
    def _register_optimizer_hooks(
        cls, optimizers, model_parts, parallelism_context
    ) -> None:
        from torchtitan.models.common.moe import register_moe_load_balancing_hook
        from torchtitan.models.kimi_k2_7.qk_clip import register_qk_clip_hook

        register_moe_load_balancing_hook(optimizers, model_parts, parallelism_context)
        register_qk_clip_hook(optimizers, model_parts, parallelism_context)

    pipeline_first_stage_module_fqns = ("vision_encoder",)

    """Kimi K2.5: DeepSeekV3 language model with a MoonViT3d vision encoder.

    Forward pass flow::

        forward(tokens, pixel_values[/videos], grid_thw, ...)
          |
          +-- tok_embeddings(tokens)               -> text embeddings
          +-- vision_encoder(pixels)               -> packed vision features
          +-- scatter at vision placeholder runs   -> multimodal embeddings
          +-- decoder layers (MLA + MoE)           -> hidden states
          +-- norm -> lm_head                      -> logits
    """

    @dataclass(kw_only=True, slots=True)
    class Config(DeepSeekV3Model.Config):
        vision_encoder: KimiK25VisionEncoder.Config | None = None

        def get_nparams_and_flops(
            self, model: nn.Module, seq_len: int
        ) -> tuple[int, int]:
            kimi_model = cast("KimiK25Model", model)
            return get_kimi_k2_7_nparams_and_flops(
                self,
                model,
                seq_len,
                modules_excluded_from_active_params=(kimi_model.vision_encoder,),
            )

        def set_sharding_(self, parallelism: ParallelismConfig) -> None:
            from .sharding import set_kimi_k2_5_sharding_config

            set_kimi_k2_5_sharding_config(
                self,
                enable_sp=parallelism.enable_sequence_parallel,
                enable_ep=parallelism.expert_parallel_degree > 1,
            )

    def __init__(self, config: Config):
        tp = spmd_mesh_size("tp")
        if (
            tp > 1
            and config.vision_encoder is not None
            and config.vision_encoder.num_heads % tp != 0
        ):
            raise ValueError(
                f"tensor parallel degree ({tp}) must divide "
                f"vision num_heads ({config.vision_encoder.num_heads})."
            )
        super().__init__(config)
        self.vision_encoder = (
            config.vision_encoder.build() if config.vision_encoder is not None else None
        )

    def parallelize(
        self,
        *,
        parallelism_context: ParallelismContext,
        training: TrainingConfig,
        parallelism: ParallelismConfig,
        local_compile_regions: list[str],
        ac_config: ActivationCheckpointingConfig | None,
        dump_folder: str,
        skip_dp: bool = False,
    ) -> KimiK25Model:
        if parallelism_context.cp_enabled:
            raise NotImplementedError(
                "Context Parallel is not yet supported for Kimi K2.5: vision "
                "scatter needs the full sequence before CP would shard it."
            )

        return super().parallelize(
            parallelism_context=parallelism_context,
            training=training,
            parallelism=parallelism,
            local_compile_regions=local_compile_regions,
            ac_config=ac_config,
            dump_folder=dump_folder,
            skip_dp=skip_dp,
        )

    def preprocess_inputs(
        self,
        input_dict: dict[str, Any],
        *,
        parallelism_context: ParallelismContext,
        parallelism: ParallelismConfig,
        max_num_documents: int | None = None,
        max_context_length: int | None = None,
        **kwargs: Any,
    ) -> tuple[torch.Tensor, torch.Tensor, dict[str, Any]]:
        """Build masks, CP-shard, SPMD-wrap, and return the batch."""
        del kwargs
        positions = input_dict.get("positions", None)
        padding_mask = input_dict.get("padding_mask", None)
        if positions is not None:
            inner = getattr(self.config.first_base_attention, "inner_attention", None)
            if isinstance(
                inner, (FlexInnerAttention.Config, VarlenInnerAttention.Config)
            ):
                input_dict["attention_metadata"] = self._get_attention_metadata(
                    positions=positions,
                    padding_mask=padding_mask,
                    max_num_documents=max_num_documents,
                    max_context_length=max_context_length,
                )

        input_shardings = {
            **decoder_input_sharding(),
            **multimodal_input_sharding(),
        }
        if parallelism_context.cp_enabled:
            input_dict = self._cp_shard(
                input_dict,
                input_shardings=input_shardings,
                parallelism_context=parallelism_context,
                parallelism=parallelism,
            )
        input_dict = annotate_input_spmd_types(
            parallelism_context, input_dict, input_shardings
        )

        inputs = input_dict.pop("input")
        labels = input_dict.pop("labels")
        input_dict["aux_loss_denominators"] = None
        return inputs, labels, input_dict

    def _prepare_multimodal_embeds(
        self,
        tokens: torch.Tensor,
        *,
        pixel_values: torch.Tensor | None,
        grid_thw: torch.Tensor | None,
        pixel_values_videos: torch.Tensor | None = None,
        grid_thw_videos: torch.Tensor | None = None,
        special_tokens: dict[str, int] | None,
    ) -> torch.Tensor:
        """Embed tokens, run the vision encoder, scatter features into text.

        With kimi's single unified placeholder, a one-modality batch's runs map
        to visual items in order. Mixing images and videos in one batch is not
        yet supported (see the TODO below).
        """
        inputs_embeds = (
            self.tok_embeddings(tokens) if self.tok_embeddings is not None else tokens
        )

        modalities = []
        if pixel_values is not None and grid_thw is not None:
            modalities.append((pixel_values, grid_thw))
        if pixel_values_videos is not None and grid_thw_videos is not None:
            modalities.append((pixel_values_videos, grid_thw_videos))

        is_dummy = not modalities
        if is_dummy:
            if self.vision_encoder is None:
                return inputs_embeds
            kernel_h, kernel_w = self.vision_encoder.merge_kernel_size
            modalities.append(
                build_dummy_vision_inputs(
                    patch_dim=self.vision_encoder.patch_embed.in_features,
                    grid_thw=(1, kernel_h, kernel_w),
                    device=inputs_embeds.device,
                )
            )
        # TODO: support mixed image+video batches. Upstream fix: when
        # image_id == video_id, emit one document-ordered vision stream so the
        # runs stay modality-agnostic and this branch goes away.
        assert len(modalities) == 1, "mixed image+video batches not yet supported"
        pixels, grid = modalities[0]
        # A non-empty modalities list means a multimodal run, so the vision
        # encoder is present (text-only configs never populate pixels).
        assert self.vision_encoder is not None

        # Patches arrive float32; match the encoder's compute dtype for the matmul.
        pixels = pixels.to(self.vision_encoder.patch_embed.weight.dtype)
        vision_embeds = self.vision_encoder(pixels, grid_thw=grid)
        if is_dummy:
            return add_zero_vision_dependency(inputs_embeds, vision_embeds)

        if special_tokens is None:
            raise ValueError("special_tokens are required for multimodal inputs.")
        placeholder_id = special_tokens["image_id"]
        assert placeholder_id == special_tokens["video_id"]
        # MoonViT collapses time (temporal pooling) and merges 2x2 spatially, so
        # the token count is (h/kh)*(w/kw), independent of t.
        kh, kw = self.vision_encoder.merge_kernel_size
        num_tokens_per_item = (grid[:, 1] // kh) * (grid[:, 2] // kw)
        vision_positions = get_vision_positions(
            tokens, num_tokens_per_item, placeholder_id
        )
        if vision_positions:
            inputs_embeds = scatter_vision_embeds(
                inputs_embeds,
                vision_embeds=vision_embeds,
                vision_positions=vision_positions,
            )
        return inputs_embeds

    def forward(  # pyrefly: ignore [bad-override]
        self,
        tokens: torch.Tensor,
        *,
        pixel_values: torch.Tensor | None = None,
        grid_thw: torch.Tensor | None = None,
        pixel_values_videos: torch.Tensor | None = None,
        grid_thw_videos: torch.Tensor | None = None,
        special_tokens: dict[str, int] | None = None,
        attention_metadata: AttentionMetadataMap | None = None,
        positions: torch.Tensor | None = None,
        padding_mask: torch.Tensor | None = None,
        aux_loss_denominators: torch.Tensor | None = None,
    ):
        """Forward pass for Kimi K2.5.

        Images and videos share one unified ``<|media_pad|>`` placeholder.

        Args:
            tokens: ``(num_tokens,)`` packed token IDs.
            pixel_values: ``(total_num_patches, patch_dim)`` packed image
                patches, or None for text-only / video-only batches.
            grid_thw: (num_images, 3) patch counts ``[t, h, w]`` per image.
            pixel_values_videos: Packed video patches, or None (mixing with
                ``pixel_values`` in one batch is not yet supported).
            grid_thw_videos: (num_videos, 3) patch counts per video.
            special_tokens: tokenizer-resolved ``image_id``/``video_id``;
                required for image/video batches, None for text-only.
            attention_metadata: Decoder attention masks.
            positions: Per-token position IDs for packed sequences.

        Returns:
            ``(num_tokens, vocab_size)`` logits.
        """
        with spmd_local_context("dp"):
            if self.tok_embeddings is not None:
                x = self._prepare_multimodal_embeds(
                    tokens,
                    pixel_values=pixel_values,
                    grid_thw=grid_thw,
                    pixel_values_videos=pixel_values_videos,
                    grid_thw_videos=grid_thw_videos,
                    special_tokens=special_tokens,
                )
            else:
                x = tokens

        if spmd.is_type_checking():
            # The scatter restores a token-aligned tensor, so text-model DP
            # resumes as global batch sharding after the multimodal region.
            spmd.assert_type(x, {"dp": spmd.S(0), "tp": spmd.R})

        with spmd.no_typecheck():
            aux_loss_denominator = (
                None if aux_loss_denominators is None else aux_loss_denominators[0]
            )
        for layer in self.layers.values():
            layer_attention_metadata = (
                None
                if attention_metadata is None
                else attention_metadata.get(
                    cast(TransformerBlock, layer).attention.attention_metadata_key
                )
            )
            x = layer(
                x,
                layer_attention_metadata,
                positions,
                padding_mask=padding_mask,
                aux_loss_denominator=aux_loss_denominator,
            )

        x = self.norm(x) if self.norm is not None else x
        if self._skip_lm_head:
            return x
        return self.lm_head(x) if self.lm_head is not None else x
