test_scheduler_plugins.py 2.15 KB
Newer Older
1
2
# SPDX-License-Identifier: Apache-2.0

3
4
import pytest

5
from vllm.core.scheduler import Scheduler
6
7
8
9
10
from vllm.engine.arg_utils import EngineArgs
from vllm.engine.llm_engine import LLMEngine
from vllm.sampling_params import SamplingParams
from vllm.v1.core.scheduler import Scheduler as V1Scheduler
from vllm.v1.engine.llm_engine import LLMEngine as V1LLMEngine
11
12


13
class DummyV0Scheduler(Scheduler):
14
15

    def schedule(self):
16
17
        raise Exception("Exception raised by DummyV0Scheduler")

18

19
class DummyV1Scheduler(V1Scheduler):
20

21
22
    def schedule(self):
        raise Exception("Exception raised by DummyV1Scheduler")
23
24


25
26
27
28
def test_scheduler_plugins_v0(monkeypatch: pytest.MonkeyPatch):
    with monkeypatch.context() as m:
        m.setenv("VLLM_USE_V1", "0")
        with pytest.raises(Exception) as exception_info:
29

30
31
32
33
34
            engine_args = EngineArgs(
                model="facebook/opt-125m",
                enforce_eager=True,  # reduce test time
                scheduler_cls=DummyV0Scheduler,
            )
35

36
            engine = LLMEngine.from_engine_args(engine_args=engine_args)
37

38
39
40
            sampling_params = SamplingParams(max_tokens=1)
            engine.add_request("0", "foo", sampling_params)
            engine.step()
41

42
43
        assert str(
            exception_info.value) == "Exception raised by DummyV0Scheduler"
44
45


46
47
48
49
50
51
def test_scheduler_plugins_v1(monkeypatch: pytest.MonkeyPatch):
    with monkeypatch.context() as m:
        m.setenv("VLLM_USE_V1", "1")
        # Explicitly turn off engine multiprocessing so
        # that the scheduler runs in this process
        m.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
52

53
        with pytest.raises(Exception) as exception_info:
54

55
56
57
58
59
            engine_args = EngineArgs(
                model="facebook/opt-125m",
                enforce_eager=True,  # reduce test time
                scheduler_cls=DummyV1Scheduler,
            )
60

61
            engine = V1LLMEngine.from_engine_args(engine_args=engine_args)
62

63
64
65
            sampling_params = SamplingParams(max_tokens=1)
            engine.add_request("0", "foo", sampling_params)
            engine.step()
66

67
68
        assert str(
            exception_info.value) == "Exception raised by DummyV1Scheduler"