Skip to content

Commit b6f31d6

Browse files
authored
fix(training): fail fast on malformed launcher rank (#5766)
Signed-off-by: yaoyu-33 <yaoyu.094@gmail.com>
1 parent dd150c1 commit b6f31d6

2 files changed

Lines changed: 14 additions & 12 deletions

File tree

src/megatron/bridge/utils/common_utils.py

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
import torch
2323
import torch.distributed
2424
from megatron.core import DistributedDataParallel as DDP
25-
from megatron.core._rank_utils import safe_get_rank as _get_rank_safe
25+
from megatron.core._rank_utils import safe_get_rank as get_rank_safe # noqa: F401
2626
from megatron.core._rank_utils import safe_get_world_size as get_world_size_safe # noqa: F401
2727
from megatron.core.transformer.module import Float16Module
2828
from megatron.core.utils import get_batch_on_this_cp_rank
@@ -43,14 +43,6 @@
4343
ALL_MODULE_WRAPPER_CLASSNAMES = (DDP, Float16Module)
4444

4545

46-
def get_rank_safe() -> int:
47-
"""Get the distributed rank, falling back to zero for malformed launcher state."""
48-
try:
49-
return _get_rank_safe()
50-
except (TypeError, ValueError):
51-
return 0
52-
53-
5446
def get_last_rank() -> int:
5547
"""Get the last rank in the distributed group"""
5648
if not torch.distributed.is_initialized():

tests/unit_tests/utils/test_common_utils.py

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -73,12 +73,22 @@ def test_uninitialized_torch_distributed_no_env_var(self, mock_is_initialized):
7373
mock_is_initialized.assert_called_once()
7474

7575
@patch("torch.distributed.is_initialized")
76-
@patch.dict(os.environ, {"RANK": "invalid"})
76+
@patch.dict(os.environ, {"RANK": "invalid"}, clear=True)
7777
def test_invalid_rank_env_var(self, mock_is_initialized):
78-
"""Test get_rank_safe falls back to 0 when RANK is malformed."""
78+
"""Test get_rank_safe propagates an actionable error when RANK is malformed."""
7979
mock_is_initialized.return_value = False
8080

81-
assert get_rank_safe() == 0
81+
with pytest.raises(ValueError, match="invalid literal for int.*invalid"):
82+
get_rank_safe()
83+
84+
@patch("torch.distributed.is_initialized")
85+
@patch.dict(os.environ, {"SLURM_NTASKS": "8", "SLURM_PROCID": "invalid"}, clear=True)
86+
def test_invalid_slurm_rank_env_var(self, mock_is_initialized):
87+
"""Test get_rank_safe propagates an actionable error when SLURM_PROCID is malformed."""
88+
mock_is_initialized.return_value = False
89+
90+
with pytest.raises(ValueError, match="invalid literal for int.*invalid"):
91+
get_rank_safe()
8292

8393

8494
class TestGetWorldSizeSafe:

0 commit comments

Comments
 (0)