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.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",
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user