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
+21 -26
View File
@@ -138,36 +138,31 @@ class InferenceScheduler:
t.task_id, t.prompt_ids, start_logical_page
)
pos_groups: Dict[int, List[Task]] = {}
for t in self._task_mgr.get_active_tasks():
pos_groups.setdefault(t.next_pos, []).append(t)
decode_tasks = self._task_mgr.get_active_tasks()
for next_pos in sorted(pos_groups.keys()):
group = sorted(pos_groups[next_pos], key=lambda t: t.task_id)
valid: List[Task] = []
for t in sorted(decode_tasks, key=lambda t: t.task_id):
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
self._task_mgr.invoke_callback(t.task_id, STOP)
valid: List[Task] = []
for t in group:
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
if valid:
next_tokens = self._executor.execute_decode(valid)
for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok)
t.output_tokens += 1
self._task_mgr.invoke_callback(
t.task_id,
self._task_mgr.tokenizer.decode([ntok]),
)
for t in valid:
if t.is_finished(stop_ids):
self._task_mgr.invoke_callback(t.task_id, STOP)
if valid:
next_tokens = self._executor.execute_decode(valid)
for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok)
t.output_tokens += 1
self._task_mgr.invoke_callback(
t.task_id,
self._task_mgr.tokenizer.decode([ntok]),
)
for t in valid:
if t.is_finished(stop_ids):
self._task_mgr.invoke_callback(t.task_id, STOP)
except Exception as e:
self._stop_event.set()
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)