From 50cfd0d555b2f72180a2109a2f0fe82f556441da Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 31 Jul 2026 14:46:03 +0800 Subject: [PATCH] perf: reduce decode overhead in scheduler and executor - Precompute page_table and decode_mask on KVCache once per step in PagePool.bind_tasks, instead of per-layer in CudaBackend/TorchNativeBackend - Skip frequency penalty history tensor construction when all penalties are 0 in Executor.execute_decode - Omit FrequencyPenaltyStrategy from sampling pipeline when penalty is 0 - Deduplicate get_active_tasks calls in scheduler loop (3 to 1), remove redundant sorted() on decode tasks - Benchmark (L20, bf16, CUDA backend): B=1 9.48->9.40ms (+1%), B=4 10.73->9.89ms (+8.6%), B=8 10.77->10.13ms (+6.4%) --- astrai/extension/attention_backend.py | 31 ++++++++++++------ astrai/inference/core/cache.py | 18 +++++++++++ astrai/inference/core/executor.py | 45 ++++++++++++++++----------- astrai/inference/core/scheduler.py | 8 +++-- astrai/inference/sample.py | 45 ++++++++++++++++++++++----- 5 files changed, 108 insertions(+), 39 deletions(-) diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index df57d1a..2372440 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -272,12 +272,18 @@ class TorchNativeBackend(AttentionBackend): kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v - max_len = kv_cache.seq_lens.max() - indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len] - pos_mask = ( - torch.arange(max_len, device=q.device)[None, :] - < kv_cache.seq_lens[:, None] - ) + max_len = kv_cache.max_len + if kv_cache.page_table is not None: + indices = kv_cache.page_table + else: + indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len] + if kv_cache.decode_mask is not None: + pos_mask = kv_cache.decode_mask + else: + pos_mask = ( + torch.arange(max_len, device=q.device)[None, :] + < kv_cache.seq_lens[:, None] + ) indices = torch.where(pos_mask, indices, torch.zeros_like(indices)) k = kv_cache.k_buffer[layer_id, indices] v = kv_cache.v_buffer[layer_id, indices] @@ -338,18 +344,25 @@ class CudaBackend(AttentionBackend): kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v - seq_lens = kv_cache.seq_lens max_len = kv_cache.max_len - page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len] + if kv_cache.page_table is not None: + page_table = kv_cache.page_table + else: + page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len] k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1) v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1) if q.size(0) == 1: mask = None + elif kv_cache.decode_mask is not None: + mask = kv_cache.decode_mask else: - mask = torch.arange(max_len, device=q.device)[None, :] < seq_lens[:, None] + mask = ( + torch.arange(max_len, device=q.device)[None, :] + < kv_cache.seq_lens[:, None] + ) out = attn_paged_decode( q, diff --git a/astrai/inference/core/cache.py b/astrai/inference/core/cache.py index f5276f4..92cfa2c 100644 --- a/astrai/inference/core/cache.py +++ b/astrai/inference/core/cache.py @@ -203,6 +203,10 @@ class KVCache: seq_lens: [batch_size] — per-request total sequence lengths out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices max_len: max(seq_lens) as Python int — avoids GPU sync in decode + page_table: [batch, max_len] — precomputed gather indices for decode; + None for prefill or when not yet computed. + decode_mask: [batch, max_len] bool — precomputed position validity + mask for decode; None for prefill or single-batch decode. """ k_buffer: Tensor @@ -212,6 +216,8 @@ class KVCache: seq_lens: Tensor out_cache_loc: Tensor max_len: int = 0 + page_table: Optional[Tensor] = None + decode_mask: Optional[Tensor] = None class PagePool: @@ -428,11 +434,21 @@ class PagePool: out_cache_loc = self._req_pool.req_to_token[ req_pool_indices, start_pos:seq_len ] + page_table = None + decode_mask = None else: write_pos = seq_lens_t - 1 out_cache_loc = self._req_pool.req_to_token[ req_pool_indices, write_pos ].unsqueeze(-1) + ml = max(seq_lens) + page_table = self._req_pool.req_to_token[req_pool_indices, :ml] + if len(task_ids) > 1: + decode_mask = ( + torch.arange(ml, device=device)[None, :] < seq_lens_t[:, None] + ) + else: + decode_mask = None return KVCache( k_buffer=self._storage.k_buffer, @@ -442,6 +458,8 @@ class PagePool: seq_lens=seq_lens_t, out_cache_loc=out_cache_loc, max_len=max(seq_lens), + page_table=page_table, + decode_mask=decode_mask, ) # ---- internals ---- diff --git a/astrai/inference/core/executor.py b/astrai/inference/core/executor.py index 8b13903..04a0799 100644 --- a/astrai/inference/core/executor.py +++ b/astrai/inference/core/executor.py @@ -105,26 +105,33 @@ class Executor: [t.frequency_penalty for t in tasks], device=self.device ) - history_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) - history_lens.append(len(ids)) + has_freq = bool((freq_penalties != 0).any()) + if has_freq: + history_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) + history_lens.append(len(ids)) - 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 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 + 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 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 + else: + padded_ids = None + padded_mask = None with torch.inference_mode(): outputs = self.model( diff --git a/astrai/inference/core/scheduler.py b/astrai/inference/core/scheduler.py index dfeadd6..9020252 100644 --- a/astrai/inference/core/scheduler.py +++ b/astrai/inference/core/scheduler.py @@ -109,9 +109,11 @@ class InferenceScheduler: self._task_mgr.wait_for_tasks(timeout=1.0) continue + active = self._task_mgr.get_active_tasks() + to_prefill = [ t - for t in self._task_mgr.get_active_tasks() + for t in active if t.output_tokens == 0 and cache.task_cached(t.task_id) < len(t.prompt_ids) ] @@ -137,10 +139,10 @@ class InferenceScheduler: t.task_id, t.prompt_ids, start_logical_page ) - decode_tasks = self._task_mgr.get_active_tasks() + decode_tasks = active valid: List[Task] = [] - for t in sorted(decode_tasks, key=lambda t: t.task_id): + for t in decode_tasks: if cache.task_extend(t.task_id, t.next_pos): valid.append(t) else: diff --git a/astrai/inference/sample.py b/astrai/inference/sample.py index 1308f01..7e72657 100644 --- a/astrai/inference/sample.py +++ b/astrai/inference/sample.py @@ -343,6 +343,10 @@ def sample( When **temperature** is exactly 0 (scalar or single-element tensor) the function short-circuits to ``argmax`` for deterministic decode. + When **frequency_penalty** is 0 (the common decode case), the entire + frequency penalty computation — including the O(batch * vocab) count + tensor allocation — is skipped. + Args: logits: Raw logits ``[batch, vocab_size]``. frequency_penalty: Penalty per occurrence for repeated tokens @@ -359,14 +363,39 @@ def sample( ``True`` — a ``(token_ids, chosen_logprobs)`` tuple where ``chosen_logprobs`` has shape ``[batch]``. """ - return SamplingPipeline( - [ - TemperatureStrategy(temperature), - TopKStrategy(top_k), - TopPStrategy(top_p), - FrequencyPenaltyStrategy(frequency_penalty), - ] - ).sample( + greedy = ( + ( + isinstance(temperature, Tensor) + and temperature.numel() == 1 + and temperature.item() == 0 + ) + if isinstance(temperature, Tensor) + else temperature == 0 + ) + + if greedy: + tokens = logits.argmax(dim=-1) + if not return_logprobs: + return tokens + log_probs = torch.log_softmax(logits.float(), dim=-1) + chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1) + return tokens, chosen + + has_freq = ( + (isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any()) + if isinstance(frequency_penalty, Tensor) + else frequency_penalty != 0 + ) + + strategies: List[BaseSamplingStrategy] = [ + TemperatureStrategy(temperature), + TopKStrategy(top_k), + TopPStrategy(top_p), + ] + if has_freq: + strategies.append(FrequencyPenaltyStrategy(frequency_penalty)) + + return SamplingPipeline(strategies).sample( logits, filter_value=filter_value, input_ids=input_ids,