Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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'")
26 changes: 26 additions & 0 deletions modelexpress_client/python/tests/test_gds_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import json
import struct
from types import SimpleNamespace
from unittest.mock import MagicMock, mock_open, patch

import pytest
Expand Down Expand Up @@ -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):
Expand Down
Loading