fix: cancel abandoned generation tasks
- Propagate stream closure and stop-sequence termination into scheduler cancellation - Defer active KV release to the scheduler owner and close metrics safely - Expose lifecycle counters and cover waiting, active, and allocation-race cleanup
This commit is contained in:
Vendored
+5
@@ -322,6 +322,11 @@ class TaskCacheManager:
|
|||||||
state.length = pos + 1
|
state.length = pos + 1
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
@property
|
||||||
|
def task_count(self) -> int:
|
||||||
|
"""Number of tasks currently holding KV request state."""
|
||||||
|
return len(self._states)
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
def task_cached(self, task_id: str) -> int:
|
||||||
state = self._states.get(task_id)
|
state = self._states.get(task_id)
|
||||||
return state.cached if state is not None else 0
|
return state.cached if state is not None else 0
|
||||||
|
|||||||
+43
-32
@@ -43,8 +43,7 @@ class GenerateResult:
|
|||||||
with self._cond:
|
with self._cond:
|
||||||
out = self.tokens.copy()
|
out = self.tokens.copy()
|
||||||
self.tokens.clear()
|
self.tokens.clear()
|
||||||
if not out:
|
self._event.clear()
|
||||||
self._event.clear()
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def wait(self, timeout: Optional[float] = None) -> bool:
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
@@ -141,25 +140,34 @@ class InferenceEngine:
|
|||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
sync_gen = self._generate(
|
request_backend = get_backend(use_default=False)
|
||||||
[prompt],
|
result = GenerateResult()
|
||||||
False,
|
task_id = self.scheduler.add_task(
|
||||||
True,
|
prompt=prompt,
|
||||||
max_tokens,
|
max_tokens=max_tokens,
|
||||||
temperature,
|
temperature=temperature,
|
||||||
top_p,
|
top_p=top_p,
|
||||||
top_k,
|
top_k=top_k,
|
||||||
frequency_penalty,
|
frequency_penalty=frequency_penalty,
|
||||||
rep_window,
|
rep_window=rep_window,
|
||||||
|
backend=request_backend,
|
||||||
|
stream_callback=result.append,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _agen():
|
async def _agen():
|
||||||
loop = asyncio.get_event_loop()
|
finished = False
|
||||||
while True:
|
try:
|
||||||
token = await loop.run_in_executor(None, next, sync_gen, None)
|
while not finished:
|
||||||
if token is None:
|
for _idx, token in result.pop_all():
|
||||||
break
|
if token is STOP:
|
||||||
yield token
|
finished = True
|
||||||
|
break
|
||||||
|
yield token
|
||||||
|
if not finished:
|
||||||
|
await asyncio.to_thread(result.wait, 0.05)
|
||||||
|
finally:
|
||||||
|
if not finished:
|
||||||
|
self.scheduler.cancel_task(task_id)
|
||||||
|
|
||||||
return _agen()
|
return _agen()
|
||||||
|
|
||||||
@@ -198,10 +206,8 @@ class InferenceEngine:
|
|||||||
result.wait_completion()
|
result.wait_completion()
|
||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
for tid in task_ids:
|
for tid in task_ids:
|
||||||
self.scheduler.remove_task(tid)
|
self.scheduler.cancel_task(tid)
|
||||||
raise
|
raise
|
||||||
for tid in task_ids:
|
|
||||||
self.scheduler.remove_task(tid)
|
|
||||||
res = result.get_results()
|
res = result.get_results()
|
||||||
return res if is_batch else res[0]
|
return res if is_batch else res[0]
|
||||||
|
|
||||||
@@ -210,17 +216,22 @@ class InferenceEngine:
|
|||||||
|
|
||||||
def gen():
|
def gen():
|
||||||
nonlocal remaining
|
nonlocal remaining
|
||||||
while remaining > 0:
|
try:
|
||||||
items = result.pop_all()
|
while remaining > 0:
|
||||||
for idx, token in items:
|
items = result.pop_all()
|
||||||
if token is STOP:
|
for idx, token in items:
|
||||||
if not finished[idx]:
|
if token is STOP:
|
||||||
finished[idx] = True
|
if not finished[idx]:
|
||||||
remaining -= 1
|
finished[idx] = True
|
||||||
else:
|
remaining -= 1
|
||||||
yield (idx, token) if is_batch else token
|
else:
|
||||||
if remaining > 0:
|
yield (idx, token) if is_batch else token
|
||||||
result.wait(timeout=0.05)
|
if remaining > 0:
|
||||||
|
result.wait(timeout=0.05)
|
||||||
|
finally:
|
||||||
|
for idx, task_id in enumerate(task_ids):
|
||||||
|
if not finished[idx]:
|
||||||
|
self.scheduler.cancel_task(task_id)
|
||||||
|
|
||||||
return gen()
|
return gen()
|
||||||
|
|
||||||
|
|||||||
+41
-33
@@ -1,5 +1,6 @@
|
|||||||
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
|
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
|
||||||
|
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
@@ -125,6 +126,7 @@ class MetricsCollector:
|
|||||||
def __init__(self, max_recent: int = 128):
|
def __init__(self, max_recent: int = 128):
|
||||||
self._timings: Dict[str, TaskTiming] = {}
|
self._timings: Dict[str, TaskTiming] = {}
|
||||||
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
|
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
self._ttft_ms_sum = 0.0
|
self._ttft_ms_sum = 0.0
|
||||||
self._ttft_ms_count = 0
|
self._ttft_ms_count = 0
|
||||||
@@ -135,18 +137,22 @@ class MetricsCollector:
|
|||||||
|
|
||||||
def register(self, task_id: str):
|
def register(self, task_id: str):
|
||||||
"""Create a timing record for a newly-created task."""
|
"""Create a timing record for a newly-created task."""
|
||||||
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time())
|
with self._lock:
|
||||||
|
self._timings[task_id] = TaskTiming(
|
||||||
|
task_id=task_id, arrival_time=time.time()
|
||||||
|
)
|
||||||
|
|
||||||
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
|
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
|
||||||
"""Close timing for a finished/aborted task and move it to completed."""
|
"""Close timing for a finished/aborted task and move it to completed."""
|
||||||
timing = self._timings.pop(task_id, None)
|
with self._lock:
|
||||||
if timing is None:
|
timing = self._timings.pop(task_id, None)
|
||||||
return
|
if timing is None:
|
||||||
timing.finish_time = time.time()
|
return
|
||||||
timing.input_tokens = input_tokens
|
timing.finish_time = time.time()
|
||||||
timing.output_tokens = output_tokens
|
timing.input_tokens = input_tokens
|
||||||
self._completed.append(timing)
|
timing.output_tokens = output_tokens
|
||||||
self._accumulate(timing)
|
self._completed.append(timing)
|
||||||
|
self._accumulate(timing)
|
||||||
|
|
||||||
# timing scopes
|
# timing scopes
|
||||||
|
|
||||||
@@ -158,34 +164,36 @@ class MetricsCollector:
|
|||||||
yield
|
yield
|
||||||
toc = time.time()
|
toc = time.time()
|
||||||
dt = toc - tic
|
dt = toc - tic
|
||||||
for tid in task_ids:
|
with self._lock:
|
||||||
t = self._timings.get(tid)
|
for tid in task_ids:
|
||||||
if t is None:
|
t = self._timings.get(tid)
|
||||||
continue
|
if t is None:
|
||||||
if phase == "prefill":
|
continue
|
||||||
t.prefill_start_time = tic
|
if phase == "prefill":
|
||||||
t.first_token_time = toc
|
t.prefill_start_time = tic
|
||||||
elif phase == "decode":
|
t.first_token_time = toc
|
||||||
t._decode_steps += 1
|
elif phase == "decode":
|
||||||
t._decode_total_s += dt
|
t._decode_steps += 1
|
||||||
|
t._decode_total_s += dt
|
||||||
|
|
||||||
# aggregate stats
|
# aggregate stats
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
stats: Dict[str, Any] = {}
|
with self._lock:
|
||||||
if self._ttft_ms_count > 0:
|
stats: Dict[str, Any] = {"in_flight_tasks": len(self._timings)}
|
||||||
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
|
if self._ttft_ms_count > 0:
|
||||||
if self._decode_tps_count > 0:
|
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
|
||||||
stats["avg_decode_tps"] = round(
|
if self._decode_tps_count > 0:
|
||||||
self._decode_tps_sum / self._decode_tps_count, 2
|
stats["avg_decode_tps"] = round(
|
||||||
)
|
self._decode_tps_sum / self._decode_tps_count, 2
|
||||||
if self._e2e_ms_count > 0:
|
)
|
||||||
stats["avg_e2e_latency_ms"] = round(
|
if self._e2e_ms_count > 0:
|
||||||
self._e2e_ms_sum / self._e2e_ms_count, 2
|
stats["avg_e2e_latency_ms"] = round(
|
||||||
)
|
self._e2e_ms_sum / self._e2e_ms_count, 2
|
||||||
if self._completed:
|
)
|
||||||
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
|
if self._completed:
|
||||||
return stats
|
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
|
||||||
|
return stats
|
||||||
|
|
||||||
# internal
|
# internal
|
||||||
|
|
||||||
|
|||||||
@@ -146,25 +146,28 @@ class ProtocolHandler:
|
|||||||
yielded = ""
|
yielded = ""
|
||||||
matched = None
|
matched = None
|
||||||
token_ids: List[int] = []
|
token_ids: List[int] = []
|
||||||
async for token in agen:
|
try:
|
||||||
body += token
|
async for token in agen:
|
||||||
|
body += token
|
||||||
|
|
||||||
new_ids = self.engine.tokenizer.encode(token)
|
new_ids = self.engine.tokenizer.encode(token)
|
||||||
token_ids.extend(new_ids)
|
token_ids.extend(new_ids)
|
||||||
|
|
||||||
matched = checker.check(body)
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
ctx.completion_tokens += 1
|
ctx.completion_tokens += 1
|
||||||
for event in self.builder.format_chunk(
|
for event in self.builder.format_chunk(
|
||||||
token,
|
token,
|
||||||
body=body,
|
body=body,
|
||||||
current_token_ids=token_ids,
|
current_token_ids=token_ids,
|
||||||
delta_token_ids=new_ids,
|
delta_token_ids=new_ids,
|
||||||
):
|
):
|
||||||
yield event
|
yield event
|
||||||
yielded += token
|
yielded += token
|
||||||
|
finally:
|
||||||
|
await agen.aclose()
|
||||||
|
|
||||||
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||||
for event in self.builder.format_stream_end(ctx, stop):
|
for event in self.builder.format_stream_end(ctx, stop):
|
||||||
@@ -184,14 +187,17 @@ class ProtocolHandler:
|
|||||||
body = ""
|
body = ""
|
||||||
matched = None
|
matched = None
|
||||||
|
|
||||||
async for token in agen:
|
try:
|
||||||
body += token
|
async for token in agen:
|
||||||
|
body += token
|
||||||
|
|
||||||
matched = checker.check(body)
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
ctx.completion_tokens += 1
|
ctx.completion_tokens += 1
|
||||||
|
finally:
|
||||||
|
await agen.aclose()
|
||||||
|
|
||||||
stop = StopInfo(matched=matched, body=body)
|
stop = StopInfo(matched=matched, body=body)
|
||||||
return self.builder.format_response(ctx, body, stop)
|
return self.builder.format_response(ctx, body, stop)
|
||||||
|
|||||||
@@ -107,12 +107,25 @@ class InferenceScheduler:
|
|||||||
def add_task(self, prompt: str, **kwargs) -> str:
|
def add_task(self, prompt: str, **kwargs) -> str:
|
||||||
return self._task_mgr.add_task(prompt, **kwargs)
|
return self._task_mgr.add_task(prompt, **kwargs)
|
||||||
|
|
||||||
def remove_task(self, task_id: str):
|
def cancel_task(self, task_id: str) -> bool:
|
||||||
for task in self._task_mgr.remove_task(task_id):
|
"""Cancel a waiting or active task without freeing in-use KV state."""
|
||||||
self._task_cache.task_free(task.task_id)
|
immediate, cancelled = self._task_mgr.cancel_task(task_id)
|
||||||
|
for task in immediate:
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
task.task_id, task.input_tokens, task.output_tokens
|
||||||
|
)
|
||||||
|
if cancelled:
|
||||||
|
self._task_mgr.wake()
|
||||||
|
return cancelled
|
||||||
|
|
||||||
|
def remove_task(self, task_id: str) -> bool:
|
||||||
|
"""Backward-compatible alias for cancellation."""
|
||||||
|
return self.cancel_task(task_id)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self._task_mgr.get_stats()
|
stats = self._task_mgr.get_stats()
|
||||||
|
stats["kv_cache_tasks"] = self._task_cache.task_count
|
||||||
|
return stats
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def backend_name(self) -> str:
|
def backend_name(self) -> str:
|
||||||
@@ -254,7 +267,13 @@ class InferenceScheduler:
|
|||||||
if self._task_cache.task_alloc(
|
if self._task_cache.task_alloc(
|
||||||
task.task_id, task.prompt_ids
|
task.task_id, task.prompt_ids
|
||||||
):
|
):
|
||||||
self._task_mgr.activate(task)
|
if not self._task_mgr.activate(task):
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
task.task_id,
|
||||||
|
task.input_tokens,
|
||||||
|
task.output_tokens,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
failed.append(task)
|
failed.append(task)
|
||||||
if failed:
|
if failed:
|
||||||
@@ -264,7 +283,11 @@ class InferenceScheduler:
|
|||||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
active = [
|
||||||
|
task
|
||||||
|
for task in self._task_mgr.get_active_tasks()
|
||||||
|
if task.status != TaskStatus.ABORTED
|
||||||
|
]
|
||||||
|
|
||||||
decoded, aborted = self._step(active)
|
decoded, aborted = self._step(active)
|
||||||
|
|
||||||
@@ -272,6 +295,8 @@ class InferenceScheduler:
|
|||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
for t in decoded:
|
for t in decoded:
|
||||||
|
if t.status == TaskStatus.ABORTED:
|
||||||
|
continue
|
||||||
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
if new_text:
|
if new_text:
|
||||||
self._task_mgr.invoke_callback(t.task_id, new_text)
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
@@ -303,13 +328,21 @@ class InferenceScheduler:
|
|||||||
|
|
||||||
def _abort_and_clear(self, free_waiting: bool):
|
def _abort_and_clear(self, free_waiting: bool):
|
||||||
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
|
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
|
||||||
for task in self._task_mgr.get_active_tasks():
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
waiting = self._task_mgr.get_waiting_tasks()
|
||||||
|
for task in active:
|
||||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
self._task_cache.task_free(task.task_id)
|
self._task_cache.task_free(task.task_id)
|
||||||
for task in self._task_mgr.get_waiting_tasks():
|
self._metrics.mark_finished(
|
||||||
|
task.task_id, task.input_tokens, task.output_tokens
|
||||||
|
)
|
||||||
|
for task in waiting:
|
||||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
if free_waiting:
|
if free_waiting:
|
||||||
self._task_cache.task_free(task.task_id)
|
self._task_cache.task_free(task.task_id)
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
task.task_id, task.input_tokens, task.output_tokens
|
||||||
|
)
|
||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
|
|
||||||
def run_batch(
|
def run_batch(
|
||||||
|
|||||||
+81
-26
@@ -4,7 +4,17 @@ import uuid
|
|||||||
from collections import deque
|
from collections import deque
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Literal, Optional
|
from typing import (
|
||||||
|
TYPE_CHECKING,
|
||||||
|
Any,
|
||||||
|
Callable,
|
||||||
|
Deque,
|
||||||
|
Dict,
|
||||||
|
List,
|
||||||
|
Literal,
|
||||||
|
Optional,
|
||||||
|
Tuple,
|
||||||
|
)
|
||||||
|
|
||||||
from tokenizers.decoders import DecodeStream
|
from tokenizers.decoders import DecodeStream
|
||||||
|
|
||||||
@@ -150,12 +160,14 @@ class TaskManager:
|
|||||||
self.waiting_queue: Deque[Task] = deque()
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
self.active_tasks: List[Task] = []
|
self.active_tasks: List[Task] = []
|
||||||
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||||
|
self._tasks: Dict[str, Task] = {}
|
||||||
|
|
||||||
self._task_event = threading.Event()
|
self._task_event = threading.Event()
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
self._total_tasks = 0
|
self._total_tasks = 0
|
||||||
self._total_tokens = 0
|
self._total_tokens = 0
|
||||||
|
self._cancelled_total = 0
|
||||||
|
|
||||||
self._metrics = metrics
|
self._metrics = metrics
|
||||||
|
|
||||||
@@ -195,6 +207,7 @@ class TaskManager:
|
|||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.append(task)
|
self.waiting_queue.append(task)
|
||||||
|
self._tasks[task_id] = task
|
||||||
self._total_tasks += 1
|
self._total_tasks += 1
|
||||||
if stream_callback:
|
if stream_callback:
|
||||||
self._callbacks[task_id] = stream_callback
|
self._callbacks[task_id] = stream_callback
|
||||||
@@ -205,28 +218,49 @@ class TaskManager:
|
|||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
return task_id
|
return task_id
|
||||||
|
|
||||||
def remove_task(self, task_id: str) -> List[Task]:
|
def cancel_task(self, task_id: str) -> Tuple[List[Task], bool]:
|
||||||
|
"""Mark a task cancelled and return tasks safe to clean immediately."""
|
||||||
with self._lock:
|
with self._lock:
|
||||||
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
task = self._tasks.get(task_id)
|
||||||
self.waiting_queue = deque(
|
|
||||||
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]
|
|
||||||
self._callbacks.pop(task_id, None)
|
self._callbacks.pop(task_id, None)
|
||||||
return removed_active
|
if task is None or task.status in (
|
||||||
|
TaskStatus.FINISHED,
|
||||||
|
TaskStatus.ABORTED,
|
||||||
|
):
|
||||||
|
return [], False
|
||||||
|
|
||||||
|
task.status = TaskStatus.ABORTED
|
||||||
|
self._cancelled_total += 1
|
||||||
|
if task in self.waiting_queue:
|
||||||
|
self.waiting_queue = deque(
|
||||||
|
waiting for waiting in self.waiting_queue if waiting is not task
|
||||||
|
)
|
||||||
|
self._tasks.pop(task_id, None)
|
||||||
|
return [task], True
|
||||||
|
return [], True
|
||||||
|
|
||||||
|
def remove_task(self, task_id: str) -> List[Task]:
|
||||||
|
"""Backward-compatible alias for cancellation."""
|
||||||
|
immediate, _ = self.cancel_task(task_id)
|
||||||
|
return immediate
|
||||||
|
|
||||||
def invoke_callback(self, task_id: str, token: str):
|
def invoke_callback(self, task_id: str, token: str):
|
||||||
cb = self._callbacks.get(task_id)
|
with self._lock:
|
||||||
|
cb = self._callbacks.get(task_id)
|
||||||
if cb:
|
if cb:
|
||||||
cb(token)
|
cb(token)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
stats: Dict[str, Any] = {
|
with self._lock:
|
||||||
"total_tasks": self._total_tasks,
|
waiting = len(self.waiting_queue)
|
||||||
"total_tokens": self._total_tokens,
|
stats: Dict[str, Any] = {
|
||||||
"active_tasks": len(self.active_tasks),
|
"total_tasks": self._total_tasks,
|
||||||
"waiting_queue": len(self.waiting_queue),
|
"total_tokens": self._total_tokens,
|
||||||
}
|
"active_tasks": len(self.active_tasks),
|
||||||
|
"waiting_tasks": waiting,
|
||||||
|
"waiting_queue": waiting,
|
||||||
|
"cancelled_total": self._cancelled_total,
|
||||||
|
}
|
||||||
if self._metrics is not None:
|
if self._metrics is not None:
|
||||||
stats.update(self._metrics.get_stats())
|
stats.update(self._metrics.get_stats())
|
||||||
return stats
|
return stats
|
||||||
@@ -242,18 +276,21 @@ class TaskManager:
|
|||||||
finished.append(task)
|
finished.append(task)
|
||||||
self._total_tokens += task.output_tokens
|
self._total_tokens += task.output_tokens
|
||||||
|
|
||||||
if self._metrics is not None:
|
|
||||||
for task in finished:
|
|
||||||
self._metrics.mark_finished(
|
|
||||||
task.task_id, task.input_tokens, task.output_tokens
|
|
||||||
)
|
|
||||||
|
|
||||||
self.active_tasks = [
|
self.active_tasks = [
|
||||||
t
|
t
|
||||||
for t in self.active_tasks
|
for t in self.active_tasks
|
||||||
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
|
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
|
||||||
]
|
]
|
||||||
return finished
|
for task in finished:
|
||||||
|
self._tasks.pop(task.task_id, None)
|
||||||
|
self._callbacks.pop(task.task_id, None)
|
||||||
|
|
||||||
|
if self._metrics is not None:
|
||||||
|
for task in finished:
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
task.task_id, task.input_tokens, task.output_tokens
|
||||||
|
)
|
||||||
|
return finished
|
||||||
|
|
||||||
def pull_candidates(self, n: int) -> List[Task]:
|
def pull_candidates(self, n: int) -> List[Task]:
|
||||||
to_add: List[Task] = []
|
to_add: List[Task] = []
|
||||||
@@ -263,18 +300,35 @@ class TaskManager:
|
|||||||
to_add.append(self.waiting_queue.popleft())
|
to_add.append(self.waiting_queue.popleft())
|
||||||
return to_add
|
return to_add
|
||||||
|
|
||||||
def activate(self, task: Task):
|
def activate(self, task: Task) -> bool:
|
||||||
task.status = TaskStatus.RUNNING
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
|
if task.status == TaskStatus.ABORTED:
|
||||||
|
self._tasks.pop(task.task_id, None)
|
||||||
|
self._callbacks.pop(task.task_id, None)
|
||||||
|
return False
|
||||||
|
task.status = TaskStatus.RUNNING
|
||||||
self.active_tasks.append(task)
|
self.active_tasks.append(task)
|
||||||
|
return True
|
||||||
|
|
||||||
def return_to_waiting(self, tasks: List[Task]):
|
def return_to_waiting(self, tasks: List[Task]):
|
||||||
|
cancelled = []
|
||||||
with self._lock:
|
with self._lock:
|
||||||
for task in reversed(tasks):
|
for task in reversed(tasks):
|
||||||
self.waiting_queue.appendleft(task)
|
if task.status == TaskStatus.ABORTED:
|
||||||
|
self._tasks.pop(task.task_id, None)
|
||||||
|
self._callbacks.pop(task.task_id, None)
|
||||||
|
cancelled.append(task)
|
||||||
|
else:
|
||||||
|
self.waiting_queue.appendleft(task)
|
||||||
|
if self._metrics is not None:
|
||||||
|
for task in cancelled:
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
task.task_id, task.input_tokens, task.output_tokens
|
||||||
|
)
|
||||||
|
|
||||||
def has_work(self) -> bool:
|
def has_work(self) -> bool:
|
||||||
return bool(self.active_tasks or self.waiting_queue)
|
with self._lock:
|
||||||
|
return bool(self.active_tasks or self.waiting_queue)
|
||||||
|
|
||||||
def wait_for_tasks(self, timeout: float = 1.0):
|
def wait_for_tasks(self, timeout: float = 1.0):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -296,6 +350,7 @@ class TaskManager:
|
|||||||
self.waiting_queue.clear()
|
self.waiting_queue.clear()
|
||||||
self.active_tasks.clear()
|
self.active_tasks.clear()
|
||||||
self._callbacks.clear()
|
self._callbacks.clear()
|
||||||
|
self._tasks.clear()
|
||||||
|
|
||||||
def wake(self):
|
def wake(self):
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
|
|||||||
@@ -157,6 +157,28 @@ def test_engine_generate_streaming_yields_tokens():
|
|||||||
assert tokens == ["t1", "t2"]
|
assert tokens == ["t1", "t2"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_engine_stream_close_cancels_unfinished_task():
|
||||||
|
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
|
||||||
|
callbacks_saved = []
|
||||||
|
|
||||||
|
def capture_cb(prompt, **kwargs):
|
||||||
|
callbacks_saved.append(kwargs["stream_callback"])
|
||||||
|
return "task-1"
|
||||||
|
|
||||||
|
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
|
||||||
|
instance = MockSched.return_value
|
||||||
|
instance.add_task.side_effect = capture_cb
|
||||||
|
|
||||||
|
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
|
||||||
|
stream = engine.generate("hello", stream=True)
|
||||||
|
|
||||||
|
callbacks_saved[0]("t1")
|
||||||
|
assert next(stream) == "t1"
|
||||||
|
stream.close()
|
||||||
|
|
||||||
|
instance.cancel_task.assert_called_once_with("task-1")
|
||||||
|
|
||||||
|
|
||||||
def test_engine_generate_async_yields_tokens_until_stop():
|
def test_engine_generate_async_yields_tokens_until_stop():
|
||||||
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
|
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
|
||||||
callbacks_saved = []
|
callbacks_saved = []
|
||||||
@@ -186,6 +208,31 @@ def test_engine_generate_async_yields_tokens_until_stop():
|
|||||||
assert asyncio.run(collect()) == ["t1", "t2"]
|
assert asyncio.run(collect()) == ["t1", "t2"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_engine_async_close_cancels_unfinished_task():
|
||||||
|
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
|
||||||
|
callbacks_saved = []
|
||||||
|
|
||||||
|
def capture_cb(prompt, **kwargs):
|
||||||
|
callbacks_saved.append(kwargs["stream_callback"])
|
||||||
|
return "task-1"
|
||||||
|
|
||||||
|
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
|
||||||
|
instance = MockSched.return_value
|
||||||
|
instance.add_task.side_effect = capture_cb
|
||||||
|
|
||||||
|
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
|
||||||
|
stream = engine.generate_async("hello")
|
||||||
|
callbacks_saved[0]("t1")
|
||||||
|
|
||||||
|
async def consume_then_close():
|
||||||
|
assert await anext(stream) == "t1"
|
||||||
|
await stream.aclose()
|
||||||
|
|
||||||
|
asyncio.run(consume_then_close())
|
||||||
|
|
||||||
|
instance.cancel_task.assert_called_once_with("task-1")
|
||||||
|
|
||||||
|
|
||||||
def test_engine_generate_non_streaming_batch():
|
def test_engine_generate_non_streaming_batch():
|
||||||
mock_model, mock_tokenizer = _make_engine_mocks(decode="r")
|
mock_model, mock_tokenizer = _make_engine_mocks(decode="r")
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Tests for scheduler concurrency."""
|
"""Tests for scheduler concurrency."""
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
@@ -259,6 +260,104 @@ def _make_real_scheduler(device):
|
|||||||
return scheduler, tokenizer, model
|
return scheduler, tokenizer, model
|
||||||
|
|
||||||
|
|
||||||
|
def test_cancel_waiting_task_storm_returns_to_baseline(device):
|
||||||
|
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||||
|
try:
|
||||||
|
task_ids = [
|
||||||
|
scheduler.add_task(f"waiting-{index}", max_tokens=32) for index in range(32)
|
||||||
|
]
|
||||||
|
|
||||||
|
assert all(scheduler.cancel_task(task_id) for task_id in task_ids)
|
||||||
|
stats = scheduler.get_stats()
|
||||||
|
assert stats["active_tasks"] == 0
|
||||||
|
assert stats["waiting_tasks"] == 0
|
||||||
|
assert stats["in_flight_tasks"] == 0
|
||||||
|
assert stats["kv_cache_tasks"] == 0
|
||||||
|
assert stats["cancelled_total"] == len(task_ids)
|
||||||
|
finally:
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
|
def test_cancel_active_task_releases_metrics_and_kv(device):
|
||||||
|
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||||
|
try:
|
||||||
|
task_id = scheduler.add_task("active", max_tokens=32)
|
||||||
|
task = scheduler._task_mgr.pull_candidates(1)[0]
|
||||||
|
assert scheduler._task_cache.task_alloc(task.task_id, task.prompt_ids)
|
||||||
|
assert scheduler._task_mgr.activate(task)
|
||||||
|
|
||||||
|
before = scheduler.get_stats()
|
||||||
|
assert before["active_tasks"] == 1
|
||||||
|
assert before["in_flight_tasks"] == 1
|
||||||
|
assert before["kv_cache_tasks"] == 1
|
||||||
|
|
||||||
|
assert scheduler.cancel_task(task_id)
|
||||||
|
scheduler.start()
|
||||||
|
deadline = time.monotonic() + 5
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
after = scheduler.get_stats()
|
||||||
|
if (
|
||||||
|
after["active_tasks"] == 0
|
||||||
|
and after["in_flight_tasks"] == 0
|
||||||
|
and after["kv_cache_tasks"] == 0
|
||||||
|
):
|
||||||
|
break
|
||||||
|
time.sleep(0.01)
|
||||||
|
|
||||||
|
assert after["active_tasks"] == 0
|
||||||
|
assert after["waiting_tasks"] == 0
|
||||||
|
assert after["in_flight_tasks"] == 0
|
||||||
|
assert after["kv_cache_tasks"] == 0
|
||||||
|
assert after["cancelled_total"] == 1
|
||||||
|
finally:
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
|
def test_cancel_during_kv_allocation_releases_metrics_and_kv(device):
|
||||||
|
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||||
|
allocation_started = threading.Event()
|
||||||
|
continue_allocation = threading.Event()
|
||||||
|
original_alloc = scheduler._task_cache.task_alloc
|
||||||
|
|
||||||
|
def blocking_alloc(*args, **kwargs):
|
||||||
|
allocation_started.set()
|
||||||
|
assert continue_allocation.wait(timeout=5)
|
||||||
|
return original_alloc(*args, **kwargs)
|
||||||
|
|
||||||
|
try:
|
||||||
|
with patch.object(
|
||||||
|
scheduler._task_cache,
|
||||||
|
"task_alloc",
|
||||||
|
side_effect=blocking_alloc,
|
||||||
|
):
|
||||||
|
scheduler.start()
|
||||||
|
task_id = scheduler.add_task("allocation-race", max_tokens=32)
|
||||||
|
assert allocation_started.wait(timeout=5)
|
||||||
|
assert scheduler.cancel_task(task_id)
|
||||||
|
continue_allocation.set()
|
||||||
|
|
||||||
|
deadline = time.monotonic() + 5
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
stats = scheduler.get_stats()
|
||||||
|
if (
|
||||||
|
stats["active_tasks"] == 0
|
||||||
|
and stats["waiting_tasks"] == 0
|
||||||
|
and stats["in_flight_tasks"] == 0
|
||||||
|
and stats["kv_cache_tasks"] == 0
|
||||||
|
):
|
||||||
|
break
|
||||||
|
time.sleep(0.01)
|
||||||
|
|
||||||
|
assert stats["active_tasks"] == 0
|
||||||
|
assert stats["waiting_tasks"] == 0
|
||||||
|
assert stats["in_flight_tasks"] == 0
|
||||||
|
assert stats["kv_cache_tasks"] == 0
|
||||||
|
assert stats["cancelled_total"] == 1
|
||||||
|
finally:
|
||||||
|
continue_allocation.set()
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
def test_run_batch_returns_token_sequences(device):
|
def test_run_batch_returns_token_sequences(device):
|
||||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -166,11 +166,15 @@ def test_messages_with_system(client, loaded_model):
|
|||||||
|
|
||||||
def test_chat_completions_stop_sequence(client, loaded_model):
|
def test_chat_completions_stop_sequence(client, loaded_model):
|
||||||
"""POST /v1/chat/completions with stop parameter truncates at stop sequence."""
|
"""POST /v1/chat/completions with stop parameter truncates at stop sequence."""
|
||||||
|
closed = []
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "Hello"
|
try:
|
||||||
yield "X"
|
yield "Hello"
|
||||||
yield "world"
|
yield "X"
|
||||||
|
yield "world"
|
||||||
|
finally:
|
||||||
|
closed.append(True)
|
||||||
|
|
||||||
get_app().state.engine = loaded_model
|
get_app().state.engine = loaded_model
|
||||||
loaded_model.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
@@ -188,15 +192,20 @@ def test_chat_completions_stop_sequence(client, loaded_model):
|
|||||||
content = data["choices"][0]["message"]["content"]
|
content = data["choices"][0]["message"]["content"]
|
||||||
assert "X" in content
|
assert "X" in content
|
||||||
assert "world" not in content
|
assert "world" not in content
|
||||||
|
assert closed == [True]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
||||||
"""POST /v1/chat/completions with stop parameter truncates SSE stream."""
|
"""POST /v1/chat/completions with stop parameter truncates SSE stream."""
|
||||||
|
closed = []
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "Hello"
|
try:
|
||||||
yield "X"
|
yield "Hello"
|
||||||
yield "world"
|
yield "X"
|
||||||
|
yield "world"
|
||||||
|
finally:
|
||||||
|
closed.append(True)
|
||||||
|
|
||||||
get_app().state.engine = loaded_model
|
get_app().state.engine = loaded_model
|
||||||
loaded_model.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
@@ -217,6 +226,7 @@ def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
|||||||
assert any(
|
assert any(
|
||||||
"finish_reason" in line for line in content.split("\n") if "stop" in line
|
"finish_reason" in line for line in content.split("\n") if "stop" in line
|
||||||
)
|
)
|
||||||
|
assert closed == [True]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_completions_real_engine(tmp_path, client):
|
def test_chat_completions_real_engine(tmp_path, client):
|
||||||
|
|||||||
@@ -73,15 +73,21 @@ def test_task_manager_remove_task():
|
|||||||
assert len(tm.waiting_queue) == 0
|
assert len(tm.waiting_queue) == 0
|
||||||
|
|
||||||
|
|
||||||
def test_task_manager_remove_active_task():
|
def test_task_manager_cancel_active_task_defers_removal():
|
||||||
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
tid = tm.add_task("test")
|
tid = tm.add_task("test")
|
||||||
tasks = tm.pull_candidates(1)
|
tasks = tm.pull_candidates(1)
|
||||||
tm.activate(tasks[0])
|
tm.activate(tasks[0])
|
||||||
assert len(tm.active_tasks) == 1
|
assert len(tm.active_tasks) == 1
|
||||||
removed = tm.remove_task(tid)
|
immediate, cancelled = tm.cancel_task(tid)
|
||||||
assert len(removed) == 1
|
assert cancelled is True
|
||||||
|
assert immediate == []
|
||||||
|
assert tm.active_tasks[0].status == TaskStatus.ABORTED
|
||||||
|
|
||||||
|
removed = tm.remove_finished_tasks([])
|
||||||
|
assert removed == tasks
|
||||||
assert len(tm.active_tasks) == 0
|
assert len(tm.active_tasks) == 0
|
||||||
|
assert tm.get_stats()["cancelled_total"] == 1
|
||||||
|
|
||||||
|
|
||||||
def test_task_manager_pull_candidates_fifo():
|
def test_task_manager_pull_candidates_fifo():
|
||||||
|
|||||||
Reference in New Issue
Block a user