CoolFace
Modelpublic

NullSense/Nanbeige4.2-3B-FP8-Dynamic

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
3likes130downloads
nanbeige.py571 linesDownload Raw Back to vllm_plugin
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