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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
3import sys
5from typing import Tuple
7from unittest.mock import MagicMock
8from unittest.mock import patch
10from tests.torch_mock import TorchMock
11from tests.diffusers_mock import DiffusersMock
13mock_torch = TorchMock()
14mock_diffusers = DiffusersMock()
16sys.path.append("wrapper")
17sys.path.append("wrapper/flux")
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())
33with patch.dict(sys.modules, mock_modules):
34 from flux.flux_xfuser import parallelize_transformer
35 import flux.flux_xfuser as _flux_xfuser_mod
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
43class DummyTransformer:
44 def __init__(self) -> None:
45 self.transformer_blocks: list[MagicMock] = []
46 self.single_transformer_blocks: list[MagicMock] = []
48 def forward(self, *args: object, **kwargs: object) -> Tuple:
49 return (mock_torch.randn(2, 4, 8), "extra")
52def test_parallelize_transformer() -> None:
53 transformer = DummyTransformer()
54 pipeline = MagicMock() # DiffusionPipeline
55 pipeline.transformer = transformer
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 ):
68 runtime_state = MagicMock()
69 runtime_state.split_text_embed_in_sp = True
70 mock_runtime.return_value = runtime_state
72 mock_sp_group.return_value.all_gather = lambda x, dim: x
73 mock_cfg_group.return_value.all_gather = lambda x, dim: x
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)
81 parallel_pipeline = parallelize_transformer(pipeline)
83 result = parallel_pipeline.transformer.forward(
84 hidden,
85 enc,
86 img_ids=img_ids,
87 txt_ids=txt_ids,
88 timestep=timestep)
90 assert isinstance(result, tuple)