From 5c180cfa90e731119865d15f5ad95beb00c62b35 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Thu, 6 Aug 2026 12:31:09 +0800 Subject: [PATCH] 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 --- astrai/inference/core/scheduler.py | 3 +++ astrai/inference/engine.py | 6 ++++++ tests/inference/test_engine.py | 34 ++++++++++++++++++++++++------ tests/inference/test_scheduler.py | 8 +++++++ 4 files changed, 44 insertions(+), 7 deletions(-) diff --git a/astrai/inference/core/scheduler.py b/astrai/inference/core/scheduler.py index f17a45b..9236da2 100644 --- a/astrai/inference/core/scheduler.py +++ b/astrai/inference/core/scheduler.py @@ -288,6 +288,9 @@ class InferenceScheduler: t_max = seq_cap - len(ids) else: t_max = min(t_max, seq_cap - len(ids)) + if t_max <= 0: + tasks.append(None) + continue task = Task( task_id=f"batch_{uuid.uuid4().hex[:8]}", prompt_ids=list(ids), diff --git a/astrai/inference/engine.py b/astrai/inference/engine.py index 7b6f412..7eb6ddb 100644 --- a/astrai/inference/engine.py +++ b/astrai/inference/engine.py @@ -146,6 +146,12 @@ class InferenceEngine: is_batch = isinstance(prompt, list) 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: return self._generate_streaming( prompts, diff --git a/tests/inference/test_engine.py b/tests/inference/test_engine.py index 9b76e8c..b240619 100644 --- a/tests/inference/test_engine.py +++ b/tests/inference/test_engine.py @@ -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() diff --git a/tests/inference/test_scheduler.py b/tests/inference/test_scheduler.py index dde02ff..0225fd4 100644 --- a/tests/inference/test_scheduler.py +++ b/tests/inference/test_scheduler.py @@ -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)