Spaces:
Paused
Paused
| 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 | |