fix : handle zero-token batch generation

- return empty results without running inference for non-positive limits
- keep scheduler batch outputs aligned with requested max_tokens
- add engine and scheduler regression coverage
This commit is contained in:
2026-08-06 12:31:09 +08:00
parent 6f09b1d2ee
commit 5c180cfa90
4 changed files with 44 additions and 7 deletions
+3
View File
@@ -288,6 +288,9 @@ class InferenceScheduler:
t_max = seq_cap - len(ids) t_max = seq_cap - len(ids)
else: else:
t_max = min(t_max, seq_cap - len(ids)) t_max = min(t_max, seq_cap - len(ids))
if t_max <= 0:
tasks.append(None)
continue
task = Task( task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}", task_id=f"batch_{uuid.uuid4().hex[:8]}",
prompt_ids=list(ids), prompt_ids=list(ids),
+6
View File
@@ -146,6 +146,12 @@ class InferenceEngine:
is_batch = isinstance(prompt, list) is_batch = isinstance(prompt, list)
prompts = prompt if is_batch else [prompt] prompts = prompt if is_batch else [prompt]
if max_tokens is not None and max_tokens <= 0:
if stream:
return iter(())
results = [""] * len(prompts)
return results if is_batch else results[0]
if stream: if stream:
return self._generate_streaming( return self._generate_streaming(
prompts, prompts,
+27 -7
View File
@@ -4,7 +4,7 @@ import threading
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from astrai.inference import STOP from astrai.inference import STOP
from astrai.inference.engine import GenerateResult from astrai.inference.engine import GenerateResult, InferenceEngine
def test_result_append_single(): def test_result_append_single():
@@ -101,8 +101,6 @@ def test_result_get_results():
def test_engine_generate_non_streaming_single(): def test_engine_generate_non_streaming_single():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock() mock_model = MagicMock()
mock_tokenizer = MagicMock() mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3] mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -126,8 +124,6 @@ def test_engine_generate_non_streaming_single():
def test_engine_generate_streaming_yields_tokens(): def test_engine_generate_streaming_yields_tokens():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock() mock_model = MagicMock()
mock_tokenizer = MagicMock() mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3] mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -157,8 +153,6 @@ def test_engine_generate_streaming_yields_tokens():
def test_engine_generate_non_streaming_batch(): def test_engine_generate_non_streaming_batch():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock() mock_model = MagicMock()
mock_tokenizer = MagicMock() mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3] mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -179,3 +173,29 @@ def test_engine_generate_non_streaming_batch():
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2) eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
results = eng.generate(["hello", "world"]) results = eng.generate(["hello", "world"])
assert results == ["r", "r"] assert results == ["r", "r"]
def test_engine_generate_zero_max_tokens_returns_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.stop_ids = [0]
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.remove_task.return_value = []
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
assert eng.generate(["hello", "world"], max_tokens=0) == ["", ""]
instance.add_task.assert_not_called()
def test_engine_generate_zero_max_tokens_stream_is_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
assert list(eng.generate("hello", stream=True, max_tokens=0)) == []
instance.add_task.assert_not_called()
+8
View File
@@ -261,6 +261,14 @@ def test_run_batch_respects_max_tokens(device):
scheduler.stop() scheduler.stop()
def test_run_batch_zero_max_tokens_returns_empty(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
assert scheduler.run_batch([[10, 20, 30]], max_tokens=0) == [[]]
finally:
scheduler.stop()
def test_run_batch_stop_id_terminates(device): def test_run_batch_stop_id_terminates(device):
"""A token matching stop_ids terminates generation for that prompt.""" """A token matching stop_ids terminates generation for that prompt."""
scheduler, _tok, _model = _make_real_scheduler(device) scheduler, _tok, _model = _make_real_scheduler(device)