Tiny remove the duplicate function in spec v2 (#14957)
This commit is contained in:
@@ -28,7 +28,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
SIMULATE_ACC_LEN,
|
SIMULATE_ACC_LEN,
|
||||||
generate_simulated_accept_index,
|
generate_simulated_accept_index,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.common import fast_topk, is_cuda, is_hip, is_npu, next_power_of_2
|
from sglang.srt.utils.common import is_cuda, is_hip, is_npu, next_power_of_2
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -47,7 +47,6 @@ if is_cuda():
|
|||||||
top_p_renorm_prob,
|
top_p_renorm_prob,
|
||||||
tree_speculative_sampling_target_only,
|
tree_speculative_sampling_target_only,
|
||||||
)
|
)
|
||||||
from sgl_kernel.top_k import fast_topk
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -371,56 +370,6 @@ class EagleVerifyInputV2Mixin:
|
|||||||
return predict, accept_length, accept_index
|
return predict, accept_length, accept_index
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, disable=_is_npu)
|
|
||||||
def select_top_k_tokens_tmp(
|
|
||||||
i: int,
|
|
||||||
topk_p: torch.Tensor,
|
|
||||||
topk_index: torch.Tensor,
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
scores: torch.Tensor,
|
|
||||||
topk: int,
|
|
||||||
):
|
|
||||||
# FIXME(lsyin): remove this duplicate code
|
|
||||||
if i == 0:
|
|
||||||
# The first step after extend
|
|
||||||
input_ids = topk_index.flatten()
|
|
||||||
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=hidden_states.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()
|
|
||||||
|
|
||||||
selected_input_index = topk_cs_index.flatten() // topk + torch.arange(
|
|
||||||
0, hidden_states.shape[0], step=topk, device=hidden_states.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
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def fill_new_verified_id(
|
def fill_new_verified_id(
|
||||||
verified_id,
|
verified_id,
|
||||||
|
|||||||
@@ -35,7 +35,6 @@ from sglang.srt.speculative.eagle_info_v2 import (
|
|||||||
assign_extend_cache_locs,
|
assign_extend_cache_locs,
|
||||||
fill_accepted_out_cache_loc,
|
fill_accepted_out_cache_loc,
|
||||||
fill_new_verified_id,
|
fill_new_verified_id,
|
||||||
select_top_k_tokens_tmp,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -44,6 +43,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
generate_token_bitmask,
|
generate_token_bitmask,
|
||||||
load_token_map,
|
load_token_map,
|
||||||
|
select_top_k_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
MultiprocessingSerializer,
|
MultiprocessingSerializer,
|
||||||
@@ -372,7 +372,7 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
# Forward multiple steps
|
# Forward multiple steps
|
||||||
scores = None
|
scores = None
|
||||||
for i in range(self.speculative_num_steps):
|
for i in range(self.speculative_num_steps):
|
||||||
input_ids, hidden_states, scores, tree_info = select_top_k_tokens_tmp(
|
input_ids, hidden_states, scores, tree_info = select_top_k_tokens(
|
||||||
i, topk_p, topk_index, hidden_states, scores, self.topk
|
i, topk_p, topk_index, hidden_states, scores, self.topk
|
||||||
)
|
)
|
||||||
score_list.append(tree_info[0])
|
score_list.append(tree_info[0])
|
||||||
|
|||||||
Reference in New Issue
Block a user