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:
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user