Coverage for tests/test_transformer_flux2.py: 100%
169 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
2"""Tests for wrapper/flux2/transformer_flux2.py."""
4import sys
6from typing import Any
7from unittest.mock import patch, MagicMock
8from tests.torch_mock import TorchMock
9from tests.diffusers_mock import DiffusersMock
11mock_torch = TorchMock()
12mock_diffusers = DiffusersMock()
14sys.path.append("wrapper")
15sys.path.append("wrapper/flux2")
17# ---------------------------------------------------------------------------
18# Stub parent classes
19# These replace the diffusers / xfuser base classes so that the real code in
20# transformer_flux2.py can be imported, instantiated, and exercised without
21# requiring the actual GPU libraries.
22# ---------------------------------------------------------------------------
25class _FakeFlux2AttnProcessor:
26 """Minimal stub for diffusers.Flux2AttnProcessor."""
28 def __init__(self) -> None:
29 pass
32class _FakeFlux2ParallelSelfAttnProcessor:
33 """Minimal stub for diffusers.Flux2ParallelSelfAttnProcessor."""
35 def __init__(self) -> None:
36 pass
39class _FakeFlux2Transformer2DModel:
40 """Minimal stub for diffusers.Flux2Transformer2DModel."""
42 def __init__(self, **kwargs: Any) -> None:
43 # Provide two mock blocks in each list so __init__ loops execute.
44 self.transformer_blocks = [MagicMock(), MagicMock()]
45 self.single_transformer_blocks = [MagicMock(), MagicMock()]
47 def forward(
48 self,
49 hidden_states: Any,
50 encoder_hidden_states: Any = None,
51 *args: Any,
52 **kwargs: Any,
53 ) -> Any:
54 # Return a plain tuple so the wrapper's return_dict logic takes the
55 # tuple branch (not the dict branch).
56 return (hidden_states,)
59class _FakeTransformerOutput:
60 """Fake non-tuple transformer output for testing the return_dict=True branch."""
62 def __init__(self, sample: Any = None) -> None:
63 self._sample = sample
65 def __getitem__(self, idx: Any) -> Any:
66 if isinstance(idx, slice):
67 # Return an empty list so that `*output[1:]` unpacking produces no
68 # extra arguments when the code under test calls
69 # `output.__class__(sample, *output[1:])` in the return_dict branch.
70 return []
71 return self._sample
74class _FakeXFuserAttentionBaseWrapper:
75 """Minimal stub for xfuser xFuserAttentionBaseWrapper."""
77 def __init__(self, attention: Any) -> None:
78 self.attention = attention
80 def forward(self, *args: Any, **kwargs: Any) -> Any:
81 return MagicMock()
84# ---------------------------------------------------------------------------
85# Decorator helpers
86# The @register(...) decorators run at import time. We must make them act as
87# identity decorators so the class bodies in transformer_flux2.py survive.
88# ---------------------------------------------------------------------------
91def _identity_register(base_cls: Any) -> Any:
92 """Return a decorator that leaves the decorated class unchanged."""
93 def decorator(cls: Any) -> Any:
94 return cls
95 return decorator
98# ---------------------------------------------------------------------------
99# Mock modules
100# ---------------------------------------------------------------------------
102_mock_attn_proc_register = MagicMock()
103_mock_attn_proc_register.register.side_effect = _identity_register
104_mock_attn_proc_register.get_processor.return_value = MagicMock(return_value=MagicMock())
106_mock_attn_proc_module = MagicMock()
107_mock_attn_proc_module.xFuserAttentionBaseWrapper = _FakeXFuserAttentionBaseWrapper
108_mock_attn_proc_module.xFuserAttentionProcessorRegister = _mock_attn_proc_register
110_mock_layer_wrappers_register = MagicMock()
111_mock_layer_wrappers_register.register.side_effect = _identity_register
113_mock_layers_module = MagicMock()
114_mock_layers_module.xFuserLayerWrappersRegister = _mock_layer_wrappers_register
116_mock_sp_group = MagicMock()
117_mock_sp_group.all_gather.side_effect = lambda x, dim: x # identity
119_mock_cfg_group = MagicMock()
120_mock_cfg_group.all_gather.side_effect = lambda x, dim: x # identity
122_mock_distributed = MagicMock()
123_mock_distributed.get_sequence_parallel_world_size.return_value = 1
124_mock_distributed.get_sequence_parallel_rank.return_value = 0
125_mock_distributed.get_classifier_free_guidance_world_size.return_value = 1
126_mock_distributed.get_classifier_free_guidance_rank.return_value = 0
127_mock_distributed.get_sp_group.return_value = _mock_sp_group
128_mock_distributed.get_cfg_group.return_value = _mock_cfg_group
130_mock_transformer_module = MagicMock()
131_mock_transformer_module.Flux2Attention = MagicMock
132_mock_transformer_module.Flux2AttnProcessor = _FakeFlux2AttnProcessor
133_mock_transformer_module.Flux2Transformer2DModel = _FakeFlux2Transformer2DModel
134_mock_transformer_module.Flux2ParallelSelfAttention = MagicMock
135_mock_transformer_module.Flux2ParallelSelfAttnProcessor = _FakeFlux2ParallelSelfAttnProcessor
136_mock_transformer_module._get_qkv_projections = MagicMock(
137 return_value=(
138 MagicMock(), MagicMock(), MagicMock(),
139 MagicMock(), MagicMock(), MagicMock(),
140 )
141)
143_mock_embeddings = MagicMock()
144_mock_embeddings.apply_rotary_emb.side_effect = lambda q, emb, sequence_dim: q # identity
146_mock_usp_module = MagicMock()
148# Build module dict from DiffusersMock and override the two entries that need
149# custom behaviour for these transformer tests.
150_diffusers_sub_modules = mock_diffusers.get_sub_modules()
151_diffusers_sub_modules["diffusers.models.transformers.transformer_flux2"] = _mock_transformer_module
152_diffusers_sub_modules["diffusers.models.embeddings"] = _mock_embeddings
154mock_modules = {
155 'torch': mock_torch,
156 'xfuser': MagicMock(),
157 'xfuser.config': MagicMock(),
158 'xfuser.core': MagicMock(),
159 'xfuser.core.distributed': _mock_distributed,
160 'xfuser.model_executor': MagicMock(),
161 'xfuser.model_executor.models': MagicMock(),
162 'xfuser.model_executor.models.transformers': MagicMock(),
163 'xfuser.model_executor.layers': _mock_layers_module,
164 'xfuser.model_executor.layers.attention_processor': _mock_attn_proc_module,
165 'xfuser.model_executor.layers.usp': _mock_usp_module,
166}
167mock_modules.update(mock_torch.get_sub_modules())
168mock_modules.update(_diffusers_sub_modules)
170with patch.dict(sys.modules, mock_modules):
171 from transformer_flux2 import (
172 xFuserFlux2AttnProcessor,
173 xFuserFlux2ParallelSelfAttnProcessor,
174 xFuserFlux2ParallelSelfAttention,
175 xFuserFlux2Transformer2DWrapper,
176 )
179# ---------------------------------------------------------------------------
180# Tests
181# ---------------------------------------------------------------------------
184def test_attn_processor_no_encoder_no_rotary() -> None:
185 """xFuserFlux2AttnProcessor: no encoder states, no rotary embedding."""
186 processor = xFuserFlux2AttnProcessor()
187 assert processor is not None
189 attn = MagicMock()
190 attn.added_kv_proj_dim = None # skip encoder branch
192 result = processor(
193 attn,
194 hidden_states=MagicMock(),
195 encoder_hidden_states=None,
196 attention_mask=None,
197 image_rotary_emb=None,
198 )
199 assert result is not None
202def test_attn_processor_with_encoder_and_rotary() -> None:
203 """xFuserFlux2AttnProcessor: encoder states + rotary embedding present."""
204 processor = xFuserFlux2AttnProcessor()
206 attn = MagicMock()
207 attn.added_kv_proj_dim = 64 # enter encoder branch
209 encoder_hidden_states = MagicMock()
210 encoder_hidden_states.shape = [1, 10, 64] # list — supports indexing
212 result = processor(
213 attn,
214 hidden_states=MagicMock(),
215 encoder_hidden_states=encoder_hidden_states,
216 attention_mask=None,
217 image_rotary_emb=MagicMock(),
218 )
219 assert result is not None
222def test_parallel_self_attn_processor_no_rotary() -> None:
223 """xFuserFlux2ParallelSelfAttnProcessor.__call__ without rotary embedding."""
224 processor = xFuserFlux2ParallelSelfAttnProcessor()
225 assert processor is not None
227 attn = MagicMock()
228 hidden_states = MagicMock()
230 # torch.split must return an iterable of exactly 2 elements.
231 qkv_mock = MagicMock()
232 qkv_mock.chunk.return_value = [MagicMock(), MagicMock(), MagicMock()]
233 mock_torch.split = MagicMock(return_value=(qkv_mock, MagicMock()))
235 result = processor(attn, hidden_states, attention_mask=None, image_rotary_emb=None)
236 assert result is not None
239def test_parallel_self_attn_processor_with_rotary() -> None:
240 """xFuserFlux2ParallelSelfAttnProcessor.__call__ with rotary embedding."""
241 processor = xFuserFlux2ParallelSelfAttnProcessor()
243 attn = MagicMock()
244 hidden_states = MagicMock()
246 qkv_mock = MagicMock()
247 qkv_mock.chunk.return_value = [MagicMock(), MagicMock(), MagicMock()]
248 mock_torch.split = MagicMock(return_value=(qkv_mock, MagicMock()))
250 result = processor(
251 attn,
252 hidden_states,
253 attention_mask=None,
254 image_rotary_emb=MagicMock(),
255 )
256 assert result is not None
259def test_parallel_self_attention_init_and_forward() -> None:
260 """xFuserFlux2ParallelSelfAttention: init sets processor; forward delegates."""
261 mock_attention = MagicMock()
262 wrapper = xFuserFlux2ParallelSelfAttention(mock_attention)
264 assert wrapper is not None
265 assert wrapper.attention is mock_attention
267 hidden_states = MagicMock()
268 result = wrapper.forward(hidden_states, attention_mask=None, image_rotary_emb=None)
269 assert result is not None
272def test_transformer_wrapper_init() -> None:
273 """xFuserFlux2Transformer2DWrapper.__init__ wires processors onto each block."""
274 wrapper = xFuserFlux2Transformer2DWrapper()
275 assert wrapper is not None
277 # Both block lists are populated by _FakeFlux2Transformer2DModel (2 each).
278 for block in wrapper.transformer_blocks:
279 assert isinstance(block.attn.processor, xFuserFlux2AttnProcessor)
280 for block in wrapper.single_transformer_blocks:
281 assert isinstance(block.attn.processor, xFuserFlux2ParallelSelfAttnProcessor)
284def test_pad_to_sp_divisible() -> None:
285 """_pad_to_sp_divisible appends zeros along the specified dimension."""
286 wrapper = xFuserFlux2Transformer2DWrapper()
288 tensor = MagicMock()
289 tensor.shape = [2, 5, 16]
290 tensor.dtype = MagicMock()
291 tensor.device = MagicMock()
293 result = wrapper._pad_to_sp_divisible(tensor, padding_length=3, dim=1)
294 assert result is not None
297def test_transformer_wrapper_forward() -> None:
298 """xFuserFlux2Transformer2DWrapper.forward runs end-to-end with sp_world_size=1."""
299 wrapper = xFuserFlux2Transformer2DWrapper()
301 # Use a list for shape so integer indexing returns real ints.
302 hidden_states = MagicMock()
303 hidden_states.shape = [1, 4, 64] # sequence_length=4, divisible by sp_world_size=1
305 encoder_hidden_states = MagicMock()
306 img_ids = MagicMock()
307 txt_ids = MagicMock()
308 timestep = MagicMock() # not a torch.Tensor instance → skips CFG-chunk branch
310 result = wrapper.forward(
311 hidden_states,
312 encoder_hidden_states=encoder_hidden_states,
313 timestep=timestep,
314 img_ids=img_ids,
315 txt_ids=txt_ids,
316 )
317 assert result is not None
320def test_transformer_wrapper_forward_with_padding() -> None:
321 """forward() pads hidden_states / img_ids when sequence length % sp_world_size != 0."""
322 wrapper = xFuserFlux2Transformer2DWrapper()
324 # sp_world_size=2, sequence_length=3 → padding_length=1, triggering the
325 # padding path: hidden_states and img_ids are padded to a length divisible
326 # by sp_world_size, and the extra padding tokens are stripped from the
327 # gathered output before returning.
328 _mock_distributed.get_sequence_parallel_world_size.return_value = 2
329 try:
330 hidden_states = MagicMock()
331 hidden_states.shape = [1, 3, 64] # seq_len=3 is not divisible by 2
333 result = wrapper.forward(
334 hidden_states,
335 encoder_hidden_states=MagicMock(),
336 timestep=MagicMock(),
337 img_ids=MagicMock(),
338 txt_ids=MagicMock(),
339 )
340 assert result is not None
341 finally:
342 _mock_distributed.get_sequence_parallel_world_size.return_value = 1
345def test_transformer_wrapper_forward_timestep_tensor() -> None:
346 """forward() chunks the timestep tensor when it is a real Tensor with ndim > 0."""
347 wrapper = xFuserFlux2Transformer2DWrapper()
349 # Build a timestep that satisfies isinstance(timestep, torch.Tensor) check.
350 timestep = mock_torch.Tensor()
351 timestep.ndim = 1
352 timestep.shape = [1] # shape[0]=1 matches hidden_states.shape[0]=1
354 hidden_states = MagicMock()
355 hidden_states.shape = [1, 4, 64]
357 result = wrapper.forward(
358 hidden_states,
359 encoder_hidden_states=MagicMock(),
360 timestep=timestep,
361 img_ids=MagicMock(),
362 txt_ids=MagicMock(),
363 )
364 assert result is not None
367def test_transformer_wrapper_forward_return_dict() -> None:
368 """forward() returns output.__class__(sample, ...) when output is not a tuple."""
369 wrapper = xFuserFlux2Transformer2DWrapper()
371 hidden_states = MagicMock()
372 hidden_states.shape = [1, 4, 64]
374 # Patch super().forward to return a _FakeTransformerOutput (not a tuple)
375 # → return_dict=True, which exercises line 298.
376 non_tuple_output = _FakeTransformerOutput(sample=MagicMock())
377 with patch.object(_FakeFlux2Transformer2DModel, "forward", return_value=non_tuple_output):
378 result = wrapper.forward(
379 hidden_states,
380 encoder_hidden_states=MagicMock(),
381 timestep=MagicMock(),
382 img_ids=MagicMock(),
383 txt_ids=MagicMock(),
384 )
385 assert result is not None