Skip to content
Merged
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
6 changes: 5 additions & 1 deletion .github/workflows/release.yml
Original file line number Diff line number Diff line change
Expand Up @@ -105,7 +105,11 @@ jobs:
[
str(runtime_python),
"-c",
"import art_vllm_runtime, torch, vllm; print('runtime imports ok')",
"import torch, vllm; "
"from art_vllm_runtime.policy_spans import "
"_patch_lora_update_coordinator; "
"_patch_lora_update_coordinator(); "
"print('runtime compatibility ok')",
],
check=True,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -248,6 +248,60 @@ def test_runtime_general_plugin_loads_full_patch_set() -> None:
assert 'art = "art_vllm_runtime.patches:apply_vllm_runtime_patches"' in pyproject


def test_lora_coordinator_supports_both_vllm_serving_layouts(
artifact_dir: Path,
) -> None:
payload = _runtime_python(
"""
import json
from types import SimpleNamespace
import art_vllm_runtime.policy_spans as policy
from vllm.entrypoints.openai.engine.serving import OpenAIServing

policy._patch_lora_update_coordinator()
legacy_patched = getattr(
OpenAIServing.__init__, "__art_lora_update_patched__", False
)

class GenerateBaseServing:
def __init__(self, models, engine_client):
self.models = models
self.engine_client = engine_client

real_import_module = policy.importlib.import_module
def import_module(name):
if name == "vllm.entrypoints.openai.engine.serving":
raise ModuleNotFoundError(name, name=name)
if name == "vllm.entrypoints.generate.base.serving":
return SimpleNamespace(GenerateBaseServing=GenerateBaseServing)
return real_import_module(name)

policy.importlib.import_module = import_module
policy._patch_lora_update_coordinator()
models = SimpleNamespace()
engine_client = SimpleNamespace()
GenerateBaseServing(models, engine_client)
print(json.dumps({
"legacy_patched": legacy_patched,
"new_patched": getattr(
GenerateBaseServing.__init__, "__art_lora_update_patched__", False
),
"shared_coordinator": (
models._art_lora_update_coordinator
is engine_client._art_lora_update_coordinator
),
}))
""",
artifact_dir,
"vllm_serving_layouts",
)
assert json.loads(payload.splitlines()[-1]) == {
"legacy_patched": True,
"new_patched": True,
"shared_coordinator": True,
}


def test_runtime_patch_adds_gemma4_moe_topk_alias(artifact_dir: Path) -> None:
payload = _runtime_python(
"import json; "
Expand Down
15 changes: 11 additions & 4 deletions vllm_runtime/src/art_vllm_runtime/policy_spans.py
Original file line number Diff line number Diff line change
Expand Up @@ -526,17 +526,24 @@ async def tracked_result_generator():


def _patch_lora_update_coordinator() -> None:
from vllm.entrypoints.openai.engine.serving import OpenAIServing

original_init = OpenAIServing.__init__
try:
module = importlib.import_module("vllm.entrypoints.openai.engine.serving")
serving_base = module.OpenAIServing
except ModuleNotFoundError as exc:
if exc.name != "vllm.entrypoints.openai.engine.serving":
raise
module = importlib.import_module("vllm.entrypoints.generate.base.serving")
serving_base = module.GenerateBaseServing

original_init = serving_base.__init__
if not getattr(original_init, "__art_lora_update_patched__", False):

def __init__(self: Any, *args: Any, **kwargs: Any) -> None:
original_init(self, *args, **kwargs)
lora_update_coordinator(self.models, self.engine_client)

__init__.__art_lora_update_patched__ = True # type: ignore[attr-defined]
OpenAIServing.__init__ = __init__ # type: ignore[method-assign]
serving_base.__init__ = __init__


def _patch_engine_request_admission() -> None:
Expand Down
Loading