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:
@@ -17,11 +17,12 @@ from astrai.inference.network import get_app, run_server
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.sample import sample
|
||||
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__ = [
|
||||
"InferenceEngine",
|
||||
"InferenceScheduler",
|
||||
"GenerationResult",
|
||||
"Executor",
|
||||
"STOP",
|
||||
"Task",
|
||||
|
||||
@@ -15,7 +15,13 @@ from astrai.extension import (
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
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.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
@@ -317,7 +323,8 @@ class InferenceScheduler:
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
return_logprobs: bool = False,
|
||||
) -> List[List[int]]:
|
||||
return_details: bool = False,
|
||||
) -> List[Any]:
|
||||
"""Synchronous batch generation without the scheduler thread.
|
||||
|
||||
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||
@@ -333,20 +340,25 @@ class InferenceScheduler:
|
||||
parameters (uniform across the batch).
|
||||
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||
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:
|
||||
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||
when ``return_logprobs`` is ``True`` —
|
||||
``List[Tuple[List[int], List[float]]]``.
|
||||
Structured results when ``return_details`` is ``True``;
|
||||
otherwise generated token IDs per prompt, or token/logprob tuples
|
||||
when ``return_logprobs`` is ``True``.
|
||||
"""
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
seq_cap = self.max_seq_len
|
||||
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:
|
||||
if len(ids) >= seq_cap:
|
||||
tasks.append(None)
|
||||
error_reasons.append("prompt_too_long")
|
||||
continue
|
||||
t_max = max_tokens
|
||||
if t_max is None:
|
||||
@@ -355,6 +367,7 @@ class InferenceScheduler:
|
||||
t_max = min(t_max, seq_cap - len(ids))
|
||||
if t_max <= 0:
|
||||
tasks.append(None)
|
||||
error_reasons.append("max_tokens_non_positive")
|
||||
continue
|
||||
task = Task(
|
||||
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):
|
||||
tasks.append(None)
|
||||
error_reasons.append("kv_cache_allocation_failed")
|
||||
continue
|
||||
task.input_tokens = len(task.prompt_ids)
|
||||
self._metrics.register(task.task_id)
|
||||
tasks.append(task)
|
||||
error_reasons.append(None)
|
||||
|
||||
runtime_errors: Dict[str, str] = {}
|
||||
try:
|
||||
live = [t for t in tasks if t is not None]
|
||||
|
||||
with self._backend_context():
|
||||
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)]
|
||||
finally:
|
||||
for t in tasks:
|
||||
@@ -389,12 +407,37 @@ class InferenceScheduler:
|
||||
)
|
||||
self._task_cache.task_free(t.task_id)
|
||||
|
||||
results: List[Any] = []
|
||||
for t in tasks:
|
||||
details: List[GenerationResult] = []
|
||||
for t, setup_error in zip(tasks, error_reasons):
|
||||
if t is None:
|
||||
results.append(([], []) if return_logprobs else [])
|
||||
elif return_logprobs:
|
||||
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||
details.append(
|
||||
GenerationResult(
|
||||
token_ids=[],
|
||||
logprobs=[],
|
||||
finish_reason="rejected",
|
||||
error_reason=setup_error,
|
||||
)
|
||||
)
|
||||
else:
|
||||
results.append(list(t.output_ids))
|
||||
return results
|
||||
runtime_error = runtime_errors.get(t.task_id)
|
||||
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]
|
||||
|
||||
@@ -2,8 +2,9 @@ import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
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
|
||||
|
||||
@@ -16,6 +17,16 @@ if TYPE_CHECKING:
|
||||
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:
|
||||
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import GenerationResult
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
@@ -171,21 +172,37 @@ class RolloutGenerator:
|
||||
frequency_penalty=self.frequency_penalty,
|
||||
rep_window=self.rep_window,
|
||||
return_logprobs=True,
|
||||
return_details=True,
|
||||
)
|
||||
if len(results) != B * G:
|
||||
raise RuntimeError(
|
||||
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
|
||||
)
|
||||
for token_ids, logprobs in results:
|
||||
if len(token_ids) != len(logprobs):
|
||||
for result in results:
|
||||
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(
|
||||
"Rollout scheduler returned misaligned token IDs and logprobs"
|
||||
)
|
||||
|
||||
# Each element is (token_ids, logprobs); pad to max length.
|
||||
max_len = 0
|
||||
for token_ids, _lp in results:
|
||||
max_len = max(max_len, len(token_ids))
|
||||
# Pad successful structured results to a uniform response length.
|
||||
max_len = max((len(result.token_ids) for result in results), default=0)
|
||||
max_len = max(max_len, 1)
|
||||
|
||||
device = self.scheduler.device
|
||||
@@ -206,7 +223,8 @@ class RolloutGenerator:
|
||||
response_texts: List[List[str]] = [[] for _ in range(B)]
|
||||
for i in range(B):
|
||||
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
|
||||
n = len(token_ids)
|
||||
if n:
|
||||
|
||||
@@ -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