[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:
co-authored by
wangtiance
Tiance Wang
gemini-code-assist[bot]
ronnie_zheng
parent
18cfeabd33
commit
57b093dc34
@@ -6,7 +6,7 @@ MindSpore is a high-performance AI framework optimized for Ascend NPUs. This doc
|
|||||||
|
|
||||||
## Requirements
|
## Requirements
|
||||||
|
|
||||||
MindSpore currently only supports Ascend NPU devices. Users need to first install CANN 8.5.
|
MindSpore currently only supports Ascend NPU devices. Users need to first install Ascend CANN 8.5.
|
||||||
The CANN software packages can be downloaded from the [Ascend Official Website](https://www.hiascend.com).
|
The CANN software packages can be downloaded from the [Ascend Official Website](https://www.hiascend.com).
|
||||||
|
|
||||||
## Supported Models
|
## Supported Models
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Any, Iterable, Optional, Tuple
|
from typing import Any, Iterable, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -197,6 +197,12 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
self.key_cache = []
|
self.key_cache = []
|
||||||
self.value_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):
|
def get_arch(self, config):
|
||||||
return _get_arch_from_config(config)
|
return _get_arch_from_config(config)
|
||||||
|
|
||||||
@@ -219,7 +225,7 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
cache = forward_batch.token_to_kv_pool.get_value_buffer(i)
|
cache = forward_batch.token_to_kv_pool.get_value_buffer(i)
|
||||||
cache_ms = tensor_torch2ms(cache)
|
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_ms = mint.unsqueeze(cache_ms, 2)
|
||||||
cache_list.append(cache_ms)
|
cache_list.append(cache_ms)
|
||||||
|
|
||||||
@@ -236,30 +242,44 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
|
|
||||||
return mutable(self.key_cache), mutable(self.value_cache)
|
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):
|
def prepare_inputs(self, input_ids, positions, forward_batch):
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
key_cache = self.get_kvcache(forward_batch)
|
key_cache = self.get_kvcache(forward_batch)
|
||||||
else:
|
else:
|
||||||
key_cache, value_cache = self.get_kvcache(forward_batch)
|
key_cache, value_cache = self.get_kvcache(forward_batch)
|
||||||
|
|
||||||
# Different processing for the mindspore attention operator
|
is_prefill = self._is_prefill(forward_batch)
|
||||||
# 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
|
|
||||||
|
|
||||||
batch_valid_length = forward_batch.seq_lens.cpu().numpy()
|
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:
|
if forward_batch.extend_seq_lens is not None:
|
||||||
q_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
|
q_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
|
||||||
else:
|
else:
|
||||||
q_seq_lens = np.ones([forward_batch.batch_size], dtype=np.int32)
|
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
|
page_size = forward_batch.token_to_kv_pool.page_size
|
||||||
block_tables = tensor_torch2ms(
|
block_tables = tensor_torch2ms(
|
||||||
(
|
(
|
||||||
forward_batch.req_to_token_pool.req_to_token[
|
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]
|
||||||
// page_size
|
// page_size
|
||||||
)
|
)
|
||||||
@@ -283,6 +303,8 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
model_inputs["value_cache"] = value_cache
|
model_inputs["value_cache"] = value_cache
|
||||||
model_inputs["block_tables"] = block_tables
|
model_inputs["block_tables"] = block_tables
|
||||||
|
# for speculative decode
|
||||||
|
model_inputs["forward_mode"] = forward_batch.forward_mode
|
||||||
return model_inputs
|
return model_inputs
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
@@ -296,10 +318,17 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
# prepare model inputs
|
# prepare model inputs
|
||||||
model_inputs = self.model.prepare_inputs(forward_batch, 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(
|
||||||
logits_result = LogitsProcessorOutput(next_token_logits=tensor_ms2torch(logits))
|
next_token_logits=tensor_ms2torch(logits),
|
||||||
|
hidden_states=tensor_ms2torch(hidden_states),
|
||||||
|
)
|
||||||
return logits_result
|
return logits_result
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -313,5 +342,22 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
except Exception:
|
except Exception:
|
||||||
return None
|
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]
|
EntryClass = [MindSporeForCausalLM]
|
||||||
|
|||||||
@@ -254,6 +254,9 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
if self.server_args.disable_cuda_graph:
|
if self.server_args.disable_cuda_graph:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if self.server_args.model_impl == "mindspore":
|
||||||
|
return
|
||||||
|
|
||||||
Device2DraftCudaGraphRunner = {
|
Device2DraftCudaGraphRunner = {
|
||||||
"npu": EAGLEDraftNpuGraphRunner,
|
"npu": EAGLEDraftNpuGraphRunner,
|
||||||
"cuda": EAGLEDraftCudaGraphRunner,
|
"cuda": EAGLEDraftCudaGraphRunner,
|
||||||
|
|||||||
Reference in New Issue
Block a user