@@ -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")
|
||||
|
||||
1541
vllm_ascend/models/deepseek_v4.py
Normal file
1541
vllm_ascend/models/deepseek_v4.py
Normal file
File diff suppressed because it is too large
Load Diff
537
vllm_ascend/models/deepseek_v4_mtp.py
Normal file
537
vllm_ascend/models/deepseek_v4_mtp.py
Normal 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)
|
||||
0
vllm_ascend/models/layer/__init__.py
Normal file
0
vllm_ascend/models/layer/__init__.py
Normal file
0
vllm_ascend/models/layer/attention/__init__.py
Normal file
0
vllm_ascend/models/layer/attention/__init__.py
Normal file
194
vllm_ascend/models/layer/attention/layer.py
Normal file
194
vllm_ascend/models/layer/attention/layer.py
Normal 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,
|
||||
)
|
||||
257
vllm_ascend/models/llama_eagle3_vwn.py
Normal file
257
vllm_ascend/models/llama_eagle3_vwn.py
Normal 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)
|
||||
Reference in New Issue
Block a user