[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:
David Zheng
2026-05-09 12:32:04 +08:00
committed by GitHub
co-authored by gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Andreas Karatzas Kunshang Ji
parent 97cc7685c4
commit df2636a9d8
4 changed files with 129 additions and 48 deletions
+56 -11
View File
@@ -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],
+60 -8
View File
@@ -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")