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
+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]