# 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 collections.abc import Iterator
from dataclasses import dataclass, field, replace
from typing import Any

import torch
from torch.distributed.pipelining.schedules import _PipelineSchedule
from torchtitan.components.data import ConcatThenSplitPackingConfig, GrainDataLoader
from torchtitan.components.data.loader import BaseDataLoader
from torchtitan.components.data.types import TrainingMicrobatch
from torchtitan.components.loss import LossFunction
from torchtitan.components.tokenizer import BaseTokenizer
from torchtitan.config import Configurable
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed import ParallelismContext, utils as dist_utils
from torchtitan.hf_datasets.text_datasets import DATASETS
from torchtitan.observability import structured_logger as sl
from torchtitan.observability.metrics import MetricsProcessor
from torchtitan.protocols.model import BaseModel
from torchtitan.tools import utils


class BaseValidator(Configurable):
    @dataclass(kw_only=True, slots=True)
    class Config(Configurable.Config):
        freq: int = 10
        """Frequency of validation"""

        def __post_init__(self) -> None:
            if self.freq <= 0:
                raise ValueError(
                    f"validation frequency must be positive, got {self.freq}"
                )

    def __init__(
        self,
        config: Config,
        **kwargs,
    ):
        self.config = config

    def validate(self, model_parts: list[BaseModel], step: int) -> None:
        raise NotImplementedError("validate method not implemented")

    def should_validate(self, step: int) -> bool:
        return step == 1 or step % self.config.freq == 0


class Validator(BaseValidator):
    """
    Simple validator focused on correctness and integration.

    Args:
        config: Validator.Config configuration
        parallelism: ParallelismConfig configuration
        dp_world_size: Data parallel world size
        dp_rank: Data parallel rank
        tokenizer: Tokenizer
        parallelism_context: Parallel dimensions
        loss_fn: Loss function to use for validation
        metrics_processor: Metrics processor
        pp_schedule: Pipeline schedule (optional)
        pp_has_first_stage: Whether this rank has the first PP stage (optional)
        pp_has_last_stage: Whether this rank has the last PP stage (optional)
    """

    @dataclass(kw_only=True, slots=True)
    class Config(BaseValidator.Config):
        steps: int = -1
        """
        Number of validation steps. -1 consumes the finite dataset once
        (dataloader repeat=False). Ranks then stop independently, so this
        requires data-parallel degree 1; otherwise validation collectives hang.
        Use a positive count when DP > 1 so every rank runs the same number of
        steps with repeat=True.
        """

        dataloader: BaseDataLoader.Config = field(
            default_factory=lambda: GrainDataLoader.Config(
                dataset=ConcatThenSplitPackingConfig(
                    dataset=DATASETS["c4_validation"],
                ),
                repeat=False,
            )
        )
        """DataLoader configuration for validation"""

        def __post_init__(self):
            BaseValidator.Config.__post_init__(self)
            if not (self.steps > 0 or self.steps == -1):
                raise ValueError(
                    f"validation steps must be positive or -1, got {self.steps}"
                )

    # TODO: improve the constructor signature
    def __init__(
        self,
        config: Config,
        *,
        parallelism: ParallelismConfig,
        dp_world_size: int,
        dp_rank: int,
        tokenizer: BaseTokenizer,
        parallelism_context: ParallelismContext,
        loss_fn: LossFunction,
        metrics_processor: MetricsProcessor,
        seq_len: int,
        num_tokens_per_microbatch: int,
        pp_schedule: _PipelineSchedule | None = None,
        pp_has_first_stage: bool | None = None,
        pp_has_last_stage: bool | None = None,
        **kwargs,
    ):
        super().__init__(config=config)
        self.parallelism = parallelism
        self.tokenizer = tokenizer
        self.parallelism_context = parallelism_context
        self.loss_fn = loss_fn
        # A bounded validation run repeats data; steps=-1 consumes one finite pass.
        self.dl_config = replace(config.dataloader, repeat=config.steps != -1)
        self.dp_world_size = dp_world_size
        self.dp_rank = dp_rank
        if config.steps == -1 and self.dp_world_size > 1:
            raise ValueError(
                "validation.steps=-1 runs one finite pass (dataloader "
                "repeat=False). With data-parallel degree > 1, ranks can exhaust "
                "at different steps and hang on validation collectives. Got "
                f"dp_world_size={self.dp_world_size}. Set validation.steps to a "
                "positive count so every rank runs the same number of steps, "
                "or run with data-parallel degree 1."
            )
        self.seq_len = seq_len
        self.num_tokens_per_microbatch = num_tokens_per_microbatch
        self.metrics_processor = metrics_processor
        self.pp_schedule = pp_schedule
        self.pp_has_first_stage = pp_has_first_stage
        self.pp_has_last_stage = pp_has_last_stage

    @sl.log_trace_span("eval")
    @torch.no_grad()
    def validate(
        self,
        model_parts: list[BaseModel],
        step: int,
    ) -> None:
        sl.add_step_tag("eval")
        self.metrics_processor.reset()
        # Set model to eval mode
        for model in model_parts:
            model.eval()

        parallelism_context = self.parallelism_context

        accumulated_loss: torch.Tensor | None = None
        device_type = utils.device_type
        total_global_valid_tokens = torch.zeros(
            (), dtype=torch.int64, device=device_type
        )
        num_steps = 0
        num_pp_microbatches = (
            self.parallelism.num_pp_microbatches
            if parallelism_context.pp_enabled
            else 1
        )

        validation_dataloader = self.dl_config.build(
            dp_world_size=self.dp_world_size,
            dp_rank=self.dp_rank,
            tokenizer=self.tokenizer,
            max_context_length=self.seq_len,
            num_tokens_per_microbatch=self.num_tokens_per_microbatch,
        )

        validation_iterator = iter(iterate_and_close_dataloader(validation_dataloader))
        while True:
            # pyrefly: ignore [missing-attribute, unsupported-operation]
            if self.config.steps != -1 and num_steps >= self.config.steps:
                break

            try:
                microbatch_group = []
                local_loss_token_count = torch.zeros((), dtype=torch.int64)
                for _ in range(num_pp_microbatches):
                    microbatch = next(validation_iterator)
                    local_loss_token_count.add_(
                        microbatch.loss_token_counts.reshape(-1)[0]
                    )
                    self.metrics_processor.ntokens_since_last_log += (
                        microbatch.labels.numel()
                    )
                    input_dict = microbatch.to_input_dict(device_type)
                    microbatch_group.append(input_dict)
            except StopIteration:
                break

            # All-reduce token count across DP ranks while keeping it on device.
            local_loss_token_count = local_loss_token_count.to(device_type)
            if parallelism_context.dp_enabled:
                dp_mesh = parallelism_context.get_mesh("dp")
                global_valid_tokens = dist_utils.dist_sum_tensor(
                    local_loss_token_count, dp_mesh, None
                )
            else:
                global_valid_tokens = local_loss_token_count

            if parallelism_context.pp_enabled:
                assert self.pp_schedule is not None
                assert self.pp_has_first_stage is not None
                assert self.pp_has_last_stage is not None

                arg_mbs: list[tuple[torch.Tensor, ...]] = []
                kwarg_mbs: list[dict[str, Any]] = []
                target_mbs: list[torch.Tensor] | None = (
                    [] if self.pp_has_last_stage else None
                )

                for input_dict in microbatch_group:
                    with self.parallelism_context.activate_spmd():
                        inputs, labels, extra_kwargs = model_parts[0].preprocess_inputs(
                            input_dict,
                            parallelism_context=self.parallelism_context,
                            parallelism=self.parallelism,
                        )
                    if self.pp_has_first_stage:
                        arg_mbs.append((inputs,))  # pyrefly: ignore[bad-argument-type]
                    kwarg_mbs.append(extra_kwargs)
                    if target_mbs is not None:
                        target_mbs.append(labels)  # pyrefly: ignore[bad-argument-type]

                with self.parallelism_context.activate_spmd():
                    losses = [] if self.pp_has_last_stage else None
                    self.pp_schedule.eval(
                        arg_mbs=arg_mbs if self.pp_has_first_stage else None,
                        kwarg_mbs=kwarg_mbs,
                        target_mbs=target_mbs,
                        losses=losses,
                    )

                # accumulate losses across pipeline microbatches
                # TODO: PP+FSDP unexpectedly puts the loss back to the CPU
                if self.pp_has_last_stage:
                    assert losses is not None
                    # using sum because loss_fn already uses reduction='sum'
                    loss_sum = torch.sum(torch.stack(losses)).to(device_type)
                else:
                    loss_sum = torch.tensor([-1.0], device=device_type)
            else:
                assert len(microbatch_group) == 1
                input_dict = microbatch_group[0]
                with self.parallelism_context.activate_spmd():
                    inputs, labels, extra_kwargs = model_parts[0].preprocess_inputs(
                        input_dict,
                        parallelism_context=self.parallelism_context,
                        parallelism=self.parallelism,
                    )
                    assert len(model_parts) == 1
                    predictions = model_parts[0](inputs, **extra_kwargs)
                    loss_sum, _ = self.loss_fn(predictions, labels)

            loss_sum = loss_sum.detach()
            if accumulated_loss is None:
                accumulated_loss = loss_sum.clone()
            else:
                accumulated_loss.add_(loss_sum)
            total_global_valid_tokens.add_(global_valid_tokens)
            num_steps += 1

        if accumulated_loss is None:
            raise ValueError(
                "Validation ran zero batches on this rank. This happens when the "
                "validation dataset supplies fewer than num_tokens_per_microbatch "
                "tokens on this rank, because concat-then-split packing drops "
                "partially filled batches. Decrease "
                "training.num_tokens_per_microbatch_per_dp_rank or use a larger "
                "validation dataset."
            )
        num_global_valid_tokens = int(total_global_valid_tokens.item())
        if num_global_valid_tokens == 0:
            raise ValueError(
                "Validation ran on zero valid tokens; cannot compute an average "
                "validation loss. Ensure the validation batches contain unmasked "
                "labels."
            )
        if parallelism_context.dp_cp_enabled:
            global_loss_sum = dist_utils.dist_sum(
                accumulated_loss, parallelism_context.get_optional_mesh("loss")
            )
        else:
            global_loss_sum = float(accumulated_loss.item())
        global_avg_loss = global_loss_sum / num_global_valid_tokens

        self.metrics_processor.log_validation(loss=global_avg_loss, step=step)

        # Set model back to train mode
        for model in model_parts:
            model.train()


def iterate_and_close_dataloader(
    dataloader: BaseDataLoader,
) -> Iterator[TrainingMicrobatch]:
    """Close a temporary dataloader when its consumer stops iterating."""
    try:
        yield from dataloader
    finally:
        dataloader.close()
