fix: 修复 remove_task 未释放 KV cache slot 导致第二轮对话死锁

- remove_task() 现在释放 KV cache slot 和 prefix cache 引用
- _refill_active_batch 中 alloc 失败时将剩余 task 推回 waiting_queue
- 主循环增加 try/except 异常兜底,发送 _STOP 给所有 task
- 重构:server.py 全局变量改为 ServerState 类;automodel.py
  使用 Registry 替代裸 dict;合并 TrainContextBuilder 的 with_*
  方法到 build()
This commit is contained in:
2026-05-08 14:53:04 +08:00
parent ffff05b2c6
commit a6f5ff3b37
8 changed files with 165 additions and 142 deletions
-3
View File
@@ -2,7 +2,6 @@
import asyncio
import gc
import logging
import threading
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
@@ -12,8 +11,6 @@ import torch.nn as nn
from astrai.inference.scheduler import _STOP, InferenceScheduler
from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__)
class GenerationRequest:
"""Request parameters for text generation.
+51 -25
View File
@@ -1,5 +1,6 @@
"""Inference scheduler for single-GPU continuous batching."""
import logging
import threading
import time
import uuid
@@ -12,6 +13,8 @@ from torch import Tensor
from astrai.model.automodel import AutoModel
from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__)
_STOP = object()
@@ -506,9 +509,20 @@ class InferenceScheduler:
task_id: The task to remove.
"""
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 task.prefix_len > 0:
prefix = tuple(task.prompt_ids[: task.prefix_len])
self.prefix_cache.release(prefix)
if task.prefix_len < len(task.prompt_ids):
self.prefix_cache.release(tuple(task.prompt_ids))
if task.slot >= 0:
self._free_slot(task.slot)
task.slot = -1
def _remove_finished_tasks(self) -> None:
"""Removes all finished tasks from the active batch.
@@ -553,7 +567,7 @@ class InferenceScheduler:
for _ in range(n):
to_add.append(self.waiting_queue.pop(0))
for task in to_add:
for i, task in enumerate(to_add):
slot = -1
reused = False
if task.prefix_len > 0:
@@ -564,6 +578,8 @@ class InferenceScheduler:
if slot < 0:
slot = self._alloc_slot()
if slot < 0:
with self._lock:
self.waiting_queue[:0] = to_add[i:]
break
task.slot = slot
task.status = TaskStatus.RUNNING
@@ -712,32 +728,42 @@ class InferenceScheduler:
Decode processes only the largest position group to ensure all
batched tasks share the same KV cache write position.
"""
while self._running:
self._remove_finished_tasks()
self._refill_active_batch()
try:
while self._running:
self._remove_finished_tasks()
self._refill_active_batch()
with self._lock:
if not self.active_tasks and not self.waiting_queue:
with self._lock:
if not self.active_tasks and not self.waiting_queue:
self._task_event.clear()
self._task_event.wait(timeout=0.01)
continue
tasks = self.active_tasks[:]
to_prefill = [t for t in tasks if t.output_tokens == 0]
if to_prefill:
self._execute_prefill(to_prefill)
pos_groups: Dict[int, List[Task]] = {}
for t in self.active_tasks:
pos_groups.setdefault(t.next_pos, []).append(t)
if pos_groups:
best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
self._execute_decode(pos_groups[best_pos], best_pos)
if not self.waiting_queue and len(self.active_tasks) <= 1:
self._task_event.wait(timeout=0.005)
self._task_event.clear()
self._task_event.wait(timeout=0.01)
continue
tasks = self.active_tasks[:]
to_prefill = [t for t in tasks if t.output_tokens == 0]
if to_prefill:
self._execute_prefill(to_prefill)
pos_groups: Dict[int, List[Task]] = {}
for t in self.active_tasks:
pos_groups.setdefault(t.next_pos, []).append(t)
if pos_groups:
best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
self._execute_decode(pos_groups[best_pos], best_pos)
if not self.waiting_queue and len(self.active_tasks) <= 1:
self._task_event.wait(timeout=0.005)
self._task_event.clear()
except Exception as e:
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self.active_tasks:
if task.stream_callback:
task.stream_callback(_STOP)
for task in self.waiting_queue:
if task.stream_callback:
task.stream_callback(_STOP)
raise
def start(self) -> None:
"""Starts the background generation loop thread."""
+55 -39
View File
@@ -23,16 +23,30 @@ from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__)
_engine: Optional[InferenceEngine] = None
_model_param: Optional[Any] = None
_project_root = Path(__file__).parent.parent.parent
_server_config: Dict[str, Any] = {
"device": "cuda",
"dtype": torch.bfloat16,
"param_path": None,
"max_batch_size": 16,
}
class ServerState:
"""Encapsulates all server runtime state.
Attributes:
engine: The inference engine instance.
model_param: The loaded model.
config: Server configuration dict.
"""
def __init__(self):
self.engine: Optional[InferenceEngine] = None
self.model_param: Optional[Any] = None
self.config: Dict[str, Any] = {
"device": "cuda",
"dtype": torch.bfloat16,
"param_path": None,
"max_batch_size": 16,
}
_state = ServerState()
def configure_server(
@@ -41,28 +55,29 @@ def configure_server(
param_path: Optional[Path] = None,
max_batch_size: int = 16,
):
_server_config["device"] = device
_server_config["dtype"] = dtype
_server_config["param_path"] = param_path
_server_config["max_batch_size"] = max_batch_size
_state.config.update(
device=device,
dtype=dtype,
param_path=param_path,
max_batch_size=max_batch_size,
)
@asynccontextmanager
async def lifespan(app: FastAPI):
global _model_param, _engine
try:
load_model(
param_path=_server_config["param_path"],
device=_server_config["device"],
dtype=_server_config["dtype"],
max_batch_size=_server_config["max_batch_size"],
param_path=_state.config["param_path"],
device=_state.config["device"],
dtype=_state.config["dtype"],
max_batch_size=_state.config["max_batch_size"],
)
except Exception as e:
logger.error(f"Failed to load model: {e}")
raise
yield
if _engine:
_engine.shutdown()
if _state.engine:
_state.engine.shutdown()
logger.info("Inference engine shutdown complete")
@@ -75,25 +90,30 @@ def load_model(
dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16,
):
global _model_param, _engine
if param_path is None:
param_path = _project_root / "params"
if not param_path.exists():
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
tokenizer = AutoTokenizer.from_pretrained(param_path)
_model_param = AutoModel.from_pretrained(param_path)
_model_param.to(device=device, dtype=dtype)
_state.model_param = AutoModel.from_pretrained(param_path)
_state.model_param.to(device=device, dtype=dtype)
logger.info(f"Model loaded on {device} with dtype {dtype}")
_engine = InferenceEngine(
model=_model_param,
_state.engine = InferenceEngine(
model=_state.model_param,
tokenizer=tokenizer,
max_batch_size=max_batch_size,
)
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
def _get_engine() -> InferenceEngine:
if _state.engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
return _state.engine
class ChatMessage(BaseModel):
role: str
content: str
@@ -121,30 +141,27 @@ class CompletionResponse(BaseModel):
async def health():
return {
"status": "ok",
"model_loaded": _model_param is not None,
"engine_ready": _engine is not None,
"model_loaded": _state.model_param is not None,
"engine_ready": _state.engine is not None,
}
@app.get("/stats")
async def get_stats():
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
return _engine.get_stats()
return _get_engine().get_stats()
@app.post("/v1/chat/completions", response_model=CompletionResponse)
async def chat_completion(request: ChatCompletionRequest):
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
engine = _get_engine()
prompt = _engine.tokenizer.apply_chat_template(
prompt = engine.tokenizer.apply_chat_template(
[{"role": m.role, "content": m.content} for m in request.messages],
tokenize=False,
)
if request.stream:
agen = _engine.generate_async(
agen = engine.generate_async(
prompt=prompt,
max_tokens=request.max_tokens,
temperature=request.temperature,
@@ -163,7 +180,7 @@ async def chat_completion(request: ChatCompletionRequest):
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
)
else:
result = _engine.generate(
result = engine.generate(
prompt=prompt,
stream=False,
max_tokens=request.max_tokens,
@@ -198,8 +215,7 @@ async def generate(
max_len: int = 2048,
stream: bool = False,
):
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
engine = _get_engine()
messages = []
if history:
@@ -209,10 +225,10 @@ async def generate(
messages.append({"role": "assistant", "content": h[1]})
messages.append({"role": "user", "content": query})
prompt = _engine.tokenizer.apply_chat_template(messages, tokenize=False)
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
if stream:
agen = _engine.generate_async(
agen = engine.generate_async(
prompt=prompt,
max_tokens=max_len,
temperature=temperature,
@@ -226,7 +242,7 @@ async def generate(
return StreamingResponse(text_stream(), media_type="text/plain")
else:
result = _engine.generate(
result = engine.generate(
prompt=prompt,
stream=False,
max_tokens=max_len,