Skip to content

Commit b0b51b2

Browse files
committed
fix(vllm): resolve native generate API at module load
Signed-off-by: Biswa Panda <biswa.panda@gmail.com>
1 parent 6c688b1 commit b0b51b2

1 file changed

Lines changed: 13 additions & 24 deletions

File tree

components/src/dynamo/vllm/engine_generate.py

Lines changed: 13 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55

66
import base64
77
import io
8-
from functools import lru_cache
98
from typing import Any
109

1110
import msgspec
@@ -19,6 +18,18 @@
1918
)
2019
from vllm.sampling_params import RequestOutputKind, SamplingParams
2120

21+
try:
22+
from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import (
23+
decode_mm_kwargs_item,
24+
)
25+
from vllm.entrypoints.scale_out.token_in_token_out.protocol import GenerateRequest
26+
except ModuleNotFoundError as exc:
27+
expected_module = "vllm.entrypoints.scale_out.token_in_token_out"
28+
if exc.name is None or not expected_module.startswith(exc.name):
29+
raise
30+
from vllm.entrypoints.serve.disagg.mm_serde import decode_mm_kwargs_item
31+
from vllm.entrypoints.serve.disagg.protocol import GenerateRequest
32+
2233
from .constants import DYNAMO_CACHE_SALT_PREFIX
2334

2435
GENERATE_CAPABILITY = "vllm_inference_v1_generate"
@@ -33,26 +44,6 @@ def serialize_routed_experts(routed_experts: Any) -> str | None:
3344
return base64.b64encode(buffer.getvalue()).decode("ascii")
3445

3546

36-
@lru_cache(maxsize=1)
37-
def _native_generate_api() -> tuple[Any, Any]:
38-
"""Load the vLLM-native endpoint adapter only when `/generate` is used."""
39-
try:
40-
from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import (
41-
decode_mm_kwargs_item,
42-
)
43-
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
44-
GenerateRequest,
45-
)
46-
except ModuleNotFoundError as exc:
47-
expected_module = "vllm.entrypoints.scale_out.token_in_token_out"
48-
if exc.name is None or not expected_module.startswith(exc.name):
49-
raise
50-
from vllm.entrypoints.serve.disagg.mm_serde import decode_mm_kwargs_item
51-
from vllm.entrypoints.serve.disagg.protocol import GenerateRequest
52-
53-
return decode_mm_kwargs_item, GenerateRequest
54-
55-
5647
def payload(request: dict[str, Any]) -> dict[str, Any] | None:
5748
extra_args = request.get("extra_args")
5849
if not isinstance(extra_args, dict):
@@ -145,7 +136,6 @@ def build_prompt(request: dict[str, Any]) -> Any:
145136

146137
mm_kwargs: dict[str, list[MultiModalKwargsItem | None]] = {}
147138
if isinstance(kwargs_data, dict):
148-
decode_mm_kwargs_item, _ = _native_generate_api()
149139
for modality, items in kwargs_data.items():
150140
mm_kwargs[modality] = [
151141
decode_mm_kwargs_item(item) if item is not None else None
@@ -173,8 +163,7 @@ def build_sampling_params(
173163
if generate_request is None:
174164
raise ValueError("extra_args.vllm_tito is missing from token-native request")
175165

176-
_, generate_request_type = _native_generate_api()
177-
parsed = generate_request_type.model_validate(generate_request)
166+
parsed = GenerateRequest.model_validate(generate_request)
178167
sampling_params = parsed.sampling_params
179168
if isinstance(sampling_params, dict):
180169
sampling_params = msgspec.convert(sampling_params, type=SamplingParams)

0 commit comments

Comments
 (0)