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

1760 lines
66 KiB
Python

#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
# Adapted from vllm-project/vllm/vllm/worker/worker.py
#
from __future__ import annotations
import functools
import json
import math
import os
from contextlib import nullcontext
from enum import Enum
from functools import lru_cache
from typing import TYPE_CHECKING, Any
import numpy as np
import regex as re
import torch
import torch_npu # noqa: F401
from packaging.version import InvalidVersion, Version
from vllm.logger import logger
from vllm.sequence import IntermediateTensors
import vllm_ascend.envs as envs_ascend
from vllm_ascend.ascend_config import WeightPrefetchConfig, get_ascend_config
if TYPE_CHECKING:
from vllm.config import VllmConfig
from vllm.v1.kv_cache_interface import AttentionSpec
else:
VllmConfig = None
COMPILATION_PASS_KEY = "graph_fusion_manager"
ASCEND_QUANTIZATION_METHOD = "ascend"
COMPRESSED_TENSORS_METHOD = "compressed-tensors"
FP8_METHOD = "fp8"
SOC_VERSION_INFERENCE_SERIES = ["Ascend310P3"]
REGISTERED_ASCEND_OPS = {}
ACL_FORMAT_FRACTAL_ND = 2
ACL_FORMAT_FRACTAL_NZ = 29
_CUSTOM_OP_ENABLED = None
_DEVICE_PRINT_OP_REGISTERED = False
_CURRENT_STREAM = None
_PREFETCH_STREAM = None
_WEIGHT_PREFETCH_METHOD = None
_GLOBAL_STREAM = None
_SHARED_EXPERTS_CALCULATION_STREAM = None
_CP_CHUNKEDPREFILL_COMM_STREAM = None
_ASCEND_CUSTOMOP_IS_REIGISTERED = False
_DEFAULT_BUFFER_SIZE = 200
_MIN_DP_BUFFER_SIZE = 50
_DYNAMIC_EPLB_BUFFER_SIZE = 100
_IS_MOE_MODEL = None
_IS_DRAFTER_MOE_MODEL = None
_IS_VL_MODEL = None
_ENABLE_SP = None
_HAS_LAYER_IDX = None
_HAS_ROPE = None
_ATNN_CALCULATION_STREAM = None
_CUSTOM_OP_VENDOR_DIR = "custom_transformer"
_CUSTOM_OP_BASE_DIR = (
os.path.dirname(__file__) if os.path.isabs(__file__) else os.path.abspath(os.path.dirname(__file__))
)
def extract_dsv4_layer_index(config: Any, layer_name: str) -> int:
"""Extract DSV4 index for config per-layer arrays.
Runtime module names keep their original MTP namespace, e.g. ``mtp.0``.
When indexing config-level arrays such as ``compress_ratios``, MTP layers
are addressed after the main model layers.
"""
from vllm.model_executor.models.utils import extract_layer_index
layer_idx = extract_layer_index(layer_name)
# TODO(zzzzwwjj): the layer idx of mtp should be aligned with vLLM
if ".mtp." in f".{layer_name}." and layer_idx < config.num_hidden_layers:
return config.num_hidden_layers + layer_idx
return layer_idx
def get_dsv4_spec_layer_idx_from_weight_name(config: Any, weight_name: str) -> int | None:
"""Return local MTP layer index for DSV4 checkpoint weight names."""
if weight_name.startswith("mtp."):
return int(weight_name.split(".")[1])
return None
def get_dsv4_compress_ratio(config: Any, layer_idx: int) -> int:
"""Return DSV4 compress ratio, treating unspecified MTP layers as dense."""
compress_ratios = getattr(config, "compress_ratios", None)
if compress_ratios is None or layer_idx >= len(compress_ratios):
return 0
return compress_ratios[layer_idx]
def model_uses_sfa_sparse(model_config: Any | None) -> bool:
hf_text_config = getattr(model_config, "hf_text_config", None)
hf_config = getattr(model_config, "hf_config", None)
return (
hf_text_config is not None
and hasattr(hf_text_config, "index_topk")
and not hasattr(hf_text_config, "compress_ratios")
and not hasattr(hf_config, "compress_ratios")
)
def enable_sfa_dcp_replicated_indexer(vllm_config: VllmConfig | None = None) -> bool:
if vllm_config is None:
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
parallel_config = vllm_config.parallel_config
return (
model_uses_sfa_sparse(vllm_config.model_config)
and parallel_config.decode_context_parallel_size > 1
and parallel_config.prefill_context_parallel_size == 1
)
def clear_enable_sp():
global _ENABLE_SP
_ENABLE_SP = None
enable_dsa_cp.cache_clear()
enable_dsa_cp_with_layer_shard.cache_clear()
enable_dsa_cp_with_o_proj_tp.cache_clear()
_libc_getenv.cache_clear()
def is_310p():
return get_ascend_device_type() == AscendDeviceType._310P
_IS_RC_DEVICE: bool | None = None
def is_rc_device() -> bool:
"""Return True if the 310P NPU runs in Root Complex (RC) mode.
RC mode (e.g. Atlas 200I Pro): host and NPU share memory. EP mode
(e.g. Atlas 300I DUO on PCIe): ``lspci`` output typically contains
``accelerators``.
"""
global _IS_RC_DEVICE
if not is_310p():
return False
if _IS_RC_DEVICE is not None:
return _IS_RC_DEVICE
try:
import subprocess
result = subprocess.run(["lspci"], capture_output=True, text=True, check=True)
_IS_RC_DEVICE = not any("accelerators" in line.strip() for line in result.stdout.splitlines())
except (subprocess.CalledProcessError, FileNotFoundError):
_IS_RC_DEVICE = False
return _IS_RC_DEVICE
def is_950():
return get_ascend_device_type() == AscendDeviceType.A5
def _mark_op_side_effectful(op: Any) -> None:
torch.fx.node.has_side_effect(op)
default_overload = getattr(op, "default", None)
if default_overload is not None:
torch.fx.node.has_side_effect(default_overload)
def _ensure_device_print_registered() -> None:
global _DEVICE_PRINT_OP_REGISTERED
if _DEVICE_PRINT_OP_REGISTERED:
return
if not enable_custom_op():
raise RuntimeError(
"device_print requires _C_ascend.device_print ops to be available "
"when custom ops are enabled in the current Ascend build."
)
try:
# Mark device_print ops side-effectful so FX/Inductor does not DCE or reorder these debug callbacks.
_mark_op_side_effectful(torch.ops._C_ascend.device_print)
_mark_op_side_effectful(torch.ops._C_ascend.device_print_tensor)
_DEVICE_PRINT_OP_REGISTERED = True
except AttributeError as exc:
raise RuntimeError(
"device_print requires _C_ascend.device_print ops to be available "
"when custom ops are enabled in the current Ascend build."
) from exc
def device_print(
value: torch.Tensor | int | float | bool | str | torch.dtype | torch.device | torch.Size,
) -> None:
"""Print one value from a device callback.
This helper is intended for debugging. To stay replay-safe under
``torch.npu.graph`` capture/replay, the underlying callback payloads are
retained instead of being reclaimed after the first host callback runs.
Avoid using it in hot paths or long-running high-frequency loops, otherwise
there may be memory issues due to too many retained payloads.
Supported usage:
>>> from vllm_ascend.utils import device_print
>>> device_print(x)
>>> device_print("already formatted text")
>>> device_print(7)
Unsupported usage:
>>> device_print("x =", x)
>>> device_print("This is ", x, "and this is ", y)
If you need device-time tensor values, pass the tensor itself. If you need
text, pass one final string that is already formatted.
Tensor values are copied to host on the current stream before the callback
prints them, so printing remains ordered with respect to the surrounding
device work.
DO NOT FORMAT A DEVICE TENSOR INTO A STRING YOURSELF AND THEN PRINT, for example:
>>> device_print(f"x = {x}")
>>> device_print("x = " + str(x))
"""
_ensure_device_print_registered()
if isinstance(value, torch.Tensor):
torch.ops._C_ascend.device_print_tensor(value)
elif isinstance(value, (str, int, float, bool, torch.dtype, torch.device, torch.Size)):
torch.ops._C_ascend.device_print(str(value))
else:
raise TypeError(
f"Unsupported device_print value type: {type(value)!r}. "
"Use exactly one argument: device_print(tensor), device_print('formatted text')."
)
def _should_trans_nz(weight: torch.Tensor) -> bool:
# FP32 cannot use NZ.
if weight.dtype == torch.float32:
return False
# meta tensor only keeps shape/dtype meta info without physical memory, it is not necessary to trans it to NZ
if weight.is_meta:
return False
# 310P always converts to NZ.
if is_310p():
return True
# Get config value instead of env
config = get_ascend_config()
nz_mode = config.weight_nz_mode
# NZ is disabled when mode is 0.
if not nz_mode:
return False
# BF16/FP16 convert only when nz_mode == 2.
if weight.dtype in {torch.bfloat16, torch.float16}:
return nz_mode == 2
# Quantized or other supported dtypes convert by default.
return True
# NZ conversion policy:
# - 310P: always convert supported weights to FRACTAL_NZ
# - non-310P: follow VLLM_ASCEND_ENABLE_NZ
# - FP32: never convert
# - meta tensor: never convert
def maybe_trans_nz(weight: torch.Tensor) -> torch.Tensor:
if not _should_trans_nz(weight):
return weight
return torch_npu.npu_format_cast(weight, ACL_FORMAT_FRACTAL_NZ)
def _round_up(x: int, align: int):
# round up x to align, for example, if align is 16, x will be rounded up to 16, 32, 48, etc.
# input: 15, 16 -> output: 16
# input: 17, 16 -> output: 32
# input: 30, 16 -> output: 32
# input: 33, 16 -> output: 48
# ...
return (x + align - 1) // align * align
def _prepend_env_path(env_name: str, path: str) -> None:
current_value = os.environ.get(env_name, "")
path_entries = [entry for entry in current_value.split(":") if entry]
if path not in path_entries:
path_entries.insert(0, path)
os.environ[env_name] = ":".join(path_entries)
def bootstrap_custom_op_env(*, include_vendor_lib: bool = False) -> None:
vendor_path = os.path.join(_CUSTOM_OP_BASE_DIR, "_cann_ops_custom", "vendors", _CUSTOM_OP_VENDOR_DIR)
if not os.path.exists(vendor_path):
return
_prepend_env_path("ASCEND_CUSTOM_OPP_PATH", vendor_path)
if include_vendor_lib:
vendor_lib_path = os.path.join(vendor_path, "op_api", "lib")
if os.path.exists(vendor_lib_path):
_prepend_env_path("LD_LIBRARY_PATH", vendor_lib_path)
def _custom_pad(x, pad_dims):
# pad the input tensor to the shape of pad_dims
# input: (13, 30), pad_dims: [0, 2, 0, 3]
# output: (16, 32)
return torch.nn.functional.pad(x, pad_dims)
def _custom_reshape(x, target_shape):
# reshape the input tensor to the shape of target_shape
# input: (16, 32), target_shape: [1, 16, 2, 16]
# output: (1, 16, 2, 16)
return x.reshape(target_shape)
def _custom_transpose(x, dim1, dim2):
# transpose the input tensor
# input: (1, 16, 2, 16), dim1: 1, dim2: 2
# output: (1, 2, 16, 16)
return x.transpose(dim1, dim2)
def nd_to_nz_2d(in_tensor: torch.Tensor) -> torch.Tensor:
# in_tensor: (13, 30)
aux_dims = [1, 0, 0, 16]
# aux_dims[1]: 16
aux_dims[1] = _round_up(in_tensor.size(0), 16)
# aux_dims[2]: 2
aux_dims[2] = _round_up(in_tensor.size(1), 16) // 16
# after: aux_dims: [1, 16, 2, 16]
pad_dims = [0, 0, 0, 0]
# pad_dims[1]: 2
pad_dims[1] = _round_up(in_tensor.size(1), 16) - in_tensor.size(1)
# pad_dims[3]: 3
pad_dims[3] = _round_up(in_tensor.size(0), 16) - in_tensor.size(0)
# after: pad_dims: [0, 2, 0, 3]
# return: (1, 2, 16, 16)
return _custom_transpose(_custom_reshape(_custom_pad(in_tensor, pad_dims), aux_dims), 1, 2).contiguous()
def nd_to_nz_spec(mask_tensor: torch.Tensor) -> torch.Tensor:
num_tokens = mask_tensor.shape[0]
max_seq_len = mask_tensor.shape[1]
tokens_pad = (num_tokens + 15) // 16 * 16
max_seq_len_pad = (max_seq_len + 15) // 16 * 16
mask_tensor_pad = torch.zeros((1, tokens_pad, max_seq_len_pad), dtype=mask_tensor.dtype, device=mask_tensor.device)
mask_tensor_pad[0][:num_tokens, :max_seq_len] = mask_tensor
mask = mask_tensor_pad.reshape((1, tokens_pad, max_seq_len_pad // 16, 16)).permute(0, 2, 1, 3)
return mask
def aligned_16(tensor: torch.Tensor):
"""Aligned tensor for 310P"""
# Get the size of the current 0th dimension
n = tensor.size(0)
# Calculate the aligned size
n_aligned = ((n + 15) // 16) * 16
# If already aligned, return the original tensor
if n == n_aligned:
return tensor
# Create a new tensor with shape (n_aligned, H, W) and fill it with zeros
new_tensor = torch.zeros(n_aligned, *tensor.shape[1:], dtype=tensor.dtype, device=tensor.device)
# Copy the original tensor to the first N positions of the new tensor
new_tensor[:n] = tensor
return new_tensor
def enable_custom_op():
"""
Enable lazy init for vllm_ascend_C to avoid early initialization of CANN's RTS component.
Ensure that ASCEND_RT_VISIBLE_DEVICES can be dynamically modified before torch.npu.set_device().
"""
import vllm.envs as envs
global _CUSTOM_OP_ENABLED
if _CUSTOM_OP_ENABLED is not None:
return _CUSTOM_OP_ENABLED
# There are some customed operators which aren't implemented
# with batch invariant in vllm-ascend, we need to disable them.
# FIXME(linfeng): Currently custom op compilation and execution are partially available
# in ASCEND950 chip, we temporarily disable all custom ops. Please refer to
# https://github.com/vllm-project/vllm-ascend/issues/7157 for latest update about custom op.
if envs.VLLM_BATCH_INVARIANT or get_ascend_device_type() == AscendDeviceType.A5:
_CUSTOM_OP_ENABLED = False
return _CUSTOM_OP_ENABLED
try:
if not torch.compiler.is_compiling():
bootstrap_custom_op_env()
# isort: off
# register custom ops into torch_library here
import vllm_ascend.vllm_ascend_C # type: ignore # noqa: F401
# register the meta implementation for custom kernel if necessary
import vllm_ascend.meta_registration # type: ignore # noqa: F401
# isort: on
_CUSTOM_OP_ENABLED = True
except ImportError as e:
# Prefer the extension's rpath for vendor op_api loading. Only fall back
# to mutating LD_LIBRARY_PATH when the import proves it is still needed.
if (not torch.compiler.is_compiling()) and "libcust_opapi.so" in str(e):
try:
bootstrap_custom_op_env(include_vendor_lib=True)
import vllm_ascend.meta_registration # type: ignore # noqa: F401
import vllm_ascend.vllm_ascend_C # type: ignore # noqa: F401
_CUSTOM_OP_ENABLED = True
except ImportError:
_CUSTOM_OP_ENABLED = False
logger.warning(
"Failed to register custom ops, all custom ops will be disabled. "
"The custom ops library might not be installed or the environment is not configured correctly. "
"Please check the custom ops installation and environment variables."
)
else:
_CUSTOM_OP_ENABLED = False
logger.warning(
"Failed to register custom ops, all custom ops will be disabled. "
"error=%s. "
"The custom ops library might not be installed or the environment is not configured correctly. "
"Please check the custom ops installation and environment variables.",
e,
)
return _CUSTOM_OP_ENABLED
def find_hccl_library() -> str:
"""
We either use the library file specified by the `HCCL_SO_PATH`
environment variable, or we find the library file brought by PyTorch.
After importing `torch`, `libhccl.so` can be
found by `ctypes` automatically.
"""
so_file = envs_ascend.HCCL_SO_PATH
# manually load the hccl library
if so_file:
logger.info("Found hccl from environment variable HCCL_SO_PATH=%s", so_file)
else:
if torch.version.cann is not None:
so_file = "libhccl.so"
else:
raise ValueError("HCCL only supports Ascend NPU backends.")
logger.info("Found hccl from library %s", so_file)
return so_file
def current_stream() -> torch.npu.Stream:
"""
replace `torch.npu.current_stream()` with `vllm.utils.current_stream()`.
it turns out that `torch.npu.current_stream()` is quite expensive,
as it will construct a new stream object at each call.
here we patch `torch.npu.set_stream` to keep track of the current stream
directly, so that we can avoid calling `torch.npu.current_stream()`.
"""
global _CURRENT_STREAM
if _CURRENT_STREAM is None:
# when this function is called before any stream is set,
# we return the default stream.
_CURRENT_STREAM = torch.npu.current_stream()
return _CURRENT_STREAM
def prefetch_stream() -> torch.npu.Stream:
global _PREFETCH_STREAM
if _PREFETCH_STREAM is None:
# when this function is called before any stream is set,
# we return the default stream.
_PREFETCH_STREAM = torch_npu.npu.Stream()
return _PREFETCH_STREAM
def set_weight_prefetch_method(weight_prefetch_config: WeightPrefetchConfig):
global _WEIGHT_PREFETCH_METHOD
if _WEIGHT_PREFETCH_METHOD is None:
from vllm_ascend.ops.weight_prefetch import WeightPrefetchMethod
_WEIGHT_PREFETCH_METHOD = WeightPrefetchMethod(weight_prefetch_config)
return _WEIGHT_PREFETCH_METHOD
def get_weight_prefetch_method():
return _WEIGHT_PREFETCH_METHOD
def global_stream() -> torch.npu.Stream:
global _GLOBAL_STREAM
if _GLOBAL_STREAM is None:
# when this function is called before any stream is set,
# we return the default stream.
_GLOBAL_STREAM = torch_npu.npu.Stream()
return _GLOBAL_STREAM
def shared_experts_calculation_stream() -> torch.npu.Stream:
global _SHARED_EXPERTS_CALCULATION_STREAM
if _SHARED_EXPERTS_CALCULATION_STREAM is None:
# when this function is called before any stream is set,
# we return the default stream.
_SHARED_EXPERTS_CALCULATION_STREAM = torch_npu.npu.Stream()
return _SHARED_EXPERTS_CALCULATION_STREAM
def cp_chunkedprefill_comm_stream() -> torch.npu.Stream:
global _CP_CHUNKEDPREFILL_COMM_STREAM
if _CP_CHUNKEDPREFILL_COMM_STREAM is None:
_CP_CHUNKEDPREFILL_COMM_STREAM = torch_npu.npu.Stream()
return _CP_CHUNKEDPREFILL_COMM_STREAM
def attention_calculation_stream() -> torch.npu.Stream:
global _ATNN_CALCULATION_STREAM
if _ATNN_CALCULATION_STREAM is None:
_ATNN_CALCULATION_STREAM = torch_npu.npu.Stream()
return _ATNN_CALCULATION_STREAM
def adapt_patch(is_global_patch: bool = False):
if is_global_patch:
from vllm_ascend.patch import platform # noqa: F401
else:
from vllm_ascend.patch import worker # noqa: F401
def setup_ascend_local_comm_res(local_rank: int, kv_transfer_config: Any | None) -> None:
"""Load the local A5 endpoint config into ASCEND_LOCAL_COMM_RES."""
if kv_transfer_config is None:
return
visible_devices = os.getenv("ASCEND_RT_VISIBLE_DEVICES")
if visible_devices is None:
from vllm_ascend.cpu_binding import DeviceInfo
devices = sorted([int(x) for x in DeviceInfo.get_npu_map_info()])
else:
devices = [int(x) for x in visible_devices.split(",") if x.strip()]
extra_config = kv_transfer_config.kv_connector_extra_config or {}
local_comm_res_path = extra_config.get("ascend_local_comm_res_path")
if not local_comm_res_path:
return
if not devices:
raise ValueError("No NPU devices found or specified in ASCEND_RT_VISIBLE_DEVICES.")
if local_rank < 0 or local_rank >= len(devices):
raise ValueError(f"local_rank {local_rank} is out of bounds for the available NPU devices: {devices}")
local_comm_res_file = os.path.join(local_comm_res_path, f"ub_endpoint_npu_{devices[local_rank]}.json")
try:
with open(local_comm_res_file) as f:
data = json.load(f)
except FileNotFoundError:
raise FileNotFoundError(
f"Endpoint config file not found: {local_comm_res_file}. "
"Please set ascend_local_comm_res_path in kv_connector_extra_config "
"to a directory containing ub_endpoint_npu_*.json endpoint configuration files."
)
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse endpoint config file: {local_comm_res_file}") from e
os.environ["ASCEND_LOCAL_COMM_RES"] = json.dumps(data, ensure_ascii=False, separators=(",", ":"))
@functools.cache
def vllm_version_is(target_vllm_version: str):
if envs_ascend.VLLM_VERSION is not None:
vllm_version = envs_ascend.VLLM_VERSION
else:
import vllm
vllm_version = vllm.__version__
try:
return Version(vllm_version) == Version(target_vllm_version)
except InvalidVersion:
raise ValueError(
f"Invalid vllm version {vllm_version} found. A dev version of vllm "
"is installed probably. Set the environment variable VLLM_VERSION "
"to control it by hand. And please make sure the value follows the "
"format of x.y.z."
)
def get_max_hidden_layers(hf_config) -> int:
cfg_dict = hf_config.to_dict()
layer_counts = []
def _rec_find(d):
if isinstance(d, dict):
for k, v in d.items():
if k == "num_hidden_layers" and isinstance(v, int):
layer_counts.append(v)
else:
_rec_find(v)
_rec_find(cfg_dict)
if not layer_counts:
raise ValueError("Not found num_hidden_layers in model config.")
return max(layer_counts)
# Update cudagraph capture sizes for vllm config
def update_cudagraph_capture_sizes(vllm_config: VllmConfig, cudagraph_capture_sizes: list[int]):
valid_max_size = cudagraph_capture_sizes[-1] if cudagraph_capture_sizes else 0
if (
vllm_config.compilation_config.max_cudagraph_capture_size is not None
and vllm_config.compilation_config.max_cudagraph_capture_size != valid_max_size
):
if vllm_config.compilation_config.cudagraph_capture_sizes is not None:
raise ValueError(
"customized max_cudagraph_capture_size"
f"(={vllm_config.compilation_config.max_cudagraph_capture_size}) "
"should be consistent with the max value of "
f"cudagraph_capture_sizes(={valid_max_size})"
)
logger.warning(
"Truncating max_cudagraph_capture_size. "
"original_size=%d, truncated_size=%d. "
"The max_cudagraph_capture_size does not match the max value of cudagraph_capture_sizes. "
"Please check the compilation_config for consistency.",
vllm_config.compilation_config.max_cudagraph_capture_size,
valid_max_size,
)
vllm_config.compilation_config.max_cudagraph_capture_size = valid_max_size
if vllm_config.compilation_config.cudagraph_capture_sizes is not None and len(cudagraph_capture_sizes) < len(
vllm_config.compilation_config.cudagraph_capture_sizes
):
logger.warning(
"cudagraph_capture_sizes specified in compilation_config is overridden. "
"compilation_config_sizes=%s, overridden_sizes=%s. "
"The sizes are adjusted based on model configuration and resource constraints.",
vllm_config.compilation_config.cudagraph_capture_sizes,
cudagraph_capture_sizes,
)
vllm_config.compilation_config.cudagraph_capture_sizes = cudagraph_capture_sizes
vllm_config.compilation_config.post_init_cudagraph_sizes()
# TODO(wxy): Move to ops module
def dispose_tensor(x: torch.Tensor):
x.set_(torch.empty((0,), device=x.device, dtype=x.dtype))
def register_ascend_customop(vllm_config: VllmConfig | None = None):
"""Register Ascend CustomOP
NOTE: if the register branch requires model type, please use `vllm.config.get_current_vllm_config`,
and ensure this will execute after model config is initilazed.
"""
global _ASCEND_CUSTOMOP_IS_REIGISTERED
if _ASCEND_CUSTOMOP_IS_REIGISTERED:
return
from vllm.model_executor.custom_op import CustomOp
from vllm_ascend.ops.activation import (
AscendQuickGELU,
AscendSiluAndMul,
AscendSiluAndMulWithClamp,
)
from vllm_ascend.ops.bailing_moe_linear_attn import AscendBailingMoELinearAttention
from vllm_ascend.ops.conv import AscendConv3dLayer
from vllm_ascend.ops.gdn import AscendGatedDeltaNetAttention
from vllm_ascend.ops.layernorm import AscendGemmaRMSNorm, AscendRMSNorm, AscendRMSNormGated
from vllm_ascend.ops.linear import (
AscendColumnParallelLinear,
AscendMergedColumnParallelLinear,
AscendQKVParallelLinear,
AscendReplicatedLinear,
AscendRowParallelLinear,
)
from vllm_ascend.ops.mla import AscendMultiHeadLatentAttention
from vllm_ascend.ops.mm_encoder_attention import AscendMMEncoderAttention
from vllm_ascend.ops.qwen2_decoder import AscendCustomQwen2Decoder
from vllm_ascend.ops.rel_pos_attention import AscendRelPosAttention
from vllm_ascend.ops.rotary_embedding import (
AscendApplyRotaryEmb,
AscendDeepseekScalingRotaryEmbedding,
AscendMRotaryEmbedding,
AscendRotaryEmbedding,
AscendYaRNRotaryEmbedding,
)
from vllm_ascend.ops.vocab_parallel_embedding import (
AscendLogitsProcessor,
AscendParallelLMHead,
AscendVocabParallelEmbedding,
)
global REGISTERED_ASCEND_OPS
REGISTERED_ASCEND_OPS = {
"QuickGELU": AscendQuickGELU,
"SiluAndMul": AscendSiluAndMul,
"SiluAndMulClamp": AscendSiluAndMulWithClamp,
"RotaryEmbedding": AscendRotaryEmbedding,
"MRotaryEmbedding": AscendMRotaryEmbedding,
"ColumnParallelLinear": AscendColumnParallelLinear,
"RowParallelLinear": AscendRowParallelLinear,
"YaRNScalingRotaryEmbedding": AscendYaRNRotaryEmbedding,
"MergedColumnParallelLinear": AscendMergedColumnParallelLinear,
"QKVParallelLinear": AscendQKVParallelLinear,
"ReplicatedLinear": AscendReplicatedLinear,
"DeepseekScalingRotaryEmbedding": AscendDeepseekScalingRotaryEmbedding,
"VocabParallelEmbedding": AscendVocabParallelEmbedding,
"ParallelLMHead": AscendParallelLMHead,
"LogitsProcessor": AscendLogitsProcessor,
"RMSNorm": AscendRMSNorm,
"GemmaRMSNorm": AscendGemmaRMSNorm,
"MultiHeadLatentAttentionWrapper": AscendMultiHeadLatentAttention,
"MMEncoderAttention": AscendMMEncoderAttention,
"ApplyRotaryEmb": AscendApplyRotaryEmb,
"RMSNormGated": AscendRMSNormGated,
"Conv3dLayer": AscendConv3dLayer,
"RelPosAttention": AscendRelPosAttention,
"CustomQwen2Decoder": AscendCustomQwen2Decoder,
"GatedDeltaNetAttention": AscendGatedDeltaNetAttention,
"BailingMoELinearAttention": AscendBailingMoELinearAttention,
}
if vllm_version_is("0.23.0"):
from vllm_ascend.ops.fused_moe.fused_moe import AscendFusedMoE
REGISTERED_ASCEND_OPS["FusedMoE"] = AscendFusedMoE
if vllm_config is None:
try:
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
except AssertionError:
vllm_config = None
if vllm_config is not None and vllm_config.model_config.is_deepseek_mla:
from vllm_ascend.ops.fused_moe.gate_linear import AscendGateLinear
REGISTERED_ASCEND_OPS["GateLinear"] = AscendGateLinear
# 310P: override selected ops with 310P implementations (keep minimal changes outside _310p)
if is_310p():
from vllm_ascend._310p.ops.activation import AscendSiluAndMul310
from vllm_ascend._310p.ops.conv import AscendConv3dLayer310
from vllm_ascend._310p.ops.fla.gdn_310 import AscendGatedDeltaNetAttention310
from vllm_ascend._310p.ops.layernorm import (
AscendGemmaRMSNorm310,
AscendRMSNorm310,
AscendRMSNormGated310,
)
from vllm_ascend._310p.ops.mm_encoder_attention import AscendMMEncoderAttention310
from vllm_ascend._310p.ops.rotary_embedding import AscendMRotaryEmbedding310, AscendRotaryEmbedding310
from vllm_ascend._310p.ops.vocab_parallel_embedding import (
AscendParallelLMHead310,
AscendVocabParallelEmbedding310,
)
REGISTERED_ASCEND_OPS.update(
{
"SiluAndMul": AscendSiluAndMul310,
"RotaryEmbedding": AscendRotaryEmbedding310,
"RMSNorm": AscendRMSNorm310,
"GemmaRMSNorm": AscendGemmaRMSNorm310,
"RMSNormGated": AscendRMSNormGated310,
"ParallelLMHead": AscendParallelLMHead310,
"VocabParallelEmbedding": AscendVocabParallelEmbedding310,
"MMEncoderAttention": AscendMMEncoderAttention310,
"Conv3dLayer": AscendConv3dLayer310,
"GatedDeltaNetAttention": AscendGatedDeltaNetAttention310,
"MRotaryEmbedding": AscendMRotaryEmbedding310,
}
)
if vllm_version_is("0.23.0"):
from vllm_ascend._310p.fused_moe.fused_moe import AscendFusedMoE310
REGISTERED_ASCEND_OPS["FusedMoE"] = AscendFusedMoE310
for name, op_cls in REGISTERED_ASCEND_OPS.items():
CustomOp.register_oot(_decorated_op_cls=op_cls, name=name)
# NOTE: Keep this at last to ensure all custom actions are registered
_ASCEND_CUSTOMOP_IS_REIGISTERED = True
class AscendDeviceType(Enum):
A2 = 0
A3 = 1
_310P = 2
A5 = 3
_ascend_device_type = None
def _init_ascend_device_type():
global _ascend_device_type
from vllm_ascend import _build_info # type: ignore
device_type = getattr(_build_info, "__device_type__", None)
if device_type is None:
soc_version = getattr(_build_info, "__soc_version__", "ASCEND910B1").upper()
device_type = "_310P" if "310P" in soc_version else "A2"
_ascend_device_type = AscendDeviceType[device_type]
def check_ascend_device_type():
global _ascend_device_type
if _ascend_device_type is None:
_init_ascend_device_type()
soc_version = torch_npu.npu.get_soc_version()
if 220 <= soc_version <= 225:
cur_device_type = AscendDeviceType.A2
elif 250 <= soc_version <= 255:
cur_device_type = AscendDeviceType.A3
elif 200 <= soc_version <= 205:
cur_device_type = AscendDeviceType._310P
elif soc_version == 260:
cur_device_type = AscendDeviceType.A5
else:
raise RuntimeError(f"Can not support soc_version: {soc_version}.")
assert _ascend_device_type == cur_device_type, (
f"Current device type: {cur_device_type} does not match the installed version's device type: "
f"{_ascend_device_type}, please check your installation package."
)
def get_ascend_device_type():
global _ascend_device_type
if _ascend_device_type is None:
_init_ascend_device_type()
return _ascend_device_type
def lmhead_tp_enable() -> bool:
return get_ascend_config().finegrained_tp_config.lmhead_tensor_parallel_size > 0
def embedding_tp_enable() -> bool:
return get_ascend_config().finegrained_tp_config.embedding_tensor_parallel_size > 0
def oproj_tp_enable() -> bool:
return get_ascend_config().finegrained_tp_config.oproj_tensor_parallel_size > 0
def olora_tp_enable() -> bool:
return get_ascend_config().finegrained_tp_config.olora_tensor_parallel_size > 1
def mlp_tp_enable() -> bool:
return get_ascend_config().finegrained_tp_config.mlp_tensor_parallel_size > 0
def matmul_allreduce_enable() -> bool:
return get_ascend_config().enable_matmul_allreduce
def enable_sp_by_pass():
return get_ascend_config().enable_sp_by_pass
def enable_sp(vllm_config=None, enable_shared_expert_dp: bool = False) -> bool:
global _ENABLE_SP
if vllm_config is None:
try:
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
except AssertionError:
vllm_config = None
additional_config = getattr(vllm_config, "additional_config", None) if vllm_config is not None else None
refresh = additional_config.get("refresh", False) if additional_config else False
if _ENABLE_SP is None or refresh:
if additional_config is not None and "enable_flashcomm1" in additional_config:
_ENABLE_SP = bool(additional_config["enable_flashcomm1"])
else:
try:
_ENABLE_SP = get_ascend_config().enable_flashcomm1
except RuntimeError:
_ENABLE_SP = envs_ascend.VLLM_ASCEND_ENABLE_FLASHCOMM1
if not _ENABLE_SP and enable_shared_expert_dp:
_ENABLE_SP = True
logger.info("shared_expert_dp requires enable_sp=True. enable_sp has been set to True.")
return bool(_ENABLE_SP)
# TODO remove it after vllm has this func
def shared_expert_dp_enabled() -> bool:
return get_ascend_config().enable_shared_expert_dp or enable_sp() or enable_sp_by_pass()
def is_moe_model(vllm_config: VllmConfig):
"""Checks if the model is a MoE model by config"""
global _IS_MOE_MODEL
if _IS_MOE_MODEL is None:
model_configs = vllm_config.model_config.hf_text_config.to_dict()
_IS_MOE_MODEL = _is_contain_expert(model_configs)
return _IS_MOE_MODEL
def is_drafter_moe_model(vllm_config: VllmConfig):
"""Checks if the drafter model is a MoE model by config"""
global _IS_DRAFTER_MOE_MODEL
if _IS_DRAFTER_MOE_MODEL is None:
speculative_config = vllm_config.speculative_config
if speculative_config.method == "extract_hidden_states":
# The extract_hidden_states drafter is a cache-only attention layer
# (never MoE), but its hf_config is copied from the possibly-MoE
# target, so the expert-key scan below would misclassify it. Skip
# the scan to keep the drafter DP sync free of a spurious
# all_reduce that idle DP ranks never match.
_IS_DRAFTER_MOE_MODEL = False
return _IS_DRAFTER_MOE_MODEL
model_configs = speculative_config.draft_model_config.hf_text_config.to_dict()
_IS_DRAFTER_MOE_MODEL = _is_contain_expert(model_configs)
if not model_configs or not model_configs.get("architectures"):
return _IS_DRAFTER_MOE_MODEL
if "Eagle3DeepseekV2ForCausalLM" in model_configs["architectures"]:
_IS_DRAFTER_MOE_MODEL = False
return _IS_DRAFTER_MOE_MODEL
def speculative_enable_dispatch_gmm_combine_decode(vllm_config: VllmConfig) -> bool:
"""When draft contains MOE Arch and non-w8a8, disable dispatch_gmm_combine_decode."""
if vllm_config.speculative_config is None:
return True
speculative_method = getattr(vllm_config.speculative_config, "method", None)
if speculative_method in [None, "ngram", "suffix"]:
return True
if speculative_method in ["eagle", "eagle3"]:
if is_drafter_moe_model(vllm_config):
draft_model_config = vllm_config.speculative_config.draft_model_config
hf_text_config = draft_model_config.hf_text_config
quant_type = getattr(hf_text_config, "moe_quantize", None)
if quant_type is None:
quant_type = getattr(hf_text_config, "quantize", None)
return quant_type == "w8a8_dynamic"
else:
return True
if speculative_method == "mtp":
mtp_quant_type = getattr(vllm_config.model_config.hf_text_config, "mtp_quantize", None)
return mtp_quant_type == "w8a8_dynamic"
return False
def _is_contain_expert(config: Any):
if isinstance(config, dict):
for k, v in config.items():
if "expert" in str(k):
return True
if _is_contain_expert(v):
return True
return False
def is_vl_model(vllm_config: VllmConfig = None):
"""Checks if the model is a VL model by config.
Uses the same criterion as vllm itself (model_config.py): a model is
multimodal when its top-level hf_config differs from its hf_text_config
(i.e. there is a separate vision sub-config). The legacy key-name checks
are kept as fallbacks for configs that override get_text_config() to return
self (rare but possible).
"""
global _IS_VL_MODEL
if vllm_config is None:
from vllm.config import get_current_vllm_config_or_none
vllm_config = get_current_vllm_config_or_none()
if _IS_VL_MODEL is None and vllm_config and vllm_config.model_config:
model_config = vllm_config.model_config
# Primary: vllm's own VL detection — hf_config is the top-level
# (multimodal) config; hf_text_config is the language-model sub-config.
# They are the same object for pure-text models.
if model_config.hf_config is not model_config.hf_text_config:
_IS_VL_MODEL = True
else:
# Fallback: check well-known config keys
hf_config = model_config.hf_config.to_dict()
if "thinker_config" in hf_config or "vision_config" in hf_config:
_IS_VL_MODEL = True
else:
_IS_VL_MODEL = False
return _IS_VL_MODEL
def has_rope(vllm_config: VllmConfig):
"""Checks if the model uses rope."""
global _HAS_ROPE
if _HAS_ROPE is None and vllm_config and vllm_config.model_config:
hf_config = vllm_config.model_config.hf_text_config.to_dict()
_HAS_ROPE = "rope_parameters" in hf_config
return _HAS_ROPE
def weak_ref_tensor(tensor: Any) -> Any:
"""
Create a weak reference to a tensor.
The new tensor will share the same data as the original tensor,
but will not keep the original tensor alive.
"""
if isinstance(tensor, torch.Tensor):
return torch_npu._C._weak_ref_tensor(tensor)
else:
return tensor
def weak_ref_tensors(
tensors: torch.Tensor | list[torch.Tensor] | tuple[torch.Tensor],
) -> torch.Tensor | list[Any] | tuple[Any] | Any:
"""
Convenience function to create weak references to tensors,
for single tensor, list of tensors or tuple of tensors.
This function should be used in the following scenario:
When a tensor is created during graph capture, and it's held by a method
that's not part of the graph, we don't really need to store it, but we
**do need** its buffer pointer. If we don't handle this, it cannot
be garbage collected, leading to a memory leak. To avoid this,
we should create a weak reference to the tensor.
"""
if isinstance(tensors, torch.Tensor):
return weak_ref_tensor(tensors)
if isinstance(tensors, list):
return [weak_ref_tensor(t) for t in tensors]
if isinstance(tensors, tuple):
return tuple(weak_ref_tensor(t) for t in tensors)
# For IntermediateTensors used in pipeline parallelism
if isinstance(tensors, IntermediateTensors):
ret = IntermediateTensors({key: weak_ref_tensor(val) for key, val in tensors.tensors.items()})
return ret
raise ValueError("Invalid type for tensors")
def npu_stream_switch(target_stream: torch.npu.Stream, *, enabled: bool = True):
"""
Switch to the target stream if enabled is True.
Otherwise, do nothing.
"""
if not enabled:
return nullcontext()
assert target_stream is not None
return torch.npu.stream(target_stream)
def create_hccl_pg_options(group_name: str):
options = torch_npu._C._distributed_c10d.ProcessGroupHCCL.Options()
hccl_config = get_hccl_config_for_pg_options(group_name) or {}
hccl_config["group_name"] = group_name
options.hccl_config = hccl_config
return options
def get_hccl_config_for_pg_options(group_name: str) -> dict | None:
"""
Get HCCL process group options for the given communication group name.
Args:
group_name: Name of the communication group
Returns:
HCCL pg_options or None for mc2 group
"""
# FIXME: Current mc2 operators only perform communication space partitioning
# based on HCCL_BUFFSIZE configuration. Using pg_options with mc2 group would
# result in memory misalignment problems.
if group_name and "mc2" in group_name:
return None
hccl_config_map = {
"dp": {"hccl_buffer_size": calculate_dp_buffer_size()},
"dynamic_eplb": {"hccl_buffer_size": _DYNAMIC_EPLB_BUFFER_SIZE},
}
return hccl_config_map.get(group_name, get_default_buffer_config())
def get_default_buffer_config() -> dict:
return {"hccl_buffer_size": _DEFAULT_BUFFER_SIZE}
def calculate_dp_buffer_size() -> int:
"""
formula of dp buffer size:
dp_size + 1 (flags: with_prefill)
"""
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
dp_size = vllm_config.parallel_config.data_parallel_size
int32_size = torch.iinfo(torch.int32).bits // 8
dp_buffer_size = math.ceil((dp_size + 1) * int32_size / (1024 * 1024))
return max(dp_buffer_size, _MIN_DP_BUFFER_SIZE)
# Currently, when in A2, setting the environment variables HCCL_INTRA_PCIE_ENABLE=1
# and HCCL_INTRA_ROCE_ENABLE=0 can reduce cross-machine communication traffic and
# significantly improve communication performance of MC2 ops dispatch/combine.
def is_hierarchical_communication_enabled():
return (
os.getenv("HCCL_INTRA_ROCE_ENABLE", "") == "0" and os.getenv("HCCL_INTRA_PCIE_ENABLE", "") == "1"
) or get_ascend_config().enable_mc2_hierarchy_comm
def is_pd_decode_recompute_scheduler_enabled(vllm_config: VllmConfig | None = None) -> bool:
"""True on PD-disaggregated decode nodes with recompute_scheduler_enable.
After KV recv, RecomputeScheduler sets num_computed_tokens to N-1 so the
decode node recomputes the last prompt token before MTP decode. Worker
metadata must not treat that step as prefill.
"""
try:
if vllm_config is None:
try:
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
except AssertionError:
vllm_config = get_ascend_config().vllm_config
if vllm_config is None:
return False
kv_cfg = vllm_config.kv_transfer_config
if kv_cfg is None or not kv_cfg.is_kv_consumer or kv_cfg.is_kv_producer:
return False
return get_ascend_config().recompute_scheduler_enable
except (RuntimeError, AttributeError):
return False
def _compute_potential_max_tokens(vllm_config) -> int:
"""Maximal decode token count, pure arithmetic over config.
The formula lives in exactly one place; it is evaluated once via
set_potential_max_tokens (model runner __init__) and then reused everywhere
via get_potential_max_tokens. Cheap (no select_moe_comm_method).
"""
compilation_config = vllm_config.compilation_config
scheduler_config = vllm_config.scheduler_config
speculative_config = vllm_config.speculative_config
uniform_decode_query_len = 1 if not speculative_config else 1 + speculative_config.num_speculative_tokens
# Use max cudagraph capture size if available, otherwise the maximal uniform
# decode token count.
if compilation_config.cudagraph_capture_sizes:
potential_max_tokens = max(
compilation_config.max_cudagraph_capture_size,
min(
scheduler_config.max_num_batched_tokens,
scheduler_config.max_num_seqs * uniform_decode_query_len,
),
)
if potential_max_tokens != compilation_config.max_cudagraph_capture_size:
logger.warning_once(
"The max_cudagraph_capture_size (%d) is smaller than the potential max tokens required for "
"decode (%d). This may lead to suboptimal performance. Consider adjusting"
"max_cudagraph_capture_size or scheduler_config (max_num_batched_tokens or max_num_seqs)"
"to ensure max_cudagraph_capture_size can accommodate the decode workload. For more details, "
"see the issue #8240(https://github.com/vllm-project/vllm-ascend/issues/8240).",
compilation_config.max_cudagraph_capture_size,
potential_max_tokens,
)
else:
potential_max_tokens = min(scheduler_config.max_num_seqs * uniform_decode_query_len, 512)
return potential_max_tokens
# potential_max_tokens is computed once in the model runner __init__ and reused by
# both the skip-allreduce decision and the o_proj static-exchange buffer sizing, so
# neither path recomputes it.
_potential_max_tokens: int | None = None
def set_potential_max_tokens(vllm_config) -> None:
"""Compute and cache potential_max_tokens once (called from model runner __init__)."""
global _potential_max_tokens
if _potential_max_tokens is not None:
return
_potential_max_tokens = _compute_potential_max_tokens(vllm_config)
def get_potential_max_tokens() -> int:
# Set once in NPUModelRunner.__init__ before any caller reads it.
assert _potential_max_tokens is not None
return _potential_max_tokens
def should_skip_allreduce_across_dp_group(vllm_config, is_draft_model: bool = False) -> bool:
"""Decide whether to skip the all-reduce across the DP group.
Skipping is applicable for all dense models and for moe models only on ranks
that act as KV consumers. We skip the DP all-reduce when either:
- Both the prefill and decode communication methods are MC2 (or FUSED_MC2), or
- Decode requires MC2 and ascend_config.recompute_scheduler_enable is True.
Skipping means each rank may have a different number of tokens, so MC2 needs
a non-zero global_bs and must NOT receive mc2_mask.
Returns False when hierarchy comm is enabled because hierarchy requires
global_bs=0 (uniform tokens), which is incompatible with skipping allreduce.
Recomputed per call (no memoization): potential_max_tokens is a set/get global
computed once in init, and select_moe_comm_method is just config lookups, so
this is cheap and avoids id-reuse / stale-cache / init-ordering hazards.
"""
if is_hierarchical_communication_enabled():
return False
# For dense models, since we don't actually need dp communication, we simply skip it.
# This usually happens when main model is moe while eagle draft model is dense.
is_context_moe_model = is_drafter_moe_model(vllm_config) if is_draft_model else is_moe_model(vllm_config)
if not is_context_moe_model:
return True
# Only applicable to MoE models on KV consumer ranks.
is_kv_consumer = vllm_config.kv_transfer_config is not None and vllm_config.kv_transfer_config.is_kv_consumer
if not is_kv_consumer:
return False
from vllm_ascend.ascend_forward_context import select_moe_comm_method
from vllm_ascend.ops.fused_moe.moe_comm_method import MoECommType
def needs_mc2(n: int) -> bool:
return select_moe_comm_method(n, vllm_config) in {MoECommType.MC2, MoECommType.FUSED_MC2}
scheduler_config = vllm_config.scheduler_config
# potential_max_tokens is read from the set/get global (computed once in init).
decode_must_use_mc2 = needs_mc2(get_potential_max_tokens())
# For prefill, use the scheduler's max_num_batched_tokens for a single batch.
prefill_must_use_mc2 = needs_mc2(scheduler_config.max_num_batched_tokens)
# Skip all-reduce if decode requires MC2 and either prefill also
# requires MC2 or recompute-based scheduler is enabled.
return decode_must_use_mc2 and (prefill_must_use_mc2 or get_ascend_config().recompute_scheduler_enable)
def has_layer_idx(model_instance: torch.nn.Module) -> bool:
if model_instance is None:
return False
global _HAS_LAYER_IDX
if _HAS_LAYER_IDX is None:
_HAS_LAYER_IDX = hasattr(model_instance, "model") and hasattr(model_instance.model, "start_layer")
return _HAS_LAYER_IDX
def flashcomm2_enable() -> bool:
config_val = get_ascend_config().enable_flashcomm2_parallel_size
return config_val > 0
def o_shard_enable() -> bool:
layer_sharding = get_ascend_config().layer_sharding
if layer_sharding is None:
return False
return "o_proj" in layer_sharding
def get_flashcomm2_config_and_validate(ascend_config, vllm_config):
flashcomm2_oproj_tp_size = ascend_config.enable_flashcomm2_parallel_size
global_tp_size = vllm_config.parallel_config.tensor_parallel_size
if ascend_config.enable_flashcomm2_parallel_size <= 0:
return 0
logger.info("Enable FLASHCOMM2 with flashcomm2_oproj_tensor_parallel_size = %s", flashcomm2_oproj_tp_size)
layer_sharding = ascend_config.layer_sharding or []
if layer_sharding:
if layer_sharding == ["o_proj"]:
logger.info_once("Enable FLASHCOMM2 with o_proj layer sharding for reduced memory consumption.")
else:
raise ValueError(
"FLASHCOMM2 only supports 'o_proj' as the sole layer sharding configuration! "
f"Found invalid layer_sharding: {layer_sharding}"
)
if not ascend_config.enable_flashcomm1:
logger.warning_once(
"It is recommended to enable FLASHCOMM1 simultaneously when starting FLASHCOMM2 for optimal performance."
)
if ascend_config.finegrained_tp_config.oproj_tensor_parallel_size > 0:
raise AssertionError(
"flashcomm2_oproj_tensor_parallel_size cannot be enabled simultaneously with oproj_tensor_parallel_size"
)
if global_tp_size <= flashcomm2_oproj_tp_size:
raise AssertionError(
f"flashcomm2_oproj_tensor_parallel_size ({flashcomm2_oproj_tp_size}) cannot exceed "
f"global tensor parallel size ({global_tp_size})"
)
if global_tp_size % flashcomm2_oproj_tp_size != 0:
raise AssertionError(
f"Global tensor parallel size ({global_tp_size}) must be divisible by "
f"flashcomm2_oproj_tensor_parallel_size ({flashcomm2_oproj_tp_size})"
)
if vllm_config.kv_transfer_config is None:
logger.warning_once(
"It is recommended to enable FLASHCOMM2 in P-scenario deployments, enable it in hybrid deployment "
"may lead to decode performance degradation."
)
if vllm_config.kv_transfer_config is not None and vllm_config.kv_transfer_config.is_kv_consumer:
raise AssertionError(
"FLASHCOMM2 primarily targets P-scenario deployments, with additional support "
"for hybrid deployment scenarios. It is not applicable in D-scenario environments."
)
return flashcomm2_oproj_tp_size
def get_flashcomm2_reorgnized_batch_ids(global_tp_size) -> list[list[int]]:
# Reorganize batch_ids so that, after the all2all and reduce-scatter operation,
# each batch_id corresponds to the rank_id within the DP domain.
# For example, when DP = [0, 1, 2, ..., 15] and flashcomm2_oproj_tensor_parallel_size = 2,
# the reorganized batch_ids will be [[batch0, batch8], [batch1, batch9], ..., [batch7, batch15]].
flashcomm2_otp_size = get_ascend_config().flashcomm2_oproj_tensor_parallel_size
num_oproj_tensor_parallel_groups: int = global_tp_size // flashcomm2_otp_size
reorgnized_batch_ids = []
for i in range(num_oproj_tensor_parallel_groups):
ranks = []
for j in range(flashcomm2_otp_size):
rank_idx = i + j * num_oproj_tensor_parallel_groups
ranks.append(rank_idx)
reorgnized_batch_ids.append(ranks)
return reorgnized_batch_ids
def refresh_block_size(vllm_config):
"""
Refresh the block size in cache config.
"""
cache_config = vllm_config.cache_config
scheduler_config = vllm_config.scheduler_config
model_config = vllm_config.model_config
if not cache_config:
return
if cache_config.block_size is None:
cache_config.block_size = 128
if not scheduler_config or not model_config:
return
if model_config.hf_config.model_type == "deepseek_v4":
if cache_config.block_size is None:
cache_config.block_size = 32
elif cache_config.block_size not in [32, 64, 128]:
logger.warning(
"For deepseek_v4 model, block size should be 32, 64 or 128. "
"Setting block size to 32 for better performance."
)
cache_config.block_size = 32
return
if model_config.is_hybrid:
# Hybrid attention+mamba models rely on the model-specific sizing
# logic rather than the generic platform default.
return
if cache_config.block_size != 128:
if cache_config.enable_prefix_caching or scheduler_config.enable_chunked_prefill:
logger.info("Block size is set to 128 if prefix cache or chunked prefill is enabled.")
cache_config.block_size = 128
return
try:
ascend_config = get_ascend_config()
except RuntimeError:
ascend_config = None
if ascend_config is not None and ascend_config.xlite_graph_config.enabled and cache_config.block_size > 128:
logger.warning(
"Setting block size for xlite compatibility. "
"original_block_size=%d, new_block_size=128. "
"xlite_graph_config requires block_size <= 128.",
cache_config.block_size,
)
cache_config.block_size = 128
def dispose_layer(layer: Any):
for attr_name in dir(layer):
attr_value = getattr(layer, attr_name)
if isinstance(attr_value, torch.Tensor):
dispose_tensor(attr_value)
def check_kv_extra_config(vllm_config):
def _check(name: str, config: dict):
tp_key = "tp_size"
dp_key = "dp_size"
if tp_key in config:
config_tp = config[tp_key]
vllm_tp = vllm_config.parallel_config.tensor_parallel_size
if config_tp != vllm_tp:
raise ValueError(
f"KV transfer '{name}' config has a conflicting tensor parallel size. "
f"Expected {vllm_tp}, but got {config_tp}."
)
if dp_key in config:
config_dp = config[dp_key]
vllm_dp = vllm_config.parallel_config.data_parallel_size
if config_dp != vllm_dp:
raise ValueError(
f"KV transfer '{name}' config has a conflicting data parallel size. "
f"Expected {vllm_dp}, but got {config_dp}."
)
if vllm_config.kv_transfer_config.is_kv_producer:
_check("prefill", vllm_config.kv_transfer_config.get_from_extra_config("prefill", {}))
if vllm_config.kv_transfer_config.is_kv_consumer:
_check("decode", vllm_config.kv_transfer_config.get_from_extra_config("decode", {}))
def is_gqa_backend(vllm_config: VllmConfig) -> bool:
model_config = getattr(vllm_config, "model_config", None)
if model_config is None:
return False
if getattr(model_config, "is_deepseek_mla", False) or getattr(model_config, "use_mla", False):
return False
model_arch_config = getattr(model_config, "model_arch_config", None)
total_num_attention_heads = getattr(model_arch_config, "total_num_attention_heads", None)
get_total_num_kv_heads = getattr(model_config, "get_total_num_kv_heads", None)
if total_num_attention_heads is None or not callable(get_total_num_kv_heads):
return False
total_num_kv_heads = get_total_num_kv_heads()
if total_num_kv_heads is None:
return False
return total_num_attention_heads != total_num_kv_heads
def uses_mooncake_connector(kv_transfer_config: Any) -> bool:
mooncake_connector_names = {"MooncakeConnector", "MooncakeConnectorV1"}
return bool(_collect_kv_connector_names(kv_transfer_config) & mooncake_connector_names)
def _collect_kv_connector_names(value: Any) -> set[str]:
connector_names: set[str] = set()
if isinstance(value, dict):
connector = value.get("kv_connector")
if isinstance(connector, str):
connector_names.add(connector)
for nested_value in value.values():
connector_names.update(_collect_kv_connector_names(nested_value))
elif isinstance(value, (list, tuple)):
for nested_value in value:
connector_names.update(_collect_kv_connector_names(nested_value))
else:
connector = getattr(value, "kv_connector", None)
if isinstance(connector, str):
connector_names.add(connector)
extra_config = getattr(value, "kv_connector_extra_config", None)
if isinstance(extra_config, (dict, list, tuple)):
connector_names.update(_collect_kv_connector_names(extra_config))
return connector_names
def singleton(cls):
instances = {}
def get_instance(*args, **kwargs):
if cls not in instances:
instances[cls] = cls(*args, **kwargs)
return instances[cls]
return get_instance
@lru_cache(maxsize=1)
def enable_dsa_cp() -> bool:
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
# DSA CP is only applicable to models with indexer (e.g., DSv3.2, DSv4).
has_indexer = hasattr(vllm_config.model_config, "hf_text_config") and hasattr(
vllm_config.model_config.hf_text_config, "index_topk"
)
if not has_indexer:
return False
dsa_cp_enable = False
additional_config = getattr(vllm_config, "additional_config", None)
if additional_config is not None and "enable_dsa_cp" in additional_config:
dsa_cp_enable = bool(additional_config["enable_dsa_cp"])
if dsa_cp_enable and not enable_sp():
raise ValueError(
"DSA CP requires SP to be enabled. Please enable SP(set VLLM_ASCEND_ENABLE_FLASHCOMM1=1) to use DSA CP."
)
return dsa_cp_enable and enable_sp()
@lru_cache(maxsize=1)
def enable_dsa_cp_with_layer_shard() -> bool:
if not enable_dsa_cp():
return False
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
kv_transfer_config = vllm_config.kv_transfer_config
# Layer sharding broadcast only pays off when it can be hidden by the
# heavier prefill-stage compute, so enable it only on the P-side instance.
is_prefill_instance = kv_transfer_config is not None and kv_transfer_config.kv_role == "kv_producer"
return is_prefill_instance
@lru_cache(maxsize=1)
def enable_dsa_cp_with_o_proj_tp() -> bool:
if not enable_dsa_cp():
return False
from vllm.config import get_current_vllm_config
vllm_config = get_current_vllm_config()
kv_transfer_config = vllm_config.kv_transfer_config
# In PD-mixed mode, keep the original TP o_proj weight when:
# 1) KV pooling is disabled, or
# 2) KV pooling is enabled with kv_role == "kv_both".
return kv_transfer_config is None or kv_transfer_config.kv_role == "kv_both"
def check_gdn_layer(vllm_config) -> bool:
"""
gdn layer is marked with `linear_attention`.
So, if `linear_attention` is detected, we think the model has gdn-attention.
"""
if not hasattr(vllm_config, "model_config"):
return False
model_config = vllm_config.model_config
if not hasattr(model_config, "hf_config"):
return False
hf_config = model_config.hf_config
# Use `or []` to prevent errors when layer_types is None
layer_types = getattr(hf_config, "layer_types", None) or []
if "linear_attention" in layer_types:
return True
text_config = getattr(hf_config, "text_config", None)
if text_config:
text_layer_types = getattr(text_config, "layer_types", None) or []
if "linear_attention" in text_layer_types:
return True
return False
def get_rope_dim(vllm_config):
model_config = vllm_config.model_config
if model_config.use_mla:
rope_dim = model_config.hf_text_config.qk_rope_head_dim
else:
rope_dim = model_config.get_head_size()
# For models using partial rope like Qwen3-Next.
if hasattr(model_config.hf_text_config, "partial_rotary_factor"):
rope_dim = int(rope_dim * model_config.hf_text_config.partial_rotary_factor)
elif hasattr(model_config.hf_text_config, "rotary_dim"):
rope_dim = int(model_config.hf_text_config.rotary_dim)
return rope_dim
def calc_split_factor(num_list: list[int]):
total = sum(num_list)
return [total / num for num in num_list]
# NOTE: The last two dimensions of ND are transferred to NZ
def trans_nd_to_nz(cache_tensor: torch.Tensor):
assert len(cache_tensor.shape) >= 2
batch = cache_tensor.shape[:-2]
a, b = cache_tensor.shape[-2], cache_tensor.shape[-1]
dtype = cache_tensor.dtype
if dtype == torch.int8:
a0, b0 = 16, 32
else:
a0, b0 = 16, 16
nz_shape = list(batch) + [math.ceil(b / b0), math.ceil(a / a0), a0, b0]
# Generate the axis order for the transpose operation.
offset = len(cache_tensor.shape) - 2
base = [2, 0, 1, 3]
array_trans = [i for i in range(offset)] + [i + offset for i in base]
# Perform shape transformation and transpose operation.
*_, n1, m1, m0, n0 = nz_shape
cache_tensor = cache_tensor.reshape(nz_shape[:-4] + [m1, m0, n1, n0])
cache_tensor = cache_tensor.permute(*array_trans)
return cache_tensor
def parse_layer_idx(prefix: str) -> int | None:
"""Extract the layer index from a module prefix string like 'model.layers.0.self_attn'."""
match = re.search(r"layers\.(\d+)", prefix)
return int(match.group(1)) if match else None
def get_compressed_pos_and_indices(
num_computed_tokens: np.ndarray,
num_scheduled_tokens: np.ndarray,
arrange_np: np.ndarray,
use_compress: bool,
kv_cache_groups,
) -> tuple[list[np.ndarray], list[np.ndarray], list[np.ndarray]]:
"""
Batch generate compressed position ids for multi-requests on DSv4.
Calculate compressed position ids independently for each single request.
Args:
num_computed_tokens: Historical processed token counts of multiple requests, shape=[num_reqs,]
num_scheduled_tokens: New scheduled token counts of multiple requests in current step, shape=[num_reqs,]
Returns:
tuple(np.ndarray, np.ndarray):
1. Flattened compressed position id array for all requests
2. Length of compressed position ids for each individual request
"""
if not use_compress:
return None, None, None # type: ignore[return-value]
# Assert input validity
assert num_computed_tokens.shape == num_scheduled_tokens.shape, (
"num_computed_tokens and num_scheduled_tokens must have the same shape"
)
assert np.all(num_computed_tokens >= 0) and np.all(num_scheduled_tokens >= 0), (
"Token count cannot be negative value"
)
positions_compressed_list = []
req_indices_compressed_list = []
num_scheduled_tokens_compressed_list = []
from vllm.v1.kv_cache_interface import UniformTypeKVCacheSpecs
for kv_cache_group_id, kv_cache_group_spec in enumerate(kv_cache_groups):
# Calculate compressed length of historical & total tokens
if isinstance(kv_cache_group_spec.kv_cache_spec, UniformTypeKVCacheSpecs):
kv_cache_spec = next(iter(kv_cache_group_spec.kv_cache_spec.kv_cache_specs.values()))
else:
kv_cache_spec = kv_cache_group_spec.kv_cache_spec
compress_ratio = getattr(kv_cache_spec, "compress_ratio", 1)
# Note(qcs): some models use compress_ratio=0 as non-compression tag.
if compress_ratio > 1:
compressed_historical_len = num_computed_tokens // compress_ratio
compressed_total_len = (num_computed_tokens + num_scheduled_tokens) // compress_ratio
else:
compressed_historical_len = num_computed_tokens
compressed_total_len = num_computed_tokens + num_scheduled_tokens
# The number of new compressed position ids for each request
num_new_compressed_pos = compressed_total_len - compressed_historical_len
# Core vectorized calculation (no for-loop)
pos_starts = compressed_historical_len
prefix_offsets = np.concatenate([[0], np.cumsum(num_new_compressed_pos[:-1])])
compressed_pos_ids = np.arange(np.sum(num_new_compressed_pos)) + np.repeat(
pos_starts - prefix_offsets, num_new_compressed_pos
)
req_indices_compressed = np.repeat(arrange_np, num_new_compressed_pos)
req_indices_compressed_list.append(req_indices_compressed)
positions_compressed_list.append(compressed_pos_ids)
num_scheduled_tokens_compressed_list.append(num_new_compressed_pos)
return positions_compressed_list, req_indices_compressed_list, num_scheduled_tokens_compressed_list
def kv_cache_spec_uses_sparse_sfa_c8(kv_cache_spec) -> bool:
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
return isinstance(kv_cache_spec, AscendMLAAttentionSpec) and bool(
getattr(kv_cache_spec, "cache_sparse_sfa_c8", False)
)
def kv_cache_spec_uses_sparse_li_c8(kv_cache_spec) -> bool:
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
return isinstance(kv_cache_spec, AscendMLAAttentionSpec) and bool(
getattr(kv_cache_spec, "cache_sparse_li_c8", False)
)
def sparse_kv_cache_has_indexer(kv_cache_spec: AttentionSpec) -> bool:
sparse_head_dim = getattr(kv_cache_spec, "sparse_head_dim", None)
return sparse_head_dim is not None and len(sparse_head_dim) == 3 and sparse_head_dim[2] > 0
def is_hidden_state_cache_spec(spec) -> bool:
"""Whether ``spec`` marks an ``extract_hidden_states`` cache-only layer."""
from vllm.v1.kv_cache_interface import HiddenStateCacheSpec
return isinstance(spec, HiddenStateCacheSpec)
@lru_cache(maxsize=1)
def _libc_getenv():
import ctypes
libc = ctypes.CDLL(None)
libc.getenv.argtypes = [ctypes.c_char_p]
libc.getenv.restype = ctypes.c_char_p
return libc.getenv
def get_c_env(name: str, encoding: str = "utf-8") -> str | None:
"""Read env via C getenv; returns None if unset."""
raw = _libc_getenv()(name.encode(encoding))
if raw is None:
return None
return raw.decode(encoding)