Files
AstrAI/astrai/inference/core/scheduler.py
T
ViperEkura 53c804e233 refactor : merge max_prompt_len into max_seq_len, replace assert with raise
- Engine/Scheduler/TaskManager: merge max_prompt_len into max_seq_len
- train.py: replace bare assert with ValueError/FileNotFoundError
- server.py: add --max_seq_len CLI option
- engine.py: remove dead page_size param
2026-07-27 08:05:11 +08:00

310 lines
12 KiB
Python

import logging
import threading
import uuid
from typing import Any, Dict, List, Optional, Tuple
import torch
from astrai.inference.core.cache import ContiguousCache, KVCache
from astrai.inference.core.executor import Executor
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
class InferenceScheduler:
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
def __init__(
self,
model: AutoModel,
tokenizer: AutoTokenizer,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
cache: Optional[KVCache] = None,
):
config = model.config
if max_seq_len is not None:
self.max_seq_len = max_seq_len
elif config.max_position_embeddings is not None:
self.max_seq_len = config.max_position_embeddings
else:
raise ValueError(
"max_seq_len must be provided either as argument "
"or in model config (config.max_position_embeddings)"
)
self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype
head_dim = config.hidden_size // config.num_attention_heads
if cache is not None:
self._cache = cache
else:
self._cache = ContiguousCache(
config.num_hidden_layers,
max_batch_size,
self.max_seq_len,
config.num_key_value_heads,
head_dim,
self.device,
self.dtype,
)
self._task_mgr = TaskManager(
tokenizer=tokenizer,
max_batch_size=max_batch_size,
max_seq_len=self.max_seq_len,
)
self._executor = Executor(
model=model,
tokenizer=tokenizer,
kv_cache=self._cache,
device=self.device,
dtype=self.dtype,
)
self._stop_event = threading.Event()
self._loop_thread: Optional[threading.Thread] = None
def add_task(self, prompt: str, **kwargs) -> str:
return self._task_mgr.add_task(prompt, **kwargs)
def remove_task(self, task_id: str):
for task in self._task_mgr.remove_task(task_id):
self._cache.task_free(task.task_id)
def get_stats(self) -> Dict[str, Any]:
return self._task_mgr.get_stats()
def _run_generation_loop(self):
stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
try:
while not self._stop_event.is_set():
finished = self._task_mgr.remove_finished_tasks(stop_ids)
for task in finished:
cache.task_free(task.task_id)
active = self._task_mgr.get_active_tasks()
available = self._task_mgr.max_batch_size - len(active)
if available > 0:
candidates = self._task_mgr.pull_candidates(available)
failed = []
for task in candidates:
if cache.task_alloc(task.task_id, task.prompt_ids):
self._task_mgr.activate(task)
else:
failed.append(task)
if failed:
self._task_mgr.return_to_waiting(failed)
if not self._task_mgr.has_work():
self._task_mgr.wait_for_tasks(timeout=1.0)
continue
to_prefill = [
t
for t in self._task_mgr.get_active_tasks()
if t.output_tokens == 0
and cache.task_cached(t.task_id) < len(t.prompt_ids)
]
if to_prefill:
for t in to_prefill:
t.input_tokens = len(t.prompt_ids)
groups: Dict[Tuple[int, int], List[Task]] = {}
for t in to_prefill:
key = (
len(t.prompt_ids),
cache.task_cached(t.task_id),
)
groups.setdefault(key, []).append(t)
for (prompt_len, start_pos), group in groups.items():
self._executor.execute_prefill(group, prompt_len, start_pos)
start_logical_page = start_pos // getattr(
cache, "page_size", 64
)
for t in group:
cache.task_record_hashes(
t.task_id, t.prompt_ids, start_logical_page
)
decode_tasks = self._task_mgr.get_active_tasks()
valid: List[Task] = []
for t in sorted(decode_tasks, key=lambda t: t.task_id):
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
self._task_mgr.invoke_callback(t.task_id, STOP)
if valid:
next_tokens = self._executor.execute_decode(valid)
for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok)
t.output_tokens += 1
new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text)
for t in valid:
if t.is_finished(stop_ids):
remaining = t.flush_remaining(self._task_mgr.tokenizer)
if remaining:
self._task_mgr.invoke_callback(t.task_id, remaining)
self._task_mgr.invoke_callback(t.task_id, STOP)
except Exception as e:
self._stop_event.set()
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_mgr.clear_queues()
def start(self):
if self._loop_thread is not None and self._loop_thread.is_alive():
return
self._stop_event.clear()
t = threading.Thread(target=self._run_generation_loop, daemon=True)
t.start()
self._loop_thread = t
def stop(self):
self._stop_event.set()
self._task_mgr.wake()
if self._loop_thread is not None:
self._loop_thread.join(timeout=2.0)
self._loop_thread = None
for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._cache.task_free(task.task_id)
self._task_mgr.clear_queues()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def run_batch(
self,
prompt_ids_list: List[List[int]],
*,
max_tokens: Optional[int] = None,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
return_logprobs: bool = False,
) -> List[List[int]]:
"""Synchronous batch generation without the scheduler thread.
Accepts already-tokenized prompts (no string round-trip) and runs
prefill + decode to completion on the calling thread. Designed for
RL rollout, where logprobs of the behaviour policy must be collected
alongside generated tokens.
Args:
prompt_ids_list: ``B`` prompts, each a list of token IDs.
max_tokens: Maximum tokens to generate per prompt. ``None``
uses ``self.max_seq_len - len(prompt_ids)``.
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
parameters (uniform across the batch).
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
Returns:
``List[List[int]]`` of generated token IDs per prompt, or —
when ``return_logprobs`` is ``True`` —
``List[Tuple[List[int], List[float]]]``.
"""
stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
seq_cap = self.max_seq_len
tasks: List[Task] = []
for ids in prompt_ids_list:
if len(ids) >= seq_cap:
tasks.append(None)
continue
t_max = max_tokens
if t_max is None:
t_max = seq_cap - len(ids)
else:
t_max = min(t_max, seq_cap - len(ids))
task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}",
prompt_ids=list(ids),
max_tokens=t_max,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
if not cache.task_alloc(task.task_id, task.prompt_ids):
tasks.append(None)
continue
task.input_tokens = len(task.prompt_ids)
tasks.append(task)
try:
live = [t for t in tasks if t is not None]
prefill_groups: Dict[Tuple[int, int], List[Task]] = {}
for t in live:
key = (len(t.prompt_ids), cache.task_cached(t.task_id))
prefill_groups.setdefault(key, []).append(t)
for (prompt_len, start_pos), group in prefill_groups.items():
self._executor.execute_prefill(group, prompt_len, start_pos)
while live:
valid: List[Task] = []
for t in sorted(live, key=lambda x: x.task_id):
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
if not valid:
break
step_out = self._executor.execute_decode(
valid, return_logprobs=return_logprobs
)
if return_logprobs:
for t, (ntok, _lp) in zip(valid, step_out):
t.output_ids.append(ntok)
t.output_tokens += 1
else:
for t, ntok in zip(valid, step_out):
t.output_ids.append(ntok)
t.output_tokens += 1
live = [t for t in valid if not t.is_finished(stop_ids)]
finally:
for t in tasks:
if t is not None:
cache.task_free(t.task_id)
results: List[Any] = []
for t in tasks:
if t is None:
results.append(([], []) if return_logprobs else [])
elif return_logprobs:
results.append((list(t.output_ids), list(t.output_logprobs)))
else:
results.append(list(t.output_ids))
return results