# 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

import torch
from torch.distributed.elastic.multiprocessing.errors import record
from torchtitan.config import ConfigLoader
from torchtitan.models.flux.inference.sampling import generate_image, save_image
from torchtitan.models.flux.trainer import FluxTrainer
from torchtitan.observability.logging import init_logger


logger = logging.getLogger(__name__)


@torch.no_grad()
@record
def inference(config: FluxTrainer.Config):
    # Reuse trainer to perform forward passes
    trainer = FluxTrainer(config)

    # Distributed processing setup: Each GPU/process handles a subset of prompts
    world_size = int(os.environ["WORLD_SIZE"])
    global_rank = int(os.environ["RANK"])
    original_prompts = open(config.inference.prompts_path).readlines()
    total_prompts = len(original_prompts)

    if total_prompts < world_size:
        raise ValueError(
            f"Number of prompts ({total_prompts}) must be >= number of ranks ({world_size}). "
            f"FSDP all-gather will hang if some ranks have no prompts to process."
        )

    # Distribute prompts across processes using round-robin assignment
    prompts = original_prompts[global_rank::world_size]

    if config.checkpointer is None:
        raise ValueError("Flux inference requires a checkpointer configuration.")
    trainer.engine.load_checkpoint()

    # Build tokenizers from the config
    tokenizer = config.tokenizer.build()

    if global_rank == 0:
        logger.info("Starting inference...")

    if prompts:
        # Generate images for this process's assigned prompts
        bs = config.inference.local_batch_size
        img_size = config.inference.img_size

        output_dir = os.path.join(
            config.dump_folder,
            config.inference.save_img_folder,
        )
        # Create mapping from local indices to global prompt indices
        global_ids = list(range(global_rank, total_prompts, world_size))

        for i in range(0, len(prompts), bs):
            with trainer.engine.parallelism_context.activate_spmd(
                typechecking=trainer.engine.config.debug.spmd_typechecking,
            ):
                images = generate_image(
                    device=trainer.engine.device,
                    dtype=trainer._dtype,
                    img_height=16 * (img_size // 16),
                    img_width=16 * (img_size // 16),
                    enable_classifier_free_guidance=config.inference.sampling.enable_classifier_free_guidance,
                    denoising_steps=config.inference.sampling.denoising_steps,
                    classifier_free_guidance_scale=config.inference.sampling.classifier_free_guidance_scale,
                    # pyrefly: ignore [bad-argument-type]
                    model=trainer.engine.model_parts[0],
                    prompt=prompts[i : i + bs],
                    autoencoder=trainer.autoencoder,
                    tokenizer=tokenizer,
                    t5_encoder=trainer.t5_encoder,
                    clip_encoder=trainer.clip_encoder,
                )
            for j in range(images.shape[0]):
                # Extract single image while preserving batch dimension [1, C, H, W]
                img = images[j : j + 1]
                global_id = global_ids[i + j]

                save_image(
                    name=f"image_prompt{global_id}_rank{str(torch.distributed.get_rank())}.png",
                    output_dir=output_dir,
                    x=img,
                    add_sampling_metadata=True,
                    prompt=prompts[i + j],
                )

    torch.distributed.destroy_process_group()


if __name__ == "__main__":
    init_logger()
    config = ConfigLoader().load()
    inference(config)  # pyrefly: ignore [bad-argument-type]
