# Copyright 2026 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks
from ..modular_pipeline_utils import OutputParam
from .before_denoise import (
    AnimaImageInputStep,
    AnimaImg2ImgPrepareLatentsStep,
    AnimaImg2ImgSetTimestepsStep,
    AnimaPrepareLatentsStep,
    AnimaSetTimestepsStep,
    AnimaTextConditioningStep,
    AnimaTextInputStep,
)
from .decoders import AnimaProcessImagesOutputStep, AnimaVaeDecoderStep
from .denoise import AnimaDenoiseStep
from .encoders import AnimaImg2ImgVaeEncoderStep, AnimaTextEncoderStep


# auto_docstring
class AnimaCoreDenoiseStep(SequentialPipelineBlocks):
    """
    Denoise block that takes encoded Anima text inputs and runs the denoising process.

      Components:
          text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler
          (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`)

      Inputs:
          qwen_prompt_embeds (`Tensor`):
              Qwen prompt embeddings generated by the text encoder step.
          qwen_attention_mask (`Tensor`):
              Qwen prompt attention mask generated by the text encoder step.
          t5_input_ids (`Tensor`):
              T5 prompt token ids generated by the text encoder step.
          t5_attention_mask (`Tensor`):
              T5 prompt attention mask generated by the text encoder step.
          negative_qwen_prompt_embeds (`Tensor`, *optional*):
              Negative Qwen prompt embeddings generated by the text encoder step.
          negative_qwen_attention_mask (`Tensor`, *optional*):
              Negative Qwen prompt attention mask generated by the text encoder step.
          negative_t5_input_ids (`Tensor`, *optional*):
              Negative T5 prompt token ids generated by the text encoder step.
          negative_t5_attention_mask (`Tensor`, *optional*):
              Negative T5 prompt attention mask generated by the text encoder step.
          num_images_per_prompt (`int`, *optional*, defaults to 1):
              The number of images to generate per prompt.
          height (`int`, *optional*):
              The height in pixels of the generated image.
          width (`int`, *optional*):
              The width in pixels of the generated image.
          latents (`Tensor`, *optional*):
              Pre-generated noisy latents for image generation.
          generator (`Generator`, *optional*):
              Torch generator for deterministic generation.
          num_inference_steps (`int`, *optional*, defaults to 50):
              The number of denoising steps.
          sigmas (`list`, *optional*):
              Custom sigmas for the denoising process.
          **denoiser_input_fields (`None`, *optional*):
              The conditional model inputs for the Anima denoiser.

      Outputs:
          latents (`Tensor`):
              Denoised latents.
    """

    block_classes = [
        AnimaTextConditioningStep,
        AnimaTextInputStep,
        AnimaPrepareLatentsStep,
        AnimaSetTimestepsStep,
        AnimaDenoiseStep,
    ]
    block_names = ["text_conditioning", "input", "prepare_latents", "set_timesteps", "denoise"]

    @property
    def description(self) -> str:
        return "Denoise block that takes encoded Anima text inputs and runs the denoising process."

    @property
    def outputs(self):
        return [OutputParam.template("latents")]


# auto_docstring
class AnimaDecodeStep(SequentialPipelineBlocks):
    """
    Decode Anima latents into generated images.

      Components:
          vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`)

      Inputs:
          latents (`Tensor`):
              Denoised Anima latents.
          output_type (`str`, *optional*, defaults to pil):
              Output format: 'pil', 'np', 'pt'.

      Outputs:
          images (`list`):
              Generated images.
    """

    block_classes = [AnimaVaeDecoderStep, AnimaProcessImagesOutputStep]
    block_names = ["decode", "postprocess"]

    @property
    def description(self) -> str:
        return "Decode Anima latents into generated images."

    @property
    def outputs(self):
        return [OutputParam.template("images")]


# auto_docstring
class AnimaImg2ImgCoreDenoiseStep(SequentialPipelineBlocks):
    """
    Denoise block for Anima image-to-image generation. Uses image_latents already in state from
    AnimaImg2ImgVaeEncoderStep.

      Components:
          text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler
          (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`)

      Inputs:
          qwen_prompt_embeds (`Tensor`):
              Qwen prompt embeddings generated by the text encoder step.
          qwen_attention_mask (`Tensor`):
              Qwen prompt attention mask generated by the text encoder step.
          t5_input_ids (`Tensor`):
              T5 prompt token ids generated by the text encoder step.
          t5_attention_mask (`Tensor`):
              T5 prompt attention mask generated by the text encoder step.
          negative_qwen_prompt_embeds (`Tensor`, *optional*):
              Negative Qwen prompt embeddings generated by the text encoder step.
          negative_qwen_attention_mask (`Tensor`, *optional*):
              Negative Qwen prompt attention mask generated by the text encoder step.
          negative_t5_input_ids (`Tensor`, *optional*):
              Negative T5 prompt token ids generated by the text encoder step.
          negative_t5_attention_mask (`Tensor`, *optional*):
              Negative T5 prompt attention mask generated by the text encoder step.
          num_images_per_prompt (`int`, *optional*, defaults to 1):
              The number of images to generate per prompt.
          image_latents (`Tensor`):
              image latents used to guide the image generation. Can be generated from vae_encoder step.
          height (`int`, *optional*):
              The height in pixels of the generated image.
          width (`int`, *optional*):
              The width in pixels of the generated image.
          num_inference_steps (`int`, *optional*, defaults to 50):
              The number of denoising steps.
          sigmas (`list`, *optional*):
              Custom sigmas for the denoising process.
          strength (`float`, *optional*, defaults to 0.9):
              Strength for img2img/inpainting.
          generator (`Generator`, *optional*):
              Torch generator for deterministic generation.
          latents (`Tensor`, *optional*):
              Pre-generated noisy latents for image generation.
          **denoiser_input_fields (`None`, *optional*):
              The conditional model inputs for the Anima denoiser.

      Outputs:
          latents (`Tensor`):
              Denoised latents.
    """

    block_classes = [
        AnimaTextConditioningStep,
        AnimaTextInputStep,
        AnimaImageInputStep,
        AnimaImg2ImgSetTimestepsStep,
        AnimaImg2ImgPrepareLatentsStep,
        AnimaDenoiseStep,
    ]
    block_names = ["text_conditioning", "input", "image_input", "set_timesteps", "prepare_latents", "denoise"]

    @property
    def description(self) -> str:
        return (
            "Denoise block for Anima image-to-image generation. "
            "Uses image_latents already in state from AnimaImg2ImgVaeEncoderStep."
        )

    @property
    def outputs(self):
        return [OutputParam.template("latents")]


# auto_docstring
class AnimaAutoCoreDenoiseStep(AutoPipelineBlocks):
    """
    Denoise step that selects between text-to-image and image-to-image denoising based on whether image_latents is
    present in state. - `AnimaCoreDenoiseStep` (text2image) is used when no image_latents are present. -
    `AnimaImg2ImgCoreDenoiseStep` (img2img) is used when image_latents are present.

      Components:
          text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler
          (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`)

      Inputs:
          qwen_prompt_embeds (`Tensor`):
              Qwen prompt embeddings generated by the text encoder step.
          qwen_attention_mask (`Tensor`):
              Qwen prompt attention mask generated by the text encoder step.
          t5_input_ids (`Tensor`):
              T5 prompt token ids generated by the text encoder step.
          t5_attention_mask (`Tensor`):
              T5 prompt attention mask generated by the text encoder step.
          negative_qwen_prompt_embeds (`Tensor`, *optional*):
              Negative Qwen prompt embeddings generated by the text encoder step.
          negative_qwen_attention_mask (`Tensor`, *optional*):
              Negative Qwen prompt attention mask generated by the text encoder step.
          negative_t5_input_ids (`Tensor`, *optional*):
              Negative T5 prompt token ids generated by the text encoder step.
          negative_t5_attention_mask (`Tensor`, *optional*):
              Negative T5 prompt attention mask generated by the text encoder step.
          num_images_per_prompt (`int`, *optional*, defaults to 1):
              The number of images to generate per prompt.
          image_latents (`Tensor`, *optional*):
              image latents used to guide the image generation. Can be generated from vae_encoder step.
          height (`int`, *optional*):
              The height in pixels of the generated image.
          width (`int`, *optional*):
              The width in pixels of the generated image.
          num_inference_steps (`int`):
              The number of denoising steps.
          sigmas (`list`, *optional*):
              Custom sigmas for the denoising process.
          strength (`float`, *optional*, defaults to 0.9):
              Strength for img2img/inpainting.
          generator (`Generator`, *optional*):
              Torch generator for deterministic generation.
          latents (`Tensor`):
              Pre-generated noisy latents for image generation.
          **denoiser_input_fields (`None`, *optional*):
              The conditional model inputs for the Anima denoiser.

      Outputs:
          latents (`Tensor`):
              Denoised latents.
    """

    block_classes = [AnimaImg2ImgCoreDenoiseStep, AnimaCoreDenoiseStep]
    block_names = ["img2img", "text2image"]
    block_trigger_inputs = ["image_latents", None]

    @property
    def description(self) -> str:
        return (
            "Denoise step that selects between text-to-image and image-to-image denoising based on whether "
            "image_latents is present in state."
            " - `AnimaCoreDenoiseStep` (text2image) is used when no image_latents are present."
            " - `AnimaImg2ImgCoreDenoiseStep` (img2img) is used when image_latents are present."
        )


# auto_docstring
class AnimaAutoVaeImageEncoderStep(AutoPipelineBlocks):
    """
    VAE Image Encoder step that encodes the input image to produce image_latents. Skipped when no image is provided
    (text-to-image workflow).

      Components:
          vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`)

      Inputs:
          image (`Image | list`, *optional*):
              Reference image(s) for denoising. Can be a single image or list of images.
          height (`int`, *optional*):
              The height in pixels of the generated image.
          width (`int`, *optional*):
              The width in pixels of the generated image.
          generator (`Generator`, *optional*):
              Torch generator for deterministic generation.

      Outputs:
          image_latents (`Tensor`):
              Encoded image latents.
          height (`int`):
              Image height used for generation.
          width (`int`):
              Image width used for generation.
    """

    block_classes = [AnimaImg2ImgVaeEncoderStep]
    block_names = ["vae_encoder"]
    block_trigger_inputs = ["image"]

    @property
    def description(self) -> str:
        return (
            "VAE Image Encoder step that encodes the input image to produce image_latents. "
            "Skipped when no image is provided (text-to-image workflow)."
        )


# auto_docstring
class AnimaAutoBlocks(SequentialPipelineBlocks):
    """
    Auto Modular pipeline for text-to-image and image-to-image generation using Anima.

      Supported workflows:
        - `text2image`: requires `prompt`
        - `img2img`: requires `image`, `prompt`

      Components:
          text_encoder (`Qwen3Model`) tokenizer (`Qwen2Tokenizer`) t5_tokenizer (`T5Tokenizer`) guider
          (`ClassifierFreeGuidance`) vae (`AutoencoderKLQwenImage`) image_processor (`VaeImageProcessor`)
          text_conditioner (`AnimaTextConditioner`) transformer (`CosmosTransformer3DModel`) scheduler
          (`FlowMatchEulerDiscreteScheduler`)

      Inputs:
          prompt (`str`):
              The prompt or prompts to guide image generation.
          negative_prompt (`str`, *optional*):
              The prompt or prompts not to guide the image generation.
          max_sequence_length (`int`, *optional*, defaults to 512):
              Maximum sequence length for prompt encoding.
          image (`Image | list`, *optional*):
              Reference image(s) for denoising. Can be a single image or list of images.
          height (`int`, *optional*):
              The height in pixels of the generated image.
          width (`int`, *optional*):
              The width in pixels of the generated image.
          generator (`Generator`, *optional*):
              Torch generator for deterministic generation.
          num_images_per_prompt (`int`, *optional*, defaults to 1):
              The number of images to generate per prompt.
          image_latents (`Tensor`, *optional*):
              image latents used to guide the image generation. Can be generated from vae_encoder step.
          num_inference_steps (`int`):
              The number of denoising steps.
          sigmas (`list`, *optional*):
              Custom sigmas for the denoising process.
          strength (`float`, *optional*, defaults to 0.9):
              Strength for img2img/inpainting.
          latents (`Tensor`):
              Pre-generated noisy latents for image generation.
          **denoiser_input_fields (`None`, *optional*):
              The conditional model inputs for the Anima denoiser.
          output_type (`str`, *optional*, defaults to pil):
              Output format: 'pil', 'np', 'pt'.

      Outputs:
          images (`list`):
              Generated images.
    """

    block_classes = [
        AnimaTextEncoderStep,
        AnimaAutoVaeImageEncoderStep,
        AnimaAutoCoreDenoiseStep,
        AnimaDecodeStep,
    ]
    block_names = ["text_encoder", "vae_encoder", "denoise", "decode"]
    _workflow_map = {
        "text2image": {"prompt": True},
        "img2img": {"image": True, "prompt": True},
    }

    @property
    def description(self) -> str:
        return "Auto Modular pipeline for text-to-image and image-to-image generation using Anima."

    @property
    def outputs(self):
        return [OutputParam.template("images")]
