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:
@@ -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")
|
||||
|
||||
@@ -4,6 +4,7 @@ import pytest
|
||||
import torch
|
||||
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import GenerationResult
|
||||
from astrai.trainer.rollout import (
|
||||
BaseRewardModel,
|
||||
RawRollout,
|
||||
@@ -173,6 +174,25 @@ def test_rollout_generator_logprobs_are_nonpositive(device):
|
||||
assert torch.all(lp <= 1e-5)
|
||||
|
||||
|
||||
def test_rollout_generator_rejects_failed_requests(device):
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||
|
||||
def failed_run_batch(*_args, **kwargs):
|
||||
assert kwargs["return_details"] is True
|
||||
return [
|
||||
GenerationResult([1], [-0.1], "length"),
|
||||
GenerationResult([], [], "rejected", "kv_cache_allocation_failed"),
|
||||
]
|
||||
|
||||
gen.scheduler.run_batch = failed_run_batch
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="Rollout generation failed: request 1: kv_cache_allocation_failed",
|
||||
):
|
||||
gen.generate(_make_instruction_batch(n=1))
|
||||
|
||||
|
||||
def test_rollout_generator_instruction_role_mapping(device):
|
||||
"""instruction -> system, input -> user, output -> assistant."""
|
||||
gen, _ = _make_generator(device, group_size=1, max_tokens=4)
|
||||
|
||||
Reference in New Issue
Block a user