-
Notifications
You must be signed in to change notification settings - Fork 2.4k
[Feat.]: Support 310P device run qwen2.5/3 dense and qwen2.5vl models #5774
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
bb7b303
76d742c
89857d5
2eeef4d
cb3e9a3
d9de8ec
a26064a
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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() | ||
|
|
||
|
|
||
| def AttentionMaskBuilder(device: torch.device) -> _AttentionMaskBuilder310P: | ||
| return _AttentionMaskBuilder310P(device) | ||
| 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
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The
Suggested change
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| 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) |
| 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) |
| 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 |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The methods
get_splitfuse_attn_maskandget_attention_maskuse 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 themodel_config, for example, by usingmodel_config.max_model_len.