perf: replace paged KV cache with contiguous ContiguousCache, decode all groups

- Add KVCache/CacheView abstract base classes in cache.py
- Add ContiguousCache (contiguous per-slot buffer, default) alongside PageCache (paged, renamed from old KVCache)
- Merge make_table_tensor + bind into bind_tasks on KVCache interface
- Remove task_cached/task_record_hashes from base class (PageCache-only)
- Scheduler: decode all position groups instead of just the largest (eliminates 63% group skip rate)
- Scheduler: accept optional cache param for swapping implementations
- Model layer type hints use CacheView base class
- Batch 1-32: 1-7% speedup from eliminating Storage.gather overhead
- All 183 inference tests pass
This commit is contained in:
2026-07-05 11:34:36 +08:00
parent 599a51f4f7
commit 5416c2e8fb
9 changed files with 215 additions and 62 deletions
+6 -4
View File
@@ -43,7 +43,6 @@ class Executor:
)
task_ids = [t.task_id for t in tasks]
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
with torch.inference_mode():
self.model(
@@ -53,7 +52,9 @@ class Executor:
)
.unsqueeze(0)
.expand(batch_sz, -1),
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
paged_cache=self.page_cache.bind_tasks(
task_ids, prompt_len, self.device
),
)
def execute_decode(self, tasks: List[Task]) -> List[int]:
@@ -72,7 +73,6 @@ class Executor:
total_len = position_ids.max().item() + 1
task_ids = [t.task_id for t in tasks]
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
@@ -81,7 +81,9 @@ class Executor:
with torch.inference_mode():
outputs = self.model(
input_ids.unsqueeze(1),
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
paged_cache=self.page_cache.bind_tasks(
task_ids, total_len, self.device
),
position_ids=position_ids.unsqueeze(1),
)
logits = outputs["logits"][:, -1, :]