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

"""Muse Glimmer model configurations used by tests."""

from dataclasses import replace

from torchtitan.components.data import ConcatThenSplitPackingConfig, GrainDataLoader
from torchtitan.components.loss import ChunkedLossWrapper, CrossEntropyLoss
from torchtitan.components.optim import (
    AdamW,
    LRSchedulersContainer,
    Optim,
    OptimizersContainer,
)
from torchtitan.components.tokenizer import MultiModalTokenizer
from torchtitan.config import TrainingConfig
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed.activation_checkpoint import FullAC, SelectiveAC
from torchtitan.hf_datasets.text_datasets import DATASETS
from torchtitan.models.common.config_utils import (
    decoder_vocab_size,
    DEFAULT_DEBUG_MODEL_SEQ_LEN,
)

from torchtitan.models.muse_glimmer import build_model_config
from torchtitan.models.muse_glimmer.model import MuseGlimmerModel
from torchtitan.observability.metrics import MetricsProcessor
from torchtitan.trainer import Trainer


# Multimodal special tokens for the Muse Glimmer debug flavor. The debug tokenizer asset
# (``./tests/assets/tokenizer``) already defines these (IDs 2004-2008, all <
# debugmodel_mm's vocab_size=2048), so no asset change is needed.
MUSE_GLIMMER_SPECIAL_TOKENS = {
    "image_token": "<|image_pad|>",
    "video_token": "<|video_pad|>",
    "vision_start_token": "<|vision_start|>",
    "vision_end_token": "<|vision_end|>",
    "pad_token": "<|endoftext|>",
}


def _muse_glimmer_mm_dataloader(
    model_config: MuseGlimmerModel.Config, dataset_name: str
) -> GrainDataLoader.Config:
    """Build the shared multimodal dataloader config, taking the vision-patch
    geometry from the model's own vision encoder.

    ``patch_size``/``temporal_patch_size``/``spatial_merge_size`` must match the
    encoder exactly: they drive the image-placeholder token count the shared
    dataset inserts, which must equal the encoder's downsampled output token
    count. Deriving them from ``model_config`` keeps loader and encoder aligned.
    ``patch_order="raster"`` matches the encoder's raster patch layout (row-major
    grid); ``build_mrope_positions=False`` since Muse Glimmer uses 1D ComplexRoPE
    on the LLM side, not MRoPE.

    NOTE: the multimodal dataloader imports (``MultiModalCollator`` /
    ``resize_to_pixel_budget``) are done lazily here because they pull in
    ``torchvision`` (via ``torchtitan.hf_datasets.multimodal.utils.image``).
    ``torchvision`` is an optional dependency (not in requirements.txt /
    pyproject.toml), so importing it at module top level breaks the text-only
    configs (``muse_glimmer_debugmodel`` / ``muse_glimmer_30b``) for users who
    have not installed it. Keeping these imports inside the multimodal-only path
    lets the text configs load without torchvision.
    """
    from torchtitan.hf_datasets.multimodal.mm_collator import MultiModalCollator
    from torchtitan.hf_datasets.multimodal.mm_datasets import (
        MM_DATASETS,
        VisionProcessor,
    )
    from torchtitan.hf_datasets.multimodal.utils.image import resize_to_pixel_budget

    encoder = model_config.vision_encoder
    if encoder is None:
        raise ValueError("Multimodal Muse Glimmer must own a vision encoder")

    base_dataset = MM_DATASETS[dataset_name]
    base_processor = base_dataset.processor
    if not isinstance(base_processor, VisionProcessor.Config):
        raise ValueError(
            f"Multimodal dataset {dataset_name!r} must use VisionProcessor.Config"
        )
    processor = replace(
        base_processor,
        patch_size=encoder.patch_size,
        temporal_patch_size=encoder.patch_temporal,
        spatial_merge_size=encoder.downsample_factor,
        resize_fn=resize_to_pixel_budget,
        min_pixels=784,
        max_pixels=3136,
        image_mean=(0.5, 0.5, 0.5),
        image_std=(0.5, 0.5, 0.5),
        max_patches=4096,
        max_patches_per_side=512,
    )
    dataset = replace(base_dataset, processor=processor)

    return GrainDataLoader.Config(
        dataset=dataset,
        collator=MultiModalCollator.Config(
            max_images_per_microbatch=8,
            patch_size=processor.patch_size,
            temporal_patch_size=processor.temporal_patch_size,
            spatial_merge_size=processor.spatial_merge_size,
            patch_order="raster",
            build_mrope_positions=False,
        ),
    )


def muse_glimmer_debugmodel(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    model_config = build_model_config(
        "debugmodel", seq_len=seq_len, attn_backend="flex"
    )
    # The output soft-cap lives in the SoftCappedLinear lm_head, so it is applied
    # per-chunk inside ChunkedLossWrapper just as it would be in the full model
    # forward.
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./tests/assets/tokenizer",
        metrics=MetricsProcessor.Config(log_freq=1),
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4_test"]),
            shuffle=False,
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(
                warmup_steps=2,
                decay_ratio=0.8,
                decay_type="linear",
                min_lr_factor=0.0,
            ),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=8 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=10,
        ),
        parallelism=ParallelismConfig(),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def muse_glimmer_debugmodel_mm(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    """Multimodal debug training config.

    Trains the ``debugmodel_mm`` flavor (debug text decoder that owns a
    scaled-down vision encoder + adapter) end-to-end on the ``cc12m-test`` local
    tar fixture. The shared Grain data pipeline emits packed ``pixel_values`` +
    ``grid_thw`` + ``special_tokens``; preprocessing builds packed-bank
    indices from the image placeholder token. Vision-placeholder positions are
    already ``IGNORE_INDEX`` in the labels, so a standard ``CrossEntropyLoss``
    (wrapped in ``ChunkedLossWrapper``) is used.

    The integration smoke suite covers the TP+CP+PP+SP path.
    """
    mm_model_spec = build_model_config(
        "debugmodel_mm", seq_len=seq_len, attn_backend="flex"
    )
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(mm_model_spec),
            ),
        ),
        hf_assets_path="./tests/assets/tokenizer",
        tokenizer=MultiModalTokenizer.Config(**MUSE_GLIMMER_SPECIAL_TOKENS),
        metrics=MetricsProcessor.Config(log_freq=1),
        model=mm_model_spec,
        dataloader=_muse_glimmer_mm_dataloader(mm_model_spec, "cc12m-test"),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(
                warmup_steps=2,
                decay_ratio=0.8,
                decay_type="linear",
                min_lr_factor=0.0,
            ),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=4 * mm_model_spec.max_context_length,
            max_context_length=mm_model_spec.max_context_length,
            steps=10,
            disable_cuda_graphs=True,
        ),
        parallelism=ParallelismConfig(),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def muse_glimmer_30b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("30B", seq_len=seq_len, attn_backend="flex")
    return Trainer.Config(
        # ChunkedLossWrapper avoids materializing the full [T, vocab] logits;
        # the soft-cap is in the SoftCappedLinear lm_head, so it is still applied
        # per-chunk.
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Muse-Glimmer-30B",
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=3e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=200),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=1 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=1000,
        ),
        parallelism=ParallelismConfig(
            data_parallel_shard_degree=-1,
            tensor_parallel_degree=1,
            context_parallel_degree=1,
            pipeline_parallel_degree=1,
        ),
        checkpointer=None,
        activation_checkpoint=FullAC.Config(),
    )


def muse_glimmer_30b_mm(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("30B_mm", seq_len=seq_len, attn_backend="flex")
    return Trainer.Config(
        # ChunkedLossWrapper avoids materializing the full [T, vocab] logits;
        # the soft-cap is in the SoftCappedLinear lm_head, so it is still applied
        # per-chunk.
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Muse-Glimmer-30B",
        tokenizer=MultiModalTokenizer.Config(**MUSE_GLIMMER_SPECIAL_TOKENS),
        model=model_config,
        dataloader=_muse_glimmer_mm_dataloader(model_config, "cc12m"),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=3e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=200),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=1 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=1000,
            disable_cuda_graphs=True,
        ),
        parallelism=ParallelismConfig(
            data_parallel_shard_degree=-1,
            tensor_parallel_degree=1,
            context_parallel_degree=1,
            pipeline_parallel_degree=1,
        ),
        checkpointer=None,
        activation_checkpoint=FullAC.Config(),
    )
