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

"""Qwen3 model configurations used by tests."""

from typing import cast

from renderers import Message

from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.components.data import (
    ConcatThenSplitPackingConfig,
    FirstFitPackingConfig,
    GrainDataLoader,
    HuggingFaceRandomAccessSource,
    SingleDatasetConfig,
)
from torchtitan.components.loss import ChunkedLossWrapper, CrossEntropyLoss
from torchtitan.components.optim import (
    Adam,
    AdamW,
    LRSchedulersContainer,
    Optim,
    OptimizersContainer,
)
from torchtitan.config import TrainingConfig
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.config.transform import (
    apply_transforms,
    NVFP4LinearConverter,
    TokenDispatcherTransform,
)
from torchtitan.distributed.activation_checkpoint import FullAC, SelectiveAC
from torchtitan.hf_datasets.text_datasets import ChatProcessor, DATASETS
from torchtitan.models.common.config_utils import (
    decoder_vocab_size,
    DEFAULT_DEBUG_MODEL_SEQ_LEN,
)
from torchtitan.models.common.token_dispatcher import DeepEPTokenDispatcher

from torchtitan.models.qwen3 import build_model_config
from torchtitan.models.qwen3.model import Qwen3Model
from torchtitan.observability.metrics import MetricsProcessor
from torchtitan.quantization.nvfp4 import nvfp4_bf16_tail_fqns
from torchtitan.trainer import Trainer


def qwen3_debugmodel(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    model_config = build_model_config("debugmodel", seq_len=seq_len)
    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(enable_sequence_parallel=True),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def qwen3_debugmodel_nvfp4(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    config = qwen3_debugmodel(seq_len=seq_len)
    # Convert every decoder-layer Linear while leaving the lm_head in bf16.
    config.model = build_model_config(
        "debugmodel",
        seq_len=seq_len,
        converters=[
            NVFP4LinearConverter.Config(
                fqns=["layers"],
            ),
        ],
    )
    return config


def qwen3_debugmodel_first_85_pct_layers_nvfp4(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    config = qwen3_debugmodel(seq_len=seq_len)
    assert config.model is not None
    # Keep the last 15% of decoder layers and the lm_head in bf16.
    num_layers = len(cast(Qwen3Model.Config, config.model).layers)
    _NVFP4_BF16_TAIL_FRACTION = 0.15
    fqns = nvfp4_bf16_tail_fqns(
        num_layers,
        _NVFP4_BF16_TAIL_FRACTION,
    )
    config.model = build_model_config(
        "debugmodel",
        seq_len=seq_len,
        converters=[
            NVFP4LinearConverter.Config(
                fqns=fqns,
            ),
        ],
    )
    return config


def qwen3_debugmodel_moe_param_groups(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    config = qwen3_moe_debug(seq_len=seq_len)
    config.optim.optimizer = OptimizersContainer.Config(
        optimizers=[
            AdamW.Config(
                pattern=r"(?:tok_embeddings|output)\.",
                lr=8e-4,
                betas=(0.9, 0.95),
                eps=1e-8,
                weight_decay=0.0,
            ),
            Adam.Config(
                pattern=r"\.router\.gate\.",
                lr=1e-4,
                betas=(0.9, 0.95),
                eps=1e-8,
            ),
            AdamW.Config(
                pattern=r".*",
                lr=8e-4,
                betas=(0.9, 0.95),
                eps=1e-8,
                weight_decay=0.1,
            ),
        ],
    )
    return config


def qwen3_debugmodel_flex_flash(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    model_config = build_model_config(
        "debugmodel", seq_len=seq_len, attn_backend="flex_flash"
    )
    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,
        ),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def qwen3_0_6b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("0.6B", seq_len=seq_len)
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-0.6B",
        metrics=MetricsProcessor.Config(log_freq=1),
        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=2),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=4 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=10,
        ),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def qwen3_1_7b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("1.7B", seq_len=seq_len)
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-1.7B",
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=20),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=4 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=100,
        ),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def qwen3_8b_first_85_pct_layers_nvfp4(seq_len: int | None = None) -> Trainer.Config:
    def process_sample(sample):
        answer = sample["answer"]
        reasoning, final_answer = answer.rsplit("####", 1)
        return [
            {"role": "user", "content": sample["question"]},
            {
                "role": "assistant",
                "reasoning_content": reasoning.strip(),
                "content": final_answer.strip(),
            },
        ]

    model_config = build_model_config("8B", seq_len=seq_len, attn_backend="varlen")
    config = Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-8B",
        model=model_config,
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=2e-5)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(
                warmup_steps=15,
                decay_ratio=0.9,
                decay_type="cosine",
                min_lr_factor=0.1,
            ),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=180,
        ),
        dataloader=GrainDataLoader.Config(
            dataset=FirstFitPackingConfig(
                dataset=SingleDatasetConfig(
                    source=HuggingFaceRandomAccessSource.Config(
                        path="openai/gsm8k",
                        name="main",
                        split="train",
                    ),
                    processor=ChatProcessor.Config(messages_fn=process_sample),
                    post_filters=(lambda sample: sample is not None,),
                ),
            ),
            max_num_documents=32,
        ),
        metrics=MetricsProcessor.Config(enable_wandb=True),
        checkpointer=CheckpointManager.Config(initial_load_in_hf=True),
        activation_checkpoint=SelectiveAC.Config(),
    )
    assert config.model is not None
    # Keep the last 15% of decoder layers and the lm_head in bf16.
    num_layers = len(cast(Qwen3Model.Config, config.model).layers)
    _NVFP4_BF16_TAIL_FRACTION = 0.15
    fqns = nvfp4_bf16_tail_fqns(
        num_layers,
        _NVFP4_BF16_TAIL_FRACTION,
    )
    config.model = build_model_config(
        "8B",
        seq_len=seq_len,
        attn_backend="varlen",
        converters=[
            NVFP4LinearConverter.Config(
                fqns=fqns,
            ),
        ],
    )
    config.loss = ChunkedLossWrapper.Config(
        loss_fn=CrossEntropyLoss.Config(
            global_vocab_size=decoder_vocab_size(config.model),
        ),
    )
    return config


def qwen3_14b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("14B", seq_len=seq_len)
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-14B",
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=600),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=4 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=3000,
        ),
        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 qwen3_30b_a3b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("30B-A3B", seq_len=seq_len)
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-30B-A3B",
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=600),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=2 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=3000,
        ),
        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 qwen3_32b(seq_len: int | None = None) -> Trainer.Config:
    model_config = build_model_config("32B", seq_len=seq_len)
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-32B",
        model=model_config,
        dataloader=GrainDataLoader.Config(
            dataset=ConcatThenSplitPackingConfig(dataset=DATASETS["c4"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=8e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=600),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=2 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=3000,
        ),
        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 qwen3_moe_debug(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    model_config = build_model_config("debugmodel_moe", seq_len=seq_len)
    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=3e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=2),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=4 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=10,
        ),
        parallelism=ParallelismConfig(
            expert_parallel_degree=1,
        ),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )


def qwen3_moe_deepep(
    seq_len: int | None = DEFAULT_DEBUG_MODEL_SEQ_LEN,
) -> Trainer.Config:
    """Qwen3 debug MoE pretraining with the DeepEP v2 backend (compact training path), EP=4.

    The MoE expert dispatch uses the DeepEP v2 ElasticBuffer all-to-all; under autograd it
    takes the compact, host-synced, backward-able path. EP=4 (4 GPUs) so the dispatch is
    actually exercised (EP=1 falls back to local); the training shape determines the fixed
    per-rank buffer capacity. Numerics match the standard all-to-all backend (step-1 bitwise,
    reduction-order drift thereafter). Needs deep_ep v2 (ElasticBuffer) in the env.

    Local devgpu (no RDMA NIC) needs these env vars so the ElasticBuffer inits NVLink-only:
      - EP_DISABLE_GIN=1            skip the NCCL GIN / RDMA requirement (no RDMA NIC)
      - EP_REUSE_NCCL_COMM=0        avoid the ElasticBuffer null-device-comm segfault
      - NVSHMEM_REMOTE_TRANSPORT=none + NVSHMEM_DISABLE_MNNVL=1   intra-node NVLink only
      - LD_LIBRARY_PATH must include the deep_ep wheels' nvshmem + nccl lib dirs
    Then launch with NGPU=4 ./run_train.sh (none of this is needed on RDMA/RoCE hosts).
    """
    model_config = build_model_config("debugmodel_moe", seq_len=seq_len)
    config = 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"]),
        ),
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=3e-4)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(warmup_steps=2),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=2 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=10,
            disable_cuda_graphs=True,
        ),
        parallelism=ParallelismConfig(expert_parallel_degree=4),
        checkpointer=None,
        activation_checkpoint=SelectiveAC.Config(),
    )
    return apply_transforms(
        config,
        [TokenDispatcherTransform(dispatcher=DeepEPTokenDispatcher)],
    )


def sft_qwen3_8b_math(seq_len: int | None = None) -> Trainer.Config:
    """Qwen3-8B SFT on GSM8K math dataset."""

    def process_sample(sample) -> list[Message]:
        answer = sample["answer"]
        reasoning, final_answer = answer.rsplit("####", 1)
        return [
            {"role": "user", "content": sample["question"]},
            {
                "role": "assistant",
                "reasoning_content": reasoning.strip(),
                "content": final_answer.strip(),
            },
        ]

    model_config = build_model_config("8B", seq_len=seq_len, attn_backend="varlen")
    return Trainer.Config(
        loss=ChunkedLossWrapper.Config(
            loss_fn=CrossEntropyLoss.Config(
                global_vocab_size=decoder_vocab_size(model_config),
            ),
        ),
        hf_assets_path="./assets/hf/Qwen3-8B",
        model=model_config,
        optim=Optim.Config(
            optimizer=OptimizersContainer.Config(
                optimizers=[AdamW.Config(pattern=r".*", lr=2e-5)]
            ),
            lr_scheduler=LRSchedulersContainer.Config(
                warmup_steps=15,
                decay_ratio=0.9,
                decay_type="cosine",
                min_lr_factor=0.1,
            ),
        ),
        training=TrainingConfig(
            num_tokens_per_microbatch_per_dp_rank=1 * model_config.max_context_length,
            max_context_length=model_config.max_context_length,
            steps=180,
        ),
        dataloader=GrainDataLoader.Config(
            dataset=FirstFitPackingConfig(
                dataset=SingleDatasetConfig(
                    source=HuggingFaceRandomAccessSource.Config(
                        path="openai/gsm8k",
                        name="main",
                        split="train",
                    ),
                    processor=ChatProcessor.Config(messages_fn=process_sample),
                    post_filters=(lambda sample: sample is not None,),
                ),
            ),
            max_num_documents=32,
        ),
        metrics=MetricsProcessor.Config(
            enable_wandb=True,
        ),
        checkpointer=CheckpointManager.Config(
            initial_load_in_hf=True,
        ),
        activation_checkpoint=SelectiveAC.Config(),
    )
