[sgl] reduce specdec cpu overhead (#23321)
This commit is contained in:
@@ -33,17 +33,21 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
|||||||
seq_lens_cpu: torch.Tensor,
|
seq_lens_cpu: torch.Tensor,
|
||||||
last_loc: torch.Tensor,
|
last_loc: torch.Tensor,
|
||||||
extend_num_tokens: int,
|
extend_num_tokens: int,
|
||||||
|
num_new_pages: int = None,
|
||||||
):
|
):
|
||||||
if self.debug_mode:
|
if self.debug_mode:
|
||||||
assert torch.all(
|
assert torch.all(
|
||||||
(last_loc + 1) % self.page_size == prefix_lens % self.page_size
|
(last_loc + 1) % self.page_size == prefix_lens % self.page_size
|
||||||
)
|
)
|
||||||
|
|
||||||
num_new_pages = (
|
if num_new_pages is None:
|
||||||
(seq_lens + self.roundup) // self.page_size
|
num_new_pages_tensor = (
|
||||||
- (prefix_lens + self.roundup) // self.page_size
|
(seq_lens + self.roundup) // self.page_size
|
||||||
).sum()
|
- (prefix_lens + self.roundup) // self.page_size
|
||||||
num_new_pages_item = num_new_pages.item()
|
).sum()
|
||||||
|
num_new_pages_item = num_new_pages_tensor.item()
|
||||||
|
else:
|
||||||
|
num_new_pages_item = num_new_pages
|
||||||
if self.need_sort and num_new_pages_item > len(self.free_pages):
|
if self.need_sort and num_new_pages_item > len(self.free_pages):
|
||||||
self.merge_and_sort_free()
|
self.merge_and_sort_free()
|
||||||
|
|
||||||
|
|||||||
@@ -410,6 +410,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
seq_lens_cpu: torch.Tensor,
|
seq_lens_cpu: torch.Tensor,
|
||||||
last_loc: torch.Tensor,
|
last_loc: torch.Tensor,
|
||||||
extend_num_tokens: int,
|
extend_num_tokens: int,
|
||||||
|
num_new_pages: int = None,
|
||||||
):
|
):
|
||||||
if self.debug_mode:
|
if self.debug_mode:
|
||||||
assert torch.all(
|
assert torch.all(
|
||||||
@@ -439,11 +440,12 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
if self.debug_mode:
|
if self.debug_mode:
|
||||||
assert len(torch.unique(out_indices)) == len(out_indices)
|
assert len(torch.unique(out_indices)) == len(out_indices)
|
||||||
|
|
||||||
num_new_pages = get_num_new_pages(
|
if num_new_pages is None:
|
||||||
seq_lens=seq_lens_cpu,
|
num_new_pages = get_num_new_pages(
|
||||||
page_size=self.page_size,
|
seq_lens=seq_lens_cpu,
|
||||||
prefix_lens=prefix_lens_cpu,
|
page_size=self.page_size,
|
||||||
)
|
prefix_lens=prefix_lens_cpu,
|
||||||
|
)
|
||||||
if num_new_pages > len(self.free_pages):
|
if num_new_pages > len(self.free_pages):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -307,6 +307,7 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
seq_lens_cpu,
|
seq_lens_cpu,
|
||||||
hisparse_last_loc,
|
hisparse_last_loc,
|
||||||
len(logical_indices),
|
len(logical_indices),
|
||||||
|
num_new_pages=num_new_pages,
|
||||||
)
|
)
|
||||||
assert (
|
assert (
|
||||||
hisparse_indices is not None
|
hisparse_indices is not None
|
||||||
|
|||||||
@@ -395,6 +395,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
seq_lens_cpu,
|
seq_lens_cpu,
|
||||||
last_loc,
|
last_loc,
|
||||||
extend_num_tokens,
|
extend_num_tokens,
|
||||||
|
num_new_pages=num_new_pages,
|
||||||
)
|
)
|
||||||
alloc_swa_indices = self.swa_attn_allocator.alloc_extend(
|
alloc_swa_indices = self.swa_attn_allocator.alloc_extend(
|
||||||
prefix_lens,
|
prefix_lens,
|
||||||
@@ -403,6 +404,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
seq_lens_cpu,
|
seq_lens_cpu,
|
||||||
swa_last_loc,
|
swa_last_loc,
|
||||||
extend_num_tokens,
|
extend_num_tokens,
|
||||||
|
num_new_pages=num_new_pages,
|
||||||
)
|
)
|
||||||
assert alloc_full_indices is not None
|
assert alloc_full_indices is not None
|
||||||
assert alloc_swa_indices is not None
|
assert alloc_swa_indices is not None
|
||||||
|
|||||||
@@ -116,23 +116,25 @@ class EagleDraftInputV2Mixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
page_size = batch.token_to_kv_pool_allocator.page_size
|
page_size = batch.token_to_kv_pool_allocator.page_size
|
||||||
cur_kv_lens_cpu = []
|
|
||||||
nxt_kv_lens_cpu = []
|
|
||||||
num_needed_tokens = 0
|
|
||||||
alloc_len_per_decode = get_alloc_len_per_decode()
|
alloc_len_per_decode = get_alloc_len_per_decode()
|
||||||
for r in batch.reqs:
|
double_alloc = alloc_len_per_decode + alloc_len_per_decode
|
||||||
# Over-allocation happens here
|
|
||||||
x = r.kv_committed_len + 2 * alloc_len_per_decode - r.kv_allocated_len
|
cur_kv_lens = [0] * bs
|
||||||
cur_kv_lens_cpu.append(r.kv_allocated_len)
|
nxt_kv_lens = [0] * bs
|
||||||
nxt_kv_lens_cpu.append(r.kv_allocated_len + x)
|
num_needed_tokens = 0
|
||||||
num_needed_tokens += x
|
for i, r in enumerate(batch.reqs):
|
||||||
r.kv_allocated_len += x
|
cur = r.kv_allocated_len
|
||||||
|
nxt = r.kv_committed_len + double_alloc
|
||||||
|
cur_kv_lens[i] = cur
|
||||||
|
nxt_kv_lens[i] = nxt
|
||||||
|
num_needed_tokens += nxt - cur
|
||||||
|
r.kv_allocated_len = nxt
|
||||||
r.decode_batch_idx += 1
|
r.decode_batch_idx += 1
|
||||||
# Pre-claim bonus slot here (like normal decode); resolve subtracts 1.
|
# Pre-claim bonus slot here (like normal decode); resolve subtracts 1.
|
||||||
r.kv_committed_len += 1
|
r.kv_committed_len += 1
|
||||||
|
|
||||||
cur_kv_lens_cpu = torch.tensor(cur_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu")
|
||||||
nxt_kv_lens_cpu = torch.tensor(nxt_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
nxt_kv_lens_cpu = torch.tensor(nxt_kv_lens, dtype=torch.int32, device="cpu")
|
||||||
|
|
||||||
if page_size == 1:
|
if page_size == 1:
|
||||||
out_cache_loc = alloc_token_slots(batch.tree_cache, num_needed_tokens)
|
out_cache_loc = alloc_token_slots(batch.tree_cache, num_needed_tokens)
|
||||||
|
|||||||
@@ -469,7 +469,63 @@ def create_num_accepted_drafts_filter(
|
|||||||
return num_accepted_drafts_filter
|
return num_accepted_drafts_filter
|
||||||
|
|
||||||
|
|
||||||
|
def _select_top_k_tokens_first(
|
||||||
|
topk_p: torch.Tensor,
|
||||||
|
topk_index: torch.Tensor,
|
||||||
|
hidden_states: Optional[torch.Tensor],
|
||||||
|
topk: int,
|
||||||
|
):
|
||||||
|
input_ids = topk_index.flatten()
|
||||||
|
if hidden_states is not None:
|
||||||
|
hidden_states = hidden_states.repeat_interleave(topk, dim=0)
|
||||||
|
|
||||||
|
tree_info = (
|
||||||
|
topk_p.unsqueeze(1), # (b, 1, topk)
|
||||||
|
topk_index, # (b, topk)
|
||||||
|
torch.arange(-1, topk, dtype=torch.long, device=input_ids.device).expand(
|
||||||
|
topk_p.shape[0], -1
|
||||||
|
), # (b, topk + 1) — expand avoids the allocation of repeat
|
||||||
|
)
|
||||||
|
return input_ids, hidden_states, topk_p, tree_info
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, disable=_is_npu)
|
@torch.compile(dynamic=True, disable=_is_npu)
|
||||||
|
def _select_top_k_tokens_later(
|
||||||
|
i: int,
|
||||||
|
topk_p: torch.Tensor,
|
||||||
|
topk_index: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
scores: torch.Tensor,
|
||||||
|
topk: int,
|
||||||
|
):
|
||||||
|
topk_sq = topk * topk
|
||||||
|
|
||||||
|
expand_scores = scores.unsqueeze(2) * topk_p.view(-1, topk, topk)
|
||||||
|
# (b, topk, 1) * (b, topk, topk) -> (b, topk, topk)
|
||||||
|
|
||||||
|
topk_cs_p, topk_cs_index = fast_topk(
|
||||||
|
expand_scores.flatten(start_dim=1), topk, dim=-1
|
||||||
|
) # (b, topk)
|
||||||
|
|
||||||
|
topk_index = topk_index.view(-1, topk_sq)
|
||||||
|
input_ids = torch.gather(topk_index, 1, topk_cs_index).flatten()
|
||||||
|
|
||||||
|
if hidden_states.shape[0] > 0:
|
||||||
|
flat_cs = topk_cs_index.flatten()
|
||||||
|
batch_offsets = torch.arange(
|
||||||
|
0, hidden_states.shape[0], step=topk, device=flat_cs.device
|
||||||
|
)
|
||||||
|
selected_input_index = flat_cs // topk + batch_offsets.repeat_interleave(topk)
|
||||||
|
hidden_states = hidden_states[selected_input_index]
|
||||||
|
|
||||||
|
tree_info = (
|
||||||
|
expand_scores, # (b, topk, topk)
|
||||||
|
topk_index, # (b, topk * topk)
|
||||||
|
topk_cs_index + (topk_sq * (i - 1) + topk), # (b, topk)
|
||||||
|
)
|
||||||
|
return input_ids, hidden_states, topk_cs_p, tree_info
|
||||||
|
|
||||||
|
|
||||||
def select_top_k_tokens(
|
def select_top_k_tokens(
|
||||||
i: int,
|
i: int,
|
||||||
topk_p: torch.Tensor,
|
topk_p: torch.Tensor,
|
||||||
@@ -479,45 +535,10 @@ def select_top_k_tokens(
|
|||||||
topk: int,
|
topk: int,
|
||||||
):
|
):
|
||||||
if i == 0:
|
if i == 0:
|
||||||
# The first step after extend
|
return _select_top_k_tokens_first(topk_p, topk_index, hidden_states, topk)
|
||||||
input_ids = topk_index.flatten()
|
return _select_top_k_tokens_later(
|
||||||
if hidden_states is not None:
|
i, topk_p, topk_index, hidden_states, scores, topk
|
||||||
hidden_states = hidden_states.repeat_interleave(topk, dim=0)
|
)
|
||||||
scores = topk_p # shape: (b, topk)
|
|
||||||
|
|
||||||
tree_info = (
|
|
||||||
topk_p.unsqueeze(1), # shape: (b, 1, topk)
|
|
||||||
topk_index, # shape: (b, topk)
|
|
||||||
torch.arange(-1, topk, dtype=torch.long, device=input_ids.device)
|
|
||||||
.unsqueeze(0)
|
|
||||||
.repeat(topk_p.shape[0], 1), # shape: (b, topk + 1)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# The later decode steps
|
|
||||||
expand_scores = torch.mul(
|
|
||||||
scores.unsqueeze(2), topk_p.reshape(-1, topk, topk)
|
|
||||||
) # (b, topk, 1) x (b, topk ,topk) -> (b, topk, topk)
|
|
||||||
topk_cs_p, topk_cs_index = fast_topk(
|
|
||||||
expand_scores.flatten(start_dim=1), topk, dim=-1
|
|
||||||
) # (b, topk)
|
|
||||||
scores = topk_cs_p # shape: (b, topk)
|
|
||||||
|
|
||||||
topk_index = topk_index.reshape(-1, topk**2)
|
|
||||||
input_ids = torch.gather(topk_index, index=topk_cs_index, dim=1).flatten()
|
|
||||||
|
|
||||||
if hidden_states.shape[0] > 0:
|
|
||||||
selected_input_index = topk_cs_index.flatten() // topk + torch.arange(
|
|
||||||
0, hidden_states.shape[0], step=topk, device=topk_index.device
|
|
||||||
).repeat_interleave(topk)
|
|
||||||
hidden_states = hidden_states[selected_input_index, :]
|
|
||||||
|
|
||||||
tree_info = (
|
|
||||||
expand_scores, # shape: (b, topk, topk)
|
|
||||||
topk_index, # shape: (b, topk * topk)
|
|
||||||
topk_cs_index + (topk**2 * (i - 1) + topk), # shape: (b, topk)
|
|
||||||
)
|
|
||||||
|
|
||||||
return input_ids, hidden_states, scores, tree_info
|
|
||||||
|
|
||||||
|
|
||||||
def generate_simulated_accept_index(
|
def generate_simulated_accept_index(
|
||||||
|
|||||||
Reference in New Issue
Block a user