[NPU] MiMo-V2-Flash Adaptation (#25455)

This commit is contained in:
iridiumine
2026-06-10 09:13:55 +08:00
committed by GitHub
parent f3ecc3688f
commit 2947781ce6
9 changed files with 419 additions and 67 deletions
@@ -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,14 +1868,76 @@ 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
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,
@@ -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
if sinks is not None:
mask = self.fia_mask
else:
max_model_len = block_tables.shape[-1] * block_size
swa_mask = self.ascend_attn_mask_builder.get_swa_mask(
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,
@@ -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."
)
@@ -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):
@@ -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,
)
@@ -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,
+6 -4
View File
@@ -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(
@@ -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:
+8 -2
View File
@@ -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):
@@ -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
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: