diff --git a/python/sglang/multimodal_gen/runtime/loader/weight_utils.py b/python/sglang/multimodal_gen/runtime/loader/weight_utils.py index d3f2ef47e..ce0e6ea5e 100644 --- a/python/sglang/multimodal_gen/runtime/loader/weight_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/weight_utils.py @@ -32,6 +32,43 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) + +def _disable_runai_streamer_rank_discovery_collective() -> None: + """RunAI Model Streamer's ``find_local_ranks()`` fires a full-world + collective on the first ``stream_files()`` of every streamer instance even + when the caller passes ``is_distributed=False`` — it only populates an env + var for the library's distributed-streaming path, which this loader never + uses (each rank loads its own full copy). Ranks reach it with divergent + timing, so it can fire out of lockstep and hang + (https://github.com/run-ai/runai-model-streamer/issues/84). + + Patch it to the single-process early return it already has; the only + behavior lost is the collective this loader never wanted. + """ + try: + from runai_model_streamer.distributed_streamer.distributed_streamer import ( + _distributedStreamerParams, + ) + except ImportError: + return + if not hasattr(_distributedStreamerParams, "find_local_ranks"): + logger.warning( + "runai_model_streamer find_local_ranks not found; skipping the " + "rank-discovery-collective workaround (multi-rank loads may hang, " + "see run-ai/runai-model-streamer#84)." + ) + return + + def _find_local_ranks_no_collective(self): + rank = torch.distributed.get_rank() if torch.distributed.is_initialized() else 0 + return 1, rank, [[rank]] + + _distributedStreamerParams.find_local_ranks = _find_local_ranks_no_collective + + +if HAS_RUNAI_MODEL_STREAMER: + _disable_runai_streamer_rank_discovery_collective() + # use system-level temp directory for file locks, so that multiple users # can share the same lock without error. # lock files in the temp directory will be automatically deleted when the diff --git a/python/sglang/multimodal_gen/test/unit/test_weight_utils.py b/python/sglang/multimodal_gen/test/unit/test_weight_utils.py new file mode 100644 index 000000000..066166046 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_weight_utils.py @@ -0,0 +1,87 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The patched find_local_ranks() must never touch torch.distributed.""" + +import unittest +from unittest.mock import patch + +from sglang.multimodal_gen.runtime.loader.weight_utils import ( + _disable_runai_streamer_rank_discovery_collective, +) + +_DIST_STREAMER_MOD = "runai_model_streamer.distributed_streamer.distributed_streamer" + + +class TestDisableRunaiStreamerRankDiscoveryCollective(unittest.TestCase): + def test_never_touches_torch_distributed_even_when_initialized(self): + from runai_model_streamer.distributed_streamer.distributed_streamer import ( + _distributedStreamerParams, + ) + + _disable_runai_streamer_rank_discovery_collective() + + with ( + patch(f"{_DIST_STREAMER_MOD}.dist.is_initialized", return_value=True), + patch(f"{_DIST_STREAMER_MOD}.dist.get_world_size", return_value=2), + patch(f"{_DIST_STREAMER_MOD}.dist.get_rank", return_value=1), + patch(f"{_DIST_STREAMER_MOD}.dist.new_group") as mock_new_group, + patch(f"{_DIST_STREAMER_MOD}.dist.all_gather_object") as mock_all_gather, + patch(f"{_DIST_STREAMER_MOD}.dist.destroy_process_group") as mock_destroy, + ): + result = _distributedStreamerParams().find_local_ranks() + + mock_new_group.assert_not_called() + mock_all_gather.assert_not_called() + mock_destroy.assert_not_called() + # rank is still reported correctly -- only the collective is gone + self.assertEqual(result, (1, 1, [[1]])) + + def test_reports_rank_zero_when_not_distributed(self): + from runai_model_streamer.distributed_streamer.distributed_streamer import ( + _distributedStreamerParams, + ) + + _disable_runai_streamer_rank_discovery_collective() + + with patch(f"{_DIST_STREAMER_MOD}.dist.is_initialized", return_value=False): + result = _distributedStreamerParams().find_local_ranks() + + self.assertEqual(result, (1, 0, [[0]])) + + def test_idempotent_across_repeated_calls(self): + # Import-time application plus any re-import/re-entry must not stack + # wrappers or otherwise change behavior. + _disable_runai_streamer_rank_discovery_collective() + _disable_runai_streamer_rank_discovery_collective() + + from runai_model_streamer.distributed_streamer.distributed_streamer import ( + _distributedStreamerParams, + ) + + with ( + patch(f"{_DIST_STREAMER_MOD}.dist.is_initialized", return_value=True), + patch(f"{_DIST_STREAMER_MOD}.dist.get_world_size", return_value=4), + patch(f"{_DIST_STREAMER_MOD}.dist.get_rank", return_value=3), + patch(f"{_DIST_STREAMER_MOD}.dist.new_group") as mock_new_group, + ): + result = _distributedStreamerParams().find_local_ranks() + + mock_new_group.assert_not_called() + self.assertEqual(result, (1, 3, [[3]])) + + def test_noop_when_library_missing_attribute(self): + # Defensive path: if a future runai_model_streamer release renames or + # removes find_local_ranks, patching must skip (with a warning) rather + # than crash import. + import sglang.multimodal_gen.runtime.loader.weight_utils as wu + + class _StubParams: + pass + + with patch(f"{_DIST_STREAMER_MOD}._distributedStreamerParams", _StubParams): + wu._disable_runai_streamer_rank_discovery_collective() # must not raise + + self.assertFalse(hasattr(_StubParams, "find_local_ranks")) + + +if __name__ == "__main__": + unittest.main()