0
vllm_ascend/_310p/attention/__init__.py
Normal file
0
vllm_ascend/_310p/attention/__init__.py
Normal file
186
vllm_ascend/_310p/attention/attention_mask.py
Normal file
186
vllm_ascend/_310p/attention/attention_mask.py
Normal file
@@ -0,0 +1,186 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
import torch
|
||||
import torch_npu
|
||||
|
||||
from vllm_ascend.attention.attention_v1 import AscendMetadata
|
||||
from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, nd_to_nz_2d, nd_to_nz_spec
|
||||
|
||||
COMPRESSED_MASK_SEQ_LEN = 2048
|
||||
PAGED_ATTENTION_COMPRESSED_MASK_VALUE = -10000.0
|
||||
|
||||
|
||||
def is_compressed_mask_supported() -> bool:
|
||||
return hasattr(torch_npu, "_npu_flash_attention_v3") and hasattr(torch_npu, "_npu_paged_attention_splitfuse_v2")
|
||||
|
||||
|
||||
class AttentionMaskBuilder310:
|
||||
chunked_prefill_attn_mask = None
|
||||
compressed_chunked_prefill_attn_mask = None
|
||||
max_seqlen = 16384
|
||||
|
||||
def __init__(self, device: torch.device, max_seqlen: int):
|
||||
"""
|
||||
Initializes the AttentionMaskBuilder for the 310P device.
|
||||
|
||||
Args:
|
||||
device (torch.device): The device on which tensors will be allocated.
|
||||
max_seqlen (int): Maximum length of a sequence (including prompt and generated text).
|
||||
"""
|
||||
AttentionMaskBuilder310.max_seqlen = max_seqlen
|
||||
self.causal_attn_mask_cache = None
|
||||
self.non_causal_attn_mask_cache = None
|
||||
self.support_compressed_mask = is_compressed_mask_supported()
|
||||
self.device = device
|
||||
|
||||
@staticmethod
|
||||
def gen_causal_additive_mask(max_seq_len: int, device: torch.device):
|
||||
"""
|
||||
Generates a standard causal lower-triangular attention mask.
|
||||
|
||||
The upper triangular part is filled with negative infinity (float("-inf"))
|
||||
to mask out future tokens, while the lower triangular part is kept as 0.
|
||||
|
||||
Args:
|
||||
max_seq_len (int): The maximum sequence length for the mask.
|
||||
device (torch.device): The target device for the tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: A float16 tensor representing the causal mask.
|
||||
"""
|
||||
tril = torch.ones((max_seq_len, max_seq_len), dtype=torch.bool, device=device).tril_()
|
||||
upper = ~tril
|
||||
mask = torch.zeros((max_seq_len, max_seq_len), dtype=torch.float16, device=device)
|
||||
mask.masked_fill_(upper, float("-inf"))
|
||||
return mask
|
||||
|
||||
@classmethod
|
||||
def get_splitfuse_mask(cls, attn_metadata: AscendMetadata, device: torch.device):
|
||||
"""
|
||||
Generates and formats the attention mask for SplitFuse (chunked prefill) decoding.
|
||||
|
||||
It calculates the specific indices required based on query start locations
|
||||
and context lengths, selects the relevant parts from the global chunked
|
||||
mask, and converts the result to the NPU-specific fractal format.
|
||||
|
||||
Args:
|
||||
attn_metadata (AscendMetadata): Metadata containing query start locations and sequence lengths.
|
||||
device (torch.device): The device to perform operations on.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The splitfuse attention mask cast to ACL_FORMAT_FRACTAL_NZ.
|
||||
"""
|
||||
if cls.chunked_prefill_attn_mask is None:
|
||||
cls.chunked_prefill_attn_mask = cls.gen_causal_additive_mask(cls.max_seqlen, device)
|
||||
qsl = attn_metadata.query_start_loc.to("cpu", dtype=torch.int32)
|
||||
qlens = qsl[1:] - qsl[:-1]
|
||||
q_list = qlens.tolist()
|
||||
context_lens = attn_metadata.seq_lens.to("cpu", dtype=torch.int32)
|
||||
c_list = context_lens.tolist()
|
||||
pos_list = [p for ql, cl in zip(q_list, c_list) for p in range(cl - ql, cl)]
|
||||
position = torch.tensor(pos_list, dtype=torch.int32, device=device)
|
||||
splitfuse_mask = cls.chunked_prefill_attn_mask.index_select(0, position)
|
||||
splitfuse_mask_nz = torch_npu.npu_format_cast(nd_to_nz_spec(splitfuse_mask).contiguous(), ACL_FORMAT_FRACTAL_NZ)
|
||||
return splitfuse_mask_nz
|
||||
|
||||
@classmethod
|
||||
def get_compressed_splitfuse_mask(cls, device: torch.device):
|
||||
"""
|
||||
Generates the fixed ND attention mask for compressed SplitFuse PA.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: A [2048, 2048] float16 ND mask on the target device.
|
||||
"""
|
||||
if (
|
||||
cls.compressed_chunked_prefill_attn_mask is None
|
||||
or cls.compressed_chunked_prefill_attn_mask.device != device
|
||||
):
|
||||
mask = torch.ones(
|
||||
size=(COMPRESSED_MASK_SEQ_LEN, COMPRESSED_MASK_SEQ_LEN),
|
||||
dtype=torch.float16,
|
||||
device=device,
|
||||
)
|
||||
mask = torch.triu(mask, diagonal=1)
|
||||
cls.compressed_chunked_prefill_attn_mask = mask.mul_(PAGED_ATTENTION_COMPRESSED_MASK_VALUE)
|
||||
return cls.compressed_chunked_prefill_attn_mask
|
||||
|
||||
def get_attention_mask(self, causal: bool, model_config) -> torch.Tensor:
|
||||
"""
|
||||
Retrieves the appropriate attention mask based on the model configuration.
|
||||
|
||||
When compressed mask is supported, the mask is generated as a fixed
|
||||
[2048, 2048] logical mask and converted to 4D FRACTAL_NZ.
|
||||
|
||||
Args:
|
||||
causal (bool): Whether to generate a causal mask.
|
||||
model_config: Configuration object containing runner details.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The causal attention mask.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the runner_type is 'pooling'.
|
||||
"""
|
||||
max_seq_len = COMPRESSED_MASK_SEQ_LEN if self.support_compressed_mask else self.max_seqlen
|
||||
if getattr(model_config, "runner_type", None) == "pooling":
|
||||
if causal:
|
||||
return self._get_causal_mask(max_seq_len)
|
||||
else:
|
||||
return self._get_non_causal_mask(max_seq_len, model_config.dtype)
|
||||
|
||||
return self._get_causal_mask(max_seq_len)
|
||||
|
||||
def _get_causal_mask(self, max_seq_len: int) -> torch.Tensor:
|
||||
"""
|
||||
Internal method to get or update the cached causal attention mask.
|
||||
|
||||
If the cache is empty, a new mask is generated and converted to the
|
||||
NPU fractal format.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The cached causal mask in ACL_FORMAT_FRACTAL_NZ.
|
||||
"""
|
||||
if self.causal_attn_mask_cache is None:
|
||||
attn_mask = self.gen_causal_additive_mask(max_seq_len, self.device)
|
||||
self.causal_attn_mask_cache = torch_npu.npu_format_cast(nd_to_nz_2d(attn_mask), ACL_FORMAT_FRACTAL_NZ)
|
||||
return self.causal_attn_mask_cache
|
||||
|
||||
def _get_non_causal_mask(self, max_seq_len: int, dtype: torch.dtype) -> torch.Tensor:
|
||||
"""
|
||||
Internal method to get or update the cached non-causal attention mask.
|
||||
|
||||
If the cache is empty, a new mask is generated and converted to the
|
||||
NPU fractal format.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The cached causal mask in ACL_FORMAT_FRACTAL_NZ.
|
||||
"""
|
||||
if self.non_causal_attn_mask_cache is not None:
|
||||
return self.non_causal_attn_mask_cache
|
||||
|
||||
attention_mask_npu = torch.zeros(
|
||||
size=(max_seq_len, max_seq_len),
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
attention_mask_npu = nd_to_nz_2d(attention_mask_npu)
|
||||
self.non_causal_attn_mask_cache = torch_npu.npu_format_cast(
|
||||
attention_mask_npu.contiguous(), ACL_FORMAT_FRACTAL_NZ
|
||||
)
|
||||
|
||||
return self.non_causal_attn_mask_cache
|
||||
345
vllm_ascend/_310p/attention/attention_v1.py
Normal file
345
vllm_ascend/_310p/attention/attention_v1.py
Normal file
@@ -0,0 +1,345 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch_npu
|
||||
from vllm.v1.attention.backends.registry import ( # type: ignore
|
||||
AttentionBackendEnum,
|
||||
register_backend,
|
||||
)
|
||||
|
||||
from vllm_ascend._310p.attention.attention_mask import (
|
||||
AttentionMaskBuilder310,
|
||||
is_compressed_mask_supported,
|
||||
)
|
||||
from vllm_ascend._310p.attention.metadata_builder import (
|
||||
AscendAttentionMetadataBuilder310,
|
||||
get_query_lens_cpu,
|
||||
)
|
||||
from vllm_ascend.attention.attention_v1 import (
|
||||
AscendAttentionBackend,
|
||||
AscendAttentionBackendImpl,
|
||||
AscendAttentionMetadataBuilder,
|
||||
AscendAttentionState,
|
||||
AscendMetadata,
|
||||
)
|
||||
|
||||
MASK_TYPE_NORM_COMPRESS_SELF_ATTENTION = 3
|
||||
MASK_TYPE_NORM_COMPRESS_PAGED_ATTENTION = 5
|
||||
|
||||
|
||||
@register_backend(AttentionBackendEnum.CUSTOM, "ASCEND")
|
||||
class AscendAttentionBackend310(AscendAttentionBackend):
|
||||
def __init__(self, *args, **kwargs):
|
||||
"""
|
||||
Initializes the 310P backend and sets up the device-specific mask builder.
|
||||
"""
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def get_kv_cache_shape(
|
||||
num_blocks: int,
|
||||
block_size: int,
|
||||
num_kv_heads: int,
|
||||
head_size: int,
|
||||
cache_type: str = "",
|
||||
):
|
||||
"""
|
||||
Determines the shape of the Key-Value (KV) cache tensor.
|
||||
|
||||
The 310P hardware requires specific memory alignment for optimal performance.
|
||||
This method defines a 5D tensor shape where the head size dimension is
|
||||
split to ensure alignment to multiples of 16.
|
||||
|
||||
Args:
|
||||
num_blocks (int): Number of memory blocks.
|
||||
block_size (int): Size of each block.
|
||||
num_kv_heads (int): Number of KV heads.
|
||||
head_size (int): Dimension size of each head.
|
||||
|
||||
Returns:
|
||||
tuple: The specific 5D shape required by the hardware
|
||||
(2, num_blocks, hidden_dim_aligned, block_size, 16).
|
||||
"""
|
||||
# Align to a multiple of 16, as required by the 310P device.
|
||||
return (2, num_blocks, (num_kv_heads * head_size) // 16, block_size, 16)
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls():
|
||||
"""
|
||||
Returns the implementation class for the attention operations.
|
||||
"""
|
||||
return AscendAttentionBackendImpl310
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AscendAttentionMetadataBuilder"]:
|
||||
"""
|
||||
Returns the metadata builder class specifically for 310P.
|
||||
"""
|
||||
return AscendAttentionMetadataBuilder310
|
||||
|
||||
@staticmethod
|
||||
def get_supported_kernel_block_sizes() -> list[int]:
|
||||
return [128, 64]
|
||||
|
||||
|
||||
class AscendAttentionBackendImpl310(AscendAttentionBackendImpl):
|
||||
"""
|
||||
Implementation of attention operations (Prefill, Decode, Chunked Prefill)
|
||||
optimized for the Ascend 310P architecture.
|
||||
"""
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.support_compressed_mask = is_compressed_mask_supported()
|
||||
|
||||
def _flash_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
seq_len: torch.Tensor,
|
||||
output: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
if not self.support_compressed_mask:
|
||||
torch_npu._npu_flash_attention(
|
||||
query=query,
|
||||
key=key,
|
||||
value=value,
|
||||
mask=mask,
|
||||
seq_len=seq_len,
|
||||
scale_value=self.scale,
|
||||
num_heads=self.num_heads,
|
||||
num_kv_heads=self.num_kv_heads,
|
||||
out=output,
|
||||
)
|
||||
return output
|
||||
|
||||
torch_npu._npu_flash_attention_v3(
|
||||
query=query,
|
||||
key=key,
|
||||
value=value,
|
||||
mask=mask,
|
||||
seq_len=seq_len,
|
||||
scale_value=self.scale,
|
||||
num_heads=self.num_heads,
|
||||
num_kv_heads=self.num_kv_heads,
|
||||
mask_type=MASK_TYPE_NORM_COMPRESS_SELF_ATTENTION,
|
||||
out=output,
|
||||
)
|
||||
return output
|
||||
|
||||
def _forward_encoder_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AscendMetadata,
|
||||
output: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return self._flash_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_metadata.attn_mask,
|
||||
attn_metadata.seq_lens,
|
||||
output,
|
||||
)
|
||||
|
||||
def forward_paged_attention(
|
||||
self,
|
||||
query: Any,
|
||||
attn_metadata: AscendMetadata,
|
||||
output: Any | None = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Executes Paged Attention (typically for the decode phase).
|
||||
|
||||
Ensures that the sequence length metadata is on the correct device
|
||||
before invoking the base implementation.
|
||||
|
||||
Args:
|
||||
query (Any): The query tensor.
|
||||
attn_metadata (AscendMetadata): Metadata associated with the attention request.
|
||||
output (Any | None): Optional output tensor.
|
||||
|
||||
Returns:
|
||||
Any: The result of the attention operation.
|
||||
"""
|
||||
if attn_metadata.seq_lens.device != query.device:
|
||||
attn_metadata.seq_lens = attn_metadata.seq_lens.to(
|
||||
device=query.device,
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
return output
|
||||
|
||||
def forward_prefill_310(self, query, key, value, attn_metadata, output):
|
||||
"""
|
||||
Executes Flash Attention for the prefill phase on 310P.
|
||||
|
||||
This method handles memory alignment padding. If the query shape implies
|
||||
padding (aligned_tokens > real_tokens), it adjusts the sequence length
|
||||
of the last request to account for the delta, ensuring the NPU operator
|
||||
processes the data correctly.
|
||||
|
||||
Args:
|
||||
query, key, value: Input tensors.
|
||||
attn_metadata (AscendMetadata): Attention metadata containing masks and seq_lens.
|
||||
output: Output tensor.
|
||||
|
||||
Returns:
|
||||
The output tensor after flash attention.
|
||||
"""
|
||||
real_tokens = int(attn_metadata.seq_lens.sum().item())
|
||||
seq_len = attn_metadata.seq_lens
|
||||
aligned_tokens = int(query.shape[0])
|
||||
delta = aligned_tokens - real_tokens
|
||||
|
||||
# Adjust sequence length if padding (alignment) was applied to the inputs
|
||||
if delta:
|
||||
seq_len = seq_len.clone()
|
||||
seq_len[-1] += delta
|
||||
|
||||
mask = attn_metadata.attn_mask
|
||||
return self._flash_attention(query, key, value, mask, seq_len, output)
|
||||
|
||||
def forward_chunked_prefill_310(self, query, attn_metadata, output):
|
||||
"""
|
||||
Executes SplitFuse (Chunked Prefill) attention on 310P.
|
||||
|
||||
This handles scenarios where the prefill is split into chunks. It prepares
|
||||
the necessary metadata (query lengths, block tables) and generates the
|
||||
specific splitfuse mask before calling the NPU operator.
|
||||
|
||||
Args:
|
||||
query: The query tensor.
|
||||
attn_metadata (AscendMetadata): Metadata containing start locations and block tables.
|
||||
output: The output tensor.
|
||||
"""
|
||||
num_actual_tokens = int(attn_metadata.num_actual_tokens)
|
||||
query = query[:num_actual_tokens]
|
||||
output_slice = output[:num_actual_tokens]
|
||||
|
||||
# Host qLens filled in AscendAttentionMetadataBuilder310.build(); eager fallback only.
|
||||
qlens = get_query_lens_cpu(attn_metadata)
|
||||
if qlens is None:
|
||||
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
|
||||
|
||||
if _EXTRA_CTX.capturing:
|
||||
raise RuntimeError(
|
||||
"310P splitfuse requires attn_metadata.query_lens_cpu during graph capture; "
|
||||
"ensure AscendAttentionMetadataBuilder310.build() ran before forward."
|
||||
)
|
||||
qsl_cpu = attn_metadata.query_start_loc.cpu()
|
||||
qlens = qsl_cpu[1:] - qsl_cpu[:-1]
|
||||
|
||||
block_table = attn_metadata.block_tables
|
||||
|
||||
if attn_metadata.seq_lens.device != query.device:
|
||||
attn_metadata.seq_lens = attn_metadata.seq_lens.to(
|
||||
device=query.device,
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
if self.support_compressed_mask:
|
||||
# splitfuse_v2 requires fixed ND [2048, 2048]; parent build() may set FRACTAL_NZ mask.
|
||||
mask = AttentionMaskBuilder310.get_compressed_splitfuse_mask(query.device)
|
||||
torch_npu._npu_paged_attention_splitfuse_v2(
|
||||
query=query,
|
||||
key_cache=self.key_cache,
|
||||
value_cache=self.value_cache,
|
||||
mask=mask,
|
||||
block_table=block_table,
|
||||
seq_len=qlens,
|
||||
context_lens=attn_metadata.seq_lens,
|
||||
num_kv_heads=self.num_kv_heads,
|
||||
num_heads=self.num_heads,
|
||||
scale_value=self.scale,
|
||||
mask_type=MASK_TYPE_NORM_COMPRESS_PAGED_ATTENTION,
|
||||
out=output_slice,
|
||||
)
|
||||
return output
|
||||
|
||||
# Generate the specific mask for splitfuse
|
||||
mask = AttentionMaskBuilder310.get_splitfuse_mask(attn_metadata, query.device)
|
||||
torch_npu._npu_paged_attention_splitfuse(
|
||||
query=query,
|
||||
key_cache=self.key_cache,
|
||||
value_cache=self.value_cache,
|
||||
mask=mask,
|
||||
block_table=block_table,
|
||||
seq_len=qlens,
|
||||
context_lens=attn_metadata.seq_lens,
|
||||
num_kv_heads=self.num_kv_heads,
|
||||
num_heads=self.num_heads,
|
||||
scale_value=self.scale,
|
||||
out=output_slice,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
def forward_impl(self, query, key, value, kv_cache, attn_metadata, output):
|
||||
"""
|
||||
Main dispatch method for attention operations.
|
||||
|
||||
Routes the execution to Decode, Prefill, or Chunked Prefill methods
|
||||
based on the current attention state found in metadata.
|
||||
|
||||
Args:
|
||||
query, key, value: Input tensors (Key/Value usually empty for decode/chunked).
|
||||
kv_cache: The KV cache structure.
|
||||
attn_metadata: Metadata determining the state (Prefill vs Decode).
|
||||
output: Tensor to write results to.
|
||||
|
||||
Returns:
|
||||
The output tensor.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the attention state is not supported on 310P.
|
||||
"""
|
||||
state = attn_metadata.attn_state
|
||||
# Condition for PrefillNoCache: No previous tokens have been processed yet
|
||||
if state == AscendAttentionState.PrefillNoCache:
|
||||
output = self.forward_prefill_310(query, key, value, attn_metadata, output)
|
||||
# Condition for DecodeOnly: Pure decoding phase where each request generates one token
|
||||
elif state == AscendAttentionState.DecodeOnly:
|
||||
output = self.forward_paged_attention(query, attn_metadata, output)
|
||||
# ChunkedPrefill / PrefillCacheHit: chunked prefill or mixed batches.
|
||||
# SpecDecoding: MTP uniform spec verify (splitfuse on 310P).
|
||||
elif (
|
||||
state in [AscendAttentionState.ChunkedPrefill, AscendAttentionState.PrefillCacheHit]
|
||||
or state == AscendAttentionState.SpecDecoding
|
||||
):
|
||||
output = self.forward_chunked_prefill_310(query, attn_metadata, output)
|
||||
else:
|
||||
raise NotImplementedError(f"AscendAttentionState: {state} is not supported for 310P currently.")
|
||||
return output
|
||||
150
vllm_ascend/_310p/attention/metadata_builder.py
Normal file
150
vllm_ascend/_310p/attention/metadata_builder.py
Normal file
@@ -0,0 +1,150 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.v1.attention.backend import CommonAttentionMetadata
|
||||
from vllm.v1.kv_cache_interface import AttentionSpec
|
||||
|
||||
from vllm_ascend._310p.attention.attention_mask import (
|
||||
AttentionMaskBuilder310,
|
||||
is_compressed_mask_supported,
|
||||
)
|
||||
from vllm_ascend.attention.attention_v1 import (
|
||||
AscendAttentionMetadataBuilder,
|
||||
AscendAttentionState,
|
||||
AscendMetadata,
|
||||
)
|
||||
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
|
||||
|
||||
QUERY_LENS_CPU_ATTR = "query_lens_cpu"
|
||||
|
||||
|
||||
def set_query_lens_cpu(attn_metadata: AscendMetadata, query_lens_cpu: torch.Tensor) -> None:
|
||||
"""Attach host qLens for ATB splitfuse without extending upstream AscendMetadata."""
|
||||
setattr(attn_metadata, QUERY_LENS_CPU_ATTR, query_lens_cpu)
|
||||
|
||||
|
||||
def get_query_lens_cpu(attn_metadata: AscendMetadata) -> torch.Tensor | None:
|
||||
value = getattr(attn_metadata, QUERY_LENS_CPU_ATTR, None)
|
||||
if value is None:
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
class AscendAttentionMetadataBuilder310(AscendAttentionMetadataBuilder):
|
||||
"""
|
||||
Metadata builder specialized for the Huawei Ascend 310P NPU.
|
||||
|
||||
This class extends the base Ascend attention metadata builder to use
|
||||
the 310P-specific attention mask builder, ensuring that masks are
|
||||
generated in the correct format (FRACTAL_NZ) and logic required by
|
||||
the 310P hardware.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
kv_cache_spec: AttentionSpec,
|
||||
layer_names: list[str],
|
||||
vllm_config: VllmConfig,
|
||||
device: torch.device,
|
||||
):
|
||||
"""
|
||||
Initializes the metadata builder and the 310P-specific mask builder.
|
||||
|
||||
Args:
|
||||
kv_cache_spec (AttentionSpec): Specification for the KV cache (block size, etc.).
|
||||
layer_names (list[str]): List of layer names in the model.
|
||||
vllm_config (VllmConfig): Global vLLM configuration object.
|
||||
device (torch.device): The device (NPU) to run operations on.
|
||||
"""
|
||||
super().__init__(kv_cache_spec, layer_names, vllm_config, device)
|
||||
|
||||
# Override the mask builder with the 310P-specific version
|
||||
max_model_len = vllm_config.model_config.max_model_len
|
||||
self.attn_mask_builder: Any = AttentionMaskBuilder310(self.device, max_model_len)
|
||||
|
||||
self._query_lens_cpu_buffer: torch.Tensor | None = None
|
||||
if device.type != "cpu":
|
||||
max_num_seqs = vllm_config.scheduler_config.max_num_seqs
|
||||
self._query_lens_cpu_buffer = torch.empty(max_num_seqs, dtype=torch.int32, device="cpu", pin_memory=True)
|
||||
|
||||
def _fill_query_lens_cpu(
|
||||
self, num_reqs: int, query_start_loc_cpu: torch.Tensor, is_drafting: bool = False
|
||||
) -> torch.Tensor:
|
||||
"""Pinned CPU per-request query lengths for ATB splitfuse (host qLensTensor)."""
|
||||
if self._query_lens_cpu_buffer is None:
|
||||
return (query_start_loc_cpu[1 : num_reqs + 1] - query_start_loc_cpu[:num_reqs]).contiguous()
|
||||
if is_drafting:
|
||||
# We are using the same buffer for multi step drafting,
|
||||
# so we have to clone the buffer or the q lens of step 0
|
||||
# will be overwritten by the following steps.
|
||||
buffer = self._query_lens_cpu_buffer[:num_reqs].clone()
|
||||
else:
|
||||
buffer = self._query_lens_cpu_buffer[:num_reqs]
|
||||
torch.sub(
|
||||
query_start_loc_cpu[1 : num_reqs + 1],
|
||||
query_start_loc_cpu[:num_reqs],
|
||||
out=buffer,
|
||||
)
|
||||
return buffer
|
||||
|
||||
def build(
|
||||
self,
|
||||
common_prefix_len: int,
|
||||
common_attn_metadata: AscendCommonAttentionMetadata,
|
||||
fast_build: bool = False,
|
||||
is_drafting: bool = False,
|
||||
) -> AscendMetadata:
|
||||
attn_metadata = super().build(common_prefix_len, common_attn_metadata, fast_build)
|
||||
|
||||
num_reqs = common_attn_metadata.num_reqs
|
||||
|
||||
splitfuse_states = (
|
||||
AscendAttentionState.SpecDecoding,
|
||||
AscendAttentionState.ChunkedPrefill,
|
||||
)
|
||||
if attn_metadata.attn_state not in splitfuse_states:
|
||||
return attn_metadata
|
||||
|
||||
query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu[: num_reqs + 1]
|
||||
# ATB splitfuse qLensTensor must be host; filled here (outside graph forward).
|
||||
set_query_lens_cpu(
|
||||
attn_metadata,
|
||||
self._fill_query_lens_cpu(num_reqs, query_start_loc_cpu, is_drafting),
|
||||
)
|
||||
|
||||
# Bind device-side views for in-place graph replay updates.
|
||||
attn_metadata.seq_lens = common_attn_metadata.seq_lens[:num_reqs]
|
||||
attn_metadata.query_start_loc = common_attn_metadata.query_start_loc[: num_reqs + 1]
|
||||
|
||||
if is_compressed_mask_supported():
|
||||
attn_metadata.attn_mask = AttentionMaskBuilder310.get_compressed_splitfuse_mask(self.device)
|
||||
|
||||
return attn_metadata
|
||||
|
||||
def build_for_drafting(
|
||||
self,
|
||||
common_attn_metadata: CommonAttentionMetadata,
|
||||
draft_index: int,
|
||||
):
|
||||
# override build_for_drafting for passing status.
|
||||
return self.build(
|
||||
common_prefix_len=0, common_attn_metadata=common_attn_metadata, fast_build=True, is_drafting=True
|
||||
)
|
||||
Reference in New Issue
Block a user