[Core] Add native ModelExpress load format (#43105)

Signed-off-by: Zheng Luo <[email protected]>
Co-authored-by: OpenAI Codex <[email protected]>
Co-authored-by: Robert Shaw <[email protected]>
This commit is contained in:
Zheng Luo
2026-05-21 16:05:01 -04:00
committed by GitHub
co-authored by OpenAI Codex Robert Shaw
parent b29cbf0652
commit 17b69828a0
5 changed files with 212 additions and 1 deletions
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import sys
from types import ModuleType, SimpleNamespace
import pytest
from torch import nn
from vllm.config import VllmConfig
from vllm.config.load import LoadConfig
from vllm.model_executor.model_loader import get_model_loader
from vllm.model_executor.model_loader.modelexpress_loader import (
ModelExpressModelLoader,
)
class FakeModelexpressLoader:
calls: list[tuple[str, tuple, dict]] = []
loaded_model: nn.Module
def __init__(self, load_config: LoadConfig):
self.load_config = load_config
def download_model(self, *args, **kwargs):
self.calls.append(("download_model", args, kwargs))
def load_weights(self, *args, **kwargs):
self.calls.append(("load_weights", args, kwargs))
def load_model(self, *args, **kwargs):
self.calls.append(("load_model", args, kwargs))
return self.loaded_model
def _install_fake_modelexpress(monkeypatch):
FakeModelexpressLoader.calls = []
FakeModelexpressLoader.loaded_model = nn.Module()
for name in [
"modelexpress",
"modelexpress.engines",
"modelexpress.engines.vllm",
]:
monkeypatch.setitem(sys.modules, name, ModuleType(name))
module = ModuleType("modelexpress.engines.vllm.loader")
module.__dict__["MxModelLoader"] = FakeModelexpressLoader
monkeypatch.setitem(sys.modules, module.__name__, module)
def test_modelexpress_load_format_resolves_to_modelexpress_loader(monkeypatch):
_install_fake_modelexpress(monkeypatch)
loader = get_model_loader(LoadConfig(load_format="modelexpress"))
assert isinstance(loader, ModelExpressModelLoader)
def test_modelexpress_loader_delegates_to_modelexpress(monkeypatch):
_install_fake_modelexpress(monkeypatch)
loader = ModelExpressModelLoader(LoadConfig(load_format="modelexpress"))
model = nn.Module()
model_config = SimpleNamespace()
vllm_config = SimpleNamespace()
loader.download_model(model_config)
loader.load_weights(model, model_config)
FakeModelexpressLoader.loaded_model.train()
result = loader.load_model(
vllm_config=vllm_config,
model_config=model_config,
prefix="model",
)
assert result is FakeModelexpressLoader.loaded_model
assert not result.training
assert FakeModelexpressLoader.calls == [
("download_model", (model_config,), {}),
("load_weights", (model, model_config), {}),
(
"load_model",
(),
{
"vllm_config": vllm_config,
"model_config": model_config,
"prefix": "model",
},
),
]
def test_modelexpress_loader_missing_modelexpress_error(monkeypatch):
import importlib
def missing_modelexpress(name):
raise ModuleNotFoundError(name=name)
monkeypatch.setattr(importlib, "import_module", missing_modelexpress)
with pytest.raises(ImportError, match="requires the ModelExpress Python package"):
ModelExpressModelLoader(LoadConfig(load_format="modelexpress"))
def test_modelexpress_loader_preserves_internal_import_errors(monkeypatch):
import importlib
def missing_dependency(name):
raise ModuleNotFoundError(name="not_modelexpress_dependency")
monkeypatch.setattr(importlib, "import_module", missing_dependency)
with pytest.raises(ModuleNotFoundError) as exc_info:
ModelExpressModelLoader(LoadConfig(load_format="modelexpress"))
assert exc_info.value.name == "not_modelexpress_dependency"
def test_modelexpress_load_format_allows_object_storage_model_weights():
model_config = SimpleNamespace(
architecture="UnknownForTest",
config_updated=False,
convert_type=None,
is_hybrid=False,
model="test-model",
model_weights="s3://bucket/model",
)
vllm_config = object.__new__(VllmConfig)
vllm_config.model_config = model_config
vllm_config.load_config = LoadConfig(load_format="modelexpress")
vllm_config.try_verify_and_update_config()
assert vllm_config.load_config.load_format == "modelexpress"
+1
View File
@@ -55,6 +55,7 @@ class LoadConfig:
https://github.com/ggml-org/ggml/blob/master/docs/gguf.md).
- "mistral" will load weights from consolidated safetensors files used by
Mistral models.
- "modelexpress" will load weights using ModelExpress.
- Other custom values can be supported via plugins.
"""
download_dir: str | None = None
+2 -1
View File
@@ -1901,12 +1901,13 @@ class VllmConfig:
)
self.load_config.load_format = "runai_streamer"
elif self.load_config.load_format not in (
"modelexpress",
"runai_streamer",
"runai_streamer_sharded",
):
raise ValueError(
f"To load a model from object storage (S3/GCS/Azure), "
f"'load_format' must be 'runai_streamer' or "
f"'load_format' must be 'modelexpress', 'runai_streamer' or "
f"'runai_streamer_sharded', "
f"but got '{self.load_config.load_format}'. "
f"Model: {self.model_config.model}"
@@ -13,6 +13,9 @@ from vllm.model_executor.model_loader.bitsandbytes_loader import BitsAndBytesMod
from vllm.model_executor.model_loader.default_loader import DefaultModelLoader
from vllm.model_executor.model_loader.dummy_loader import DummyModelLoader
from vllm.model_executor.model_loader.gguf_loader import GGUFModelLoader
from vllm.model_executor.model_loader.modelexpress_loader import (
ModelExpressModelLoader,
)
from vllm.model_executor.model_loader.runai_streamer_loader import (
RunaiModelStreamerLoader,
)
@@ -37,6 +40,7 @@ LoadFormats = Literal[
"gguf",
"instanttensor",
"mistral",
"modelexpress",
"npcache",
"pt",
"runai_streamer",
@@ -54,6 +58,7 @@ _LOAD_FORMAT_TO_MODEL_LOADER: dict[str, type[BaseModelLoader]] = {
"gguf": GGUFModelLoader,
"instanttensor": DefaultModelLoader,
"mistral": DefaultModelLoader,
"modelexpress": ModelExpressModelLoader,
"npcache": DefaultModelLoader,
"pt": DefaultModelLoader,
"runai_streamer": RunaiModelStreamerLoader,
@@ -150,6 +155,7 @@ __all__ = [
"BaseModelLoader",
"BitsAndBytesModelLoader",
"GGUFModelLoader",
"ModelExpressModelLoader",
"DefaultModelLoader",
"DummyModelLoader",
"RunaiModelStreamerLoader",
@@ -0,0 +1,70 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from __future__ import annotations
import importlib
from torch import nn
from vllm.config import ModelConfig, VllmConfig
from vllm.config.load import LoadConfig
from vllm.model_executor.model_loader.base_loader import BaseModelLoader
from vllm.tracing import instrument
_MODELEXPRESS_LOADER_MODULE = "modelexpress.engines.vllm.loader"
_MISSING_MODELEXPRESS_MODULES = frozenset(
{
"modelexpress",
"modelexpress.engines",
"modelexpress.engines.vllm",
_MODELEXPRESS_LOADER_MODULE,
}
)
def _missing_modelexpress_error() -> ImportError:
return ImportError(
"The 'modelexpress' load format requires the ModelExpress Python package. "
"Install it with `pip install modelexpress`."
)
class ModelExpressModelLoader(BaseModelLoader):
"""Thin vLLM loader wrapper for ModelExpress."""
def __init__(self, load_config: LoadConfig):
super().__init__(load_config)
self._loader = self._load_modelexpress_loader(load_config)
@staticmethod
def _load_modelexpress_loader(load_config: LoadConfig) -> BaseModelLoader:
try:
module = importlib.import_module(_MODELEXPRESS_LOADER_MODULE)
except ModuleNotFoundError as exc:
if exc.name not in _MISSING_MODELEXPRESS_MODULES:
raise
raise _missing_modelexpress_error() from exc
ModelExpressVllmLoader = module.MxModelLoader
return ModelExpressVllmLoader(load_config)
def download_model(self, model_config: ModelConfig) -> None:
self._loader.download_model(model_config)
def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None:
self._loader.load_weights(model, model_config)
@instrument(span_name="Load model")
def load_model(
self,
vllm_config: VllmConfig,
model_config: ModelConfig,
prefix: str = "",
) -> nn.Module:
model = self._loader.load_model(
vllm_config=vllm_config,
model_config=model_config,
prefix=prefix,
)
return model.eval()