diff --git a/tests/pipelines/flux2/test_pipeline_flux2_klein_kv.py b/tests/pipelines/flux2/test_pipeline_flux2_klein_kv.py index 4f77579af6d6..141814b92b54 100644 --- a/tests/pipelines/flux2/test_pipeline_flux2_klein_kv.py +++ b/tests/pipelines/flux2/test_pipeline_flux2_klein_kv.py @@ -1,6 +1,4 @@ -import unittest - -import numpy as np +import pytest import torch from PIL import Image from transformers import Qwen2TokenizerFast, Qwen3Config, Qwen3ForCausalLM @@ -12,18 +10,19 @@ Flux2Transformer2DModel, ) -from ...testing_utils import torch_device -from ..test_pipelines_common import PipelineTesterMixin, check_qkv_fused_layers_exist +from ...testing_utils import assert_tensors_close, torch_device +from ..testing_utils import ( + BasePipelineTesterConfig, + MemoryTesterMixin, + PipelineTesterMixin, + check_qkv_fused_layers_exist, +) -class Flux2KleinKVPipelineFastTests(PipelineTesterMixin, unittest.TestCase): +class Flux2KleinKVPipelineTesterConfig(BasePipelineTesterConfig): pipeline_class = Flux2KleinKVPipeline - params = frozenset(["prompt", "height", "width", "prompt_embeds", "image"]) - batch_params = frozenset(["prompt"]) - - test_xformers_attention = False - test_layerwise_casting = True - test_group_offloading = True + required_input_params_in_call_signature = frozenset(["prompt", "height", "width", "prompt_embeds", "image"]) + batch_input_params = frozenset(["prompt"]) def get_dummy_components(self, num_layers: int = 1, num_single_layers: int = 1): torch.manual_seed(0) @@ -83,67 +82,70 @@ def get_dummy_components(self, num_layers: int = 1, num_single_layers: int = 1): "vae": vae, } - def get_dummy_inputs(self, device, seed=0): - if str(device).startswith("mps"): - generator = torch.manual_seed(seed) - else: - generator = torch.Generator(device="cpu").manual_seed(seed) - + def get_dummy_inputs(self): inputs = { "prompt": "a dog is dancing", "image": Image.new("RGB", (64, 64)), - "generator": generator, + "generator": self.get_generator(0), "num_inference_steps": 2, "height": 8, "width": 8, "max_sequence_length": 64, - "output_type": "np", "text_encoder_out_layers": (1,), + # Request torch outputs so tests compare torch tensors directly (see `BasePipelineTesterConfig`). + "output_type": "pt", } return inputs + +class TestFlux2KleinKVPipeline(Flux2KleinKVPipelineTesterConfig, PipelineTesterMixin): def test_fused_qkv_projections(self): - device = "cpu" # ensure determinism for the device-dependent torch.Generator - components = self.get_dummy_components() - pipe = self.pipeline_class(**components) - pipe = pipe.to(device) - pipe.set_progress_bar_config(disable=None) + # Run on CPU to keep the slice comparisons deterministic. + pipe = self.get_pipeline() - inputs = self.get_dummy_inputs(device) + inputs = self.get_dummy_inputs() image = pipe(**inputs).images - original_image_slice = image[0, -3:, -3:, -1] + original_image_slice = image[0, -1, -3:, -3:] pipe.transformer.fuse_qkv_projections() - self.assertTrue( - check_qkv_fused_layers_exist(pipe.transformer, ["to_qkv"]), - ("Something wrong with the fused attention layers. Expected all the attention projections to be fused."), + assert check_qkv_fused_layers_exist(pipe.transformer, ["to_qkv"]), ( + "Something wrong with the fused attention layers. Expected all the attention projections to be fused." ) - inputs = self.get_dummy_inputs(device) + inputs = self.get_dummy_inputs() image = pipe(**inputs).images - image_slice_fused = image[0, -3:, -3:, -1] + image_slice_fused = image[0, -1, -3:, -3:] pipe.transformer.unfuse_qkv_projections() - inputs = self.get_dummy_inputs(device) + inputs = self.get_dummy_inputs() image = pipe(**inputs).images - image_slice_disabled = image[0, -3:, -3:, -1] - - self.assertTrue( - np.allclose(original_image_slice, image_slice_fused, atol=1e-3, rtol=1e-3), - ("Fusion of QKV projections shouldn't affect the outputs."), + image_slice_disabled = image[0, -1, -3:, -3:] + + assert_tensors_close( + original_image_slice, + image_slice_fused, + atol=1e-3, + rtol=1e-3, + msg="Fusion of QKV projections shouldn't affect the outputs.", ) - self.assertTrue( - np.allclose(image_slice_fused, image_slice_disabled, atol=1e-3, rtol=1e-3), - ("Outputs, with QKV projection fusion enabled, shouldn't change when fused QKV projections are disabled."), + assert_tensors_close( + image_slice_fused, + image_slice_disabled, + atol=1e-3, + rtol=1e-3, + msg="Outputs, with QKV projection fusion enabled, shouldn't change when fused QKV projections are disabled.", ) - self.assertTrue( - np.allclose(original_image_slice, image_slice_disabled, atol=1e-2, rtol=1e-2), - ("Original outputs should match when fused QKV projections are disabled."), + assert_tensors_close( + original_image_slice, + image_slice_disabled, + atol=1e-2, + rtol=1e-2, + msg="Original outputs should match when fused QKV projections are disabled.", ) def test_image_output_shape(self): - pipe = self.pipeline_class(**self.get_dummy_components()).to(torch_device) - inputs = self.get_dummy_inputs(torch_device) + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() height_width_pairs = [(32, 32), (72, 57)] for height, width in height_width_pairs: @@ -152,21 +154,22 @@ def test_image_output_shape(self): inputs.update({"height": height, "width": width}) image = pipe(**inputs).images[0] - output_height, output_width, _ = image.shape - self.assertEqual( - (output_height, output_width), - (expected_height, expected_width), - f"Output shape {image.shape} does not match expected shape {(expected_height, expected_width)}", + _, output_height, output_width = image.shape + assert (output_height, output_width) == (expected_height, expected_width), ( + f"Output shape {image.shape} does not match expected shape {(expected_height, expected_width)}" ) def test_without_image(self): - device = "cpu" - pipe = self.pipeline_class(**self.get_dummy_components()).to(device) - inputs = self.get_dummy_inputs(device) + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() del inputs["image"] image = pipe(**inputs).images - self.assertEqual(image.shape, (1, 8, 8, 3)) + assert image.shape == (1, 3, 8, 8) - @unittest.skip("Needs to be revisited") + @pytest.mark.skip("Needs to be revisited") def test_encode_prompt_works_in_isolation(self): pass + + +class TestFlux2KleinKVPipelineMemory(Flux2KleinKVPipelineTesterConfig, MemoryTesterMixin): + """Memory optimization tests (CPU offload, group offload, layerwise casting) for the Flux2 Klein KV pipeline."""