Files
enginex-ascend-910-vllm/vllm_ascend/_310p/attention/attention_mask.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

187 lines
7.5 KiB
Python

#
# 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