diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index d0a4beb2d..533e331dc 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -45,7 +45,9 @@ def _reshape_kv_for_fia_nz( logger = logging.getLogger(__name__) -SWA_INT_MAX = 2147483647 + +# default max value of full attention window size +FULL_ATTENTION_WINDOW = 2147483647 @dataclass @@ -1044,6 +1046,7 @@ class AscendAttnBackend(AttentionBackend): save_kv_cache, q_rope=q_rope, k_rope=k_rope, + sinks=sinks, ) if not self.use_mla: @@ -1106,10 +1109,12 @@ class AscendAttnBackend(AttentionBackend): pre_tokens=( layer.sliding_window_size if layer.sliding_window_size != -1 - else SWA_INT_MAX + else FULL_ATTENTION_WINDOW ), next_tokens=( - 0 if layer.sliding_window_size != -1 else SWA_INT_MAX + 0 + if layer.sliding_window_size != -1 + else FULL_ATTENTION_WINDOW ), atten_mask=self.fia_mask, block_table=block_tables, @@ -1637,6 +1642,7 @@ class AscendAttnBackend(AttentionBackend): save_kv_cache: bool, q_rope: Optional[torch.Tensor] = None, k_rope: Optional[torch.Tensor] = None, + sinks: Optional[torch.Tensor] = None, ): if save_kv_cache: if self.use_mla: @@ -1658,16 +1664,22 @@ class AscendAttnBackend(AttentionBackend): -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim ) query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous() + if not self.graph_mode: num_token_padding = query.shape[0] query = query[: forward_batch.num_token_non_padded_cpu] + if self.forward_metadata.seq_lens_cpu_int is None: actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list else: actual_seq_lengths_kv = ( self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() ) - if forward_batch.forward_mode.is_draft_extend(): + + if ( + forward_batch.forward_mode.is_draft_extend() + or forward_batch.forward_mode.is_draft_extend_v2() + ): actual_seq_lengths = ( np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist() ) @@ -1677,27 +1689,43 @@ class AscendAttnBackend(AttentionBackend): self.speculative_num_draft_tokens + query.shape[0], self.speculative_num_draft_tokens, ) + + is_swa_layer = layer.sliding_window_size != -1 + if ( + is_swa_layer + and self.is_hybrid_swa + and hasattr(self.forward_metadata, "block_tables_swa") + ): + block_table = self.forward_metadata.block_tables_swa + else: + block_table = self.forward_metadata.block_tables + if layer.attn_type == AttentionType.ENCODER_ONLY: mask = None sparse_mode = 0 else: mask = self.mtp_mask - sparse_mode = 3 + sparse_mode = 4 if is_swa_layer else 3 - attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score( + attn_output, _ = torch_npu.npu_fused_infer_attention_score_v2( query, k_cache, v_cache, - block_table=self.forward_metadata.block_tables, + block_table=block_table, block_size=self.page_size, - num_heads=layer.tp_q_head_num, + num_query_heads=layer.tp_q_head_num, num_key_value_heads=layer.tp_k_head_num, input_layout="TND", atten_mask=mask, - scale=layer.scaling, - actual_seq_lengths=actual_seq_lengths, - actual_seq_lengths_kv=actual_seq_lengths_kv, + softmax_scale=layer.scaling, + actual_seq_qlen=actual_seq_lengths, + actual_seq_kvlen=actual_seq_lengths_kv, sparse_mode=sparse_mode, + pre_tokens=( + layer.sliding_window_size if is_swa_layer else FULL_ATTENTION_WINDOW + ), + next_tokens=0 if is_swa_layer else FULL_ATTENTION_WINDOW, + learnable_sink=sinks, ) attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim) if ( @@ -1840,27 +1868,89 @@ class AscendAttnBackend(AttentionBackend): ) if sinks is not None: - k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) - v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) - # Use SWA block tables if hybrid SWA is enabled for this layer if self.is_hybrid_swa and layer.sliding_window_size != -1: block_tables = self.forward_metadata.block_tables_swa else: block_tables = self.forward_metadata.block_tables - attn_out = attention_sinks_triton( - q, - k_cache, - v_cache, - sinks, - block_tables, - self.forward_metadata.seq_lens, - layer.scaling, - layer.sliding_window_size, - layer.tp_q_head_num, - layer.tp_k_head_num, - ) - return attn_out + if self.use_fia: + k_cache = ( + self.token_to_kv_pool.get_key_buffer(layer.layer_id) + .view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim) + .contiguous() + ) + v_cache = ( + self.token_to_kv_pool.get_value_buffer(layer.layer_id) + .view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim) + .contiguous() + ) + query = q.reshape( + -1, layer.tp_q_head_num, layer.qk_head_dim + ).contiguous() + + if self.forward_metadata.seq_lens_cpu_int is None: + actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list + else: + actual_seq_lengths_kv = ( + self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() + ) + seq_lens_list = ( + self.forward_metadata.seq_lens_cpu_list + if self.forward_metadata.seq_lens_cpu_int is None + else self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() + ) + actual_seq_lengths = ( + torch.tensor([1] * len(seq_lens_list), dtype=torch.int32) + .cumsum(dim=0) + .tolist() + ) + if layer.sliding_window_size != -1: + sparse_mode = 4 + else: + sparse_mode = 3 + + attn_output, _ = torch_npu.npu_fused_infer_attention_score_v2( + query, + k_cache, + v_cache, + num_query_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + input_layout="TND", + pre_tokens=( + layer.sliding_window_size + if layer.sliding_window_size != -1 + else FULL_ATTENTION_WINDOW + ), + next_tokens=0, + atten_mask=self.fia_mask.to(torch.int8), + sparse_mode=sparse_mode, + softmax_scale=layer.scaling, + block_table=block_tables, + block_size=self.page_size, + actual_seq_qlen=actual_seq_lengths, + actual_seq_kvlen=actual_seq_lengths_kv, + learnable_sink=sinks, + ) + attn_output = attn_output.view( + -1, layer.tp_q_head_num * layer.v_head_dim + ) + return attn_output + else: + k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) + v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) + attn_out = attention_sinks_triton( + q, + k_cache, + v_cache, + sinks, + block_tables, + self.forward_metadata.seq_lens, + layer.scaling, + layer.sliding_window_size, + layer.tp_q_head_num, + layer.tp_k_head_num, + ) + return attn_out if not self.use_mla: k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view( @@ -1876,6 +1966,60 @@ class AscendAttnBackend(AttentionBackend): actual_seq_len_kv = ( self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() ) + + if (layer.qk_head_dim != layer.v_head_dim) and ( + self.is_hybrid_swa and layer.sliding_window_size == -1 + ): + query_v2 = q.reshape( + -1, layer.tp_q_head_num, layer.qk_head_dim + ).contiguous() + actual_seq_qlen = ( + torch.tensor([1] * len(actual_seq_len_kv), dtype=torch.int32) + .cumsum(dim=0) + .tolist() + ) + common_kwargs = dict( + num_query_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + input_layout="TND", + pre_tokens=FULL_ATTENTION_WINDOW, + next_tokens=0, + atten_mask=self.fia_mask.to(torch.int8), + sparse_mode=3, + softmax_scale=layer.scaling, + block_table=self.forward_metadata.block_tables, + block_size=self.page_size, + actual_seq_qlen=actual_seq_qlen, + actual_seq_kvlen=actual_seq_len_kv, + ) + workspace = ( + torch_npu._npu_fused_infer_attention_score_v2_get_max_workspace( + query_v2, + k_cache.contiguous(), + v_cache.contiguous(), + **common_kwargs, + ) + ) + attn_output = torch.empty( + ( + query_v2.shape[0], + layer.tp_q_head_num, + layer.v_head_dim, + ), + dtype=q.dtype, + device=q.device, + ) + softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device) + torch_npu.npu_fused_infer_attention_score_v2.out( + query_v2, + k_cache.contiguous(), + v_cache.contiguous(), + **common_kwargs, + workspace=workspace, + out=[attn_output, softmax_lse], + ) + return attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim) + num_tokens = query.shape[0] workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace( query, @@ -2072,12 +2216,17 @@ class AscendAttnBackend(AttentionBackend): self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() ) block_size = self.page_size - max_model_len = block_tables.shape[-1] * block_size - swa_mask = self.ascend_attn_mask_builder.get_swa_mask( - self.forward_metadata.seq_lens, - max_model_len, - layer.sliding_window_size, - ) + + if sinks is not None: + mask = self.fia_mask + else: + max_model_len = block_tables.shape[-1] * block_size + mask = self.ascend_attn_mask_builder.get_swa_mask( + self.forward_metadata.seq_lens, + max_model_len, + layer.sliding_window_size, + ) + attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2( q.view( forward_batch.batch_size, @@ -2095,14 +2244,14 @@ class AscendAttnBackend(AttentionBackend): num_key_value_heads=layer.tp_k_head_num, input_layout="BSND", block_size=block_size, - atten_mask=( - swa_mask if layer.sliding_window_size != -1 else None - ), + atten_mask=(mask if layer.sliding_window_size != -1 else None), sparse_mode=4 if layer.sliding_window_size != -1 else 0, softmax_scale=layer.scaling, block_table=block_tables, actual_seq_qlen=[1] * len(self.forward_metadata.seq_lens), actual_seq_kvlen=actual_seq_len_kv, + pre_tokens=layer.sliding_window_size, + next_tokens=0, learnable_sink=sinks, ) attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim) @@ -2142,7 +2291,7 @@ class AscendAttnBackend(AttentionBackend): -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim ), v_cache.view( - -1, self.page_size, layer.tp_v_head_num * layer.qk_head_dim + -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim ), num_heads=layer.tp_q_head_num, num_key_value_heads=layer.tp_k_head_num, diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/multi_layer_eagle_draft_extend_npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/multi_layer_eagle_draft_extend_npu_graph_runner.py new file mode 100644 index 000000000..b0bbd995a --- /dev/null +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/multi_layer_eagle_draft_extend_npu_graph_runner.py @@ -0,0 +1,159 @@ +# Copyright 2024-2025 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Run the multi-layer eagle draft extend model with npu graph.""" + +from __future__ import annotations + +import logging +import threading +import time +from typing import TYPE_CHECKING, List, Optional + +import torch + +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import ( + MultiLayerEagleDraftExtendCudaGraphRunner, + MultiLayerEagleMultiStepDraftExtendCudaGraphRunner, +) +from sglang.srt.utils import get_available_gpu_memory + +logger = logging.getLogger(__name__) + +if TYPE_CHECKING: + from sglang.srt.speculative.multi_layer_eagle_worker import ( + MultiLayerEagleDraftWorker, + ) + + +class MultiLayerEagleDraftExtendNpuGraphRunner( + MultiLayerEagleDraftExtendCudaGraphRunner +): + def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int): + super().__init__(eagle_worker, step) + + def _create_graph(self): + return torch.npu.NPUGraph() + + def _capture_init(self, run_once_fn): + for _ in range(2): + torch.npu.synchronize() + self.model_runner.tp_group.barrier() + run_once_fn() + + def _capture_graph(self, graph, pool, stream, run_once_fn): + with torch.npu.graph( + graph, + pool=pool, + stream=stream, + auto_dispatch_capture=True, + ): + out = run_once_fn() + return out + + def _replay(self, forward_batch: ForwardBatch): + seq_lens = self.buffers.seq_lens_cpu[: self.raw_bs].tolist() + [0] * ( + self.bs - self.raw_bs + ) + thread = threading.Thread( + target=self.graphs[self.bs].update, + kwargs={"cpu_update_input": [{"actual_seq_kvlen": seq_lens}]}, + ) + thread.start() + self.graphs[self.bs].replay() + thread.join() + + +class MultiLayerEagleMultiStepDraftExtendNpuGraphRunner( + MultiLayerEagleMultiStepDraftExtendCudaGraphRunner +): + def __init__(self, eagle_worker: MultiLayerEagleDraftWorker): + super().__init__(eagle_worker) + + def _init_and_capture(self): + if self.eagle_worker.server_args.disable_cuda_graph: + self.runners = [None] * self.speculative_num_steps + return + + self.runners: List[Optional[MultiLayerEagleDraftExtendNpuGraphRunner]] = [] + buffer_len_list: List[int] = [] + + for step in range(self.speculative_num_steps): + if self.draft_extend_attn_backend_list[step]: + runner = MultiLayerEagleDraftExtendNpuGraphRunner( + self.eagle_worker, step + ) + self.runners.append(runner) + + self.seq_len_fill_value = runner.seq_len_fill_value + self.max_bs = runner.max_bs + buffer_len_list.append(runner.max_num_token) + self.offsets.append(self.offsets[-1] + runner.max_num_token) + else: + self.runners.append(None) + + self.cuda_graph_buffers["seq_lens_cpu"] = torch.full( + (self.max_bs,), + self.seq_len_fill_value, + dtype=torch.int32, + ) + + with torch.device(self.device): + self.cuda_graph_buffers["input_ids"] = torch.zeros( + (self.offsets[-1],), dtype=torch.int64 + ) + self.cuda_graph_buffers["out_cache_loc"] = torch.ones( + (self.offsets[-1],), dtype=torch.int64 + ) + self.cuda_graph_buffers["positions"] = torch.zeros( + (self.offsets[-1],), dtype=torch.int64 + ) + + self.cuda_graph_buffers["seq_lens"] = torch.full( + (self.max_bs,), + self.seq_len_fill_value, + dtype=torch.int32, + ) + self.cuda_graph_buffers["req_pool_indices"] = torch.zeros( + (self.max_bs,), dtype=torch.int64 + ) + self.cuda_graph_buffers["num_correct_drafts"] = torch.full( + (self.max_bs,), 1, dtype=torch.int32 + ) + self.cuda_graph_buffers["num_accept_tokens"] = torch.full( + (self.max_bs,), 1, dtype=torch.int32 + ) + + for step in range(self.speculative_num_steps - 1, -1, -1): + if self.runners[step] is not None: + tic = time.perf_counter() + before_mem = get_available_gpu_memory(self.device, self.gpu_id) + logger.info( + f"Capture draft extend cuda graph begin (step {step}). This can take up to several minutes. avail mem={before_mem:.2f} GB" + ) + + self.runners[step].init_buffers_and_capture( + self.cuda_graph_buffers, + self.offsets[step], + ( + self.runners[step + 1] + if step + 1 < self.speculative_num_steps + else None + ), + ) + + after_mem = get_available_gpu_memory(self.device, self.gpu_id) + logger.info( + f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB." + ) diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py index 9717414ef..1e1df9598 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py @@ -94,21 +94,28 @@ class NPUGraphRunner(CudaGraphRunner): self.model_runner = model_runner self._init_arch_map() self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") + self.if_use_v2 = any( + arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM") + for arch in (model_runner.model_config.hf_config.architectures or []) + ) def _init_arch_map(self): if self.is_dllm: self.attr_name: Dict[str, str] = { AttentionArch.MLA: "actual_seq_lengths_kv", AttentionArch.MHA: "actual_seq_lengths_kv", + "TARGET_VERIFY": "actual_seq_kvlen", } else: self.attr_name: Dict[str, str] = { AttentionArch.MLA: "actual_seq_lengths_kv", AttentionArch.MHA: "context_lens", + "TARGET_VERIFY": "actual_seq_kvlen", } self.attr_type: Dict[str, Union[list, torch.Tensor]] = { AttentionArch.MLA: [], AttentionArch.MHA: torch.Tensor(), + "TARGET_VERIFY": [], } def _create_device_graph(self): @@ -133,9 +140,13 @@ class NPUGraphRunner(CudaGraphRunner): return out def _get_update_attr_name(self): + if self.if_use_v2: + return self.attr_name["TARGET_VERIFY"] return self.attr_name[AttentionArch.MLA] def _get_update_attr_type(self): + if self.if_use_v2: + return self.attr_type["TARGET_VERIFY"] return self.attr_type[AttentionArch.MLA] def _update_inputs(self, seq_lens): diff --git a/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py b/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py index c9f55ba12..92ba98667 100644 --- a/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py +++ b/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py @@ -54,6 +54,10 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): layer_num: int, device: str, enable_memory_saver: bool, + v_head_dim: Optional[int] = None, + swa_head_num: Optional[int] = None, + swa_head_dim: Optional[int] = None, + swa_v_head_dim: Optional[int] = None, start_layer: Optional[int] = None, end_layer: Optional[int] = None, enable_alt_stream: bool = True, @@ -69,6 +73,10 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): layer_num=layer_num, device=device, enable_memory_saver=enable_memory_saver, + v_head_dim=v_head_dim, + swa_head_num=swa_head_num, + swa_head_dim=swa_head_dim, + swa_v_head_dim=swa_v_head_dim, start_layer=start_layer, end_layer=end_layer, enable_alt_stream=enable_alt_stream, @@ -81,9 +89,8 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): # The padded slot 0 is used for writing dummy outputs from padded tokens. # Continuous memory improves the efficiency of Ascend`s transmission backend, # while other backends remain unchanged. - self.kv_buffer = torch.zeros( + self.k_buffer = torch.zeros( ( - 2, self.layer_num, self.size // self.page_size + 1, self.page_size, @@ -93,21 +100,31 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): dtype=self.store_dtype, device=self.device, ) - self.k_buffer = self.kv_buffer[0] - self.v_buffer = self.kv_buffer[1] + self.v_buffer = torch.zeros( + ( + self.layer_num, + self.size // self.page_size + 1, + self.page_size, + self.head_num, + self.v_head_dim, + ), + dtype=self.store_dtype, + device=self.device, + ) if self.use_fia: - self.k_buffer = [] - self.v_buffer = [] - for i in range(self.layer_num): - k_buffer_layer = self.kv_buffer[0][i].view( - -1, 1, self.head_num, self.head_dim - ) - v_buffer_layer = self.kv_buffer[1][i].view( - -1, 1, self.head_num, self.head_dim - ) - self.k_buffer.append(k_buffer_layer) - self.v_buffer.append(v_buffer_layer) + # Use per-layer Python lists to avoid torch.compile capturing + # the entire multi-layer tensor (OOM during graph capture). + # Each layer view: [P*ps, 1, H, D], sharing the contiguous + # storage allocated above. + self.k_buffer = [ + self.k_buffer[i].view(-1, 1, self.head_num, self.head_dim) + for i in range(self.layer_num) + ] + self.v_buffer = [ + self.v_buffer[i].view(-1, 1, self.head_num, self.v_head_dim) + for i in range(self.layer_num) + ] def _init_kv_copy_and_warmup(self): # implementation relies on self.data_strides / self.data_ptrs, which the @@ -188,7 +205,7 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): torch_npu.npu_scatter_nd_update_( v_buffer_layer, loc.view(-1, 1), - cache_v.view(-1, 1, self.head_num, self.head_dim), + cache_v.view(-1, 1, self.head_num, self.v_head_dim), ) else: loc = loc.to(torch.int32) @@ -199,7 +216,7 @@ class NPUMHATokenToKVPool(MHATokenToKVPool): -1, self.page_size, self.head_num, self.head_dim ), value_cache=self.v_buffer[layer_id - self.start_layer].view( - -1, self.page_size, self.head_num, self.head_dim + -1, self.page_size, self.head_num, self.v_head_dim ), slot_indices=loc, ) diff --git a/python/sglang/srt/hardware_backend/npu/quantization/fused_moe_method_npu.py b/python/sglang/srt/hardware_backend/npu/quantization/fused_moe_method_npu.py index 5e4d3c423..d3b18e727 100644 --- a/python/sglang/srt/hardware_backend/npu/quantization/fused_moe_method_npu.py +++ b/python/sglang/srt/hardware_backend/npu/quantization/fused_moe_method_npu.py @@ -741,7 +741,7 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase): hidden_states = torch.ops.npu.npu_grouped_matmul( x=[hidden_states], weight=[layer.w2_weight], - scale=[layer.w2_weight_scale.to(output_dtype)], + scale=[layer.w2_weight_scale_bf16], per_token_scale=[swiglu_out_scale], split_item=2, group_list_type=group_list_type, diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index 2d4a240be..cf264f7f5 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -268,11 +268,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) assert alloc_swa_indices is not None - self.full_to_swa_index_mapping[alloc_full_indices[-swa_tail_len:]] = ( - alloc_swa_indices - ) + self.full_to_swa_index_mapping[ + alloc_full_indices[-swa_tail_len:].to(torch.int64) + ] = alloc_swa_indices.to(torch.int64) if swa_tail_len < extend_num_tokens: - self.full_to_swa_index_mapping[alloc_full_indices[:-swa_tail_len]] = 0 + self.full_to_swa_index_mapping[ + alloc_full_indices[:-swa_tail_len].to(torch.int64) + ] = 0 return alloc_full_indices def alloc_decode( diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index e57dfc003..c32433edd 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -505,9 +505,6 @@ class ModelRunnerKVCacheMixin: enable_kvcache_transpose=False, device=self.device, token_to_kv_pool_class=NPUMHATokenToKVPool, - enable_kv_cache_copy=( - self.server_args.speculative_algorithm is not None - ), **kwargs, ) elif self.use_mla_backend: diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index 3b0244b32..4952148d9 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -282,7 +282,11 @@ class MiMoV2MoE(nn.Module): ) # todo : implement tbo forward needed - if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake(): + if ( + get_moe_a2a_backend().is_deepep() + or get_moe_a2a_backend().is_mooncake() + or get_moe_a2a_backend().is_ascend_fuseep() + ): # TODO: we will support tp < ep in the future self.ep_size = get_moe_expert_parallel_world_size() self.num_experts = ( @@ -299,7 +303,9 @@ class MiMoV2MoE(nn.Module): ) self._enable_a2a_moe = ( - get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake() + get_moe_a2a_backend().is_deepep() + or get_moe_a2a_backend().is_mooncake() + or get_moe_a2a_backend().is_ascend_fuseep() ) def get_moe_weights(self): diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 04b3841a2..56bf32705 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -19,6 +19,9 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import torch from sglang.srt.environ import envs +from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import ( + MultiLayerEagleMultiStepDraftExtendNpuGraphRunner, +) from sglang.srt.layers.moe.utils import speculative_moe_backend_context from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs from sglang.srt.managers.io_struct import ( @@ -52,6 +55,7 @@ from sglang.srt.speculative.spec_utils import ( record_stream_for_v2_verify, select_top_k_tokens, ) +from sglang.srt.utils import is_npu from sglang.srt.utils.async_probe import ( maybe_detect_inf, maybe_detect_nan, @@ -59,6 +63,8 @@ from sglang.srt.utils.async_probe import ( ) from sglang.srt.utils.common import empty_context, fast_topk +_is_npu = is_npu() + if TYPE_CHECKING: from sglang.srt.model_executor.model_runner import ModelRunner, ModelRunnerOutput @@ -227,9 +233,14 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): if self.server_args.disable_cuda_graph: return - self.cuda_graph_runner_for_draft_extend = ( - MultiLayerEagleMultiStepDraftExtendCudaGraphRunner(self) - ) + if not _is_npu: + self.cuda_graph_runner_for_draft_extend = ( + MultiLayerEagleMultiStepDraftExtendCudaGraphRunner(self) + ) + else: + self.cuda_graph_runner_for_draft_extend = ( + MultiLayerEagleMultiStepDraftExtendNpuGraphRunner(self) + ) def reset_cuda_graph_buffers(self, forward_batch, batch_result): if self.cuda_graph_runner_for_draft_extend: