init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -15,5 +15,75 @@
# limitations under the License.
#
from vllm_ascend.patch.worker import patch_common # noqa: F401
from vllm_ascend.patch.worker import patch_main # noqa: F401
from vllm.triton_utils import HAS_TRITON
from vllm_ascend.utils import is_310p, vllm_version_is
# The v2 model runner is intentionally NOT made compatible with the v0.23.0
# release. vLLM v0.23.0 and the verified main commit are diverged, and the v2
# worker patches target main-only APIs; rather than maintain a separate v0.23.0
# compatibility path we keep v2 main-only. With v0.23.0 installed this flag is
# False, so none of the patch_v2.* / routed-experts-capture patches below are
# imported and the v2 worker stays dormant (the release uses the v1 runner).
if vllm_version_is("0.23.0"):
_V2_MODEL_RUNNER_SUPPORTED = False
else:
_V2_MODEL_RUNNER_SUPPORTED = True
if HAS_TRITON:
import vllm_ascend.patch.worker.patch_triton
if _V2_MODEL_RUNNER_SUPPORTED:
import vllm_ascend.patch.worker.patch_v2.patch_triton # noqa
import vllm_ascend.patch.worker.patch_process_weights_after_loading # noqa
import vllm_ascend.patch.worker.patch_weight_utils # noqa
import vllm_ascend.patch.worker.patch_distributed # noqa
import vllm_ascend.patch.worker.patch_minimax_m2 # noqa
import vllm_ascend.patch.worker.patch_minimax_m2_linear_attn # noqa
import vllm_ascend.patch.worker.patch_mamba_utils # noqa
import vllm_ascend.patch.worker.patch_qwen3_next_mtp # noqa
if not is_310p():
import vllm_ascend.patch.worker.patch_qwen3_5 # noqa
import vllm_ascend.patch.worker.patch_qwen3_dflash # noqa
import vllm_ascend.patch.worker.patch_qwen3vl # noqa
else:
import vllm_ascend.patch.worker.patch_idex_310 # noqa
import vllm_ascend.patch.worker.patch_rejection_sampler # noqa
# torchair/npugraph_ex is only available on NPU; silently skip when missing
# so that CPU-only environments (e.g. UT runners without torch_npu) can still
# import this module without crashing.
try: # noqa: SIM105
import vllm_ascend.patch.worker.patch_npugraph_ex_triton # noqa
except ImportError:
pass
import vllm_ascend.patch.worker.patch_kimi_k25 # noqa
import vllm_ascend.patch.worker.patch_draft_quarot # noqa
import vllm_ascend.patch.worker.patch_eagle3_init # noqa
import vllm_ascend.patch.worker.patch_cudagraph # noqa
import vllm_ascend.patch.worker.patch_deepseek_mtp # noqa
import vllm_ascend.patch.worker.patch_deepseek_v2 # noqa
import vllm_ascend.patch.worker.patch_gqa_c8 # noqa
# vLLM's use_v2_model_runner may enable the v2 runner without the
# VLLM_USE_V2_MODEL_RUNNER env var (e.g. based on model architecture).
# We always patch it so that on Ascend the v2 runner is enabled only
# when the env var is explicitly set.
import vllm_ascend.patch.worker.patch_v2.patch_use_v2_model_runner # noqa
if not vllm_version_is("0.23.0"):
import vllm_ascend.patch.worker.patch_fused_moe # noqa
if _V2_MODEL_RUNNER_SUPPORTED:
import vllm_ascend.patch.worker.patch_v2.patch_uva # noqa
import vllm_ascend.patch.worker.patch_v2.patch_input_batch # noqa
import vllm_ascend.patch.worker.patch_v2.patch_model_state # noqa
import vllm_ascend.patch.worker.patch_v2.patch_block_table # noqa
import vllm_ascend.patch.worker.patch_v2.patch_attn_utils # noqa
# only patch routed experts capture in main2main.
if _V2_MODEL_RUNNER_SUPPORTED:
import vllm_ascend.patch.worker.patch_routed_experts_capture # noqa

View File

@@ -0,0 +1,295 @@
#
# 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.
#
from __future__ import annotations
import logging
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from datetime import timedelta
from threading import Lock
from typing import cast
logger = logging.getLogger(__name__)
_AUDITED_PG_OPTION_FIELDS = ("hccl_config",)
# These fields are populated by torch_npu/new_group at runtime and are either
# already represented elsewhere in the reuse key or intentionally excluded.
_REDUNDANT_PG_OPTION_FIELDS = (
"global_ranks_in_group",
"group_id",
"group_name",
)
_KNOWN_PG_OPTION_DEFAULTS = {
"backend": "hccl",
"global_ranks_in_group": (),
"group_id": "",
"group_name": "",
"hccl_config": {},
"is_high_priority_stream": False,
"op_timeout": timedelta(seconds=10),
}
_OPTION_DEFAULT_NON_AUDITED = (None, False, 0, 0.0)
_NON_GROUP_MEMBER = object()
_NON_GROUP_MEMBER_SET = False
@dataclass(frozen=True)
class HcclPgKey:
backend: str
ranks: tuple[int, ...]
options_key: tuple[tuple[str, object], ...]
reuse_domain: str
@dataclass
class RegistryEntry:
handle: object
refcount: int
def make_hccl_pg_key(
ranks: list[int] | tuple[int, ...],
backend: str,
pg_options: object,
reuse_domain: str,
) -> HcclPgKey | None:
"""
Return a hashable key that identifies a shared HCCL process group.
Unknown non-default pg option fields cause fail-closed behavior (returns None),
which disables process-group reuse for this configuration.
"""
if backend != "hccl":
return None
normalized_options = _normalize_hccl_pg_options(pg_options)
if normalized_options is None:
return None
if not _global_ranks_match_requested_ranks(ranks, pg_options):
return None
return HcclPgKey(
backend=backend,
ranks=tuple(ranks),
options_key=normalized_options,
reuse_domain=reuse_domain,
)
class HcclPgRegistry:
"""
HCCL process-group reuse registry.
Cross-key process-group creation is intentionally not a full concurrent factory:
callers still need to serialize creation by design, and this helper keeps lock
scope to registry lookup/refcount mutation only.
"""
def __init__(self):
self._entries: dict[HcclPgKey, RegistryEntry] = {}
self._registry_lock = Lock()
def acquire(
self,
*,
ranks,
backend,
pg_options,
reuse_domain,
create_fn,
) -> object:
key = make_hccl_pg_key(ranks, backend, pg_options, reuse_domain)
if key is None:
return create_fn()
with self._registry_lock:
entry = self._entries.get(key)
if entry is not None:
entry.refcount += 1
return entry.handle
handle = create_fn()
with self._registry_lock:
existing = self._entries.get(key)
if existing is None:
self._entries[key] = RegistryEntry(handle=handle, refcount=1)
return handle
existing.refcount += 1
if not _is_non_group_member(handle):
_destroy_process_group(handle)
return existing.handle
def release(self, key: HcclPgKey) -> object | None:
with self._registry_lock:
entry = self._entries.get(key)
if entry is None:
return None
if entry.refcount > 1:
entry.refcount -= 1
return None
del self._entries[key]
if _is_non_group_member(entry.handle):
return None
_destroy_process_group(entry.handle)
return entry.handle
def clear(self):
with self._registry_lock:
self._entries.clear()
# Full reinitialization path already destroys process groups; clear
# only removes stale registry metadata.
def _normalize_hccl_pg_options(
pg_options: object,
) -> tuple[tuple[str, object], ...] | None:
if pg_options is None:
return ()
options_dict = dict(pg_options) if isinstance(pg_options, Mapping) else None
if _has_unknown_non_default_fields(pg_options):
return None
normalized_items: list[tuple[str, object]] = []
for field_name in _AUDITED_PG_OPTION_FIELDS:
default_value = _KNOWN_PG_OPTION_DEFAULTS[field_name]
if options_dict is not None:
actual_value = options_dict.get(field_name, default_value)
else:
actual_value = getattr(pg_options, field_name, default_value)
if _is_default_option_value(field_name, actual_value):
continue
normalized_items.append((field_name, _freeze_for_key(actual_value)))
return tuple(sorted(normalized_items))
def _has_unknown_non_default_fields(pg_options: object) -> bool:
options_dict = None
if isinstance(pg_options, Mapping):
options_dict = dict(pg_options)
else:
options_dict = vars(pg_options) if hasattr(pg_options, "__dict__") else None
if options_dict is not None:
field_names: list[str] = list(options_dict.keys())
else:
field_names = [name for name in dir(pg_options) if not name.startswith("_")]
for name in field_names:
if name in _AUDITED_PG_OPTION_FIELDS:
continue
if name in _REDUNDANT_PG_OPTION_FIELDS:
continue
try:
if options_dict is not None:
value = options_dict[name]
else:
value = getattr(pg_options, name)
except Exception:
continue
if callable(value):
continue
if _is_default_option_value(name, value):
continue
logger.warning(
"Disabling HCCL process-group reuse because pg_options has non-default field '%s'",
name,
)
return True
return False
def _global_ranks_match_requested_ranks(
ranks: list[int] | tuple[int, ...],
pg_options: object,
) -> bool:
if isinstance(pg_options, Mapping):
value = pg_options.get("global_ranks_in_group", ())
else:
value = getattr(pg_options, "global_ranks_in_group", ())
if value is None:
return True
value_tuple = tuple(value)
if not value_tuple:
return True
ranks_tuple = tuple(ranks)
if value_tuple == ranks_tuple:
return True
logger.warning(
"Disabling HCCL process-group reuse because pg_options.global_ranks_in_group=%s "
"does not match requested ranks=%s",
value_tuple,
ranks_tuple,
)
return False
def _freeze_for_key(value: object) -> object:
if isinstance(value, dict):
return tuple(
(str(key), _freeze_for_key(val)) for key, val in sorted(value.items(), key=lambda item: str(item[0]))
)
if isinstance(value, (list, tuple)):
return tuple(_freeze_for_key(item) for item in value)
if isinstance(value, set):
return tuple(_freeze_for_key(item) for item in sorted(value, key=lambda item: str(item)))
return value
def _is_default_option_value(name: str, value: object) -> bool:
if name in _KNOWN_PG_OPTION_DEFAULTS:
default_value = _KNOWN_PG_OPTION_DEFAULTS[name]
if name in ("global_ranks_in_group",):
default_ranks = cast(tuple[object, ...], default_value)
if isinstance(value, Iterable) and not isinstance(value, (str, bytes, dict)):
return tuple(value) == default_ranks
return value == default_value
if name == "hccl_config":
return value in (None, {}, default_value)
return value == default_value
if name in ("_rank", "_backend"):
return True
return value in _OPTION_DEFAULT_NON_AUDITED
def _is_non_group_member(handle: object) -> bool:
global _NON_GROUP_MEMBER
global _NON_GROUP_MEMBER_SET
if not _NON_GROUP_MEMBER_SET:
_NON_GROUP_MEMBER = _load_non_group_member_sentinel()
_NON_GROUP_MEMBER_SET = True
return handle is _NON_GROUP_MEMBER
def _load_non_group_member_sentinel() -> object:
try:
from torch.distributed.distributed_c10d import GroupMember
return GroupMember.NON_GROUP_MEMBER
except Exception:
return object()
def _destroy_process_group(handle: object):
from torch.distributed import destroy_process_group
destroy_process_group(handle)

View File

@@ -0,0 +1,38 @@
from vllm.config import CUDAGraphMode
from vllm.forward_context import BatchDescriptor
from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher
def _create_padded_batch_descriptor(
self,
num_tokens: int,
uniform_decode: bool,
has_lora: bool,
num_active_loras: int = 0,
) -> BatchDescriptor:
max_num_seqs = self.vllm_config.scheduler_config.max_num_seqs
uniform_decode_query_len = self.uniform_decode_query_len
num_tokens_padded = self._bs_to_padded_graph_size[num_tokens]
# FULL mode should not be treated as uniform decode
if (
uniform_decode
and self.cudagraph_mode.has_mode(CUDAGraphMode.FULL)
and self.cudagraph_mode != CUDAGraphMode.FULL
):
num_reqs = min(num_tokens_padded // uniform_decode_query_len, max_num_seqs)
assert num_tokens_padded % uniform_decode_query_len == 0
else:
uniform_decode = False
num_reqs = min(num_tokens_padded, max_num_seqs)
return BatchDescriptor(
num_tokens=num_tokens_padded,
num_reqs=num_reqs,
uniform=uniform_decode,
has_lora=has_lora,
num_active_loras=num_active_loras,
)
CudagraphDispatcher._create_padded_batch_descriptor = _create_padded_batch_descriptor

View File

@@ -0,0 +1,80 @@
import torch
import torch.nn as nn
import vllm
from transformers import DeepseekV2Config, DeepseekV3Config
from vllm.config import VllmConfig
from vllm.model_executor.models.deepseek_mtp import DeepSeekMTP, DeepSeekMultiTokenPredictorLayer
from vllm.model_executor.models.deepseek_v2 import GlmMoeDsaForCausalLM
from vllm.model_executor.models.utils import AutoWeightsLoader
MTP_ROT_WEIGHT_NAME = "rot.weight"
def get_spec_layer_idx_from_weight_name(config: DeepseekV2Config | DeepseekV3Config, weight_name: str) -> int | None:
if hasattr(config, "num_nextn_predict_layers") and config.num_nextn_predict_layers > 0:
layer_idx = config.num_hidden_layers
for i in range(config.num_nextn_predict_layers):
if (
weight_name.startswith(f"model.layers.{layer_idx + i}.")
or weight_name.startswith(MTP_ROT_WEIGHT_NAME)
or weight_name.startswith(f"layers.{layer_idx + i}.")
):
return layer_idx + i
return None
class AscendDeepSeekMultiTokenPredictorLayer(DeepSeekMultiTokenPredictorLayer):
def __init__(self, vllm_config: VllmConfig, prefix: str) -> None:
super().__init__(vllm_config, prefix)
quant_description = getattr(vllm_config.quant_config, "quant_description", None)
self.is_rot_used = quant_description.get("is_rot_used", False) if quant_description is not None else False
self.target_model_type = vllm_config.speculative_config.target_model_config.hf_text_config.model_type
if self.is_rot_used and self.target_model_type == "glm_moe_dsa":
self.rot = nn.Linear(self.config.hidden_size, self.config.hidden_size, bias=False)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
previous_hidden_states: torch.Tensor,
inputs_embeds: torch.Tensor | None = None,
spec_step_index: int = 0,
) -> torch.Tensor:
assert inputs_embeds is not None
# masking inputs at position 0, as not needed by MTP
inputs_embeds = torch.where(positions.unsqueeze(-1) == 0, 0, inputs_embeds)
inputs_embeds = self.enorm(inputs_embeds)
if self.is_rot_used and self.target_model_type == "glm_moe_dsa":
previous_hidden_states = self.rot(previous_hidden_states)
previous_hidden_states = self.hnorm(previous_hidden_states)
hidden_states = self.eh_proj(torch.cat([inputs_embeds, previous_hidden_states], dim=-1))
hidden_states, residual = self.mtp_block(positions=positions, hidden_states=hidden_states, residual=None)
hidden_states = residual + hidden_states # pre-final-norm (logits hidden)
# Recycle the post-final-norm hidden into the next draft step.
# compute_logits applies shared_head (== final norm) to the pre-norm
# element, so logits and the recycle each get exactly one final-norm.
# Matches SGLang's deepseek_nextn.
return hidden_states, self.shared_head(hidden_states)
class AscendDeepSeekMTP(DeepSeekMTP):
def _rewrite_spec_layer_name(self, spec_layer: int, name: str) -> str:
if name != MTP_ROT_WEIGHT_NAME:
return super()._rewrite_spec_layer_name(spec_layer, name)
else:
return f"model.layers.{spec_layer}.rot.weight"
class AscendGlmMoeDsaForCausalLM(GlmMoeDsaForCausalLM):
def load_weights(self, weights):
loader = AutoWeightsLoader(self, skip_prefixes=[MTP_ROT_WEIGHT_NAME])
return loader.load_weights(weights)
vllm.model_executor.models.deepseek_v2.get_spec_layer_idx_from_weight_name = get_spec_layer_idx_from_weight_name
vllm.model_executor.models.deepseek_mtp.get_spec_layer_idx_from_weight_name = get_spec_layer_idx_from_weight_name
vllm.model_executor.models.deepseek_mtp.DeepSeekMultiTokenPredictorLayer = AscendDeepSeekMultiTokenPredictorLayer
vllm.model_executor.models.deepseek_mtp.DeepSeekMTP = AscendDeepSeekMTP
vllm.model_executor.models.deepseek_v2.GlmMoeDsaForCausalLM = AscendGlmMoeDsaForCausalLM

View File

@@ -0,0 +1,279 @@
import torch
from torch import nn
from transformers import DeepseekV2Config, DeepseekV3Config
from vllm.config import CacheConfig, VllmConfig
from vllm.distributed import get_tensor_model_parallel_world_size
from vllm.model_executor.layers.layernorm import RMSNorm
from vllm.model_executor.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
)
from vllm.model_executor.layers.mla import (
MLAModules,
MultiHeadLatentAttentionWrapper,
)
from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.model_executor.layers.rotary_embedding import get_rope
from vllm.model_executor.models.deepseek_v2 import (
DeepSeekV2FusedQkvAProjLinear,
DeepseekV2MLAAttention,
Indexer,
yarn_get_mscale,
)
from vllm.model_executor.models.utils import extract_layer_index
def _should_skip_indexer_init(
config: DeepseekV2Config | DeepseekV3Config,
prefix: str,
skip_topk: bool,
) -> bool:
if not skip_topk:
return False
layer_id = extract_layer_index(prefix)
num_hidden_layers = getattr(config, "num_hidden_layers", None)
if num_hidden_layers is not None and layer_id >= num_hidden_layers:
return False
# GLM-5.2 describes checkpoint-level shared indexers explicitly. Runtime
# IndexCache overrides on GLM-5.1 only skip top-k computation; its
# checkpoint still contains an Indexer for every layer.
indexer_types = getattr(config, "indexer_types", None)
indexer_type = indexer_types[layer_id] if indexer_types is not None and layer_id < len(indexer_types) else None
return isinstance(indexer_type, str) and indexer_type.lower() == "shared"
def _deepseek_v2_mla_attention_init(
self,
vllm_config: VllmConfig,
config: DeepseekV2Config | DeepseekV3Config,
hidden_size: int,
num_heads: int,
qk_nope_head_dim: int,
qk_rope_head_dim: int,
v_head_dim: int,
q_lora_rank: int | None,
kv_lora_rank: int,
max_position_embeddings: int = 8192,
cache_config: CacheConfig | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
topk_indices_buffer: torch.Tensor | None = None,
input_size: int | None = None,
) -> None:
# 这里不能使用 super().__init__(),因为当前函数定义在原类之外,
# 最后通过赋值的方式替换 DeepseekV2MLAAttention.__init__。
nn.Module.__init__(self)
self.hidden_size = hidden_size
self.qk_nope_head_dim = qk_nope_head_dim
self.qk_rope_head_dim = qk_rope_head_dim
self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
self.v_head_dim = v_head_dim
self.q_lora_rank = q_lora_rank
self.kv_lora_rank = kv_lora_rank
self.num_heads = num_heads
tp_size = get_tensor_model_parallel_world_size()
assert num_heads % tp_size == 0
self.num_local_heads = num_heads // tp_size
self.scaling = self.qk_head_dim**-0.5
self.max_position_embeddings = max_position_embeddings
# Use input_size for projection input dimensions if provided,
# otherwise default to hidden_size (used in Eagle3 Deepseek with MLA).
proj_input_size = input_size if input_size is not None else self.hidden_size
if self.q_lora_rank is not None:
self.fused_qkv_a_proj = DeepSeekV2FusedQkvAProjLinear(
proj_input_size,
[
self.q_lora_rank,
self.kv_lora_rank + self.qk_rope_head_dim,
],
quant_config=quant_config,
prefix=f"{prefix}.fused_qkv_a_proj",
)
else:
self.kv_a_proj_with_mqa = ReplicatedLinear(
proj_input_size,
self.kv_lora_rank + self.qk_rope_head_dim,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.kv_a_proj_with_mqa",
)
if self.q_lora_rank is not None:
self.q_a_layernorm = RMSNorm(
self.q_lora_rank,
eps=config.rms_norm_eps,
)
self.q_b_proj = ColumnParallelLinear(
self.q_lora_rank,
self.num_heads * self.qk_head_dim,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.q_b_proj",
)
else:
self.q_proj = ColumnParallelLinear(
proj_input_size,
self.num_heads * self.qk_head_dim,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.q_proj",
)
self.kv_a_layernorm = RMSNorm(
self.kv_lora_rank,
eps=config.rms_norm_eps,
)
self.kv_b_proj = ColumnParallelLinear(
self.kv_lora_rank,
self.num_heads * (self.qk_nope_head_dim + self.v_head_dim),
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.kv_b_proj",
)
self.o_proj = RowParallelLinear(
self.num_heads * self.v_head_dim,
self.hidden_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.o_proj",
)
if config.rope_parameters["rope_type"] != "default":
config.rope_parameters["rope_type"] = (
"deepseek_yarn"
if config.rope_parameters.get(
"apply_yarn_scaling",
True,
)
else "deepseek_llama_scaling"
)
self.rotary_emb = get_rope(
qk_rope_head_dim,
max_position=max_position_embeddings,
rope_parameters=config.rope_parameters,
is_neox_style=False,
)
if config.rope_parameters["rope_type"] != "default" and config.rope_parameters["rope_type"] == "deepseek_yarn":
mscale_all_dim = config.rope_parameters.get(
"mscale_all_dim",
False,
)
scaling_factor = config.rope_parameters["factor"]
mscale = yarn_get_mscale(
scaling_factor,
float(mscale_all_dim),
)
self.scaling = self.scaling * mscale * mscale
self.is_v32 = hasattr(config, "index_topk")
# IndexCache config.
#
# skip_topk controls top-k reuse. Indexer initialization is skipped only
# when the checkpoint marks this layer as sharing another layer's Indexer.
_skip_topk = False
_index_topk_freq = getattr(
config,
"index_topk_freq",
1,
)
_index_topk_pattern = getattr(
config,
"index_topk_pattern",
None,
)
_index_skip_topk_offset = getattr(
config,
"index_skip_topk_offset",
2,
)
layer_id = extract_layer_index(prefix)
if _index_topk_pattern is None:
_skip_topk = (
max(
layer_id - _index_skip_topk_offset + 1,
0,
)
% _index_topk_freq
!= 0
)
elif 0 <= layer_id < len(_index_topk_pattern):
_skip_topk = _index_topk_pattern[layer_id] == "S"
skip_indexer_init = _should_skip_indexer_init(config, prefix, _skip_topk)
if self.is_v32 and not skip_indexer_init:
self.indexer_rope_emb = get_rope(
qk_rope_head_dim,
max_position=max_position_embeddings,
rope_parameters=config.rope_parameters,
is_neox_style=not getattr(
config,
"indexer_rope_interleave",
False,
),
)
self.indexer = Indexer(
vllm_config,
config,
hidden_size,
q_lora_rank,
quant_config,
cache_config,
topk_indices_buffer,
f"{prefix}.indexer",
is_inplace_rope=self.indexer_rope_emb.enabled(),
)
else:
self.indexer_rope_emb = None
self.indexer = None
mla_modules = MLAModules(
kv_a_layernorm=self.kv_a_layernorm,
kv_b_proj=self.kv_b_proj,
rotary_emb=self.rotary_emb,
o_proj=self.o_proj,
fused_qkv_a_proj=(self.fused_qkv_a_proj if self.q_lora_rank is not None else None),
kv_a_proj_with_mqa=(self.kv_a_proj_with_mqa if self.q_lora_rank is None else None),
q_a_layernorm=(self.q_a_layernorm if self.q_lora_rank is not None else None),
q_b_proj=(self.q_b_proj if self.q_lora_rank is not None else None),
q_proj=(self.q_proj if self.q_lora_rank is None else None),
indexer=self.indexer,
indexer_rotary_emb=self.indexer_rope_emb,
is_sparse=self.is_v32,
topk_indices_buffer=topk_indices_buffer,
)
self.mla_attn = MultiHeadLatentAttentionWrapper(
self.hidden_size,
self.num_local_heads,
self.scaling,
self.qk_nope_head_dim,
self.qk_rope_head_dim,
self.v_head_dim,
self.q_lora_rank,
self.kv_lora_rank,
mla_modules,
cache_config,
quant_config,
prefix,
skip_topk=_skip_topk,
)
DeepseekV2MLAAttention.__init__ = _deepseek_v2_mla_attention_init

View File

@@ -0,0 +1,269 @@
#
# 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.
#
from __future__ import annotations
import logging
from functools import wraps
from typing import Any, cast
import torch
import vllm
from torch.distributed import Backend
from vllm.distributed.parallel_state import GroupCoordinator, _get_unique_name, _register_group
from vllm_ascend.distributed.device_communicators.npu_communicator import NPUCommunicator
from vllm_ascend.patch.worker._hccl_pg_registry import HcclPgKey, HcclPgRegistry, make_hccl_pg_key
from vllm_ascend.utils import create_hccl_pg_options
_HCCL_PG_REGISTRY = HcclPgRegistry()
logger = logging.getLogger(__name__)
def _normalize_backend(backend: str | Backend) -> str:
return str(backend)
def _resolve_reuse_domain(group_name: str) -> str:
group_base_name = group_name.split(":")[0]
if "eplb" in group_base_name or group_base_name == "mc2":
return group_base_name
return "shared"
def _create_device_group(
ranks: list[int],
backend: str,
hccl_pg_options: object,
):
return torch.distributed.new_group(
ranks,
backend=backend,
pg_options=hccl_pg_options,
)
def _acquire_hccl_group(
*,
ranks: list[int],
backend: str,
hccl_pg_options: object,
reuse_domain: str,
):
# Coordinator construction must remain process-serial and globally ordered:
# new_group is collective, and the registry only deduplicates equivalent
# HCCL groups within that ordering contract. It is not a concurrent PG factory.
hccl_key = make_hccl_pg_key(ranks, backend, hccl_pg_options, reuse_domain)
device_group = _HCCL_PG_REGISTRY.acquire(
ranks=ranks,
backend=backend,
pg_options=hccl_pg_options,
reuse_domain=reuse_domain,
create_fn=lambda: _create_device_group(ranks, backend, hccl_pg_options),
)
return device_group, hccl_key
def _wrap_destroy_distributed_environment(destroy_fn):
if getattr(cast(Any, destroy_fn), "_hccl_registry_clearing_wrapped", False) is True:
return destroy_fn
@wraps(destroy_fn)
def wrapped(*args, **kwargs):
try:
return destroy_fn(*args, **kwargs)
finally:
_HCCL_PG_REGISTRY.clear()
cast(Any, wrapped)._hccl_registry_clearing_wrapped = True
return wrapped
def _patch_destroy_distributed_environment():
destroy_fn = _wrap_destroy_distributed_environment(vllm.distributed.parallel_state.destroy_distributed_environment)
vllm.distributed.parallel_state.destroy_distributed_environment = destroy_fn
vllm.distributed.destroy_distributed_environment = destroy_fn
class GroupCoordinatorPatch(GroupCoordinator):
def __init__(
self,
group_ranks: list[list[int]],
local_rank: int,
torch_distributed_backend: str | Backend,
use_device_communicator: bool, # whether to use device communicator
use_message_queue_broadcaster: bool = False,
group_name: str | None = None,
):
group_name = group_name or "anonymous"
self.unique_name = _get_unique_name(group_name)
_register_group(self)
self.rank = torch.distributed.get_rank()
self.local_rank = local_rank
self.backend = _normalize_backend(torch_distributed_backend)
self._acquired_hccl_keys: list[HcclPgKey] = []
self._unshared_hccl_groups: list[object] = []
self.use_device_communicator = use_device_communicator
self.device_communicator: NPUCommunicator | None = None
self.mq_broadcaster = None
self.cpu_group = None
self.device_group = None
self.device = None
self.use_custom_op_call = True
self.use_cpu_custom_send_recv = False
self.group_name = group_name
self.group_ranks = group_ranks
try:
self._init_device_groups(create_cpu_group=True)
assert self.cpu_group is not None
assert self.device_group is not None
self._init_device_communicator()
from vllm.distributed.device_communicators.shm_broadcast import MessageQueue
if use_message_queue_broadcaster and self.world_size > 1:
self.mq_broadcaster = MessageQueue.create_from_process_group(
self.cpu_group,
1 << 22,
6,
)
except Exception:
try:
self.destroy()
except Exception:
logger.exception("Failed to clean up partially initialized GroupCoordinatorPatch")
raise
def _init_device_groups(self, create_cpu_group: bool) -> None:
reuse_domain = _resolve_reuse_domain(self.group_name)
self_device_group = None
for ranks in self.group_ranks:
hccl_pg_options = create_hccl_pg_options(self.group_name)
device_group, hccl_key = _acquire_hccl_group(
ranks=ranks,
backend=self.backend,
hccl_pg_options=hccl_pg_options,
reuse_domain=reuse_domain,
)
if hccl_key is not None:
self._acquired_hccl_keys.append(hccl_key)
elif self.backend == "hccl" and self.rank in ranks:
self._unshared_hccl_groups.append(device_group)
cpu_group = torch.distributed.new_group(ranks, backend="gloo") if create_cpu_group else None
if self.rank in ranks:
if create_cpu_group:
self.ranks = ranks
self.world_size = len(ranks)
self.rank_in_group = ranks.index(self.rank)
self.cpu_group = cpu_group
self_device_group = device_group
if self_device_group is not None:
self.device_group = self_device_group
def _init_device_communicator(self) -> None:
self.device = torch.npu.current_device()
if self.use_device_communicator and self.world_size > 1:
self.device_communicator = NPUCommunicator(
cpu_group=self.cpu_group,
device=self.device,
device_group=self.device_group,
unique_name=self.unique_name,
)
def _release_hccl_resources(self) -> bool:
destroyed = False
device_communicator = getattr(self, "device_communicator", None)
if device_communicator is not None:
device_communicator.destroy()
self.device_communicator = None
destroyed = True
if hasattr(self, "_acquired_hccl_keys"):
for hccl_key in reversed(self._acquired_hccl_keys):
_HCCL_PG_REGISTRY.release(hccl_key)
self._acquired_hccl_keys = []
destroyed = True
if hasattr(self, "_unshared_hccl_groups"):
for device_group in reversed(self._unshared_hccl_groups):
torch.distributed.destroy_process_group(device_group)
self._unshared_hccl_groups = []
destroyed = True
return destroyed
def destroy(self):
if getattr(self, "mq_broadcaster", None) is not None:
self.mq_broadcaster = None
self._release_hccl_resources()
device_group = getattr(self, "device_group", None)
if device_group is not None and self.backend != "hccl":
torch.distributed.destroy_process_group(device_group)
if hasattr(self, "device_group"):
del self.device_group
cpu_group = getattr(self, "cpu_group", None)
if cpu_group is not None:
torch.distributed.destroy_process_group(cpu_group)
if hasattr(self, "cpu_group"):
del self.cpu_group
def destroy_hccl(self) -> bool:
"""Release the HCCL process group."""
destroyed = self._release_hccl_resources()
if hasattr(self, "device_group"):
self.device_group = None
return destroyed
def restore_hccl(self) -> bool:
"""Recreate the HCCL process group in place after sleep mode."""
if self.device_group is not None:
return False
self._init_device_groups(create_cpu_group=False)
assert self.device_group is not None
self._init_device_communicator()
return True
def all_to_all(
self,
input_: torch.Tensor,
scatter_dim: int = 0,
gather_dim: int = -1,
scatter_sizes: list[int] | None = None,
gather_sizes: list[int] | None = None,
) -> torch.Tensor:
if self.world_size == 1:
return input_
assert -input_.dim() <= scatter_dim < input_.dim(), (
f"Invalid scatter dim ({scatter_dim}) for input tensor with shape {input_.size()}"
)
assert -input_.dim() <= gather_dim < input_.dim(), (
f"Invalid gather dim ({gather_dim}) for input tensor with shape {input_.size()}"
)
assert self.device_communicator is not None, "device_communicator should be initialized when world_size > 1"
return self.device_communicator.all_to_all(input_, scatter_dim, gather_dim, scatter_sizes, gather_sizes)
vllm.distributed.parallel_state.GroupCoordinator = GroupCoordinatorPatch
_patch_destroy_distributed_environment()

View File

@@ -0,0 +1,145 @@
import logging
import os
from collections.abc import Iterable
from pathlib import Path
import torch
from safetensors.torch import load_file
from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM
from vllm.model_executor.models.utils import (
AutoWeightsLoader,
process_eagle_weight,
)
logger = logging.getLogger(__name__)
def get_embedding_tensor(directory_path):
"""
Scans the directory and returns the first tensor found that contains 'embed' in its key.
Returns the tensor if found, otherwise None.
"""
if not os.path.isdir(directory_path):
return None
# List files and filter for .safetensors
for filename in os.listdir(directory_path):
if filename.endswith(".safetensors"):
file_path = os.path.join(directory_path, filename)
# Load the file
state_dict = load_file(file_path)
# Search for the first matching key
for key, tensor in state_dict.items():
if "embed" in key.lower():
# Return immediately once found
return tensor
return None
def get_rotation_path(target_vllm_config):
"""
Gets the path of the rotation matrix, returns None if the target model is not a quarot model.
"""
target_model_path = target_vllm_config.model_config.model
try:
quant_description = target_vllm_config.quant_config.quant_description
rotation_relative_path = quant_description["optional"]["quarot"]["rotation_map"]["global_rotation"]
except KeyError:
return None
return Path(target_model_path) / rotation_relative_path
def get_rotataion_matrix(rotation_path):
"""
Anti-rotate maxtrix.
"""
try:
safetensor_data = load_file(rotation_path)
Q = safetensor_data["global_rotation"]
return Q
except Exception as e:
logger.error(
"Failed to load rotation weight from '%s'. If you want to use quarot model with eagle3, take a check.",
rotation_path,
)
raise e
def compute_rotataion_matrix3(Q):
"""
Anti-rotate matrix for 3 layers of hidden_states.
"""
return torch.block_diag(Q, Q, Q)
def patch_load_weights(target_vllm_config):
target_model_path = Path(target_vllm_config.model_config.model)
rotation_path = get_rotation_path(target_vllm_config)
# if rotation path is not found, then quarot is not in use.
if rotation_path is None:
return
Eagle3LlamaForCausalLM.load_weights = make_load_weights(target_model_path, rotation_path)
def make_load_weights(target_model_path, rotation_path):
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
Q = get_rotataion_matrix(rotation_path)
Q3 = compute_rotataion_matrix3(Q)
if isinstance(self.config.dtype, str):
embed_dtype = getattr(torch, self.config.dtype)
else:
embed_dtype = self.config.dtype
model_weights = {}
includes_draft_id_mapping = False
includes_embed_tokens = False
for name, loaded_weight in weights:
if "t2d" in name:
continue
if "d2t" in name:
name = name.replace("d2t", "draft_id_to_target_id")
includes_draft_id_mapping = True
elif "lm_head" not in name:
name = "model." + name
if "fc." in name:
# anti-rotate fc
dtype = loaded_weight.dtype
loaded_weight = (loaded_weight.to(torch.float32) @ Q3.to(torch.float32)).to(dtype)
if "embed_tokens" in name:
includes_embed_tokens = True
model_weights[name] = loaded_weight
process_eagle_weight(self, name)
# process embedding if drafter does not have embedding
if not includes_embed_tokens:
name = "model.embed_tokens.weight"
loaded_weight = (get_embedding_tensor(target_model_path).to(torch.float32) @ Q.T.to(torch.float32)).to(
embed_dtype
)
model_weights[name] = loaded_weight
includes_embed_tokens = True
process_eagle_weight(self, name)
skip_substrs = []
if not includes_draft_id_mapping:
skip_substrs.append("draft_id_to_target_id")
if not includes_embed_tokens:
skip_substrs.append("embed_tokens")
if not self.model.use_aux_hidden_state:
skip_substrs.append("fc.")
loader = AutoWeightsLoader(
self,
skip_prefixes=None,
skip_substrs=skip_substrs,
)
loader.load_weights(model_weights.items())
return load_weights

View File

@@ -0,0 +1,129 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
"""
Patch: fix target_layer_num for Eagle3 draft models under Pipeline Parallelism.
Upstream Eagle3 draft models (Eagle3LlamaForCausalLM, Eagle3DeepseekV2ForCausalLM)
compute ``target_layer_num`` via ``model_config.get_num_layers(parallel_config)``
which, under PP, returns the **per-PP-stage** count. This value feeds into the
draft model's ``start_layer_id`` (used to build parameter name prefixes like
``model.layers.<start_layer_id + i>``). With PP>1 the prefixes collide with
the checkpoint (e.g. a 61-layer target + 2-way PP builds prefixes 31..34 while
the checkpoint expects 61..64), breaking weight loading. Additionally,
``config.target_layer_count`` (used to index ``layer_types`` for draft
attention) ends up wrong.
Fix: use ``get_total_num_hidden_layers()`` instead. This matches the
checkpoint's global layer indices and keeps ``target_layer_count`` correct.
Currently patches:
- Eagle3LlamaForCausalLM (Qwen, LLaMA-based Eagle3 targets)
- Eagle3DeepseekV2ForCausalLM / Eagle3DeepseekV3ForCausalLM (DeepSeek-V2/V3,
Kimi K2/K2.6)
"""
import logging
import torch
import torch.nn as nn
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead
from vllm.model_executor.models.deepseek_eagle3 import (
DeepseekV2Eagle3Model,
Eagle3DeepseekV2ForCausalLM,
)
from vllm.model_executor.models.llama_eagle3 import (
Eagle3LlamaForCausalLM,
LlamaModel,
get_draft_quant_config,
)
from vllm.model_executor.models.utils import maybe_prefix
logger = logging.getLogger(__name__)
def _patched_eagle3_llama_init(self, *, vllm_config, prefix: str = ""):
nn.Module.__init__(self)
self.config = vllm_config.speculative_config.draft_model_config.hf_config
if getattr(self.config, "draft_vocab_size", None) is None:
base_vocab_size = getattr(self.config, "vocab_size", None)
self.config.draft_vocab_size = base_vocab_size
target_layer_num = vllm_config.model_config.get_total_num_hidden_layers()
self.config.target_layer_count = target_layer_num
self.model = LlamaModel(vllm_config=vllm_config, prefix="model", start_layer_id=target_layer_num)
logit_scale = getattr(self.config, "logit_scale", 1.0)
self.lm_head = ParallelLMHead(
self.config.draft_vocab_size,
self.config.hidden_size,
quant_config=get_draft_quant_config(vllm_config),
prefix=maybe_prefix(prefix, "lm_head"),
)
self.logits_processor = LogitsProcessor(self.config.draft_vocab_size, scale=logit_scale)
self.draft_id_to_target_id = nn.Parameter(
torch.zeros(self.config.draft_vocab_size, dtype=torch.long),
requires_grad=False,
)
self.use_parallel_drafting = vllm_config.speculative_config.parallel_drafting
if self.use_parallel_drafting:
self.register_buffer(
"mask_hidden",
torch.zeros(
1,
(3 if self.model.use_aux_hidden_state else 1) * self.config.hidden_size,
),
persistent=False,
)
def _patched_eagle3_deepseek_v2_init(self, *, vllm_config, prefix: str = ""):
nn.Module.__init__(self)
self.config = vllm_config.speculative_config.draft_model_config.hf_config
if getattr(self.config, "draft_vocab_size", None) is None:
base_vocab_size = getattr(self.config, "vocab_size", None)
self.config.draft_vocab_size = base_vocab_size
target_layer_num = vllm_config.model_config.get_total_num_hidden_layers()
self.config.target_layer_count = target_layer_num
self.model = DeepseekV2Eagle3Model(vllm_config=vllm_config, prefix="model", start_layer_id=target_layer_num)
logit_scale = getattr(self.config, "logit_scale", 1.0)
self.lm_head = ParallelLMHead(
self.config.draft_vocab_size,
self.config.hidden_size,
prefix=maybe_prefix(prefix, "lm_head"),
)
self.logits_processor = LogitsProcessor(self.config.draft_vocab_size, scale=logit_scale)
self.draft_id_to_target_id = nn.Parameter(
torch.zeros(self.config.draft_vocab_size, dtype=torch.long),
requires_grad=False,
)
Eagle3LlamaForCausalLM.__init__ = _patched_eagle3_llama_init
Eagle3DeepseekV2ForCausalLM.__init__ = _patched_eagle3_deepseek_v2_init
logger.info(
"Patched Eagle3LlamaForCausalLM and Eagle3DeepseekV2ForCausalLM "
"__init__ to use get_total_num_hidden_layers() for target_layer_num."
)

View File

@@ -0,0 +1,238 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
"""
Patch: Propagate Eagle3 aux hidden states through PP pipeline.
In Eagle3 speculative decoding with Pipeline Parallelism (PP), auxiliary
hidden states are collected from specific target model layers (e.g., layers
2, N/2, N-3). When these layers span multiple PP stages, the last PP rank
(where the drafter runs) only sees a subset of aux states, causing
combine_hidden_states to fail with k-axis shape mismatch.
This patch wraps the inner model's forward and make_empty_intermediate_tensors
to transparently pass aux hidden states through IntermediateTensors across PP
stages. Each PP stage carries forward all aux states from previous stages,
and the last PP rank merges them into a single list for the drafter.
Currently supports:
- DeepseekV2Model (used by Kimi K2/K2.6, DeepSeek-V2/V3)
- EagleModelMixin-based models (MiniMaxM2, Llama, Qwen2, etc.)
"""
import logging
from itertools import islice
import torch
import torch.nn as nn
from vllm.distributed.parallel_state import get_pp_group
from vllm.sequence import IntermediateTensors
from vllm.v1.attention.backend import AttentionMetadata
logger = logging.getLogger(__name__)
_AUX_KEY_PREFIX = "aux_layer_"
def _extract_aux_from_intermediate(
intermediate_tensors: "IntermediateTensors | None",
) -> list[torch.Tensor]:
if intermediate_tensors is None:
return []
aux_keys = sorted(
(k for k in intermediate_tensors.tensors if k.startswith(_AUX_KEY_PREFIX)),
key=lambda k: int(k.split("_")[-1]),
)
return [intermediate_tensors.tensors[k] for k in aux_keys]
def _make_deepseek_v2_forward():
def pp_eagle3_forward(
self,
input_ids: "torch.Tensor | None",
positions: torch.Tensor,
kv_caches: list[torch.Tensor],
attn_metadata: "AttentionMetadata",
intermediate_tensors: "IntermediateTensors | None" = None,
inputs_embeds: "torch.Tensor | None" = None,
):
pp_group = get_pp_group()
prev_aux_list = _extract_aux_from_intermediate(intermediate_tensors)
if pp_group.is_first_rank:
if inputs_embeds is not None:
hidden_states = inputs_embeds
else:
if input_ids is None:
raise ValueError("Either input_ids or inputs_embeds must be provided to DeepseekV2Model.forward")
hidden_states = self.embed_input_ids(input_ids)
residual = None
else:
assert intermediate_tensors is not None
hidden_states = intermediate_tensors["hidden_states"]
residual = intermediate_tensors["residual"]
llama_4_scaling_config = getattr(self.config, "llama_4_scaling", None)
llama_4_scaling: torch.Tensor | None = None
if llama_4_scaling_config is not None:
from vllm.model_executor.models.deepseek_v2 import _get_llama_4_scaling
llama_4_scaling = _get_llama_4_scaling(
original_max_position_embeddings=llama_4_scaling_config["original_max_position_embeddings"],
scaling_beta=llama_4_scaling_config["beta"],
positions=positions,
)
aux_hidden_states: list[torch.Tensor] = list(prev_aux_list)
for idx, layer in enumerate(
islice(self.layers, self.start_layer, self.end_layer),
start=self.start_layer,
):
if idx in self.aux_hidden_state_layers:
aux_hidden_states.append(hidden_states + residual if residual is not None else hidden_states)
hidden_states, residual = layer(
positions,
hidden_states,
residual,
kv_caches[idx - self.start_layer],
attn_metadata,
llama_4_scaling,
)
if not pp_group.is_last_rank:
result = IntermediateTensors(
{
"hidden_states": hidden_states,
"residual": residual,
}
)
for i, t in enumerate(aux_hidden_states):
result.tensors[f"{_AUX_KEY_PREFIX}{i}"] = t
return result
hidden_states, _ = self.norm(hidden_states, residual)
if len(aux_hidden_states) > 0:
return hidden_states, aux_hidden_states
return hidden_states
return pp_eagle3_forward
def _make_eagle_mixin_forward():
def pp_eagle3_forward(
self,
input_ids: "torch.Tensor | None",
positions: torch.Tensor,
intermediate_tensors: "IntermediateTensors | None" = None,
inputs_embeds: "torch.Tensor | None" = None,
):
pp_group = get_pp_group()
prev_aux_list = _extract_aux_from_intermediate(intermediate_tensors)
if pp_group.is_first_rank:
if inputs_embeds is not None:
hidden_states = inputs_embeds
else:
hidden_states = self.embed_input_ids(input_ids)
residual = None
else:
assert intermediate_tensors is not None
hidden_states = intermediate_tensors["hidden_states"]
residual = intermediate_tensors["residual"]
aux_hidden_states = self._maybe_add_hidden_state(list(prev_aux_list), 0, hidden_states, residual)
for idx, layer in enumerate(
islice(self.layers, self.start_layer, self.end_layer),
start=self.start_layer,
):
hidden_states, residual = layer(positions, hidden_states, residual)
self._maybe_add_hidden_state(aux_hidden_states, idx + 1, hidden_states, residual)
if not pp_group.is_last_rank:
result = IntermediateTensors(
{
"hidden_states": hidden_states,
"residual": residual,
}
)
for i, t in enumerate(aux_hidden_states):
result.tensors[f"{_AUX_KEY_PREFIX}{i}"] = t
return result
hidden_states, _ = self.norm(hidden_states, residual)
if len(aux_hidden_states) > 0:
return hidden_states, aux_hidden_states
return hidden_states
return pp_eagle3_forward
def _patch_make_empty_intermediate_tensors(inner_model: nn.Module) -> None:
if getattr(inner_model, "_eagle3_pp_aux_make_empty_patched", False):
return
original_make_empty = inner_model.make_empty_intermediate_tensors
def pp_make_empty_intermediate_tensors(batch_size, dtype, device):
result = original_make_empty(batch_size, dtype, device)
aux_layers = getattr(inner_model, "aux_hidden_state_layers", ())
# A non-first PP rank only receives aux hidden states produced by
# earlier pipeline stages. Local aux states are appended during forward.
num_incoming_aux_layers = sum(layer_idx < inner_model.start_layer for layer_idx in aux_layers)
hidden_size = inner_model.config.hidden_size
for i in range(num_incoming_aux_layers):
result.tensors[f"{_AUX_KEY_PREFIX}{i}"] = torch.zeros(
(batch_size, hidden_size),
dtype=dtype,
device=device,
)
return result
inner_model.make_empty_intermediate_tensors = pp_make_empty_intermediate_tensors
inner_model._eagle3_pp_aux_make_empty_patched = True
def patch_eagle3_pp_aux_propagation(inner_model: nn.Module) -> bool:
from vllm.model_executor.models.deepseek_v2 import DeepseekV2Model
from vllm.model_executor.models.interfaces import EagleModelMixin
if isinstance(inner_model, DeepseekV2Model):
make_forward = _make_deepseek_v2_forward
elif isinstance(inner_model, EagleModelMixin):
make_forward = _make_eagle_mixin_forward
else:
logger.warning(
"Eagle3 PP aux propagation is only supported for DeepseekV2Model "
"or EagleModelMixin-based models, got %s. Skipping patch.",
type(inner_model).__name__,
)
return False
if not getattr(inner_model, "_eagle3_pp_aux_forward_patched", False):
inner_model.forward = make_forward().__get__(inner_model, type(inner_model))
inner_model._eagle3_pp_aux_forward_patched = True
_patch_make_empty_intermediate_tensors(inner_model)
logger.info(
"Applied Eagle3 PP aux propagation patch to %s (aux_layers=%s, start_layer=%d, end_layer=%d).",
type(inner_model).__name__,
inner_model.aux_hidden_state_layers,
inner_model.start_layer,
inner_model.end_layer,
)
return True

View File

@@ -0,0 +1,20 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# Reuse the platform patch. Keeping the monkey patch in one module avoids
# wrapping an already patched FusedMoE factory during worker initialization.
import vllm_ascend.patch.platform.patch_fused_moe # noqa: F401

View File

@@ -0,0 +1,78 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
import logging
from collections.abc import Callable, Iterable
import torch
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
from vllm.model_executor.models.glm4_moe import Glm4MoeForCausalLM
from vllm.model_executor.models.minimax_m2 import MiniMaxM2ForCausalLM
from vllm.model_executor.models.qwen3 import Qwen3ForCausalLM
logger = logging.getLogger(__name__)
_orig_qwen3_causal_lm_load_weights = Qwen3ForCausalLM.load_weights
_orig_Glm4_causal_lm_load_weights = Glm4MoeForCausalLM.load_weights
_orig_Minimax_m2_causal_lm_load_weights = MiniMaxM2ForCausalLM.load_weights
def _patched_causal_lm_load_weights(
self, weights: Iterable[tuple[str, torch.Tensor]], original_load_weights: Callable
) -> set[str]:
quant_config = self.quant_config
if quant_config is None or not callable(getattr(quant_config, "get_cache_scale", None)):
return original_load_weights(self, weights)
params_dict = dict(self.named_parameters())
c8_loaded_params: set[str] = set()
def _intercept_c8_scales(
raw_weights: Iterable[tuple[str, torch.Tensor]],
) -> Iterable[tuple[str, torch.Tensor]]:
for name, loaded_weight in raw_weights:
scale_name = quant_config.get_cache_scale(name)
if scale_name is not None:
if scale_name in params_dict:
param = params_dict[scale_name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight.squeeze())
c8_loaded_params.add(scale_name)
else:
logger.warning(
"Cache scale %s found in quant_config for weight %s "
"but not found in model parameters; weight will be skipped.",
scale_name,
name,
)
else:
yield name, loaded_weight
loaded_params = original_load_weights(self, _intercept_c8_scales(weights))
loaded_params.update(c8_loaded_params)
return loaded_params
Qwen3ForCausalLM.load_weights = lambda self, weights: _patched_causal_lm_load_weights(
self, weights, _orig_qwen3_causal_lm_load_weights
)
Glm4MoeForCausalLM.load_weights = lambda self, weights: _patched_causal_lm_load_weights(
self, weights, _orig_Glm4_causal_lm_load_weights
)
MiniMaxM2ForCausalLM.load_weights = lambda self, weights: _patched_causal_lm_load_weights(
self, weights, _orig_Minimax_m2_causal_lm_load_weights
)

View File

@@ -0,0 +1,54 @@
import vllm
from vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn import QwenGatedDeltaNetAttention
from vllm_ascend._310p.ops.fla.gdn_310 import AscendGatedDeltaNetAttention310
from vllm_ascend._310p.ops.fla.idex import (
prepare_chunk_indices_310,
prepare_chunk_offsets_310,
)
from vllm_ascend._310p.spec_decode.llm_base_proposer_310 import AscendSpecDecodeBaseProposer310
from vllm_ascend.ops.gdn import AscendGatedDeltaNetAttention
from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer
from vllm_ascend.utils import is_rc_device
vllm.model_executor.layers.fla.ops.index.prepare_chunk_indices = prepare_chunk_indices_310
vllm.model_executor.layers.fla.ops.index.prepare_chunk_offsets = prepare_chunk_offsets_310
# 310P: protect tail slot during MTP input_ids shift to avoid GatherV2 corruption
# caused by the NPU slice-assign writing one element past the intended range
# on the persistent drafter input_ids buffer.
AscendSpecDecodeBaseProposer.set_inputs_first_pass = ( # type: ignore[method-assign]
AscendSpecDecodeBaseProposer310.set_inputs_first_pass
)
AscendSpecDecodeBaseProposer._run_merged_draft = ( # type: ignore[method-assign]
AscendSpecDecodeBaseProposer310._run_merged_draft
)
# Patch _warmup_prefill_kernels to no-op on 310P: triton.next_power_of_2 does
# not exist in the triton version used on 310P CI, and NPU does not use these
# CUDA warmup kernel anyway.
QwenGatedDeltaNetAttention._warmup_prefill_kernels = lambda self, qkv_or_qkvz, v_dim: None # type: ignore[method-assign]
QwenGatedDeltaNetAttention._split_ba_for_tp = AscendGatedDeltaNetAttention._split_ba_for_tp
QwenGatedDeltaNetAttention.get_state_shape = AscendGatedDeltaNetAttention.get_state_shape
QwenGatedDeltaNetAttention._forward_core = AscendGatedDeltaNetAttention310._forward_core
QwenGatedDeltaNetAttention.get_state_dtype = AscendGatedDeltaNetAttention310.get_state_dtype
# 310P: make Qwen GDN use the 310P attention backend, including the
# MTP ACL graph padding replay fixes provided by gdn_attn_builder_310.py.
QwenGatedDeltaNetAttention.get_attn_backend = AscendGatedDeltaNetAttention310.get_attn_backend
if is_rc_device():
from vllm.model_executor.models.qwen3_vl import Qwen3_VisionTransformer
from vllm.v1.attention.backends.gdn_attn import GDNAttentionBackend
from vllm_ascend._310p.ops.gdn_attn_builder_310 import GDNAttentionMetadataBuilder310
from vllm_ascend._310p.ops.qwen3vl_310 import rot_pos_emb_310
# 310P RC: use blocking H2D in rot_pos_emb to avoid race with subsequent indexing.
Qwen3_VisionTransformer.rot_pos_emb = rot_pos_emb_310 # type: ignore[method-assign]
# Qwen3.5 on 310P RC uses upstream GDNAttentionBackend via MambaBase.get_attn_backend().
GDNAttentionBackend.get_builder_cls = staticmethod( # type: ignore[method-assign]
lambda: GDNAttentionMetadataBuilder310
)

View File

@@ -0,0 +1,95 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
import torch
import torch.nn as nn
import torch.nn.functional as F
from vllm.model_executor.models.kimi_k25_vit import (
Learnable2DInterpPosEmbDivided_fixed,
MoonViT3dPretrainedModel,
get_rope_shape_decorate,
)
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
@get_rope_shape_decorate
def get_rope_shape(org, interpolation_mode, shape):
return (
F.interpolate(
org.permute((2, 0, 1)).unsqueeze(0),
size=shape,
mode=interpolation_mode,
)
.squeeze(0)
.permute((1, 2, 0))
.flatten(end_dim=1)
)
class AscendLearnable2DInterpPosEmbDivided_fixed(nn.Module):
def forward(self, x: torch.Tensor, grid_thws: torch.Tensor | list) -> torch.Tensor:
pos_embs = []
if isinstance(grid_thws, torch.Tensor):
grid_list = grid_thws.tolist()
else:
grid_list = grid_thws
for t, h, w in grid_list:
assert t <= self.num_frames, (
f"[vllm-ascend/patch_kimi_k25] Invalid frame count. t={t}, num_frames={self.num_frames}"
)
if (h, w) == self.weight.shape[:-1]:
pos_emb_2d = self.weight.flatten(end_dim=1)
else:
pos_emb_2d = get_rope_shape(
self.weight,
interpolation_mode=self.interpolation_mode,
shape=(h, w),
)
if t == 1:
pos_emb_3d = pos_emb_2d
else:
pos_emb_3d = pos_emb_2d.unsqueeze(0).repeat(t, 1, 1) + self.time_weight[0:t]
pos_embs.append(pos_emb_3d.reshape(-1, pos_emb_3d.shape[-1]))
out = x + torch.cat(pos_embs)
return out
Learnable2DInterpPosEmbDivided_fixed.forward = AscendLearnable2DInterpPosEmbDivided_fixed.forward
# Patch MoonViT3dPretrainedModel.to() to ignore the `dtype` argument.
# When KimiK25ForConditionalGeneration.__init__ calls:
# self.vision_tower = self.vision_tower.to(device=..., dtype=model_config.dtype)
# the `dtype=model_config.dtype` (e.g. bf16) would overwrite the fp8 parameters
# created by the Ascend quantization scheme, causing a dtype mismatch later
# in weight_loader when the checkpoint's fp8 weights are loaded.
if get_ascend_device_type() == AscendDeviceType.A5:
_original_moonvit_to = MoonViT3dPretrainedModel.to
def _patched_moonvit_to(self, *args, **kwargs):
# Filter out dtype from positional arguments and remove from kwargs
# to prevent overriding quantized weight dtypes on A5.
new_args = tuple(a for a in args if not isinstance(a, torch.dtype))
kwargs.pop("dtype", None)
return _original_moonvit_to(self, *new_args, **kwargs)
MoonViT3dPretrainedModel.to = _patched_moonvit_to

View File

@@ -0,0 +1,284 @@
# mypy: ignore-errors
import itertools
from typing import Any
import torch
from vllm.config import CacheConfig
from vllm.model_executor.layers.mamba.mamba_utils import MambaStateCopyFunc
from vllm.utils.math_utils import cdiv
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker import mamba_utils
from vllm.v1.worker.gpu_input_batch import CachedRequestState
from vllm.v1.worker.lora_model_runner_mixin import GPUInputBatch
from vllm.v1.worker.mamba_utils import MambaCopyBuffers
from vllm_ascend.ops.triton.batch_memcpy import batch_memcpy_kernel
from vllm_ascend.ops.triton.mamba.postprocess import postprocess_mamba_fused_kernel
from vllm_ascend.utils import is_310p
def _can_launch_triton_batch_memcpy() -> bool:
return not is_310p()
def _batch_memcpy_triton(src_ptrs, dst_ptrs, sizes):
batch = src_ptrs.shape[0]
assert dst_ptrs.shape[0] == batch
assert sizes.shape[0] == batch
grid = (batch,)
# using larger block_size to accelerate copy.
BLOCK_SIZE = 8192
batch_memcpy_kernel[grid](src_ptrs, dst_ptrs, sizes, BLOCK_SIZE=BLOCK_SIZE)
def _tensor_view_from_data_ptr(state: torch.Tensor, start_addr: int, num_elements: int) -> torch.Tensor:
byte_offset = start_addr - state.data_ptr()
element_size = state.element_size()
if byte_offset < 0 or byte_offset % element_size != 0:
raise RuntimeError("Invalid Mamba state copy pointer.")
element_offset = byte_offset // element_size
flat_state = state.view(-1)
if element_offset + num_elements > flat_state.numel():
raise RuntimeError("Mamba state copy range exceeds tensor storage.")
return flat_state.narrow(0, element_offset, num_elements)
def _get_tensor_copy_pairs(copy_bufs: mamba_utils.MambaCopyBuffers) -> list[tuple[torch.Tensor, torch.Tensor]]:
if copy_bufs.offset == 0 or not hasattr(copy_bufs, "_tensor_copy_pairs"):
copy_bufs._tensor_copy_pairs = []
return copy_bufs._tensor_copy_pairs
def _collect_mamba_copy_meta_torch(
copy_bufs: mamba_utils.MambaCopyBuffers,
kv_cache_config,
mamba_state_copy_funcs,
mamba_group_ids: list[int],
src_block_idx: int,
dest_block_idx: int,
accept_token_bias: int,
req_state,
forward_context: dict[str, Any],
) -> None:
if src_block_idx == dest_block_idx and accept_token_bias == 0:
return
tensor_copy_pairs = _get_tensor_copy_pairs(copy_bufs)
sizes_np = copy_bufs.sizes.np
offset = copy_bufs.offset
for mamba_group_id in mamba_group_ids:
block_ids = req_state.block_ids[mamba_group_id]
dest_block_id = block_ids[dest_block_idx]
layer_names = kv_cache_config.kv_cache_groups[mamba_group_id].layer_names
for layer_name in layer_names:
attention = forward_context[layer_name]
kv_caches: list[torch.Tensor] = attention.kv_cache
for state, state_copy_func in zip(kv_caches, mamba_state_copy_funcs):
copy_spec = state_copy_func(state, block_ids, src_block_idx, accept_token_bias + 1)
src_state = _tensor_view_from_data_ptr(state, copy_spec.start_addr, copy_spec.num_elements)
dst_state = _tensor_view_from_data_ptr(state, state[dest_block_id].data_ptr(), copy_spec.num_elements)
tensor_copy_pairs.append((src_state, dst_state))
sizes_np[offset] = copy_spec.num_elements * state.element_size()
offset += 1
copy_bufs.offset = offset
def _do_mamba_copy_block_torch(copy_bufs: mamba_utils.MambaCopyBuffers):
n = copy_bufs.offset
if n == 0:
if hasattr(copy_bufs, "_tensor_copy_pairs"):
copy_bufs._tensor_copy_pairs = []
return
tensor_copy_pairs = getattr(copy_bufs, "_tensor_copy_pairs", None)
if tensor_copy_pairs is None or len(tensor_copy_pairs) != n:
raise RuntimeError("Mamba tensor copy metadata is incomplete.")
for src_state, dst_state in tensor_copy_pairs:
dst_state.copy_(src_state.clone())
copy_bufs._tensor_copy_pairs = []
def _postprocess_mamba_align_gpu_cpu_fallback(
*,
bufs: "mamba_utils.MambaBuffers",
num_reqs: int,
num_accepted_tokens_gpu: torch.Tensor,
num_accepted_tokens_cpu_tensor: torch.Tensor,
input_batch: GPUInputBatch,
kv_cache_config: KVCacheConfig,
forward_context: dict[str, Any],
mamba_state_copy_funcs: tuple[MambaStateCopyFunc, ...],
) -> None:
"""CPU fallback for 310P where the Triton fused postprocess is unavailable."""
ctx = bufs.postprocess_align
assert ctx is not None
assert ctx.mamba_state_idx_buf is not None
assert ctx.num_scheduled_tokens_buf is not None
assert ctx.num_computed_tokens_buf is not None
assert ctx.num_draft_tokens_buf is not None
# stage_postprocess_inputs_to_gpu has already materialized the same
# per-request values into the CpuGpuBuffer numpy views. 310P cannot use the
# Triton fused kernel, so reuse the CPU views to mirror its decision logic.
mamba_state_idx = ctx.mamba_state_idx_buf.np
num_scheduled_tokens = ctx.num_scheduled_tokens_buf.np
num_computed_tokens = ctx.num_computed_tokens_buf.np
num_draft_tokens = ctx.num_draft_tokens_buf.np
block_size = ctx.block_size
# Upstream initializes num_accepted_tokens_out from the real accepted-token
# counts, then only overwrites entries where src and dest are the same
# block. Preserve that default so the next preprocess keeps the right
# accept_token_bias when multiple draft tokens were accepted.
num_accepted_tokens_cpu_tensor[:num_reqs].copy_(num_accepted_tokens_gpu[:num_reqs])
num_accepted_tokens = input_batch.num_accepted_tokens_cpu
for i in range(num_reqs):
num_tokens_running_state = num_computed_tokens[i] + num_scheduled_tokens[i] - num_draft_tokens[i]
new_num_computed_tokens = num_tokens_running_state + num_accepted_tokens[i] - 1
aligned_new_computed_tokens = new_num_computed_tokens // block_size * block_size
if aligned_new_computed_tokens < num_tokens_running_state:
continue
src_block_idx = mamba_state_idx[i]
dest_block_idx = aligned_new_computed_tokens // block_size - 1
accept_token_bias = aligned_new_computed_tokens - num_tokens_running_state
if src_block_idx == dest_block_idx:
# Match the fused kernel: once the running state remains in the
# same block, the next preprocess should start from token bias 0.
num_accepted_tokens_cpu_tensor[i] = 1
if accept_token_bias == 0:
continue
# The upstream fused kernel also copies Mamba state in this postprocess
# step. Do the same with tensor views so 310P avoids Triton without
# changing where conv/temporal state lands before the next iteration.
for mamba_group_id in ctx.mamba_group_ids:
block_ids = input_batch.block_table[mamba_group_id].get_numpy_array()[i]
dest_block_id = block_ids[dest_block_idx]
layer_names = kv_cache_config.kv_cache_groups[mamba_group_id].layer_names
for layer_name in layer_names:
attention = forward_context[layer_name]
kv_caches: list[torch.Tensor] = attention.kv_cache
for state, state_copy_func in zip(kv_caches, mamba_state_copy_funcs):
copy_spec = state_copy_func(state, block_ids, src_block_idx, accept_token_bias + 1)
src_state = _tensor_view_from_data_ptr(state, copy_spec.start_addr, copy_spec.num_elements)
dst_state = _tensor_view_from_data_ptr(
state, state[dest_block_id].data_ptr(), copy_spec.num_elements
)
dst_state.copy_(src_state.clone())
def _batch_memcpy_unavailable(src_ptrs, dst_ptrs, sizes):
raise RuntimeError(
"Pointer-based Mamba batch memcpy requires Triton and is not available "
"on 310P. Use the tensor-copy fallback path instead."
)
if _can_launch_triton_batch_memcpy():
mamba_utils.batch_memcpy_kernel = batch_memcpy_kernel
mamba_utils.batch_memcpy = _batch_memcpy_triton
mamba_utils.postprocess_mamba_fused_kernel = postprocess_mamba_fused_kernel
else:
mamba_utils.batch_memcpy = _batch_memcpy_unavailable
mamba_utils.collect_mamba_copy_meta = _collect_mamba_copy_meta_torch
mamba_utils.do_mamba_copy_block = _do_mamba_copy_block_torch
mamba_utils.postprocess_mamba_align_gpu = _postprocess_mamba_align_gpu_cpu_fallback
# Ascend NPU does not support DT_UINT64 in aclnnInplaceZero.
# MambaCopyBuffers.create() uses torch.uint64 for src_ptrs/dst_ptrs,
# which triggers a runtime error. Remap to int64 at the source.
_original_create = MambaCopyBuffers.create
@classmethod
def _patched_create(cls, max_num_reqs, kv_cache_config, copy_funcs, make_buffer):
return _original_create(
max_num_reqs,
kv_cache_config,
copy_funcs,
lambda n, dtype: make_buffer(n, dtype=torch.int64 if dtype == torch.uint64 else dtype),
)
MambaCopyBuffers.create = _patched_create
def preprocess_mamba(
scheduler_output: SchedulerOutput,
kv_cache_config: KVCacheConfig,
cache_config: CacheConfig,
mamba_state_idx: dict[str, int],
input_batch: GPUInputBatch,
requests: dict[str, CachedRequestState],
forward_context: dict[str, Any],
mamba_state_copy_funcs: tuple[MambaStateCopyFunc, ...],
copy_bufs: MambaCopyBuffers,
):
"""
Copy the mamba state of previous step to the last
(1 + num_speculative_blocks) block.
"""
mamba_group_ids = copy_bufs.mamba_group_ids
mamba_spec = copy_bufs.mamba_spec
num_speculative_blocks = mamba_spec.num_speculative_blocks
# TODO(Chen): we need to optimize this function a lot
# assert cache_config.enable_prefix_caching
block_size = mamba_spec.block_size
finished_req_ids = scheduler_output.finished_req_ids
preempted_req_ids = scheduler_output.preempted_req_ids or set()
resumed_req_ids = scheduler_output.scheduled_cached_reqs.resumed_req_ids
for req_id in itertools.chain(finished_req_ids, preempted_req_ids, resumed_req_ids):
mamba_state_idx.pop(req_id, None)
copy_bufs.offset = 0
for i, req_id in enumerate(input_batch.req_ids):
req_state = requests[req_id]
prev_state_idx = mamba_state_idx.get(req_id)
if prev_state_idx is None:
# new / resumed request, no previous state
# if num_computed_tokens is 0, prev_state_idx will be -1
prev_state_idx = (req_state.num_computed_tokens - 1) // block_size
num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id]
num_blocks: int = (
cdiv(req_state.num_computed_tokens + num_scheduled_tokens, block_size) + num_speculative_blocks
)
# We always save the current running state at the last
# (1 + num_speculative_blocks) block.
# A corner case worth mention here: assume we have block_size = 4 and
# num_speculative_tokens = 2. The request is [A, B, C] and contains 2 draft
# tokens [draft 1, draft 2]. Then we will have:
# Block 0: [A, B, C, draft 1]
# Block 1: [draft 2, TOFILL, TOFILL, TOFILL]
# Block 2: speculative block
# Block 3: speculative block
# And use block 1 to save the running state.
curr_state_idx = num_blocks - 1 - num_speculative_blocks
mamba_state_idx[req_id] = curr_state_idx
if prev_state_idx != -1 and prev_state_idx != curr_state_idx:
mamba_utils.collect_mamba_copy_meta(
copy_bufs,
kv_cache_config,
mamba_state_copy_funcs,
mamba_group_ids,
prev_state_idx,
curr_state_idx,
input_batch.num_accepted_tokens_cpu[i] - 1,
req_state,
forward_context,
)
input_batch.num_accepted_tokens_cpu[i] = 1
# do not copy here, since kv_transfer still not load
# do_mamba_copy_block(copy_bufs)
mamba_utils.preprocess_mamba = preprocess_mamba

View File

@@ -0,0 +1,181 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# MiniMax-M2 on Ascend: MoE router logits, fused attention, fp8 load dequant.
#
from collections.abc import Iterable
import torch
from vllm.model_executor.models.minimax_m2 import (
MiniMaxM2Attention,
MiniMaxM2Model,
MiniMaxM2MoE,
)
from vllm.platforms import current_platform
from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_slice
FP8_DTYPES = tuple(
getattr(torch, dtype_name)
for dtype_name in (
"float8_e4m3fn",
"float8_e4m3fnuz",
"float8_e5m2",
"float8_e5m2fnuz",
"float8_e8m0fnu",
)
if hasattr(torch, dtype_name)
)
# ---------------------------------------------------------------------------
# MiniMaxM2MoE.forward: keep router logits in fp32 on NPU.
# ---------------------------------------------------------------------------
def _patched_moe_forward(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor:
num_tokens, hidden_dim = hidden_states.shape
hidden_states = hidden_states.view(-1, hidden_dim)
# router_logits: (num_tokens, n_experts)
router_logits, _ = self.gate(hidden_states.to(torch.float32))
final_hidden_states = self.experts(hidden_states=hidden_states, router_logits=router_logits)
return final_hidden_states.view(num_tokens, hidden_dim)
MiniMaxM2MoE.forward = _patched_moe_forward
# ---------------------------------------------------------------------------
# MiniMaxM2Attention: fused qkv split, rmsnorm, and rope on NPU.
# ---------------------------------------------------------------------------
def _patch_forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
) -> torch.Tensor:
qkv, _ = self.qkv_proj(hidden_states)
cos, sin = get_cos_and_sin_slice()
q, k, v = torch.ops.vllm.split_qkv_tp_rmsnorm_rope(
input=qkv,
q_weight=self.q_norm.weight,
k_weight=self.k_norm.weight,
q_hidden_size=self.q_size,
kv_hidden_size=self.kv_size,
head_dim=self.head_dim,
rotary_dim=getattr(self.rotary_emb, "rotary_dim", self.head_dim),
eps=self.q_norm.variance_epsilon,
tp_world=self.q_norm.tp_world,
cos=cos,
sin=sin,
)
attn_output = self.attn(q, k, v)
output, _ = self.o_proj(attn_output)
return output
MiniMaxM2Attention.forward = _patch_forward
# ---------------------------------------------------------------------------
# MiniMaxM2Model: fp8 dequant helpers and load_weights wrapper
# ---------------------------------------------------------------------------
def _need_dequantize_fp8_weights(self) -> bool:
quant_cfg = getattr(self.config, "quantization_config", None)
return (
isinstance(quant_cfg, dict) and quant_cfg.get("quant_method") == "fp8" and current_platform.device_name == "npu"
)
def _dequantize_fp8_block_weight(
fp8_weight: torch.Tensor,
weight_scale_inv: torch.Tensor,
block_size: tuple[int, int],
) -> torch.Tensor:
block_n, block_k = block_size
n, k = fp8_weight.shape
n_tiles = (n + block_n - 1) // block_n
k_tiles = (k + block_k - 1) // block_k
if tuple(weight_scale_inv.shape) != (n_tiles, k_tiles):
raise ValueError(
"Unexpected fp8 scale shape: "
f"weight={tuple(fp8_weight.shape)}, "
f"scale={tuple(weight_scale_inv.shape)}, "
f"block_size={block_size}"
)
expanded_scale = weight_scale_inv.repeat_interleave(block_n, dim=0).repeat_interleave(block_k, dim=1)
expanded_scale = expanded_scale[:n, :k].to(dtype=torch.bfloat16)
return fp8_weight.to(dtype=torch.bfloat16) * expanded_scale
def _fp8_dequant_weight_iter(
self: "MiniMaxM2Model",
weights: Iterable[tuple[str, torch.Tensor]],
) -> Iterable[tuple[str, torch.Tensor]]:
quant_cfg = getattr(self.config, "quantization_config", {})
block_cfg = quant_cfg.get("weight_block_size", [128, 128])
weight_block_size: tuple[int, int] = (128, 128)
if isinstance(block_cfg, list) and len(block_cfg) == 2:
weight_block_size = (int(block_cfg[0]), int(block_cfg[1]))
pending_fp8_weights: dict[str, torch.Tensor] = {}
pending_fp8_scales: dict[str, torch.Tensor] = {}
for name, loaded_weight in weights:
if name.endswith(".weight_scale_inv"):
paired_weight_name = name[: -len("_scale_inv")]
pending_weight = pending_fp8_weights.pop(paired_weight_name, None)
if pending_weight is None:
pending_fp8_scales[name] = loaded_weight
continue
loaded_weight = self._dequantize_fp8_block_weight(pending_weight, loaded_weight, weight_block_size)
name = paired_weight_name
elif loaded_weight.dtype in FP8_DTYPES and name.endswith(".weight"):
scale_name = f"{name}_scale_inv"
pending_scale = pending_fp8_scales.pop(scale_name, None)
if pending_scale is None:
pending_fp8_weights[name] = loaded_weight
continue
loaded_weight = self._dequantize_fp8_block_weight(loaded_weight, pending_scale, weight_block_size)
yield name, loaded_weight
if pending_fp8_weights or pending_fp8_scales:
raise ValueError(
"Unpaired fp8 MiniMax-M2 weight/scale tensors detected: "
f"pending_weights={len(pending_fp8_weights)}, "
f"pending_scales={len(pending_fp8_scales)}"
)
MiniMaxM2Model._need_dequantize_fp8_weights = _need_dequantize_fp8_weights
MiniMaxM2Model._dequantize_fp8_block_weight = staticmethod(_dequantize_fp8_block_weight)
MiniMaxM2Model._fp8_dequant_weight_iter = _fp8_dequant_weight_iter
_original_load_weights = MiniMaxM2Model.load_weights
def _patched_load_weights(
self: "MiniMaxM2Model",
weights: Iterable[tuple[str, torch.Tensor]],
) -> set[str]:
if self._need_dequantize_fp8_weights():
weights = self._fp8_dequant_weight_iter(weights)
return _original_load_weights(self, weights)
MiniMaxM2Model.load_weights = _patched_load_weights

View File

@@ -0,0 +1,154 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# MiniMax-M2 linear attention: MiniMaxText01RMSNormTP weight sharding and NPU q/k norm path.
#
import logging
from functools import partial
import torch
import torch.nn as nn
from vllm.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce,
)
from vllm.model_executor.custom_op import CustomOp
from vllm.model_executor.layers.minimax_rms_norm import ( # type: ignore[import-not-found]
MiniMaxText01RMSNormTP,
)
from vllm.platforms import current_platform
logger = logging.getLogger(__name__)
_ORIG_QK_METHOD_NAME: str | None = None
_original_qk_method = None
_qk_is_staticmethod = False
if hasattr(MiniMaxText01RMSNormTP, "forward_qk"):
_ORIG_QK_METHOD_NAME = "forward_qk"
_original_qk_method = getattr(MiniMaxText01RMSNormTP, _ORIG_QK_METHOD_NAME)
elif hasattr(MiniMaxText01RMSNormTP, "_normalize_qk"):
# Older vLLM versions
_ORIG_QK_METHOD_NAME = "_normalize_qk"
_original_qk_method = getattr(MiniMaxText01RMSNormTP, _ORIG_QK_METHOD_NAME)
if _ORIG_QK_METHOD_NAME is not None:
# Detect whether upstream defined it as a staticmethod (some versions do).
_orig_desc = MiniMaxText01RMSNormTP.__dict__.get(_ORIG_QK_METHOD_NAME)
_qk_is_staticmethod = isinstance(_orig_desc, staticmethod)
else:
logger.warning(
"Neither forward_qk nor _normalize_qk found on MiniMaxText01RMSNormTP; "
"MiniMax-M2 linear attention patching is a no-op. "
"This may indicate a vLLM API change."
)
def _patched_qk(
q_norm: "MiniMaxText01RMSNormTP",
k_norm: "MiniMaxText01RMSNormTP",
q: torch.Tensor,
k: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
# NPU fast path: kernelized local RMSNorm for q/k, then TP-global rstd correction.
if current_platform.device_name == "npu":
q, q_inv_rms = torch.ops.npu.npu_rms_norm(q, q_norm.weight, q_norm.variance_epsilon)
k, k_inv_rms = torch.ops.npu.npu_rms_norm(k, k_norm.weight, k_norm.variance_epsilon)
if q_norm.tp_world > 1:
q_local_inv_rms = q_inv_rms.to(torch.float32)
if q_local_inv_rms.shape[-1] != 1:
q_local_inv_rms = q_local_inv_rms.mean(dim=-1, keepdim=True)
q_local_var = (q_local_inv_rms.reciprocal().pow(2) - q_norm.variance_epsilon).clamp_min_(0.0)
k_local_inv_rms = k_inv_rms.to(torch.float32)
if k_local_inv_rms.shape[-1] != 1:
k_local_inv_rms = k_local_inv_rms.mean(dim=-1, keepdim=True)
k_local_var = (k_local_inv_rms.reciprocal().pow(2) - k_norm.variance_epsilon).clamp_min_(0.0)
qk_var = torch.cat([q_local_var, k_local_var], dim=-1)
qk_var = tensor_model_parallel_all_reduce(qk_var) / q_norm.tp_world
q_global_var, k_global_var = qk_var.chunk(2, dim=-1)
q_local_rstd = torch.rsqrt(q_local_var + q_norm.variance_epsilon)
k_local_rstd = torch.rsqrt(k_local_var + k_norm.variance_epsilon)
q_global_rstd = torch.rsqrt(q_global_var + q_norm.variance_epsilon)
k_global_rstd = torch.rsqrt(k_global_var + k_norm.variance_epsilon)
q = q * (q_global_rstd / q_local_rstd).to(q.dtype)
k = k * (k_global_rstd / k_local_rstd).to(k.dtype)
return q, k
assert _original_qk_method is not None
# We install the patch as a staticmethod below, so prefer the static calling
# convention for the original as well.
return _original_qk_method(q_norm, k_norm, q, k)
def _patched_weight_loader(
param: nn.Parameter,
loaded_weight: torch.Tensor,
shard_world_size: int | None = None,
shard_rank: int | None = None,
) -> None:
if shard_world_size is None:
shard_world_size = get_tensor_model_parallel_world_size()
if shard_rank is None:
shard_rank = get_tensor_model_parallel_rank()
shard_size = loaded_weight.shape[0] // shard_world_size
shard = slice(shard_rank * shard_size, (shard_rank + 1) * shard_size)
param.data.copy_(loaded_weight[shard])
def _patched_init(
self: "MiniMaxText01RMSNormTP",
hidden_size: int,
eps: float = 1e-6,
*,
weight_shard_world_size: int | None = None,
weight_shard_rank: int | None = None,
) -> None:
CustomOp.__init__(self)
self.tp_world = get_tensor_model_parallel_world_size()
self.tp_rank = get_tensor_model_parallel_rank()
self.weight_shard_world = weight_shard_world_size or self.tp_world
self.weight_shard_rank = self.tp_rank if weight_shard_rank is None else weight_shard_rank
if hidden_size % self.weight_shard_world != 0:
raise ValueError(
"MiniMaxText01RMSNormTP hidden_size must be divisible by "
f"weight_shard_world_size, got hidden_size={hidden_size}, "
f"weight_shard_world_size={self.weight_shard_world}"
)
self.weight = nn.Parameter(torch.ones(int(hidden_size / self.weight_shard_world)))
self.weight.weight_loader = partial(
_patched_weight_loader,
shard_world_size=self.weight_shard_world,
shard_rank=self.weight_shard_rank,
)
self.variance_epsilon = eps
MiniMaxText01RMSNormTP.__init__ = _patched_init
MiniMaxText01RMSNormTP.weight_loader = staticmethod(_patched_weight_loader)
if _ORIG_QK_METHOD_NAME is not None:
# Force staticmethod style, as requested.
setattr(MiniMaxText01RMSNormTP, _ORIG_QK_METHOD_NAME, staticmethod(_patched_qk))

View File

@@ -0,0 +1,129 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
import importlib
import sys
import torch
from torch._subclasses.fake_tensor import FakeTensor
try:
import npugraph_ex as nge
from npugraph_ex.core._concrete_graph import _is_symlist
from npugraph_ex.npu_fx_compiler import _unpack_meta_list
_USE_NPUGRAPH_EX = True
except ImportError:
import torchair as nge
from torchair.core._concrete_graph import _is_symlist
from torchair.npu_fx_compiler import _unpack_meta_list
_USE_NPUGRAPH_EX = False
class ValuePack:
def __init__(self, meta, npu_meta=None) -> None:
self._meta = meta
self._npu_meta = meta if npu_meta is None else npu_meta
@property
def meta(self):
return self._meta
@property
def npu(self):
return self._npu_meta
def __getitem__(self, key):
if isinstance(self._meta, dict):
return self._meta.get(key)
raise ValueError(f"Unsupported meta type for ValuePack __getitem__, key:{key}, type: {type(self._meta)}")
def __repr__(self) -> str:
if isinstance(self._meta, FakeTensor):
meta_str = f"FakeTensor(dtype={self._meta.dtype}, size={list(self._meta.size())}"
elif isinstance(self._meta, torch.Tensor):
meta_str = f"torch.Tensor(dtype={self._meta.dtype}, size={list(self._meta.size())}"
elif isinstance(self._meta, torch.SymInt):
meta_str = f"torch.SymInt({self._meta})"
else:
try:
meta_str = f"{type(self._meta)}({self._meta})"
except Exception:
meta_str = f"{type(self._meta)}"
return f"Pack(meta:{meta_str} npu:{self._npu_meta})"
def _unpack_meta(args, kwargs):
unpacked_args = []
unpacked_kwargs = {}
def _get_meta_part(arg):
if isinstance(arg, (list, tuple)) and any(isinstance(v, ValuePack) for v in arg):
return _unpack_meta_list(arg)
elif isinstance(arg, dict):
return {k: v.meta if isinstance(v, ValuePack) else v for k, v in arg.items()}
elif isinstance(arg, ValuePack):
return arg.meta
else:
return arg
for arg in args:
unpacked_args.append(_get_meta_part(arg))
for key, value in kwargs.items():
unpacked_kwargs[key] = _get_meta_part(value)
return list(unpacked_args), unpacked_kwargs
def _unpack_npu(self, args, kwargs):
unpacked = []
unpacked_kwargs = {}
def _get_npu_part(arg):
if isinstance(arg, (list, tuple)) and len(arg):
if _is_symlist(arg):
arg = self._graph.parse_symlist(arg)
else:
arg = [(v.npu if isinstance(v, ValuePack) else v) for v in arg]
return arg
elif isinstance(arg, dict):
return {k: v.npu if isinstance(v, ValuePack) else v for k, v in arg.items()}
elif isinstance(arg, ValuePack):
return arg.npu
else:
return arg
for arg in args:
unpacked.append(_get_npu_part(arg))
for key, value in kwargs.items():
unpacked_kwargs[key] = _get_npu_part(value)
return unpacked, unpacked_kwargs
nge.core._concrete_graph.ValuePack = ValuePack
# The ValuePack class is referenced in the npu_fx_compiler module (and fx_summary for torchair),
# and after the patch, these modules need to be reloaded.
if not _USE_NPUGRAPH_EX:
importlib.reload(sys.modules["torchair.fx_summary"])
pkg_prefix = "npugraph_ex" if _USE_NPUGRAPH_EX else "torchair"
importlib.reload(sys.modules[f"{pkg_prefix}.npu_fx_compiler"])
nge.npu_fx_compiler._unpack_meta = _unpack_meta
nge.npu_fx_compiler._NpuGraphConverter._unpack_npu = _unpack_npu

View File

@@ -0,0 +1,65 @@
import sys
import torch
from torch import nn
from vllm.config import ModelConfig
from vllm.model_executor.layers.attention import (
Attention,
MLAAttention,
MMEncoderAttention,
)
from vllm.model_executor.layers.quantization.base_config import (
QuantizeMethodBase,
)
from vllm.model_executor.model_loader import base_loader, utils
from vllm.model_executor.model_loader.reload import set_torchao_reload_attrs
from vllm.model_executor.model_loader.utils import device_loading_context
def _is_dsa_attention(module: nn.Module) -> bool:
module_cls = type(module)
return module_cls.__module__ == "vllm_ascend.models.layer.attention.layer" and module_cls.__name__ == "DSAAttention"
def ascend_process_weights_after_loading(
model: nn.Module, model_config: ModelConfig, target_device: torch.device
) -> None:
for _, module in model.named_modules():
quant_method = getattr(module, "quant_method", None)
if isinstance(quant_method, QuantizeMethodBase):
# When quant methods need to process weights after loading
# (for repacking, quantizing, etc), they expect parameters
# to be on the global target device. This scope is for the
# case where cpu offloading is used, where we will move the
# parameters onto device for processing and back off after.
with device_loading_context(module, target_device):
quant_method.process_weights_after_loading(module)
# Initialize post-load attention weights for Attention, MLA, and MM encoder.
# NOTE: Happens after other modules so we can easily decompress weights.
for _, module in model.named_modules():
if (isinstance(module, (Attention, MLAAttention, MMEncoderAttention)) or _is_dsa_attention(module)) and hasattr(
module, "process_weights_after_loading"
):
# TODO(lucas): see if there is a way to unify the signatures
# of process_weights_after_loading
with device_loading_context(module, target_device):
module.process_weights_after_loading(model_config.dtype)
# Needed for torchao model reloading via model.reload_weights
# @kylesayrs @jerryzh168 this can be removed if callers move to `reload_weights`
if model_config.quantization == "torchao":
set_torchao_reload_attrs(model, model_config)
utils.process_weights_after_loading = ascend_process_weights_after_loading
base_loader.process_weights_after_loading = ascend_process_weights_after_loading
vllm_ascend_loaders = [
"vllm_ascend.model_loader.netloader.netloader",
"vllm_ascend.model_loader.rfork.rfork_loader",
]
for loader_module in vllm_ascend_loaders:
loader = sys.modules.get(loader_module)
if loader is not None:
loader.__dict__["process_weights_after_loading"] = ascend_process_weights_after_loading

View File

@@ -0,0 +1,209 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
# from collections.abc import Iterable
# mypy: ignore-errors
import torch
from vllm.distributed import get_tensor_model_parallel_world_size
from vllm.distributed.parallel_state import get_pp_group
from vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn import QwenGatedDeltaNetAttention as _GDNBaseCls
from vllm.model_executor.models.qwen3_5 import Qwen3_5DecoderLayer
try:
from vllm.model_executor.models.qwen3_5_mtp import Qwen3_5MultiTokenPredictor
from vllm.sequence import IntermediateTensors
except ImportError:
Qwen3_5MultiTokenPredictor = None
IntermediateTensors = None
from vllm.model_executor.models.qwen3_next import Qwen3NextAttention
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.ops.gdn import AscendGatedDeltaNetAttention
from vllm_ascend.utils import is_310p
_GDN_PATCH_TARGET = _GDNBaseCls
class AscendQwen3NextAttention(Qwen3NextAttention):
def forward(self, positions: torch.Tensor, output: torch.Tensor, hidden_states: torch.Tensor):
qkv, _ = self.qkv_proj(hidden_states)
if "qwen3_5" in self.config.model_type:
cos_sin = self.rotary_emb.cos_sin_cache[positions]
if cos_sin.device != qkv.device:
cos_sin = cos_sin.to(qkv.device)
if cos_sin.dtype != qkv.dtype:
cos_sin = cos_sin.to(qkv.dtype)
q, k, v, gate = torch.ops.vllm.triton_split_qkv_rmsnorm_mrope(
qkv=qkv,
q_weight=1.0 + self.q_norm.weight,
k_weight=1.0 + self.k_norm.weight,
cos_sin=cos_sin,
num_q_heads=self.num_heads,
num_kv_heads=self.num_kv_heads,
head_size=self.head_dim,
eps=self.config.rms_norm_eps,
mrope_section=self.rotary_emb.mrope_section,
is_interleaved=self.rotary_emb.mrope_interleaved,
rope_dim=self.rotary_emb.rotary_dim,
has_gate=self.attn_output_gate,
)
else:
if self.attn_output_gate:
q_gate, k, v = qkv.split([self.q_size * 2, self.kv_size, self.kv_size], dim=-1)
orig_shape = q_gate.shape[:-1]
q_gate = q_gate.view(*orig_shape, self.num_heads, -1)
q, gate = torch.chunk(q_gate, 2, dim=-1)
q = q.reshape(*orig_shape, -1)
gate = gate.reshape(*orig_shape, -1)
else:
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q = self.q_norm(q.view(-1, self.num_heads, self.head_dim)).view(-1, self.num_heads * self.head_dim)
k = self.k_norm(k.view(-1, self.num_kv_heads, self.head_dim)).view(-1, self.num_kv_heads * self.head_dim)
q, k = self.rotary_emb(positions, q, k)
attn_output = self.attn(q, k, v)
if self.attn_output_gate:
gate = torch.sigmoid(gate)
attn_output = attn_output * gate
output[:], _ = self.o_proj(attn_output)
class AscendQwen3_5DecoderLayer(Qwen3_5DecoderLayer):
def forward(
self,
hidden_states: torch.Tensor,
residual: torch.Tensor | None,
positions: torch.Tensor = None,
**kwargs: object,
):
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(hidden_states, residual)
if self.layer_idx == 0 and _EXTRA_CTX.flash_comm_v1_enabled:
tp_size = get_tensor_model_parallel_world_size()
n_out = (hidden_states.shape[0] + tp_size - 1) // tp_size
hidden_dim = hidden_states.shape[-1]
self_attention_output = torch.empty(
(n_out, hidden_dim), dtype=hidden_states.dtype, device=hidden_states.device
)
else:
self_attention_output = torch.empty_like(hidden_states)
if self.layer_type == "linear_attention":
self.linear_attn(
hidden_states=hidden_states,
output=self_attention_output,
)
elif self.layer_type == "full_attention":
self.self_attn(
hidden_states=hidden_states,
output=self_attention_output,
positions=positions,
)
else:
raise ValueError("Invalid layer_type")
hidden_states = self_attention_output
if self.layer_scale:
if len(hidden_states.shape) == 2:
hidden_states = hidden_states * (self.attn_layer_scale.to(hidden_states.dtype)[0] + 1)
else:
hidden_states = hidden_states * (self.attn_layer_scale.to(hidden_states.dtype) + 1)
# Fully Connected
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
hidden_states = self.mlp(hidden_states)
if self.layer_scale:
if len(hidden_states.shape) == 2:
hidden_states = hidden_states * (self.ffn_layer_scale.to(hidden_states.dtype)[0] + 1)
else:
assert len(hidden_states.shape) == len(self.ffn_layer_scale.shape), (
f"shape must be the same {len(hidden_states.shape)}, {len(self.ffn_layer_scale.shape)}"
)
hidden_states = hidden_states * (self.ffn_layer_scale.to(hidden_states.dtype) + 1)
return hidden_states, residual
if Qwen3_5MultiTokenPredictor is not None:
def qwen3_5_mtp_forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
hidden_states: torch.Tensor,
intermediate_tensors: IntermediateTensors | None = None,
inputs_embeds: torch.Tensor | None = None,
spec_step_idx: int = 0,
) -> torch.Tensor:
# Backport upstream Qwen3.5 MTP behavior: the local drafter runs on the
# last PP stage and should always combine token embeddings with the
# target hidden states instead of consuming PP intermediate tensors.
if inputs_embeds is None:
inputs_embeds = self.embed_input_ids(input_ids)
assert hidden_states.shape[-1] == inputs_embeds.shape[-1]
inputs_embeds = self.pre_fc_norm_embedding(inputs_embeds)
hidden_states = self.pre_fc_norm_hidden(hidden_states)
hidden_states = torch.cat([inputs_embeds, hidden_states], dim=-1)
hidden_states = self.fc(hidden_states)
residual = None
current_step_idx = spec_step_idx % self.num_mtp_layers
hidden_states, residual = self.layers[current_step_idx](
positions=positions,
hidden_states=hidden_states,
residual=residual,
)
if not get_pp_group().is_last_rank:
return IntermediateTensors(
{
"hidden_states": hidden_states,
"residual": residual,
}
)
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
Qwen3_5MultiTokenPredictor.forward = qwen3_5_mtp_forward
Qwen3_5DecoderLayer.forward = AscendQwen3_5DecoderLayer.forward
Qwen3NextAttention.forward = AscendQwen3NextAttention.forward
_GDN_PATCH_TARGET._split_ba_for_tp = AscendGatedDeltaNetAttention._split_ba_for_tp
_GDN_PATCH_TARGET.get_state_shape = AscendGatedDeltaNetAttention.get_state_shape
_GDN_PATCH_TARGET.get_attn_backend = AscendGatedDeltaNetAttention.get_attn_backend
if is_310p():
from vllm_ascend._310p.ops.fla.gdn_310 import AscendGatedDeltaNetAttention310
_GDN_PATCH_TARGET._forward_core = AscendGatedDeltaNetAttention310._forward_core
_GDN_PATCH_TARGET.get_state_dtype = AscendGatedDeltaNetAttention310.get_state_dtype
else:
_GDN_PATCH_TARGET.forward = AscendGatedDeltaNetAttention.forward
_GDN_PATCH_TARGET._forward_core = AscendGatedDeltaNetAttention._forward_core
_GDN_PATCH_TARGET._warmup_prefill_kernels = AscendGatedDeltaNetAttention._warmup_prefill_kernels

View File

@@ -0,0 +1,62 @@
import torch
import torch.nn.functional as F
from vllm.model_executor.models.qwen3_dflash import DFlashQwen3Model
def precompute_and_store_context_kv(
self,
context_states: torch.Tensor,
context_positions: torch.Tensor,
context_slot_mapping: torch.Tensor | None = None,
) -> None:
if not hasattr(self, "_num_attn_layers"):
self._build_fused_kv_buffers()
num_ctx = context_states.shape[0]
L = self._num_attn_layers
kv = self._kv_size
hd = self._head_dim
nkv = self._num_kv_heads
# --- Fused KV projection (one GEMM for all layers) ---
normed_context_states = self.hidden_norm(context_states)
all_kv_flat = F.linear(normed_context_states, self._fused_kv_weight, self._fused_kv_bias)
# Single contiguous copy that separates K/V and transposes to
# layer-major layout. Result: [2, L, num_ctx, nkv, hd] contiguous.
# Indexing dim-0 gives contiguous [L, num_ctx, nkv, hd] for K and V.
all_kv = all_kv_flat.view(num_ctx, L, 2, nkv, hd).permute(2, 1, 0, 3, 4).contiguous()
all_k = all_kv[0] # [L, num_ctx, nkv, hd], contiguous
all_v = all_kv[1] # [L, num_ctx, nkv, hd], contiguous
# --- Per-layer RMSNorm K (3D: [num_ctx, nkv, hd] per layer) ---
all_k_normed = torch.empty_like(all_k)
for i in range(L):
k_norm_layer = self.layers[i].self_attn.k_norm
all_k_normed[i] = k_norm_layer(all_k[i])
# --- Fused RoPE across all layers ---
# View as [L * num_ctx, kv] so RoPE sees one big batch (no copy).
# In-place RoPE: pass K as the "query" arg with key=None.
all_k_flat = all_k_normed.view(L * num_ctx, kv)
positions_repeated = context_positions.repeat(L)
tmpv = all_k_flat.clone()
self.layers[0].self_attn.rotary_emb(positions_repeated, all_k_flat, tmpv)
if context_slot_mapping is None:
return
# --- Per-layer cache insert ---
all_k_final = all_k_flat.view(L, num_ctx, nkv, hd)
for i in range(L):
attn = self._attn_layers[i]
kv_cache = attn.kv_cache
attn.impl.do_kv_cache_update(
attn,
all_k_final[i],
all_v[i],
kv_cache,
context_slot_mapping,
)
DFlashQwen3Model.precompute_and_store_context_kv = precompute_and_store_context_kv

View File

@@ -0,0 +1,50 @@
import torch
import vllm.v1.worker.utils as utils
from vllm.model_executor.layers.attention import Attention
from vllm.v1.worker.utils import defaultdict, extract_layer_index
# Without this patch, it will raise an exception when initialize kv_cache.
# TODO To remove the patch, we need check why the original bind_kv_cache raises an NotImplementedError.
def bind_kv_cache(
kv_caches: dict[str, torch.Tensor],
forward_context: dict[str, Attention],
runner_kv_caches: list[torch.Tensor],
num_attn_module: int = 1,
) -> None:
"""
Bind the allocated KV cache to both ModelRunner and forward context so
that the KV cache can be used in the forward pass.
This function:
1) Fills the ModelRunner's kv cache list (`runner_kv_caches`) with
kv_caches.
2) Associates each attention layer in the `forward_context` with its
corresponding KV cache in kv_caches.
Args:
kv_caches: The allocated kv_caches with layer names as keys.
forward_context: The global forward context containing all Attention
layers with layer names as keys.
runner_kv_caches: The kv_cache declared by ModelRunner.
"""
# Bind kv_caches to ModelRunner
assert len(runner_kv_caches) == 0
# Convert kv_caches dict to a list of tensors in the order of layer_index.
index2name = defaultdict(list)
for layer_name in kv_caches:
index2name[extract_layer_index(layer_name, num_attn_module)].append(layer_name)
for layer_index in sorted(index2name.keys()):
layer_names = index2name[layer_index]
# remove some codes for the typical case of encoder-decoder model, e.g., bart.
layer_name = layer_names[0]
runner_kv_caches.append(kv_caches[layer_name])
# Bind kv_caches to forward context
for layer_name, kv_cache in kv_caches.items():
forward_context[layer_name].kv_cache = kv_cache
utils.bind_kv_cache = bind_kv_cache

View File

@@ -0,0 +1,110 @@
import torch
from vllm.distributed import get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size
from vllm.model_executor.models.qwen3 import Qwen3Attention
from vllm.model_executor.models.qwen3_moe import Qwen3MoeAttention
from vllm.model_executor.models.qwen3_vl import (
Qwen3_VisionTransformer,
Qwen3VLForConditionalGeneration,
pos_embed_interpolate_native,
)
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.ops.rotary_embedding import AscendMRotaryEmbedding
def tensor_parallel_wrap(func):
def wrap(*args, **kwargs):
deepstack_input_embeds = func(*args, **kwargs)
if deepstack_input_embeds is None:
return deepstack_input_embeds
try:
flash_comm_v1_enabled = _EXTRA_CTX.flash_comm_v1_enabled
except (AssertionError, AttributeError, KeyError):
flash_comm_v1_enabled = False
if flash_comm_v1_enabled:
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
deepstack_input_embeds.tensors = {
k: v.chunk(tp_size)[tp_rank] for k, v in deepstack_input_embeds.tensors.items()
}
return deepstack_input_embeds
return wrap
def forward_with_split_qkv_rmsnorm_mrope(self, positions: torch.Tensor, hidden_states: torch.Tensor):
qkv, _ = self.qkv_proj(hidden_states)
if isinstance(self.rotary_emb, AscendMRotaryEmbedding):
cos_sin = self.rotary_emb.cos_sin_cache[positions]
if cos_sin.device != qkv.device:
cos_sin = cos_sin.to(qkv.device)
if cos_sin.dtype != qkv.dtype:
cos_sin = cos_sin.to(qkv.dtype)
q, k, v, _ = torch.ops.vllm.triton_split_qkv_rmsnorm_mrope(
qkv=qkv,
q_weight=self.q_norm.weight,
k_weight=self.k_norm.weight,
cos_sin=cos_sin,
num_q_heads=self.num_heads,
num_kv_heads=self.num_kv_heads,
head_size=self.head_dim,
eps=self.q_norm.variance_epsilon,
mrope_section=self.rotary_emb.mrope_section,
is_interleaved=self.rotary_emb.mrope_interleaved,
rope_dim=self.rotary_emb.rotary_dim,
)
else:
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q_by_head = q.view(*q.shape[:-1], q.shape[-1] // self.head_dim, self.head_dim)
q_by_head = self.q_norm(q_by_head)
q = q_by_head.view(q.shape)
k_by_head = k.view(*k.shape[:-1], k.shape[-1] // self.head_dim, self.head_dim)
k_by_head = self.k_norm(k_by_head)
k = k_by_head.view(k.shape)
q, k = self.rotary_emb(positions, q, k)
attn_output = self.attn(q, k, v)
output, _ = self.o_proj(attn_output)
return output
Qwen3Attention.forward = forward_with_split_qkv_rmsnorm_mrope
Qwen3MoeAttention.forward = forward_with_split_qkv_rmsnorm_mrope
Qwen3VLForConditionalGeneration._get_deepstack_input_embeds = tensor_parallel_wrap(
Qwen3VLForConditionalGeneration._get_deepstack_input_embeds
)
def _fast_pos_embed_interpolate(self, grid_thw: list[list[int]]) -> torch.Tensor:
outputs = []
for t, h, w in grid_thw:
outputs.append(
pos_embed_interpolate_native(
self.pos_embed.weight,
t,
h,
w,
self.num_grid_per_side,
self.spatial_merge_size,
self.dtype,
)
)
return torch.cat(outputs, dim=0)
Qwen3_VisionTransformer.fast_pos_embed_interpolate = _fast_pos_embed_interpolate
def patch_qwen3_vl_moe_pp_layer_range():
try:
from vllm.model_executor.models.qwen3_vl_moe import Qwen3MoeLLMForCausalLM
except Exception:
return
if not hasattr(Qwen3MoeLLMForCausalLM, "start_layer"):
Qwen3MoeLLMForCausalLM.start_layer = property(lambda self: self.model.start_layer)
if not hasattr(Qwen3MoeLLMForCausalLM, "end_layer"):
Qwen3MoeLLMForCausalLM.end_layer = property(lambda self: self.model.end_layer)
patch_qwen3_vl_moe_pp_layer_range()

View File

@@ -0,0 +1,9 @@
import vllm.v1.sample.rejection_sampler as rs
from vllm_ascend.sample.rejection_sampler import apply_sampling_constraints, expand_batch_to_tokens, rejection_sample
# TODO: delete this patch after apply_sampling_constraints and rejection_sample
# are extracted to as class func of RejectionSampler
rs.apply_sampling_constraints = apply_sampling_constraints
rs.rejection_sample = rejection_sample
rs.expand_batch_to_tokens = expand_batch_to_tokens

View File

@@ -0,0 +1,188 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/fused_moe/routed_experts_capturer.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
from __future__ import annotations
import logging
import torch
import torch.distributed as dist
from vllm.distributed.parallel_state import get_tp_group
from vllm.forward_context import get_forward_context
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import RoutedExpertsCapturer
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType
logger = logging.getLogger(__name__)
def capture(self, layer_id: int, topk_ids: torch.Tensor) -> None:
"""Capture expert routing decisions for a specific layer.
Under data parallelism, ``topk_ids`` may have four different batch
layouts depending on where the DP combine happens and whether
Sequence Parallelism (SP) is active for the MoE layer:
- ``n == total`` (naive dispatch): all DP ranks' tokens are
concatenated before routing; we slice out this rank's span
using the cumulative per-rank counts.
- ``n == token_num_per_dp`` (modular-kernel path): DP combine
happens inside ``quant_method.apply``; ``select_experts`` only
ever sees this rank's tokens, so we take the whole tensor.
- ``n == total_with_padding`` (padded all-gather path): tokens are
padded to max_tokens before all-gather across DP group; each
DP rank occupies a contiguous block of size max_tokens, and we
extract only the actual tokens for this rank (skip padding).
When all DP ranks have equal token counts, ``total == total_with_padding``,
so the naive dispatch branch fires instead (equivalent result).
- ``n == ceil(token_num_per_dp / tp_size)`` (SP + modular-kernel
path): tokens were split along dim=0 across the TP group by
``_sequence_parallel_context``
(``moe_runner_base.py:_sequence_parallel_context``), so each
TP rank only sees its shard. We all-gather along dim=0 to
reconstruct this DP rank's full routing tensor. SP pads with
ceil-div (see ``_compute_sp_num_tokens`` in
``forward_context.py``), so the gathered tensor may contain a
few trailing padding rows which are trimmed by the downstream
``[:token_num_per_dp]`` slice.
Args:
layer_id: The layer index.
topk_ids: Tensor of shape (batch_size, num_routed_experts).
"""
ctx = get_forward_context()
if ctx.dp_metadata is None: # single dp
start_loc = 0
end_loc = topk_ids.shape[0]
token_num_per_dp = topk_ids.shape[0]
else: # multi dp
num_tokens_dp = ctx.dp_metadata.num_tokens_across_dp_cpu
token_num_per_dp = int(num_tokens_dp[self.dp_rank].item())
total = int(num_tokens_dp.sum().item())
n = topk_ids.shape[0]
# Calculate total with padding for all-gather scenario.
# When tokens are padded to max_tokens before all-gather across DP group,
# the total size becomes max_tokens * dp_size.
# Example: DP0 has 5 tokens, DP1 has 7 tokens, max_tokens=7.
# After padding: DP0 has 7 tokens, DP1 has 7 tokens.
# After all-gather: total_with_padding = 7 * 2 = 14.
max_tokens = int(num_tokens_dp.max().item())
total_with_padding = max_tokens * len(num_tokens_dp)
if n == total:
# Naive dispatch: all DP ranks' tokens concatenated
# before routing. This rank owns tokens
# [end_loc - token_num_per_dp, end_loc).
cumsum = torch.cumsum(num_tokens_dp, dim=0)
end_loc = int(cumsum[self.dp_rank].item())
start_loc = end_loc - token_num_per_dp
elif n == token_num_per_dp:
# Modular-kernel path: DP combine happens inside
# quant_method.apply; select_experts only sees this
# rank's tokens, take the whole tensor.
start_loc = 0
end_loc = token_num_per_dp
elif n == total_with_padding:
# NOTE(Ronald1995): When all DP ranks have equal token counts,
# total == total_with_padding, so the first branch (n == total)
# fires instead. This overlap is intentional since both branches
# produce equivalent results in that case.
# Padded all-gather path: tokens are padded to max_tokens before
# all-gather across DP group. Each DP rank occupies a contiguous
# block of size max_tokens. Extract only the actual tokens for
# this rank (skip padding).
# Example: dp_rank=0, max_tokens=7, token_num_per_dp=5.
# start_loc = 0 * 7 = 0
# end_loc = 0 + 5 = 5 (only first 5 tokens are valid)
start_loc = self.dp_rank * max_tokens
end_loc = start_loc + token_num_per_dp
elif (
self.tp_size > 1
and n != token_num_per_dp
and (
# all2all scenario use tensor split, different tp rank have different
# size of tokens.
n == (token_num_per_dp + self.tp_size - 1) // self.tp_size
or n == token_num_per_dp // self.tp_size
# mc2 scenario will pad dp tokens to max_tokens and then ceil-div.
or n == (max_tokens + self.tp_size - 1) // self.tp_size
)
):
# SP + modular-kernel path. All-gather across the TP
# group along dim=0 to reconstruct the full per-DP-rank
# tensor; keep only the first ``token_num_per_dp`` rows
# (trailing rows are SP ceil-div padding). The TP group
# is always initialized on real rollout workers, and
# every rank in the group reaches this branch in
# lockstep (bind is per-FusedMoE layer, SP is a global
# condition), so a bare all_gather here will not
# deadlock -- let it raise if the precondition is
# violated rather than skip silently.
#
# ``topk_ids`` is already whatever the router produced
# (typically int32/int64, both supported by NCCL); the
# downstream ``device_buffer[...] = topk_ids[...]``
# setitem narrows into int32 automatically.
# NOTE(Ronald1995): if total_num_per_dp == max_tokens,
# it will be both all2all and mc2 scenario.
# but we fires all2all scenario first.
# the result will be the same.
# all2all scenario in vllm-ascend.
if _EXTRA_CTX.moe_comm_type == MoECommType.ALLTOALL:
gather_topk_ids_shape = (
(token_num_per_dp, topk_ids.shape[1])
if token_num_per_dp >= self.tp_size
else (self.tp_size, topk_ids.shape[1])
)
# mc2 scenario in vllm-ascend
else:
gather_topk_ids_shape = (n * self.tp_size, topk_ids.shape[1])
gather_topk_ids = torch.empty(
gather_topk_ids_shape,
dtype=topk_ids.dtype,
device=topk_ids.device,
)
split_topk_ids = torch.tensor_split(gather_topk_ids, self.tp_size, dim=0)
dist.all_gather(list(split_topk_ids), topk_ids, get_tp_group().device_group)
topk_ids = gather_topk_ids
start_loc = 0
end_loc = token_num_per_dp
else:
sp_expected = (token_num_per_dp + self.tp_size - 1) // self.tp_size if self.tp_size > 0 else -1
raise AssertionError(
"RoutedExpertsCapturer: unexpected topk_ids batch "
f"dim {n} (expected {total}, {token_num_per_dp}, "
f"{total_with_padding}, or {sp_expected} for "
f"dp_rank={self.dp_rank}, tp_size={self.tp_size})"
)
# Defensive: model may expose more layers than the capture buffer
# was sized for (unusual, but guards against miss-config).
if layer_id >= self.device_buffer.shape[1]:
return
self.device_buffer[:token_num_per_dp, layer_id, :] = topk_ids[start_loc:end_loc, :]
RoutedExpertsCapturer.capture = capture

View File

@@ -0,0 +1,126 @@
import vllm.model_executor.layers.fla.ops
import vllm.model_executor.layers.mamba.ops.causal_conv1d
import vllm.v1.worker.gpu.sample.gumbel
from vllm.triton_utils import HAS_TRITON, triton
from vllm.utils.math_utils import next_power_of_2
from vllm_ascend.ops.triton.fla.chunk import chunk_gated_delta_rule
from vllm_ascend.ops.triton.fla.layernorm_guard import LayerNormFn
from vllm_ascend.ops.triton.fla.sigmoid_gating import fused_recurrent_gated_delta_rule_fwd_kernel
from vllm_ascend.ops.triton.mamba.causal_conv1d import causal_conv1d_update_npu
triton.next_power_of_2 = next_power_of_2
vllm.model_executor.layers.mamba.ops.causal_conv1d.causal_conv1d_update = causal_conv1d_update_npu
vllm.model_executor.layers.fla.ops.fused_recurrent.fused_recurrent_gated_delta_rule_fwd_kernel = (
fused_recurrent_gated_delta_rule_fwd_kernel
)
vllm.model_executor.layers.fla.ops.layernorm_guard.LayerNormFn = LayerNormFn
vllm.model_executor.layers.fla.ops.chunk_gated_delta_rule = chunk_gated_delta_rule
# On NPU platforms without an active Triton backend (e.g. 310P), replace the
# Triton-based fused_post_conv_prep with a pure-PyTorch fallback so that
# qwen_gdn_linear_attn's from-import picks up the replacement before model
# load.
if not HAS_TRITON:
import torch
import torch.nn.functional as _F
def _fused_post_conv_prep_pytorch(
conv_output,
a,
b,
A_log,
dt_bias,
num_k_heads,
head_k_dim,
head_v_dim,
apply_l2norm=True,
output_g_exp=False,
):
L = conv_output.shape[0]
H, K, V = num_k_heads, head_k_dim, head_v_dim
HV = A_log.shape[0]
q = conv_output[:, : H * K].reshape(L, H, K)
k = conv_output[:, H * K : 2 * H * K].reshape(L, H, K)
v = conv_output[:, 2 * H * K :].reshape(L, HV, V)
if apply_l2norm:
# x / sqrt(sum(x^2) + eps) — matches Triton kernel, in fp32
def _l2norm(t):
t_f = t.float()
return (t_f / torch.sqrt((t_f * t_f).sum(-1, keepdim=True) + 1e-6)).to(t.dtype)
q, k = _l2norm(q), _l2norm(k)
q, k, v = q.contiguous(), k.contiguous(), v.contiguous()
x = (a + dt_bias.unsqueeze(0)).float()
g = -torch.exp(A_log.float().unsqueeze(0)) * _F.softplus(x)
if output_g_exp:
g = torch.exp(g)
return q, k, v, g, torch.sigmoid(b.float())
vllm.model_executor.layers.fla.ops.fused_post_conv_prep = _fused_post_conv_prep_pytorch
def _fused_recurrent_packed_decode_pytorch(
mixed_qkv,
a,
b,
A_log,
dt_bias,
scale,
initial_state,
out,
ssm_state_indices,
use_qk_l2norm_in_kernel=False,
):
B = mixed_qkv.shape[0]
HV, V, K = initial_state.shape[-3:]
H = (mixed_qkv.shape[1] - HV * V) // (2 * K)
ratio = HV // H
q = mixed_qkv[:, : H * K].reshape(B, H, K)
k = mixed_qkv[:, H * K : 2 * H * K].reshape(B, H, K)
v = mixed_qkv[:, 2 * H * K :].reshape(B, HV, V)
SOFTPLUS_THRESHOLD = 20.0
x = (a + dt_bias.unsqueeze(0)).float()
softplus_x = torch.where(x <= SOFTPLUS_THRESHOLD, torch.log1p(torch.exp(x)), x)
g = -torch.exp(A_log.float().unsqueeze(0)) * softplus_x # [B, HV]
beta = torch.sigmoid(b.float()) # [B, HV]
for n in range(B):
state_idx = int(ssm_state_indices[n].item())
if state_idx <= 0:
out[n, 0] = 0
continue
h = initial_state[state_idx].float() # [HV, V, K]
q_n = q[n].float().repeat_interleave(ratio, dim=0) # [HV, K]
k_n = k[n].float().repeat_interleave(ratio, dim=0) # [HV, K]
v_n = v[n].float() # [HV, V]
if use_qk_l2norm_in_kernel:
def _l2norm(t):
t_f = t.float()
return t_f / torch.sqrt((t_f * t_f).sum(-1, keepdim=True) + 1e-6)
q_n, k_n = _l2norm(q_n), _l2norm(k_n)
q_n = q_n * scale
h = h * torch.exp(g[n]).view(HV, 1, 1)
v_n = v_n - torch.einsum("hvk,hk->hv", h, k_n)
v_n = v_n * beta[n].view(HV, 1)
h = h + torch.einsum("hv,hk->hvk", v_n, k_n)
out[n, 0] = torch.einsum("hvk,hk->hv", h, q_n).to(out.dtype)
initial_state[state_idx] = h.to(initial_state.dtype)
return out, initial_state
vllm.model_executor.layers.fla.ops.fused_recurrent.fused_recurrent_gated_delta_rule_packed_decode = (
_fused_recurrent_packed_decode_pytorch
)

View File

@@ -0,0 +1,11 @@
import vllm
from vllm_ascend.worker.v2.attn_utils import (
_allocate_kv_cache,
_reshape_kv_cache_v2,
get_kv_cache_spec,
)
vllm.v1.worker.gpu.attn_utils._allocate_kv_cache = _allocate_kv_cache
vllm.v1.worker.gpu.attn_utils._reshape_kv_cache = _reshape_kv_cache_v2
vllm.v1.worker.gpu.model_runner.get_kv_cache_spec = get_kv_cache_spec

View File

@@ -0,0 +1,25 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/block_table.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
from vllm.v1.worker.gpu import model_runner
from vllm_ascend.worker.v2.block_table import AscendBlockTables
# vllm-ascend need to initialize slot mapping as torch.int32 dtype,
# but vllm default is torch.int64 dtype.
model_runner.BlockTables = AscendBlockTables

View File

@@ -0,0 +1,27 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/input_batch.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
# 显式导入模块,确保模块被加载后再进行 patch
from vllm.v1.worker.gpu import cudagraph_utils, model_runner
from vllm_ascend.worker.v2.input_batch import AscendInputBatch
cudagraph_utils.InputBatch = AscendInputBatch
model_runner.InputBatch = AscendInputBatch

View File

@@ -0,0 +1,26 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/model_states/default.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
from vllm.v1.worker.gpu import model_runner
from vllm_ascend.worker.v2.model_states import init_asecnd_model_state
# prepare_attn in AscendModelState is different from vllm,
# we need to override init_model_state.
model_runner.init_model_state = init_asecnd_model_state

View File

@@ -0,0 +1,34 @@
from vllm.v1.worker.gpu import input_batch, model_runner, structured_outputs
from vllm.v1.worker.gpu.sample import bad_words, gumbel, logprob, penalties, prompt_logprob, sampler, states
from vllm.v1.worker.gpu.spec_decode import rejection_sampler, rejection_sampler_utils
from vllm.v1.worker.gpu.spec_decode.eagle import speculator
from vllm_ascend.worker.v2.input_batch import post_update
from vllm_ascend.worker.v2.sample.bad_words import apply_bad_words
from vllm_ascend.worker.v2.sample.gumbel import apply_temperature, gumbel_sample
from vllm_ascend.worker.v2.sample.logprob import compute_token_logprobs, compute_topk_logprobs
from vllm_ascend.worker.v2.sample.min_p import apply_min_p
from vllm_ascend.worker.v2.sample.penalties import apply_penalties, bincount
from vllm_ascend.worker.v2.spec_decode.rejection_sampler_utils import (
rejection_sample as npu_rejection_sample,
)
from vllm_ascend.worker.v2.structured_outputs import _apply_grammar_bitmask_kernel
penalties.apply_penalties = apply_penalties
# because sampler.py and speculator.py are imported before this patch, they must be overridden
sampler.gumbel_sample = gumbel_sample
input_batch.post_update = post_update
prompt_logprob.compute_topk_logprobs = compute_topk_logprobs
sampler.compute_topk_logprobs = compute_topk_logprobs
rejection_sampler.compute_topk_logprobs = compute_topk_logprobs
states.apply_min_p = apply_min_p
penalties.bincount = bincount
speculator.gumbel_sample = gumbel_sample
model_runner.post_update = post_update
bad_words.apply_bad_words = apply_bad_words
gumbel.apply_temperature = apply_temperature
states.apply_temperature = apply_temperature
logprob.compute_token_logprobs = compute_token_logprobs
structured_outputs._apply_grammar_bitmask_kernel = _apply_grammar_bitmask_kernel
rejection_sampler_utils.rejection_sample = npu_rejection_sample
rejection_sampler.rejection_sample = npu_rejection_sample

View File

@@ -0,0 +1,3 @@
# Reuse the platform patch. EngineCore subprocesses only load global/platform
# patches, while workers also import this compatibility module.
import vllm_ascend.patch.platform.patch_use_v2_model_runner # noqa: F401

View File

@@ -0,0 +1,161 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/block_table.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import os
from collections.abc import Callable, Sequence
from importlib.metadata import version
import numpy as np
import torch
import vllm.v1.worker.gpu.buffer_utils
from vllm.logger import logger
def check_triton_ascend_version_valid() -> bool:
"""
Check triton-ascend version and warn about UVA feature disablement.
If the installed version isn't affected by the UVA issue, return True.
"""
# Triton Ascend versions affected by the UVA pointer validation issue.
UVA_INCOMPATIBLE_VERSIONS = ("3.2.1", "3.2.2")
installed_version = version("triton-ascend")
if installed_version in UVA_INCOMPATIBLE_VERSIONS:
logger.warning(
"triton-ascend %s disables the UVA feature.\n"
"Related bug issue: https://github.com/triton-lang/triton-ascend/issues/783",
installed_version,
)
return False
return True
def is_uva_available() -> bool:
"""check if uva feature is supported in this environment"""
# FIXME(chenboxun): Some triton-ascend versions reject pinned CPU tensors.
# Thus UVA is disabled for affected versions.
# (Related bug issue link: https://github.com/triton-lang/triton-ascend/issues/783)
return (
"pinned_mem_register:True" in os.environ.get("PYTORCH_NPU_ALLOC_CONF", {})
and check_triton_ascend_version_valid()
)
def get_row_indices_from_key(key: int | slice | tuple, dim_size: int) -> set[int]:
"""get the set of row indices involved in the given key."""
if isinstance(key, int):
# parse index such as np[1]
key = key if key >= 0 else dim_size + key
# handle negative index
if key < 0 or key >= dim_size:
raise IndexError(f"row index {key} out of [0, {dim_size})")
return {key}
elif isinstance(key, slice):
# parse slice such as np[1:3]
start, stop, step = key.indices(dim_size)
return set(range(start, stop, step))
elif isinstance(key, tuple):
# parse row slice such as np[1,:100]
if len(key) == 0:
return set(range(dim_size))
return get_row_indices_from_key(key[0], dim_size)
else:
# for other types such as list/ndarray, we return all rows.
return set(range(dim_size))
class MonitoredNumPyArray:
"""A wrapper around a NumPy array that monitors modifications."""
def __init__(self, array: np.ndarray, callback: Callable):
self._array = array
self._callback = callback
def __setitem__(self, key, value):
self._array[key] = value
dim_size = self._array.shape[0]
row_indices = get_row_indices_from_key(key, dim_size)
for row in row_indices:
self._callback(row)
def __getitem__(self, key):
return self._array[key]
def __getattr__(self, name):
return getattr(self._array, name)
class MonitoredTorchTensor:
"""A wrapper around a torch tensor that monitors modifications."""
def __init__(self, tensor: torch.Tensor, callback: Callable):
self._tensor = tensor
self._callback = callback
def __setitem__(self, key, value):
self._tensor[key] = value
dim_size = self._tensor.size(0)
row_indices = get_row_indices_from_key(key, dim_size)
for row in row_indices:
self._callback(row)
def __getitem__(self, key):
return self._tensor[key]
def __getattr__(self, name):
return getattr(self._tensor, name)
class UvaBufferWrapper:
"""
Ascend NPU doesn't support UVA tensors directly.
This is a wrapper class that provides CPU and NPU views of a UVA tensor.
However if users add environment parameter below, UVA feature is Supported.
os.environ['PYTORCH_NPU_ALLOC_CONF'] = 'pinned_mem_register:True'
"""
def __init__(self, size: int | Sequence[int], dtype: torch.dtype):
self._cpu: torch.Tensor = torch.zeros(size, dtype=dtype, device="cpu", pin_memory=True)
self._np: np.ndarray = self._cpu.numpy()
self._modified_indices: set[int] = set()
self._uva: torch.Tensor = self._cpu if is_uva_available() else torch.zeros_like(self._cpu, device="npu")
def _mark_cpu_modified(self, key: int):
self._modified_indices.add(key)
@property
def cpu(self):
return self._cpu if is_uva_available() else MonitoredTorchTensor(self._cpu, self._mark_cpu_modified)
@property
def np(self):
return self._np if is_uva_available() else MonitoredNumPyArray(self._np, self._mark_cpu_modified)
@property
def uva(self):
"""Get the device data of the buffer."""
if not is_uva_available() and self._modified_indices:
# Sort for better memory access locality
dirty_rows = sorted(self._modified_indices)
# can't use copy_ method, because copy_ for index tensor
# will malloc new memory.
self._uva[dirty_rows] = self._cpu[dirty_rows].to(device="npu", non_blocking=True)
self._modified_indices.clear()
return self._uva
vllm.v1.worker.gpu.buffer_utils.UvaBuffer = UvaBufferWrapper

View File

@@ -0,0 +1,93 @@
import sys
from typing import Any
from vllm.logger import logger
from vllm.model_executor.model_loader.weight_utils import maybe_remap_kv_scale_name
class ImportPatchDecorator:
"""Import patch decorator"""
_patches: dict[str, Any] = {}
@classmethod
def register(cls, module_name):
"""Decorator for registering module patches"""
def decorator(func):
cls._patches[module_name] = func
return func
return decorator
@classmethod
def apply_patches(cls):
"""Apply all patches"""
for module_name, patch_func in cls._patches.items():
if module_name in sys.modules:
module = sys.modules[module_name]
try:
patch_func(module)
except Exception as e:
logger.error("Patch application failed %s: %s", module_name, e)
@ImportPatchDecorator.register("vllm.model_executor.models.deepseek_v2")
def patch_deepseek(module):
ori_maybe_remap_kv_scale_name = maybe_remap_kv_scale_name
def new_remap(name: str, params_dict: dict):
name = ori_maybe_remap_kv_scale_name(name, params_dict)
replace_scale_names = [
"fa_q.scale",
"fa_k.scale",
"fa_v.scale",
"fa_q.offset",
"fa_k.offset",
"fa_v.offset",
"indexer.q_rot",
"indexer.k_rot",
]
for scale_name in replace_scale_names:
if name.endswith(scale_name):
remap_name = name.replace(scale_name, f"mla_attn.mla_attn.{scale_name}")
if remap_name in params_dict:
return remap_name
else:
return remap_name.replace(".mla_attn", "")
return name
if hasattr(module, "maybe_remap_kv_scale_name"):
module._original_maybe_remap_kv_scale_name = module.maybe_remap_kv_scale_name
module.maybe_remap_kv_scale_name = new_remap
@ImportPatchDecorator.register("vllm.model_executor.model_loader.weight_utils")
def patch_weight_utils(module):
if "vllm.model_executor.models.deepseek_v2" in sys.modules:
deepseek = sys.modules["vllm.model_executor.models.deepseek_v2"]
if hasattr(deepseek, "maybe_remap_kv_scale_name"):
module.maybe_remap_kv_scale_name = deepseek.maybe_remap_kv_scale_name
original_import = __builtins__["__import__"] # type: ignore
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
module = original_import(name, globals, locals, fromlist, level)
if name in ImportPatchDecorator._patches:
try:
ImportPatchDecorator._patches[name](module)
except Exception as e:
logger.error("Patch application failed during import %s: %s", name, e)
return module
__builtins__["__import__"] = patched_import
ImportPatchDecorator.apply_patches()