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

import logging
import os
from dataclasses import dataclass, field, replace

import torch
from torch.distributed.pipelining.schedules import _PipelineSchedule

from torchtitan.components.data import GrainDataLoader
from torchtitan.components.loss import LossFunction
from torchtitan.components.tokenizer import BaseTokenizer
from torchtitan.components.validate import iterate_and_close_dataloader, Validator
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed import (
    context_parallel,
    ParallelismContext,
    utils as dist_utils,
)
from torchtitan.observability.metrics import MetricsProcessor
from torchtitan.protocols.model import BaseModel

from .configs import SamplingConfig
from .flux_datasets import FluxValidationDatasetConfig
from .inference.sampling import generate_image, save_image
from .model.autoencoder import AutoEncoder
from .model.hf_embedder import FluxEmbedder
from .sharding import flux_input_sharding
from .tokenizer import FluxTokenizerContainer
from .utils import create_position_encoding_for_latents, pack_latents, preprocess_data


logger = logging.getLogger(__name__)


class FluxValidator(Validator):
    """
    Flux model validator focused on correctness and integration.

    Args:
        config: FluxValidator.Config configuration
        parallelism: Parallelism 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
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Validator.Config):
        dataloader: GrainDataLoader.Config  # pyrefly: ignore [bad-override]
        """DataLoader configuration for Flux validation"""

        all_timesteps: bool = False
        """Generate all 8 timesteps for each sample instead of round-robin"""

        save_img_count: int = -1
        """Number of images to save during validation (-1 for unlimited)"""

        save_img_folder: str = "validation_images"
        """Folder to save validation images"""

        sampling: SamplingConfig = field(default_factory=SamplingConfig)
        """Sampling configuration for validation image generation"""

    def __init__(
        self,
        config: Config,
        *,
        parallelism: ParallelismConfig,
        dp_world_size: int,
        dp_rank: int,
        tokenizer: BaseTokenizer,
        parallelism_context: ParallelismContext,
        loss_fn: LossFunction,
        seq_len: int,
        num_tokens_per_microbatch: int,
        metrics_processor: MetricsProcessor | None = None,
        pp_schedule: _PipelineSchedule | None = None,
        pp_has_first_stage: bool | None = None,
        pp_has_last_stage: bool | None = None,
        **kwargs,
    ):
        self.config = config
        self.parallelism = parallelism
        self.tokenizer = tokenizer
        self.parallelism_context = parallelism_context
        self.loss_fn = loss_fn
        self.all_timesteps = config.all_timesteps

        assert isinstance(tokenizer, FluxTokenizerContainer)

        dataset = config.dataloader.dataset
        if isinstance(dataset, FluxValidationDatasetConfig):
            dataset = dataset.dataset
        # A bounded validation run repeats data; steps=-1 consumes one finite pass.
        self.dl_config = replace(
            config.dataloader,
            dataset=(
                dataset
                if config.all_timesteps
                else FluxValidationDatasetConfig(dataset=dataset)
            ),
            repeat=config.steps != -1,
        )
        self.dp_world_size = dp_world_size
        self.dp_rank = dp_rank
        self.seq_len = seq_len
        self.num_tokens_per_microbatch = num_tokens_per_microbatch
        # pyrefly: ignore [bad-assignment]
        self.metrics_processor = metrics_processor

        if config.steps == -1:
            logger.warning(
                "Setting validation steps to -1 might cause hangs because of "
                "unequal sample counts across ranks when dataset is exhausted."
            )

    def flux_init(
        self,
        device: torch.device,
        _dtype: torch.dtype,
        autoencoder: AutoEncoder,
        t5_encoder: FluxEmbedder,
        clip_encoder: FluxEmbedder,
        dump_folder: str,
    ):
        # pyrefly: ignore [read-only]
        self.device = device
        self._dtype = _dtype
        self.autoencoder = autoencoder
        self.t5_encoder = t5_encoder
        self.clip_encoder = clip_encoder
        self.dump_folder = dump_folder

    @torch.no_grad()
    def validate(
        self,
        model_parts: list[BaseModel],
        step: int,
    ) -> None:
        # Set model to eval mode
        # TODO: currently does not support pipeline parallelism
        model = model_parts[0]
        model.eval()
        self.metrics_processor.reset()

        assert isinstance(self.config, FluxValidator.Config)
        max_saved_images = self.config.save_img_count
        image_idx = 0

        parallelism_context = self.parallelism_context

        accumulated_loss: torch.Tensor | None = None
        device_type = dist_utils.device_type
        total_local_elements = torch.zeros((), dtype=torch.int64, device=device_type)
        num_steps = 0

        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,
        )

        for microbatch in iterate_and_close_dataloader(validation_dataloader):
            if self.config.steps != -1 and num_steps >= self.config.steps:
                break

            input_dict = microbatch.to_input_dict(self.device)
            labels = input_dict.pop("labels")
            prompt = input_dict.pop("prompt")
            if not isinstance(prompt, list):
                prompt = [prompt]
            img_height, img_width = labels.shape[-2:]
            for p in prompt:
                assert isinstance(p, str), f"prompt must be a string, got {type(p)}"
                if max_saved_images != -1 and image_idx >= max_saved_images:
                    break
                with self.parallelism_context.activate_spmd():
                    image = generate_image(
                        device=self.device,
                        dtype=self._dtype,
                        img_height=img_height,
                        img_width=img_width,
                        enable_classifier_free_guidance=self.config.sampling.enable_classifier_free_guidance,
                        denoising_steps=self.config.sampling.denoising_steps,
                        classifier_free_guidance_scale=self.config.sampling.classifier_free_guidance_scale,
                        # pyrefly: ignore [bad-argument-type]
                        model=model,
                        prompt=p,
                        autoencoder=self.autoencoder,
                        # pyrefly: ignore [bad-argument-type]
                        tokenizer=self.tokenizer,
                        t5_encoder=self.t5_encoder,
                        clip_encoder=self.clip_encoder,
                    )

                save_image(
                    name=(
                        f"image_rank{torch.distributed.get_rank()}_step{step}_"
                        f"{image_idx:06d}.png"
                    ),
                    output_dir=os.path.join(
                        self.dump_folder,
                        self.config.save_img_folder,
                    ),
                    x=image,
                    add_sampling_metadata=True,
                    prompt=p,
                )
                image_idx += 1

            # generate t5 and clip embeddings
            input_dict["image"] = labels
            input_dict = preprocess_data(
                device=self.device,
                dtype=self._dtype,
                autoencoder=self.autoencoder,
                clip_encoder=self.clip_encoder,
                t5_encoder=self.t5_encoder,
                batch=input_dict,
            )
            labels = input_dict["img_encodings"].to(device_type)
            clip_encodings = input_dict["clip_encodings"]
            t5_encodings = input_dict["t5_encodings"]

            bsz = labels.shape[0]

            # If using all_timesteps we generate all 8 timesteps and expand our batch inputs here
            if self.all_timesteps:
                stratified_timesteps = torch.tensor(
                    [1 / 8 * (i + 0.5) for i in range(8)],
                    dtype=torch.float32,
                    device=self.device,
                ).repeat(bsz)
                clip_encodings = clip_encodings.repeat_interleave(8, dim=0)
                t5_encodings = t5_encodings.repeat_interleave(8, dim=0)
                labels = labels.repeat_interleave(8, dim=0)
            else:
                stratified_timesteps = input_dict.pop("timestep")

            # Count full latent elements before CP shards the sequence.
            total_local_elements += labels.numel()

            # Note the tps may be inaccurate due to the generating image step not being counted
            self.metrics_processor.ntokens_since_last_log += labels.numel()

            # Apply timesteps here and update our bsz to efficiently compute all timesteps and samples in a single forward pass
            with torch.no_grad(), torch.device(self.device):
                noise = torch.randn_like(labels)
                timesteps = stratified_timesteps.to(labels)
                sigmas = timesteps.view(-1, 1, 1, 1)
                latents = (1 - sigmas) * labels + sigmas * noise

            bsz, _, latent_height, latent_width = latents.shape

            POSITION_DIM = 3  # constant for Flux flow model
            with torch.no_grad(), torch.device(self.device):
                # Create positional encodings
                latent_pos_enc = create_position_encoding_for_latents(
                    bsz, latent_height, latent_width, POSITION_DIM
                )
                text_pos_enc = torch.zeros(bsz, t5_encodings.shape[1], POSITION_DIM)

                # Patchify: Convert latent into a sequence of patches
                latents = pack_latents(latents)
                target = pack_latents(noise - labels)

            # Apply CP sharding if enabled
            if parallelism_context.cp_enabled:
                cp_inputs = {
                    "img": latents,
                    "img_ids": latent_pos_enc,
                    "txt": t5_encodings,
                    "txt_ids": text_pos_enc,
                    "target": target,
                }
                input_sharding = flux_input_sharding()
                with self.parallelism_context.activate_spmd():
                    load_balancer_config = (
                        self.parallelism.context_parallel_load_balancer
                    )
                    load_balancer = (
                        load_balancer_config.build(
                            seq_len=context_parallel.get_cp_input_seq_len(
                                cp_inputs, input_shardings=input_sharding
                            ),
                            attention_metadata=None,
                        )
                        if load_balancer_config is not None
                        else None
                    )
                    permutation = (
                        load_balancer.generate_permutation()
                        if load_balancer is not None
                        else None
                    )
                    cp_inputs = context_parallel.shard_tensors(
                        cp_inputs,
                        input_shardings=input_sharding,
                        permutation=permutation,
                    )
                latents = cp_inputs["img"]
                latent_pos_enc = cp_inputs["img_ids"]
                t5_encodings = cp_inputs["txt"]
                text_pos_enc = cp_inputs["txt_ids"]
                target = cp_inputs["target"]

            with self.parallelism_context.activate_spmd():
                latent_noise_pred = model(
                    img=latents,
                    img_ids=latent_pos_enc,
                    txt=t5_encodings,
                    txt_ids=text_pos_enc,
                    y=clip_encodings,
                    timesteps=timesteps,
                )

                loss, _ = self.loss_fn(latent_noise_pred, target)

            del noise, target, latent_noise_pred, latents

            loss = loss.detach()
            if accumulated_loss is None:
                accumulated_loss = loss.clone()
            else:
                accumulated_loss.add_(loss)

            num_steps += 1

        assert accumulated_loss is not None

        # CP ranks shard the same full latent tensor, so only DP contributes
        # additional elements to the denominator.
        if parallelism_context.dp_enabled:
            total_global_elements = dist_utils.dist_sum_tensor(
                total_local_elements, parallelism_context.get_mesh("dp")
            )
        else:
            total_global_elements = total_local_elements

        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 / int(total_global_elements.item())

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

        # Set model back to train mode
        model.train()
