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

1#!/usr/bin/env python3 

2"""Tests for wrapper/flux2/transformer_flux2.py.""" 

3 

4import sys 

5 

6from typing import Any 

7from unittest.mock import patch, MagicMock 

8from tests.torch_mock import TorchMock 

9from tests.diffusers_mock import DiffusersMock 

10 

11mock_torch = TorchMock() 

12mock_diffusers = DiffusersMock() 

13 

14sys.path.append("wrapper") 

15sys.path.append("wrapper/flux2") 

16 

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# --------------------------------------------------------------------------- 

23 

24 

25class _FakeFlux2AttnProcessor: 

26 """Minimal stub for diffusers.Flux2AttnProcessor.""" 

27 

28 def __init__(self) -> None: 

29 pass 

30 

31 

32class _FakeFlux2ParallelSelfAttnProcessor: 

33 """Minimal stub for diffusers.Flux2ParallelSelfAttnProcessor.""" 

34 

35 def __init__(self) -> None: 

36 pass 

37 

38 

39class _FakeFlux2Transformer2DModel: 

40 """Minimal stub for diffusers.Flux2Transformer2DModel.""" 

41 

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()] 

46 

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,) 

57 

58 

59class _FakeTransformerOutput: 

60 """Fake non-tuple transformer output for testing the return_dict=True branch.""" 

61 

62 def __init__(self, sample: Any = None) -> None: 

63 self._sample = sample 

64 

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 

72 

73 

74class _FakeXFuserAttentionBaseWrapper: 

75 """Minimal stub for xfuser xFuserAttentionBaseWrapper.""" 

76 

77 def __init__(self, attention: Any) -> None: 

78 self.attention = attention 

79 

80 def forward(self, *args: Any, **kwargs: Any) -> Any: 

81 return MagicMock() 

82 

83 

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# --------------------------------------------------------------------------- 

89 

90 

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 

96 

97 

98# --------------------------------------------------------------------------- 

99# Mock modules 

100# --------------------------------------------------------------------------- 

101 

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()) 

105 

106_mock_attn_proc_module = MagicMock() 

107_mock_attn_proc_module.xFuserAttentionBaseWrapper = _FakeXFuserAttentionBaseWrapper 

108_mock_attn_proc_module.xFuserAttentionProcessorRegister = _mock_attn_proc_register 

109 

110_mock_layer_wrappers_register = MagicMock() 

111_mock_layer_wrappers_register.register.side_effect = _identity_register 

112 

113_mock_layers_module = MagicMock() 

114_mock_layers_module.xFuserLayerWrappersRegister = _mock_layer_wrappers_register 

115 

116_mock_sp_group = MagicMock() 

117_mock_sp_group.all_gather.side_effect = lambda x, dim: x # identity 

118 

119_mock_cfg_group = MagicMock() 

120_mock_cfg_group.all_gather.side_effect = lambda x, dim: x # identity 

121 

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 

129 

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) 

142 

143_mock_embeddings = MagicMock() 

144_mock_embeddings.apply_rotary_emb.side_effect = lambda q, emb, sequence_dim: q # identity 

145 

146_mock_usp_module = MagicMock() 

147 

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 

153 

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) 

169 

170with patch.dict(sys.modules, mock_modules): 

171 from transformer_flux2 import ( 

172 xFuserFlux2AttnProcessor, 

173 xFuserFlux2ParallelSelfAttnProcessor, 

174 xFuserFlux2ParallelSelfAttention, 

175 xFuserFlux2Transformer2DWrapper, 

176 ) 

177 

178 

179# --------------------------------------------------------------------------- 

180# Tests 

181# --------------------------------------------------------------------------- 

182 

183 

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 

188 

189 attn = MagicMock() 

190 attn.added_kv_proj_dim = None # skip encoder branch 

191 

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 

200 

201 

202def test_attn_processor_with_encoder_and_rotary() -> None: 

203 """xFuserFlux2AttnProcessor: encoder states + rotary embedding present.""" 

204 processor = xFuserFlux2AttnProcessor() 

205 

206 attn = MagicMock() 

207 attn.added_kv_proj_dim = 64 # enter encoder branch 

208 

209 encoder_hidden_states = MagicMock() 

210 encoder_hidden_states.shape = [1, 10, 64] # list — supports indexing 

211 

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 

220 

221 

222def test_parallel_self_attn_processor_no_rotary() -> None: 

223 """xFuserFlux2ParallelSelfAttnProcessor.__call__ without rotary embedding.""" 

224 processor = xFuserFlux2ParallelSelfAttnProcessor() 

225 assert processor is not None 

226 

227 attn = MagicMock() 

228 hidden_states = MagicMock() 

229 

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())) 

234 

235 result = processor(attn, hidden_states, attention_mask=None, image_rotary_emb=None) 

236 assert result is not None 

237 

238 

239def test_parallel_self_attn_processor_with_rotary() -> None: 

240 """xFuserFlux2ParallelSelfAttnProcessor.__call__ with rotary embedding.""" 

241 processor = xFuserFlux2ParallelSelfAttnProcessor() 

242 

243 attn = MagicMock() 

244 hidden_states = MagicMock() 

245 

246 qkv_mock = MagicMock() 

247 qkv_mock.chunk.return_value = [MagicMock(), MagicMock(), MagicMock()] 

248 mock_torch.split = MagicMock(return_value=(qkv_mock, MagicMock())) 

249 

250 result = processor( 

251 attn, 

252 hidden_states, 

253 attention_mask=None, 

254 image_rotary_emb=MagicMock(), 

255 ) 

256 assert result is not None 

257 

258 

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) 

263 

264 assert wrapper is not None 

265 assert wrapper.attention is mock_attention 

266 

267 hidden_states = MagicMock() 

268 result = wrapper.forward(hidden_states, attention_mask=None, image_rotary_emb=None) 

269 assert result is not None 

270 

271 

272def test_transformer_wrapper_init() -> None: 

273 """xFuserFlux2Transformer2DWrapper.__init__ wires processors onto each block.""" 

274 wrapper = xFuserFlux2Transformer2DWrapper() 

275 assert wrapper is not None 

276 

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) 

282 

283 

284def test_pad_to_sp_divisible() -> None: 

285 """_pad_to_sp_divisible appends zeros along the specified dimension.""" 

286 wrapper = xFuserFlux2Transformer2DWrapper() 

287 

288 tensor = MagicMock() 

289 tensor.shape = [2, 5, 16] 

290 tensor.dtype = MagicMock() 

291 tensor.device = MagicMock() 

292 

293 result = wrapper._pad_to_sp_divisible(tensor, padding_length=3, dim=1) 

294 assert result is not None 

295 

296 

297def test_transformer_wrapper_forward() -> None: 

298 """xFuserFlux2Transformer2DWrapper.forward runs end-to-end with sp_world_size=1.""" 

299 wrapper = xFuserFlux2Transformer2DWrapper() 

300 

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 

304 

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 

309 

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 

318 

319 

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() 

323 

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 

332 

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 

343 

344 

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() 

348 

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 

353 

354 hidden_states = MagicMock() 

355 hidden_states.shape = [1, 4, 64] 

356 

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 

365 

366 

367def test_transformer_wrapper_forward_return_dict() -> None: 

368 """forward() returns output.__class__(sample, ...) when output is not a tuple.""" 

369 wrapper = xFuserFlux2Transformer2DWrapper() 

370 

371 hidden_states = MagicMock() 

372 hidden_states.shape = [1, 4, 64] 

373 

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