| |
| |
|
|
| from collections.abc import Iterable |
| from dataclasses import replace |
| from itertools import islice |
| from typing import Any |
|
|
| import torch |
| from torch import nn |
| from nanbeige_vllm_plugin.nanbeige_config import NanbeigeConfig |
|
|
| from vllm.compilation.decorators import support_torch_compile |
| from vllm.config import CacheConfig, VllmConfig |
| from vllm.distributed import get_pp_group, get_tensor_model_parallel_world_size |
| from vllm.model_executor.layers.activation import SiluAndMul |
| from vllm.model_executor.layers.attention import ( |
| Attention, |
| EncoderOnlyAttention, |
| ) |
| from vllm.model_executor.layers.layernorm import RMSNorm |
| from vllm.model_executor.layers.linear import ( |
| MergedColumnParallelLinear, |
| QKVParallelLinear, |
| RowParallelLinear, |
| ) |
| from vllm.model_executor.layers.logits_processor import LogitsProcessor |
| from vllm.model_executor.layers.quantization import QuantizationConfig |
| from vllm.model_executor.layers.rotary_embedding import get_rope |
| 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.sequence import IntermediateTensors |
| from vllm.transformers_utils.config import is_interleaved, set_default_rope_theta |
| from vllm.v1.attention.backend import AttentionType |
|
|
| from vllm.model_executor.models.interfaces import ( |
| EagleModelMixin, |
| SupportsEagle, |
| SupportsEagle3, |
| SupportsLoRA, |
| SupportsPP, |
| ) |
| from vllm.model_executor.models.utils import ( |
| AutoWeightsLoader, |
| PPMissingLayer, |
| extract_layer_index, |
| is_pp_missing_parameter, |
| make_empty_intermediate_tensors_factory, |
| make_layers, |
| maybe_prefix, |
| ) |
|
|
|
|
| class NanbeigeMLP(nn.Module): |
| def __init__( |
| self, |
| hidden_size: int, |
| intermediate_size: int, |
| hidden_act: str, |
| quant_config: QuantizationConfig | None = None, |
| prefix: str = "", |
| ) -> None: |
| super().__init__() |
| self.gate_up_proj = MergedColumnParallelLinear( |
| hidden_size, |
| [intermediate_size] * 2, |
| bias=False, |
| quant_config=quant_config, |
| prefix=f"{prefix}.gate_up_proj", |
| ) |
| self.down_proj = RowParallelLinear( |
| intermediate_size, |
| hidden_size, |
| bias=False, |
| quant_config=quant_config, |
| prefix=f"{prefix}.down_proj", |
| ) |
| if hidden_act != "silu": |
| raise ValueError( |
| f"Unsupported activation: {hidden_act}. Only silu is supported for now." |
| ) |
| self.act_fn = SiluAndMul() |
|
|
| def forward(self, x): |
| gate_up, _ = self.gate_up_proj(x) |
| x = self.act_fn(gate_up) |
| x, _ = self.down_proj(x) |
| return x |
|
|
|
|
| class NanbeigeAttention(nn.Module): |
| def __init__( |
| self, |
| config: NanbeigeConfig, |
| hidden_size: int, |
| num_heads: int, |
| num_kv_heads: int, |
| rope_parameters: dict[str, Any], |
| max_position: int = 4096 * 32, |
| cache_config: CacheConfig | None = None, |
| quant_config: QuantizationConfig | None = None, |
| prefix: str = "", |
| attn_type: str = AttentionType.DECODER, |
| dual_chunk_attention_config: dict[str, Any] | None = None, |
| qk_norm: bool = False, |
| rms_norm_eps: float = 1e-6, |
| ) -> None: |
| super().__init__() |
| self.hidden_size = hidden_size |
| tp_size = get_tensor_model_parallel_world_size() |
| self.total_num_heads = num_heads |
| assert self.total_num_heads % tp_size == 0 |
| self.num_heads = self.total_num_heads // tp_size |
| self.total_num_kv_heads = num_kv_heads |
| if self.total_num_kv_heads >= tp_size: |
| |
| |
| assert self.total_num_kv_heads % tp_size == 0 |
| else: |
| |
| |
| assert tp_size % self.total_num_kv_heads == 0 |
| self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) |
| self.head_dim = hidden_size // self.total_num_heads |
| self.head_dim = getattr(config, "head_dim", hidden_size // self.total_num_heads) |
| self.q_size = self.num_heads * self.head_dim |
| self.kv_size = self.num_kv_heads * self.head_dim |
| self.scaling = self.head_dim**-0.5 |
| self.dual_chunk_attention_config = dual_chunk_attention_config |
| self.qk_norm = qk_norm |
|
|
| self.qkv_proj = QKVParallelLinear( |
| hidden_size, |
| self.head_dim, |
| self.total_num_heads, |
| self.total_num_kv_heads, |
| bias=False, |
| quant_config=quant_config, |
| prefix=f"{prefix}.qkv_proj", |
| ) |
| self.o_proj = RowParallelLinear( |
| self.total_num_heads * self.head_dim, |
| hidden_size, |
| bias=False, |
| quant_config=quant_config, |
| prefix=f"{prefix}.o_proj", |
| ) |
|
|
| |
| if self.qk_norm: |
| self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) |
| self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) |
|
|
| self.rotary_emb = get_rope( |
| self.head_dim, |
| max_position=max_position, |
| rope_parameters=rope_parameters, |
| dual_chunk_attention_config=dual_chunk_attention_config, |
| ) |
|
|
| self.loops_num = getattr(config, "num_loops", 1) |
| total_layers = config.num_hidden_layers |
| self.attn = nn.ModuleList() |
|
|
| for loop_idx in range(self.loops_num): |
| base_layer_idx = extract_layer_index(prefix) |
| unique_layer_idx = loop_idx * total_layers + base_layer_idx |
| unique_prefix = prefix.replace( |
| f"layers.{base_layer_idx}", f"layers.{unique_layer_idx}" |
| ) |
| self.attn.append( |
| Attention( |
| self.num_heads, |
| self.head_dim, |
| self.scaling, |
| num_kv_heads=self.num_kv_heads, |
| cache_config=cache_config, |
| quant_config=quant_config, |
| attn_type=attn_type, |
| prefix=f"{unique_prefix}.attn", |
| **{ |
| "layer_idx": unique_layer_idx, |
| "dual_chunk_attention_config": dual_chunk_attention_config, |
| } |
| if dual_chunk_attention_config and loop_idx == 0 |
| else {}, |
| ) |
| ) |
|
|
| def forward( |
| self, |
| positions: torch.Tensor, |
| hidden_states: torch.Tensor, |
| loop_idx: int, |
| ) -> torch.Tensor: |
| qkv, _ = self.qkv_proj(hidden_states) |
| q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) |
|
|
| |
| if self.qk_norm: |
| |
| |
| total_tokens = q.shape[0] |
| q = q.view(total_tokens, self.num_heads, self.head_dim) |
| k = k.view(total_tokens, self.num_kv_heads, self.head_dim) |
|
|
| |
| q = self.q_norm(q) |
| k = self.k_norm(k) |
|
|
| |
| q = q.view(total_tokens, self.q_size) |
| k = k.view(total_tokens, self.kv_size) |
|
|
| q, k = self.rotary_emb(positions, q, k) |
| attn_output = self.attn[loop_idx](q, k, v) |
| output, _ = self.o_proj(attn_output) |
| return output |
|
|
|
|
| class NanbeigeDecoderLayer(nn.Module): |
| def __init__( |
| self, |
| config: NanbeigeConfig, |
| cache_config: CacheConfig | None = None, |
| quant_config: QuantizationConfig | None = None, |
| prefix: str = "", |
| ) -> None: |
| super().__init__() |
| self.hidden_size = config.hidden_size |
| set_default_rope_theta(config, default_theta=1000000) |
| dual_chunk_attention_config = getattr( |
| config, "dual_chunk_attention_config", None |
| ) |
|
|
| if getattr(config, "is_causal", True): |
| attn_type = AttentionType.DECODER |
| else: |
| attn_type = AttentionType.ENCODER_ONLY |
|
|
| |
| qk_norm = getattr(config, "qk_norm", False) |
|
|
| self.self_attn = NanbeigeAttention( |
| config=config, |
| hidden_size=self.hidden_size, |
| num_heads=config.num_attention_heads, |
| max_position=config.max_position_embeddings, |
| num_kv_heads=config.num_key_value_heads, |
| cache_config=cache_config, |
| quant_config=quant_config, |
| rope_parameters=config.rope_parameters, |
| prefix=f"{prefix}.self_attn", |
| attn_type=attn_type, |
| dual_chunk_attention_config=dual_chunk_attention_config, |
| qk_norm=qk_norm, |
| rms_norm_eps=config.rms_norm_eps, |
| ) |
| self.mlp = NanbeigeMLP( |
| hidden_size=self.hidden_size, |
| intermediate_size=config.intermediate_size, |
| hidden_act=config.hidden_act, |
| quant_config=quant_config, |
| prefix=f"{prefix}.mlp", |
| ) |
| self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| self.post_attention_layernorm = RMSNorm( |
| config.hidden_size, eps=config.rms_norm_eps |
| ) |
|
|
| def forward( |
| self, |
| positions: torch.Tensor, |
| hidden_states: torch.Tensor, |
| residual: torch.Tensor | None, |
| loop_idx: int = 0, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| |
| if residual is None: |
| residual = hidden_states |
| hidden_states = self.input_layernorm(hidden_states) |
| else: |
| hidden_states, residual = self.input_layernorm(hidden_states, residual) |
| hidden_states = self.self_attn( |
| positions=positions, |
| hidden_states=hidden_states, |
| loop_idx=loop_idx, |
| ) |
|
|
| |
| hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) |
| hidden_states = self.mlp(hidden_states) |
| return hidden_states, residual |
|
|
|
|
| @support_torch_compile( |
| dynamic_arg_dims={ |
| "input_ids": {0: "b"}, |
| "positions": {-1: "b"}, |
| "intermediate_tensors": {0: "b"}, |
| "inputs_embeds": {0: "b"}, |
| } |
| ) |
| class NanbeigeModel(nn.Module, EagleModelMixin): |
| def __init__( |
| self, |
| *, |
| vllm_config: VllmConfig, |
| prefix: str = "", |
| decoder_layer_type: type[nn.Module] = NanbeigeDecoderLayer, |
| ): |
| super().__init__() |
|
|
| config = vllm_config.model_config.hf_config.get_text_config() |
| cache_config = vllm_config.cache_config |
| quant_config = vllm_config.quant_config |
|
|
| |
| if is_interleaved(vllm_config.model_config.hf_text_config): |
| assert config.max_window_layers == config.num_hidden_layers, ( |
| "Sliding window for some but all layers is not supported. " |
| "This model uses sliding window but `max_window_layers` = {} " |
| "is less than `num_hidden_layers` = {}. Please open an issue " |
| "to discuss this feature.".format( |
| config.max_window_layers, |
| config.num_hidden_layers, |
| ) |
| ) |
|
|
| self.config = config |
| self.quant_config = quant_config |
| self.vocab_size = config.vocab_size |
| self.loops_num = getattr(config, "num_loops", 1) |
| self.skip_loop_final_norm = getattr(config, "skip_loop_final_norm", False) |
|
|
| if get_pp_group().is_first_rank or ( |
| config.tie_word_embeddings and get_pp_group().is_last_rank |
| ): |
| self.embed_tokens = VocabParallelEmbedding( |
| config.vocab_size, |
| config.hidden_size, |
| quant_config=quant_config, |
| prefix=f"{prefix}.embed_tokens", |
| ) |
| else: |
| self.embed_tokens = PPMissingLayer() |
|
|
| self.start_layer, self.end_layer, self.layers = make_layers( |
| config.num_hidden_layers, |
| lambda prefix: decoder_layer_type( |
| config=config, |
| cache_config=cache_config, |
| quant_config=quant_config, |
| prefix=prefix, |
| ), |
| prefix=f"{prefix}.layers", |
| ) |
|
|
| self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory( |
| ["hidden_states", "residual"], config.hidden_size |
| ) |
| if get_pp_group().is_last_rank: |
| self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| else: |
| self.norm = PPMissingLayer() |
|
|
| def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: |
| return self.embed_tokens(input_ids) |
|
|
| def forward( |
| self, |
| input_ids: torch.Tensor | None, |
| positions: torch.Tensor, |
| intermediate_tensors: IntermediateTensors | None = None, |
| inputs_embeds: torch.Tensor | None = None, |
| ) -> torch.Tensor | IntermediateTensors: |
| if get_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([], 0, hidden_states, residual) |
| for loop_idx in range(self.loops_num): |
| for idx, layer in enumerate( |
| islice(self.layers, self.start_layer, self.end_layer) |
| ): |
| hidden_states, residual = layer(positions, hidden_states, residual, loop_idx=loop_idx) |
| self._maybe_add_hidden_state( |
| aux_hidden_states, idx + 1, hidden_states, residual |
| ) |
|
|
| if loop_idx < self.loops_num - 1: |
| if residual is not None: |
| hidden_states = hidden_states + residual |
| residual = None |
| if not self.skip_loop_final_norm: |
| hidden_states = self.norm(hidden_states) |
|
|
| if not get_pp_group().is_last_rank: |
| return IntermediateTensors( |
| {"hidden_states": hidden_states, "residual": residual} |
| ) |
|
|
| hidden_states, _ = self.norm(hidden_states, residual) |
|
|
| if len(aux_hidden_states) > 0: |
| return hidden_states, aux_hidden_states |
|
|
| return hidden_states |
|
|
| def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: |
| stacked_params_mapping = [ |
| |
| ("qkv_proj", "q_proj", "q"), |
| ("qkv_proj", "k_proj", "k"), |
| ("qkv_proj", "v_proj", "v"), |
| ("gate_up_proj", "gate_proj", 0), |
| ("gate_up_proj", "up_proj", 1), |
| ] |
| params_dict = dict(self.named_parameters(remove_duplicate=False)) |
| 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 ( |
| scale_name := ( |
| self.quant_config.get_cache_scale(name) |
| if hasattr(self.quant_config, "get_cache_scale") |
| else None |
| ) |
| ): |
| |
| param = params_dict[scale_name] |
| weight_loader = getattr(param, "weight_loader", default_weight_loader) |
| loaded_weight = ( |
| loaded_weight if loaded_weight.dim() == 0 else loaded_weight[0] |
| ) |
| weight_loader(param, loaded_weight) |
| loaded_params.add(scale_name) |
| continue |
| for param_name, weight_name, shard_id in stacked_params_mapping: |
| if weight_name not in name: |
| continue |
| name = name.replace(weight_name, param_name) |
| |
| if name.endswith(".bias") and name not in params_dict: |
| continue |
| if is_pp_missing_parameter(name, self): |
| continue |
| if name.endswith("scale"): |
| |
| name = maybe_remap_kv_scale_name(name, params_dict) |
| if name is None: |
| continue |
| param = params_dict[name] |
| weight_loader = getattr(param, "weight_loader", default_weight_loader) |
| if weight_loader == default_weight_loader: |
| weight_loader(param, loaded_weight) |
| else: |
| weight_loader(param, loaded_weight, shard_id) |
| break |
| else: |
| |
| 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 |
| if is_pp_missing_parameter(name, self): |
| continue |
| if name not in params_dict: |
| continue |
| param = params_dict[name] |
| weight_loader = getattr(param, "weight_loader", default_weight_loader) |
| weight_loader(param, loaded_weight) |
| loaded_params.add(name) |
| return loaded_params |
|
|
|
|
| class NanbeigeForCausalLM( |
| nn.Module, SupportsLoRA, SupportsPP, SupportsEagle, SupportsEagle3 |
| ): |
| packed_modules_mapping = { |
| "qkv_proj": [ |
| "q_proj", |
| "k_proj", |
| "v_proj", |
| ], |
| "gate_up_proj": [ |
| "gate_proj", |
| "up_proj", |
| ], |
| } |
|
|
| def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): |
| super().__init__() |
| config = vllm_config.model_config.hf_config.get_text_config() |
| quant_config = vllm_config.quant_config |
|
|
| self.config = config |
|
|
| self.quant_config = quant_config |
| self.model = NanbeigeModel( |
| vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") |
| ) |
|
|
| if get_pp_group().is_last_rank: |
| if config.tie_word_embeddings: |
| self.lm_head = self.model.embed_tokens |
| else: |
| self.lm_head = ParallelLMHead( |
| config.vocab_size, |
| config.hidden_size, |
| quant_config=quant_config, |
| prefix=maybe_prefix(prefix, "lm_head"), |
| ) |
| else: |
| self.lm_head = PPMissingLayer() |
|
|
| self.logits_processor = LogitsProcessor(config.vocab_size) |
|
|
| self.make_empty_intermediate_tensors = ( |
| self.model.make_empty_intermediate_tensors |
| ) |
|
|
| 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 | None, |
| positions: torch.Tensor, |
| intermediate_tensors: IntermediateTensors | None = None, |
| inputs_embeds: torch.Tensor | None = None, |
| ) -> torch.Tensor | IntermediateTensors: |
| hidden_states = self.model( |
| input_ids, positions, intermediate_tensors, inputs_embeds |
| ) |
| return hidden_states |
|
|
| def compute_logits( |
| self, |
| hidden_states: torch.Tensor, |
| ) -> torch.Tensor | None: |
| logits = self.logits_processor(self.lm_head, hidden_states) |
| return logits |
|
|
| def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: |
| loader = AutoWeightsLoader( |
| self, |
| skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), |
| ) |
| return loader.load_weights(weights) |
|
|
|
|