|
7 | 7 | import vllm_metal.v1.model_runner as mr |
8 | 8 |
|
9 | 9 |
|
10 | | -def _make_runner(mr_module): |
11 | | - runner = mr_module.MetalModelRunner.__new__(mr_module.MetalModelRunner) |
12 | | - runner.model = object() |
13 | | - runner.tokenizer = object() |
14 | | - return runner |
15 | | - |
16 | | - |
17 | | -def test_generate_accumulates_streamed_segments(monkeypatch) -> None: |
18 | | - captured: dict[str, object] = {} |
19 | | - |
20 | | - def fake_stream_generate(model, tokenizer, prompt, max_tokens=256, **kwargs): |
21 | | - captured["prompt"] = prompt |
22 | | - captured["max_tokens"] = max_tokens |
23 | | - captured["kwargs"] = kwargs |
24 | | - yield SimpleNamespace(text="hello") |
25 | | - yield SimpleNamespace(text=" ") |
26 | | - yield SimpleNamespace(text="world") |
27 | | - |
28 | | - monkeypatch.setattr(mr, "stream_generate", fake_stream_generate) |
29 | | - |
30 | | - runner = _make_runner(mr) |
31 | | - out = runner.generate("p", max_tokens=3, temperature=0.0) |
32 | | - |
33 | | - assert out == "hello world" |
34 | | - assert captured["prompt"] == "p" |
35 | | - assert captured["max_tokens"] == 3 |
36 | | - # mlx_lm 0.29+ uses sampler parameter instead of temp |
37 | | - assert "sampler" in captured["kwargs"] |
38 | | - assert callable(captured["kwargs"]["sampler"]) |
39 | | - |
40 | | - |
41 | | -def test_generate_passes_sampler_for_temperature_sampling(monkeypatch) -> None: |
42 | | - captured: dict[str, object] = {} |
43 | | - |
44 | | - def fake_stream_generate(model, tokenizer, prompt, max_tokens=256, **kwargs): |
45 | | - captured["kwargs"] = kwargs |
46 | | - assert "sampler" in kwargs |
47 | | - assert callable(kwargs["sampler"]) |
48 | | - yield SimpleNamespace(text="a") |
49 | | - yield SimpleNamespace(text="b") |
50 | | - |
51 | | - monkeypatch.setattr(mr, "stream_generate", fake_stream_generate) |
52 | | - |
53 | | - runner = _make_runner(mr) |
54 | | - out = runner.generate("p", max_tokens=2, temperature=0.5) |
55 | | - |
56 | | - assert out == "ab" |
57 | | - assert "sampler" in captured["kwargs"] |
| 10 | +class TestV1MetalModelRunnerGenerate: |
| 11 | + def _make_runner(self) -> mr.MetalModelRunner: |
| 12 | + runner = mr.MetalModelRunner.__new__(mr.MetalModelRunner) |
| 13 | + runner.model = object() |
| 14 | + runner.tokenizer = object() |
| 15 | + return runner |
| 16 | + |
| 17 | + def test_accumulates_streamed_segments(self, monkeypatch) -> None: |
| 18 | + captured: dict[str, object] = {} |
| 19 | + |
| 20 | + def fake_stream_generate(model, tokenizer, prompt, max_tokens=256, **kwargs): |
| 21 | + captured["prompt"] = prompt |
| 22 | + captured["max_tokens"] = max_tokens |
| 23 | + captured["kwargs"] = kwargs |
| 24 | + yield SimpleNamespace(text="hello") |
| 25 | + yield SimpleNamespace(text=" ") |
| 26 | + yield SimpleNamespace(text="world") |
| 27 | + |
| 28 | + monkeypatch.setattr(mr, "stream_generate", fake_stream_generate) |
| 29 | + |
| 30 | + runner = self._make_runner() |
| 31 | + out = runner.generate("p", max_tokens=3, temperature=0.0) |
| 32 | + |
| 33 | + assert out == "hello world" |
| 34 | + assert captured["prompt"] == "p" |
| 35 | + assert captured["max_tokens"] == 3 |
| 36 | + # mlx_lm 0.29+ uses sampler parameter instead of temp |
| 37 | + assert "sampler" in captured["kwargs"] |
| 38 | + assert callable(captured["kwargs"]["sampler"]) |
| 39 | + |
| 40 | + def test_passes_sampler_for_temperature_sampling(self, monkeypatch) -> None: |
| 41 | + captured: dict[str, object] = {} |
| 42 | + |
| 43 | + def fake_stream_generate(model, tokenizer, prompt, max_tokens=256, **kwargs): |
| 44 | + captured["kwargs"] = kwargs |
| 45 | + assert "sampler" in kwargs |
| 46 | + assert callable(kwargs["sampler"]) |
| 47 | + yield SimpleNamespace(text="a") |
| 48 | + yield SimpleNamespace(text="b") |
| 49 | + |
| 50 | + monkeypatch.setattr(mr, "stream_generate", fake_stream_generate) |
| 51 | + |
| 52 | + runner = self._make_runner() |
| 53 | + out = runner.generate("p", max_tokens=2, temperature=0.5) |
| 54 | + |
| 55 | + assert out == "ab" |
| 56 | + assert "sampler" in captured["kwargs"] |
0 commit comments