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
+27 -7
View File
@@ -4,7 +4,7 @@ import threading
from unittest.mock import MagicMock, patch
from astrai.inference import STOP
from astrai.inference.engine import GenerateResult
from astrai.inference.engine import GenerateResult, InferenceEngine
def test_result_append_single():
@@ -101,8 +101,6 @@ def test_result_get_results():
def test_engine_generate_non_streaming_single():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
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():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
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():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
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)
results = eng.generate(["hello", "world"])
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()
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):
"""A token matching stop_ids terminates generation for that prompt."""
scheduler, _tok, _model = _make_real_scheduler(device)