LittleDesignSolution/Kimi-K2.5
017
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)