openPangu-R-7B-2512 / inference /vllm-ascend_v0.11.0rc0.patch
drizzlezyk's picture
Upload inference/vllm-ascend_v0.11.0rc0.patch with huggingface_hub
cb4f943 verified
Raw
History Blame Contribute Delete
42.9 kB
diff --git a/vllm_ascend/attention/attention_v1.py b/vllm_ascend/attention/attention_v1.py
index d289bb4..0357b50 100644
--- a/vllm_ascend/attention/attention_v1.py
+++ b/vllm_ascend/attention/attention_v1.py
@@ -21,6 +21,7 @@ from typing import ClassVar, List, Optional, Tuple, Type
import torch
import torch.nn as nn
+import torch.nn.functional as F
import torch_npu
from vllm.attention.backends.abstract import (AttentionBackend, AttentionImpl,
AttentionLayer, AttentionType)
@@ -30,6 +31,7 @@ from vllm.utils import cdiv, direct_register_custom_op
from vllm.v1.attention.backends.utils import AttentionCGSupport
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import AttentionSpec
+from vllm.model_executor.models.utils import extract_layer_index
from vllm_ascend.attention.utils import (AscendCommonAttentionMetadata,
maybe_save_kv_layer_to_connector,
@@ -39,6 +41,9 @@ from vllm_ascend.ops.attention import vanilla_chunked_prefill
from vllm_ascend.utils import (ACL_FORMAT_FRACTAL_NZ, aligned_16, is_310p,
nd_to_nz_2d, nd_to_nz_spec)
+if torch.version.cann.startswith("8.3"):
+ import omni_custom_ops
+
class AscendAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@@ -115,6 +120,7 @@ class AscendAttentionBackend(AttentionBackend):
return [64]
+
class AscendAttentionState(Enum):
PrefillNoCache = 0
PrefillCacheHit = 1
@@ -135,8 +141,8 @@ class AscendMetadata:
num_actual_tokens: int = 0
# The sequence length per sequence. Sequence length means the computed
- # tokens + new tokens (is None if it is a decoding).
- # (batch_size,)
+ # tokens + new tokens (is None if it is a decoding).(batch_size,)
+
seq_lens: torch.Tensor = None
query_start_loc: torch.Tensor = None
@@ -145,20 +151,25 @@ class AscendMetadata:
max_query_len: Optional[int] = None
# ********************** KV Cache Related Properties ********************* #
- # Block addresses per sequence (Seq id -> list of physical block).
- # (batch_size, max_blocks_per_seq)
+ # Block addresses per sequence (Seq id -> list of physical block).(batch_size, max_blocks_per_seq)
+
block_tables: torch.Tensor = None
# The indices of the token slots that input tokens will be stored into.
# E.g., if `slot_mapping` is [35, 2, 17] and the block size is 16, the
# three tokens are stored in the 3rd slot in block 2, 2nd slot in block 0,
- # and 1st slot in block 1, respectively.
- # (num_tokens,)
+ # and 1st slot in block 1, respectively. (num_tokens,)
+
slot_mapping: torch.Tensor = None
# *************************** Other Properties *************************** #
enable_dbo_across_dp: bool = False
+ # Patch for param sink
+ sink_block_tables: Optional[List[torch.Tensor]] = None
+ sink_attn_mask: Optional[torch.Tensor] = None
+ sink_seq_kvlens: torch.Tensor = None
+ swa_seq_qlens: torch.Tensor = None
class AscendAttentionMetadataBuilder:
# Does this backend/builder support ACL Graphs for attention (default: no).
@@ -182,6 +193,7 @@ class AscendAttentionMetadataBuilder:
self.max_num_blocks_per_req = cdiv(
self.model_config.max_model_len,
AscendAttentionBackend.get_supported_block_size()[0])
+ self.param_sink_number = self.model_config.hf_config.param_sink_number
def reorder_batch(self, input_batch,
scheduler_output: "SchedulerOutput") -> bool:
@@ -210,6 +222,33 @@ class AscendAttentionMetadataBuilder:
query_start_loc = query_start_loc_cpu.to(self.device,
non_blocking=True)
+ num_input_tokens = common_attn_metadata.num_input_tokens
+
+
+ if num_input_tokens > num_reqs and attn_state == AscendAttentionState.DecodeOnly:
+ tokens_gap_num = num_input_tokens-num_reqs
+
+ sink_block_tables = F.pad(block_table, (1, 0, 0, tokens_gap_num), value=0)
+
+ sink_seq_kvlens = seq_lens + self.param_sink_number
+ sink_seq_kvlens = torch.cat([sink_seq_kvlens, torch.full((tokens_gap_num,), \
+ self.param_sink_number, dtype=torch.int32)], dim=0)
+
+ gap_query_lens = torch.cat([query_lens, torch.ones(tokens_gap_num, dtype=torch.int32)], dim=0)
+ swa_seq_qlens = torch.cumsum(gap_query_lens, dim=0).to(dtype=torch.int32)
+ else:
+ sink_block_tables = F.pad(block_table, (1, 0, 0, 0), value=0)
+ sink_seq_kvlens = seq_lens + self.param_sink_number
+ swa_seq_qlens = torch.cumsum(query_lens, dim=0).to(dtype=torch.int32)
+
+
+ if attn_mask is not None:
+ sink_attn_mask = F.pad(attn_mask, (self.param_sink_number, 0, 0, 0), value=0)
+ else:
+ sink_attn_mask = None
+
+
+
if is_310p():
if attn_state == AscendAttentionState.PrefillNoCache:
mask_nz = nd_to_nz_2d(attn_mask)
@@ -230,7 +269,12 @@ class AscendAttentionMetadataBuilder:
slot_mapping=slot_mapping,
attn_mask=attn_mask,
attn_state=attn_state,
- enable_dbo_across_dp=common_attn_metadata.enable_dbo_across_dp)
+ enable_dbo_across_dp=common_attn_metadata.enable_dbo_across_dp,
+ sink_block_tables=sink_block_tables,
+ sink_attn_mask=sink_attn_mask,
+ sink_seq_kvlens=sink_seq_kvlens,
+ swa_seq_qlens=swa_seq_qlens
+ )
return attn_metadata
def build_for_graph_capture(
@@ -265,6 +309,7 @@ class AscendAttentionBackendImpl(AttentionImpl):
kv_cache_dtype: str,
logits_soft_cap: Optional[float],
attn_type: str,
+ layer_name: str,
kv_sharing_target_layer_name: Optional[str],
**kwargs,
) -> None:
@@ -287,6 +332,13 @@ class AscendAttentionBackendImpl(AttentionImpl):
self.key_cache = None
self.value_cache = None
+ self.layer_idx = extract_layer_index(layer_name)
+
+ # Patch for Sink
+ self.sink_cached = False
+ self.attn_mask = torch.ones((2048, 2048), dtype=torch.int8, device="npu").triu_(diagonal=1)
+ self.attn_mask = self.attn_mask.to(torch.bool)
+
def _forward_prefill_no_cache(
self,
query: torch.Tensor,
@@ -295,6 +347,7 @@ class AscendAttentionBackendImpl(AttentionImpl):
attn_metadata: AscendMetadata,
output: Optional[torch.Tensor] = None,
num_tokens=0,
+ param_sink_number: Optional[int] = 0
) -> torch.Tensor:
assert attn_metadata is not None
assert attn_metadata.attn_mask is not None
@@ -311,18 +364,72 @@ class AscendAttentionBackendImpl(AttentionImpl):
mask = mask.repeat(attn_metadata.seq_lens.size(0), 1, 1, 1)
mask = torch_npu.npu_format_cast(mask.contiguous(),
ACL_FORMAT_FRACTAL_NZ)
+ if torch.version.cann.startswith("8.3"):
+ mask = torch.ones((2048, 2048), dtype=torch.int8, device=mask.device).triu_(diagonal=1)
+ # TODO: nocache swa
+ if param_sink_number > 0:
+ query_lens = attn_metadata.query_lens
+ seq_lens = attn_metadata.seq_lens + param_sink_number
+ output, _ = torch.ops.custom.npu_fused_infer_attention_sink(
+ query,
+ key,
+ value,
+ atten_mask=mask,
+ actual_seq_qlen=query_lens,
+ actual_seq_kvlen=seq_lens,
+ num_query_heads=self.num_heads,
+ num_key_value_heads=self.num_kv_heads,
+ input_layout='TND',
+ sparse_mode=3,
+ sink_number=param_sink_number,
+ softmax_scale=self.scale,
+ )
+ else:
+ output, _ = torch_npu.npu_fused_infer_attention_score(
+ query=query,
+ key=key,
+ value=value,
+ atten_mask=mask,
+ input_layout="TND",
+ actual_seq_lengths=attn_metadata.query_start_loc[1:],
+ actual_seq_lengths_kv=attn_metadata.seq_lens,
+ num_key_value_heads=self.num_kv_heads,
+ num_heads=self.num_heads,
+ scale=self.scale,
+ sparse_mode=3,
+ )
+ return output
+ # Patch for sink on CANN8.2
+ if param_sink_number > 0:
+ seq_lens = attn_metadata.seq_lens + param_sink_number
+ # TODO: _npu_flash_attention only allows qlen==kvlen,
+ mask_elem = mask[0, -1]
+ sink_mask = torch.full((mask.size(0) + param_sink_number,
+ mask.size(1) + param_sink_number),
+ mask_elem, dtype=mask.dtype, device=mask.device)
+ sink_mask[param_sink_number:, :param_sink_number] = 0.0
+ sink_mask[param_sink_number:, param_sink_number:] = mask
+ sink_mask[:param_sink_number, :param_sink_number].triu_(diagonal=1)
+ mask = sink_mask
+
+ output = torch.zeros((output.size(0) + param_sink_number,
+ output.size(1), output.size(2)),
+ dtype=output.dtype,
+ device=output.device)
+ else:
+ seq_lens = attn_metadata.seq_lens
torch_npu._npu_flash_attention(query=query,
key=key,
value=value,
mask=mask,
- seq_len=attn_metadata.seq_lens,
+ seq_len=seq_lens,
scale_value=self.scale,
num_heads=self.num_heads,
num_kv_heads=self.num_kv_heads,
out=output)
assert output is not None
- return output[:num_tokens, :, :]
+ return output[param_sink_number:param_sink_number + num_tokens, :, :]
def _forward_prefill_cache_hit(
self,
@@ -356,6 +463,8 @@ class AscendAttentionBackendImpl(AttentionImpl):
query: torch.Tensor,
attn_metadata: AscendMetadata,
output: Optional[torch.Tensor] = None,
+ layer: AttentionLayer = None,
+ param_sink_number: Optional[int] = 0
) -> torch.Tensor:
if is_310p():
# seq_lens_tensor needs to be transferred to the device for 310P.
@@ -426,16 +535,46 @@ class AscendAttentionBackendImpl(AttentionImpl):
handle = torch.npu.graph_task_group_end(stream)
graph_params.handles[num_tokens].append(handle)
else:
- torch_npu._npu_paged_attention(
- query=query,
- key_cache=self.key_cache,
- value_cache=self.value_cache,
- num_kv_heads=self.num_kv_heads,
- num_heads=self.num_heads,
- scale_value=self.scale,
- block_table=attn_metadata.block_tables,
- context_lens=attn_metadata.seq_lens,
- out=output)
+ # Patch for Sparse KV cache of SWA.
+ num_block, block_size, _, _ = self.key_cache.shape # type: ignore
+ key = self.key_cache.view( # type: ignore
+ num_block, block_size, -1)
+ value = self.value_cache.view( # type: ignore
+ num_block, block_size, -1)
+ block_tables = attn_metadata.sink_block_tables
+ use_swa = (self.layer_idx % 2 == 0)
+ seq_kvlens = attn_metadata.sink_seq_kvlens
+ if use_swa:
+ attn_mask = self.attn_mask.to(query.device, non_blocking=True)
+
+ output, _ = torch.ops.custom.npu_fused_infer_attention_sink(
+ query,
+ key,
+ value,
+ atten_mask=attn_mask,
+ actual_seq_qlen=attn_metadata.swa_seq_qlens,
+ actual_seq_kvlen=seq_kvlens,
+ block_table=block_tables,
+ pre_tokens=128,
+ next_tokens=0,
+ num_query_heads=self.num_heads,
+ num_key_value_heads=self.num_kv_heads,
+ input_layout='TND',
+ sparse_mode=4,
+ block_size=block_size,
+ sink_number=param_sink_number,
+ softmax_scale=self.scale)
+ else:
+ torch_npu._npu_paged_attention(
+ query=query,
+ key_cache=self.key_cache,
+ value_cache=self.value_cache,
+ num_kv_heads=self.num_kv_heads,
+ num_heads=self.num_heads,
+ scale_value=self.scale,
+ block_table=block_tables,
+ context_lens=seq_kvlens,
+ out=output)
return output
def _forward_v1_style(
@@ -443,6 +582,7 @@ class AscendAttentionBackendImpl(AttentionImpl):
query: torch.Tensor,
attn_metadata: AscendMetadata,
output: Optional[torch.Tensor] = None,
+ param_sink_number: Optional[int] = 0
) -> torch.Tensor:
# Use chunked prefill for head size 192 scenario, like deepseek
# paged_attention_splitfuse maybe crash at such scenario.
@@ -485,34 +625,87 @@ class AscendAttentionBackendImpl(AttentionImpl):
value = self.value_cache.view( # type: ignore
num_block, block_size, -1)
- output, _ = torch_npu.npu_fused_infer_attention_score(
- query=query,
- key=key,
- value=value,
- atten_mask=attn_metadata.attn_mask,
- block_table=attn_metadata.block_tables,
- input_layout="TND",
- block_size=block_size,
- actual_seq_lengths=attn_metadata.query_start_loc[1:],
- actual_seq_lengths_kv=attn_metadata.seq_lens,
- num_key_value_heads=self.num_kv_heads,
- num_heads=self.num_heads,
- scale=self.scale,
- sparse_mode=3,
- )
+ #TODO: swa层,window长度 传参
+ use_swa = (self.layer_idx % 2 == 0)
+ sparse_mode = 4 if use_swa else 3
+ if param_sink_number > 0:
+ if sparse_mode == 4:
+ output, _ = torch.ops.custom.npu_fused_infer_attention_sink(
+ query,
+ key,
+ value,
+ atten_mask=self.attn_mask,
+ actual_seq_qlen=attn_metadata.swa_seq_qlens,
+ actual_seq_kvlen=attn_metadata.sink_seq_kvlens,
+ block_table=attn_metadata.sink_block_tables,
+ pre_tokens=128,
+ next_tokens=0,
+ num_query_heads=self.num_heads,
+ num_key_value_heads=self.num_kv_heads,
+ input_layout='TND',
+ sparse_mode=4,
+ block_size=block_size,
+ sink_number=param_sink_number,
+ softmax_scale=self.scale
+ )
+ elif sparse_mode == 3:
+ output, _ = torch.ops.custom.npu_fused_infer_attention_sink(
+ query,
+ key,
+ value,
+ atten_mask=self.attn_mask,
+ actual_seq_qlen=attn_metadata.swa_seq_qlens,
+ actual_seq_kvlen=attn_metadata.sink_seq_kvlens,
+ block_table=attn_metadata.sink_block_tables,
+ num_query_heads=self.num_heads,
+ num_key_value_heads=self.num_kv_heads,
+ input_layout='TND',
+ sparse_mode=3,
+ block_size=block_size,
+ sink_number=param_sink_number,
+ softmax_scale=self.scale
+ )
+
+ else:
+ output, _ = torch_npu.npu_fused_infer_attention_score(
+ query=query,
+ key=key,
+ value=value,
+ atten_mask=attn_metadata.attn_mask,
+ block_table=attn_metadata.block_tables,
+ input_layout="TND",
+ block_size=block_size,
+ actual_seq_lengths=attn_metadata.query_start_loc[1:],
+ actual_seq_lengths_kv=attn_metadata.seq_lens,
+ num_key_value_heads=self.num_kv_heads,
+ num_heads=self.num_heads,
+ scale=self.scale,
+ sparse_mode=3,
+ )
else:
+ # Patch for sink on CANN 8.2
+ if param_sink_number > 0:
+ seq_kvlens = attn_metadata.seq_lens + param_sink_number
+ block_tables = F.pad(attn_metadata.block_tables, (1, 0, 0, 0), value=0)
+ mask = F.pad(attn_metadata.attn_mask, (param_sink_number, 0, 0, 0), value=0)
+ else:
+ seq_kvlens = attn_metadata.seq_lens
+ block_tables = attn_metadata.block_tables
+ mask = attn_metadata.attn_mask
+
torch_npu._npu_paged_attention_splitfuse(
query=query,
key_cache=self.key_cache,
value_cache=self.value_cache,
- mask=attn_metadata.attn_mask,
- block_table=attn_metadata.block_tables,
+ mask=mask,
+ block_table=block_tables,
seq_len=attn_metadata.query_lens,
- context_lens=attn_metadata.seq_lens,
+ context_lens=seq_kvlens,
num_kv_heads=self.num_kv_heads,
num_heads=self.num_heads,
scale_value=self.scale,
out=output)
+
return output
def forward(
@@ -525,6 +718,10 @@ class AscendAttentionBackendImpl(AttentionImpl):
attn_metadata: AscendMetadata,
output: Optional[torch.Tensor] = None,
trace_flag: bool = True,
+ sink_query: Optional[torch.Tensor] = None,
+ sink_key: Optional[torch.Tensor] = None,
+ sink_value: Optional[torch.Tensor] = None,
+ v_head_size: Optional[int] = None,
) -> torch.Tensor:
"""Forward pass with Ascend attention.
Args:
@@ -556,7 +753,12 @@ class AscendAttentionBackendImpl(AttentionImpl):
key=key,
value=value,
output=output,
- layer_name=layer.layer_name)
+ layer_name=layer.layer_name,
+ sink_query=sink_query,
+ sink_key=sink_key,
+ sink_value=sink_value,
+ v_head_size=v_head_size
+ )
elif hasattr(layer, 'quant_method') and use_kv_cache_int8:
output = layer.quant_method.apply(layer, query, key, value,
@@ -575,10 +777,13 @@ class AscendAttentionBackendImpl(AttentionImpl):
"encoder/decoder cross-attention "
"are not implemented for "
"PallasAttentionBackendImpl")
+ sink_key_flag = (sink_key is not None)
+ param_sink_number = sink_key.shape[0] if sink_key_flag else 0
# View q k v to BSH.
query = query.view(-1, self.num_heads, self.head_size)
key = key.view(-1, self.num_kv_heads, self.head_size)
- value = value.view(-1, self.num_kv_heads, self.head_size)
+ value = value.view(-1, self.num_kv_heads,
+ v_head_size if v_head_size is not None else self.head_size)
# TODO: Remove this contiguous in the future.
value = value.contiguous()
@@ -586,33 +791,63 @@ class AscendAttentionBackendImpl(AttentionImpl):
if self.key_cache is None:
self.key_cache, self.value_cache = kv_cache[0], kv_cache[1]
slots = attn_metadata.slot_mapping
+
torch_npu._npu_reshape_and_cache(
key=key[:num_actual_tokens],
value=value[:num_actual_tokens],
key_cache=self.key_cache,
value_cache=self.value_cache,
slot_indices=slots)
-
+ if sink_key_flag and not self.sink_cached:
+ # kv cache start from block 1 and slots 128, so we store sink in block 0.
+ slots = torch.arange(0, param_sink_number,
+ dtype=attn_metadata.slot_mapping.dtype,
+ device=attn_metadata.slot_mapping.device)
+ torch_npu._npu_reshape_and_cache(
+ key=sink_key,
+ value=sink_value,
+ key_cache=self.key_cache,
+ value_cache=self.value_cache,
+ slot_indices=slots)
+ self.sink_cached = True
+
+ # TODO: 暂不进PrefillCacheHit分支,不更新sink实现
+ if attn_metadata.attn_state == AscendAttentionState.PrefillNoCache and sink_key_flag:
+ attn_metadata.attn_state = AscendAttentionState.ChunkedPrefill
# V0-Style scheduler situation.
if attn_metadata.attn_state == AscendAttentionState.PrefillNoCache:
+ if torch.version.cann.startswith("8.3"):
+ # npu_fused_infer_attention_score and npu_fused_infer_attention_sink
+ # does not support cases where query.shape[0] != actual_seq_lengths
+ # Thus we need unpad it here.
+ num_tokens = attn_metadata.query_start_loc[-1]
+ query = query[:num_tokens]
+ key = key[:num_tokens]
+ value = value[:num_tokens]
+ elif sink_key_flag:
+ query = torch.cat([sink_query, query], dim=0)
+ if sink_key_flag:
+ key = torch.cat([sink_key, key], dim=0)
+ value = torch.cat([sink_value, value], dim=0)
output = self._forward_prefill_no_cache(
- query, key, value, attn_metadata, output, num_tokens)
+ query, key, value, attn_metadata, output, num_tokens,
+ param_sink_number
+ )
elif attn_metadata.attn_state == \
AscendAttentionState.PrefillCacheHit:
output = self._forward_prefill_cache_hit(
query, attn_metadata, output)
elif attn_metadata.attn_state == AscendAttentionState.DecodeOnly:
output = self._forward_decode_only(query, attn_metadata,
- output)
+ output, layer,
+ param_sink_number)
# Normal V1 situation.
else:
if torch.version.cann.startswith("8.3"):
- # npu_fused_infer_attention_score does not support cases
- # where query.shape[0] != attn_metadata.query_start_loc[-1].
- # Thus we need unpad it here.
num_tokens = attn_metadata.query_start_loc[-1]
query = query[:num_tokens]
- output = self._forward_v1_style(query, attn_metadata, output)
+ output = self._forward_v1_style(query, attn_metadata, output,
+ param_sink_number)
# to make in-place change to the output tensor
if hasattr(layer, 'quant_method') and use_kv_cache_int8:
@@ -627,6 +862,10 @@ def unified_ascend_attention_with_output(
value: torch.Tensor,
output: torch.Tensor,
layer_name: str,
+ sink_query: Optional[torch.Tensor] = None,
+ sink_key: Optional[torch.Tensor] = None,
+ sink_value: Optional[torch.Tensor] = None,
+ v_head_size: Optional[int] = None,
) -> None:
wait_for_kv_layer_from_connector(layer_name)
forward_context: ForwardContext = get_forward_context()
@@ -642,7 +881,11 @@ def unified_ascend_attention_with_output(
kv_cache,
attn_metadata,
output,
- trace_flag=False)
+ trace_flag=False,
+ sink_query=sink_query,
+ sink_key=sink_key,
+ sink_value=sink_value,
+ v_head_size=v_head_size)
maybe_save_kv_layer_to_connector(layer_name, kv_cache)
return
@@ -653,6 +896,11 @@ def unified_attention_with_output_fake(
value: torch.Tensor,
output: torch.Tensor,
layer_name: str,
+ # patch for pangu with attention sink
+ sink_query: Optional[torch.Tensor] = None,
+ sink_key: Optional[torch.Tensor] = None,
+ sink_value: Optional[torch.Tensor] = None,
+ v_head_size: Optional[int] = None,
) -> None:
return
diff --git a/vllm_ascend/attention/utils.py b/vllm_ascend/attention/utils.py
index 519cde0..93e1c95 100644
--- a/vllm_ascend/attention/utils.py
+++ b/vllm_ascend/attention/utils.py
@@ -63,6 +63,8 @@ class AscendCommonAttentionMetadata:
graph_pad_size: int = -1
+ num_input_tokens: int = -1
+
def split_decodes_and_prefills(
common_attn_metadata: AscendCommonAttentionMetadata,
diff --git a/vllm_ascend/platform.py b/vllm_ascend/platform.py
index f1581df..b690bcb 100644
--- a/vllm_ascend/platform.py
+++ b/vllm_ascend/platform.py
@@ -216,6 +216,9 @@ class NPUPlatform(Platform):
if compilation_config.cudagraph_mode == CUDAGraphMode.FULL_AND_PIECEWISE:
compilation_config.cudagraph_mode = CUDAGraphMode.PIECEWISE
+ if compilation_config.cudagraph_mode == CUDAGraphMode.PIECEWISE:
+ compilation_config.cudagraph_mode = CUDAGraphMode.PIECEWISE_DECODE_ONLY
+
if compilation_config.cudagraph_mode == CUDAGraphMode.NONE:
compilation_config.level = CompilationLevel.NO_COMPILATION
# TODO: Currently MLA does not support FULL_DECODE_ONLY, remove the second condition
@@ -223,7 +226,8 @@ class NPUPlatform(Platform):
elif compilation_config.cudagraph_mode == CUDAGraphMode.PIECEWISE or (
compilation_config.cudagraph_mode
== CUDAGraphMode.FULL_DECODE_ONLY and model_config is not None
- and model_config.use_mla):
+ and model_config.use_mla) or (
+ compilation_config.cudagraph_mode == CUDAGraphMode.PIECEWISE_DECODE_ONLY):
logger.info(
"PIECEWISE compilation enabled on NPU. use_inductor not supported - "
"using only ACL Graph mode")
@@ -232,7 +236,8 @@ class NPUPlatform(Platform):
compilation_config.set_splitting_ops_for_v1()
compilation_config.use_inductor = False
compilation_config.splitting_ops.extend([
- "vllm.unified_ascend_attention_with_output", "vllm.mla_forward"
+ "vllm.unified_ascend_attention_with_output", "vllm.mla_forward",
+ "vllm.aggregate_hiddden",
])
update_aclgraph_sizes(vllm_config)
elif compilation_config.cudagraph_mode == CUDAGraphMode.FULL_DECODE_ONLY:
diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py
index 9281dd7..34808ec 100644
--- a/vllm_ascend/worker/model_runner_v1.py
+++ b/vllm_ascend/worker/model_runner_v1.py
@@ -281,6 +281,7 @@ class NPUModelRunner(LoRAModelRunnerMixin):
self.encoder_cache: Dict[str, torch.Tensor] = {}
self.attn_mask = None
self.attn_state = None
+ self.with_prefill = False
self.requests: Dict[str, CachedRequestState] = {}
self.intermediate_tensors: Optional[IntermediateTensors] = None
self.runner_only_attn_layers: set[str] = set()
@@ -509,6 +510,48 @@ class NPUModelRunner(LoRAModelRunnerMixin):
self.num_draft_tokens = self._make_buffer(self.max_num_reqs,
dtype=torch.int32)
+ # Patch for conv cache
+ self.router_sliding_window = getattr(self.model_config.hf_text_config, "router_sliding_window", 0)
+ if self.router_sliding_window > 1:
+ self.cache_length = self.router_sliding_window - 1
+ self.req_cache_map = {}
+ self.occupied_cache = [0]*(self.max_num_reqs)
+ self.q_offsets = torch.arange(-self.cache_length, 0, device=self.device)
+ self.cache_slot_id = torch.empty(self.max_num_reqs,
+ dtype=torch.long, device=self.device)
+ self.is_first_chunk = torch.empty(self.max_num_reqs, dtype=torch.bool, device=self.device) # For chunked prefill
+
+ def _build_conv_context(self, with_prefill:bool = False, dummy:bool = False, num_tokens:int = 0):
+ # conv cache slot & prefill hiddenstates loc
+ cache_slot_id = self.cache_slot_id[:self.input_batch.num_reqs]
+ query_start_loc = self.query_start_loc[:self.input_batch.num_reqs + 1]
+ is_first_chunk = self.is_first_chunk[:self.input_batch.num_reqs]
+
+ if with_prefill:
+ for idx, req_id in enumerate(self.input_batch.req_ids):
+ if req_id in self.req_cache_map:
+ cache_id = self.req_cache_map[req_id]
+ cache_slot_id[idx] = cache_id
+ is_first_chunk[idx] = False
+ else:
+ # new request with the first chunk
+ new_cahce_id = self.occupied_cache.index(0)
+ self.occupied_cache[new_cahce_id] = 1
+ self.req_cache_map[req_id] = new_cahce_id
+ cache_slot_id[idx] = new_cahce_id
+ is_first_chunk[idx] = True
+ else:
+ for idx, req_id in enumerate(self.input_batch.req_ids):
+ cache_id = self.req_cache_map[req_id]
+ cache_slot_id[idx] = cache_id
+ is_first_chunk[idx] = False
+
+ forward_context = get_forward_context()
+ forward_context.cache_slot_id = cache_slot_id
+ forward_context.is_first_chunk = is_first_chunk
+ forward_context.query_start_loc = query_start_loc
+
+
def _make_buffer(self,
*size: Union[int, torch.SymInt],
dtype: torch.dtype,
@@ -548,12 +591,16 @@ class NPUModelRunner(LoRAModelRunnerMixin):
self.input_batch.num_accepted_tokens_cpu[i] = num_tokens
def _use_aclgraph(self) -> bool:
- return self.compilation_config.cudagraph_mode != CUDAGraphMode.NONE and self.compilation_config.level == CompilationLevel.PIECEWISE and not self.model_config.enforce_eager
+ return self.compilation_config.cudagraph_mode != CUDAGraphMode.NONE and \
+ self.compilation_config.level == CompilationLevel.PIECEWISE and not self.model_config.enforce_eager
def _update_states(self, scheduler_output: "SchedulerOutput") -> None:
# Remove finished requests from the cached states.
for req_id in scheduler_output.finished_req_ids:
- self.requests.pop(req_id, None)
+ self.requests.pop(req_id, None)
+ if self.router_sliding_window > 1 and req_id in self.req_cache_map:
+ cache_id = self.req_cache_map.pop(req_id)
+ self.occupied_cache[cache_id] = 0
# Remove the finished requests from the persistent batch.
# NOTE(woosuk): There could be an edge case where finished_req_ids and
@@ -891,7 +938,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
def _make_attention_mask(self, seq_lens, position,
attn_state) -> torch.Tensor:
# Chunk Prefill situation.
- if attn_state == AscendAttentionState.ChunkedPrefill and not self.vllm_config.model_config.use_mla and not self.ascend_config.use_sfa:
+ if attn_state == AscendAttentionState.ChunkedPrefill and not self.vllm_config.model_config.use_mla and not \
+ self.ascend_config.use_sfa:
if torch.version.cann.startswith("8.3"):
return self.attn_mask_builder.get_splitfuse_attn_mask()
else:
@@ -942,7 +990,7 @@ class NPUModelRunner(LoRAModelRunnerMixin):
src_end = num_computed_tokens + prompt_part_len
self.mrope_positions_cpu[:, dst_start:dst_end] = \
- req.mrope_positions[:,src_start:src_end]
+ req.mrope_positions[:, src_start:src_end]
mrope_pos_ptr += prompt_part_len
@@ -1126,9 +1174,9 @@ class NPUModelRunner(LoRAModelRunnerMixin):
cumsum_dtype: Optional[np.dtype] = None,
) -> tuple[np.ndarray, np.ndarray]:
"""Get the cumulative sum and batched arange of the given array.
- # E.g., [2, 5, 3] -> ([2, 7, 10], [0, 1, 0, 1, 2, 3, 4, 0, 1, 2])
- # Equivalent to but faster than:
- # np.concatenate([np.arange(n) for n in num_tokens])
+ E.g., [2, 5, 3] -> ([2, 7, 10], [0, 1, 0, 1, 2, 3, 4, 0, 1, 2])
+ Equivalent to but faster than:
+ np.concatenate([np.arange(n) for n in num_tokens])
"""
# Step 1. [2, 5, 3] -> [2, 7, 10]
cu_num_tokens = np.cumsum(num_tokens, dtype=cumsum_dtype)
@@ -1518,6 +1566,7 @@ class NPUModelRunner(LoRAModelRunnerMixin):
max_query_len=max_num_scheduled_tokens,
graph_pad_size=self.graph_pad_size,
decode_token_per_req=self.decode_token_per_req,
+ num_input_tokens=num_input_tokens
)
if self.speculative_config and \
@@ -1964,6 +2013,9 @@ class NPUModelRunner(LoRAModelRunnerMixin):
model_instance=self.model):
self.maybe_setup_kv_connector(scheduler_output)
+ if self.router_sliding_window > 1:
+ self._build_conv_context(self.with_prefill)
+
hidden_states = self._generate_process_reqs_hidden_states(
attn_metadata, self.with_prefill, maybe_padded_num_tokens,
input_ids, positions, intermediate_tensors, inputs_embeds)
@@ -2339,7 +2391,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
) -> torch.Tensor:
# only support eager mode and piecewise graph now
assert aclgraph_runtime_mode in {
- CUDAGraphMode.NONE, CUDAGraphMode.PIECEWISE, CUDAGraphMode.FULL
+ CUDAGraphMode.NONE, CUDAGraphMode.PIECEWISE, CUDAGraphMode.FULL,
+ CUDAGraphMode.PIECEWISE_DECODE_ONLY
}
# Padding for DP
@@ -2472,6 +2525,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
batch_descriptor=batch_descriptor,
prefetch_stream=self.prefetch_stream,
model_instance=self.model):
+ if self.router_sliding_window > 1:
+ self._build_conv_context(with_prefill, dummy=True, num_tokens=num_tokens)
hidden_states = self._generate_dummy_run_hidden_states(
with_prefill, is_torchair_compile, input_ids, positions,
attn_metadata, num_tokens, intermediate_tensors,
@@ -2789,8 +2844,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
# In order to transfer kv cache through the reigster_memory api from llmdatadist, the memory
# address should be aligned by 2M. In most case, torch_npu can allocate 2M aligned memory, but
- # we found there are also some exceptions during test, so we manual align those memory here, this part
- # of code may consume 2M * 2 * elem_size memory every layer.
+ # we found there are also some exceptions during test, so we manual align those memory here,
+ # this part of code may consume 2M * 2 * elem_size memory every layer.
nope_allocate_shape = num_blocks * block_size * num_kv_heads * nope_dim
nope_allocate_shape_alignment = nope_allocate_shape + alignment
rope_allocate_shape = num_blocks * block_size * num_kv_heads * rope_dim
@@ -2888,8 +2943,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
# In order to transfer kv cache through the reigster_memory api from llmdatadist, the memory
# address should be aligned by 2M. In most case, torch_npu can allocate 2M aligned memory, but
- # we found there are also some exceptions during test, so we manual align those memory here, this part
- # of code may consume 2M * 2 * elem_size memory every layer.
+ # we found there are also some exceptions during test, so we manual align those memory here,
+ # this part of code may consume 2M * 2 * elem_size memory every layer.
nope_allocate_shape = num_blocks * block_size * num_kv_heads * nope_dim
nope_allocate_shape_alignment = nope_allocate_shape + alignment
rope_allocate_shape = num_blocks * block_size * num_kv_heads * rope_dim
@@ -3432,6 +3487,7 @@ class NPUModelRunner(LoRAModelRunnerMixin):
and all(op in self.compilation_config.splitting_ops for op in [
"vllm.unified_ascend_attention_with_output",
"vllm.mla_forward",
+ "vllm.aggregate_hiddden",
]))
# Flexible resolve the aclgraph mode
@@ -3495,7 +3551,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
uniform_decode: bool):
assert aclgraph_runtime_mode != CUDAGraphMode.NONE and \
aclgraph_runtime_mode in [CUDAGraphMode.FULL,
- CUDAGraphMode.PIECEWISE]
+ CUDAGraphMode.PIECEWISE,
+ CUDAGraphMode.PIECEWISE_DECODE_ONLY]
# Only rank 0 should print progress bar during capture
if is_global_first_rank():
@@ -3519,10 +3576,12 @@ class NPUModelRunner(LoRAModelRunnerMixin):
# attention while `PIECEWISE` implies no attention.
force_attention = (aclgraph_runtime_mode == CUDAGraphMode.FULL)
self._dummy_run(num_tokens,
+ with_prefill = (uniform_decode == False),
aclgraph_runtime_mode=CUDAGraphMode.NONE,
force_attention=force_attention,
uniform_decode=uniform_decode)
self._dummy_run(num_tokens,
+ with_prefill = (uniform_decode == False),
aclgraph_runtime_mode=aclgraph_runtime_mode,
force_attention=force_attention,
uniform_decode=uniform_decode)
@@ -3556,7 +3615,8 @@ class NPUModelRunner(LoRAModelRunnerMixin):
logger.error(
f"ACLgraph sizes capture fail: {type(e).__name__}:\n"
"ACLgraph has insufficient available streams to capture the configured number of sizes. "
- "Please verify both the availability of adequate streams and the appropriateness of the configured size count.\n\n"
+ "Please verify both the availability of adequate streams "
+ "and the appropriateness of the configured size count.\n\n"
"Recommended solutions:\n"
"1. Manually configure the compilation_config parameter "
"with a reduced set of sizes: '{\"cudagraph_capture_sizes\":[size1, size2, size3, ...]}'.\n"
@@ -3564,8 +3624,9 @@ class NPUModelRunner(LoRAModelRunnerMixin):
f"{str(e)}")
raise
- if aclgraph_mode.decode_mode() == CUDAGraphMode.FULL and \
- aclgraph_mode.separate_routine():
+ if aclgraph_mode.separate_routine() and \
+ (aclgraph_mode.decode_mode() == CUDAGraphMode.FULL or \
+ aclgraph_mode.decode_mode() == CUDAGraphMode.PIECEWISE):
max_num_tokens = self.scheduler_config.max_num_seqs * \
self.uniform_decode_query_len
decode_cudagraph_batch_sizes = [
@@ -3576,7 +3637,7 @@ class NPUModelRunner(LoRAModelRunnerMixin):
reversed(decode_cudagraph_batch_sizes))
self._capture_aclgraphs(
compilation_cases=compilation_cases_decode,
- aclgraph_runtime_mode=CUDAGraphMode.FULL,
+ aclgraph_runtime_mode=aclgraph_mode.decode_mode(),
uniform_decode=True)
# Disable aclgraph capturing globally, so any unexpected aclgraph