refactor: TaskManager 剥离页管理,STOP 移至 task.py

- TaskManager 移除 page_cache/page_size 依赖,增 pull_candidates/activate/return_to_waiting
- Executor 增 allocate_pages_for_activation/free_task_pages,承接全部页操作
- STOP 从 cache.py 移至 task.py
- scheduler loop 显式装配: 清理→释页 / 拉取→分配→激活
- sampling.py → sample.py
This commit is contained in:
2026-05-11 14:04:31 +08:00
parent 317ed90bac
commit 73d6cc0f26
7 changed files with 78 additions and 79 deletions
+16 -62
View File
@@ -5,11 +5,12 @@ import uuid
from enum import Enum
from typing import Any, Callable, Dict, List, Optional
from astrai.inference.cache import STOP, PagedCache
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
STOP = object()
class TaskStatus(Enum):
PENDING = "pending"
@@ -64,18 +65,14 @@ class TaskManager:
def __init__(
self,
tokenizer: AutoTokenizer,
page_cache: PagedCache,
max_batch_size: int = 16,
max_seq_len: int = 8192,
max_prompt_len: int = 512,
page_size: int = 64,
):
self.tokenizer = tokenizer
self.page_cache = page_cache
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.max_prompt_len = max_prompt_len
self.page_size = page_size
self.waiting_queue: List[Task] = []
self.active_tasks: List[Task] = []
@@ -124,18 +121,12 @@ class TaskManager:
self._task_event.set()
return task_id
def remove_task(self, task_id: str) -> None:
def remove_task(self, task_id: str) -> List[Task]:
with self._lock:
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id]
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
for task in removed_active:
if not task._pages_freed:
self._free_pages(task.page_table)
task.page_table.clear()
task.n_pages = 0
task._pages_freed = True
return removed_active
def get_stats(self) -> Dict[str, Any]:
return {
@@ -145,7 +136,7 @@ class TaskManager:
"waiting_queue": len(self.waiting_queue),
}
def remove_finished_tasks(self, stop_ids: List[int]) -> None:
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
finished = []
for task in self.active_tasks:
if task.status == TaskStatus.ABORTED:
@@ -157,58 +148,28 @@ class TaskManager:
finished.append(task)
self._total_tokens += task.output_tokens
for task in finished:
if not task._pages_freed:
self._free_pages(task.page_table)
task.page_table.clear()
task.n_pages = 0
task._pages_freed = True
self.active_tasks = [
t
for t in self.active_tasks
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
]
return finished
def refill_active_batch(self) -> None:
available = self.max_batch_size - len(self.active_tasks)
if available <= 0:
return
def pull_candidates(self, n: int) -> List[Task]:
to_add: List[Task] = []
with self._lock:
n = min(available, len(self.waiting_queue))
for _ in range(n):
take = min(n, len(self.waiting_queue))
for _ in range(take):
to_add.append(self.waiting_queue.pop(0))
return to_add
failed: List[Task] = []
for task in to_add:
prompt_len = len(task.prompt_ids)
def activate(self, task: Task) -> None:
task.status = TaskStatus.RUNNING
self.active_tasks.append(task)
hit_pages = self.page_cache.lookup_prefix(task.prompt_ids)
cached_tokens = len(hit_pages) * self.page_size
for p in hit_pages:
self.page_cache.inc_ref(p)
remaining = prompt_len - cached_tokens
n_new = self._n_pages_for(remaining) if remaining > 0 else 0
new_pages = self.page_cache.alloc_n(n_new) if n_new > 0 else []
if remaining > 0 and not new_pages:
for p in hit_pages:
self.page_cache.free(p)
failed.append(task)
continue
task.page_table = hit_pages + new_pages
task.n_pages = len(task.page_table)
task._prefix_cached_tokens = cached_tokens
task.status = TaskStatus.RUNNING
self.active_tasks.append(task)
if failed:
with self._lock:
self.waiting_queue[:0] = failed
def return_to_waiting(self, tasks: List[Task]) -> None:
with self._lock:
self.waiting_queue[:0] = tasks
def has_work(self) -> bool:
return bool(self.active_tasks or self.waiting_queue)
@@ -219,10 +180,3 @@ class TaskManager:
def wake(self) -> None:
self._task_event.set()
def _n_pages_for(self, n_tokens: int) -> int:
return (n_tokens + self.page_size - 1) // self.page_size
def _free_pages(self, indices: List[int]) -> None:
for idx in indices:
self.page_cache.free(idx)