mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-24 22:50:15 +00:00
[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:
co-authored by
OpenAI Codex
Robert Shaw
parent
b29cbf0652
commit
17b69828a0
@@ -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"
|
||||
@@ -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
@@ -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()
|
||||
Reference in New Issue
Block a user