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

"""Configurations for the feature integration-test suite."""

import logging
import os
from collections.abc import Iterator
from dataclasses import dataclass, fields, replace
from typing import Any, cast

import torch
import torch.distributed as dist
from renderers import Message, Qwen3RendererConfig
from torch.distributed.tensor import DTensor
from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.components.data import (
    FirstFitPackingConfig,
    GrainDataLoader,
    SingleDatasetConfig,
)
from torchtitan.components.data.types import TrainingMicrobatch
from torchtitan.components.optim import AdamW, OptimizersContainer
from torchtitan.components.renderer import from_renderers
from torchtitan.components.validate import Validator
from torchtitan.config.transform import apply_transforms, ContextParallelTransform

from torchtitan.distributed.activation_checkpoint import FullAC, SelectiveAC
from torchtitan.hf_datasets.text_datasets import ChatProcessor

from torchtitan.models.common.attention import FlexInnerAttention, VarlenInnerAttention
from torchtitan.models.common.attention.cp_attention import (
    KVAllGatherCPFlexInnerAttention,
    UlyssesCPFlexInnerAttention,
    UlyssesCPVarlenInnerAttention,
)
from torchtitan.observability.sdc_replayer import SDCReplayer, SDCReplayMismatch
from torchtitan.protocols import BaseModel
from torchtitan.trainer import Trainer
from torchtitan.training_engine import TrainingEngine

from torchtitan_recipes.tests.models.deepseek_v3 import deepseek_v3_debugmodel
from torchtitan_recipes.tests.models.llama3 import (
    llama3_debugmodel,
    llama3_debugmodel_ce_loss,
    llama3_debugmodel_lora,
    llama3_debugmodel_varlen_attn,
    sft_debugmodel,
)
from torchtitan_recipes.tests.models.muse_glimmer import muse_glimmer_debugmodel

from .. import _set_spmd_typechecking


logger = logging.getLogger(__name__)


class SDCReplayMismatchTrainingEngine(TrainingEngine):
    """Inject a gradient mismatch into the first SDC replay execution."""

    def __init__(
        self,
        config: TrainingEngine.Config,
        *,
        model_config: BaseModel.Config,
        max_num_documents: int | None,
        output_dir: str,
    ) -> None:
        super().__init__(
            config,
            model_config=model_config,
            max_num_documents=max_num_documents,
            output_dir=output_dir,
        )
        self._num_forward_backward_calls = 0

    def _non_pp_forward_backward_microbatch(
        self,
        *,
        inputs: torch.Tensor | tuple[torch.Tensor, ...],
        labels: torch.Tensor | tuple[torch.Tensor, ...],
        model_kwargs: dict[str, Any],
        loss_kwargs: dict[str, Any],
    ) -> torch.Tensor:
        loss = super()._non_pp_forward_backward_microbatch(
            inputs=inputs,
            labels=labels,
            model_kwargs=model_kwargs,
            loss_kwargs=loss_kwargs,
        )
        self._num_forward_backward_calls += 1
        if self._num_forward_backward_calls != 2 or dist.get_rank() != 0:
            return loss

        with torch.no_grad():
            for model_part in self.model_parts:
                for parameter in model_part.parameters():
                    if parameter.grad is None:
                        continue
                    grad = parameter.grad
                    local_grad = grad.to_local() if isinstance(grad, DTensor) else grad
                    if local_grad.numel() > 0:
                        # Corrupt the first replay before SDC captures its signature.
                        local_grad[(0,) * local_grad.ndim].add_(1)
                        return loss
        raise AssertionError("Could not find a local gradient to corrupt.")


class SDCReplayMismatchTrainer(Trainer):
    """Inject a replay-only gradient mismatch and verify it is fatal."""

    @dataclass(kw_only=True, slots=True)
    class Config(Trainer.Config):
        pass

    engine_cls = SDCReplayMismatchTrainingEngine
    engine: SDCReplayMismatchTrainingEngine

    def train_step(
        self,
        data_iterator: Iterator[TrainingMicrobatch],
    ) -> None:
        try:
            super().train_step(data_iterator)
        except SDCReplayMismatch as error:
            assert error.step == 1
            assert error.local_step == 1
            assert error.replay == 1
            assert error.rank == 0
            assert error.signature_mismatch is not None
            assert error.signature_mismatch.startswith("gradient:0:")
            assert self.engine.sdc_replayer is not None
            assert self.engine.sdc_replayer.steps_since_reset == 0
            logger.info("Detected expected %s", error)
            self.engine.num_completed_steps = error.step
            return
        raise AssertionError("Expected SDC replay to detect the injected mismatch.")


def deepseek_v3_debugmodel_sdc_replay_mismatch() -> Trainer.Config:
    base_config = deepseek_v3_debugmodel(seq_len=2048)
    config = SDCReplayMismatchTrainer.Config(
        **{
            config_field.name: getattr(base_config, config_field.name)
            for config_field in fields(base_config)
            if config_field.init
        }
    )
    config.debug.deterministic = True
    config.debug.seed = 42
    config.training.disable_cuda_graphs = True
    config.training.steps = 1
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.expert_parallel_degree = 2
    config.sdc_replayer = SDCReplayer.Config()
    return config


def llama3_debugmodel_sdc_replay_cuda_graph() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    config.debug.deterministic = True
    config.debug.seed = 42
    config.training.steps = 3
    config.sdc_replayer = SDCReplayer.Config()
    return config


def llama3_debugmodel_default() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.profiler.enable_profiling = True
    config.metrics.enable_tensorboard = True
    return config


def llama3_debugmodel_tp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.tensor_parallel_degree = 2
    return config


def llama3_debugmodel_ce_loss_tp2() -> Trainer.Config:
    config = llama3_debugmodel_ce_loss(seq_len=2048)
    # Non-chunked CE loss does not pass SPMD type checking yet.
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.tensor_parallel_degree = 2
    return config


def llama3_debugmodel_tp2_no_sp() -> Trainer.Config:
    config = llama3_debugmodel_tp2()
    config.parallelism.enable_sequence_parallel = False
    return config


def llama3_debugmodel_full_checkpoint_save() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.checkpointer = CheckpointManager.Config(
        interval=10,
        last_save_model_only=False,
    )
    return config


def llama3_debugmodel_full_checkpoint_load() -> Trainer.Config:
    config = llama3_debugmodel_full_checkpoint_save()
    config.training.steps = 20
    return config


def llama3_debugmodel_hf_checkpoint_save() -> Trainer.Config:
    config = llama3_debugmodel_full_checkpoint_save()
    config.checkpointer.folder = "hf_checkpoint"
    config.checkpointer.last_save_model_only = True
    config.checkpointer.last_save_in_hf = True
    return config


def llama3_debugmodel_hf_checkpoint_load() -> Trainer.Config:
    """Loads what ``llama3_debugmodel_hf_checkpoint_save`` wrote.

    The integration runner supplies the per-test output directory to each run.
    """
    config = llama3_debugmodel_full_checkpoint_save()
    test_output_dir = os.getenv(
        "TORCHTITAN_TEST_OUTPUT_DIR",
        os.path.join(
            os.getenv("RUNNER_TEMP", ""),
            "artifacts-to-be-uploaded/model_only_hf_checkpoint",
        ),
    )
    config.checkpointer.initial_load_path = os.path.join(
        test_output_dir,
        "hf_checkpoint/step-10/",
    )
    config.checkpointer.initial_load_model_only = True
    config.checkpointer.initial_load_in_hf = True
    return config


def llama3_debugmodel_last_save_model_only_bf16() -> Trainer.Config:
    config = llama3_debugmodel_full_checkpoint_save()
    config.checkpointer.last_save_model_only = True
    config.checkpointer.export_dtype = "bfloat16"
    return config


def llama3_debugmodel_pp2_1f1b() -> Trainer.Config:
    """PP-only, so it leaves SPMD type checking off.

    Type checking needs at least one SPMD axis greater than 1; collapsing
    every dense SPMD axis to size 1 trips DTensor's rejection of a Shard
    placement on a degenerate axis.
    """
    config = llama3_debugmodel(seq_len=2048)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "1F1B"
    config.parallelism.data_parallel_shard_degree = 1
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def llama3_debugmodel_fsdp2_pp2_1f1b() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "1F1B"
    config.parallelism.data_parallel_shard_degree = 2
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def muse_glimmer_debugmodel_fsdp2_pp2_deferred_gradient_reduction() -> Trainer.Config:
    config = muse_glimmer_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "Interleaved1F1B"
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.fsdp_defer_gradient_reduction = True
    config.parallelism.fsdp_reshard_after_forward = "never"
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.training.num_tokens_per_train_step = 65536
    return config


def _muse_glimmer_debugmodel_pp2_max_outstanding_sends(
    schedule: str,
) -> Trainer.Config:
    config = muse_glimmer_debugmodel(seq_len=2048)
    config.comm.backend = "real_pp_fake_spmd"
    config.debug.deterministic = True
    config.debug.seed = 42
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.data_parallel_shard_degree = 1
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = schedule
    config.parallelism.pipeline_parallel_max_outstanding_sends = 2
    config.activation_checkpoint = None
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.training.disable_cuda_graphs = True
    return config


def muse_glimmer_debugmodel_pp2_looped_bfs_send_budget() -> Trainer.Config:
    return _muse_glimmer_debugmodel_pp2_max_outstanding_sends("LoopedBFS")


def muse_glimmer_debugmodel_pp2_interleaved_1f1b_send_budget() -> Trainer.Config:
    return _muse_glimmer_debugmodel_pp2_max_outstanding_sends("Interleaved1F1B")


def muse_glimmer_debugmodel_fsdp2_pp2_optimizer_cuda_graph() -> Trainer.Config:
    config = muse_glimmer_debugmodel_fsdp2_pp2_deferred_gradient_reduction()
    config.optim.enable_cuda_graph = True
    return config


def muse_glimmer_debugmodel_fsdp2_optimizer_cuda_graph() -> Trainer.Config:
    config = muse_glimmer_debugmodel_fsdp2_deferred_gradient_reduction()
    config.optim.enable_cuda_graph = True
    return config


def llama3_debugmodel_fsdp2_pp2_1f1b_layers_per_stage() -> Trainer.Config:
    config = llama3_debugmodel_fsdp2_pp2_1f1b()
    config.parallelism.pipeline_parallel_layers_per_stage = 4
    return config


def llama3_debugmodel_tp2_pp2_gpipe() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "GPipe"
    config.parallelism.tensor_parallel_degree = 2
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def llama3_debugmodel_fsdp2_tp2_pp2_save() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.checkpointer = CheckpointManager.Config(
        interval=10,
        last_save_model_only=False,
    )
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.tensor_parallel_degree = 2
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def llama3_debugmodel_fsdp2_tp2_pp2_load() -> Trainer.Config:
    config = llama3_debugmodel_fsdp2_tp2_pp2_save()
    config.training.steps = 20
    return config


def llama3_debugmodel_fsdp2_tp2_pp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.tensor_parallel_degree = 2
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def llama3_debugmodel_pp4_interleaved_1f1b() -> Trainer.Config:
    """PP-only; see ``llama3_debugmodel_pp2_1f1b`` for why type checking is off."""
    config = llama3_debugmodel(seq_len=2048)
    config.parallelism.pipeline_parallel_degree = 4
    config.parallelism.num_pp_microbatches = 8
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.debug.deterministic = True
    config.debug.seed = 42
    return config


def llama3_debugmodel_pp4_interleaved_1f1b_layers_per_stage() -> Trainer.Config:
    config = llama3_debugmodel_pp4_interleaved_1f1b()
    config.parallelism.pipeline_parallel_layers_per_stage = 1
    return config


def llama3_debugmodel_pp4_zero_bubble() -> Trainer.Config:
    config = llama3_debugmodel_varlen_attn(seq_len=2048)
    config.parallelism.pipeline_parallel_degree = 4
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "InterleavedZeroBubble"
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.activation_checkpoint = FullAC.Config()
    config.debug.deterministic = True
    config.debug.seed = 42
    return config


def llama3_debugmodel_pp2_zbv() -> Trainer.Config:
    config = llama3_debugmodel_varlen_attn(seq_len=2048)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "ZBVZeroBubble"
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.activation_checkpoint = FullAC.Config()
    config.debug.deterministic = True
    config.debug.seed = 42
    return config


def llama3_debugmodel_pp2_custom_csv() -> Trainer.Config:
    config = llama3_debugmodel_varlen_attn(seq_len=2048)
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.parallelism.pipeline_parallel_schedule = "PipelineScheduleMulti"
    config.parallelism.pipeline_parallel_schedule_csv = (
        "./tests/assets/custom_schedule.csv"
    )
    config.activation_checkpoint = FullAC.Config()
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.debug.deterministic = True
    config.debug.seed = 42
    return config


def muse_glimmer_debugmodel_optimizer_bf16_states() -> Trainer.Config:
    config = muse_glimmer_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.training.mixed_precision_reduce = "float32"
    config.optim.optimizer = OptimizersContainer.Config(
        optimizers=[AdamW.Config(pattern=r".*", moment_dtype="bfloat16")]
    )
    return config


def llama3_debugmodel_ddp4() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.data_parallel_shard_degree = 1
    config.parallelism.data_parallel_replicate_degree = 4
    return config


def llama3_debugmodel_hsdp2x2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.data_parallel_replicate_degree = 2
    return config


def llama3_debugmodel_cp4() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.context_parallel_degree = 4
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_ulysses_cp2() -> Trainer.Config:
    """Attention reshards the CP axis onto the head dimension."""
    config = llama3_debugmodel()
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.context_parallel_degree = 2
    # Head-sharded attention has no per-rank sequence imbalance to balance.
    config.parallelism.context_parallel_load_balancer = None
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={FlexInnerAttention: UlyssesCPFlexInnerAttention}
            )
        ],
    )


def llama3_debugmodel_ulysses_cp2_varlen() -> Trainer.Config:
    """Llama 3 with varlen Ulysses CP."""
    config = llama3_debugmodel_varlen_attn()
    # Packed varlen metadata lacks SPMD annotations.
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.context_parallel_degree = 2
    # Ulysses does not support token reordering.
    config.parallelism.context_parallel_load_balancer = None
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    VarlenInnerAttention: UlyssesCPVarlenInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_hsdp2x2_tp2() -> Trainer.Config:
    config = llama3_debugmodel_hsdp2x2()
    config.parallelism.tensor_parallel_degree = 2
    return config


def llama3_debugmodel_fsdp2_cp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.data_parallel_shard_degree = 2
    config.parallelism.context_parallel_degree = 2
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_ddp2_cp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.data_parallel_shard_degree = 1
    config.parallelism.data_parallel_replicate_degree = 2
    config.parallelism.context_parallel_degree = 2
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_hsdp2x2_cp2() -> Trainer.Config:
    config = llama3_debugmodel_hsdp2x2()
    config.debug.spmd_typechecking = False
    config.parallelism.context_parallel_degree = 2
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_fsdp2_tp2_cp2() -> Trainer.Config:
    config = llama3_debugmodel_fsdp2_cp2()
    config.parallelism.tensor_parallel_degree = 2
    return config


def llama3_debugmodel_fsdp_reshard_always() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.fsdp_reshard_after_forward = "always"
    return config


def llama3_debugmodel_optional_checkpoint_save() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.checkpointer = CheckpointManager.Config(
        interval=10,
        last_save_model_only=False,
    )
    return config


def llama3_debugmodel_optional_checkpoint_load_tp2() -> Trainer.Config:
    """Loads a ``[dp:4]`` checkpoint at ``[dp:2, tp:2]``.

    The dataloader is excluded from loading to avoid errors caused by the
    mismatched dp degree.
    """
    config = llama3_debugmodel_optional_checkpoint_save()
    config.checkpointer.exclude_from_loading = [
        "lr_scheduler",
        "dataloader",
        "optimizer",
    ]
    config.parallelism.tensor_parallel_degree = 2
    config.training.steps = 20
    return config


def llama3_debugmodel_gradient_accumulation() -> Trainer.Config:
    """Two gradient accumulation steps on 2 GPUs."""
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.training.disable_cuda_graphs = True
    config.training.num_tokens_per_microbatch_per_dp_rank = 16384
    config.training.num_tokens_per_train_step = 65536
    return config


def muse_glimmer_debugmodel_fsdp2_per_group_cuda_graph() -> Trainer.Config:
    """Three fixed accumulation groups with FSDP degree 2.

    GPU unit tests cover changing group counts.
    """
    config = muse_glimmer_debugmodel(seq_len=2048)
    config.debug.deterministic = True
    config.debug.seed = 42
    config.training.cuda_graph_per_accumulation_group = True
    config.training.num_tokens_per_microbatch_per_dp_rank = 16384
    config.training.num_tokens_per_train_step = 98304
    config.parallelism.data_parallel_shard_degree = 2
    return config


def muse_glimmer_debugmodel_fsdp2_deferred_gradient_reduction() -> Trainer.Config:
    config = muse_glimmer_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.parallelism.fsdp_defer_gradient_reduction = True
    config.training.num_tokens_per_microbatch_per_dp_rank = 16384
    config.training.num_tokens_per_train_step = 65536
    return config


def llama3_debugmodel_validation_tp2_cp2_pp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    assert isinstance(config.dataloader, GrainDataLoader.Config)
    config.validator = Validator.Config(
        freq=5,
        steps=10,
        dataloader=replace(config.dataloader),
    )
    config.parallelism.tensor_parallel_degree = 2
    config.parallelism.context_parallel_degree = 2
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    config.training.disable_cuda_graphs = True
    return apply_transforms(
        config,
        [
            ContextParallelTransform(
                inner_attention_map={
                    FlexInnerAttention: KVAllGatherCPFlexInnerAttention
                }
            )
        ],
    )


def llama3_debugmodel_fused_swiglu_tp2() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.override.imports = ["torchtitan_recipes.overrides.fused_swiglu.fused_swiglu"]
    config.parallelism.tensor_parallel_degree = 2
    return config


def deepseek_v3_debugmodel_fused_swiglu_tp2_ep4() -> Trainer.Config:
    config = deepseek_v3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.override.imports = ["torchtitan_recipes.overrides.fused_swiglu.fused_swiglu"]
    config.parallelism.tensor_parallel_degree = 2
    config.parallelism.expert_parallel_degree = 4
    config.training.disable_cuda_graphs = True
    return config


def llama3_debugmodel_varlen_attn_fsdp4_sac() -> Trainer.Config:
    config = llama3_debugmodel_varlen_attn(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.data_parallel_shard_degree = 4
    config.activation_checkpoint = SelectiveAC.Config()
    return config


def llama3_debugmodel_lora_tp2_pp2() -> Trainer.Config:
    config = llama3_debugmodel_lora(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=False)
    config.parallelism.tensor_parallel_degree = 2
    config.parallelism.pipeline_parallel_degree = 2
    config.parallelism.num_pp_microbatches = 8
    config.training.num_tokens_per_microbatch_per_dp_rank = 2048
    return config


def llama3_debugmodel_sft() -> Trainer.Config:
    config = sft_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    return config


def llama3_debugmodel_sft_multiturn() -> Trainer.Config:
    config = llama3_debugmodel_sft()

    def messages(sample) -> list[Message]:
        return [
            {"role": "user", "content": sample["question"]},
            {"role": "assistant", "content": sample["answer"]},
            {"role": "user", "content": "Repeat your answer."},
            {"role": "assistant", "content": sample["answer"]},
        ]

    dataloader = cast(GrainDataLoader.Config, config.dataloader)
    packing = cast(FirstFitPackingConfig, dataloader.dataset)
    dataset = cast(SingleDatasetConfig, packing.dataset)
    processor = cast(ChatProcessor.Config, dataset.processor)
    processor.messages_fn = messages
    processor.renderer = from_renderers(Qwen3RendererConfig())
    return config


def llama3_debugmodel_seed_checkpoint() -> Trainer.Config:
    config = llama3_debugmodel(seq_len=2048)
    _set_spmd_typechecking(config, typechecking=True)
    config.checkpointer = CheckpointManager.Config()
    config.create_seed_checkpoint = True
    config.training.disable_cuda_graphs = True
    return config
