mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-07 06:18:06 +00:00
[Bugfix] Fix LOGITPROC_SOURCE_ENTRYPOINT test to use spawn-compatible dist-info registration for XPU/ROCm (#42040)
Signed-off-by: dqzhengAP <[email protected]> Signed-off-by: David Zheng <[email protected]> Signed-off-by: Andreas Karatzas <[email protected]> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: Andreas Karatzas <[email protected]> Co-authored-by: Kunshang Ji <[email protected]>
This commit is contained in:
co-authored by
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Andreas Karatzas
Kunshang Ji
parent
97cc7685c4
commit
df2636a9d8
+56
-11
@@ -20,9 +20,9 @@ import time
|
||||
import warnings
|
||||
from collections.abc import Callable, Iterable, Sequence
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from multiprocessing import Process
|
||||
from multiprocessing import Process, get_context
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, cast
|
||||
from unittest.mock import patch
|
||||
|
||||
import anthropic
|
||||
@@ -127,6 +127,21 @@ ROCM_ENGINE_KWARGS: dict = (
|
||||
)
|
||||
|
||||
|
||||
def requires_spawn_multiprocessing() -> bool:
|
||||
"""Whether this platform requires spawn instead of fork for test processes."""
|
||||
return current_platform.is_rocm() or current_platform.is_xpu()
|
||||
|
||||
|
||||
def _run_in_new_process_group(
|
||||
child_process_fxn: Callable[[dict[str, str] | None, str, list[str]], None],
|
||||
env_dict: dict[str, str] | None,
|
||||
model: str,
|
||||
vllm_serve_args: list[str],
|
||||
) -> None:
|
||||
os.setsid()
|
||||
child_process_fxn(env_dict, model, vllm_serve_args)
|
||||
|
||||
|
||||
class RemoteVLLMServer:
|
||||
"""Base class for launching vLLM server subprocesses for testing.
|
||||
|
||||
@@ -738,8 +753,11 @@ class RemoteOpenAIServerCustom(RemoteOpenAIServer):
|
||||
def _start_server(
|
||||
self, model: str, vllm_serve_args: list[str], env_dict: dict[str, str] | None
|
||||
) -> None:
|
||||
self.proc: Process = Process(
|
||||
target=self.child_process_fxn, args=(env_dict, model, vllm_serve_args)
|
||||
method = "spawn" if requires_spawn_multiprocessing() else "fork"
|
||||
ctx = get_context(method)
|
||||
self.proc: Process = cast(Any, ctx).Process(
|
||||
target=_run_in_new_process_group,
|
||||
args=(self.child_process_fxn, env_dict, model, vllm_serve_args),
|
||||
) # type: ignore[assignment]
|
||||
self.proc.start()
|
||||
|
||||
@@ -769,12 +787,40 @@ class RemoteOpenAIServerCustom(RemoteOpenAIServer):
|
||||
def _poll(self) -> int | None:
|
||||
return self.proc.exitcode
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
self.proc.terminate()
|
||||
self.proc.join(8)
|
||||
def _terminate_process_tree(self) -> None:
|
||||
pid = self.proc.pid
|
||||
if pid is None:
|
||||
return
|
||||
|
||||
pgid: int | None
|
||||
try:
|
||||
pgid = os.getpgid(pid)
|
||||
# _run_in_new_process_group should make the child the group
|
||||
# leader. Avoid signaling pytest's process group if startup failed
|
||||
# before os.setsid() ran.
|
||||
if pgid != pid:
|
||||
pgid = None
|
||||
except (ProcessLookupError, OSError):
|
||||
pgid = None
|
||||
|
||||
with contextlib.suppress(ProcessLookupError, OSError):
|
||||
self.proc.terminate()
|
||||
print(f"[RemoteOpenAIServerCustom] Sent SIGTERM to process {pid}")
|
||||
|
||||
self.proc.join(15)
|
||||
if self.proc.is_alive():
|
||||
# force kill if needed
|
||||
self.proc.kill()
|
||||
print(
|
||||
f"[RemoteOpenAIServerCustom] Server {pid} did not respond "
|
||||
"to SIGTERM, sending SIGKILL to process group"
|
||||
)
|
||||
if pgid is not None:
|
||||
with contextlib.suppress(ProcessLookupError, OSError):
|
||||
os.killpg(pgid, signal.SIGKILL)
|
||||
else:
|
||||
self.proc.kill()
|
||||
self.proc.join(10)
|
||||
|
||||
self._kill_process_group_survivors(pgid)
|
||||
|
||||
|
||||
def _test_completion(
|
||||
@@ -1633,8 +1679,7 @@ def create_new_process_for_each_test(
|
||||
A decorator to run test functions in separate processes.
|
||||
"""
|
||||
if method is None:
|
||||
use_spawn = current_platform.is_rocm() or current_platform.is_xpu()
|
||||
method = "spawn" if use_spawn else "fork"
|
||||
method = "spawn" if requires_spawn_multiprocessing() else "fork"
|
||||
|
||||
assert method in ["spawn", "fork"], "Method must be either 'spawn' or 'fork'"
|
||||
|
||||
|
||||
@@ -16,8 +16,8 @@ from tests.v1.logits_processors.utils import (
|
||||
DummyLogitsProcessor,
|
||||
WrappedPerReqLogitsProcessor,
|
||||
prompts,
|
||||
setup_fake_entrypoint,
|
||||
)
|
||||
from tests.v1.logits_processors.utils import entry_points as fake_entry_points
|
||||
from vllm import LLM, SamplingParams
|
||||
from vllm.v1.sample.logits_processor import (
|
||||
STR_POOLING_REJECTS_LOGITSPROCS,
|
||||
@@ -145,13 +145,9 @@ def test_custom_logitsprocs(monkeypatch, logitproc_source: CustomLogitprocSource
|
||||
|
||||
if logitproc_source == CustomLogitprocSource.LOGITPROC_SOURCE_ENTRYPOINT:
|
||||
# Scenario: vLLM loads a logitproc from a preconfigured entrypoint
|
||||
# To that end, mock a dummy logitproc entrypoint
|
||||
import importlib.metadata
|
||||
|
||||
importlib.metadata.entry_points = fake_entry_points # type: ignore
|
||||
|
||||
# fork is required for workers to see entrypoint patch
|
||||
monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork")
|
||||
# To that end, register a real dist-info package so spawned
|
||||
# workers can discover the entrypoint via PYTHONPATH
|
||||
setup_fake_entrypoint(monkeypatch)
|
||||
_run_test({}, logitproc_loaded=True)
|
||||
return
|
||||
|
||||
@@ -266,14 +262,9 @@ def test_rejects_custom_logitsprocs(
|
||||
# Scenario: vLLM loads a model and ignores a logitproc that is
|
||||
# available at a preconfigured entrypoint
|
||||
|
||||
# Patch in dummy logitproc entrypoint
|
||||
import importlib.metadata
|
||||
|
||||
importlib.metadata.entry_points = fake_entry_points # type: ignore
|
||||
|
||||
# fork is required for entrypoint patch to be visible to workers,
|
||||
# although they should ignore the entrypoint patch anyway
|
||||
monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork")
|
||||
# Register real dist-info package so spawned workers can
|
||||
# discover the entrypoint via PYTHONPATH (spawn-compatible)
|
||||
setup_fake_entrypoint(monkeypatch)
|
||||
|
||||
llm = LLM(**llm_kwargs)
|
||||
# Require that no custom logitsprocs have been loaded
|
||||
|
||||
@@ -18,8 +18,8 @@ from tests.v1.logits_processors.utils import (
|
||||
MODEL_NAME,
|
||||
TEMP_GREEDY,
|
||||
prompts,
|
||||
setup_fake_entrypoint,
|
||||
)
|
||||
from tests.v1.logits_processors.utils import entry_points as fake_entry_points
|
||||
|
||||
|
||||
def _server_with_logitproc_entrypoint(
|
||||
@@ -27,16 +27,9 @@ def _server_with_logitproc_entrypoint(
|
||||
model: str,
|
||||
vllm_serve_args: list[str],
|
||||
) -> None:
|
||||
"""Start vLLM server, inject dummy logitproc entrypoint"""
|
||||
|
||||
# Patch `entry_points` to inject logitproc entrypoint
|
||||
import importlib.metadata
|
||||
|
||||
importlib.metadata.entry_points = fake_entry_points # type: ignore
|
||||
"""Start vLLM server with dummy logitproc entrypoint."""
|
||||
from vllm.entrypoints.cli import main
|
||||
|
||||
# fork is required for workers to see entrypoint patch
|
||||
os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "fork"
|
||||
if env_dict is not None:
|
||||
os.environ.update(env_dict)
|
||||
|
||||
@@ -50,7 +43,7 @@ def _server_with_logitproc_fqcn(
|
||||
model: str,
|
||||
vllm_serve_args: list[str],
|
||||
) -> None:
|
||||
"""Start vLLM server, inject module with dummy logitproc"""
|
||||
"""Start vLLM server with dummy logitproc specified by FQCN."""
|
||||
from vllm.entrypoints.cli import main
|
||||
|
||||
if env_dict is not None:
|
||||
@@ -80,8 +73,8 @@ def default_server_args():
|
||||
def server(default_server_args, request, monkeypatch):
|
||||
"""Consider two server configurations:
|
||||
(1) --logits-processors cli arg specifies dummy logits processor via fully-
|
||||
qualified class name (FQCN); patch in a dummy logits processor module
|
||||
(2) No --logits-processors cli arg; patch in a dummy logits processor
|
||||
qualified class name (FQCN)
|
||||
(2) No --logits-processors cli arg; inject a dummy logits processor
|
||||
entrypoint
|
||||
"""
|
||||
|
||||
@@ -94,6 +87,7 @@ def server(default_server_args, request, monkeypatch):
|
||||
_server_fxn = _server_with_logitproc_fqcn
|
||||
else:
|
||||
# Launch server, inject dummy logitproc entrypoint
|
||||
setup_fake_entrypoint(monkeypatch)
|
||||
args = default_server_args
|
||||
_server_fxn = _server_with_logitproc_entrypoint
|
||||
|
||||
@@ -119,7 +113,6 @@ api_keyword_args = {
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[MODEL_NAME],
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import types
|
||||
import os
|
||||
import tempfile
|
||||
from enum import Enum, auto
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from tests.utils import requires_spawn_multiprocessing
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.logger import init_logger
|
||||
from vllm.sampling_params import SamplingParams
|
||||
@@ -102,11 +105,6 @@ class DummyLogitsProcessor(LogitsProcessor):
|
||||
return logits
|
||||
|
||||
|
||||
"""Dummy module with dummy logitproc class"""
|
||||
dummy_module = types.ModuleType(DUMMY_LOGITPROC_MODULE)
|
||||
dummy_module.DummyLogitsProcessor = DummyLogitsProcessor # type: ignore
|
||||
|
||||
|
||||
class EntryPoint:
|
||||
"""Dummy entrypoint class for logitsprocs testing"""
|
||||
|
||||
@@ -187,5 +185,59 @@ class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor):
|
||||
return DummyPerReqLogitsProcessor(target_token)
|
||||
|
||||
|
||||
"""Fake version of importlib.metadata.entry_points"""
|
||||
entry_points = lambda group: EntryPoints(group)
|
||||
def register_fake_entrypoint(monkeypatch) -> str:
|
||||
"""Register the dummy logitsproc entrypoint in a way that is visible
|
||||
to spawned subprocesses by creating a real dist-info directory on disk.
|
||||
|
||||
Unlike monkey-patching importlib.metadata.entry_points (which only works
|
||||
with fork), this approach writes a real dist-info package that
|
||||
importlib.metadata can discover in any subprocess via PYTHONPATH.
|
||||
|
||||
Returns the temp directory path.
|
||||
"""
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix="dummy-logitproc-"))
|
||||
dist_info = tmpdir / "dummy_logitproc-0.1.dist-info"
|
||||
dist_info.mkdir()
|
||||
|
||||
# Write METADATA file (required by importlib.metadata)
|
||||
(dist_info / "METADATA").write_text(
|
||||
"Metadata-Version: 2.1\nName: dummy-logitproc\nVersion: 0.1\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Write entry_points.txt
|
||||
(dist_info / "entry_points.txt").write_text(
|
||||
f"[{LOGITSPROCS_GROUP}]\n"
|
||||
f"{DUMMY_LOGITPROC_ENTRYPOINT} = {DUMMY_LOGITPROC_FQCN}\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Add to PYTHONPATH so spawned subprocesses can discover it
|
||||
existing = os.environ.get("PYTHONPATH", "")
|
||||
monkeypatch.setenv(
|
||||
"PYTHONPATH", str(tmpdir) + (os.pathsep + existing if existing else "")
|
||||
)
|
||||
|
||||
# Also update sys.path for the current process so the driver can
|
||||
# discover the entrypoint.
|
||||
monkeypatch.syspath_prepend(str(tmpdir))
|
||||
|
||||
return str(tmpdir)
|
||||
|
||||
|
||||
def fake_entry_points(group: str) -> EntryPoints:
|
||||
"""Fake version of importlib.metadata.entry_points."""
|
||||
return EntryPoints(group)
|
||||
|
||||
|
||||
def setup_fake_entrypoint(monkeypatch) -> None:
|
||||
"""Expose the dummy logitproc entrypoint for the current platform."""
|
||||
if requires_spawn_multiprocessing():
|
||||
register_fake_entrypoint(monkeypatch)
|
||||
monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn")
|
||||
return
|
||||
|
||||
import importlib.metadata
|
||||
|
||||
monkeypatch.setattr(importlib.metadata, "entry_points", fake_entry_points)
|
||||
monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork")
|
||||
|
||||
Reference in New Issue
Block a user