CoolFace
Modelpublic

LittleDesignSolution/Kimi-K2.5

sourceHugging Faceotherupdated 5mo agoView on Hugging Face
0likes17downloads
modeling_deepseek.py1809 linesDownload Raw Back to root
1# coding=utf-82# Copyright 2023 DeepSeek-AI and The HuggingFace Inc. team. All rights reserved.3#4# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX5# and OPT implementations in this library. It has been modified from its6# original forms to accommodate minor architectural differences compared7# to GPT-NeoX and OPT used by the Meta AI team that trained the model.8#9# Licensed under the Apache License, Version 2.0 (the "License");10# you may not use this file except in compliance with the License.11# You may obtain a copy of the License at12#13#     http://www.apache.org/licenses/LICENSE-2.014#15# Unless required by applicable law or agreed to in writing, software16# distributed under the License is distributed on an "AS IS" BASIS,17# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.18# See the License for the specific language governing permissions and19# limitations under the License.20""" PyTorch DeepSeek model."""21import math22import warnings23from typing import List, Optional, Tuple, Union24 25import numpy as np26import torch27import torch.distributed as dist28import torch.nn.functional as F29import torch.utils.checkpoint30from torch import nn31from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss32from transformers.activations import ACT2FN33from transformers.cache_utils import Cache, DynamicCache34from transformers.modeling_attn_mask_utils import \35    _prepare_4d_causal_attention_mask36from transformers.modeling_outputs import (BaseModelOutputWithPast,37                                           CausalLMOutputWithPast,38                                           SequenceClassifierOutputWithPast)39from transformers.modeling_utils import PreTrainedModel40from transformers.pytorch_utils import (ALL_LAYERNORM_LAYERS,41                                        is_torch_greater_or_equal_than_1_13)42from transformers.utils import (add_start_docstrings,43                                add_start_docstrings_to_model_forward,44                                is_flash_attn_2_available,45                                is_flash_attn_greater_or_equal_2_10, logging,46                                replace_return_docstrings)47from transformers.utils.import_utils import is_torch_fx_available48 49from .configuration_deepseek import DeepseekV3Config50 51if is_flash_attn_2_available():52    from flash_attn import flash_attn_func, flash_attn_varlen_func53    from flash_attn.bert_padding import pad_input  # noqa54    from flash_attn.bert_padding import index_first_axis, unpad_input55 56# This makes `_prepare_4d_causal_attention_mask` a leaf function in the FX graph.57# It means that the function will not be traced through and simply appear as a node in the graph.58if is_torch_fx_available():59    if not is_torch_greater_or_equal_than_1_13:60        import torch.fx61 62    _prepare_4d_causal_attention_mask = torch.fx.wrap(63        _prepare_4d_causal_attention_mask)64 65logger = logging.get_logger(__name__)66 67_CONFIG_FOR_DOC = "DeepseekV3Config"68 69 70def _get_unpad_data(attention_mask):71    seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)72    indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()73    max_seqlen_in_batch = seqlens_in_batch.max().item()74    cu_seqlens = F.pad(75        torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))76    return (77        indices,78        cu_seqlens,79        max_seqlen_in_batch,80    )81 82 83# code modified from transformers 4.48.3 to amend breaks in newer transformers versions84def get_usable_length(past_key_value,85                      new_seq_length: int,86                      layer_idx: Optional[int] = 0) -> int:87    max_length = past_key_value.get_max_cache_shape()88    previous_seq_length = past_key_value.get_seq_length(layer_idx)89    if max_length is not None and max_length > 0 and previous_seq_length + new_seq_length > max_length:90        return max_length - new_seq_length91    return previous_seq_length92 93 94class DeepseekV3RMSNorm(nn.Module):95 96    def __init__(self, hidden_size, eps=1e-6):97        """98        DeepseekV3RMSNorm is equivalent to T5LayerNorm99        """100        super().__init__()101        self.weight = nn.Parameter(torch.ones(hidden_size))102        self.variance_epsilon = eps103 104    def forward(self, hidden_states):105        input_dtype = hidden_states.dtype106        hidden_states = hidden_states.to(torch.float32)107        variance = hidden_states.pow(2).mean(-1, keepdim=True)108        hidden_states = hidden_states * torch.rsqrt(variance +109                                                    self.variance_epsilon)110        return self.weight * hidden_states.to(input_dtype)111 112 113ALL_LAYERNORM_LAYERS.append(DeepseekV3RMSNorm)114 115 116class DeepseekV3RotaryEmbedding(nn.Module):117 118    def __init__(self,119                 dim,120                 max_position_embeddings=2048,121                 base=10000,122                 device=None):123        super().__init__()124 125        self.dim = dim126        self.max_position_embeddings = max_position_embeddings127        self.base = base128        inv_freq = 1.0 / (self.base**(129            torch.arange(0, self.dim, 2).float().to(device) / self.dim))130        self.register_buffer("inv_freq", inv_freq, persistent=False)131 132        # Build here to make `torch.jit.trace` work.133        self._set_cos_sin_cache(134            seq_len=max_position_embeddings,135            device=self.inv_freq.device,136            dtype=torch.get_default_dtype(),137        )138        self.max_seq_len_cached = None139 140    def _set_cos_sin_cache(self, seq_len, device, dtype):141        self.max_seq_len_cached = seq_len142        t = torch.arange(self.max_seq_len_cached,143                         device=device,144                         dtype=self.inv_freq.dtype)145 146        freqs = torch.outer(t, self.inv_freq.to(t.device))147        # Different from paper, but it uses a different permutation in order to obtain the same calculation148        emb = torch.cat((freqs, freqs), dim=-1)149        self.register_buffer("cos_cached",150                             emb.cos().to(dtype),151                             persistent=False)152        self.register_buffer("sin_cached",153                             emb.sin().to(dtype),154                             persistent=False)155 156    def forward(self, x, seq_len=None):157        # x: [bs, num_attention_heads, seq_len, head_size]158        if self.max_seq_len_cached is None or seq_len > self.max_seq_len_cached:159            self._set_cos_sin_cache(seq_len=seq_len,160                                    device=x.device,161                                    dtype=x.dtype)162 163        return (164            self.cos_cached[:seq_len].to(dtype=x.dtype),165            self.sin_cached[:seq_len].to(dtype=x.dtype),166        )167 168 169# Copied from transformers.models.llama.modeling_llama.LlamaLinearScalingRotaryEmbedding with Llama->DeepseekV3170class DeepseekV3LinearScalingRotaryEmbedding(DeepseekV3RotaryEmbedding):171    """DeepseekV3RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""172 173    def __init__(174        self,175        dim,176        max_position_embeddings=2048,177        base=10000,178        device=None,179        scaling_factor=1.0,180    ):181        self.scaling_factor = scaling_factor182        super().__init__(dim, max_position_embeddings, base, device)183 184    def _set_cos_sin_cache(self, seq_len, device, dtype):185        self.max_seq_len_cached = seq_len186        t = torch.arange(self.max_seq_len_cached,187                         device=device,188                         dtype=self.inv_freq.dtype)189        t = t / self.scaling_factor190 191        freqs = torch.outer(t, self.inv_freq)192        # Different from paper, but it uses a different permutation in order to obtain the same calculation193        emb = torch.cat((freqs, freqs), dim=-1)194        self.register_buffer("cos_cached",195                             emb.cos().to(dtype),196                             persistent=False)197        self.register_buffer("sin_cached",198                             emb.sin().to(dtype),199                             persistent=False)200 201 202# Copied from transformers.models.llama.modeling_llama.LlamaDynamicNTKScalingRotaryEmbedding with Llama->DeepseekV3203class DeepseekV3DynamicNTKScalingRotaryEmbedding(DeepseekV3RotaryEmbedding):204    """DeepseekV3RotaryEmbedding extended with Dynamic NTK scaling. Credits to the Reddit users /u/bloc97 and /u/emozilla"""205 206    def __init__(207        self,208        dim,209        max_position_embeddings=2048,210        base=10000,211        device=None,212        scaling_factor=1.0,213    ):214        self.scaling_factor = scaling_factor215        super().__init__(dim, max_position_embeddings, base, device)216 217    def _set_cos_sin_cache(self, seq_len, device, dtype):218        self.max_seq_len_cached = seq_len219 220        if seq_len > self.max_position_embeddings:221            base = self.base * ((self.scaling_factor * seq_len /222                                 self.max_position_embeddings) -223                                (self.scaling_factor - 1))**(self.dim /224                                                             (self.dim - 2))225            inv_freq = 1.0 / (base**(226                torch.arange(0, self.dim, 2).float().to(device) / self.dim))227            self.register_buffer("inv_freq", inv_freq, persistent=False)228 229        t = torch.arange(self.max_seq_len_cached,230                         device=device,231                         dtype=self.inv_freq.dtype)232 233        freqs = torch.outer(t, self.inv_freq)234        # Different from paper, but it uses a different permutation in order to obtain the same calculation235        emb = torch.cat((freqs, freqs), dim=-1)236        self.register_buffer("cos_cached",237                             emb.cos().to(dtype),238                             persistent=False)239        self.register_buffer("sin_cached",240                             emb.sin().to(dtype),241                             persistent=False)242 243 244# Inverse dim formula to find dim based on number of rotations245def yarn_find_correction_dim(num_rotations,246                             dim,247                             base=10000,248                             max_position_embeddings=2048):249    return (dim * math.log(max_position_embeddings /250                           (num_rotations * 2 * math.pi))) / (2 *251                                                              math.log(base))252 253 254# Find dim range bounds based on rotations255def yarn_find_correction_range(low_rot,256                               high_rot,257                               dim,258                               base=10000,259                               max_position_embeddings=2048):260    low = math.floor(261        yarn_find_correction_dim(low_rot, dim, base, max_position_embeddings))262    high = math.ceil(263        yarn_find_correction_dim(high_rot, dim, base, max_position_embeddings))264    return max(low, 0), min(high, dim - 1)  # Clamp values just in case265 266 267def yarn_get_mscale(scale=1, mscale=1):268    if scale <= 1:269        return 1.0270    return 0.1 * mscale * math.log(scale) + 1.0271 272 273def yarn_linear_ramp_mask(min, max, dim):274    if min == max:275        max += 0.001  # Prevent singularity276 277    linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min)278    ramp_func = torch.clamp(linear_func, 0, 1)279    return ramp_func280 281 282class DeepseekV3YarnRotaryEmbedding(DeepseekV3RotaryEmbedding):283 284    def __init__(285        self,286        dim,287        max_position_embeddings=2048,288        base=10000,289        device=None,290        scaling_factor=1.0,291        original_max_position_embeddings=4096,292        beta_fast=32,293        beta_slow=1,294        mscale=1,295        mscale_all_dim=0,296    ):297        self.scaling_factor = scaling_factor298        self.original_max_position_embeddings = original_max_position_embeddings299        self.beta_fast = beta_fast300        self.beta_slow = beta_slow301        self.mscale = mscale302        self.mscale_all_dim = mscale_all_dim303        super().__init__(dim, max_position_embeddings, base, device)304 305    def _set_cos_sin_cache(self, seq_len, device, dtype):306        self.max_seq_len_cached = seq_len307        dim = self.dim308 309        freq_extra = 1.0 / (self.base**(310            torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))311        freq_inter = 1.0 / (self.scaling_factor * self.base**(312            torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))313 314        low, high = yarn_find_correction_range(315            self.beta_fast,316            self.beta_slow,317            dim,318            self.base,319            self.original_max_position_embeddings,320        )321        inv_freq_mask = 1.0 - yarn_linear_ramp_mask(low, high, dim // 2).to(322            device=device, dtype=torch.float32)323        inv_freq = freq_inter * (1 -324                                 inv_freq_mask) + freq_extra * inv_freq_mask325        self.register_buffer("inv_freq", inv_freq, persistent=False)326 327        t = torch.arange(seq_len, device=device, dtype=torch.float32)328 329        freqs = torch.outer(t, inv_freq)330 331        _mscale = float(332            yarn_get_mscale(self.scaling_factor, self.mscale) /333            yarn_get_mscale(self.scaling_factor, self.mscale_all_dim))334 335        emb = torch.cat((freqs, freqs), dim=-1)336        self.register_buffer("cos_cached", (emb.cos() * _mscale).to(dtype),337                             persistent=False)338        self.register_buffer("sin_cached", (emb.sin() * _mscale).to(dtype),339                             persistent=False)340 341 342# Copied from transformers.models.llama.modeling_llama.rotate_half343def rotate_half(x):344    """Rotates half the hidden dims of the input."""345    x1 = x[..., :x.shape[-1] // 2]346    x2 = x[..., x.shape[-1] // 2:]347    return torch.cat((-x2, x1), dim=-1)348 349 350# Copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb351def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):352    """Applies Rotary Position Embedding to the query and key tensors.353 354    Args:355        q (`torch.Tensor`): The query tensor.356        k (`torch.Tensor`): The key tensor.357        cos (`torch.Tensor`): The cosine part of the rotary embedding.358        sin (`torch.Tensor`): The sine part of the rotary embedding.359        position_ids (`torch.Tensor`):360            The position indices of the tokens corresponding to the query and key tensors. For example, this can be361            used to pass offsetted position ids when working with a KV-cache.362        unsqueeze_dim (`int`, *optional*, defaults to 1):363            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and364            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note365            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and366            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes367            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have368            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.369    Returns:370        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.371    """372    cos = cos[position_ids].unsqueeze(unsqueeze_dim)373    sin = sin[position_ids].unsqueeze(unsqueeze_dim)374 375    b, h, s, d = q.shape376    q = q.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)377 378    b, h, s, d = k.shape379    k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)380 381    q_embed = (q * cos) + (rotate_half(q) * sin)382    k_embed = (k * cos) + (rotate_half(k) * sin)383    return q_embed, k_embed384 385 386class DeepseekV3MLP(nn.Module):387 388    def __init__(self, config, hidden_size=None, intermediate_size=None):389        super().__init__()390        self.config = config391        self.hidden_size = config.hidden_size if hidden_size is None else hidden_size392        self.intermediate_size = (config.intermediate_size if intermediate_size393                                  is None else intermediate_size)394 395        self.gate_proj = nn.Linear(self.hidden_size,396                                   self.intermediate_size,397                                   bias=False)398        self.up_proj = nn.Linear(self.hidden_size,399                                 self.intermediate_size,400                                 bias=False)401        self.down_proj = nn.Linear(self.intermediate_size,402                                   self.hidden_size,403                                   bias=False)404        self.act_fn = ACT2FN[config.hidden_act]405 406    def forward(self, x):407        down_proj = self.down_proj(408            self.act_fn(self.gate_proj(x)) * self.up_proj(x))409        return down_proj410 411 412class MoEGate(nn.Module):413 414    def __init__(self, config):415        super().__init__()416        self.config = config417        self.top_k = config.num_experts_per_tok418        self.n_routed_experts = config.n_routed_experts419        self.routed_scaling_factor = config.routed_scaling_factor420        self.scoring_func = config.scoring_func421        self.seq_aux = config.seq_aux422        self.topk_method = config.topk_method423        self.n_group = config.n_group424        self.topk_group = config.topk_group425 426        # topk selection algorithm427        self.norm_topk_prob = config.norm_topk_prob428        self.gating_dim = config.hidden_size429        self.weight = nn.Parameter(430            torch.empty((self.n_routed_experts, self.gating_dim)))431        if self.topk_method == "noaux_tc":432            self.e_score_correction_bias = nn.Parameter(433                torch.empty((self.n_routed_experts)))434        self.reset_parameters()435 436    def reset_parameters(self) -> None:437        import torch.nn.init as init438 439        init.kaiming_uniform_(self.weight, a=math.sqrt(5))440 441    def forward(self, hidden_states):442        bsz, seq_len, h = hidden_states.shape443        ### compute gating score444        hidden_states = hidden_states.view(-1, h)445        logits = F.linear(hidden_states.type(torch.float32),446                          self.weight.type(torch.float32), None)447        if self.scoring_func == "sigmoid":448            scores = logits.sigmoid()449        else:450            raise NotImplementedError(451                f"insupportable scoring function for MoE gating: {self.scoring_func}"452            )453 454        ### select top-k experts455        if self.topk_method == "noaux_tc":456            assert not self.training457            scores_for_choice = scores.view(458                bsz * seq_len, -1) + self.e_score_correction_bias.unsqueeze(0)459            group_scores = (scores_for_choice.view(460                bsz * seq_len, self.n_group,461                -1).topk(2, dim=-1)[0].sum(dim=-1))  # [n, n_group]462            group_idx = torch.topk(group_scores,463                                   k=self.topk_group,464                                   dim=-1,465                                   sorted=False)[1]  # [n, top_k_group]466            group_mask = torch.zeros_like(group_scores)  # [n, n_group]467            group_mask.scatter_(1, group_idx, 1)  # [n, n_group]468            score_mask = (group_mask.unsqueeze(-1).expand(469                bsz * seq_len, self.n_group,470                self.n_routed_experts // self.n_group).reshape(471                    bsz * seq_len, -1))  # [n, e]472            tmp_scores = scores_for_choice.masked_fill(~score_mask.bool(),473                                                       0.0)  # [n, e]474            _, topk_idx = torch.topk(tmp_scores,475                                     k=self.top_k,476                                     dim=-1,477                                     sorted=False)478            topk_weight = scores.gather(1, topk_idx)479        else:480            raise NotImplementedError(481                f"insupportable TopK function for MoE gating: {self.topk_method}"482            )483 484        ### norm gate to sum 1485        if self.top_k > 1 and self.norm_topk_prob:486            denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20487            topk_weight = topk_weight / denominator488        topk_weight = topk_weight * self.routed_scaling_factor  # must multiply the scaling factor489 490        return topk_idx, topk_weight491 492 493class DeepseekV3MoE(nn.Module):494    """495    A mixed expert module containing shared experts.496    """497 498    def __init__(self, config):499        super().__init__()500        self.config = config501        self.num_experts_per_tok = config.num_experts_per_tok502 503        if hasattr(config, "ep_size") and config.ep_size > 1:504            assert config.ep_size == dist.get_world_size()505            self.ep_size = config.ep_size506            self.experts_per_rank = config.n_routed_experts // config.ep_size507            self.ep_rank = dist.get_rank()508            self.experts = nn.ModuleList([509                (DeepseekV3MLP(config,510                               intermediate_size=config.moe_intermediate_size)511                 if i >= self.ep_rank * self.experts_per_rank512                 and i < (self.ep_rank + 1) * self.experts_per_rank else None)513                for i in range(config.n_routed_experts)514            ])515        else:516            self.ep_size = 1517            self.experts_per_rank = config.n_routed_experts518            self.ep_rank = 0519            self.experts = nn.ModuleList([520                DeepseekV3MLP(config,521                              intermediate_size=config.moe_intermediate_size)522                for i in range(config.n_routed_experts)523            ])524        self.gate = MoEGate(config)525        if config.n_shared_experts is not None:526            intermediate_size = config.moe_intermediate_size * config.n_shared_experts527            self.shared_experts = DeepseekV3MLP(528                config=config, intermediate_size=intermediate_size)529 530    def forward(self, hidden_states):531        identity = hidden_states532        orig_shape = hidden_states.shape533        topk_idx, topk_weight = self.gate(hidden_states)534        hidden_states = hidden_states.view(-1, hidden_states.shape[-1])535        flat_topk_idx = topk_idx.view(-1)536        if not self.training:537            y = self.moe_infer(hidden_states, topk_idx,538                               topk_weight).view(*orig_shape)539        if self.config.n_shared_experts is not None:540            y = y + self.shared_experts(identity)541        return y542 543    @torch.no_grad()544    def moe_infer(self, x, topk_ids, topk_weight):545        cnts = topk_ids.new_zeros((topk_ids.shape[0], len(self.experts)))546        cnts.scatter_(1, topk_ids, 1)547        tokens_per_expert = cnts.sum(dim=0)548        idxs = topk_ids.view(-1).argsort()549        sorted_tokens = x[idxs // topk_ids.shape[1]]550        sorted_tokens_shape = sorted_tokens.shape551        if self.ep_size > 1:552            tokens_per_ep_rank = tokens_per_expert.view(self.ep_size,553                                                        -1).sum(dim=1)554            tokens_per_expert_group = tokens_per_expert.new_empty(555                tokens_per_expert.shape[0])556            dist.all_to_all_single(tokens_per_expert_group, tokens_per_expert)557            output_splits = (tokens_per_expert_group.view(558                self.ep_size, -1).sum(1).cpu().numpy().tolist())559            gathered_tokens = sorted_tokens.new_empty(560                tokens_per_expert_group.sum(dim=0).cpu().item(),561                sorted_tokens.shape[1])562            input_split_sizes = tokens_per_ep_rank.cpu().numpy().tolist()563            dist.all_to_all(564                list(gathered_tokens.split(output_splits)),565                list(sorted_tokens.split(input_split_sizes)),566            )567            tokens_per_expert_post_gather = tokens_per_expert_group.view(568                self.ep_size, self.experts_per_rank).sum(dim=0)569            gatherd_idxs = np.zeros(shape=(gathered_tokens.shape[0], ),570                                    dtype=np.int32)571            s = 0572            for i, k in enumerate(tokens_per_expert_group.cpu().numpy()):573                gatherd_idxs[s:s + k] = i % self.experts_per_rank574                s += k575            gatherd_idxs = gatherd_idxs.argsort()576            sorted_tokens = gathered_tokens[gatherd_idxs]577            tokens_per_expert = tokens_per_expert_post_gather578        tokens_per_expert = tokens_per_expert.cpu().numpy()579 580        outputs = []581        start_idx = 0582        for i, num_tokens in enumerate(tokens_per_expert):583            end_idx = start_idx + num_tokens584            if num_tokens == 0:585                continue586            expert = self.experts[i + self.ep_rank * self.experts_per_rank]587            tokens_for_this_expert = sorted_tokens[start_idx:end_idx]588            expert_out = expert(tokens_for_this_expert)589            outputs.append(expert_out)590            start_idx = end_idx591 592        outs = torch.cat(outputs,593                         dim=0) if len(outputs) else sorted_tokens.new_empty(0)594        if self.ep_size > 1:595            new_x = torch.empty_like(outs)596            new_x[gatherd_idxs] = outs597            gathered_tokens = new_x.new_empty(*sorted_tokens_shape)598            dist.all_to_all(599                list(gathered_tokens.split(input_split_sizes)),600                list(new_x.split(output_splits)),601            )602            outs = gathered_tokens603 604        new_x = torch.empty_like(outs)605        new_x[idxs] = outs606        final_out = (new_x.view(607            *topk_ids.shape, -1).type(topk_weight.dtype).mul_(608                topk_weight.unsqueeze(dim=-1)).sum(dim=1).type(new_x.dtype))609        return final_out610 611 612# Copied from transformers.models.llama.modeling_llama.repeat_kv613def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:614    """615    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,616    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)617    """618    batch, num_key_value_heads, slen, head_dim = hidden_states.shape619    if n_rep == 1:620        return hidden_states621    hidden_states = hidden_states[:, :,622                                  None, :, :].expand(batch,623                                                     num_key_value_heads,624                                                     n_rep, slen, head_dim)625    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen,626                                 head_dim)627 628 629# Copied from transformers.models.llama.modeling_llama.LlamaAttention with Llama->DeepseekV3630class DeepseekV3Attention(nn.Module):631    """Multi-headed attention from 'Attention Is All You Need' paper"""632 633    def __init__(self,634                 config: DeepseekV3Config,635                 layer_idx: Optional[int] = None):636        super().__init__()637        self.config = config638        self.layer_idx = layer_idx639        if layer_idx is None:640            logger.warning_once(641                f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will "642                "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` "643                "when creating this class.")644 645        self.attention_dropout = config.attention_dropout646        self.hidden_size = config.hidden_size647        self.num_heads = config.num_attention_heads648 649        self.max_position_embeddings = config.max_position_embeddings650        self.rope_theta = config.rope_theta651        self.q_lora_rank = config.q_lora_rank652        self.qk_rope_head_dim = config.qk_rope_head_dim653        self.kv_lora_rank = config.kv_lora_rank654        self.v_head_dim = config.v_head_dim655        self.qk_nope_head_dim = config.qk_nope_head_dim656        self.q_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim657 658        self.is_causal = True659 660        if self.q_lora_rank is None:661            self.q_proj = nn.Linear(self.hidden_size,662                                    self.num_heads * self.q_head_dim,663                                    bias=False)664        else:665            self.q_a_proj = nn.Linear(self.hidden_size,666                                      config.q_lora_rank,667                                      bias=config.attention_bias)668            self.q_a_layernorm = DeepseekV3RMSNorm(config.q_lora_rank)669            self.q_b_proj = nn.Linear(config.q_lora_rank,670                                      self.num_heads * self.q_head_dim,671                                      bias=False)672 673        self.kv_a_proj_with_mqa = nn.Linear(674            self.hidden_size,675            config.kv_lora_rank + config.qk_rope_head_dim,676            bias=config.attention_bias,677        )678        self.kv_a_layernorm = DeepseekV3RMSNorm(config.kv_lora_rank)679        self.kv_b_proj = nn.Linear(680            config.kv_lora_rank,681            self.num_heads *682            (self.q_head_dim - self.qk_rope_head_dim + self.v_head_dim),683            bias=False,684        )685 686        self.o_proj = nn.Linear(687            self.num_heads * self.v_head_dim,688            self.hidden_size,689            bias=config.attention_bias,690        )691        self._init_rope()692 693        self.softmax_scale = self.q_head_dim**(-0.5)694        if self.config.rope_scaling is not None:695            mscale_all_dim = self.config.rope_scaling.get("mscale_all_dim", 0)696            scaling_factor = self.config.rope_scaling["factor"]697            if mscale_all_dim:698                mscale = yarn_get_mscale(scaling_factor, mscale_all_dim)699                self.softmax_scale = self.softmax_scale * mscale * mscale700 701    def _init_rope(self):702        if self.config.rope_scaling is None:703            self.rotary_emb = DeepseekV3RotaryEmbedding(704                self.qk_rope_head_dim,705                max_position_embeddings=self.max_position_embeddings,706                base=self.rope_theta,707            )708        else:709            scaling_type = self.config.rope_scaling["type"]710            scaling_factor = self.config.rope_scaling["factor"]711            if scaling_type == "linear":712                self.rotary_emb = DeepseekV3LinearScalingRotaryEmbedding(713                    self.qk_rope_head_dim,714                    max_position_embeddings=self.max_position_embeddings,715                    scaling_factor=scaling_factor,716                    base=self.rope_theta,717                )718            elif scaling_type == "dynamic":719                self.rotary_emb = DeepseekV3DynamicNTKScalingRotaryEmbedding(720                    self.qk_rope_head_dim,721                    max_position_embeddings=self.max_position_embeddings,722                    scaling_factor=scaling_factor,723                    base=self.rope_theta,724                )725            elif scaling_type == "yarn":726                kwargs = {727                    key: self.config.rope_scaling[key]728                    for key in [729                        "original_max_position_embeddings",730                        "beta_fast",731                        "beta_slow",732                        "mscale",733                        "mscale_all_dim",734                    ] if key in self.config.rope_scaling735                }736                self.rotary_emb = DeepseekV3YarnRotaryEmbedding(737                    self.qk_rope_head_dim,738                    max_position_embeddings=self.max_position_embeddings,739                    scaling_factor=scaling_factor,740                    base=self.rope_theta,741                    **kwargs,742                )743            else:744                raise ValueError(f"Unknown RoPE scaling type {scaling_type}")745 746    def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):747        return (tensor.view(bsz, seq_len, self.num_heads,748                            self.v_head_dim).transpose(1, 2).contiguous())749 750    def forward(751        self,752        hidden_states: torch.Tensor,753        attention_mask: Optional[torch.Tensor] = None,754        position_ids: Optional[torch.LongTensor] = None,755        past_key_value: Optional[Cache] = None,756        output_attentions: bool = False,757        use_cache: bool = False,758        **kwargs,759    ) -> Tuple[torch.Tensor, Optional[torch.Tensor],760               Optional[Tuple[torch.Tensor]]]:761        if "padding_mask" in kwargs:762            warnings.warn(763                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"764            )765        bsz, q_len, _ = hidden_states.size()766 767        if self.q_lora_rank is None:768            q = self.q_proj(hidden_states)769        else:770            q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))771        q = q.view(bsz, q_len, self.num_heads, self.q_head_dim).transpose(1, 2)772        q_nope, q_pe = torch.split(773            q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)774 775        compressed_kv = self.kv_a_proj_with_mqa(hidden_states)776        compressed_kv, k_pe = torch.split(777            compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)778        k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2)779        kv = (self.kv_b_proj(self.kv_a_layernorm(compressed_kv)).view(780            bsz, q_len, self.num_heads,781            self.qk_nope_head_dim + self.v_head_dim).transpose(1, 2))782 783        k_nope, value_states = torch.split(784            kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)785        kv_seq_len = value_states.shape[-2]786        if past_key_value is not None:787            if self.layer_idx is None:788                raise ValueError(789                    f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "790                    "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "791                    "with a layer index.")792            kv_seq_len += get_usable_length(past_key_value, kv_seq_len,793                                            self.layer_idx)794        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)795 796        q_pe, k_pe = apply_rotary_pos_emb(q_pe, k_pe, cos, sin, position_ids)797 798        query_states = k_pe.new_empty(bsz, self.num_heads, q_len,799                                      self.q_head_dim)800        query_states[:, :, :, :self.qk_nope_head_dim] = q_nope801        query_states[:, :, :, self.qk_nope_head_dim:] = q_pe802 803        key_states = k_pe.new_empty(bsz, self.num_heads, q_len,804                                    self.q_head_dim)805        key_states[:, :, :, :self.qk_nope_head_dim] = k_nope806        key_states[:, :, :, self.qk_nope_head_dim:] = k_pe807        if past_key_value is not None:808            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models809            key_states, value_states = past_key_value.update(810                key_states, value_states, self.layer_idx, cache_kwargs)811 812        attn_weights = (813            torch.matmul(query_states, key_states.transpose(2, 3)) *814            self.softmax_scale)815 816        if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):817            raise ValueError(818                f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"819                f" {attn_weights.size()}")820        assert attention_mask is not None821        if attention_mask is not None:822            if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):823                raise ValueError(824                    f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"825                )826            attn_weights = attn_weights + attention_mask827 828        # upcast attention to fp32829        attn_weights = nn.functional.softmax(attn_weights,830                                             dim=-1,831                                             dtype=torch.float32).to(832                                                 query_states.dtype)833        attn_weights = nn.functional.dropout(attn_weights,834                                             p=self.attention_dropout,835                                             training=self.training)836        attn_output = torch.matmul(attn_weights, value_states)837 838        if attn_output.size() != (bsz, self.num_heads, q_len, self.v_head_dim):839            raise ValueError(840                f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.v_head_dim)}, but is"841                f" {attn_output.size()}")842 843        attn_output = attn_output.transpose(1, 2).contiguous()844 845        attn_output = attn_output.reshape(bsz, q_len,846                                          self.num_heads * self.v_head_dim)847 848        attn_output = self.o_proj(attn_output)849 850        if not output_attentions:851            attn_weights = None852 853        return attn_output, attn_weights, past_key_value854 855 856# Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2 with Llama->DeepseekV3857class DeepseekV3FlashAttention2(DeepseekV3Attention):858    """859    DeepseekV3 flash attention module. This module inherits from `DeepseekV3Attention` as the weights of the module stays860    untouched. The only required change would be on the forward pass where it needs to correctly call the public API of861    flash attention and deal with padding tokens in case the input contains any of them.862    """863 864    def __init__(self, *args, **kwargs):865        super().__init__(*args, **kwargs)866 867        # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.868        # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignment, that was made default for flash_attn>=2.1. This attribute is used to handle this difference. Reference: https://github.com/Dao-AILab/flash-attention/releases/tag/v2.1.0.869        # Beware that with flash_attn<2.1, using q_seqlen != k_seqlen (except for the case q_seqlen == 1) produces a wrong mask (top-left).870        self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10(871        )872 873    def forward(874        self,875        hidden_states: torch.Tensor,876        attention_mask: Optional[torch.LongTensor] = None,877        position_ids: Optional[torch.LongTensor] = None,878        past_key_value: Optional[Cache] = None,879        output_attentions: bool = False,880        use_cache: bool = False,881        **kwargs,882    ) -> Tuple[torch.Tensor, Optional[torch.Tensor],883               Optional[Tuple[torch.Tensor]]]:884        # DeepseekV3FlashAttention2 attention does not support output_attentions885        if "padding_mask" in kwargs:886            warnings.warn(887                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"888            )889 890            # overwrite attention_mask with padding_mask891            attention_mask = kwargs.pop("padding_mask")892 893        output_attentions = False894 895        bsz, q_len, _ = hidden_states.size()896 897        if self.q_lora_rank is None:898            q = self.q_proj(hidden_states)899        else:900            q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))901        q = q.view(bsz, q_len, self.num_heads, self.q_head_dim).transpose(1, 2)902        q_nope, q_pe = torch.split(903            q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)904 905        # Flash attention requires the input to have the shape906        # batch_size x seq_length x head_dim x hidden_dim907        # therefore we just need to keep the original shape908        compressed_kv = self.kv_a_proj_with_mqa(hidden_states)909        compressed_kv, k_pe = torch.split(910            compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)911        k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2)912        kv = (self.kv_b_proj(self.kv_a_layernorm(compressed_kv)).view(913            bsz, q_len, self.num_heads,914            self.qk_nope_head_dim + self.v_head_dim).transpose(1, 2))915 916        k_nope, value_states = torch.split(917            kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)918        kv_seq_len = value_states.shape[-2]919 920        kv_seq_len = value_states.shape[-2]921        if past_key_value is not None:922            kv_seq_len += get_usable_length(past_key_value, kv_seq_len,923                                            self.layer_idx)924 925        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)926        q_pe, k_pe = apply_rotary_pos_emb(q_pe, k_pe, cos, sin, position_ids)927 928        query_states = k_pe.new_empty(bsz, self.num_heads, q_len,929                                      self.q_head_dim)930        query_states[:, :, :, :self.qk_nope_head_dim] = q_nope931        query_states[:, :, :, self.qk_nope_head_dim:] = q_pe932 933        key_states = k_pe.new_empty(bsz, self.num_heads, q_len,934                                    self.q_head_dim)935        key_states[:, :, :, :self.qk_nope_head_dim] = k_nope936        key_states[:, :, :, self.qk_nope_head_dim:] = k_pe937 938        if self.q_head_dim != self.v_head_dim:939            value_states = F.pad(value_states,940                                 [0, self.q_head_dim - self.v_head_dim])941 942        if past_key_value is not None:943            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models944            key_states, value_states = past_key_value.update(945                key_states, value_states, self.layer_idx, cache_kwargs)946 947        # TODO: These transpose are quite inefficient but Flash Attention requires the layout [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache948        # to be able to avoid many of these transpose/reshape/view.949        query_states = query_states.transpose(1, 2)950        key_states = key_states.transpose(1, 2)951        value_states = value_states.transpose(1, 2)952 953        dropout_rate = self.attention_dropout if self.training else 0.0954 955        # In PEFT, usually we cast the layer norms in float32 for training stability reasons956        # therefore the input hidden states gets silently casted in float32. Hence, we need957        # cast them back in the correct dtype just to be sure everything works as expected.958        # This might slowdown training & inference so it is recommended to not cast the LayerNorms959        # in fp32. (DeepseekV3RMSNorm handles it correctly)960 961        input_dtype = query_states.dtype962        if input_dtype == torch.float32:963            # Handle the case where the model is quantized964            if hasattr(self.config, "_pre_quantization_dtype"):965                target_dtype = self.config._pre_quantization_dtype966            elif torch.is_autocast_enabled():967                target_dtype = torch.get_autocast_gpu_dtype()968            else:969                target_dtype = (self.q_proj.weight.dtype if self.q_lora_rank970                                is None else self.q_a_proj.weight.dtype)971 972            logger.warning_once(973                f"The input hidden states seems to be silently casted in float32, this might be related to"974                f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"975                f" {target_dtype}.")976 977            query_states = query_states.to(target_dtype)978            key_states = key_states.to(target_dtype)979            value_states = value_states.to(target_dtype)980 981        attn_output = self._flash_attention_forward(982            query_states,983            key_states,984            value_states,985            attention_mask,986            q_len,987            dropout=dropout_rate,988            softmax_scale=self.softmax_scale,989        )990        if self.q_head_dim != self.v_head_dim:991            attn_output = attn_output[:, :, :, :self.v_head_dim]992 993        attn_output = attn_output.reshape(bsz, q_len, self.num_heads *994                                          self.v_head_dim).contiguous()995        attn_output = self.o_proj(attn_output)996 997        if not output_attentions:998            attn_weights = None999 1000        return attn_output, attn_weights, past_key_value1001 1002    def _flash_attention_forward(1003        self,1004        query_states,1005        key_states,1006        value_states,1007        attention_mask,1008        query_length,1009        dropout=0.0,1010        softmax_scale=None,1011    ):1012        """1013        Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token1014        first unpad the input, then computes the attention scores and pad the final attention scores.1015 1016        Args:1017            query_states (`torch.Tensor`):1018                Input query states to be passed to Flash Attention API1019            key_states (`torch.Tensor`):1020                Input key states to be passed to Flash Attention API1021            value_states (`torch.Tensor`):1022                Input value states to be passed to Flash Attention API1023            attention_mask (`torch.Tensor`):1024                The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the1025                position of padding tokens and 1 for the position of non-padding tokens.1026            dropout (`int`, *optional*):1027                Attention dropout1028            softmax_scale (`float`, *optional*):1029                The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)1030        """1031        if not self._flash_attn_uses_top_left_mask:1032            causal = self.is_causal1033        else:1034            # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1. For details, please see the comment in DeepseekV3FlashAttention2 __init__.1035            causal = self.is_causal and query_length != 11036 1037        # Contains at least one padding token in the sequence1038        if attention_mask is not None:1039            batch_size = query_states.shape[0]1040            (1041                query_states,1042                key_states,1043                value_states,1044                indices_q,1045                cu_seq_lens,1046                max_seq_lens,1047            ) = self._upad_input(query_states, key_states, value_states,1048                                 attention_mask, query_length)1049 1050            cu_seqlens_q, cu_seqlens_k = cu_seq_lens1051            max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens1052 1053            attn_output_unpad = flash_attn_varlen_func(1054                query_states,1055                key_states,1056                value_states,1057                cu_seqlens_q=cu_seqlens_q,1058                cu_seqlens_k=cu_seqlens_k,1059                max_seqlen_q=max_seqlen_in_batch_q,1060                max_seqlen_k=max_seqlen_in_batch_k,1061                dropout_p=dropout,1062                softmax_scale=softmax_scale,1063                causal=causal,1064            )1065 1066            attn_output = pad_input(attn_output_unpad, indices_q, batch_size,1067                                    query_length)1068        else:1069            attn_output = flash_attn_func(1070                query_states,1071                key_states,1072                value_states,1073                dropout,1074                softmax_scale=softmax_scale,1075                causal=causal,1076            )1077 1078        return attn_output1079 1080    def _upad_input(self, query_layer, key_layer, value_layer, attention_mask,1081                    query_length):1082        indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(1083            attention_mask)1084        batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape1085 1086        key_layer = index_first_axis(1087            key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads,1088                              head_dim),1089            indices_k,1090        )1091        value_layer = index_first_axis(1092            value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads,1093                                head_dim),1094            indices_k,1095        )1096        if query_length == kv_seq_len:1097            query_layer = index_first_axis(1098                query_layer.reshape(batch_size * kv_seq_len, self.num_heads,1099                                    head_dim),1100                indices_k,1101            )1102            cu_seqlens_q = cu_seqlens_k1103            max_seqlen_in_batch_q = max_seqlen_in_batch_k1104            indices_q = indices_k1105        elif query_length == 1:1106            max_seqlen_in_batch_q = 11107            cu_seqlens_q = torch.arange(1108                batch_size + 1, dtype=torch.int32, device=query_layer.device1109            )  # There is a memcpy here, that is very bad.1110            indices_q = cu_seqlens_q[:-1]1111            query_layer = query_layer.squeeze(1)1112        else:1113            # The -q_len: slice assumes left padding.1114            attention_mask = attention_mask[:, -query_length:]1115            query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(1116                query_layer, attention_mask)1117 1118        return (1119            query_layer,1120            key_layer,1121            value_layer,1122            indices_q,1123            (cu_seqlens_q, cu_seqlens_k),1124            (max_seqlen_in_batch_q, max_seqlen_in_batch_k),1125        )1126 1127 1128ATTENTION_CLASSES = {1129    "eager": DeepseekV3Attention,1130    "flash_attention_2": DeepseekV3FlashAttention2,1131}1132 1133 1134class DeepseekV3DecoderLayer(nn.Module):1135 1136    def __init__(self, config: DeepseekV3Config, layer_idx: int):1137        super().__init__()1138        self.hidden_size = config.hidden_size1139 1140        self.self_attn = ATTENTION_CLASSES[config._attn_implementation](1141            config=config, layer_idx=layer_idx)1142 1143        self.mlp = (DeepseekV3MoE(config) if1144                    (config.n_routed_experts is not None1145                     and layer_idx >= config.first_k_dense_replace1146                     and layer_idx % config.moe_layer_freq == 0) else1147                    DeepseekV3MLP(config))1148        self.input_layernorm = DeepseekV3RMSNorm(config.hidden_size,1149                                                 eps=config.rms_norm_eps)1150        self.post_attention_layernorm = DeepseekV3RMSNorm(1151            config.hidden_size, eps=config.rms_norm_eps)1152 1153    def forward(1154        self,1155        hidden_states: torch.Tensor,1156        attention_mask: Optional[torch.Tensor] = None,1157        position_ids: Optional[torch.LongTensor] = None,1158        past_key_value: Optional[Tuple[torch.Tensor]] = None,1159        output_attentions: Optional[bool] = False,1160        use_cache: Optional[bool] = False,1161        **kwargs,1162    ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor,1163                                                 torch.FloatTensor]]]:1164        """1165        Args:1166            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`1167            attention_mask (`torch.FloatTensor`, *optional*):1168                attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,1169                query_sequence_length, key_sequence_length)` if default attention is used.1170            output_attentions (`bool`, *optional*):1171                Whether or not to return the attentions tensors of all attention layers. See `attentions` under1172                returned tensors for more detail.1173            use_cache (`bool`, *optional*):1174                If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding1175                (see `past_key_values`).1176            past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states1177        """1178        if "padding_mask" in kwargs:1179            warnings.warn(1180                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"1181            )1182        residual = hidden_states1183 1184        hidden_states = self.input_layernorm(hidden_states)1185 1186        # Self Attention1187        hidden_states, self_attn_weights, present_key_value = self.self_attn(1188            hidden_states=hidden_states,1189            attention_mask=attention_mask,1190            position_ids=position_ids,1191            past_key_value=past_key_value,1192            output_attentions=output_attentions,1193            use_cache=use_cache,1194            **kwargs,1195        )1196        hidden_states = residual + hidden_states1197 1198        # Fully Connected1199        residual = hidden_states1200        hidden_states = self.post_attention_layernorm(hidden_states)

Showing the first 1,200 of 1809 lines. Download the file for the rest.