[Feature] Support DeepSeek MTP on NPU (#11897)
Co-authored-by: liupeng374 <liupeng374@huawei.com>
This commit is contained in:
@@ -24,12 +24,13 @@ from sglang.srt.speculative.eagle_info_v2 import (
|
||||
EagleDraftInputV2Mixin,
|
||||
EagleVerifyInputV2Mixin,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
||||
from sglang.srt.speculative.spec_utils import (
|
||||
SIMULATE_ACC_LEN,
|
||||
TREE_SPEC_KERNEL_AVAILABLE,
|
||||
align_evict_mask_to_page_size,
|
||||
assign_req_to_token_pool,
|
||||
assign_req_to_token_pool_func,
|
||||
create_accept_length_filter,
|
||||
create_extend_after_decode_spec_info,
|
||||
filter_finished_cache_loc_kernel,
|
||||
@@ -37,17 +38,16 @@ from sglang.srt.speculative.spec_utils import (
|
||||
get_src_tgt_cache_loc,
|
||||
get_target_cache_loc,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda, is_hip, next_power_of_2
|
||||
from sglang.srt.utils import is_cuda, is_npu, next_power_of_2
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
if is_cuda():
|
||||
from sgl_kernel import (
|
||||
top_k_renorm_prob,
|
||||
top_p_renorm_prob,
|
||||
tree_speculative_sampling_target_only,
|
||||
verify_tree_greedy,
|
||||
)
|
||||
elif is_hip():
|
||||
from sgl_kernel import verify_tree_greedy
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -77,18 +77,22 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
|
||||
@classmethod
|
||||
def create_idle_input(cls, topk: int, spec_steps: int, num_verify_tokens: int):
|
||||
if not _is_npu:
|
||||
device = "cuda"
|
||||
else:
|
||||
device = "npu"
|
||||
return cls(
|
||||
draft_token=torch.empty((0,), dtype=torch.long, device="cuda"),
|
||||
custom_mask=torch.full((0,), True, dtype=torch.bool, device="cuda"),
|
||||
positions=torch.empty((0,), dtype=torch.int64, device="cuda"),
|
||||
draft_token=torch.empty((0,), dtype=torch.long, device=device),
|
||||
custom_mask=torch.full((0,), True, dtype=torch.bool, device=device),
|
||||
positions=torch.empty((0,), dtype=torch.int64, device=device),
|
||||
retrive_index=torch.full(
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device="cuda"
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device=device
|
||||
),
|
||||
retrive_next_token=torch.full(
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device="cuda"
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device=device
|
||||
),
|
||||
retrive_next_sibling=torch.full(
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device="cuda"
|
||||
(0, num_verify_tokens), -1, dtype=torch.long, device=device
|
||||
),
|
||||
retrive_cum_len=None,
|
||||
topk=topk,
|
||||
@@ -134,14 +138,13 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
self.last_loc = last_loc
|
||||
|
||||
bs = batch.batch_size()
|
||||
assign_req_to_token_pool[(bs,)](
|
||||
assign_req_to_token_pool_func(
|
||||
batch.req_pool_indices,
|
||||
batch.req_to_token_pool.req_to_token,
|
||||
batch.seq_lens,
|
||||
end_offset,
|
||||
batch.out_cache_loc,
|
||||
batch.req_to_token_pool.req_to_token.shape[1],
|
||||
next_power_of_2(bs),
|
||||
bs,
|
||||
)
|
||||
|
||||
def generate_attn_arg_prefill(
|
||||
@@ -151,16 +154,17 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
paged_kernel_lens_sum: int,
|
||||
req_to_token: torch.Tensor,
|
||||
):
|
||||
device = req_pool_indices.device
|
||||
batch_size = len(req_pool_indices)
|
||||
qo_indptr = torch.arange(
|
||||
0,
|
||||
(1 + batch_size) * self.draft_token_num,
|
||||
step=self.draft_token_num,
|
||||
dtype=torch.int32,
|
||||
device="cuda",
|
||||
device=device,
|
||||
)
|
||||
cum_kv_seq_len = torch.zeros(
|
||||
(batch_size + 1,), dtype=torch.int32, device="cuda"
|
||||
(batch_size + 1,), dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
paged_kernel_lens = paged_kernel_lens + self.draft_token_num
|
||||
@@ -169,7 +173,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
kv_indices = torch.empty(
|
||||
paged_kernel_lens_sum + self.draft_token_num * batch_size,
|
||||
dtype=torch.int32,
|
||||
device="cuda",
|
||||
device=device,
|
||||
)
|
||||
create_flashinfer_kv_indices_triton[(batch_size,)](
|
||||
req_to_token,
|
||||
@@ -226,11 +230,11 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
|
||||
predict_shape = list(logits_output.next_token_logits.shape)[:-1]
|
||||
predict_shape[-1] += 1
|
||||
predict = torch.empty(predict_shape, dtype=torch.int32, device="cuda")
|
||||
predict = torch.empty(predict_shape, dtype=torch.int32, device=batch.device)
|
||||
accept_index = torch.full(
|
||||
(bs, self.spec_steps + 1), -1, dtype=torch.int32, device="cuda"
|
||||
(bs, self.spec_steps + 1), -1, dtype=torch.int32, device=batch.device
|
||||
)
|
||||
accept_length = torch.empty((bs,), dtype=torch.int32, device="cuda")
|
||||
accept_length = torch.empty((bs,), dtype=torch.int32, device=batch.device)
|
||||
|
||||
if bs != len(sampling_info):
|
||||
sampling_info = copy.deepcopy(sampling_info)
|
||||
@@ -254,7 +258,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
linear_penalty = torch.zeros(
|
||||
(bs, logits_output.next_token_logits.shape[1]),
|
||||
dtype=torch.float32,
|
||||
device="cuda",
|
||||
device=batch.device,
|
||||
)
|
||||
sampling_info.apply_logits_bias(linear_penalty)
|
||||
logits_output.next_token_logits.add_(
|
||||
@@ -276,11 +280,10 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
"Falling back to greedy verification."
|
||||
)
|
||||
|
||||
if is_all_greedy or not TREE_SPEC_KERNEL_AVAILABLE:
|
||||
if is_all_greedy or not TREE_SPEC_KERNEL_AVAILABLE or _is_npu:
|
||||
target_predict = torch.argmax(logits_output.next_token_logits, dim=-1)
|
||||
target_predict = target_predict.reshape(bs, self.draft_token_num)
|
||||
|
||||
verify_tree_greedy(
|
||||
predict, accept_index, accept_length = verify_tree_greedy_func(
|
||||
predicts=predict, # mutable
|
||||
accept_index=accept_index, # mutable
|
||||
accept_token_num=accept_length, # mutable
|
||||
@@ -289,7 +292,9 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
retrive_next_token=self.retrive_next_token,
|
||||
retrive_next_sibling=self.retrive_next_sibling,
|
||||
target_predict=target_predict,
|
||||
topk=self.topk,
|
||||
)
|
||||
|
||||
else:
|
||||
# apply temperature and get target probs
|
||||
expanded_temperature = torch.repeat_interleave(
|
||||
@@ -315,14 +320,16 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
target_probs = target_probs.reshape(bs, self.draft_token_num, -1)
|
||||
|
||||
draft_probs = torch.zeros(
|
||||
target_probs.shape, dtype=torch.float32, device="cuda"
|
||||
target_probs.shape, dtype=torch.float32, device=batch.device
|
||||
)
|
||||
|
||||
# coins for rejection sampling
|
||||
coins = torch.rand_like(candidates, dtype=torch.float32, device="cuda")
|
||||
coins = torch.rand_like(
|
||||
candidates, dtype=torch.float32, device=batch.device
|
||||
)
|
||||
# coins for final sampling
|
||||
coins_for_final_sampling = torch.rand(
|
||||
(bs,), dtype=torch.float32, device="cuda"
|
||||
(bs,), dtype=torch.float32, device=batch.device
|
||||
)
|
||||
tree_speculative_sampling_target_only(
|
||||
predicts=predict, # mutable
|
||||
@@ -468,14 +475,13 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
if not has_finished:
|
||||
if page_size == 1 or self.topk == 1:
|
||||
batch.out_cache_loc = batch.out_cache_loc[accept_index]
|
||||
assign_req_to_token_pool[(bs,)](
|
||||
assign_req_to_token_pool_func(
|
||||
batch.req_pool_indices,
|
||||
batch.req_to_token_pool.req_to_token,
|
||||
batch.seq_lens,
|
||||
batch.seq_lens + accept_length + 1,
|
||||
batch.out_cache_loc,
|
||||
batch.req_to_token_pool.req_to_token.shape[1],
|
||||
next_power_of_2(bs),
|
||||
bs,
|
||||
)
|
||||
else:
|
||||
batch.out_cache_loc = tgt_cache_loc
|
||||
@@ -501,14 +507,13 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
||||
)
|
||||
else:
|
||||
if page_size == 1 or self.topk == 1:
|
||||
assign_req_to_token_pool[(bs,)](
|
||||
assign_req_to_token_pool_func(
|
||||
batch.req_pool_indices,
|
||||
batch.req_to_token_pool.req_to_token,
|
||||
batch.seq_lens,
|
||||
batch.seq_lens + accept_length + 1,
|
||||
batch.out_cache_loc[accept_index],
|
||||
batch.req_to_token_pool.req_to_token.shape[1],
|
||||
next_power_of_2(bs),
|
||||
bs,
|
||||
)
|
||||
batch.seq_lens.add_(accept_length + 1)
|
||||
batch.seq_lens_cpu.add_(accept_length_cpu + 1)
|
||||
@@ -695,17 +700,18 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
||||
paged_kernel_lens_sum: int,
|
||||
req_to_token: torch.Tensor,
|
||||
):
|
||||
device = req_pool_indices.device
|
||||
bs = self.accept_length.numel()
|
||||
qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device="cuda")
|
||||
qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
||||
qo_indptr[1:] = torch.cumsum(self.accept_length, dim=0)
|
||||
cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device="cuda")
|
||||
cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
||||
cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0)
|
||||
|
||||
if paged_kernel_lens_sum is None:
|
||||
paged_kernel_lens_sum = cum_kv_seq_len[-1]
|
||||
|
||||
kv_indices = torch.empty(
|
||||
paged_kernel_lens_sum, dtype=torch.int32, device="cuda"
|
||||
paged_kernel_lens_sum, dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
create_flashinfer_kv_indices_triton[(bs,)](
|
||||
|
||||
Reference in New Issue
Block a user