[NPU] MiMo-V2-Flash Adaptation (#25455)
This commit is contained in:
@@ -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,
|
||||
|
||||
+159
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user