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