From 577c913e3da43749159973ef6129f4c2fdac236c Mon Sep 17 00:00:00 2001 From: Rohan Patnaik Date: Sun, 27 Sep 2026 21:08:06 +0530 Subject: [PATCH] Use launcher world size before distributed init Signed-off-by: Rohan Patnaik --- deepspeed/runtime/config.py | 6 +++-- tests/unit/runtime/test_ds_config_dict.py | 32 +++++++++++++++++++++++ 2 files changed, 36 insertions(+), 2 deletions(-) diff --git a/deepspeed/runtime/config.py b/deepspeed/runtime/config.py index 390a0be5c559..4d2cc0165cc4 100755 --- a/deepspeed/runtime/config.py +++ b/deepspeed/runtime/config.py @@ -470,8 +470,10 @@ def __init__(self, config: Union[str, dict], mpu=None, mesh_device=None): else: self.world_size = dist.get_world_size() except (RuntimeError, AssertionError, AttributeError): - self.global_rank = 0 - self.world_size = 1 + self.global_rank = int(os.environ.get("RANK", "0")) + self.world_size = int(os.environ.get("WORLD_SIZE", "1")) + if "sequence_parallel_size" in self._param_dict: + self.world_size /= self._param_dict["sequence_parallel_size"] logger.info(f"Config mesh_device {mesh_device} world_size = {self.world_size}") # Pass a copy so that the user json is unmodified, e.g. for logging. param_dict = copy.copy(self._param_dict) diff --git a/tests/unit/runtime/test_ds_config_dict.py b/tests/unit/runtime/test_ds_config_dict.py index 9837907fac23..a4f8d803fade 100644 --- a/tests/unit/runtime/test_ds_config_dict.py +++ b/tests/unit/runtime/test_ds_config_dict.py @@ -149,6 +149,38 @@ def test_gradient_allreduce_op_default(): assert config.gradient_allreduce_op == "mean" +@pytest.mark.parametrize("rank,world_size,sequence_parallel_size,expected_rank,expected_world_size", + [(None, None, None, 0, 1), (1, 2, None, 1, 2), (3, 4, 2, 3, 2)]) +def test_config_uses_launcher_environment_before_distributed_initialization(monkeypatch, rank, world_size, + sequence_parallel_size, expected_rank, + expected_world_size): + if rank is not None: + monkeypatch.setenv("RANK", str(rank)) + else: + monkeypatch.delenv("RANK", raising=False) + if world_size is not None: + monkeypatch.setenv("WORLD_SIZE", str(world_size)) + else: + monkeypatch.delenv("WORLD_SIZE", raising=False) + + def get_rank_before_distributed_initialization(): + raise RuntimeError + + monkeypatch.setattr(dist, "get_rank", get_rank_before_distributed_initialization) + config_dict = { + "train_batch_size": expected_world_size, + "train_micro_batch_size_per_gpu": 1, + "gradient_accumulation_steps": 1, + } + if sequence_parallel_size is not None: + config_dict["sequence_parallel_size"] = sequence_parallel_size + + config = DeepSpeedConfig(config_dict) + + assert config.global_rank == expected_rank + assert config.world_size == expected_world_size + + def test_disable_python_gc_config_default(): config = DeepSpeedConfig({"train_batch_size": 1}) assert config.disable_python_gc is False