[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__)
|
logger = logging.getLogger(__name__)
|
||||||
SWA_INT_MAX = 2147483647
|
|
||||||
|
# default max value of full attention window size
|
||||||
|
FULL_ATTENTION_WINDOW = 2147483647
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -1044,6 +1046,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
save_kv_cache,
|
save_kv_cache,
|
||||||
q_rope=q_rope,
|
q_rope=q_rope,
|
||||||
k_rope=k_rope,
|
k_rope=k_rope,
|
||||||
|
sinks=sinks,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
@@ -1106,10 +1109,12 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
pre_tokens=(
|
pre_tokens=(
|
||||||
layer.sliding_window_size
|
layer.sliding_window_size
|
||||||
if layer.sliding_window_size != -1
|
if layer.sliding_window_size != -1
|
||||||
else SWA_INT_MAX
|
else FULL_ATTENTION_WINDOW
|
||||||
),
|
),
|
||||||
next_tokens=(
|
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,
|
atten_mask=self.fia_mask,
|
||||||
block_table=block_tables,
|
block_table=block_tables,
|
||||||
@@ -1637,6 +1642,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
save_kv_cache: bool,
|
save_kv_cache: bool,
|
||||||
q_rope: Optional[torch.Tensor] = None,
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
k_rope: Optional[torch.Tensor] = None,
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
|
sinks: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
@@ -1658,16 +1664,22 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
-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()
|
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous()
|
||||||
|
|
||||||
if not self.graph_mode:
|
if not self.graph_mode:
|
||||||
num_token_padding = query.shape[0]
|
num_token_padding = query.shape[0]
|
||||||
query = query[: forward_batch.num_token_non_padded_cpu]
|
query = query[: forward_batch.num_token_non_padded_cpu]
|
||||||
|
|
||||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||||
actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list
|
actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list
|
||||||
else:
|
else:
|
||||||
actual_seq_lengths_kv = (
|
actual_seq_lengths_kv = (
|
||||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
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 = (
|
actual_seq_lengths = (
|
||||||
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
|
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 + query.shape[0],
|
||||||
self.speculative_num_draft_tokens,
|
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:
|
if layer.attn_type == AttentionType.ENCODER_ONLY:
|
||||||
mask = None
|
mask = None
|
||||||
sparse_mode = 0
|
sparse_mode = 0
|
||||||
else:
|
else:
|
||||||
mask = self.mtp_mask
|
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,
|
query,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
block_table=self.forward_metadata.block_tables,
|
block_table=block_table,
|
||||||
block_size=self.page_size,
|
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,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
input_layout="TND",
|
input_layout="TND",
|
||||||
atten_mask=mask,
|
atten_mask=mask,
|
||||||
scale=layer.scaling,
|
softmax_scale=layer.scaling,
|
||||||
actual_seq_lengths=actual_seq_lengths,
|
actual_seq_qlen=actual_seq_lengths,
|
||||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
actual_seq_kvlen=actual_seq_lengths_kv,
|
||||||
sparse_mode=sparse_mode,
|
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)
|
attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
if (
|
if (
|
||||||
@@ -1840,27 +1868,89 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if sinks is not None:
|
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
|
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||||
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
block_tables = self.forward_metadata.block_tables_swa
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
else:
|
else:
|
||||||
block_tables = self.forward_metadata.block_tables
|
block_tables = self.forward_metadata.block_tables
|
||||||
attn_out = attention_sinks_triton(
|
if self.use_fia:
|
||||||
q,
|
k_cache = (
|
||||||
k_cache,
|
self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
v_cache,
|
.view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim)
|
||||||
sinks,
|
.contiguous()
|
||||||
block_tables,
|
)
|
||||||
self.forward_metadata.seq_lens,
|
v_cache = (
|
||||||
layer.scaling,
|
self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
layer.sliding_window_size,
|
.view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
||||||
layer.tp_q_head_num,
|
.contiguous()
|
||||||
layer.tp_k_head_num,
|
)
|
||||||
)
|
query = q.reshape(
|
||||||
return attn_out
|
-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:
|
if not self.use_mla:
|
||||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
|
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 = (
|
actual_seq_len_kv = (
|
||||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
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]
|
num_tokens = query.shape[0]
|
||||||
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
||||||
query,
|
query,
|
||||||
@@ -2072,12 +2216,17 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||||
)
|
)
|
||||||
block_size = self.page_size
|
block_size = self.page_size
|
||||||
max_model_len = block_tables.shape[-1] * block_size
|
|
||||||
swa_mask = self.ascend_attn_mask_builder.get_swa_mask(
|
if sinks is not None:
|
||||||
self.forward_metadata.seq_lens,
|
mask = self.fia_mask
|
||||||
max_model_len,
|
else:
|
||||||
layer.sliding_window_size,
|
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(
|
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||||
q.view(
|
q.view(
|
||||||
forward_batch.batch_size,
|
forward_batch.batch_size,
|
||||||
@@ -2095,14 +2244,14 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
num_key_value_heads=layer.tp_k_head_num,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
input_layout="BSND",
|
input_layout="BSND",
|
||||||
block_size=block_size,
|
block_size=block_size,
|
||||||
atten_mask=(
|
atten_mask=(mask if layer.sliding_window_size != -1 else None),
|
||||||
swa_mask if layer.sliding_window_size != -1 else None
|
|
||||||
),
|
|
||||||
sparse_mode=4 if layer.sliding_window_size != -1 else 0,
|
sparse_mode=4 if layer.sliding_window_size != -1 else 0,
|
||||||
softmax_scale=layer.scaling,
|
softmax_scale=layer.scaling,
|
||||||
block_table=block_tables,
|
block_table=block_tables,
|
||||||
actual_seq_qlen=[1] * len(self.forward_metadata.seq_lens),
|
actual_seq_qlen=[1] * len(self.forward_metadata.seq_lens),
|
||||||
actual_seq_kvlen=actual_seq_len_kv,
|
actual_seq_kvlen=actual_seq_len_kv,
|
||||||
|
pre_tokens=layer.sliding_window_size,
|
||||||
|
next_tokens=0,
|
||||||
learnable_sink=sinks,
|
learnable_sink=sinks,
|
||||||
)
|
)
|
||||||
attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
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
|
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||||
),
|
),
|
||||||
v_cache.view(
|
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_heads=layer.tp_q_head_num,
|
||||||
num_key_value_heads=layer.tp_k_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.model_runner = model_runner
|
||||||
self._init_arch_map()
|
self._init_arch_map()
|
||||||
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
|
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):
|
def _init_arch_map(self):
|
||||||
if self.is_dllm:
|
if self.is_dllm:
|
||||||
self.attr_name: Dict[str, str] = {
|
self.attr_name: Dict[str, str] = {
|
||||||
AttentionArch.MLA: "actual_seq_lengths_kv",
|
AttentionArch.MLA: "actual_seq_lengths_kv",
|
||||||
AttentionArch.MHA: "actual_seq_lengths_kv",
|
AttentionArch.MHA: "actual_seq_lengths_kv",
|
||||||
|
"TARGET_VERIFY": "actual_seq_kvlen",
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
self.attr_name: Dict[str, str] = {
|
self.attr_name: Dict[str, str] = {
|
||||||
AttentionArch.MLA: "actual_seq_lengths_kv",
|
AttentionArch.MLA: "actual_seq_lengths_kv",
|
||||||
AttentionArch.MHA: "context_lens",
|
AttentionArch.MHA: "context_lens",
|
||||||
|
"TARGET_VERIFY": "actual_seq_kvlen",
|
||||||
}
|
}
|
||||||
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
|
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
|
||||||
AttentionArch.MLA: [],
|
AttentionArch.MLA: [],
|
||||||
AttentionArch.MHA: torch.Tensor(),
|
AttentionArch.MHA: torch.Tensor(),
|
||||||
|
"TARGET_VERIFY": [],
|
||||||
}
|
}
|
||||||
|
|
||||||
def _create_device_graph(self):
|
def _create_device_graph(self):
|
||||||
@@ -133,9 +140,13 @@ class NPUGraphRunner(CudaGraphRunner):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
def _get_update_attr_name(self):
|
def _get_update_attr_name(self):
|
||||||
|
if self.if_use_v2:
|
||||||
|
return self.attr_name["TARGET_VERIFY"]
|
||||||
return self.attr_name[AttentionArch.MLA]
|
return self.attr_name[AttentionArch.MLA]
|
||||||
|
|
||||||
def _get_update_attr_type(self):
|
def _get_update_attr_type(self):
|
||||||
|
if self.if_use_v2:
|
||||||
|
return self.attr_type["TARGET_VERIFY"]
|
||||||
return self.attr_type[AttentionArch.MLA]
|
return self.attr_type[AttentionArch.MLA]
|
||||||
|
|
||||||
def _update_inputs(self, seq_lens):
|
def _update_inputs(self, seq_lens):
|
||||||
|
|||||||
@@ -54,6 +54,10 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
layer_num: int,
|
layer_num: int,
|
||||||
device: str,
|
device: str,
|
||||||
enable_memory_saver: bool,
|
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,
|
start_layer: Optional[int] = None,
|
||||||
end_layer: Optional[int] = None,
|
end_layer: Optional[int] = None,
|
||||||
enable_alt_stream: bool = True,
|
enable_alt_stream: bool = True,
|
||||||
@@ -69,6 +73,10 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
layer_num=layer_num,
|
layer_num=layer_num,
|
||||||
device=device,
|
device=device,
|
||||||
enable_memory_saver=enable_memory_saver,
|
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,
|
start_layer=start_layer,
|
||||||
end_layer=end_layer,
|
end_layer=end_layer,
|
||||||
enable_alt_stream=enable_alt_stream,
|
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.
|
# The padded slot 0 is used for writing dummy outputs from padded tokens.
|
||||||
# Continuous memory improves the efficiency of Ascend`s transmission backend,
|
# Continuous memory improves the efficiency of Ascend`s transmission backend,
|
||||||
# while other backends remain unchanged.
|
# while other backends remain unchanged.
|
||||||
self.kv_buffer = torch.zeros(
|
self.k_buffer = torch.zeros(
|
||||||
(
|
(
|
||||||
2,
|
|
||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.size // self.page_size + 1,
|
self.size // self.page_size + 1,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
@@ -93,21 +100,31 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
dtype=self.store_dtype,
|
dtype=self.store_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
self.k_buffer = self.kv_buffer[0]
|
self.v_buffer = torch.zeros(
|
||||||
self.v_buffer = self.kv_buffer[1]
|
(
|
||||||
|
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:
|
if self.use_fia:
|
||||||
self.k_buffer = []
|
# Use per-layer Python lists to avoid torch.compile capturing
|
||||||
self.v_buffer = []
|
# the entire multi-layer tensor (OOM during graph capture).
|
||||||
for i in range(self.layer_num):
|
# Each layer view: [P*ps, 1, H, D], sharing the contiguous
|
||||||
k_buffer_layer = self.kv_buffer[0][i].view(
|
# storage allocated above.
|
||||||
-1, 1, self.head_num, self.head_dim
|
self.k_buffer = [
|
||||||
)
|
self.k_buffer[i].view(-1, 1, self.head_num, self.head_dim)
|
||||||
v_buffer_layer = self.kv_buffer[1][i].view(
|
for i in range(self.layer_num)
|
||||||
-1, 1, self.head_num, self.head_dim
|
]
|
||||||
)
|
self.v_buffer = [
|
||||||
self.k_buffer.append(k_buffer_layer)
|
self.v_buffer[i].view(-1, 1, self.head_num, self.v_head_dim)
|
||||||
self.v_buffer.append(v_buffer_layer)
|
for i in range(self.layer_num)
|
||||||
|
]
|
||||||
|
|
||||||
def _init_kv_copy_and_warmup(self):
|
def _init_kv_copy_and_warmup(self):
|
||||||
# implementation relies on self.data_strides / self.data_ptrs, which the
|
# implementation relies on self.data_strides / self.data_ptrs, which the
|
||||||
@@ -188,7 +205,7 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
torch_npu.npu_scatter_nd_update_(
|
torch_npu.npu_scatter_nd_update_(
|
||||||
v_buffer_layer,
|
v_buffer_layer,
|
||||||
loc.view(-1, 1),
|
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:
|
else:
|
||||||
loc = loc.to(torch.int32)
|
loc = loc.to(torch.int32)
|
||||||
@@ -199,7 +216,7 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
-1, self.page_size, self.head_num, self.head_dim
|
-1, self.page_size, self.head_num, self.head_dim
|
||||||
),
|
),
|
||||||
value_cache=self.v_buffer[layer_id - self.start_layer].view(
|
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,
|
slot_indices=loc,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -741,7 +741,7 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
||||||
x=[hidden_states],
|
x=[hidden_states],
|
||||||
weight=[layer.w2_weight],
|
weight=[layer.w2_weight],
|
||||||
scale=[layer.w2_weight_scale.to(output_dtype)],
|
scale=[layer.w2_weight_scale_bf16],
|
||||||
per_token_scale=[swiglu_out_scale],
|
per_token_scale=[swiglu_out_scale],
|
||||||
split_item=2,
|
split_item=2,
|
||||||
group_list_type=group_list_type,
|
group_list_type=group_list_type,
|
||||||
|
|||||||
@@ -268,11 +268,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
assert alloc_swa_indices is not None
|
assert alloc_swa_indices is not None
|
||||||
|
|
||||||
self.full_to_swa_index_mapping[alloc_full_indices[-swa_tail_len:]] = (
|
self.full_to_swa_index_mapping[
|
||||||
alloc_swa_indices
|
alloc_full_indices[-swa_tail_len:].to(torch.int64)
|
||||||
)
|
] = alloc_swa_indices.to(torch.int64)
|
||||||
if swa_tail_len < extend_num_tokens:
|
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
|
return alloc_full_indices
|
||||||
|
|
||||||
def alloc_decode(
|
def alloc_decode(
|
||||||
|
|||||||
@@ -505,9 +505,6 @@ class ModelRunnerKVCacheMixin:
|
|||||||
enable_kvcache_transpose=False,
|
enable_kvcache_transpose=False,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
token_to_kv_pool_class=NPUMHATokenToKVPool,
|
token_to_kv_pool_class=NPUMHATokenToKVPool,
|
||||||
enable_kv_cache_copy=(
|
|
||||||
self.server_args.speculative_algorithm is not None
|
|
||||||
),
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
elif self.use_mla_backend:
|
elif self.use_mla_backend:
|
||||||
|
|||||||
@@ -282,7 +282,11 @@ class MiMoV2MoE(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# todo : implement tbo forward needed
|
# 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
|
# TODO: we will support tp < ep in the future
|
||||||
self.ep_size = get_moe_expert_parallel_world_size()
|
self.ep_size = get_moe_expert_parallel_world_size()
|
||||||
self.num_experts = (
|
self.num_experts = (
|
||||||
@@ -299,7 +303,9 @@ class MiMoV2MoE(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self._enable_a2a_moe = (
|
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):
|
def get_moe_weights(self):
|
||||||
|
|||||||
@@ -19,6 +19,9 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
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.moe.utils import speculative_moe_backend_context
|
||||||
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
@@ -52,6 +55,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
record_stream_for_v2_verify,
|
record_stream_for_v2_verify,
|
||||||
select_top_k_tokens,
|
select_top_k_tokens,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils import is_npu
|
||||||
from sglang.srt.utils.async_probe import (
|
from sglang.srt.utils.async_probe import (
|
||||||
maybe_detect_inf,
|
maybe_detect_inf,
|
||||||
maybe_detect_nan,
|
maybe_detect_nan,
|
||||||
@@ -59,6 +63,8 @@ from sglang.srt.utils.async_probe import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.utils.common import empty_context, fast_topk
|
from sglang.srt.utils.common import empty_context, fast_topk
|
||||||
|
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner, ModelRunnerOutput
|
from sglang.srt.model_executor.model_runner import ModelRunner, ModelRunnerOutput
|
||||||
|
|
||||||
@@ -227,9 +233,14 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
if self.server_args.disable_cuda_graph:
|
if self.server_args.disable_cuda_graph:
|
||||||
return
|
return
|
||||||
|
|
||||||
self.cuda_graph_runner_for_draft_extend = (
|
if not _is_npu:
|
||||||
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner(self)
|
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):
|
def reset_cuda_graph_buffers(self, forward_batch, batch_result):
|
||||||
if self.cuda_graph_runner_for_draft_extend:
|
if self.cuda_graph_runner_for_draft_extend:
|
||||||
|
|||||||
Reference in New Issue
Block a user