NullSense/Nanbeige4.2-3B-FP8-Dynamic
3130
1# SPDX-License-Identifier: Apache-2.02# SPDX-FileCopyrightText: Copyright contributors to the vLLM project3 4from collections.abc import Iterable5from dataclasses import replace6from itertools import islice7from typing import Any8 9import torch10from torch import nn11from nanbeige_vllm_plugin.nanbeige_config import NanbeigeConfig12 13from vllm.compilation.decorators import support_torch_compile14from vllm.config import CacheConfig, VllmConfig15from vllm.distributed import get_pp_group, get_tensor_model_parallel_world_size16from vllm.model_executor.layers.activation import SiluAndMul17from vllm.model_executor.layers.attention import (18 Attention,19 EncoderOnlyAttention,20)21from vllm.model_executor.layers.layernorm import RMSNorm22from vllm.model_executor.layers.linear import (23 MergedColumnParallelLinear,24 QKVParallelLinear,25 RowParallelLinear,26)27from vllm.model_executor.layers.logits_processor import LogitsProcessor28from vllm.model_executor.layers.quantization import QuantizationConfig29from vllm.model_executor.layers.rotary_embedding import get_rope30from vllm.model_executor.layers.vocab_parallel_embedding import (31 ParallelLMHead,32 VocabParallelEmbedding,33)34from vllm.model_executor.model_loader.weight_utils import (35 default_weight_loader,36 maybe_remap_kv_scale_name,37)38from vllm.sequence import IntermediateTensors39from vllm.transformers_utils.config import is_interleaved, set_default_rope_theta40from vllm.v1.attention.backend import AttentionType41 42from vllm.model_executor.models.interfaces import (43 EagleModelMixin,44 SupportsEagle,45 SupportsEagle3,46 SupportsLoRA,47 SupportsPP,48)49from vllm.model_executor.models.utils import (50 AutoWeightsLoader,51 PPMissingLayer,52 extract_layer_index,53 is_pp_missing_parameter,54 make_empty_intermediate_tensors_factory,55 make_layers,56 maybe_prefix,57)58 59 60class NanbeigeMLP(nn.Module):61 def __init__(62 self,63 hidden_size: int,64 intermediate_size: int,65 hidden_act: str,66 quant_config: QuantizationConfig | None = None,67 prefix: str = "",68 ) -> None:69 super().__init__()70 self.gate_up_proj = MergedColumnParallelLinear(71 hidden_size,72 [intermediate_size] * 2,73 bias=False,74 quant_config=quant_config,75 prefix=f"{prefix}.gate_up_proj",76 )77 self.down_proj = RowParallelLinear(78 intermediate_size,79 hidden_size,80 bias=False,81 quant_config=quant_config,82 prefix=f"{prefix}.down_proj",83 )84 if hidden_act != "silu":85 raise ValueError(86 f"Unsupported activation: {hidden_act}. Only silu is supported for now."87 )88 self.act_fn = SiluAndMul()89 90 def forward(self, x):91 gate_up, _ = self.gate_up_proj(x)92 x = self.act_fn(gate_up)93 x, _ = self.down_proj(x)94 return x95 96 97class NanbeigeAttention(nn.Module):98 def __init__(99 self,100 config: NanbeigeConfig,101 hidden_size: int,102 num_heads: int,103 num_kv_heads: int,104 rope_parameters: dict[str, Any],105 max_position: int = 4096 * 32,106 cache_config: CacheConfig | None = None,107 quant_config: QuantizationConfig | None = None,108 prefix: str = "",109 attn_type: str = AttentionType.DECODER,110 dual_chunk_attention_config: dict[str, Any] | None = None,111 qk_norm: bool = False,112 rms_norm_eps: float = 1e-6,113 ) -> None:114 super().__init__()115 self.hidden_size = hidden_size116 tp_size = get_tensor_model_parallel_world_size()117 self.total_num_heads = num_heads118 assert self.total_num_heads % tp_size == 0119 self.num_heads = self.total_num_heads // tp_size120 self.total_num_kv_heads = num_kv_heads121 if self.total_num_kv_heads >= tp_size:122 # Number of KV heads is greater than TP size, so we partition123 # the KV heads across multiple tensor parallel GPUs.124 assert self.total_num_kv_heads % tp_size == 0125 else:126 # Number of KV heads is less than TP size, so we replicate127 # the KV heads across multiple tensor parallel GPUs.128 assert tp_size % self.total_num_kv_heads == 0129 self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)130 self.head_dim = hidden_size // self.total_num_heads131 self.head_dim = getattr(config, "head_dim", hidden_size // self.total_num_heads)132 self.q_size = self.num_heads * self.head_dim133 self.kv_size = self.num_kv_heads * self.head_dim134 self.scaling = self.head_dim**-0.5135 self.dual_chunk_attention_config = dual_chunk_attention_config136 self.qk_norm = qk_norm137 138 self.qkv_proj = QKVParallelLinear(139 hidden_size,140 self.head_dim,141 self.total_num_heads,142 self.total_num_kv_heads,143 bias=False,144 quant_config=quant_config,145 prefix=f"{prefix}.qkv_proj",146 )147 self.o_proj = RowParallelLinear(148 self.total_num_heads * self.head_dim,149 hidden_size,150 bias=False,151 quant_config=quant_config,152 prefix=f"{prefix}.o_proj",153 )154 155 # QK Normalization support (used in BAGEL and some other models)156 if self.qk_norm:157 self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)158 self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)159 160 self.rotary_emb = get_rope(161 self.head_dim,162 max_position=max_position,163 rope_parameters=rope_parameters,164 dual_chunk_attention_config=dual_chunk_attention_config,165 )166 167 self.loops_num = getattr(config, "num_loops", 1)168 total_layers = config.num_hidden_layers169 self.attn = nn.ModuleList()170 171 for loop_idx in range(self.loops_num):172 base_layer_idx = extract_layer_index(prefix)173 unique_layer_idx = loop_idx * total_layers + base_layer_idx174 unique_prefix = prefix.replace(175 f"layers.{base_layer_idx}", f"layers.{unique_layer_idx}"176 )177 self.attn.append(178 Attention(179 self.num_heads,180 self.head_dim,181 self.scaling,182 num_kv_heads=self.num_kv_heads,183 cache_config=cache_config,184 quant_config=quant_config,185 attn_type=attn_type,186 prefix=f"{unique_prefix}.attn",187 **{188 "layer_idx": unique_layer_idx,189 "dual_chunk_attention_config": dual_chunk_attention_config,190 }191 if dual_chunk_attention_config and loop_idx == 0192 else {},193 )194 )195 196 def forward(197 self,198 positions: torch.Tensor,199 hidden_states: torch.Tensor,200 loop_idx: int,201 ) -> torch.Tensor:202 qkv, _ = self.qkv_proj(hidden_states)203 q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)204 205 # Apply QK normalization if enabled (before RoPE)206 if self.qk_norm:207 # Reshape to apply per-head normalization208 # q shape: (total_tokens, q_size) -> (total_tokens, num_heads, head_dim)209 total_tokens = q.shape[0]210 q = q.view(total_tokens, self.num_heads, self.head_dim)211 k = k.view(total_tokens, self.num_kv_heads, self.head_dim)212 213 # Apply normalization214 q = self.q_norm(q)215 k = self.k_norm(k)216 217 # Reshape back218 q = q.view(total_tokens, self.q_size)219 k = k.view(total_tokens, self.kv_size)220 221 q, k = self.rotary_emb(positions, q, k)222 attn_output = self.attn[loop_idx](q, k, v)223 output, _ = self.o_proj(attn_output)224 return output225 226 227class NanbeigeDecoderLayer(nn.Module):228 def __init__(229 self,230 config: NanbeigeConfig,231 cache_config: CacheConfig | None = None,232 quant_config: QuantizationConfig | None = None,233 prefix: str = "",234 ) -> None:235 super().__init__()236 self.hidden_size = config.hidden_size237 set_default_rope_theta(config, default_theta=1000000)238 dual_chunk_attention_config = getattr(239 config, "dual_chunk_attention_config", None240 )241 242 if getattr(config, "is_causal", True):243 attn_type = AttentionType.DECODER244 else:245 attn_type = AttentionType.ENCODER_ONLY246 247 # Check if QK normalization is enabled (used in BAGEL and some other models)248 qk_norm = getattr(config, "qk_norm", False)249 250 self.self_attn = NanbeigeAttention(251 config=config,252 hidden_size=self.hidden_size,253 num_heads=config.num_attention_heads,254 max_position=config.max_position_embeddings,255 num_kv_heads=config.num_key_value_heads,256 cache_config=cache_config,257 quant_config=quant_config,258 rope_parameters=config.rope_parameters,259 prefix=f"{prefix}.self_attn",260 attn_type=attn_type,261 dual_chunk_attention_config=dual_chunk_attention_config,262 qk_norm=qk_norm,263 rms_norm_eps=config.rms_norm_eps,264 )265 self.mlp = NanbeigeMLP(266 hidden_size=self.hidden_size,267 intermediate_size=config.intermediate_size,268 hidden_act=config.hidden_act,269 quant_config=quant_config,270 prefix=f"{prefix}.mlp",271 )272 self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)273 self.post_attention_layernorm = RMSNorm(274 config.hidden_size, eps=config.rms_norm_eps275 )276 277 def forward(278 self,279 positions: torch.Tensor,280 hidden_states: torch.Tensor,281 residual: torch.Tensor | None,282 loop_idx: int = 0,283 ) -> tuple[torch.Tensor, torch.Tensor]:284 # Self Attention285 if residual is None:286 residual = hidden_states287 hidden_states = self.input_layernorm(hidden_states)288 else:289 hidden_states, residual = self.input_layernorm(hidden_states, residual)290 hidden_states = self.self_attn(291 positions=positions,292 hidden_states=hidden_states,293 loop_idx=loop_idx,294 )295 296 # Fully Connected297 hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)298 hidden_states = self.mlp(hidden_states)299 return hidden_states, residual300 301 302@support_torch_compile(303 dynamic_arg_dims={304 "input_ids": {0: "b"},305 "positions": {-1: "b"},306 "intermediate_tensors": {0: "b"},307 "inputs_embeds": {0: "b"},308 }309)310class NanbeigeModel(nn.Module, EagleModelMixin):311 def __init__(312 self,313 *,314 vllm_config: VllmConfig,315 prefix: str = "",316 decoder_layer_type: type[nn.Module] = NanbeigeDecoderLayer,317 ):318 super().__init__()319 320 config = vllm_config.model_config.hf_config.get_text_config()321 cache_config = vllm_config.cache_config322 quant_config = vllm_config.quant_config323 324 # TODO (@robertgshaw2): see if this can be moved out325 if is_interleaved(vllm_config.model_config.hf_text_config):326 assert config.max_window_layers == config.num_hidden_layers, (327 "Sliding window for some but all layers is not supported. "328 "This model uses sliding window but `max_window_layers` = {} "329 "is less than `num_hidden_layers` = {}. Please open an issue "330 "to discuss this feature.".format(331 config.max_window_layers,332 config.num_hidden_layers,333 )334 )335 336 self.config = config337 self.quant_config = quant_config338 self.vocab_size = config.vocab_size339 self.loops_num = getattr(config, "num_loops", 1)340 self.skip_loop_final_norm = getattr(config, "skip_loop_final_norm", False)341 342 if get_pp_group().is_first_rank or (343 config.tie_word_embeddings and get_pp_group().is_last_rank344 ):345 self.embed_tokens = VocabParallelEmbedding(346 config.vocab_size,347 config.hidden_size,348 quant_config=quant_config,349 prefix=f"{prefix}.embed_tokens",350 )351 else:352 self.embed_tokens = PPMissingLayer()353 354 self.start_layer, self.end_layer, self.layers = make_layers(355 config.num_hidden_layers,356 lambda prefix: decoder_layer_type(357 config=config,358 cache_config=cache_config,359 quant_config=quant_config,360 prefix=prefix,361 ),362 prefix=f"{prefix}.layers",363 )364 365 self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory(366 ["hidden_states", "residual"], config.hidden_size367 )368 if get_pp_group().is_last_rank:369 self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)370 else:371 self.norm = PPMissingLayer()372 373 def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:374 return self.embed_tokens(input_ids)375 376 def forward(377 self,378 input_ids: torch.Tensor | None,379 positions: torch.Tensor,380 intermediate_tensors: IntermediateTensors | None = None,381 inputs_embeds: torch.Tensor | None = None,382 ) -> torch.Tensor | IntermediateTensors:383 if get_pp_group().is_first_rank:384 if inputs_embeds is not None:385 hidden_states = inputs_embeds386 else:387 hidden_states = self.embed_input_ids(input_ids)388 residual = None389 else:390 assert intermediate_tensors is not None391 hidden_states = intermediate_tensors["hidden_states"]392 residual = intermediate_tensors["residual"]393 394 aux_hidden_states = self._maybe_add_hidden_state([], 0, hidden_states, residual)395 for loop_idx in range(self.loops_num):396 for idx, layer in enumerate(397 islice(self.layers, self.start_layer, self.end_layer)398 ):399 hidden_states, residual = layer(positions, hidden_states, residual, loop_idx=loop_idx)400 self._maybe_add_hidden_state(401 aux_hidden_states, idx + 1, hidden_states, residual402 )403 404 if loop_idx < self.loops_num - 1:405 if residual is not None:406 hidden_states = hidden_states + residual407 residual = None408 if not self.skip_loop_final_norm:409 hidden_states = self.norm(hidden_states)410 411 if not get_pp_group().is_last_rank:412 return IntermediateTensors(413 {"hidden_states": hidden_states, "residual": residual}414 )415 416 hidden_states, _ = self.norm(hidden_states, residual)417 418 if len(aux_hidden_states) > 0:419 return hidden_states, aux_hidden_states420 421 return hidden_states422 423 def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:424 stacked_params_mapping = [425 # (param_name, shard_name, shard_id)426 ("qkv_proj", "q_proj", "q"),427 ("qkv_proj", "k_proj", "k"),428 ("qkv_proj", "v_proj", "v"),429 ("gate_up_proj", "gate_proj", 0),430 ("gate_up_proj", "up_proj", 1),431 ]432 params_dict = dict(self.named_parameters(remove_duplicate=False))433 loaded_params: set[str] = set()434 for name, loaded_weight in weights:435 if "rotary_emb.inv_freq" in name:436 continue437 # LOCAL PATCH: get_cache_scale doesn't exist on CompressedTensorsConfig in438 # v0.25.1 (fork targets newer main). Our FP8-dynamic checkpoint ships no KV439 # scales, so skipping the branch when the method is absent is lossless.440 if self.quant_config is not None and (441 scale_name := (442 self.quant_config.get_cache_scale(name)443 if hasattr(self.quant_config, "get_cache_scale")444 else None445 )446 ):447 # Loading kv cache quantization scales448 param = params_dict[scale_name]449 weight_loader = getattr(param, "weight_loader", default_weight_loader)450 loaded_weight = (451 loaded_weight if loaded_weight.dim() == 0 else loaded_weight[0]452 )453 weight_loader(param, loaded_weight)454 loaded_params.add(scale_name)455 continue456 for param_name, weight_name, shard_id in stacked_params_mapping:457 if weight_name not in name:458 continue459 name = name.replace(weight_name, param_name)460 # Skip loading extra bias for GPTQ models.461 if name.endswith(".bias") and name not in params_dict:462 continue463 if is_pp_missing_parameter(name, self):464 continue465 if name.endswith("scale"):466 # Remapping the name of FP8 kv-scale.467 name = maybe_remap_kv_scale_name(name, params_dict)468 if name is None:469 continue470 param = params_dict[name]471 weight_loader = getattr(param, "weight_loader", default_weight_loader)472 if weight_loader == default_weight_loader:473 weight_loader(param, loaded_weight)474 else:475 weight_loader(param, loaded_weight, shard_id)476 break477 else:478 # Skip loading extra bias for GPTQ models.479 if name.endswith(".bias") and name not in params_dict:480 continue481 # Remapping the name of FP8 kv-scale.482 name = maybe_remap_kv_scale_name(name, params_dict)483 if name is None:484 continue485 if is_pp_missing_parameter(name, self):486 continue487 if name not in params_dict:488 continue489 param = params_dict[name]490 weight_loader = getattr(param, "weight_loader", default_weight_loader)491 weight_loader(param, loaded_weight)492 loaded_params.add(name)493 return loaded_params494 495 496class NanbeigeForCausalLM(497 nn.Module, SupportsLoRA, SupportsPP, SupportsEagle, SupportsEagle3498):499 packed_modules_mapping = {500 "qkv_proj": [501 "q_proj",502 "k_proj",503 "v_proj",504 ],505 "gate_up_proj": [506 "gate_proj",507 "up_proj",508 ],509 }510 511 def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):512 super().__init__()513 config = vllm_config.model_config.hf_config.get_text_config()514 quant_config = vllm_config.quant_config515 516 self.config = config517 518 self.quant_config = quant_config519 self.model = NanbeigeModel(520 vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")521 )522 523 if get_pp_group().is_last_rank:524 if config.tie_word_embeddings:525 self.lm_head = self.model.embed_tokens526 else:527 self.lm_head = ParallelLMHead(528 config.vocab_size,529 config.hidden_size,530 quant_config=quant_config,531 prefix=maybe_prefix(prefix, "lm_head"),532 )533 else:534 self.lm_head = PPMissingLayer()535 536 self.logits_processor = LogitsProcessor(config.vocab_size)537 538 self.make_empty_intermediate_tensors = (539 self.model.make_empty_intermediate_tensors540 )541 542 def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:543 return self.model.embed_input_ids(input_ids)544 545 def forward(546 self,547 input_ids: torch.Tensor | None,548 positions: torch.Tensor,549 intermediate_tensors: IntermediateTensors | None = None,550 inputs_embeds: torch.Tensor | None = None,551 ) -> torch.Tensor | IntermediateTensors:552 hidden_states = self.model(553 input_ids, positions, intermediate_tensors, inputs_embeds554 )555 return hidden_states556 557 def compute_logits(558 self,559 hidden_states: torch.Tensor,560 ) -> torch.Tensor | None:561 logits = self.logits_processor(self.lm_head, hidden_states)562 return logits563 564 def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:565 loader = AutoWeightsLoader(566 self,567 skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None),568 )569 return loader.load_weights(weights)570 571 