[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:
Xuhao Zhang
2026-03-11 09:11:19 +03:00
committed by GitHub
co-authored by wangtiance Tiance Wang gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> ronnie_zheng
parent 18cfeabd33
commit 57b093dc34
3 changed files with 63 additions and 14 deletions
@@ -6,7 +6,7 @@ MindSpore is a high-performance AI framework optimized for Ascend NPUs. This doc
## 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).
## Supported Models
+59 -13
View File
@@ -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,