perf: eliminate GPU syncs in contiguous cache write/gather hot paths

- Replace .tolist() calls with _total_len in gather(); move _slot_len updates from per-layer write to once-per-step bind_tasks

- Use torch.as_tensor instead of torch.tensor in decode penalty history construction
This commit is contained in:
2026-07-21 16:40:56 +08:00
parent f1b4b05d08
commit ccf728a1b7
2 changed files with 11 additions and 22 deletions
+7 -10
View File
@@ -104,28 +104,25 @@ class Executor:
)
history_lists = []
mask_lists = []
history_lens = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
mask_lists.append([True] * len(ids))
history_lens.append(len(ids))
max_len = max(len(h) for h in history_lists)
max_len = max(history_lens) if history_lens else 0
padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device
)
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
padded_ids[i, : len(h)] = torch.tensor(
h, dtype=torch.long, device=self.device
)
padded_mask[i, : len(m)] = torch.tensor(
m, dtype=torch.bool, device=self.device
)
for i, h in enumerate(history_lists):
L = history_lens[i]
padded_ids[i, :L] = torch.as_tensor(h, dtype=torch.long, device=self.device)
padded_mask[i, :L] = True
with torch.inference_mode():
outputs = self.model(