Team Ai
Modelpublic

llmware/slim-boolean-phi-3

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes20downloads
modeling_phi3.py1646 linesDownload Raw Back to root
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,

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