From 57b093dc340fb1412a919632b11160eedf625276 Mon Sep 17 00:00:00 2001 From: Xuhao Zhang <41683930+zzzzzzzxh@users.noreply.github.com> Date: Wed, 11 Mar 2026 14:11:19 +0800 Subject: [PATCH] [NPU]MindSpore backend support eagle3 (#17098) Co-authored-by: wangtiance Co-authored-by: Tiance Wang Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: ronnie_zheng --- .../extending/mindspore_models.md | 2 +- python/sglang/srt/models/mindspore.py | 72 +++++++++++++++---- .../sglang/srt/speculative/eagle_worker_v2.py | 3 + 3 files changed, 63 insertions(+), 14 deletions(-) diff --git a/docs/supported_models/extending/mindspore_models.md b/docs/supported_models/extending/mindspore_models.md index ce82fec68..1e5583293 100644 --- a/docs/supported_models/extending/mindspore_models.md +++ b/docs/supported_models/extending/mindspore_models.md @@ -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 diff --git a/python/sglang/srt/models/mindspore.py b/python/sglang/srt/models/mindspore.py index 0514756f8..da95ab139 100644 --- a/python/sglang/srt/models/mindspore.py +++ b/python/sglang/srt/models/mindspore.py @@ -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] diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index f3418db13..24003af5f 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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,