perf: merge decode batch for 10x throughput

- merge all active decode tasks into single forward pass (was grouped by next_pos)
- add per-task write_positions to ContiguousCacheView for correct KV writes
- override ContiguousCache.task_cached (base returned 0, caused prefill loops)
- add --cache_len/--frequency_penalty/--rep_window to generate.py
- chunked batch processing with tqdm progress

bench (1.2B model, 128 prompts, 64 tok, batch=128):
  before: 77.2s, ~111 tok/s
  after:   7.1s, ~1210 tok/s (10.9x)
This commit is contained in:
2026-07-18 08:50:46 +08:00
parent f7df02f9a3
commit a24a7b4da5
4 changed files with 162 additions and 69 deletions
+6 -1
View File
@@ -106,7 +106,12 @@ class Executor:
with torch.inference_mode():
outputs = self.model(
input_ids.unsqueeze(1),
paged_cache=self.kv_cache.bind_tasks(task_ids, total_len, self.device),
paged_cache=self.kv_cache.bind_tasks(
task_ids,
total_len,
self.device,
write_positions=position_ids,
),
position_ids=position_ids.unsqueeze(1),
)
logits = outputs["logits"][:, -1, :]