llmware/slim-boolean-phi-3
020
1# coding=utf-82# Copyright 2024 Microsoft and the HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15 16""" PyTorch Phi-3 model."""17 18import inspect19import math20import warnings21from typing import List, Optional, Tuple, Union22 23import torch24import torch.nn.functional as F25import torch.utils.checkpoint26from torch import nn27from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss28 29from transformers.activations import ACT2FN30from transformers.cache_utils import Cache, DynamicCache31from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask32from transformers.modeling_outputs import (33 BaseModelOutputWithPast,34 CausalLMOutputWithPast,35 SequenceClassifierOutputWithPast,36 TokenClassifierOutput,37)38from transformers.modeling_utils import PreTrainedModel39from transformers.utils import (40 add_code_sample_docstrings,41 add_start_docstrings,42 add_start_docstrings_to_model_forward,43 is_flash_attn_2_available,44 is_flash_attn_greater_or_equal_2_10,45 logging,46 replace_return_docstrings,47)48from .configuration_phi3 import Phi3Config49 50 51logger = logging.get_logger(__name__)52 53# Transformers scans dependencies in the modeling file, causing issues on conditional loading. The regex only ignores try/catch blocks, but not if statements54# if is_flash_attn_2_available():55_flash_supports_window_size = False56try:57 from flash_attn import flash_attn_func, flash_attn_varlen_func58 from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa59 60 _flash_supports_window_size = "window_size" in list(inspect.signature(flash_attn_func).parameters)61except ImportError as error:62 logger.warning(63 f"`flash-attention` package not found, consider installing for better performance: {error}."64 )65 if not _flash_supports_window_size:66 logger.warning(67 "Current `flash-attenton` does not support `window_size`. Either upgrade or use `attn_implementation='eager'`."68 )69 70_CHECKPOINT_FOR_DOC = "microsoft/Phi-3-mini-4k-instruct"71_CONFIG_FOR_DOC = "Phi3Config"72 73PHI3_PRETRAINED_MODEL_ARCHIVE_LIST = [74 "microsoft/Phi-3-mini-4k-instruct",75 "microsoft/Phi-3-mini-128k-instruct",76 # See all Phi-3 models at https://huggingface.co/models?filter=Phi-377]78 79 80# Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->Phi381class Phi3RMSNorm(nn.Module):82 def __init__(self, hidden_size, eps=1e-6):83 """84 Phi3RMSNorm is equivalent to T5LayerNorm85 """86 super().__init__()87 self.weight = nn.Parameter(torch.ones(hidden_size))88 self.variance_epsilon = eps89 90 def forward(self, hidden_states):91 input_dtype = hidden_states.dtype92 hidden_states = hidden_states.to(torch.float32)93 variance = hidden_states.pow(2).mean(-1, keepdim=True)94 hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)95 return self.weight * hidden_states.to(input_dtype)96 97 98# Copied from transformers.models.llama.modeling_llama._get_unpad_data99def _get_unpad_data(attention_mask):100 seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)101 indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()102 max_seqlen_in_batch = seqlens_in_batch.max().item()103 cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))104 return (105 indices,106 cu_seqlens,107 max_seqlen_in_batch,108 )109 110 111# Copied from transformers.models.gemma.modeling_gemma.GemmaRotaryEmbedding with gemma->phi3, Gemma->Phi3112class Phi3RotaryEmbedding(nn.Module):113 def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):114 super().__init__()115 116 self.dim = dim117 self.max_position_embeddings = max_position_embeddings118 self.base = base119 self.register_buffer("inv_freq", None, persistent=False)120 121 @torch.no_grad()122 def forward(self, x, position_ids, seq_len=None):123 # x: [bs, num_attention_heads, seq_len, head_size]124 if self.inv_freq is None:125 self.inv_freq = 1.0 / (126 self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device=x.device).float() / self.dim)127 )128 inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)129 position_ids_expanded = position_ids[:, None, :].float()130 # Force float32 since bfloat16 loses precision on long contexts131 # See https://github.com/huggingface/transformers/pull/29285132 device_type = x.device.type133 device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"134 with torch.autocast(device_type=device_type, enabled=False):135 freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)136 emb = torch.cat((freqs, freqs), dim=-1)137 cos = emb.cos()138 sin = emb.sin()139 return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)140 141 142class Phi3SuScaledRotaryEmbedding(Phi3RotaryEmbedding):143 def __init__(144 self,145 dim,146 short_factor,147 long_factor,148 original_max_position_embeddings=2048,149 max_position_embeddings=2048,150 base=10000,151 device=None,152 ):153 super().__init__(dim, max_position_embeddings, base, device)154 155 self.short_factor = short_factor156 self.long_factor = long_factor157 self.original_max_position_embeddings = original_max_position_embeddings158 159 def _calc_scaling_factor(self, scale):160 if scale <= 1.0:161 return 1.0162 return math.sqrt(1 + math.log(scale) / math.log(self.original_max_position_embeddings))163 164 @torch.no_grad()165 def forward(self, x, position_ids, seq_len=None):166 seq_len = torch.max(position_ids) + 1167 if seq_len > self.original_max_position_embeddings:168 ext_factors = torch.tensor(self.long_factor, dtype=torch.float32, device=x.device)169 else:170 ext_factors = torch.tensor(self.short_factor, dtype=torch.float32, device=x.device)171 172 self.inv_freq = 1.0 / (173 ext_factors174 * self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device=x.device).float() / self.dim)175 )176 inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)177 position_ids_expanded = position_ids[:, None, :].float()178 179 # Force float32 since bfloat16 loses precision on long contexts180 # See https://github.com/huggingface/transformers/pull/29285181 device_type = x.device.type182 device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"183 with torch.autocast(device_type=device_type, enabled=False):184 freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)185 scaling_factor = self._calc_scaling_factor(186 self.max_position_embeddings / self.original_max_position_embeddings187 )188 emb = torch.cat((freqs, freqs), dim=-1)189 cos = emb.cos() * scaling_factor190 sin = emb.sin() * scaling_factor191 return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)192 193 194class Phi3YarnScaledRotaryEmbedding(Phi3RotaryEmbedding):195 def __init__(196 self,197 dim,198 short_factor,199 long_factor,200 original_max_position_embeddings=2048,201 max_position_embeddings=2048,202 base=10000,203 device=None,204 ):205 super().__init__(dim, max_position_embeddings, base, device)206 207 self.short_factor = short_factor208 self.long_factor = long_factor209 self.original_max_position_embeddings = original_max_position_embeddings210 211 def _calc_scaling_factor(self, scale):212 if scale <= 1.0:213 return 1.0214 return 0.1 * math.log(scale) + 1.0215 216 @torch.no_grad()217 def forward(self, x, position_ids, seq_len=None):218 seq_len = torch.max(position_ids) + 1219 if seq_len > self.original_max_position_embeddings:220 ext_factors = torch.tensor(self.long_factor, dtype=torch.float32, device=x.device)221 else:222 ext_factors = torch.tensor(self.short_factor, dtype=torch.float32, device=x.device)223 224 self.inv_freq = 1.0 / (225 ext_factors226 * self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device=x.device).float() / self.dim)227 )228 inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)229 position_ids_expanded = position_ids[:, None, :].float()230 231 # Force float32 since bfloat16 loses precision on long contexts232 # See https://github.com/huggingface/transformers/pull/29285233 device_type = x.device.type234 device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"235 with torch.autocast(device_type=device_type, enabled=False):236 freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)237 scaling_factor = self._calc_scaling_factor(238 self.max_position_embeddings / self.original_max_position_embeddings239 )240 emb = torch.cat((freqs, freqs), dim=-1)241 cos = emb.cos() * scaling_factor242 sin = emb.sin() * scaling_factor243 return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)244 245 246# Copied from transformers.models.llama.modeling_llama.rotate_half247def rotate_half(x):248 """Rotates half the hidden dims of the input."""249 x1 = x[..., : x.shape[-1] // 2]250 x2 = x[..., x.shape[-1] // 2 :]251 return torch.cat((-x2, x1), dim=-1)252 253 254# Copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb255def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):256 """Applies Rotary Position Embedding to the query and key tensors.257 258 Args:259 q (`torch.Tensor`): The query tensor.260 k (`torch.Tensor`): The key tensor.261 cos (`torch.Tensor`): The cosine part of the rotary embedding.262 sin (`torch.Tensor`): The sine part of the rotary embedding.263 position_ids (`torch.Tensor`, *optional*):264 Deprecated and unused.265 unsqueeze_dim (`int`, *optional*, defaults to 1):266 The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and267 sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note268 that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and269 k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes270 cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have271 the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.272 Returns:273 `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.274 """275 cos = cos.unsqueeze(unsqueeze_dim)276 sin = sin.unsqueeze(unsqueeze_dim)277 q_embed = (q * cos) + (rotate_half(q) * sin)278 k_embed = (k * cos) + (rotate_half(k) * sin)279 return q_embed, k_embed280 281 282class Phi3MLP(nn.Module):283 def __init__(self, config):284 super().__init__()285 286 self.config = config287 self.gate_up_proj = nn.Linear(config.hidden_size, 2 * config.intermediate_size, bias=False)288 self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)289 290 self.activation_fn = ACT2FN[config.hidden_act]291 292 def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:293 up_states = self.gate_up_proj(hidden_states)294 295 gate, up_states = up_states.chunk(2, dim=-1)296 up_states = up_states * self.activation_fn(gate)297 298 return self.down_proj(up_states)299 300 301# Copied from transformers.models.llama.modeling_llama.repeat_kv with llama->phi302def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:303 """304 This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,305 num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)306 """307 batch, num_key_value_heads, slen, head_dim = hidden_states.shape308 if n_rep == 1:309 return hidden_states310 hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)311 return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)312 313 314class Phi3Attention(nn.Module):315 """Multi-headed attention from 'Attention Is All You Need' paper"""316 317 def __init__(self, config: Phi3Config, layer_idx: Optional[int] = None):318 super().__init__()319 self.config = config320 self.layer_idx = layer_idx321 if layer_idx is None:322 logger.warning_once(323 f"Instantiating {self.__class__.__name__} without passing a `layer_idx` is not recommended and will "324 "lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` "325 "when creating this class."326 )327 328 self.attention_dropout = config.attention_dropout329 self.hidden_size = config.hidden_size330 self.num_heads = config.num_attention_heads331 self.head_dim = self.hidden_size // self.num_heads332 self.num_key_value_heads = config.num_key_value_heads333 self.num_key_value_groups = self.num_heads // self.num_key_value_heads334 self.max_position_embeddings = config.max_position_embeddings335 self.original_max_position_embeddings = config.original_max_position_embeddings336 self.rope_theta = config.rope_theta337 self.rope_scaling = config.rope_scaling338 self.is_causal = True339 340 if (self.head_dim * self.num_heads) != self.hidden_size:341 raise ValueError(342 f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}"343 f" and `num_heads`: {self.num_heads})."344 )345 346 op_size = self.num_heads * self.head_dim + 2 * (self.num_key_value_heads * self.head_dim)347 self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)348 self.qkv_proj = nn.Linear(self.hidden_size, op_size, bias=False)349 self._init_rope()350 351 def _init_rope(self):352 if self.rope_scaling is None:353 self.rotary_emb = Phi3RotaryEmbedding(354 self.head_dim,355 max_position_embeddings=self.max_position_embeddings,356 base=self.rope_theta,357 )358 else:359 scaling_type = self.config.rope_scaling["type"]360 short_factor = self.config.rope_scaling["short_factor"]361 long_factor = self.config.rope_scaling["long_factor"]362 363 if scaling_type == "su":364 self.rotary_emb = Phi3SuScaledRotaryEmbedding(365 self.head_dim,366 short_factor,367 long_factor,368 max_position_embeddings=self.max_position_embeddings,369 original_max_position_embeddings=self.original_max_position_embeddings,370 base=self.rope_theta,371 )372 elif scaling_type == "yarn":373 self.rotary_emb = Phi3YarnScaledRotaryEmbedding(374 self.head_dim,375 short_factor,376 long_factor,377 max_position_embeddings=self.max_position_embeddings,378 original_max_position_embeddings=self.original_max_position_embeddings,379 base=self.rope_theta,380 )381 else:382 raise ValueError(f"Unknown RoPE scaling type {scaling_type}")383 384 def forward(385 self,386 hidden_states: torch.Tensor,387 attention_mask: Optional[torch.Tensor] = None,388 position_ids: Optional[torch.LongTensor] = None,389 past_key_value: Optional[Cache] = None,390 output_attentions: bool = False,391 use_cache: bool = False,392 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:393 logger.warning_once("You are not running the flash-attention implementation, expect numerical differences.")394 395 bsz, q_len, _ = hidden_states.size()396 397 qkv = self.qkv_proj(hidden_states)398 query_pos = self.num_heads * self.head_dim399 query_states = qkv[..., :query_pos]400 key_states = qkv[..., query_pos : query_pos + self.num_key_value_heads * self.head_dim]401 value_states = qkv[..., query_pos + self.num_key_value_heads * self.head_dim :]402 403 query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)404 key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)405 value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)406 407 kv_seq_len = key_states.shape[-2]408 if past_key_value is not None:409 if self.layer_idx is None:410 raise ValueError(411 f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "412 "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "413 "with a layer index."414 )415 kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)416 cos, sin = self.rotary_emb(value_states, position_ids, seq_len=kv_seq_len)417 418 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)419 420 if past_key_value is not None:421 cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models422 key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)423 424 # repeat k/v heads if n_kv_heads < n_heads425 key_states = repeat_kv(key_states, self.num_key_value_groups)426 value_states = repeat_kv(value_states, self.num_key_value_groups)427 428 attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)429 430 if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):431 raise ValueError(432 f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"433 f" {attn_weights.size()}"434 )435 436 if attention_mask is not None:437 if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):438 raise ValueError(439 f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"440 )441 attn_weights = attn_weights + attention_mask442 443 # upcast attention to fp32444 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(value_states.dtype)445 attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)446 447 attn_output = torch.matmul(attn_weights, value_states)448 449 if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):450 raise ValueError(451 f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"452 f" {attn_output.size()}"453 )454 455 attn_output = attn_output.transpose(1, 2).contiguous()456 attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)457 458 attn_output = self.o_proj(attn_output)459 460 if not output_attentions:461 attn_weights = None462 463 return attn_output, attn_weights, past_key_value464 465 466class Phi3FlashAttention2(Phi3Attention):467 """468 Phi-3 flash attention module. This module inherits from `Phi3Attention` as the weights of the module stays469 untouched. The only required change would be on the forward pass where it needs to correctly call the public API of470 flash attention and deal with padding tokens in case the input contains any of them.471 """472 473 # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2.__init__474 def __init__(self, *args, **kwargs):475 super().__init__(*args, **kwargs)476 477 # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.478 # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignement, 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.479 # 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).480 self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()481 482 def forward(483 self,484 hidden_states: torch.Tensor,485 attention_mask: Optional[torch.LongTensor] = None,486 position_ids: Optional[torch.LongTensor] = None,487 past_key_value: Optional[Cache] = None,488 output_attentions: bool = False,489 use_cache: bool = False,490 **kwargs,491 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:492 # Phi3FlashAttention2 attention does not support output_attentions493 494 if not _flash_supports_window_size:495 logger.warning_once(496 "The current flash attention version does not support sliding window attention. Please use `attn_implementation='eager'` or upgrade flash-attn library."497 )498 raise ValueError("The current flash attention version does not support sliding window attention.")499 500 output_attentions = False501 502 if "padding_mask" in kwargs:503 warnings.warn(504 "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"505 )506 507 # overwrite attention_mask with padding_mask508 attention_mask = kwargs.pop("padding_mask")509 510 bsz, q_len, _ = hidden_states.size()511 512 qkv = self.qkv_proj(hidden_states)513 query_pos = self.num_heads * self.head_dim514 query_states = qkv[..., :query_pos]515 key_states = qkv[..., query_pos : query_pos + self.num_key_value_heads * self.head_dim]516 value_states = qkv[..., query_pos + self.num_key_value_heads * self.head_dim :]517 518 # Flash attention requires the input to have the shape519 # batch_size x seq_length x head_dim x hidden_dim520 # therefore we just need to keep the original shape521 query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)522 key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)523 value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)524 525 kv_seq_len = key_states.shape[-2]526 if past_key_value is not None:527 if self.layer_idx is None:528 raise ValueError(529 f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "530 "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "531 "with a layer index."532 )533 kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)534 535 # Because the input can be padded, the absolute sequence length depends on the max position id.536 rotary_seq_len = max(kv_seq_len, position_ids[:, -1].max().item()) + 1537 cos, sin = self.rotary_emb(value_states, position_ids, seq_len=rotary_seq_len)538 539 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)540 541 use_sliding_windows = (542 _flash_supports_window_size543 and getattr(self.config, "sliding_window", None) is not None544 and kv_seq_len > self.config.sliding_window545 )546 547 if past_key_value is not None:548 # Activate slicing cache only if the config has a value `sliding_windows` attribute549 cache_has_contents = past_key_value.get_seq_length(self.layer_idx) > 0550 if (551 getattr(self.config, "sliding_window", None) is not None552 and kv_seq_len > self.config.sliding_window553 and cache_has_contents554 ):555 slicing_tokens = 1 - self.config.sliding_window556 557 past_key = past_key_value[self.layer_idx][0]558 past_value = past_key_value[self.layer_idx][1]559 560 past_key = past_key[:, :, slicing_tokens:, :].contiguous()561 past_value = past_value[:, :, slicing_tokens:, :].contiguous()562 563 if past_key.shape[-2] != self.config.sliding_window - 1:564 raise ValueError(565 f"past key must have a shape of (`batch_size, num_heads, self.config.sliding_window-1, head_dim`), got"566 f" {past_key.shape}"567 )568 569 if attention_mask is not None:570 attention_mask = attention_mask[:, slicing_tokens:]571 attention_mask = torch.cat([attention_mask, torch.ones_like(attention_mask[:, -1:])], dim=-1)572 573 cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models574 key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)575 576 # repeat k/v heads if n_kv_heads < n_heads577 key_states = repeat_kv(key_states, self.num_key_value_groups)578 value_states = repeat_kv(value_states, self.num_key_value_groups)579 580 attn_dropout = self.attention_dropout if self.training else 0.0581 582 # In PEFT, usually we cast the layer norms in float32 for training stability reasons583 # therefore the input hidden states gets silently casted in float32. Hence, we need584 # cast them back in the correct dtype just to be sure everything works as expected.585 # This might slowdown training & inference so it is recommended to not cast the LayerNorms586 # in fp32.587 588 if query_states.dtype == torch.float32:589 if torch.is_autocast_enabled():590 target_dtype = torch.get_autocast_gpu_dtype()591 # Handle the case where the model is quantized592 elif hasattr(self.config, "_pre_quantization_dtype"):593 target_dtype = self.config._pre_quantization_dtype594 else:595 target_dtype = self.qkv_proj.weight.dtype596 597 logger.warning_once(598 f"The input hidden states seems to be silently casted in float32, this might be related to"599 f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"600 f" {target_dtype}."601 )602 603 query_states = query_states.to(target_dtype)604 key_states = key_states.to(target_dtype)605 value_states = value_states.to(target_dtype)606 607 # Reashape to the expected shape for Flash Attention608 query_states = query_states.transpose(1, 2)609 key_states = key_states.transpose(1, 2)610 value_states = value_states.transpose(1, 2)611 612 attn_output = self._flash_attention_forward(613 query_states,614 key_states,615 value_states,616 attention_mask,617 q_len,618 dropout=attn_dropout,619 use_sliding_windows=use_sliding_windows,620 )621 622 attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()623 attn_output = self.o_proj(attn_output)624 625 if not output_attentions:626 attn_weights = None627 628 return attn_output, attn_weights, past_key_value629 630 # Copied from transformers.models.mistral.modeling_mistral.MistralFlashAttention2._flash_attention_forward631 def _flash_attention_forward(632 self,633 query_states,634 key_states,635 value_states,636 attention_mask,637 query_length,638 dropout=0.0,639 softmax_scale=None,640 use_sliding_windows=False,641 ):642 """643 Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token644 first unpad the input, then computes the attention scores and pad the final attention scores.645 646 Args:647 query_states (`torch.Tensor`):648 Input query states to be passed to Flash Attention API649 key_states (`torch.Tensor`):650 Input key states to be passed to Flash Attention API651 value_states (`torch.Tensor`):652 Input value states to be passed to Flash Attention API653 attention_mask (`torch.Tensor`):654 The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the655 position of padding tokens and 1 for the position of non-padding tokens.656 dropout (`float`):657 Attention dropout658 softmax_scale (`float`, *optional*):659 The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)660 use_sliding_windows (`bool`, *optional*):661 Whether to activate sliding window attention.662 """663 if not self._flash_attn_uses_top_left_mask:664 causal = self.is_causal665 else:666 # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1. For details, please see the comment in LlamaFlashAttention2 __init__.667 causal = self.is_causal and query_length != 1668 669 # Contains at least one padding token in the sequence670 if attention_mask is not None:671 batch_size = query_states.shape[0]672 query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(673 query_states, key_states, value_states, attention_mask, query_length674 )675 676 cu_seqlens_q, cu_seqlens_k = cu_seq_lens677 max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens678 679 if not use_sliding_windows:680 attn_output_unpad = flash_attn_varlen_func(681 query_states,682 key_states,683 value_states,684 cu_seqlens_q=cu_seqlens_q,685 cu_seqlens_k=cu_seqlens_k,686 max_seqlen_q=max_seqlen_in_batch_q,687 max_seqlen_k=max_seqlen_in_batch_k,688 dropout_p=dropout,689 softmax_scale=softmax_scale,690 causal=causal,691 )692 else:693 attn_output_unpad = flash_attn_varlen_func(694 query_states,695 key_states,696 value_states,697 cu_seqlens_q=cu_seqlens_q,698 cu_seqlens_k=cu_seqlens_k,699 max_seqlen_q=max_seqlen_in_batch_q,700 max_seqlen_k=max_seqlen_in_batch_k,701 dropout_p=dropout,702 softmax_scale=softmax_scale,703 causal=causal,704 window_size=(self.config.sliding_window, self.config.sliding_window),705 )706 707 attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)708 else:709 if not use_sliding_windows:710 attn_output = flash_attn_func(711 query_states,712 key_states,713 value_states,714 dropout,715 softmax_scale=softmax_scale,716 causal=causal,717 )718 else:719 attn_output = flash_attn_func(720 query_states,721 key_states,722 value_states,723 dropout,724 softmax_scale=softmax_scale,725 causal=causal,726 window_size=(self.config.sliding_window, self.config.sliding_window),727 )728 729 return attn_output730 731 # Copied from transformers.models.mistral.modeling_mistral.MistralFlashAttention2._upad_input732 def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):733 batch_size, kv_seq_len, num_heads, head_dim = key_layer.shape734 735 # On the first iteration we need to properly re-create the padding mask736 # by slicing it on the proper place737 if kv_seq_len != attention_mask.shape[-1]:738 attention_mask_num_tokens = attention_mask.shape[-1]739 attention_mask = attention_mask[:, attention_mask_num_tokens - kv_seq_len :]740 741 indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)742 743 key_layer = index_first_axis(key_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)744 value_layer = index_first_axis(value_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)745 746 if query_length == kv_seq_len:747 query_layer = index_first_axis(748 query_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k749 )750 cu_seqlens_q = cu_seqlens_k751 max_seqlen_in_batch_q = max_seqlen_in_batch_k752 indices_q = indices_k753 elif query_length == 1:754 max_seqlen_in_batch_q = 1755 cu_seqlens_q = torch.arange(756 batch_size + 1, dtype=torch.int32, device=query_layer.device757 ) # There is a memcpy here, that is very bad.758 indices_q = cu_seqlens_q[:-1]759 query_layer = query_layer.squeeze(1)760 else:761 # The -q_len: slice assumes left padding.762 attention_mask = attention_mask[:, -query_length:]763 query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)764 765 return (766 query_layer,767 key_layer,768 value_layer,769 indices_q,770 (cu_seqlens_q, cu_seqlens_k),771 (max_seqlen_in_batch_q, max_seqlen_in_batch_k),772 )773 774 775# copied from transformers.models.llama.modeling_llama.LlamaSdpaAttention with Llama->Phi3776# TODO @Arthur no longer copied from LLama after static cache777class Phi3SdpaAttention(Phi3Attention):778 """779 Phi3 attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from780 `Phi3Attention` as the weights of the module stays untouched. The only changes are on the forward pass to adapt to781 SDPA API.782 """783 784 # Adapted from Phi3Attention.forward785 def forward(786 self,787 hidden_states: torch.Tensor,788 attention_mask: Optional[torch.Tensor] = None,789 position_ids: Optional[torch.LongTensor] = None,790 past_key_value: Optional[Cache] = None,791 output_attentions: bool = False,792 use_cache: bool = False,793 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:794 if output_attentions:795 # TODO: Improve this warning with e.g. `model.config.attn_implementation = "manual"` once this is implemented.796 logger.warning_once(797 "Phi3Model is using Phi3SdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, "798 'but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.'799 )800 return super().forward(801 hidden_states=hidden_states,802 attention_mask=attention_mask,803 position_ids=position_ids,804 past_key_value=past_key_value,805 output_attentions=output_attentions,806 use_cache=use_cache,807 )808 809 bsz, q_len, _ = hidden_states.size()810 811 qkv = self.qkv_proj(hidden_states)812 query_pos = self.num_heads * self.head_dim813 query_states = qkv[..., :query_pos]814 key_states = qkv[..., query_pos : query_pos + self.num_key_value_heads * self.head_dim]815 value_states = qkv[..., query_pos + self.num_key_value_heads * self.head_dim :]816 817 query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)818 key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)819 value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)820 821 kv_seq_len = key_states.shape[-2]822 if past_key_value is not None:823 kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)824 cos, sin = self.rotary_emb(value_states, position_ids, seq_len=kv_seq_len)825 826 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)827 828 if past_key_value is not None:829 cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models830 key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)831 832 key_states = repeat_kv(key_states, self.num_key_value_groups)833 value_states = repeat_kv(value_states, self.num_key_value_groups)834 835 if attention_mask is not None:836 if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):837 raise ValueError(838 f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"839 )840 841 # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with custom attn_mask,842 # Reference: https://github.com/pytorch/pytorch/issues/112577.843 if query_states.device.type == "cuda" and attention_mask is not None:844 query_states = query_states.contiguous()845 key_states = key_states.contiguous()846 value_states = value_states.contiguous()847 848 attn_output = torch.nn.functional.scaled_dot_product_attention(849 query_states,850 key_states,851 value_states,852 attn_mask=attention_mask,853 dropout_p=self.attention_dropout if self.training else 0.0,854 # The q_len > 1 is necessary to match with AttentionMaskConverter.to_causal_4d that does not create a causal mask in case q_len == 1.855 is_causal=self.is_causal and attention_mask is None and q_len > 1,856 )857 858 attn_output = attn_output.transpose(1, 2).contiguous()859 attn_output = attn_output.view(bsz, q_len, self.hidden_size)860 861 attn_output = self.o_proj(attn_output)862 863 return attn_output, None, past_key_value864 865 866PHI3_ATTENTION_CLASSES = {867 "eager": Phi3Attention,868 "flash_attention_2": Phi3FlashAttention2,869 "sdpa": Phi3SdpaAttention,870}871 872 873class Phi3DecoderLayer(nn.Module):874 def __init__(self, config: Phi3Config, layer_idx: int):875 super().__init__()876 877 self.config = config878 self.self_attn = PHI3_ATTENTION_CLASSES[config._attn_implementation](config, layer_idx=layer_idx)879 880 self.mlp = Phi3MLP(config)881 self.input_layernorm = Phi3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)882 883 self.resid_attn_dropout = nn.Dropout(config.resid_pdrop)884 self.resid_mlp_dropout = nn.Dropout(config.resid_pdrop)885 self.post_attention_layernorm = Phi3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)886 887 def forward(888 self,889 hidden_states: torch.Tensor,890 attention_mask: Optional[torch.Tensor] = None,891 position_ids: Optional[torch.LongTensor] = None,892 past_key_value: Optional[Tuple[torch.Tensor]] = None,893 output_attentions: Optional[bool] = False,894 use_cache: Optional[bool] = False,895 **kwargs,896 ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:897 if "padding_mask" in kwargs:898 warnings.warn(899 "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"900 )901 """902 Args:903 hidden_states (`torch.FloatTensor`):904 input to the layer of shape `(batch, seq_len, embed_dim)`905 attention_mask (`torch.FloatTensor`, *optional*): attention mask of size906 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.907 position_ids (`torch.LongTensor` of shape `({0})`, *optional*):908 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range909 `[0, config.n_positions - 1]`. [What are position IDs?](../glossary#position-ids)910 output_attentions (`bool`, *optional*):911 Whether or not to return the attentions tensors of all attention layers. See `attentions` under912 returned tensors for more detail.913 use_cache (`bool`, *optional*):914 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding915 (see `past_key_values`).916 past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states917 """918 919 residual = hidden_states920 921 hidden_states = self.input_layernorm(hidden_states)922 923 # Self Attention924 attn_outputs, self_attn_weights, present_key_value = self.self_attn(925 hidden_states=hidden_states,926 attention_mask=attention_mask,927 position_ids=position_ids,928 past_key_value=past_key_value,929 output_attentions=output_attentions,930 use_cache=use_cache,931 )932 933 hidden_states = residual + self.resid_attn_dropout(attn_outputs)934 935 residual = hidden_states936 hidden_states = self.post_attention_layernorm(hidden_states)937 hidden_states = self.mlp(hidden_states)938 hidden_states = residual + self.resid_mlp_dropout(hidden_states)939 940 outputs = (hidden_states,)941 942 if output_attentions:943 outputs += (self_attn_weights,)944 945 if use_cache:946 outputs += (present_key_value,)947 948 return outputs949 950 951PHI3_START_DOCSTRING = r"""952 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the953 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads954 etc.)955 956 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.957 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage958 and behavior.959 960 Parameters:961 config ([`Phi3Config`]):962 Model configuration class with all the parameters of the model. Initializing with a config file does not963 load the weights associated with the model, only the configuration. Check out the964 [`~PreTrainedModel.from_pretrained`] method to load the model weights.965"""966 967 968@add_start_docstrings(969 "The bare Phi-3 model outputting raw hidden-states without any specific head on top.",970 PHI3_START_DOCSTRING,971)972class Phi3PreTrainedModel(PreTrainedModel):973 config_class = Phi3Config974 base_model_prefix = "model"975 supports_gradient_checkpointing = True976 _no_split_modules = ["Phi3DecoderLayer"]977 _skip_keys_device_placement = "past_key_values"978 _supports_flash_attn_2 = True979 _supports_sdpa = False980 _supports_cache_class = True981 982 _version = "0.0.5"983 984 def _init_weights(self, module):985 std = self.config.initializer_range986 if isinstance(module, nn.Linear):987 module.weight.data.normal_(mean=0.0, std=std)988 if module.bias is not None:989 module.bias.data.zero_()990 elif isinstance(module, nn.Embedding):991 module.weight.data.normal_(mean=0.0, std=std)992 if module.padding_idx is not None:993 module.weight.data[module.padding_idx].zero_()994 995 996PHI3_INPUTS_DOCSTRING = r"""997 Args:998 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):999 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide1000 it.1001 1002 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1003 [`PreTrainedTokenizer.__call__`] for details.1004 1005 [What are input IDs?](../glossary#input-ids)1006 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1007 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1008 1009 - 1 for tokens that are **not masked**,1010 - 0 for tokens that are **masked**.1011 1012 [What are attention masks?](../glossary#attention-mask)1013 1014 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1015 [`PreTrainedTokenizer.__call__`] for details.1016 1017 If `past_key_values` is used, optionally only the last `input_ids` have to be input (see1018 `past_key_values`).1019 1020 If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]1021 and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more1022 information on the default strategy.1023 1024 - 1 indicates the head is **not masked**,1025 - 0 indicates the head is **masked**.1026 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1027 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,1028 config.n_positions - 1]`.1029 1030 [What are position IDs?](../glossary#position-ids)1031 past_key_values (`Cache` or `tuple(tuple(torch.FloatTensor))`, *optional*):1032 Pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention1033 blocks) that can be used to speed up sequential decoding. This typically consists in the `past_key_values`1034 returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.1035 1036 Two formats are allowed:1037 - a [`~cache_utils.Cache`] instance;1038 - Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of1039 shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`). This is also known as the legacy1040 cache format.1041 1042 The model will output the same cache format that is fed as input. If no `past_key_values` are passed, the1043 legacy cache format will be returned.1044 1045 If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't1046 have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`1047 of shape `(batch_size, sequence_length)`.1048 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):1049 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This1050 is useful if you want more control over how to convert `input_ids` indices into associated vectors than the1051 model's internal embedding lookup matrix.1052 use_cache (`bool`, *optional*):1053 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see1054 `past_key_values`).1055 output_attentions (`bool`, *optional*):1056 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1057 tensors for more detail.1058 output_hidden_states (`bool`, *optional*):1059 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1060 more detail.1061 return_dict (`bool`, *optional*):1062 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1063"""1064 1065 1066@add_start_docstrings(1067 "The bare Phi-3 model outputting raw hidden-states without any specific head on top.",1068 PHI3_START_DOCSTRING,1069)1070class Phi3Model(Phi3PreTrainedModel):1071 """1072 Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Phi3DecoderLayer`]1073 1074 Args:1075 config: Phi3Config1076 """1077 1078 def __init__(self, config: Phi3Config):1079 super().__init__(config)1080 self.padding_idx = config.pad_token_id1081 self.vocab_size = config.vocab_size1082 1083 self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)1084 self.embed_dropout = nn.Dropout(config.embd_pdrop)1085 self.layers = nn.ModuleList(1086 [Phi3DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]1087 )1088 self._attn_implementation = config._attn_implementation1089 self.norm = Phi3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)1090 1091 self.gradient_checkpointing = False1092 # Initialize weights and apply final processing1093 self.post_init()1094 1095 def get_input_embeddings(self):1096 return self.embed_tokens1097 1098 def set_input_embeddings(self, value):1099 self.embed_tokens = value1100 1101 @add_start_docstrings_to_model_forward(PHI3_INPUTS_DOCSTRING)1102 def forward(1103 self,1104 input_ids: torch.LongTensor = None,1105 attention_mask: Optional[torch.Tensor] = None,1106 position_ids: Optional[torch.LongTensor] = None,1107 past_key_values: Optional[List[torch.FloatTensor]] = None,1108 inputs_embeds: Optional[torch.FloatTensor] = None,1109 use_cache: Optional[bool] = None,1110 output_attentions: Optional[bool] = None,1111 output_hidden_states: Optional[bool] = None,1112 return_dict: Optional[bool] = None,1113 ) -> Union[Tuple, BaseModelOutputWithPast]:1114 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions1115 output_hidden_states = (1116 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1117 )1118 use_cache = use_cache if use_cache is not None else self.config.use_cache1119 1120 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1121 1122 # retrieve input_ids and inputs_embeds1123 if input_ids is not None and inputs_embeds is not None:1124 raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")1125 elif input_ids is not None:1126 batch_size, seq_length = input_ids.shape[:2]1127 elif inputs_embeds is not None:1128 batch_size, seq_length = inputs_embeds.shape[:2]1129 else:1130 raise ValueError("You have to specify either input_ids or inputs_embeds")1131 1132 past_key_values_length = 01133 1134 if self.gradient_checkpointing and self.training:1135 if use_cache:1136 logger.warning_once(1137 "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."1138 )1139 use_cache = False1140 1141 if use_cache:1142 use_legacy_cache = not isinstance(past_key_values, Cache)1143 if use_legacy_cache:1144 past_key_values = DynamicCache.from_legacy_cache(past_key_values)1145 past_key_values_length = past_key_values.get_usable_length(seq_length)1146 1147 if position_ids is None:1148 device = input_ids.device if input_ids is not None else inputs_embeds.device1149 position_ids = torch.arange(1150 past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device1151 )1152 position_ids = position_ids.unsqueeze(0).view(-1, seq_length)1153 else:1154 position_ids = position_ids.view(-1, seq_length).long()1155 1156 if inputs_embeds is None:1157 inputs_embeds = self.embed_tokens(input_ids)1158 1159 if attention_mask is not None and self._attn_implementation == "flash_attention_2" and use_cache:1160 is_padding_right = attention_mask[:, -1].sum().item() != batch_size1161 if is_padding_right:1162 raise ValueError(1163 "You are attempting to perform batched generation with padding_side='right'"1164 " this may lead to unexpected behaviour for Flash Attention version of Phi3. Make sure to "1165 " call `tokenizer.padding_side = 'left'` before tokenizing the input. "1166 )1167 1168 if self._attn_implementation == "flash_attention_2":1169 # 2d mask is passed through the layers1170 attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None1171 else:1172 # 4d mask is passed through the layers1173 attention_mask = _prepare_4d_causal_attention_mask(1174 attention_mask,1175 (batch_size, seq_length),1176 inputs_embeds,1177 past_key_values_length,1178 sliding_window=self.config.sliding_window,1179 )1180 1181 hidden_states = inputs_embeds1182 1183 # decoder layers1184 all_hidden_states = () if output_hidden_states else None1185 all_self_attns = () if output_attentions else None1186 next_decoder_cache = None1187 1188 for decoder_layer in self.layers:1189 if output_hidden_states:1190 all_hidden_states += (hidden_states,)1191 1192 if self.gradient_checkpointing and self.training:1193 layer_outputs = self._gradient_checkpointing_func(1194 decoder_layer.__call__,1195 hidden_states,1196 attention_mask,1197 position_ids,1198 past_key_values,1199 output_attentions,1200 use_cache,