init v0.23.0

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

View File

@@ -1,60 +1,10 @@
from vllm import ModelRegistry
import vllm_ascend.envs as envs_ascend
def register_model():
ModelRegistry.register_model(
"Qwen2VLForConditionalGeneration",
"vllm_ascend.models.qwen2_vl:AscendQwen2VLForConditionalGeneration")
ModelRegistry.register_model("DeepseekV4ForCausalLM", "vllm_ascend.models.deepseek_v4:AscendDeepseekV4ForCausalLM")
ModelRegistry.register_model("DeepSeekV4MTPModel", "vllm_ascend.models.deepseek_v4_mtp:DeepSeekV4MTP")
ModelRegistry.register_model(
"Qwen3VLMoeForConditionalGeneration",
"vllm_ascend.models.qwen2_5_vl_without_padding:AscendQwen3VLMoeForConditionalGeneration"
"LlamaForCausalLMVwnEagle3", "vllm_ascend.models.llama_eagle3_vwn:Eagle3VwnLlamaForCausalLM"
)
ModelRegistry.register_model(
"Qwen3VLForConditionalGeneration",
"vllm_ascend.models.qwen2_5_vl_without_padding:AscendQwen3VLForConditionalGeneration"
)
if envs_ascend.USE_OPTIMIZED_MODEL:
ModelRegistry.register_model(
"Qwen2_5_VLForConditionalGeneration",
"vllm_ascend.models.qwen2_5_vl:AscendQwen2_5_VLForConditionalGeneration"
)
else:
ModelRegistry.register_model(
"Qwen2_5_VLForConditionalGeneration",
"vllm_ascend.models.qwen2_5_vl_without_padding:AscendQwen2_5_VLForConditionalGeneration_Without_Padding"
)
ModelRegistry.register_model(
"DeepseekV2ForCausalLM",
"vllm_ascend.models.deepseek_v2:CustomDeepseekV2ForCausalLM")
ModelRegistry.register_model(
"DeepseekV3ForCausalLM",
"vllm_ascend.models.deepseek_v2:CustomDeepseekV3ForCausalLM")
ModelRegistry.register_model(
"DeepseekV32ForCausalLM",
"vllm_ascend.models.deepseek_v2:CustomDeepseekV3ForCausalLM")
ModelRegistry.register_model(
"DeepSeekMTPModel",
"vllm_ascend.models.deepseek_mtp:CustomDeepSeekMTP")
ModelRegistry.register_model(
"Qwen3MoeForCausalLM",
"vllm_ascend.models.qwen3_moe:CustomQwen3MoeForCausalLM")
# There is no PanguProMoEForCausalLM in vLLM, so we should register it before vLLM config initialization
# to make sure the model can be loaded correctly. This register step can be removed once vLLM support PanguProMoEForCausalLM.
ModelRegistry.register_model(
"PanguProMoEForCausalLM",
"vllm_ascend.torchair.models.torchair_pangu_moe:PanguProMoEForCausalLM"
)
ModelRegistry.register_model(
"Qwen3NextForCausalLM",
"vllm_ascend.models.qwen3_next:CustomQwen3NextForCausalLM")

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,537 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import typing
from collections.abc import Callable, Iterable
import torch
import torch.nn as nn
from transformers import PretrainedConfig
from vllm._aiter_ops import rocm_aiter_ops
from vllm.compilation.decorators import support_torch_compile
from vllm.config import VllmConfig
from vllm.distributed import get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size
from vllm.model_executor.layers.fused_moe import FusedMoE
from vllm.model_executor.layers.layernorm import RMSNorm
from vllm.model_executor.layers.linear import ReplicatedLinear
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead, VocabParallelEmbedding
from vllm.model_executor.model_loader.weight_utils import default_weight_loader, maybe_remap_kv_scale_name
from vllm.model_executor.models.interfaces import SupportsPP
from vllm.model_executor.models.utils import PPMissingLayer, maybe_prefix
from vllm.platforms import current_platform
from vllm.sequence import IntermediateTensors
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.utils import enable_dsa_cp, vllm_version_is
if not vllm_version_is("0.23.0"):
from vllm.model_executor.layers.fused_moe import fused_moe_make_expert_params_mapping
from .deepseek_v4 import (
DeepseekV2DecoderLayer,
DeepseekV2MixtureOfExperts,
DeepseekV4MoE,
get_spec_layer_idx_from_weight_name,
)
class SharedHead(nn.Module):
def __init__(
self,
config: PretrainedConfig,
prefix: str,
quant_config: QuantizationConfig | None = None,
) -> None:
super().__init__()
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
prefix=maybe_prefix(prefix, "head"),
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
return self.norm(hidden_states)
class DeepSeekMultiTokenPredictorLayer(nn.Module):
def __init__(self, vllm_config: VllmConfig, prefix: str) -> None:
super().__init__()
config = vllm_config.speculative_config.draft_model_config.hf_config
self.config = config
quant_config = vllm_config.quant_config
self.e_proj = ReplicatedLinear(
config.hidden_size, config.hidden_size, bias=False, quant_config=quant_config, return_bias=False
)
self.h_proj = ReplicatedLinear(
config.hidden_size, config.hidden_size, bias=False, quant_config=quant_config, return_bias=False
)
self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.device = current_platform.device_type
self.is_v32 = hasattr(config, "index_topk")
if self.is_v32:
topk_tokens = config.index_topk
topk_indices_buffer = torch.empty(
vllm_config.scheduler_config.max_num_batched_tokens,
topk_tokens,
dtype=torch.int32,
device=self.device,
)
else:
topk_indices_buffer = None
self.shared_head = SharedHead(config=config, prefix=prefix, quant_config=quant_config)
self.mtp_block = DeepseekV2DecoderLayer(
vllm_config,
prefix,
config=self.config,
topk_indices_buffer=topk_indices_buffer,
is_draft_layer=True,
)
self.hc_eps = config.hc_eps
self.hc_mult = hc_mult = config.hc_mult
hc_dim = hc_mult * config.hidden_size
self.hc_head_fn = nn.Parameter(torch.empty(hc_mult, hc_dim, dtype=torch.float32))
self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32))
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
self.norm_eps = config.rms_norm_eps
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)
previous_hidden_states = previous_hidden_states.view(-1, self.hc_mult, self.config.hidden_size)
previous_hidden_states = self.hnorm(previous_hidden_states)
hidden_states = self.e_proj(inputs_embeds).unsqueeze(-2) + self.h_proj(previous_hidden_states)
hidden_states, residual = self.mtp_block(positions=positions, hidden_states=hidden_states, residual=None)
# hidden_states = self.hc_head(hidden_states, self.hc_head_fn,
# self.hc_head_scale, self.hc_head_base)
return hidden_states
def hc_head(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor):
shape, dtype = x.size(), x.dtype
x = x.flatten(1).float()
rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps)
mixes = torch.nn.functional.linear(x, hc_fn) * rsqrt
pre = torch.sigmoid(mixes * hc_scale + hc_base) + self.hc_eps
y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=1)
return y.to(dtype)
class DeepSeekMultiTokenPredictor(nn.Module):
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
config = vllm_config.model_config.hf_config
self.mtp_start_layer_idx = config.num_hidden_layers
self.num_mtp_layers = getattr(config, "num_nextn_predict_layers", 1)
# to map the exact layer index from weights
self.layers = torch.nn.ModuleDict(
{
str(idx): DeepSeekMultiTokenPredictorLayer(vllm_config, f"{prefix}.{idx}")
for idx in range(
0,
self.num_mtp_layers,
)
}
)
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
)
self.logits_processor = LogitsProcessor(config.vocab_size)
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.embed_tokens(input_ids)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
previous_hidden_states: torch.Tensor,
inputs_embeds: torch.Tensor | None = None,
spec_step_idx: int = 0,
) -> torch.Tensor:
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
current_step_idx = spec_step_idx % self.num_mtp_layers
return self.layers[str(current_step_idx)](
input_ids,
positions,
previous_hidden_states,
inputs_embeds,
current_step_idx,
)
def compute_logits(
self,
hidden_states: torch.Tensor,
spec_step_idx: int = 0,
) -> torch.Tensor:
current_step_idx = spec_step_idx % self.num_mtp_layers
mtp_layer = self.layers[str(current_step_idx)]
hidden_states = hidden_states.view(-1, mtp_layer.hc_mult, mtp_layer.config.hidden_size)
hidden_states = mtp_layer.hc_head(
hidden_states, mtp_layer.hc_head_fn, mtp_layer.hc_head_scale, mtp_layer.hc_head_base
)
logits = self.logits_processor(mtp_layer.shared_head.head, mtp_layer.shared_head(hidden_states))
return logits
@support_torch_compile
class DeepSeekV4MTP(nn.Module, SupportsPP, DeepseekV2MixtureOfExperts):
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
super().__init__()
self.config = vllm_config.model_config.hf_config
self.quant_config = vllm_config.quant_config
self.model = DeepSeekMultiTokenPredictor(vllm_config=vllm_config, prefix=maybe_prefix(prefix, "mtp"))
# Set MoE hyperparameters
self.set_moe_parameters()
def set_moe_parameters(self):
self.expert_weights = []
self.num_expert_groups = getattr(self.config, "n_group", 1)
self.moe_layers = []
self.moe_mlp_layers = []
example_moe = None
for layer in self.model.layers.values():
if isinstance(layer, PPMissingLayer):
continue
assert isinstance(layer, DeepSeekMultiTokenPredictorLayer)
layer = layer.mtp_block
assert isinstance(layer, DeepseekV2DecoderLayer)
if isinstance(layer.mlp, DeepseekV4MoE):
# Pick last one layer since the first ones may be dense layers.
example_moe = layer.mlp
self.moe_mlp_layers.append(layer.mlp)
self.moe_layers.append(layer.mlp.experts)
self.extract_moe_parameters(example_moe)
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.model.embed_input_ids(input_ids)
def 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:
hidden_states = self.model(input_ids, positions, hidden_states, inputs_embeds, spec_step_idx)
return hidden_states
def compute_logits(
self,
hidden_states: torch.Tensor,
spec_step_idx: int = 0,
) -> torch.Tensor | None:
return self.model.compute_logits(hidden_states, spec_step_idx)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
rocm_aiter_moe_shared_expert_enabled = rocm_aiter_ops.is_fusion_moe_shared_experts_enabled()
rocm_aiter_moe_shared_expert_enabled = getattr(get_ascend_config(), "mix_placement", False)
stacked_params_mapping = [
("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1),
]
if vllm_version_is("0.23.0"):
expert_params_mapping = FusedMoE.make_expert_params_mapping(
model=self.model,
ckpt_gate_proj_name="gate_proj",
ckpt_down_proj_name="down_proj",
ckpt_up_proj_name="up_proj",
num_experts=self.config.n_routed_experts
+ (self.config.n_shared_experts if rocm_aiter_moe_shared_expert_enabled else 0),
num_redundant_experts=self.num_redundant_experts,
)
else:
expert_params_mapping = fused_moe_make_expert_params_mapping(
model=self.model,
ckpt_gate_proj_name="gate_proj",
ckpt_down_proj_name="down_proj",
ckpt_up_proj_name="up_proj",
num_experts=self.config.n_routed_experts
+ (self.config.n_shared_experts if rocm_aiter_moe_shared_expert_enabled else 0),
num_redundant_experts=self.num_redundant_experts,
)
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
# Attention heads per rank
heads_per_rank = self.config.num_attention_heads // tp_size
head_start = tp_rank * heads_per_rank
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
if self.quant_config is not None and self.quant_config.get_name() == "fp8":
if name == "embed.weight":
name = "mtp.0.emb.tok_emb.weight"
if name == "head.weight":
name = "mtp.0.head.weight"
spec_layer = get_spec_layer_idx_from_weight_name(self.config, name)
if spec_layer is None:
continue
assert "mtp.0." in name
if ".emb.tok_emb." in name:
name = name.replace("mtp.0.", "model.")
elif self.no_mtp_block_in_name(name):
name = name.replace("mtp.0.", "model.layers.0.")
else:
name = name.replace("mtp.0.", "model.layers.0.mtp_block.")
if ".w1." in name:
name = name.replace(".w1.", ".gate_proj.")
if ".w2." in name:
name = name.replace(".w2.", ".down_proj.")
if ".w3." in name:
name = name.replace(".w3.", ".up_proj.")
if name.endswith(".scale"):
name = name.replace(".scale", ".weight_scale")
if ".head." in name:
name = name.replace(".head.", ".shared_head.head.")
if ".norm." in name:
name = name.replace(".norm.", ".shared_head.norm.")
if ".emb.tok_emb." in name:
name = name.replace(".emb.tok_emb.", ".embed_tokens.")
if "attn" in name and "self_attn" not in name:
name = name.replace(".attn.", ".self_attn.")
if ".ffn." in name:
name = name.replace(".ffn.", ".mlp.")
if ".ffn_norm." in name:
name = name.replace(".ffn_norm.", ".post_attention_layernorm.")
if ".attn_norm." in name:
name = name.replace(".attn_norm.", ".input_layernorm.")
if ".gate.bias" in name:
name = name.replace(".gate.bias", ".gate.e_score_correction_bias")
if "sink" in name:
param = params_dict[name]
if enable_dsa_cp():
param.data.copy_(loaded_weight)
else:
# Handle attention sinks (distributed across ranks)
narrow_weight = loaded_weight.narrow(0, head_start, heads_per_rank)
param.data.copy_(narrow_weight)
loaded_params.add(name)
continue
is_fusion_moe_shared_experts_layer = rocm_aiter_moe_shared_expert_enabled and ("mlp.shared_experts" in name)
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if weight_name not in name:
continue
# We have mlp.experts[0].gate_proj in the checkpoint.
# Since we handle the experts below in expert_params_mapping,
# we need to skip here BEFORE we update the name, otherwise
# name will be updated to mlp.experts[0].gate_up_proj, which
# will then be updated below in expert_params_mapping
# for mlp.experts[0].gate_gate_up_proj, which breaks load.
if ("mlp.experts." in name) and name not in params_dict:
continue
if is_fusion_moe_shared_experts_layer:
continue
name_mapped = name.replace(weight_name, param_name)
# QKV fusion is optional, fall back to normal
# weight loading if it's not enabled
if (param_name == "fused_qkv_a_proj") and name_mapped not in params_dict:
continue
else:
name = name_mapped
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
# Special handling: when AITER fusion_shared_experts is enabled,
# checkpoints may provide a single widened shared_experts tensor
# without explicit expert indices
# (e.g. ...mlp.shared_experts.gate_proj.weight).
# For models with multiple shared experts, split that tensor
# evenly into per-shared-expert slices and load them into
# appended expert slots mlp.experts.{n_routed_experts + j}.*
# accordingly.
num_chunks = 1
if is_fusion_moe_shared_experts_layer:
num_chunks = getattr(self.config, "n_shared_experts", 1) or 1
# Determine split axis based on op type
# gate/up: ColumnParallel → split along dim 0
# down: RowParallel → split along dim 1
split_dim = 1 if "down_proj.weight" in name else 0
total = loaded_weight.shape[split_dim]
assert total % num_chunks == 0, (
f"Shared expert weight dim {total} not divisible by num_chunks {num_chunks}"
)
chunk_size = total // num_chunks
for j in range(num_chunks):
chunk_name = name
weight_to_load = loaded_weight
if is_fusion_moe_shared_experts_layer:
if split_dim == 0:
weight_to_load = loaded_weight[j * chunk_size : (j + 1) * chunk_size, :]
else:
weight_to_load = loaded_weight[:, j * chunk_size : (j + 1) * chunk_size]
# Synthesize an expert-style name so expert mapping
# can route it
chunk_name = name.replace(
"mlp.shared_experts",
f"mlp.experts.{self.config.n_routed_experts + j}",
)
# Use expert_params_mapping to locate the destination
# param and delegate to its expert-aware weight_loader
# with expert_id.
is_expert_weight = False
for mapping in expert_params_mapping:
param_name, weight_name, expert_id, shard_id = mapping
if weight_name not in chunk_name:
continue
# Anyway, this is an expert weight and should not be
# attempted to load as other weights later
is_expert_weight = True
# Do not modify `name` since the loop may continue here
# Instead, create a new variable
name_mapped = chunk_name.replace(weight_name, param_name)
param = params_dict[name_mapped]
# We should ask the weight loader to return success or
# not here since otherwise we may skip experts with
# other available replicas.
weight_loader = typing.cast(Callable[..., bool], param.weight_loader)
success = weight_loader(
param,
weight_to_load,
name_mapped,
shard_id=shard_id,
expert_id=expert_id,
return_success=True,
)
if success:
if not is_fusion_moe_shared_experts_layer:
name = name_mapped
else:
loaded_params.add(name_mapped)
break
else:
if is_expert_weight:
# We've checked that this is an expert weight
# However it's not mapped locally to this rank
# So we simply skip it
continue
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
name = maybe_remap_kv_scale_name(name, params_dict)
if name is None:
continue
# # According to DeepSeek-V3 Technical Report, MTP modules
# # shares embedding layer. We only load the first weights.
# if (
# spec_layer != self.model.mtp_start_layer_idx
# and ".layers" not in name
# ):
# continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
if not is_fusion_moe_shared_experts_layer:
loaded_params.add(name)
return loaded_params
def _rewrite_spec_layer_name(self, spec_layer: int, name: str) -> str:
"""
Rewrite the weight name to match the format of the original model.
Add .mtp_block for modules in transformer layer block for spec layer
and rename shared layer weights to be top level.
"""
spec_layer_weight_names = [
"embed_tokens",
"enorm",
"hnorm",
"eh_proj",
"shared_head",
]
shared_weight_names = ["embed_tokens"]
spec_layer_weight = False
shared_weight = False
for weight_name in spec_layer_weight_names:
if weight_name in name:
spec_layer_weight = True
if weight_name in shared_weight_names:
shared_weight = True
break
if not spec_layer_weight:
# treat rest weights as weights for transformer layer block
name = name.replace(f"model.layers.{spec_layer}.", f"model.layers.{spec_layer}.mtp_block.")
elif shared_weight:
# treat shared weights as top level weights
name = name.replace(f"model.layers.{spec_layer}.", "model.")
return name
def no_mtp_block_in_name(self, layer_name: str) -> bool:
names = [
".hc_head_fn",
".hc_head_base",
".hc_head_scale",
".e_proj.",
".h_proj.",
".enorm.",
".hnorm.",
".norm.",
".head.",
".emb.tok_emb.",
]
return any(name in layer_name for name in names)

View File

View File

@@ -0,0 +1,194 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Attention layer."""
from typing import cast
import torch
import torch.nn as nn
import vllm.envs as envs
from vllm.config import CacheConfig, get_current_vllm_config
from vllm.config.vllm import VllmConfig
from vllm.model_executor.layers.attention.attention import _init_kv_cache_quant
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
# from vllm.model_executor.layers.batch_invariant import vllm_is_batch_invariant
from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.platforms import current_platform
from vllm.utils.torch_utils import kv_cache_dtype_str_to_dtype
from vllm.v1.attention.backend import AttentionBackend
from vllm.v1.attention.backends.mla.sparse_swa import DeepseekV4SWACache
from vllm.v1.kv_cache_interface import KVCacheSpec
from vllm_ascend.attention.abstract import DSAAttentionImpl
from vllm_ascend.attention.dsa_v1 import AscendDSABackend
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
from vllm_ascend.utils import (
AscendDeviceType,
get_ascend_device_type,
)
def get_dsv4_block_sizes():
# cache_config.block_size: [mla, swa, c4 state, c128 state], [page_size_padded_t1, page_size_padded_t2]
_DSV4_BLOCK_SIZES = {
128: [[128, 128, 8, 32], [16640, 131072]],
64: [[64, 64, 4, 16], [8320, 65536]],
32: [[32, 32, 2, 8], [4160, 32768]],
}
_DSV4_BLOCK_SIZES_A5 = {
128: [[128, 128, 8, 16], [16896, 81920]],
64: [[64, 64, 4, 8], [8448, 40960]],
32: [[32, 32, 2, 4], [4224, 20480]],
}
if get_ascend_device_type() in {AscendDeviceType.A5}:
return _DSV4_BLOCK_SIZES_A5
else:
return _DSV4_BLOCK_SIZES
DSV4_BLOCK_SIZES = get_dsv4_block_sizes()
class DSAAttention(nn.Module, AttentionLayerBase):
"""Multi-Head Latent Attention layer.
This class takes query, and compressed key/value tensors as input.
The class does the following:
1. Store the input key and value tensors in the KV cache.
2. Perform (multi-head/multi-query/grouped-query) attention.
3. Return the output tensor.
"""
def __init__(
self,
dim: int,
n_heads: int,
scale: float,
n_local_heads: int,
q_lora_rank: int,
o_lora_rank: int,
head_dim: int,
rope_head_dim: int | None,
nope_head_dim: int,
n_groups: int,
n_local_groups: int,
window_size: int,
compress_ratio: int,
cache_config: CacheConfig | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
**extra_impl_args,
):
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.scale = scale
self.n_local_heads = n_local_heads
self.q_lora_rank = q_lora_rank
self.o_lora_rank = o_lora_rank
self.head_dim = head_dim
self.rope_head_dim = rope_head_dim
self.nope_head_dim = nope_head_dim
self.n_groups = n_groups
self.n_local_groups = n_local_groups
self.window_size = window_size
self.compress_ratio = compress_ratio
self.layer_name = prefix
self.head_size = self.head_dim
self.swa_cache_layer: DeepseekV4SWACache = extra_impl_args.get("swa_cache_layer")
assert self.swa_cache_layer is not None
if cache_config is not None:
kv_cache_dtype = cache_config.cache_dtype
else:
kv_cache_dtype = "auto"
# Initialize KV cache quantization attributes
_init_kv_cache_quant(self, quant_config, prefix)
self.attn_backend = AscendDSABackend
# NOTE(zxr): vllm_is_batch_invariant is delete during updating to v0.20.1
if (
cache_config is not None
and cache_config.enable_prefix_caching
and (self.attn_backend.get_name() == "TRITON_MLA" or self.attn_backend.get_name() == "FLASHINFER")
):
cache_config.enable_prefix_caching = False
impl_cls = cast(type[DSAAttentionImpl], self.attn_backend.get_impl_cls())
self.impl = impl_cls(
dim=self.dim,
n_heads=self.n_heads,
scale=self.scale,
n_local_heads=self.n_local_heads,
q_lora_rank=self.q_lora_rank,
o_lora_rank=self.o_lora_rank,
head_dim=self.head_dim,
rope_head_dim=self.rope_head_dim,
nope_head_dim=self.nope_head_dim,
n_groups=self.n_groups,
n_local_groups=self.n_local_groups,
window_size=self.window_size,
compress_ratio=self.compress_ratio,
**extra_impl_args,
)
self.use_direct_call = not current_platform.opaque_attention_op()
compilation_config = get_current_vllm_config().compilation_config
if prefix in compilation_config.static_forward_context:
raise ValueError(f"Duplicate layer name: {prefix}")
compilation_config.static_forward_context[prefix] = self
self.kv_cache = [
torch.tensor([]) for _ in range(get_current_vllm_config().parallel_config.pipeline_parallel_size)
]
self.kv_cache_dtype = kv_cache_dtype
self.use_sparse = True
# Initialize q/k/v range constants.
self.q_range = torch.tensor(envs.Q_SCALE_CONSTANT, dtype=torch.float32)
self.k_range = torch.tensor(envs.K_SCALE_CONSTANT, dtype=torch.float32)
self.v_range = torch.tensor(envs.V_SCALE_CONSTANT, dtype=torch.float32)
def forward(
self,
q: torch.Tensor,
kv_c_normed: torch.Tensor,
k_pe: torch.Tensor,
output_shape: torch.Size | None = None,
) -> torch.Tensor:
return q
def process_weights_after_loading(self, act_dtype: torch.dtype):
if hasattr(self.impl, "process_weights_after_loading"):
self.impl.process_weights_after_loading(act_dtype)
def get_attn_backend(self) -> type[AttentionBackend]:
return self.attn_backend
def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec:
if self.compress_ratio <= 1: # SWA part. Allocated separately as DeepseekV4SWACache.
return None
kv_cache_dtype = kv_cache_dtype_str_to_dtype(self.kv_cache_dtype, vllm_config.model_config)
if get_ascend_device_type() in {AscendDeviceType.A5}:
kv_cache_dtype = torch.float8_e4m3fn
vllm_config.cache_config.cache_dtype = "float8_e4m3fn"
cached_head_size = (
(self.head_size + 128) if get_ascend_device_type() in {AscendDeviceType.A5} else self.head_size
)
return AscendMLAAttentionSpec(
block_size=DSV4_BLOCK_SIZES[vllm_config.cache_config.block_size][0][0],
num_kv_heads=1,
head_size=cached_head_size,
dtype=kv_cache_dtype,
model_version="deepseek_v4",
compress_ratio=self.compress_ratio,
cache_dtype_str=vllm_config.cache_config.cache_dtype,
)

View File

@@ -0,0 +1,257 @@
import torch
import torch.nn as nn
from vllm.compilation.decorators import support_torch_compile
from vllm.config import get_current_vllm_config
from vllm.model_executor.layers.layernorm import RMSNorm
from vllm.model_executor.layers.linear import QKVParallelLinear, ReplicatedLinear
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)
from vllm.model_executor.models.llama_eagle3 import (
Eagle3LlamaForCausalLM,
)
from vllm.model_executor.models.llama_eagle3 import (
LlamaDecoderLayer as Eagle3LlamaDecoderLayer,
)
from vllm.model_executor.models.llama_eagle3 import (
LlamaModel as Eagle3LlamaModel,
)
from vllm.model_executor.models.utils import get_draft_quant_config, maybe_prefix
def _linear(inp, out, vc, qc, pfx):
return ReplicatedLinear(
input_size=inp,
output_size=out,
bias=False,
params_dtype=vc.model_config.dtype,
quant_config=qc,
prefix=pfx,
return_bias=False,
)
class PreVwnLayerV1(nn.Module):
def __init__(self, vllm_config, prefix="", config=None, quant_config=None):
super().__init__()
cfg = config or vllm_config.model_config.hf_config
hs, m, r = cfg.hidden_size, getattr(cfg, "vwn_m", 1), getattr(cfg, "vwn_r", 1)
wd = int(hs * r)
self.m, self.hidden_size, self.wider_dim = m, hs, wd
self.input_layernorm = RMSNorm(hs, eps=cfg.rms_norm_eps)
self.hidden_norm = RMSNorm(hs, eps=cfg.rms_norm_eps)
self.fc = _linear(2 * hs, hs, vllm_config, quant_config, maybe_prefix(prefix, "fc"))
self.upward = _linear(hs // m, wd // m, vllm_config, quant_config, maybe_prefix(prefix, "upward"))
def forward(self, embeds, hidden_states):
x = self.fc(torch.cat([self.input_layernorm(embeds), self.hidden_norm(hidden_states)], dim=-1))
return self.upward(x.view(-1, self.hidden_size // self.m)).view(-1, self.wider_dim)
class VwnLlamaDecoderLayer(Eagle3LlamaDecoderLayer):
def __init__(self, vllm_config, prefix="", config=None, layer_idx=0):
super().__init__(vllm_config, prefix=prefix, config=config, layer_idx=layer_idx)
cfg = config or vllm_config.model_config.hf_config
qc = self.get_quant_config(vllm_config)
m, r = getattr(cfg, "vwn_m", 1), getattr(cfg, "vwn_r", 1)
hs, wd = self.hidden_size, int(self.hidden_size * r)
self.m, self.wider_dim, self.layer_idx = m, wd, layer_idx
if layer_idx == 0:
self.self_attn.qkv_proj = QKVParallelLinear(
hs,
self.self_attn.head_dim,
self.self_attn.total_num_heads,
self.self_attn.total_num_kv_heads,
bias=getattr(cfg, "attention_bias", False),
quant_config=qc,
prefix=maybe_prefix(prefix, "self_attn.qkv_proj"),
)
mp = maybe_prefix
self.pre_vwn_layer = PreVwnLayerV1(vllm_config, mp(prefix, "layers.pre_vwn_layer"), cfg, qc)
self.downward_and_forgot = _linear(wd // m, (hs + wd) // m, vllm_config, qc, mp(prefix, "downward_and_forgot"))
self.pre_attention_layernorm = RMSNorm(hs, eps=cfg.rms_norm_eps)
self.upward_after_attn = _linear(hs // m, wd // m, vllm_config, qc, mp(prefix, "upward_after_attn"))
self.downward_and_forgot_after_attn = _linear(
wd // m, (hs + wd) // m, vllm_config, qc, mp(prefix, "downward_and_forgot_after_attn")
)
self.post_attention_layernorm = RMSNorm(hs, eps=cfg.rms_norm_eps)
self.upward_after_mlp = _linear(hs // m, wd // m, vllm_config, qc, mp(prefix, "upward_after_mlp"))
self.downward = _linear(wd // m, hs // m, vllm_config, qc, mp(prefix, "downward"))
def forward(self, positions, embeds, hidden_states, residual):
if self.layer_idx == 0:
hs, wd, m = self.hidden_size, self.wider_dim, self.m
wider = self.pre_vwn_layer(embeds, hidden_states)
# Attention
out = self.downward_and_forgot(wider.view(-1, wd // m)).view(-1, hs + wd)
hidden, res = out.split([hs, wd], dim=-1)
hidden = self.self_attn(positions=positions, hidden_states=self.pre_attention_layernorm(hidden))
wider = self.upward_after_attn(hidden.view(-1, hs // m)).view(-1, wd) + res
# MLP
out = self.downward_and_forgot_after_attn(wider.view(-1, wd // m)).view(-1, hs + wd)
hidden, res = out.split([hs, wd], dim=-1)
wider = (
self.upward_after_mlp(self.mlp(self.post_attention_layernorm(hidden)).view(-1, hs // m)).view(-1, wd)
+ res
)
# Downward
hidden_states = self.downward(wider.view(-1, wd // m)).view(-1, hs)
return hidden_states, residual
@support_torch_compile(dynamic_arg_dims={"input_ids": 0, "positions": -1, "hidden_states": 0, "input_embeds": 0})
class VwnLlamaModel(Eagle3LlamaModel):
def __init__(self, *, vllm_config, start_layer_id=0, prefix=""):
nn.Module.__init__(self)
self.config = vllm_config.speculative_config.draft_model_config.hf_config
self.vocab_size = self.config.vocab_size
self.quant_config = get_draft_quant_config(vllm_config)
eagle_config = getattr(self.config, "eagle_config", None)
if eagle_config is not None and "use_aux_hidden_state" in eagle_config:
self.use_aux_hidden_state = eagle_config["use_aux_hidden_state"]
else:
self.use_aux_hidden_state = True
self.norm_before_fc = getattr(self.config, "norm_before_fc", False)
vc = get_current_vllm_config()
self.embed_tokens = VocabParallelEmbedding(
self.config.vocab_size,
self.config.hidden_size,
prefix=maybe_prefix(prefix, "embed_tokens"),
)
self.layers = nn.ModuleList(
[
VwnLlamaDecoderLayer(vc, maybe_prefix(prefix, f"layers.{i + start_layer_id}"), self.config, layer_idx=i)
for i in range(self.config.num_hidden_layers)
]
)
if self.use_aux_hidden_state:
if hasattr(self.config, "target_hidden_size"):
fc_input_size = self.config.target_hidden_size * 3
else:
fc_input_size = self.config.hidden_size * 3
if self.norm_before_fc:
self.input_norm = RMSNorm(
fc_input_size,
eps=self.config.rms_norm_eps,
)
else:
self.input_norm = None
self.fc_norm = None
self.num_aux_hidden_states = 3
self.fc = ReplicatedLinear(
input_size=fc_input_size,
output_size=self.config.hidden_size,
bias=False,
params_dtype=vllm_config.model_config.dtype,
quant_config=self.quant_config,
prefix=maybe_prefix(prefix, "fc"),
return_bias=False,
)
self.norm = RMSNorm(
self.config.hidden_size,
eps=self.config.rms_norm_eps,
)
def forward(self, input_ids, positions, hidden_states, input_embeds=None):
if input_embeds is None:
input_embeds = self.embed_input_ids(input_ids)
residual = None
for layer in self.layers:
hidden_states, residual = layer(
positions=positions, embeds=input_embeds, hidden_states=hidden_states, residual=residual
)
return self.norm(hidden_states, residual), hidden_states
class Eagle3VwnLlamaForCausalLM(Eagle3LlamaForCausalLM):
def __init__(self, *, vllm_config, prefix=""):
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
n = vllm_config.model_config.get_num_layers(vllm_config.parallel_config)
self.config.target_layer_count = n
self.model = VwnLlamaModel(vllm_config=vllm_config, prefix="model", start_layer_id=n)
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 compute_logits(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor | None:
logits = self.logits_processor(self.lm_head, hidden_states)
if self.draft_id_to_target_id is None:
assert logits.shape[1] == self.config.vocab_size, (
f"Expected logits to have shape (*, {self.config.vocab_size}), but got {logits.shape}"
)
return logits
base = torch.arange(self.config.draft_vocab_size, device=logits.device)
targets = base + self.draft_id_to_target_id
logits_new = logits.new_full(
(
logits.shape[0],
self.config.vocab_size,
),
float("-inf"),
)
logits_new[:, targets] = logits
return logits_new
def combine_hidden_states(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor:
if not self.model.use_aux_hidden_state:
return hidden_states
# combine multiple auxiliary hidden states returned by eagle3
if self.model.norm_before_fc:
hidden_states = self.model.input_norm(hidden_states)
# `norm_before_fc` adds a single RMSNorm before the FC layer, whereas `fc_norm`
# applies separate RMSNorms to each chunk of the hidden states.
if self.model.fc_norm is not None:
chunks = hidden_states.chunk(self.model.num_aux_hidden_states, dim=-1)
hidden_states = torch.cat(
[norm(chunk) for norm, chunk in zip(self.model.fc_norm, chunks)],
dim=-1,
)
return self.model.fc(hidden_states)