refactor: remove inference redundancy and fix cache leaks

- drop Executor unused tokenizer field, _head_dim, stale metrics docstring
- unify greedy sampling via SamplingPipeline.sample, drop top-level duplicate
- drop Task.flush_remaining no-op and unreachable prompt-length branch
- drop ProtocolHandler redundant chunks list (reuse body)
- fix page_size=1 token-slot leak on task_free
- clear _task_pages/_task_slots on alloc-failure paths
- reset _bind_state on task_free to avoid stale steady-state reuse
- remove unreachable contiguous branches in paged-only helpers
This commit is contained in:
2026-08-08 21:45:01 +08:00
parent d9240ab149
commit ca50fe4721
7 changed files with 12 additions and 52 deletions
+1 -4
View File
@@ -181,12 +181,10 @@ class ProtocolHandler:
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
) -> Dict[str, Any]:
checker = StopChecker(stop_sequences)
chunks: List[str] = []
body = ""
matched = None
async for token in agen:
chunks.append(token)
body += token
matched = checker.check(body)
@@ -195,6 +193,5 @@ class ProtocolHandler:
ctx.completion_tokens += 1
content = "".join(chunks)
stop = StopInfo(matched=matched, body=body)
return self.builder.format_response(ctx, content, stop)
return self.builder.format_response(ctx, body, stop)
+9 -4
View File
@@ -413,6 +413,8 @@ class PagePool:
if slots is None:
for p in self._task_pages[task_id]:
self._alloc.free(p)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
@@ -427,6 +429,8 @@ class PagePool:
self._alloc.free(hp)
for np_ in new_pages:
self._alloc.free(np_)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
@@ -442,6 +446,7 @@ class PagePool:
req_idx = self._task_req.pop(task_id, None)
if req_idx is None:
return
self._bind_state = None
self._task_len.pop(req_idx, None)
self._task_cached.pop(task_id, None)
@@ -455,6 +460,9 @@ class PagePool:
else:
for p in self._task_pages.get(task_id, []):
self._alloc.free(p)
if self.page_size == 1:
for slot in self._task_slots.get(task_id, []):
self._alloc.free(slot)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
@@ -503,7 +511,7 @@ class PagePool:
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
if self._prefix is None or self.contiguous:
if self._prefix is None:
return
pages = self._task_pages.get(task_id, [])
full_pages = len(prompt_ids) // self.page_size
@@ -635,9 +643,6 @@ class PagePool:
req_idx = self._task_req[task_id]
total = len(prompt_ids)
if self.contiguous:
return
if self.page_size == 1:
slots = self._task_slots.get(task_id, [])
all_slots = slots[: total - cached]
-4
View File
@@ -19,7 +19,6 @@ from astrai.inference.core.task import Task
from astrai.inference.core.workspace import InferenceWorkspace
from astrai.inference.sample import sample
from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
@@ -184,13 +183,11 @@ class Executor:
def __init__(
self,
model: AutoModel,
tokenizer: AutoTokenizer,
kv_cache: PagePool,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
):
self.model = model
self.tokenizer = tokenizer
self.kv_cache = kv_cache
self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype
@@ -207,7 +204,6 @@ class Executor:
config = model.config
max_q_heads = config.num_attention_heads
head_dim = config.hidden_size // config.num_attention_heads
self._head_dim = head_dim
self._graph_supported = CudaBackend.supports(head_dim=head_dim)
self._workspace = InferenceWorkspace(
max_batch_size=kv_cache.max_batch_size,
+2 -4
View File
@@ -12,7 +12,7 @@ class TaskTiming:
"""Timestamp snapshots and computed metrics for one generation task.
Created by :class:`MetricsCollector` at task-registration time;
updated via ``prefill_scope`` / ``mark_finished``.
updated via ``record`` / ``mark_finished``.
"""
task_id: str
@@ -27,8 +27,6 @@ class TaskTiming:
# derived metrics
# derived metrics
@property
def queue_wait_ms(self) -> Optional[float]:
if self.prefill_start_time is not None:
@@ -116,7 +114,7 @@ class MetricsCollector:
metrics = MetricsCollector()
metrics.register(task_id, arrival_time)
with metrics.prefill_scope(task_ids):
with metrics.record(task_ids, "prefill"):
run_prefill(...)
metrics.mark_finished(task_id, input_tokens, output_tokens)
-4
View File
@@ -68,7 +68,6 @@ class InferenceScheduler:
self._executor = Executor(
model=model,
tokenizer=tokenizer,
kv_cache=self._cache,
device=self.device,
dtype=self.dtype,
@@ -210,9 +209,6 @@ class InferenceScheduler:
if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text)
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:
-18
View File
@@ -1,4 +1,3 @@
import logging
import threading
import time
import uuid
@@ -11,8 +10,6 @@ from tokenizers.decoders import DecodeStream
from astrai.inference.core.metrics import MetricsCollector
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
STOP = object()
@@ -92,16 +89,6 @@ class Task:
self._decoder = StreamDecoder(tokenizer)
return self._decoder.push(self.output_ids[-1])
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
"""Emit any text still buffered in the decoder.
With the Rust-native DecodeStream, the stream is always in a
correct state — any completed text was already emitted by the
last ``push``. A trailing incomplete multi-byte sequence has no
valid text to emit, so this is a no-op.
"""
return ""
@property
def next_pos(self) -> int:
# The first output is sampled from prefill and enters KV on the next step.
@@ -157,11 +144,6 @@ class TaskManager:
if len(prompt_ids) > self.max_seq_len:
prompt_ids = prompt_ids[-self.max_seq_len :]
if len(prompt_ids) > self.max_seq_len:
if stream_callback:
stream_callback(STOP)
return task_id
if max_tokens is None:
max_tokens = self.max_seq_len - len(prompt_ids)
else:
-14
View File
@@ -363,20 +363,6 @@ def sample(
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
``chosen_logprobs`` has shape ``[batch]``.
"""
greedy = (
bool((temperature == 0).all())
if isinstance(temperature, Tensor)
else temperature == 0
)
if greedy:
tokens = logits.argmax(dim=-1)
if not return_logprobs:
return tokens
log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
has_freq = (
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
if isinstance(frequency_penalty, Tensor)