# Copyright 2026 The HuggingFace Team.
#
# 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.

import tempfile
import unittest

import numpy as np
import PIL.Image
import torch

from diffusers import (
    AnimaAutoBlocks,
    AnimaModularPipeline,
    AnimaTextConditioner,
    CosmosTransformer3DModel,
)

from ...testing_utils import enable_full_determinism, is_lora, require_peft_backend
from ..testing_utils import (
    BaseModularPipelineOutputMixin,
    BaseModularPipelineTesterConfig,
    ModularGuiderTesterMixin,
    ModularLoadingTesterMixin,
    ModularMemoryTesterMixin,
    ModularPipelineTesterMixin,
    ModularWorkflowTesterMixin,
)


enable_full_determinism()


ANIMA_TEXT2IMAGE_WORKFLOWS = {
    "text2image": [
        ("text_encoder", "AnimaTextEncoderStep"),
        ("denoise.text_conditioning", "AnimaTextConditioningStep"),
        ("denoise.input", "AnimaTextInputStep"),
        ("denoise.prepare_latents", "AnimaPrepareLatentsStep"),
        ("denoise.set_timesteps", "AnimaSetTimestepsStep"),
        ("denoise.denoise", "AnimaDenoiseStep"),
        ("decode.decode", "AnimaVaeDecoderStep"),
        ("decode.postprocess", "AnimaProcessImagesOutputStep"),
    ],
}

ANIMA_IMG2IMG_WORKFLOWS = {
    "img2img": [
        ("text_encoder", "AnimaTextEncoderStep"),
        ("vae_encoder", "AnimaImg2ImgVaeEncoderStep"),
        ("denoise.text_conditioning", "AnimaTextConditioningStep"),
        ("denoise.input", "AnimaTextInputStep"),
        ("denoise.image_input", "AnimaImageInputStep"),
        ("denoise.set_timesteps", "AnimaImg2ImgSetTimestepsStep"),
        ("denoise.prepare_latents", "AnimaImg2ImgPrepareLatentsStep"),
        ("denoise.denoise", "AnimaDenoiseStep"),
        ("decode.decode", "AnimaVaeDecoderStep"),
        ("decode.postprocess", "AnimaProcessImagesOutputStep"),
    ],
}


def get_dummy_image(height=32, width=32):
    image_array = np.random.randint(0, 256, (height, width, 3), dtype=np.uint8)
    return PIL.Image.fromarray(image_array)


class AnimaTextConditionerFastTests(unittest.TestCase):
    def test_conditioner_output_shape_and_padding(self):
        conditioner = AnimaTextConditioner(
            source_dim=16,
            target_dim=16,
            model_dim=16,
            num_layers=2,
            num_attention_heads=4,
            target_vocab_size=128,
            min_sequence_length=8,
        )
        source_hidden_states = torch.randn(2, 5, 16)
        target_input_ids = torch.randint(0, 128, (2, 4))
        source_attention_mask = torch.ones(2, 5)
        target_attention_mask = torch.ones(2, 4)
        target_attention_mask[1, -1] = 0

        output = conditioner(
            source_hidden_states=source_hidden_states,
            target_input_ids=target_input_ids,
            source_attention_mask=source_attention_mask,
            target_attention_mask=target_attention_mask,
        )

        self.assertEqual(output.shape, (2, 8, 16))
        self.assertTrue(torch.allclose(output[1, 3], torch.zeros_like(output[1, 3]), atol=1e-5))
        self.assertTrue(torch.allclose(output[:, 4:], torch.zeros_like(output[:, 4:]), atol=1e-5))


class AnimaModularPipelineTesterConfig(BaseModularPipelineTesterConfig):
    pipeline_class = AnimaModularPipeline
    pipeline_blocks_class = AnimaAutoBlocks
    pretrained_model_name_or_path = "hf-internal-testing/tiny-anima-modular-pipe"
    params = frozenset(["prompt", "height", "width", "negative_prompt"])
    batch_params = frozenset(["prompt", "negative_prompt"])
    expected_workflow_blocks = ANIMA_TEXT2IMAGE_WORKFLOWS

    def get_dummy_inputs(self, seed=0):
        generator = torch.Generator(device="cpu").manual_seed(seed)
        return {
            "prompt": "dance monkey",
            "negative_prompt": "bad quality",
            "generator": generator,
            "num_inference_steps": 2,
            "height": 32,
            "width": 32,
            "max_sequence_length": 16,
            "output_type": "pt",
        }


class TestAnimaModularPipelineFast(AnimaModularPipelineTesterConfig, ModularPipelineTesterMixin):
    def test_inference_empty_negative_prompt(self):
        pipe = self.get_pipeline()

        inputs = self.get_dummy_inputs()
        inputs["negative_prompt"] = ""
        output = pipe(**inputs).images

        assert output.shape == (1, 3, 32, 32)
        assert not torch.isnan(output).any()

    def test_inference_batch_single_identical(self):
        super().test_inference_batch_single_identical(expected_max_diff=5e-4)


class TestAnimaModularPipelineLoading(AnimaModularPipelineTesterConfig, ModularLoadingTesterMixin):
    def test_save_load_components(self):
        pipe = self.get_pipeline()

        with tempfile.TemporaryDirectory() as tmpdir:
            pipe.save_pretrained(tmpdir, safe_serialization=True)
            pipe = self.pipeline_class.from_pretrained(tmpdir)
            pipe.load_components()

        assert isinstance(pipe.text_conditioner, AnimaTextConditioner)
        assert isinstance(pipe.transformer, CosmosTransformer3DModel)


@is_lora
class TestAnimaModularPipelineLoRA(AnimaModularPipelineTesterConfig, BaseModularPipelineOutputMixin):
    def test_lora_state_dict_conversion(self):
        state_dict = {
            "diffusion_model.blocks.0.self_attn.q_proj.lora_A.weight": torch.randn(2, 32),
            "diffusion_model.blocks.0.self_attn.q_proj.lora_B.weight": torch.randn(32, 2),
            "diffusion_model.blocks.0.adaln_modulation_cross_attn.1.lora_A.weight": torch.randn(2, 32),
            "diffusion_model.blocks.0.adaln_modulation_cross_attn.1.lora_B.weight": torch.randn(4, 2),
            "diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_A.weight": torch.randn(2, 16),
            "diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_B.weight": torch.randn(16, 2),
        }

        converted_state_dict = self.pipeline_class.lora_state_dict(state_dict)

        assert "transformer.transformer_blocks.0.attn1.to_q.lora_A.weight" in converted_state_dict
        assert "transformer.transformer_blocks.0.norm2.linear_1.lora_B.weight" in converted_state_dict
        assert "text_conditioner.blocks.0.self_attn.q_proj.lora_A.weight" in converted_state_dict

    @require_peft_backend
    def test_load_lora_weights(self):
        pipe = self.get_pipeline()
        state_dict = {
            "diffusion_model.blocks.0.self_attn.q_proj.lora_A.weight": torch.randn(2, 32),
            "diffusion_model.blocks.0.self_attn.q_proj.lora_B.weight": torch.randn(32, 2),
            "diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_A.weight": torch.randn(2, 16),
            "diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_B.weight": torch.randn(16, 2),
        }

        pipe.load_lora_weights(state_dict, adapter_name="dummy")

        assert "dummy" in pipe.transformer.peft_config
        assert "dummy" in pipe.text_conditioner.peft_config


class TestAnimaModularPipelineWorkflow(AnimaModularPipelineTesterConfig, ModularWorkflowTesterMixin):
    pass


class TestAnimaModularPipelineMemory(AnimaModularPipelineTesterConfig, ModularMemoryTesterMixin):
    pass


class TestAnimaModularPipelineGuider(AnimaModularPipelineTesterConfig, ModularGuiderTesterMixin):
    pass


class AnimaImg2ImgModularPipelineTesterConfig(BaseModularPipelineTesterConfig):
    pipeline_class = AnimaModularPipeline
    pipeline_blocks_class = AnimaAutoBlocks
    pretrained_model_name_or_path = "hf-internal-testing/tiny-anima-modular-pipe"
    params = frozenset(["prompt", "image", "strength", "height", "width", "negative_prompt"])
    batch_params = frozenset(["prompt", "negative_prompt"])
    expected_workflow_blocks = ANIMA_IMG2IMG_WORKFLOWS

    def get_dummy_inputs(self, seed=0):
        generator = torch.Generator(device="cpu").manual_seed(seed)
        return {
            "prompt": "dance monkey",
            "negative_prompt": "bad quality",
            "image": get_dummy_image(32, 32),
            "strength": 0.8,
            "generator": generator,
            "num_inference_steps": 2,
            "height": 32,
            "width": 32,
            "max_sequence_length": 16,
            "output_type": "pt",
        }


class TestAnimaImg2ImgModularPipelineFast(AnimaImg2ImgModularPipelineTesterConfig, ModularPipelineTesterMixin):
    def test_inference_basic(self):
        pipe = self.get_pipeline()
        inputs = self.get_dummy_inputs()
        output = pipe(**inputs).images

        assert output.shape == (1, 3, 32, 32)
        assert not torch.isnan(output).any()

    def test_inference_strength_low(self):
        pipe = self.get_pipeline()
        inputs = self.get_dummy_inputs()
        inputs["strength"] = 0.3
        output = pipe(**inputs).images

        assert output.shape == (1, 3, 32, 32)
        assert not torch.isnan(output).any()

    def test_inference_strength_high(self):
        pipe = self.get_pipeline()
        inputs = self.get_dummy_inputs()
        inputs["strength"] = 0.95
        output = pipe(**inputs).images

        assert output.shape == (1, 3, 32, 32)
        assert not torch.isnan(output).any()

    def test_inference_empty_negative_prompt(self):
        pipe = self.get_pipeline()
        inputs = self.get_dummy_inputs()
        inputs["negative_prompt"] = ""
        output = pipe(**inputs).images

        assert output.shape == (1, 3, 32, 32)
        assert not torch.isnan(output).any()

    def test_inference_batch_single_identical(self):
        super().test_inference_batch_single_identical(expected_max_diff=5e-4)


class TestAnimaImg2ImgModularPipelineLoading(AnimaImg2ImgModularPipelineTesterConfig, ModularLoadingTesterMixin):
    pass


class TestAnimaImg2ImgModularPipelineWorkflow(AnimaImg2ImgModularPipelineTesterConfig, ModularWorkflowTesterMixin):
    pass


class TestAnimaImg2ImgModularPipelineMemory(AnimaImg2ImgModularPipelineTesterConfig, ModularMemoryTesterMixin):
    pass
