@@ -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
|
||||
|
||||
295
vllm_ascend/patch/worker/_hccl_pg_registry.py
Normal file
295
vllm_ascend/patch/worker/_hccl_pg_registry.py
Normal 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)
|
||||
38
vllm_ascend/patch/worker/patch_cudagraph.py
Normal file
38
vllm_ascend/patch/worker/patch_cudagraph.py
Normal 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
|
||||
80
vllm_ascend/patch/worker/patch_deepseek_mtp.py
Normal file
80
vllm_ascend/patch/worker/patch_deepseek_mtp.py
Normal 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
|
||||
279
vllm_ascend/patch/worker/patch_deepseek_v2.py
Normal file
279
vllm_ascend/patch/worker/patch_deepseek_v2.py
Normal 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
|
||||
269
vllm_ascend/patch/worker/patch_distributed.py
Normal file
269
vllm_ascend/patch/worker/patch_distributed.py
Normal 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()
|
||||
145
vllm_ascend/patch/worker/patch_draft_quarot.py
Normal file
145
vllm_ascend/patch/worker/patch_draft_quarot.py
Normal 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
|
||||
129
vllm_ascend/patch/worker/patch_eagle3_init.py
Normal file
129
vllm_ascend/patch/worker/patch_eagle3_init.py
Normal 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."
|
||||
)
|
||||
238
vllm_ascend/patch/worker/patch_eagle3_pp_aux.py
Normal file
238
vllm_ascend/patch/worker/patch_eagle3_pp_aux.py
Normal 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
|
||||
20
vllm_ascend/patch/worker/patch_fused_moe.py
Normal file
20
vllm_ascend/patch/worker/patch_fused_moe.py
Normal 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
|
||||
78
vllm_ascend/patch/worker/patch_gqa_c8.py
Normal file
78
vllm_ascend/patch/worker/patch_gqa_c8.py
Normal 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
|
||||
)
|
||||
54
vllm_ascend/patch/worker/patch_idex_310.py
Normal file
54
vllm_ascend/patch/worker/patch_idex_310.py
Normal 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
|
||||
)
|
||||
95
vllm_ascend/patch/worker/patch_kimi_k25.py
Normal file
95
vllm_ascend/patch/worker/patch_kimi_k25.py
Normal 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
|
||||
284
vllm_ascend/patch/worker/patch_mamba_utils.py
Normal file
284
vllm_ascend/patch/worker/patch_mamba_utils.py
Normal 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
|
||||
181
vllm_ascend/patch/worker/patch_minimax_m2.py
Normal file
181
vllm_ascend/patch/worker/patch_minimax_m2.py
Normal 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
|
||||
154
vllm_ascend/patch/worker/patch_minimax_m2_linear_attn.py
Normal file
154
vllm_ascend/patch/worker/patch_minimax_m2_linear_attn.py
Normal 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))
|
||||
129
vllm_ascend/patch/worker/patch_npugraph_ex_triton.py
Normal file
129
vllm_ascend/patch/worker/patch_npugraph_ex_triton.py
Normal 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
|
||||
@@ -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
|
||||
209
vllm_ascend/patch/worker/patch_qwen3_5.py
Normal file
209
vllm_ascend/patch/worker/patch_qwen3_5.py
Normal 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
|
||||
62
vllm_ascend/patch/worker/patch_qwen3_dflash.py
Normal file
62
vllm_ascend/patch/worker/patch_qwen3_dflash.py
Normal 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
|
||||
50
vllm_ascend/patch/worker/patch_qwen3_next_mtp.py
Normal file
50
vllm_ascend/patch/worker/patch_qwen3_next_mtp.py
Normal 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
|
||||
110
vllm_ascend/patch/worker/patch_qwen3vl.py
Normal file
110
vllm_ascend/patch/worker/patch_qwen3vl.py
Normal 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()
|
||||
9
vllm_ascend/patch/worker/patch_rejection_sampler.py
Normal file
9
vllm_ascend/patch/worker/patch_rejection_sampler.py
Normal 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
|
||||
188
vllm_ascend/patch/worker/patch_routed_experts_capture.py
Normal file
188
vllm_ascend/patch/worker/patch_routed_experts_capture.py
Normal 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
|
||||
126
vllm_ascend/patch/worker/patch_triton.py
Normal file
126
vllm_ascend/patch/worker/patch_triton.py
Normal 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
|
||||
)
|
||||
0
vllm_ascend/patch/worker/patch_v2/__init__.py
Normal file
0
vllm_ascend/patch/worker/patch_v2/__init__.py
Normal file
11
vllm_ascend/patch/worker/patch_v2/patch_attn_utils.py
Normal file
11
vllm_ascend/patch/worker/patch_v2/patch_attn_utils.py
Normal 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
|
||||
25
vllm_ascend/patch/worker/patch_v2/patch_block_table.py
Normal file
25
vllm_ascend/patch/worker/patch_v2/patch_block_table.py
Normal 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
|
||||
27
vllm_ascend/patch/worker/patch_v2/patch_input_batch.py
Normal file
27
vllm_ascend/patch/worker/patch_v2/patch_input_batch.py
Normal 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
|
||||
26
vllm_ascend/patch/worker/patch_v2/patch_model_state.py
Normal file
26
vllm_ascend/patch/worker/patch_v2/patch_model_state.py
Normal 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
|
||||
34
vllm_ascend/patch/worker/patch_v2/patch_triton.py
Normal file
34
vllm_ascend/patch/worker/patch_v2/patch_triton.py
Normal 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
|
||||
@@ -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
|
||||
161
vllm_ascend/patch/worker/patch_v2/patch_uva.py
Normal file
161
vllm_ascend/patch/worker/patch_v2/patch_uva.py
Normal 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
|
||||
93
vllm_ascend/patch/worker/patch_weight_utils.py
Normal file
93
vllm_ascend/patch/worker/patch_weight_utils.py
Normal 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()
|
||||
Reference in New Issue
Block a user