diff --git a/tests/pipelines/wan/test_wan_image_to_video.py b/tests/pipelines/wan/test_wan_image_to_video.py index 28c58156b104..6feb1a454e7f 100644 --- a/tests/pipelines/wan/test_wan_image_to_video.py +++ b/tests/pipelines/wan/test_wan_image_to_video.py @@ -12,10 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -import tempfile -import unittest -import numpy as np import torch from PIL import Image from transformers import ( @@ -29,31 +26,20 @@ from diffusers import AutoencoderKLWan, FlowMatchEulerDiscreteScheduler, WanImageToVideoPipeline, WanTransformer3DModel -from ...testing_utils import enable_full_determinism, torch_device -from ..pipeline_params import TEXT_TO_IMAGE_BATCH_PARAMS, TEXT_TO_IMAGE_IMAGE_PARAMS, TEXT_TO_IMAGE_PARAMS -from ..test_pipelines_common import PipelineTesterMixin +from ...testing_utils import assert_tensors_close, torch_device +from ..testing_utils import BasePipelineTesterConfig, MemoryTesterMixin, PipelineTesterMixin -enable_full_determinism() - - -class WanImageToVideoPipelineFastTests(PipelineTesterMixin, unittest.TestCase): +class WanImageToVideoPipelineTesterConfig(BasePipelineTesterConfig): pipeline_class = WanImageToVideoPipeline - params = TEXT_TO_IMAGE_PARAMS - {"cross_attention_kwargs", "height", "width"} - batch_params = TEXT_TO_IMAGE_BATCH_PARAMS - image_params = TEXT_TO_IMAGE_IMAGE_PARAMS - image_latents_params = TEXT_TO_IMAGE_IMAGE_PARAMS - required_optional_params = frozenset( - [ - "num_inference_steps", - "generator", - "latents", - "return_dict", - "callback_on_step_end", - "callback_on_step_end_tensor_inputs", - ] + required_input_params_in_call_signature = frozenset( + ["image", "prompt", "negative_prompt", "guidance_scale", "prompt_embeds", "negative_prompt_embeds"] + ) + batch_input_params = frozenset(["prompt"]) + # Wan is a video pipeline: it exposes `num_videos_per_prompt`, not the base default `num_images_per_prompt`. + optional_input_params = frozenset( + ["num_inference_steps", "num_videos_per_prompt", "generator", "latents", "output_type", "return_dict"] ) - test_xformers_attention = False def get_dummy_components(self): torch.manual_seed(0) @@ -104,7 +90,7 @@ def get_dummy_components(self): torch.manual_seed(0) image_processor = CLIPImageProcessor(crop_size=32, size=32) - components = { + return { "transformer": transformer, "vae": vae, "scheduler": scheduler, @@ -114,43 +100,36 @@ def get_dummy_components(self): "image_processor": image_processor, "transformer_2": None, } - return components - def get_dummy_inputs(self, device, seed=0): - if str(device).startswith("mps"): - generator = torch.manual_seed(seed) - else: - generator = torch.Generator(device=device).manual_seed(seed) + def get_dummy_inputs(self): image_height = 16 image_width = 16 image = Image.new("RGB", (image_width, image_height)) - inputs = { + return { "image": image, "prompt": "dance monkey", "negative_prompt": "negative", # TODO "height": image_height, "width": image_width, - "generator": generator, + "generator": self.get_generator(0), "num_inference_steps": 2, "guidance_scale": 6.0, "num_frames": 9, "max_sequence_length": 16, + # Request torch outputs so tests compare torch tensors directly (see `BasePipelineTesterConfig`). "output_type": "pt", } - return inputs - def test_inference(self): - device = "cpu" - components = self.get_dummy_components() - pipe = self.pipeline_class(**components) - pipe.to(device) - pipe.set_progress_bar_config(disable=None) +class TestWanImageToVideoPipeline(WanImageToVideoPipelineTesterConfig, PipelineTesterMixin): + def test_inference(self): + # Run on CPU: the expected slice below is CPU-specific. + pipe = self.get_pipeline() - inputs = self.get_dummy_inputs(device) + inputs = self.get_dummy_inputs() video = pipe(**inputs).frames generated_video = video[0] - self.assertEqual(generated_video.shape, (9, 3, 16, 16)) + assert generated_video.shape == (9, 3, 16, 16) # fmt: off expected_slice = torch.tensor([0.4528, 0.4525, 0.4493, 0.4537, 0.4521, 0.4532, 0.4543, 0.4536, 0.5084, 0.5252, 0.5211, 0.5120, 0.5419, 0.5355, 0.5169, 0.5213]) @@ -158,73 +137,50 @@ def test_inference(self): generated_slice = generated_video.flatten() generated_slice = torch.cat([generated_slice[:8], generated_slice[-8:]]) - self.assertTrue(torch.allclose(generated_slice, expected_slice, atol=1e-3)) - - @unittest.skip("Test not supported") - def test_attention_slicing_forward_pass(self): - pass - - @unittest.skip("TODO: revisit failing as it requires a very high threshold to pass") - def test_inference_batch_single_identical(self): - pass - - # _optional_components include transformer, transformer_2 and image_encoder, image_processor, but only transformer_2 is optional for wan2.1 i2v pipeline - def test_save_load_optional_components(self, expected_max_difference=1e-4): - optional_component = "transformer_2" - - components = self.get_dummy_components() - components[optional_component] = None - pipe = self.pipeline_class(**components) - for component in pipe.components.values(): - if hasattr(component, "set_default_attn_processor"): - component.set_default_attn_processor() - pipe.to(torch_device) - pipe.set_progress_bar_config(disable=None) - - generator_device = "cpu" - inputs = self.get_dummy_inputs(generator_device) - torch.manual_seed(0) + assert torch.allclose(generated_slice, expected_slice, atol=1e-3) + + def test_save_load_optional_components(self, tmp_path, expected_max_difference=1e-4): + # `_optional_components` lists `transformer`, `transformer_2`, `image_encoder` and `image_processor`, but only + # `transformer_2` is optional for this wan2.1 i2v pipeline. The base test nulls every optional component, which + # would drop the required `transformer` and leave no denoiser, so restrict this to `transformer_2`. + pipe = self.get_pipeline().to(torch_device) + pipe.transformer_2 = None + + inputs = self.get_dummy_inputs() output = pipe(**inputs)[0] - with tempfile.TemporaryDirectory() as tmpdir: - pipe.save_pretrained(tmpdir, safe_serialization=False) - pipe_loaded = self.pipeline_class.from_pretrained(tmpdir) - for component in pipe_loaded.components.values(): - if hasattr(component, "set_default_attn_processor"): - component.set_default_attn_processor() - pipe_loaded.to(torch_device) - pipe_loaded.set_progress_bar_config(disable=None) - - self.assertTrue( - getattr(pipe_loaded, optional_component) is None, - f"`{optional_component}` did not stay set to None after loading.", - ) + pipe.save_pretrained(tmp_path, safe_serialization=False) + pipe_loaded = self.pipeline_class.from_pretrained(tmp_path) + pipe_loaded.to(torch_device) + pipe_loaded.set_progress_bar_config(disable=None) - inputs = self.get_dummy_inputs(generator_device) - torch.manual_seed(0) + assert pipe_loaded.transformer_2 is None, "`transformer_2` did not stay set to None after loading." + + inputs = self.get_dummy_inputs() output_loaded = pipe_loaded(**inputs)[0] - max_diff = np.abs(output.detach().cpu().numpy() - output_loaded.detach().cpu().numpy()).max() - self.assertLess(max_diff, expected_max_difference) + assert_tensors_close( + output_loaded, + output, + atol=expected_max_difference, + msg="Output changed after dropping the optional component.", + ) + +class TestWanImageToVideoPipelineMemory(WanImageToVideoPipelineTesterConfig, MemoryTesterMixin): + pass -class WanFLFToVideoPipelineFastTests(PipelineTesterMixin, unittest.TestCase): + +class WanFLFToVideoPipelineTesterConfig(BasePipelineTesterConfig): pipeline_class = WanImageToVideoPipeline - params = TEXT_TO_IMAGE_PARAMS - {"cross_attention_kwargs", "height", "width"} - batch_params = TEXT_TO_IMAGE_BATCH_PARAMS - image_params = TEXT_TO_IMAGE_IMAGE_PARAMS - image_latents_params = TEXT_TO_IMAGE_IMAGE_PARAMS - required_optional_params = frozenset( - [ - "num_inference_steps", - "generator", - "latents", - "return_dict", - "callback_on_step_end", - "callback_on_step_end_tensor_inputs", - ] + required_input_params_in_call_signature = frozenset( + ["image", "prompt", "negative_prompt", "guidance_scale", "prompt_embeds", "negative_prompt_embeds"] + ) + batch_input_params = frozenset(["prompt"]) + # Wan is a video pipeline: it exposes `num_videos_per_prompt`, not the base default `num_images_per_prompt`. + optional_input_params = frozenset( + ["num_inference_steps", "num_videos_per_prompt", "generator", "latents", "output_type", "return_dict"] ) - test_xformers_attention = False def get_dummy_components(self): torch.manual_seed(0) @@ -276,7 +232,7 @@ def get_dummy_components(self): torch.manual_seed(0) image_processor = CLIPImageProcessor(crop_size=4, size=4) - components = { + return { "transformer": transformer, "vae": vae, "scheduler": scheduler, @@ -286,45 +242,38 @@ def get_dummy_components(self): "image_processor": image_processor, "transformer_2": None, } - return components - def get_dummy_inputs(self, device, seed=0): - if str(device).startswith("mps"): - generator = torch.manual_seed(seed) - else: - generator = torch.Generator(device=device).manual_seed(seed) + def get_dummy_inputs(self): image_height = 16 image_width = 16 image = Image.new("RGB", (image_width, image_height)) last_image = Image.new("RGB", (image_width, image_height)) - inputs = { + return { "image": image, "last_image": last_image, "prompt": "dance monkey", "negative_prompt": "negative", "height": image_height, "width": image_width, - "generator": generator, + "generator": self.get_generator(0), "num_inference_steps": 2, "guidance_scale": 6.0, "num_frames": 9, "max_sequence_length": 16, + # Request torch outputs so tests compare torch tensors directly (see `BasePipelineTesterConfig`). "output_type": "pt", } - return inputs - def test_inference(self): - device = "cpu" - components = self.get_dummy_components() - pipe = self.pipeline_class(**components) - pipe.to(device) - pipe.set_progress_bar_config(disable=None) +class TestWanFLFToVideoPipeline(WanFLFToVideoPipelineTesterConfig, PipelineTesterMixin): + def test_inference(self): + # Run on CPU: the expected slice below is CPU-specific. + pipe = self.get_pipeline() - inputs = self.get_dummy_inputs(device) + inputs = self.get_dummy_inputs() video = pipe(**inputs).frames generated_video = video[0] - self.assertEqual(generated_video.shape, (9, 3, 16, 16)) + assert generated_video.shape == (9, 3, 16, 16) # fmt: off expected_slice = torch.tensor([0.4525, 0.4525, 0.4497, 0.4537, 0.4520, 0.4529, 0.4540, 0.4535, 0.5157, 0.5449, 0.5201, 0.5192, 0.5398, 0.5374, 0.5162, 0.5112]) @@ -332,51 +281,35 @@ def test_inference(self): generated_slice = generated_video.flatten() generated_slice = torch.cat([generated_slice[:8], generated_slice[-8:]]) - self.assertTrue(torch.allclose(generated_slice, expected_slice, atol=1e-3)) - - @unittest.skip("Test not supported") - def test_attention_slicing_forward_pass(self): - pass - - @unittest.skip("TODO: revisit failing as it requires a very high threshold to pass") - def test_inference_batch_single_identical(self): - pass - - # _optional_components include transformer, transformer_2 and image_encoder, image_processor, but only transformer_2 is optional for wan2.1 FLFT2V pipeline - def test_save_load_optional_components(self, expected_max_difference=1e-4): - optional_component = "transformer_2" - - components = self.get_dummy_components() - components[optional_component] = None - pipe = self.pipeline_class(**components) - for component in pipe.components.values(): - if hasattr(component, "set_default_attn_processor"): - component.set_default_attn_processor() - pipe.to(torch_device) - pipe.set_progress_bar_config(disable=None) - - generator_device = "cpu" - inputs = self.get_dummy_inputs(generator_device) - torch.manual_seed(0) + assert torch.allclose(generated_slice, expected_slice, atol=1e-3) + + def test_save_load_optional_components(self, tmp_path, expected_max_difference=1e-4): + # `_optional_components` lists `transformer`, `transformer_2`, `image_encoder` and `image_processor`, but only + # `transformer_2` is optional for this wan2.1 FLFT2V pipeline. The base test nulls every optional component, + # which would drop the required `transformer` and leave no denoiser, so restrict this to `transformer_2`. + pipe = self.get_pipeline().to(torch_device) + pipe.transformer_2 = None + + inputs = self.get_dummy_inputs() output = pipe(**inputs)[0] - with tempfile.TemporaryDirectory() as tmpdir: - pipe.save_pretrained(tmpdir, safe_serialization=False) - pipe_loaded = self.pipeline_class.from_pretrained(tmpdir) - for component in pipe_loaded.components.values(): - if hasattr(component, "set_default_attn_processor"): - component.set_default_attn_processor() - pipe_loaded.to(torch_device) - pipe_loaded.set_progress_bar_config(disable=None) - - self.assertTrue( - getattr(pipe_loaded, optional_component) is None, - f"`{optional_component}` did not stay set to None after loading.", - ) + pipe.save_pretrained(tmp_path, safe_serialization=False) + pipe_loaded = self.pipeline_class.from_pretrained(tmp_path) + pipe_loaded.to(torch_device) + pipe_loaded.set_progress_bar_config(disable=None) - inputs = self.get_dummy_inputs(generator_device) - torch.manual_seed(0) + assert pipe_loaded.transformer_2 is None, "`transformer_2` did not stay set to None after loading." + + inputs = self.get_dummy_inputs() output_loaded = pipe_loaded(**inputs)[0] - max_diff = np.abs(output.detach().cpu().numpy() - output_loaded.detach().cpu().numpy()).max() - self.assertLess(max_diff, expected_max_difference) + assert_tensors_close( + output_loaded, + output, + atol=expected_max_difference, + msg="Output changed after dropping the optional component.", + ) + + +class TestWanFLFToVideoPipelineMemory(WanFLFToVideoPipelineTesterConfig, MemoryTesterMixin): + pass