fix: report failed rollout requests

- Return structured finish and error reasons for synchronous generation
- Reject failed online rollout batches instead of training on empty responses
- Verify allocation and extension failures release metrics and KV state
This commit is contained in:
0z5a
2026-09-02 12:52:50 +08:00
parent c36846c8a4
commit 90de5bc1bd
6 changed files with 185 additions and 24 deletions
+69 -1
View File
@@ -8,7 +8,7 @@ import pytest
import torch
from astrai.extension import CudaBackend, TorchNativeBackend, get_backend
from astrai.inference import InferenceScheduler
from astrai.inference import GenerationResult, InferenceScheduler
from astrai.inference.metrics import MetricsCollector
from astrai.inference.runtime.executor import DecodeSteadyState, Executor
from astrai.inference.task import Task
@@ -372,6 +372,74 @@ def test_run_batch_too_long_prompt_skipped(device):
scheduler.stop()
def test_run_batch_details_distinguish_rejection_from_success(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
long_prompt = list(range(100))
results = scheduler.run_batch(
[long_prompt, [10, 20]],
max_tokens=2,
temperature=0,
return_logprobs=True,
return_details=True,
)
assert results[0] == GenerationResult(
token_ids=[],
logprobs=[],
finish_reason="rejected",
error_reason="prompt_too_long",
)
assert results[1].finish_reason in ("stop", "length")
assert results[1].error_reason is None
assert len(results[1].token_ids) == len(results[1].logprobs)
finally:
scheduler.stop()
def test_run_batch_details_report_non_positive_max_tokens(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
result = scheduler.run_batch([[10, 20]], max_tokens=0, return_details=True)[0]
assert result.finish_reason == "rejected"
assert result.error_reason == "max_tokens_non_positive"
finally:
scheduler.stop()
def test_run_batch_details_report_allocation_failure(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
with patch.object(scheduler._task_cache, "task_alloc", return_value=False):
result = scheduler.run_batch([[10, 20]], max_tokens=2, return_details=True)[
0
]
assert result.finish_reason == "rejected"
assert result.error_reason == "kv_cache_allocation_failed"
finally:
scheduler.stop()
def test_run_batch_details_report_extension_failure_and_cleanup(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
with patch.object(
scheduler,
"_step",
side_effect=lambda tasks, **_kwargs: ([], list(tasks)),
):
result = scheduler.run_batch([[10, 20]], max_tokens=2, return_details=True)[
0
]
assert result.finish_reason == "rejected"
assert result.error_reason == "kv_cache_extension_failed"
assert scheduler._task_cache._states == {}
assert scheduler._metrics._timings == {}
finally:
scheduler.stop()
def test_decode_does_not_reuse_previous_batch_state():
executor = object.__new__(Executor)
executor.device = torch.device("cpu")