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:
@@ -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, :]
|
||||
|
||||
Reference in New Issue
Block a user