Skip to content

Commit ae6a1cb

Browse files
authored
Fix O(n²) text concatenation in MetalModelRunner.generate() (#60)
This PR is: - To avoid O(n²) string concatenation in non‑v1 vllm_metal/model_runner.py::MetalModelRunner.generate() by collecting streamed segments and joining once - To add unit tests for streamed segment accumulation (v1 and legacy runners) and refactor the tests into clearer, separated files Note: similar to one of the previous PR I submitted; just changing str += to .join() --------- Signed-off-by: Yuan Lik Xun <lxyuan0420@gmail.com>
1 parent 5dd2f58 commit ae6a1cb

3 files changed

Lines changed: 92 additions & 51 deletions

File tree

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,42 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
3+
from __future__ import annotations
4+
5+
from types import SimpleNamespace
6+
7+
import vllm_metal.model_runner as mr
8+
9+
10+
class TestMetalModelRunnerGenerate:
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_make_sampler(*, temp: float):
21+
captured["temp"] = temp
22+
return object()
23+
24+
def fake_stream_generate(model, tokenizer, prompt, max_tokens=256, **kwargs):
25+
captured["prompt"] = prompt
26+
captured["max_tokens"] = max_tokens
27+
captured["kwargs"] = kwargs
28+
yield SimpleNamespace(text="hello")
29+
yield SimpleNamespace(text=" ")
30+
yield SimpleNamespace(text="world")
31+
32+
monkeypatch.setattr(mr, "make_sampler", fake_make_sampler)
33+
monkeypatch.setattr(mr, "stream_generate", fake_stream_generate)
34+
35+
runner = self._make_runner()
36+
out = runner.generate("p", max_tokens=3, temperature=0.7)
37+
38+
assert out == "hello world"
39+
assert captured["prompt"] == "p"
40+
assert captured["max_tokens"] == 3
41+
assert captured["temp"] == 0.7
42+
assert "sampler" in captured["kwargs"]

tests/test_v1_model_runner_generate.py

Lines changed: 47 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -7,51 +7,50 @@
77
import vllm_metal.v1.model_runner as mr
88

99

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"]

vllm_metal/model_runner.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -229,7 +229,7 @@ def generate(
229229
raise RuntimeError(msg)
230230

231231
# Generate tokens using stream_generate
232-
generated_text = ""
232+
segments: list[str] = []
233233

234234
# Create sampler with temperature
235235
sampler = make_sampler(temp=temperature)
@@ -242,9 +242,9 @@ def generate(
242242
sampler=sampler,
243243
):
244244
# Accumulate incremental text from each token
245-
generated_text += response.text
245+
segments.append(response.text)
246246

247-
return generated_text
247+
return "".join(segments)
248248

249249
def __del__(self) -> None:
250250
"""Cleanup model resources."""

0 commit comments

Comments
 (0)