- Precompute page_table and decode_mask on KVCache once per step in PagePool.bind_tasks, instead of per-layer in CudaBackend/TorchNativeBackend - Skip frequency penalty history tensor construction when all penalties are 0 in Executor.execute_decode - Omit FrequencyPenaltyStrategy from sampling pipeline when penalty is 0 - Deduplicate get_active_tasks calls in scheduler loop (3 to 1), remove redundant sorted() on decode tasks - Benchmark (L20, bf16, CUDA backend): B=1 9.48->9.40ms (+1%), B=4 10.73->9.89ms (+8.6%), B=8 10.77->10.13ms (+6.4%)
312 lines
12 KiB
Python
312 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 PagePool
|
|
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[PagePool] = 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 = PagePool(
|
|
n_layers=config.num_hidden_layers,
|
|
n_kv_heads=config.num_key_value_heads,
|
|
head_dim=head_dim,
|
|
max_batch_size=max_batch_size,
|
|
max_seq_len=self.max_seq_len,
|
|
device=self.device,
|
|
dtype=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
|
|
|
|
active = self._task_mgr.get_active_tasks()
|
|
|
|
to_prefill = [
|
|
t
|
|
for t in active
|
|
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 = active
|
|
|
|
valid: List[Task] = []
|
|
for t in decode_tasks:
|
|
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
|