import sys from unittest.mock import MagicMock # Define dummy objects to satisfy imports class MockModule(MagicMock): def __getattr__(self, name): # Specific overrides for expected function returns if name == "get_sequence_parallel_world_size": return lambda: 1 if name == "get_sequence_parallel_rank": return lambda: 0 if name == "get_sp_group": return MagicMock return MagicMock() # Register modules in sys.modules sys.modules['xfuser'] = MockModule() sys.modules['xfuser.core'] = MockModule() sys.modules['xfuser.core.distributed'] = MockModule() sys.modules['yunchang'] = MockModule() # Functions inside xfuser.core.distributed def get_sp_group(): mock_group = MagicMock() mock_group.ranks = {0: 0} # Mock broadcast def broadcast(tensor, src=0): return tensor mock_group.broadcast = broadcast # Mock all_gather def all_gather(x, dim=1): return x mock_group.all_gather = all_gather return mock_group def initialize_model_parallel(*args, **kwargs): pass def init_distributed_environment(*args, **kwargs): pass def get_sequence_parallel_world_size(): return 1 def get_sequence_parallel_rank(): return 0 # Bind functions so they can be imported directly sys.modules['xfuser.core.distributed'].get_sp_group = get_sp_group sys.modules['xfuser.core.distributed'].initialize_model_parallel = initialize_model_parallel sys.modules['xfuser.core.distributed'].init_distributed_environment = init_distributed_environment sys.modules['xfuser.core.distributed'].get_sequence_parallel_world_size = get_sequence_parallel_world_size sys.modules['xfuser.core.distributed'].get_sequence_parallel_rank = get_sequence_parallel_rank