[NPU]MindSpore backend support eagle3 (#17098)
Co-authored-by: wangtiance <tiancew@qq.com> Co-authored-by: Tiance Wang <wangtiance@gmail.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, Iterable, Optional, Tuple
|
||||
from typing import Any, Iterable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -197,6 +197,12 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
self.key_cache = []
|
||||
self.value_cache = []
|
||||
|
||||
@property
|
||||
def hot_token_id(self):
|
||||
if hasattr(self.model, "hot_token_id"):
|
||||
return tensor_ms2torch(self.model.hot_token_id)
|
||||
return None
|
||||
|
||||
def get_arch(self, config):
|
||||
return _get_arch_from_config(config)
|
||||
|
||||
@@ -219,7 +225,7 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
else:
|
||||
cache = forward_batch.token_to_kv_pool.get_value_buffer(i)
|
||||
cache_ms = tensor_torch2ms(cache)
|
||||
if cache_ms.ndim == 3:
|
||||
if self.use_mla and cache_ms.ndim == 3:
|
||||
cache_ms = mint.unsqueeze(cache_ms, 2)
|
||||
cache_list.append(cache_ms)
|
||||
|
||||
@@ -236,30 +242,44 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
|
||||
return mutable(self.key_cache), mutable(self.value_cache)
|
||||
|
||||
def _is_prefill(self, forward_batch: ForwardBatch):
|
||||
# Different processing for the mindspore attention operator
|
||||
# Without any prefix cache => Use FlashAttentionScore
|
||||
# With cache => Use PagedAttention, no matter the query length is 1 or not
|
||||
is_prefill = (
|
||||
forward_batch.forward_mode.is_extend()
|
||||
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||
and not forward_batch.forward_mode.is_draft_extend()
|
||||
and not forward_batch.forward_mode.is_target_verify()
|
||||
)
|
||||
if forward_batch.extend_prefix_lens is not None:
|
||||
is_prefill = (
|
||||
is_prefill and forward_batch.extend_prefix_lens.sum().item() == 0
|
||||
)
|
||||
return is_prefill
|
||||
|
||||
def prepare_inputs(self, input_ids, positions, forward_batch):
|
||||
if self.use_mla:
|
||||
key_cache = self.get_kvcache(forward_batch)
|
||||
else:
|
||||
key_cache, value_cache = self.get_kvcache(forward_batch)
|
||||
|
||||
# Different processing for the mindspore attention operator
|
||||
# Without any prefix cache => Use FlashAttentionScore
|
||||
# With cache => Use PagedAttention, no matter the query length is 1 or not
|
||||
is_prefill = forward_batch.forward_mode.is_extend()
|
||||
is_prefill = is_prefill and forward_batch.extend_prefix_lens.sum().item() == 0
|
||||
|
||||
is_prefill = self._is_prefill(forward_batch)
|
||||
batch_valid_length = forward_batch.seq_lens.cpu().numpy()
|
||||
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
batch_valid_length += forward_batch.spec_info.num_tokens_per_req
|
||||
if forward_batch.extend_seq_lens is not None:
|
||||
q_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
|
||||
else:
|
||||
q_seq_lens = np.ones([forward_batch.batch_size], dtype=np.int32)
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
q_seq_lens = q_seq_lens * forward_batch.spec_info.num_tokens_per_req
|
||||
|
||||
page_size = forward_batch.token_to_kv_pool.page_size
|
||||
block_tables = tensor_torch2ms(
|
||||
(
|
||||
forward_batch.req_to_token_pool.req_to_token[
|
||||
forward_batch.req_pool_indices, : forward_batch.seq_lens.max()
|
||||
forward_batch.req_pool_indices, : batch_valid_length.max()
|
||||
][:, ::page_size]
|
||||
// page_size
|
||||
)
|
||||
@@ -283,6 +303,8 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
if not self.use_mla:
|
||||
model_inputs["value_cache"] = value_cache
|
||||
model_inputs["block_tables"] = block_tables
|
||||
# for speculative decode
|
||||
model_inputs["forward_mode"] = forward_batch.forward_mode
|
||||
return model_inputs
|
||||
|
||||
def forward(
|
||||
@@ -296,10 +318,17 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
# prepare model inputs
|
||||
model_inputs = self.model.prepare_inputs(forward_batch, model_inputs)
|
||||
|
||||
logits = self.model(**model_inputs)
|
||||
# Used by speculative decoding (EAGLE)
|
||||
if self.model.capture_aux_hidden_states:
|
||||
logits, hidden_states = self.model(**model_inputs)
|
||||
else:
|
||||
logits = self.model(**model_inputs)
|
||||
hidden_states = None
|
||||
|
||||
# TODO: npu tensor ms2torch error to be fix, remain issues of torch_npu to get tensor from dlpack
|
||||
logits_result = LogitsProcessorOutput(next_token_logits=tensor_ms2torch(logits))
|
||||
logits_result = LogitsProcessorOutput(
|
||||
next_token_logits=tensor_ms2torch(logits),
|
||||
hidden_states=tensor_ms2torch(hidden_states),
|
||||
)
|
||||
return logits_result
|
||||
|
||||
@classmethod
|
||||
@@ -313,5 +342,22 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
# The following methods are used for speculative decoding
|
||||
def get_embed_and_head(self):
|
||||
embed, head = self.model.get_embed_and_head()
|
||||
return tensor_ms2torch(embed), tensor_ms2torch(head)
|
||||
|
||||
def set_embed_and_head(self, embed, head):
|
||||
self.model.set_embed_and_head(tensor_torch2ms(embed), tensor_torch2ms(head))
|
||||
|
||||
def get_embed(self):
|
||||
return tensor_ms2torch(self.model.get_embed())
|
||||
|
||||
def set_embed(self, embed):
|
||||
self.model.set_embed(tensor_torch2ms(embed))
|
||||
|
||||
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
|
||||
self.model.set_eagle3_layers_to_capture(layer_ids)
|
||||
|
||||
|
||||
EntryClass = [MindSporeForCausalLM]
|
||||
|
||||
@@ -254,6 +254,9 @@ class EagleDraftWorker(BaseDraftWorker):
|
||||
if self.server_args.disable_cuda_graph:
|
||||
return
|
||||
|
||||
if self.server_args.model_impl == "mindspore":
|
||||
return
|
||||
|
||||
Device2DraftCudaGraphRunner = {
|
||||
"npu": EAGLEDraftNpuGraphRunner,
|
||||
"cuda": EAGLEDraftCudaGraphRunner,
|
||||
|
||||
Reference in New Issue
Block a user