Files
wehub-resource-sync eec33d25b2
pre-commit / pre-commit (push) Failing after 1s
Build Wheel / build (3.11) (push) Failing after 1s
Build Wheel / build (3.12) (push) Failing after 0s
chore: import upstream snapshot with attribution
2026-07-13 12:29:08 +08:00

426 lines
16 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Unit tests for TeaCache extractor functions.
This module provides a generic testing framework for model-specific extractor functions
used by TeaCache. Each model's extractor can be tested by:
1. Creating a fixture that returns model module
2. Creating a fixture that returns sample inputs for that model
3. Creating a test class that inherits from BaseExtractorTest
4. Implementing any model-specific test methods
Currently implemented:
- TestFlux2KleinExtractor: Flux2Klein model extractor
- TestFlux2Extractor: Flux2 model extractor
- TestFluxExtractor: Flux model extractor
"""
from abc import ABC, abstractmethod
from unittest.mock import MagicMock, Mock, patch
import pytest
import torch
from tests.helpers.mark import hardware_test
from vllm_omni.diffusion.cache.teacache.extractors import (
extract_flux2_context,
extract_flux2_klein_context,
extract_flux_context,
)
from vllm_omni.diffusion.models.flux.flux_transformer import FluxTransformer2DModel
from vllm_omni.diffusion.models.flux2_klein.flux2_klein_transformer import (
Flux2Transformer2DModel,
)
pytestmark = [pytest.mark.core_model]
@pytest.fixture(scope="function", autouse=True)
def setup_tp_group():
"""Set up TP group for each test function"""
with patch("vllm.model_executor.layers.linear.get_tensor_model_parallel_world_size", return_value=1):
with patch("vllm.distributed.parallel_state.get_tp_group") as mock_get_tp_group:
mock_tp_group = MagicMock()
mock_tp_group.world_size = 1
mock_get_tp_group.return_value = mock_tp_group
yield
class BaseExtractorTest(ABC):
"""Base class for testing TeaCache extractors.
Subclasses should implement:
- get_extractor(): Return extractor function
- get_module(): Return model module
- get_sample_inputs(): Return sample inputs for model
"""
@abstractmethod
def get_extractor(self):
"""Return extractor function to test."""
pass
@abstractmethod
def get_module(self):
"""Return model module instance."""
pass
@abstractmethod
def get_sample_inputs(self):
"""Return sample inputs for model."""
pass
class TestFlux2KleinExtractor(BaseExtractorTest):
"""Test extract_flux2_klein_context function."""
def get_extractor(self):
return extract_flux2_klein_context
@pytest.fixture
def flux2_klein_module(self):
"""Create a minimal Flux2Transformer2DModel for testing."""
model = Flux2Transformer2DModel(
num_layers=2,
num_single_layers=2,
num_attention_heads=48,
attention_head_dim=128,
joint_attention_dim=15360,
)
return model
def get_module(self, flux2_klein_module):
return flux2_klein_module
@pytest.fixture
def sample_inputs(self):
"""Create sample input tensors for Flux2Klein.
Note: hidden_states uses in_channels=128 (default for Flux2Klein),
not inner_dim=6144. The x_embedder projects from 128 -> 6144.
encoder_hidden_states uses joint_attention_dim=15360 (model default),
which then gets projected to inner_dim=6144 by context_embedder.
"""
batch_size = 1
img_seq_len = 1024
txt_seq_len = 512
in_channels = 128 # Model default in_channels
txt_dim = 15360 # Model default joint_attention_dim
return {
"hidden_states": torch.randn(batch_size, img_seq_len, in_channels),
"encoder_hidden_states": torch.randn(batch_size, txt_seq_len, txt_dim),
"timestep": torch.tensor([500]),
"img_ids": torch.randint(0, 64, (batch_size, img_seq_len, 4)),
"txt_ids": torch.randint(0, 64, (batch_size, txt_seq_len, 4)),
"guidance": torch.tensor([3.5]),
}
def get_sample_inputs(self, sample_inputs):
return sample_inputs
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_modulated_input_shape(self, flux2_klein_module, sample_inputs):
"""Test that modulated_input has correct shape matching the model's inner_dim.
Note: After x_embedder projection, hidden_states are projected from
in_channels (128) to inner_dim (6144), so modulated_input should match
the projected shape, not the input shape.
"""
context = extract_flux2_klein_context(flux2_klein_module, **sample_inputs)
batch_size, img_seq_len, _ = sample_inputs["hidden_states"].shape
inner_dim = flux2_klein_module.inner_dim
assert context.modulated_input.shape == (batch_size, img_seq_len, inner_dim)
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_run_transformer_blocks_callable(self, flux2_klein_module, sample_inputs):
"""Test that run_transformer_blocks is callable."""
context = extract_flux2_klein_context(flux2_klein_module, **sample_inputs)
assert callable(context.run_transformer_blocks)
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_postprocess_callable(self, flux2_klein_module, sample_inputs):
"""Test that postprocess is callable."""
context = extract_flux2_klein_context(flux2_klein_module, **sample_inputs)
assert callable(context.postprocess)
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_extra_states_contains_full_transformer(self, flux2_klein_module, sample_inputs):
"""Test that extra_states contains run_flux2_full_transformer_with_single."""
context = extract_flux2_klein_context(flux2_klein_module, **sample_inputs)
assert context.extra_states is not None
assert "run_flux2_full_transformer_with_single" in context.extra_states
assert callable(context.extra_states["run_flux2_full_transformer_with_single"])
def test_without_guidance(self, flux2_klein_module, sample_inputs):
"""Test context extraction works without guidance (no CFG)."""
inputs = sample_inputs.copy()
inputs["guidance"] = None
context = extract_flux2_klein_context(flux2_klein_module, **inputs)
assert context is not None
assert context.temb is not None
@pytest.mark.cpu
def test_invalid_module_raises_error(self):
"""Test that invalid module without transformer_blocks raises ValueError."""
invalid_module = Mock()
invalid_module.transformer_blocks = []
with pytest.raises(ValueError, match="Module must have transformer_blocks"):
extract_flux2_klein_context(
invalid_module,
hidden_states=torch.randn(1, 1024, 6144),
encoder_hidden_states=torch.randn(1, 512, 15360),
timestep=torch.tensor([500]),
img_ids=torch.randint(0, 64, (1, 1024, 4)),
txt_ids=torch.randint(0, 64, (1, 512, 4)),
)
class TestFlux2Extractor(BaseExtractorTest):
"""Test extract_flux2_context function."""
def get_extractor(self):
return extract_flux2_context
@pytest.fixture
def flux2_module(self):
"""Create a minimal Flux2Transformer2DModel for testing."""
from vllm_omni.diffusion.models.flux2.flux2_transformer import Flux2Transformer2DModel
model = Flux2Transformer2DModel(
num_layers=2,
num_single_layers=2,
num_attention_heads=48,
attention_head_dim=128,
joint_attention_dim=15360,
)
return model
def get_module(self, flux2_module):
return flux2_module
@pytest.fixture
def sample_inputs(self):
"""Create sample input tensors for Flux2.
Note: hidden_states uses in_channels=128 (default for Flux2),
not inner_dim=6144. The x_embedder projects from 128 -> 6144.
encoder_hidden_states uses joint_attention_dim=15360 (model default),
which then gets projected to inner_dim=6144 by context_embedder.
"""
batch_size = 1
img_seq_len = 1024
txt_seq_len = 512
in_channels = 128 # Model default in_channels
txt_dim = 15360 # Model default joint_attention_dim
return {
"hidden_states": torch.randn(batch_size, img_seq_len, in_channels),
"encoder_hidden_states": torch.randn(batch_size, txt_seq_len, txt_dim),
"timestep": torch.tensor([500]),
"img_ids": torch.randint(0, 64, (batch_size, img_seq_len, 4)),
"txt_ids": torch.randint(0, 64, (batch_size, txt_seq_len, 4)),
"guidance": torch.tensor([3.5]),
}
def get_sample_inputs(self, sample_inputs):
return sample_inputs
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_modulated_input_shape(self, flux2_module, sample_inputs):
"""Test that modulated_input has correct shape matching the model's inner_dim.
Note: After x_embedder projection, hidden_states are projected from
in_channels (128) to inner_dim (6144), so modulated_input should match
the projected shape, not the input shape.
"""
context = extract_flux2_context(flux2_module, **sample_inputs)
batch_size, img_seq_len, _ = sample_inputs["hidden_states"].shape
inner_dim = flux2_module.inner_dim
assert context.modulated_input.shape == (batch_size, img_seq_len, inner_dim)
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_run_transformer_blocks_callable(self, flux2_module, sample_inputs):
"""Test that run_transformer_blocks is callable."""
context = extract_flux2_context(flux2_module, **sample_inputs)
assert callable(context.run_transformer_blocks)
@hardware_test(res={"cuda": "L4"}, num_cards=1)
def test_postprocess_callable(self, flux2_module, sample_inputs):
"""Test that postprocess is callable."""
context = extract_flux2_context(flux2_module, **sample_inputs)
assert callable(context.postprocess)
def test_without_guidance(self, flux2_module, sample_inputs):
"""Test context extraction works without guidance (no CFG)."""
inputs = sample_inputs.copy()
inputs["guidance"] = None
context = extract_flux2_context(flux2_module, **inputs)
assert context is not None
assert context.temb is not None
@pytest.mark.cpu
def test_invalid_module_raises_error(self):
"""Test that invalid module without transformer_blocks raises ValueError."""
invalid_module = Mock()
invalid_module.transformer_blocks = []
with pytest.raises(ValueError, match="Module must have transformer_blocks"):
extract_flux2_context(
invalid_module,
hidden_states=torch.randn(1, 1024, 6144),
encoder_hidden_states=torch.randn(1, 512, 15360),
timestep=torch.tensor([500]),
img_ids=torch.randint(0, 64, (1, 1024, 4)),
txt_ids=torch.randint(0, 64, (1, 512, 4)),
)
@pytest.mark.cpu
class TestFluxExtractor(BaseExtractorTest):
"""Test extract_flux_context function."""
@pytest.fixture(autouse=True)
def cpu_vllm_config(self):
"""Force CPU custom-op dispatch for this test class."""
from vllm.config import DeviceConfig, VllmConfig, set_current_vllm_config
with set_current_vllm_config(VllmConfig(device_config=DeviceConfig(device="cpu"))):
yield
@pytest.fixture(autouse=True)
def mock_flux_attention_backend(self):
"""Use the SDPA backend so FLUX can be instantiated in CPU tests."""
from vllm_omni.diffusion.attention.backends.sdpa import SDPABackend
with patch(
"vllm_omni.diffusion.attention.layer.get_attn_backend_for_role",
return_value=(SDPABackend, None),
):
yield
def get_extractor(self):
return extract_flux_context
@pytest.fixture
def flux_module(self):
"""Create a minimal FluxTransformer2DModel for testing."""
return FluxTransformer2DModel(
num_layers=2,
num_single_layers=2,
num_attention_heads=2,
attention_head_dim=16,
joint_attention_dim=32,
pooled_projection_dim=16,
axes_dims_rope=(4, 4, 8),
)
@pytest.fixture
def flux_module_without_guidance(self):
"""Create a minimal non-guidance-distilled FLUX transformer."""
return FluxTransformer2DModel(
num_layers=2,
num_single_layers=2,
num_attention_heads=2,
attention_head_dim=16,
joint_attention_dim=32,
pooled_projection_dim=16,
guidance_embeds=False,
axes_dims_rope=(4, 4, 8),
)
def get_module(self, flux_module):
return flux_module
@pytest.fixture
def sample_inputs(self):
"""Create sample input tensors for Flux."""
batch_size = 1
img_seq_len = 16
txt_seq_len = 8
in_channels = 64 # Flux default in_channels
txt_dim = 32
pooled_dim = 16
return {
"hidden_states": torch.randn(batch_size, img_seq_len, in_channels),
"encoder_hidden_states": torch.randn(batch_size, txt_seq_len, txt_dim),
"pooled_projections": torch.randn(batch_size, pooled_dim),
"timestep": torch.tensor([500]),
"img_ids": torch.randint(0, 64, (batch_size, img_seq_len, 3)),
"txt_ids": torch.randint(0, 64, (batch_size, txt_seq_len, 3)),
"guidance": torch.tensor([3.5]),
}
def get_sample_inputs(self, sample_inputs):
return sample_inputs
def test_modulated_input_shape(self, flux_module, sample_inputs):
"""Test that modulated_input has the projected FLUX inner dimension."""
context = extract_flux_context(flux_module, **sample_inputs)
batch_size, img_seq_len, _ = sample_inputs["hidden_states"].shape
assert context.modulated_input.shape == (batch_size, img_seq_len, flux_module.inner_dim)
def test_run_transformer_blocks_callable(self, flux_module, sample_inputs):
"""Test that run_transformer_blocks is callable."""
context = extract_flux_context(flux_module, **sample_inputs)
assert callable(context.run_transformer_blocks)
def test_postprocess_callable(self, flux_module, sample_inputs):
"""Test that postprocess is callable."""
context = extract_flux_context(flux_module, **sample_inputs)
assert callable(context.postprocess)
def test_postprocess_output_shape(self, flux_module, sample_inputs):
"""Test that postprocess projects back to the input channel width."""
context = extract_flux_context(flux_module, **sample_inputs)
output = context.postprocess(context.hidden_states)
assert output.sample.shape == sample_inputs["hidden_states"].shape
def test_postprocess_return_tuple_when_return_dict_false(self, flux_module, sample_inputs):
"""Test that postprocess honors return_dict=False."""
context = extract_flux_context(flux_module, **sample_inputs, return_dict=False)
output = context.postprocess(context.hidden_states)
assert isinstance(output, tuple)
assert len(output) == 1
assert output[0].shape == sample_inputs["hidden_states"].shape
def test_without_guidance(self, flux_module_without_guidance, sample_inputs):
"""Test context extraction works for FLUX variants without guidance embeddings."""
inputs = sample_inputs.copy()
inputs["guidance"] = None
context = extract_flux_context(flux_module_without_guidance, **inputs)
assert context is not None
assert context.temb is not None
def test_invalid_module_raises_error(self):
"""Test that invalid module without transformer_blocks raises ValueError."""
invalid_module = Mock()
invalid_module.transformer_blocks = []
with pytest.raises(ValueError, match="Module must have transformer_blocks"):
extract_flux_context(
invalid_module,
hidden_states=torch.randn(1, 16, 64),
encoder_hidden_states=torch.randn(1, 8, 32),
pooled_projections=torch.randn(1, 16),
timestep=torch.tensor([500]),
img_ids=torch.randint(0, 64, (1, 16, 3)),
txt_ids=torch.randint(0, 64, (1, 8, 3)),
)