Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -62,14 +62,17 @@ set(VLLM_ASCEND_CUSTOM_OP
)

set(VLLM_ASCEND_CUSTOM_OP_EXCLUDE
${KERNEL_FILES}/bgmv_expand.cpp
${KERNEL_FILES}/bgmv_shrink.cpp
${KERNEL_FILES}/sgmv_expand.cpp
${KERNEL_FILES}/sgmv_shrink.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/kernels/bgmv_expand.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/kernels/bgmv_shrink.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/kernels/sgmv_expand.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/kernels/sgmv_shrink.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/mla_preprocess/op_kernel/mla_preprocess_kernel.cpp
${CMAKE_CURRENT_SOURCE_DIR}/csrc/batch_matmul_transpose/op_kernel/batch_matmul_transpose_kernel.cpp
)

if(SOC_VERSION STREQUAL "ASCEND310P3")
message(STATUS "310P hardware detected: disabling MLAPO operators")
message(STATUS "310P hardware detected: excluding batch_matmul_transpose operators")
list(REMOVE_ITEM VLLM_ASCEND_CUSTOM_OP ${VLLM_ASCEND_CUSTOM_OP_EXCLUDE})
endif()

Expand Down
Empty file added vllm_ascend/_310p/__init__.py
Empty file.
Empty file.
61 changes: 61 additions & 0 deletions vllm_ascend/_310p/attention/attention_mask.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
from typing import Any, Callable, Optional

import torch

import vllm_ascend.attention.attention_mask as _base_mask

_BASE_BUILDER: Callable[[torch.device], Any] = _base_mask.AttentionMaskBuilder


def _gen_causal_additive_mask_fp16(max_seq_len: int,
device: torch.device) -> torch.Tensor:
tril = torch.ones((max_seq_len, max_seq_len),
dtype=torch.bool,
device=device).tril_()
upper = ~tril
m = torch.zeros((max_seq_len, max_seq_len),
dtype=torch.float16,
device=device)
m.masked_fill_(upper, float("-inf"))
return m


class _AttentionMaskBuilder310P:

def __init__(self, device: torch.device):
self._base = _BASE_BUILDER(device)

self._fp16_mask_cache: Optional[torch.Tensor] = None
self._fp16_mask_cached_len: int = 0

def __getattr__(self, name: str) -> Any:
return getattr(self._base, name)

@property
def device(self) -> torch.device:
return self._base.device

def _get_fp16_mask(self, max_seq_len: int) -> torch.Tensor:
if self._fp16_mask_cache is None or max_seq_len > self._fp16_mask_cached_len:
self._fp16_mask_cache = _gen_causal_additive_mask_fp16(
max_seq_len, self.device)
self._fp16_mask_cached_len = max_seq_len
assert self._fp16_mask_cache is not None
return self._fp16_mask_cache[:max_seq_len, :max_seq_len].contiguous()

def get_attn_mask(self, max_seq_len: int, dtype: torch.dtype):
if dtype == torch.float16:
return self._get_fp16_mask(max_seq_len)
return self._base.get_attn_mask(max_seq_len, dtype)

def get_splitfuse_attn_mask(self) -> torch.Tensor:
return self._get_fp16_mask(2048)

def get_attention_mask(self, model_config) -> torch.Tensor:
if getattr(model_config, "runner_type", None) == "pooling":
return self._base.get_attn_mask(2048, torch.bool)
return self.get_splitfuse_attn_mask()
Comment on lines +51 to +57

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The methods get_splitfuse_attn_mask and get_attention_mask use a hardcoded sequence length of 2048. This could lead to incorrect behavior or errors for models with a different maximum sequence length. It would be more robust to derive this value from the model_config, for example, by using model_config.max_model_len.

Suggested change
def get_splitfuse_attn_mask(self) -> torch.Tensor:
return self._get_fp16_mask(2048)
def get_attention_mask(self, model_config) -> torch.Tensor:
if getattr(model_config, "runner_type", None) == "pooling":
return self._base.get_attn_mask(2048, torch.bool)
return self.get_splitfuse_attn_mask()
def get_splitfuse_attn_mask(self, max_seq_len: int) -> torch.Tensor:
return self._get_fp16_mask(max_seq_len)
def get_attention_mask(self, model_config) -> torch.Tensor:
# Fallback to 2048 if max_model_len is not available.
max_seq_len = getattr(model_config, "max_model_len", 2048)
if getattr(model_config, "runner_type", None) == "pooling":
return self._base.get_attn_mask(max_seq_len, torch.bool)
return self.get_splitfuse_attn_mask(max_seq_len)



def AttentionMaskBuilder(device: torch.device) -> _AttentionMaskBuilder310P:
return _AttentionMaskBuilder310P(device)
113 changes: 113 additions & 0 deletions vllm_ascend/_310p/attention/attention_v1.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,113 @@
import torch
import torch_npu

from vllm_ascend._310p.attention.attention_mask import AttentionMaskBuilder
from vllm_ascend._310p.attention.metadata_builder import \
AscendAttentionMetadataBuilder310P
from vllm_ascend.attention.attention_v1 import \
AscendAttentionBackend as _BaseBackend
from vllm_ascend.attention.attention_v1 import \
AscendAttentionBackendImpl as _BaseImpl
from vllm_ascend.attention.attention_v1 import (AscendAttentionMetadataBuilder,
AscendAttentionState)
from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, aligned_16, nd_to_nz_2d


class AscendAttentionBackend310(_BaseBackend):

def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.attn_mask_builder = AttentionMaskBuilder(self.device)

@staticmethod
def get_kv_cache_shape(num_blocks: int, block_size: int, num_kv_heads: int,
head_size: int):
return (2, num_blocks, (num_kv_heads * head_size) // 16, block_size,
16)

@staticmethod
def get_impl_cls():
return AscendAttentionBackendImpl310

@staticmethod
def get_builder_cls() -> type["AscendAttentionMetadataBuilder"]:
return AscendAttentionMetadataBuilder310P


class AscendMLABackend310(AscendAttentionBackend310):
pass


class AscendSFABackend310(AscendAttentionBackend310):
pass


class AscendAttentionBackendImpl310(_BaseImpl):

def forward_paged_attention(self, query, attn_metadata, output):
if attn_metadata.seq_lens.device != query.device:
attn_metadata.seq_lens = attn_metadata.seq_lens.to(
device=query.device, non_blocking=True)
return super().forward_paged_attention(query, attn_metadata, output)

def _forward_prefill_310p_fallback(self, query, key, value, attn_metadata,
output):
real_tokens = int(attn_metadata.seq_lens.sum().item())

query, key, value, output = (aligned_16(t)
for t in (query, key, value, output))

seq_len = attn_metadata.seq_lens
if seq_len.dtype != torch.int32:
seq_len = seq_len.to(torch.int32)

aligned_tokens = int(query.shape[0])
delta = aligned_tokens - real_tokens
if delta:
seq_len = seq_len.clone()
seq_len[-1] += delta

mask = attn_metadata.attn_mask
if mask is not None and mask.dim() == 2:
max_len = int(seq_len.max().item())
aligned_len = ((max_len + 15) // 16) * 16

mask2d = mask[:aligned_len, :aligned_len].contiguous()
mask2d = mask2d.to(torch.float16)
mask_nz = nd_to_nz_2d(mask2d).contiguous()

bsz = int(seq_len.numel())
if bsz > 1:
mask_nz = mask_nz.repeat(bsz, 1, 1, 1).contiguous()

mask = torch_npu.npu_format_cast(mask_nz, ACL_FORMAT_FRACTAL_NZ)

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,
)

out_real = output[:real_tokens, :, :]
return out_real

def forward_impl(self, query, key, value, kv_cache, attn_metadata, output):
if attn_metadata.attn_state == AscendAttentionState.DecodeOnly:
output = self.forward_paged_attention(query, attn_metadata, output)

if attn_metadata.attn_state == AscendAttentionState.PrefillNoCache:
num_tokens = query.shape[0]
q = query[:num_tokens]
k = key[:num_tokens]
v = value[:num_tokens]
out = self._forward_prefill_310p_fallback(q, k, v, attn_metadata,
output)
output[:num_tokens] = out

return output
Comment on lines +100 to +113

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

The forward_impl method only handles DecodeOnly and PrefillNoCache attention states. Other states from the AscendAttentionState enum, such as PrefillCacheHit, ChunkedPrefill, and SpecDecoding, are not handled. This will cause incorrect behavior when these attention states occur, as the function will return the output tensor without processing it. This is a critical bug that needs to be addressed. Additionally, the use of two separate if statements is incorrect; an if/elif structure should be used to ensure only one path is taken.

Suggested change
def forward_impl(self, query, key, value, kv_cache, attn_metadata, output):
if attn_metadata.attn_state == AscendAttentionState.DecodeOnly:
output = self.forward_paged_attention(query, attn_metadata, output)
if attn_metadata.attn_state == AscendAttentionState.PrefillNoCache:
num_tokens = query.shape[0]
q = query[:num_tokens]
k = key[:num_tokens]
v = value[:num_tokens]
out = self._forward_prefill_310p_fallback(q, k, v, attn_metadata, output)
output[:num_tokens] = out
return output
def forward_impl(self, query, key, value, kv_cache, attn_metadata, output):
if attn_metadata.attn_state == AscendAttentionState.DecodeOnly:
output = self.forward_paged_attention(query, attn_metadata, output)
elif attn_metadata.attn_state == AscendAttentionState.PrefillNoCache:
num_tokens = query.shape[0]
q = query[:num_tokens]
k = key[:num_tokens]
v = value[:num_tokens]
out = self._forward_prefill_310p_fallback(q, k, v, attn_metadata, output)
output[:num_tokens] = out
else:
raise NotImplementedError(
f"Attention state {attn_metadata.attn_state} is not yet supported in AscendAttentionBackendImpl310."
)
return output

25 changes: 25 additions & 0 deletions vllm_ascend/_310p/attention/metadata_builder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
from __future__ import annotations

from typing import Any

import torch
from vllm.config import VllmConfig
from vllm.v1.kv_cache_interface import AttentionSpec

from vllm_ascend._310p.attention.attention_mask import AttentionMaskBuilder
from vllm_ascend.attention.attention_v1 import \
AscendAttentionMetadataBuilder as _BaseBuilder


class AscendAttentionMetadataBuilder310P(_BaseBuilder):

def __init__(
self,
kv_cache_spec: AttentionSpec,
layer_names: list[str],
vllm_config: VllmConfig,
device: torch.device,
):
super().__init__(kv_cache_spec, layer_names, vllm_config, device)

self.attn_mask_builder: Any = AttentionMaskBuilder(self.device)
108 changes: 108 additions & 0 deletions vllm_ascend/_310p/modelrunner_310p.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
from __future__ import annotations

from typing import Any, Dict

import torch
import torch_npu
from vllm.logger import logger
from vllm.v1.kv_cache_interface import KVCacheConfig

from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ
from vllm_ascend.worker.model_runner_v1 import NPUModelRunner


class NPUModelRunner310(NPUModelRunner):

def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._acl_format = ACL_FORMAT_FRACTAL_NZ

def _num_attn_module(self) -> int:
return 2 if self.model_config.hf_config.model_type == "longcat_flash" else 1

def _initialize_kv_cache_tensors_310p(
self, kv_cache_config: "KVCacheConfig") -> dict[str, Any]:
from vllm.v1.kv_cache_interface import FullAttentionSpec
from vllm.v1.worker.utils import bind_kv_cache

if self.vllm_config.kv_transfer_config is not None:
raise ValueError("KV cache transfer is not supported for 310P.")

kv_cache_sizes: dict[str, int] = {}
for kv_cache_tensor in kv_cache_config.kv_cache_tensors:
assert len(kv_cache_tensor.shared_by) == 1, (
"KV cache tensor shared by multiple layers is not supported in 310P."
)
kv_cache_sizes[kv_cache_tensor.shared_by[0]] = kv_cache_tensor.size

kv_caches: Dict[str, Any] = {}

for group in self._kv_cache_spec_attn_group_iterator():
kv_cache_spec = group.kv_cache_spec
attn_backend = group.backend

if not isinstance(kv_cache_spec, FullAttentionSpec):
raise ValueError("Unknown KV cache spec type.")

for layer_name in group.layer_names:
if layer_name in self.runner_only_attn_layers:
continue

tensor_size = kv_cache_sizes[layer_name]
assert tensor_size % kv_cache_spec.page_size_bytes == 0
num_blocks = tensor_size // kv_cache_spec.page_size_bytes
assert num_blocks >= kv_cache_config.num_blocks

if self.vllm_config.additional_config.get(
"kv_cache_dtype", None) == "int8":
kv_cache_shape = attn_backend.get_bsh_kv_cache_shape(
num_blocks,
kv_cache_spec.block_size,
kv_cache_spec.num_kv_heads,
kv_cache_spec.head_size,
)
elif hasattr(
attn_backend,
"get_supported_block_size") and self.use_hybrid_blocks:
block_size = attn_backend.get_supported_block_size()[0]
block_size_chunk = kv_cache_spec.block_size // block_size
kv_cache_shape = attn_backend.get_kv_cache_shape(
num_blocks * block_size_chunk,
block_size,
kv_cache_spec.num_kv_heads,
kv_cache_spec.head_size,
)
else:
kv_cache_shape = attn_backend.get_kv_cache_shape(
num_blocks,
kv_cache_spec.block_size,
kv_cache_spec.num_kv_heads,
kv_cache_spec.head_size,
)

dtype = kv_cache_spec.dtype

if "attn" in layer_name:
k_tensor = torch.zeros(kv_cache_shape[1:],
dtype=dtype,
device=self.device)
v_tensor = torch.zeros(kv_cache_shape[1:],
dtype=dtype,
device=self.device)
k_cache = torch_npu.npu_format_cast(
k_tensor, self._acl_format)
v_cache = torch_npu.npu_format_cast(
v_tensor, self._acl_format)
kv_caches[layer_name] = (k_cache, v_cache)

bind_kv_cache(
kv_caches,
self.compilation_config.static_forward_context,
self.kv_caches,
self._num_attn_module(),
)
return kv_caches

def initialize_kv_cache_tensors(
self, kv_cache_config: "KVCacheConfig") -> dict[str, Any]:
return self._initialize_kv_cache_tensors_310p(kv_cache_config)
Empty file.
15 changes: 15 additions & 0 deletions vllm_ascend/_310p/ops/activation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
import torch
import torch.nn.functional as F

from vllm_ascend.ops.activation import AscendSiluAndMul as _Base


class AscendSiluAndMul310(_Base):

def forward(self, x: torch.Tensor) -> torch.Tensor:
torch.ops.vllm.maybe_prefetch_mlp_down_proj(x)
h = x.shape[-1] // 2
out = (F.silu(x[..., :h].to(torch.float32)) *
x[..., h:].to(torch.float32)).to(torch.float16)
torch.ops.vllm.maybe_wait_prefetch_done(out)
return out
Loading
Loading