Coverage for tests/test_flux_parallelize.py: 100%

41 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1#!/usr/bin/env python3 

2 

3import sys 

4 

5from typing import Tuple 

6 

7from unittest.mock import MagicMock 

8from unittest.mock import patch 

9 

10from tests.torch_mock import TorchMock 

11from tests.diffusers_mock import DiffusersMock 

12 

13mock_torch = TorchMock() 

14mock_diffusers = DiffusersMock() 

15 

16sys.path.append("wrapper") 

17sys.path.append("wrapper/flux") 

18 

19mock_modules = { 

20 'xfuser': MagicMock(), 

21 'xfuser.config': MagicMock(), 

22 'xfuser.core': MagicMock(), 

23 'xfuser.core.distributed': MagicMock(), 

24 'xfuser.model_executor': MagicMock(), 

25 'xfuser.model_executor.models': MagicMock(), 

26 'xfuser.model_executor.models.transformers.transformer_flux': MagicMock(), 

27 'xfuser.model_executor.layers': MagicMock(), 

28 'xfuser.model_executor.layers.attention_processor': MagicMock(), 

29} 

30mock_modules.update(mock_torch.get_sub_modules()) 

31mock_modules.update(mock_diffusers.get_sub_modules()) 

32 

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

34 from flux.flux_xfuser import parallelize_transformer 

35 import flux.flux_xfuser as _flux_xfuser_mod 

36 

37# Re-register the module so patch() targets the same object as 

38# parallelize_transformer.__globals__ (patch.dict restores sys.modules 

39# on exit, removing the module entry that was added during import). 

40sys.modules['flux.flux_xfuser'] = _flux_xfuser_mod 

41 

42 

43class DummyTransformer: 

44 def __init__(self) -> None: 

45 self.transformer_blocks: list[MagicMock] = [] 

46 self.single_transformer_blocks: list[MagicMock] = [] 

47 

48 def forward(self, *args: object, **kwargs: object) -> Tuple: 

49 return (mock_torch.randn(2, 4, 8), "extra") 

50 

51 

52def test_parallelize_transformer() -> None: 

53 transformer = DummyTransformer() 

54 pipeline = MagicMock() # DiffusionPipeline 

55 pipeline.transformer = transformer 

56 

57 # Patch xfuser distributed utils so they return trivial values 

58 with ( 

59 patch("flux.flux_xfuser.get_classifier_free_guidance_world_size", return_value=1), 

60 patch("flux.flux_xfuser.get_classifier_free_guidance_rank", return_value=0), 

61 patch("flux.flux_xfuser.get_sequence_parallel_world_size", return_value=1), 

62 patch("flux.flux_xfuser.get_sequence_parallel_rank", return_value=0), 

63 patch("flux.flux_xfuser.get_runtime_state") as mock_runtime, 

64 patch("flux.flux_xfuser.get_sp_group") as mock_sp_group, 

65 patch("flux.flux_xfuser.get_cfg_group") as mock_cfg_group 

66 ): 

67 

68 runtime_state = MagicMock() 

69 runtime_state.split_text_embed_in_sp = True 

70 mock_runtime.return_value = runtime_state 

71 

72 mock_sp_group.return_value.all_gather = lambda x, dim: x 

73 mock_cfg_group.return_value.all_gather = lambda x, dim: x 

74 

75 hidden = mock_torch.randn(2, 4, 8) 

76 enc = mock_torch.randn(2, 4, 8) 

77 img_ids = mock_torch.ones(2, 4, 1) 

78 txt_ids = mock_torch.ones(2, 4, 1) 

79 timestep = mock_torch.tensor(1) 

80 

81 parallel_pipeline = parallelize_transformer(pipeline) 

82 

83 result = parallel_pipeline.transformer.forward( 

84 hidden, 

85 enc, 

86 img_ids=img_ids, 

87 txt_ids=txt_ids, 

88 timestep=timestep) 

89 

90 assert isinstance(result, tuple)