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%)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user