Skip to content
Open
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
115 changes: 59 additions & 56 deletions tests/pipelines/flux2/test_pipeline_flux2_klein_kv.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,4 @@
import unittest

import numpy as np
import pytest
import torch
from PIL import Image
from transformers import Qwen2TokenizerFast, Qwen3Config, Qwen3ForCausalLM
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand All @@ -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."""
Loading