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