diff --git a/modelexpress_client/python/modelexpress/load_strategy/gds_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/gds_strategy.py index b40c73ab9..b09f3bb59 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/gds_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/gds_strategy.py @@ -46,7 +46,9 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: use_tqdm = getattr(ctx.load_config, "use_tqdm_on_load", True) revision = getattr(ctx.model_config, "revision", None) weights_iter = gds_loader.load_iter( - ctx.model_config.model, use_tqdm=use_tqdm, revision=revision + _get_model_path(ctx.model_config), + use_tqdm=use_tqdm, + revision=revision, ) except Exception as e: logger.warning( @@ -68,3 +70,12 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: register_tensors(result, ctx) return result + + +def _get_model_path(model_config) -> str: + """Return the checkpoint path used by vLLM or SGLang model configs.""" + for attribute in ("model", "model_path"): + value = getattr(model_config, attribute, None) + if value: + return str(value) + raise AttributeError("model config has neither 'model' nor 'model_path'") diff --git a/modelexpress_client/python/tests/test_gds_loader.py b/modelexpress_client/python/tests/test_gds_loader.py index c508600b3..e7c3d525e 100644 --- a/modelexpress_client/python/tests/test_gds_loader.py +++ b/modelexpress_client/python/tests/test_gds_loader.py @@ -5,6 +5,7 @@ import json import struct +from types import SimpleNamespace from unittest.mock import MagicMock, mock_open, patch import pytest @@ -261,6 +262,31 @@ def test_gds_success(self, _mock_nixl, _mock_pub, mock_gds_cls, _mock_avail): model.load_weights.assert_called_once() mock_gds.shutdown.assert_called_once() + @patch("modelexpress.gds_transfer.is_gds_available", return_value=True) + @patch("modelexpress.gds_loader.MxGdsLoader") + def test_gds_uses_sglang_model_path(self, mock_gds_cls, _mock_avail): + from modelexpress.load_strategy.gds_strategy import GdsStrategy + + mock_gds = MagicMock() + mock_gds.load_iter.return_value = iter([("w", torch.zeros(1))]) + mock_gds_cls.return_value = mock_gds + + ctx = self._make_context() + ctx.model_config = SimpleNamespace( + model_path="test-sglang-model", + revision="test-revision", + ) + ctx.load_config.use_tqdm_on_load = False + + GdsStrategy().load(MagicMock(), ctx) + + mock_gds.load_iter.assert_called_once_with( + "test-sglang-model", + use_tqdm=False, + revision="test-revision", + ) + mock_gds.shutdown.assert_called_once() + @patch("modelexpress.gds_transfer.is_gds_available", return_value=True) @patch("modelexpress.gds_loader.MxGdsLoader") def test_gds_failure_raises_strategy_failed(self, mock_gds_cls, _mock_avail):