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
+2 -1
View File
@@ -17,11 +17,12 @@ from astrai.inference.network import get_app, run_server
from astrai.inference.runtime.executor import Executor from astrai.inference.runtime.executor import Executor
from astrai.inference.runtime.sample import sample from astrai.inference.runtime.sample import sample
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus from astrai.inference.task import STOP, GenerationResult, Task, TaskManager, TaskStatus
__all__ = [ __all__ = [
"InferenceEngine", "InferenceEngine",
"InferenceScheduler", "InferenceScheduler",
"GenerationResult",
"Executor", "Executor",
"STOP", "STOP",
"Task", "Task",
+57 -14
View File
@@ -15,7 +15,13 @@ from astrai.extension import (
from astrai.inference.cache import PagePool, TaskCacheManager from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.metrics import MetricsCollector from astrai.inference.metrics import MetricsCollector
from astrai.inference.runtime.executor import Executor from astrai.inference.runtime.executor import Executor
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus from astrai.inference.task import (
STOP,
GenerationResult,
Task,
TaskManager,
TaskStatus,
)
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
@@ -317,7 +323,8 @@ class InferenceScheduler:
frequency_penalty: float = 0.0, frequency_penalty: float = 0.0,
rep_window: int = 64, rep_window: int = 64,
return_logprobs: bool = False, return_logprobs: bool = False,
) -> List[List[int]]: return_details: bool = False,
) -> List[Any]:
"""Synchronous batch generation without the scheduler thread. """Synchronous batch generation without the scheduler thread.
Accepts already-tokenized prompts (no string round-trip) and runs Accepts already-tokenized prompts (no string round-trip) and runs
@@ -333,20 +340,25 @@ class InferenceScheduler:
parameters (uniform across the batch). parameters (uniform across the batch).
return_logprobs: If ``True``, return ``(token_ids, logprobs)`` return_logprobs: If ``True``, return ``(token_ids, logprobs)``
tuples per prompt (logprobs aligned 1-to-1 with token_ids). tuples per prompt (logprobs aligned 1-to-1 with token_ids).
return_details: If ``True``, return a structured result per prompt
with terminal and error reasons. Logprobs are populated when
``return_logprobs`` is also ``True``.
Returns: Returns:
``List[List[int]]`` of generated token IDs per prompt, or — Structured results when ``return_details`` is ``True``;
when ``return_logprobs`` is ``True`` — otherwise generated token IDs per prompt, or token/logprob tuples
``List[Tuple[List[int], List[float]]]``. when ``return_logprobs`` is ``True``.
""" """
stop_ids = self._task_mgr.tokenizer.stop_ids stop_ids = self._task_mgr.tokenizer.stop_ids
seq_cap = self.max_seq_len seq_cap = self.max_seq_len
request_backend = get_backend(use_default=False) request_backend = get_backend(use_default=False)
tasks: List[Task] = [] tasks: List[Optional[Task]] = []
error_reasons: List[Optional[str]] = []
for ids in prompt_ids_list: for ids in prompt_ids_list:
if len(ids) >= seq_cap: if len(ids) >= seq_cap:
tasks.append(None) tasks.append(None)
error_reasons.append("prompt_too_long")
continue continue
t_max = max_tokens t_max = max_tokens
if t_max is None: if t_max is None:
@@ -355,6 +367,7 @@ class InferenceScheduler:
t_max = min(t_max, seq_cap - len(ids)) t_max = min(t_max, seq_cap - len(ids))
if t_max <= 0: if t_max <= 0:
tasks.append(None) tasks.append(None)
error_reasons.append("max_tokens_non_positive")
continue continue
task = Task( task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}", task_id=f"batch_{uuid.uuid4().hex[:8]}",
@@ -369,17 +382,22 @@ class InferenceScheduler:
) )
if not self._task_cache.task_alloc(task.task_id, task.prompt_ids): if not self._task_cache.task_alloc(task.task_id, task.prompt_ids):
tasks.append(None) tasks.append(None)
error_reasons.append("kv_cache_allocation_failed")
continue continue
task.input_tokens = len(task.prompt_ids) task.input_tokens = len(task.prompt_ids)
self._metrics.register(task.task_id) self._metrics.register(task.task_id)
tasks.append(task) tasks.append(task)
error_reasons.append(None)
runtime_errors: Dict[str, str] = {}
try: try:
live = [t for t in tasks if t is not None] live = [t for t in tasks if t is not None]
with self._backend_context(): with self._backend_context():
while live: while live:
decoded, _ = self._step(live, return_logprobs=return_logprobs) decoded, aborted = self._step(live, return_logprobs=return_logprobs)
for task in aborted:
runtime_errors[task.task_id] = "kv_cache_extension_failed"
live = [t for t in decoded if not t.is_finished(stop_ids)] live = [t for t in decoded if not t.is_finished(stop_ids)]
finally: finally:
for t in tasks: for t in tasks:
@@ -389,12 +407,37 @@ class InferenceScheduler:
) )
self._task_cache.task_free(t.task_id) self._task_cache.task_free(t.task_id)
results: List[Any] = [] details: List[GenerationResult] = []
for t in tasks: for t, setup_error in zip(tasks, error_reasons):
if t is None: if t is None:
results.append(([], []) if return_logprobs else []) details.append(
elif return_logprobs: GenerationResult(
results.append((list(t.output_ids), list(t.output_logprobs))) token_ids=[],
logprobs=[],
finish_reason="rejected",
error_reason=setup_error,
)
)
else: else:
results.append(list(t.output_ids)) runtime_error = runtime_errors.get(t.task_id)
return results stopped = bool(t.output_ids and t.output_ids[-1] in stop_ids)
if runtime_error:
finish_reason = "rejected"
elif stopped:
finish_reason = "stop"
else:
finish_reason = "length"
details.append(
GenerationResult(
token_ids=list(t.output_ids),
logprobs=list(t.output_logprobs),
finish_reason=finish_reason,
error_reason=runtime_error,
)
)
if return_details:
return details
if return_logprobs:
return [(result.token_ids, result.logprobs) for result in details]
return [result.token_ids for result in details]
+12 -1
View File
@@ -2,8 +2,9 @@ import threading
import time import time
import uuid import uuid
from collections import deque from collections import deque
from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Optional from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Literal, Optional
from tokenizers.decoders import DecodeStream from tokenizers.decoders import DecodeStream
@@ -16,6 +17,16 @@ if TYPE_CHECKING:
STOP = object() STOP = object()
@dataclass(frozen=True)
class GenerationResult:
"""Structured terminal result for one synchronous generation request."""
token_ids: List[int]
logprobs: List[float]
finish_reason: Literal["stop", "length", "cancelled", "rejected"]
error_reason: Optional[str] = None
class StreamDecoder: class StreamDecoder:
"""Incremental decoder backed by the tokenizers library's DecodeStream. """Incremental decoder backed by the tokenizers library's DecodeStream.
+25 -7
View File
@@ -21,6 +21,7 @@ import torch
from torch import Tensor from torch import Tensor
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import GenerationResult
@dataclass(kw_only=True) @dataclass(kw_only=True)
@@ -171,21 +172,37 @@ class RolloutGenerator:
frequency_penalty=self.frequency_penalty, frequency_penalty=self.frequency_penalty,
rep_window=self.rep_window, rep_window=self.rep_window,
return_logprobs=True, return_logprobs=True,
return_details=True,
) )
if len(results) != B * G: if len(results) != B * G:
raise RuntimeError( raise RuntimeError(
f"Rollout scheduler returned {len(results)} results, expected {B * G}" f"Rollout scheduler returned {len(results)} results, expected {B * G}"
) )
for token_ids, logprobs in results: for result in results:
if len(token_ids) != len(logprobs): if not isinstance(result, GenerationResult):
raise RuntimeError("Rollout scheduler returned an invalid result type")
failures = [
(index, result)
for index, result in enumerate(results)
if result.error_reason is not None
or result.finish_reason in ("cancelled", "rejected")
]
if failures:
reasons = ", ".join(
f"request {index}: {result.error_reason or result.finish_reason}"
for index, result in failures
)
raise RuntimeError(f"Rollout generation failed: {reasons}")
for result in results:
if len(result.token_ids) != len(result.logprobs):
raise RuntimeError( raise RuntimeError(
"Rollout scheduler returned misaligned token IDs and logprobs" "Rollout scheduler returned misaligned token IDs and logprobs"
) )
# Each element is (token_ids, logprobs); pad to max length. # Pad successful structured results to a uniform response length.
max_len = 0 max_len = max((len(result.token_ids) for result in results), default=0)
for token_ids, _lp in results:
max_len = max(max_len, len(token_ids))
max_len = max(max_len, 1) max_len = max(max_len, 1)
device = self.scheduler.device device = self.scheduler.device
@@ -206,7 +223,8 @@ class RolloutGenerator:
response_texts: List[List[str]] = [[] for _ in range(B)] response_texts: List[List[str]] = [[] for _ in range(B)]
for i in range(B): for i in range(B):
for g in range(G): for g in range(G):
token_ids, lps = results[flat_idx] result = results[flat_idx]
token_ids, lps = result.token_ids, result.logprobs
flat_idx += 1 flat_idx += 1
n = len(token_ids) n = len(token_ids)
if n: if n:
+69 -1
View File
@@ -8,7 +8,7 @@ import pytest
import torch import torch
from astrai.extension import CudaBackend, TorchNativeBackend, get_backend 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.metrics import MetricsCollector
from astrai.inference.runtime.executor import DecodeSteadyState, Executor from astrai.inference.runtime.executor import DecodeSteadyState, Executor
from astrai.inference.task import Task from astrai.inference.task import Task
@@ -372,6 +372,74 @@ def test_run_batch_too_long_prompt_skipped(device):
scheduler.stop() 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(): def test_decode_does_not_reuse_previous_batch_state():
executor = object.__new__(Executor) executor = object.__new__(Executor)
executor.device = torch.device("cpu") executor.device = torch.device("cpu")
+20
View File
@@ -4,6 +4,7 @@ import pytest
import torch import torch
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import GenerationResult
from astrai.trainer.rollout import ( from astrai.trainer.rollout import (
BaseRewardModel, BaseRewardModel,
RawRollout, RawRollout,
@@ -173,6 +174,25 @@ def test_rollout_generator_logprobs_are_nonpositive(device):
assert torch.all(lp <= 1e-5) 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): def test_rollout_generator_instruction_role_mapping(device):
"""instruction -> system, input -> user, output -> assistant.""" """instruction -> system, input -> user, output -> assistant."""
gen, _ = _make_generator(device, group_size=1, max_tokens=4) gen, _ = _make_generator(device, group_size=1, max_tokens=4)