Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
255 changes: 94 additions & 161 deletions tests/pipelines/wan/test_wan_image_to_video.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -114,117 +100,87 @@ 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])
# fmt: on

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)
Expand Down Expand Up @@ -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,
Expand All @@ -286,97 +242,74 @@ 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])
# fmt: on

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
Loading