Team Ai
Modelpublic

LeroyDyer/SpydazWebAI_VisionEncoderDecoderModel_Mini3b

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes13downloads
modeling_mistral_advanced.py4064 linesDownload Raw Back to root
1# coding=utf-82# Copyright 2023 Mistral AI and the HuggingFace Inc. team. All rights reserved.3#4# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX5# and OPT implementations in this library. It has been modified from its6# original forms to accommodate minor architectural differences compared7# to GPT-NeoX and OPT used by the Meta AI team that trained the model.8#9# Licensed under the Apache License, Version 2.0 (the "License");10# you may not use this file except in compliance with the License.11# You may obtain a copy of the License at12#13#     http://www.apache.org/licenses/LICENSE-2.014#15# Unless required by applicable law or agreed to in writing, software16# distributed under the License is distributed on an "AS IS" BASIS,17# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.18# See the License for the specific language governing permissions and19# limitations under the License.20""" PyTorch Mistral model."""21from termcolor import colored22from tqdm import tqdm23import pandas as pd24import seaborn as sns25import matplotlib.pyplot as plt26import inspect27import math28import copy29import time30import warnings31from typing import List, Optional, Tuple, Union32import gc33import os34import tempfile35import random36import numpy as np37import warnings38import torch39import torch.nn.functional as F40import torch.utils.checkpoint41from matplotlib.colors import LinearSegmentedColormap, LogNorm42from torch import nn43from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss44from configuration_mistral_advanced import VisionEncoderDecoderConfig45from configuration_mistral_advanced import EncoderDecoderConfig46from transformers.auto.configuration_auto import AutoConfig47from transformers.auto.modeling_auto import AutoModel, AutoModelForCausalLM48from collections import defaultdict49from transformers.activations import ACT2FN50from transformers.cache_utils import Cache, DynamicCache51from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask, _prepare_4d_causal_attention_mask_for_sdpa52from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast, SequenceClassifierOutputWithPast53from transformers.modeling_utils import PreTrainedModel54from transformers.utils import (55    add_start_docstrings,56    add_start_docstrings_to_model_forward,57    is_flash_attn_2_available,58    is_flash_attn_greater_or_equal_2_10,59    logging,60    replace_return_docstrings,61)62 63 64 65from configuration_mistral_advanced import MistralConfig66 67 68if is_flash_attn_2_available():69    from flash_attn import flash_attn_func, flash_attn_varlen_func70    from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input  # noqa71 72    _flash_supports_window_size = "window_size" in list(inspect.signature(flash_attn_func).parameters)73 74 75logger = logging.get_logger(__name__)76 77_CONFIG_FOR_DOC = "MistralConfig"78 79 80# Copied from transformers.models.llama.modeling_llama._get_unpad_data81def _get_unpad_data(attention_mask):82    seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)83    indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()84    max_seqlen_in_batch = seqlens_in_batch.max().item()85    cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))86    return (87        indices,88        cu_seqlens,89        max_seqlen_in_batch,90    )91 92 93 94# Copied from transformers.models.bart.modeling_bart._make_causal_mask95def _make_causal_mask(96    input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 097):98    """99    Make causal mask used for bi-directional self-attention.100    """101    bsz, tgt_len = input_ids_shape102    mask = torch.full((tgt_len, tgt_len), torch.finfo(dtype).min, device=device)103    mask_cond = torch.arange(mask.size(-1), device=device)104    mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)105    mask = mask.to(dtype)106 107    if past_key_values_length > 0:108        mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)109    return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)110 111def _make_sliding_window_causal_mask(112    input_ids_shape: torch.Size,113    dtype: torch.dtype,114    device: torch.device,115    past_key_values_length: int = 0,116    sliding_window: int = 4096,117):118    """119    Make causal mask used for sliding window attention120    """121    bsz, tgt_len = input_ids_shape122 123    tensor = torch.full(124        (tgt_len, tgt_len),125        fill_value=1,126        device=device,127    )128    mask = torch.tril(tensor, diagonal=0)129    # make the mask banded to account for sliding window130    mask = torch.triu(mask, diagonal=-sliding_window)131    mask = torch.log(mask).to(dtype)132 133    if past_key_values_length > 0:134        mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)135    return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)136 137 138# Copied from transformers.models.bart.modeling_bart._expand_mask139def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):140    """141    Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.142    """143    bsz, src_len = mask.size()144    tgt_len = tgt_len if tgt_len is not None else src_len145 146    expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)147 148    inverted_mask = 1.0 - expanded_mask149 150    return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)151 152# Inverse dim formula to find dim based on number of rotations153def _yarn_find_correction_dim(num_rotations, dim, base=10000, max_position_embeddings=2048):154    return (dim * math.log(max_position_embeddings/(num_rotations * 2 * math.pi)))/(2 * math.log(base))155 156# Find dim range bounds based on rotations157def _yarn_find_correction_range(low_rot, high_rot, dim, base=10000, max_position_embeddings=2048):158    low = math.floor(_yarn_find_correction_dim(159        low_rot, dim, base, max_position_embeddings))160    high = math.ceil(_yarn_find_correction_dim(161        high_rot, dim, base, max_position_embeddings))162    return max(low, 0), min(high, dim-1)  # Clamp values just in case163 164def _yarn_linear_ramp_mask(min, max, dim):165    if min == max:166        max += 0.001  # Prevent singularity167 168    linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min)169    ramp_func = torch.clamp(linear_func, 0, 1)170    return ramp_func171 172def _yarn_get_mscale(scale=1):173    if scale <= 1:174        return 1.0175    return 0.07 * math.log(scale) + 1.0176 177 178# Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->Mistral179class MistralRMSNorm(nn.Module):180    def __init__(self, hidden_size, eps=1e-6):181        """182        MistralRMSNorm is equivalent to T5LayerNorm183        """184        super().__init__()185        self.weight = nn.Parameter(torch.ones(hidden_size))186        self.variance_epsilon = eps187 188    def forward(self, hidden_states):189        input_dtype = hidden_states.dtype190        hidden_states = hidden_states.to(torch.float32)191        variance = hidden_states.pow(2).mean(-1, keepdim=True)192        hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)193        return self.weight * hidden_states.to(input_dtype)194 195 196# copied from transformers.models.llama.modeling_llama.LlamaRotaryEmbedding with Llama->Mistral197# TODO @Arthur no longer copied from LLama after static cache198class MistralRotaryEmbedding(nn.Module):199    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):200        super().__init__()201 202        self.dim = dim203        self.max_position_embeddings = max_position_embeddings204        self.base = base205        inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float().to(device) / self.dim))206        self.register_buffer("inv_freq", inv_freq, persistent=False)207 208        # Build here to make `torch.jit.trace` work.209        self._set_cos_sin_cache(210            seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()211        )212 213    def _set_cos_sin_cache(self, seq_len, device, dtype):214        self.max_seq_len_cached = seq_len215        t = torch.arange(self.max_seq_len_cached, device=device, dtype=torch.int64).type_as(self.inv_freq)216 217        freqs = torch.outer(t, self.inv_freq)218        # Different from paper, but it uses a different permutation in order to obtain the same calculation219        emb = torch.cat((freqs, freqs), dim=-1)220        self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)221        self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)222 223    def forward(self, x, seq_len=None):224        # x: [bs, num_attention_heads, seq_len, head_size]225        if seq_len > self.max_seq_len_cached:226            self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype)227 228        return (229            self.cos_cached[:seq_len].to(dtype=x.dtype),230            self.sin_cached[:seq_len].to(dtype=x.dtype),231        )232 233### YARN ADDITIONS234 235class MistralDynamicNTKScalingRotaryEmbedding(MistralRotaryEmbedding):236    """MistralRotaryEmbedding extended with Dynamic NTK scaling. Credits to the Reddit users /u/bloc97 and /u/emozilla"""237 238    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):239        self.scaling_factor = scaling_factor240        super().__init__(dim, max_position_embeddings, base, device)241 242    def _set_cos_sin_cache(self, seq_len, device, dtype):243        self.max_seq_len_cached = seq_len244 245        if seq_len > self.max_position_embeddings:246            base = self.base * (247                (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)248            ) ** (self.dim / (self.dim - 2))249            inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))250            self.register_buffer("inv_freq", inv_freq, persistent=False)251 252        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)253 254        freqs = torch.einsum("i,j->ij", t, self.inv_freq)255        # Different from paper, but it uses a different permutation in order to obtain the same calculation256        emb = torch.cat((freqs, freqs), dim=-1)257        self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)258        self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)259 260class MistralLinearScalingRotaryEmbedding(MistralRotaryEmbedding):261    """MistralRotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""262 263    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):264        self.scaling_factor = scaling_factor265        super().__init__(dim, max_position_embeddings, base, device)266 267    def _set_cos_sin_cache(self, seq_len, device, dtype):268        self.max_seq_len_cached = seq_len269        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)270        t = t / self.scaling_factor271 272        freqs = torch.einsum("i,j->ij", t, self.inv_freq)273        # Different from paper, but it uses a different permutation in order to obtain the same calculation274        emb = torch.cat((freqs, freqs), dim=-1)275        self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)276        self.register_buffer("sin_cached", emb.cos().to(dtype), persistent=False)277 278class MistralYaRNScaledRotaryEmbedding(torch.nn.Module):279    """MistralRotaryEmbedding extended with YaRN. See: https://arxiv.org/abs/2309.00071"""280    def __init__(self, dim, max_position_embeddings=2048, base=10000, scale=1, original_max_position_embeddings=2048,281                 extrapolation_factor=1, attn_factor=1, beta_fast=128, beta_slow=2, finetuned=False, device=None):282        super().__init__()283 284        self.dim = dim285        self.max_position_embeddings = max_position_embeddings286        self.base = base287        self.scale = scale288        self.original_max_position_embeddings = original_max_position_embeddings289        self.extrapolation_factor = extrapolation_factor290        self.attn_factor = attn_factor291        self.beta_fast = beta_fast292        self.beta_slow = beta_slow293 294        self.yarn(device)295 296        # Build here to make `torch.jit.trace` work.297        self.max_seq_len_cached = max_position_embeddings298        t = torch.arange(self.max_seq_len_cached, device=self.inv_freq.device, dtype=self.inv_freq.dtype)299        freqs = torch.einsum("i,j->ij", t, self.inv_freq)300        # Different from paper, but it uses a different permutation in order to obtain the same calculation301        emb = torch.cat((freqs, freqs), dim=-1)302        dtype = torch.get_default_dtype()303 304        self.register_buffer("cos_cached", (emb.cos() * self.mscale).to(dtype), persistent=False)305        self.register_buffer("sin_cached", (emb.sin() * self.mscale).to(dtype), persistent=False)306 307    def forward(self, x, seq_len=None):308        # x: [bs, num_attention_heads, seq_len, head_size]309        # This `if` block is unlikely to be run after we build sin/cos in `__init__`. Keep the logic here just in case.310        if seq_len > self.max_seq_len_cached:311            self.max_seq_len_cached = seq_len312 313            t = torch.arange(self.max_seq_len_cached, device=x.device, dtype=self.inv_freq.dtype)314            freqs = torch.einsum("i,j->ij", t, self.inv_freq)315            # Different from paper, but it uses a different permutation in order to obtain the same calculation316            emb = torch.cat((freqs, freqs), dim=-1).to(x.device)317 318            self.register_buffer("cos_cached", (emb.cos() * self.mscale).to(x.dtype), persistent=False)319            self.register_buffer("sin_cached", (emb.sin() * self.mscale).to(x.dtype), persistent=False)320        return (321            self.cos_cached[:seq_len].to(dtype=x.dtype),322            self.sin_cached[:seq_len].to(dtype=x.dtype),323        )324 325    def yarn(self, device):326        pos_freqs = self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim)327        inv_freq_extrapolation = 1.0 / pos_freqs328        inv_freq_interpolation = 1.0 / (self.scale * pos_freqs)329 330        low, high = _yarn_find_correction_range(self.beta_fast, self.beta_slow, self.dim, self.base, self.original_max_position_embeddings)331        inv_freq_mask = (1 - _yarn_linear_ramp_mask(low, high, self.dim // 2).float().to(device)) * self.extrapolation_factor # Get n-d rotational scaling corrected for extrapolation332        inv_freq = inv_freq_interpolation * (1 - inv_freq_mask) + inv_freq_extrapolation * inv_freq_mask333 334        self.register_buffer("inv_freq", inv_freq, persistent=False)335        self.mscale = float(_yarn_get_mscale(self.scale) * self.attn_factor) # Get n-d magnitude scaling corrected for interpolation336 337class MistralDynamicYaRNScaledRotaryEmbedding(torch.nn.Module):338    """MistralRotaryEmbedding extended with Dynamic YaRN. See: https://arxiv.org/abs/2309.00071"""339    def __init__(self, dim, max_position_embeddings=2048, base=10000, original_max_position_embeddings=2048,340                 extrapolation_factor=1, attn_factor=1, beta_fast=128, beta_slow=2, finetuned=False, device=None):341        super().__init__()342 343        self.dim = dim344        self.max_position_embeddings = max_position_embeddings345        self.base = base346        self.original_max_position_embeddings = original_max_position_embeddings347        self.extrapolation_factor = extrapolation_factor348        self.attn_factor = attn_factor349        self.beta_fast = beta_fast350        self.beta_slow = beta_slow351 352        if finetuned:353            self.yarn(self.max_position_embeddings / self.original_max_position_embeddings, device)354        else:355            inv_freq = 1.0 / \356                (base ** (torch.arange(0, dim, 2).float().to(device) / dim))357            self.register_buffer("inv_freq", inv_freq, persistent=False)358            self.mscale = 1359 360        # Build here to make `torch.jit.trace` work.361        self.max_seq_len_cached = max_position_embeddings362        t = torch.arange(self.max_seq_len_cached, device=self.inv_freq.device, dtype=torch.float32)363        freqs = torch.einsum("i,j->ij", t, self.inv_freq)364        # Different from paper, but it uses a different permutation in order to obtain the same calculation365        emb = torch.cat((freqs, freqs), dim=-1)366        dtype = torch.get_default_dtype()367 368        self.register_buffer("cos_cached", (emb.cos() * self.mscale).to(dtype), persistent=False)369        self.register_buffer("sin_cached", (emb.sin() * self.mscale).to(dtype), persistent=False)370 371    def forward(self, x, seq_len=None):372        # x: [bs, num_attention_heads, seq_len, head_size]373        # This `if` block is unlikely to be run after we build sin/cos in `__init__`. Keep the logic here just in case.374        if seq_len > self.max_seq_len_cached:375            self.max_seq_len_cached = seq_len376 377            self.yarn(seq_len / self.max_position_embeddings, x.device)378 379            t = torch.arange(self.max_seq_len_cached, device=x.device, dtype=self.inv_freq.dtype)380            freqs = torch.einsum("i,j->ij", t, self.inv_freq)381            # Different from paper, but it uses a different permutation in order to obtain the same calculation382            emb = torch.cat((freqs, freqs), dim=-1).to(x.device)383 384            self.register_buffer("cos_cached", (emb.cos() * self.mscale).to(x.dtype), persistent=False)385            self.register_buffer("sin_cached", (emb.sin() * self.mscale).to(x.dtype), persistent=False)386        return (387            self.cos_cached[:seq_len].to(dtype=x.dtype),388            self.sin_cached[:seq_len].to(dtype=x.dtype),389        )390 391    def yarn(self, scale, device):392        pos_freqs = self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim)393        inv_freq_extrapolation = 1.0 / pos_freqs394        inv_freq_interpolation = 1.0 / (scale * pos_freqs)395 396        low, high = _yarn_find_correction_range(self.beta_fast, self.beta_slow, self.dim, self.base, self.original_max_position_embeddings)397        inv_freq_mask = (1 - _yarn_linear_ramp_mask(low, high, self.dim // 2).float().to(device)) * self.extrapolation_factor # Get n-d rotational scaling corrected for extrapolation398        inv_freq = inv_freq_interpolation * (1 - inv_freq_mask) + inv_freq_extrapolation * inv_freq_mask399 400        self.register_buffer("inv_freq", inv_freq, persistent=False)401        self.mscale = float(_yarn_get_mscale(scale) * self.attn_factor) # Get n-d magnitude scaling corrected for interpolation402###################403 404 405 406 407# Copied from transformers.models.llama.modeling_llama.rotate_half408def rotate_half(x):409    """Rotates half the hidden dims of the input."""410    x1 = x[..., : x.shape[-1] // 2]411    x2 = x[..., x.shape[-1] // 2 :]412    return torch.cat((-x2, x1), dim=-1)413 414 415# copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb416# TODO @Arthur no longer copied from LLama after static cache417def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):418    """Applies Rotary Position Embedding to the query and key tensors.419 420    Args:421        q (`torch.Tensor`): The query tensor.422        k (`torch.Tensor`): The key tensor.423        cos (`torch.Tensor`): The cosine part of the rotary embedding.424        sin (`torch.Tensor`): The sine part of the rotary embedding.425        position_ids (`torch.Tensor`):426            The position indices of the tokens corresponding to the query and key tensors. For example, this can be427            used to pass offsetted position ids when working with a KV-cache.428        unsqueeze_dim (`int`, *optional*, defaults to 1):429            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and430            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note431            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and432            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes433            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have434            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.435    Returns:436        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.437    """438    cos = cos[position_ids].unsqueeze(unsqueeze_dim)439    sin = sin[position_ids].unsqueeze(unsqueeze_dim)440    q_embed = (q * cos) + (rotate_half(q) * sin)441    k_embed = (k * cos) + (rotate_half(k) * sin)442    return q_embed, k_embed443 444 445class MistralMLP(nn.Module):446    def __init__(self, config):447        super().__init__()448        self.config = config449        self.hidden_size = config.hidden_size450        self.intermediate_size = config.intermediate_size451        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)452        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)453        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)454        self.act_fn = ACT2FN[config.hidden_act]455 456    def forward(self, x):457        return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))458 459 460# Copied from transformers.models.llama.modeling_llama.repeat_kv461def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:462    """463    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,464    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)465    """466    batch, num_key_value_heads, slen, head_dim = hidden_states.shape467    if n_rep == 1:468        return hidden_states469    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)470    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)471 472 473class MistralAttention(nn.Module):474    """475    Multi-headed attention from 'Attention Is All You Need' paper. Modified to use sliding window attention: Longformer476    and "Generating Long Sequences with Sparse Transformers".477    """478 479    def __init__(self, config: MistralConfig, layer_idx: Optional[int] = None):480        super().__init__()481        self.config = config482        self.layer_idx = layer_idx483        if layer_idx is None:484            logger.warning_once(485                f"Instantiating {self.__class__.__name__} without passing a `layer_idx` is not recommended and will "486                "lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` "487                "when creating this class."488            )489 490        self.hidden_size = config.hidden_size491        self.num_heads = config.num_attention_heads492        self.head_dim = self.hidden_size // self.num_heads493        self.num_key_value_heads = config.num_key_value_heads494        self.num_key_value_groups = self.num_heads // self.num_key_value_heads495        self.max_position_embeddings = config.max_position_embeddings496        self.rope_theta = config.rope_theta497        self.is_causal = True498        self.attention_dropout = config.attention_dropout499 500        if (self.head_dim * self.num_heads) != self.hidden_size:501            raise ValueError(502                f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}"503                f" and `num_heads`: {self.num_heads})."504            )505        self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False)506        self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)507        self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)508        self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)509 510        self.rotary_emb = MistralRotaryEmbedding(511            self.head_dim,512            max_position_embeddings=self.max_position_embeddings,513            base=self.rope_theta,514        )515 516    def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):517        return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()518 519    def forward(520        self,521        hidden_states: torch.Tensor,522        attention_mask: Optional[torch.Tensor] = None,523        position_ids: Optional[torch.LongTensor] = None,524        past_key_value: Optional[Cache] = None,525        output_attentions: bool = False,526        use_cache: bool = False,527        **kwargs,528    ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:529        if "padding_mask" in kwargs:530            warnings.warn(531                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"532            )533        bsz, q_len, _ = hidden_states.size()534 535        query_states = self.q_proj(hidden_states)536        key_states = self.k_proj(hidden_states)537        value_states = self.v_proj(hidden_states)538 539        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)540        key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)541        value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)542 543        kv_seq_len = key_states.shape[-2]544        if past_key_value is not None:545            if self.layer_idx is None:546                raise ValueError(547                    f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "548                    "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "549                    "with a layer index."550                )551            kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)552        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)553        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)554 555        if past_key_value is not None:556            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models557            key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)558 559        # repeat k/v heads if n_kv_heads < n_heads560        key_states = repeat_kv(key_states, self.num_key_value_groups)561        value_states = repeat_kv(value_states, self.num_key_value_groups)562 563        attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)564 565        if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):566            raise ValueError(567                f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"568                f" {attn_weights.size()}"569            )570 571        if attention_mask is not None:572            if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):573                raise ValueError(574                    f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"575                )576 577            attn_weights = attn_weights + attention_mask578 579        # upcast attention to fp32580        attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)581        attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)582        attn_output = torch.matmul(attn_weights, value_states)583 584        if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):585            raise ValueError(586                f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"587                f" {attn_output.size()}"588            )589 590        attn_output = attn_output.transpose(1, 2).contiguous()591        attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)592 593        attn_output = self.o_proj(attn_output)594 595        if not output_attentions:596            attn_weights = None597 598        return attn_output, attn_weights, past_key_value599 600 601class MistralFlashAttention2(MistralAttention):602    """603    Mistral flash attention module. This module inherits from `MistralAttention` as the weights of the module stays604    untouched. The only required change would be on the forward pass where it needs to correctly call the public API of605    flash attention and deal with padding tokens in case the input contains any of them.606    """607 608    # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2.__init__609    def __init__(self, *args, **kwargs):610        super().__init__(*args, **kwargs)611 612        # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.613        # 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.614        # 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).615        self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()616 617    def forward(618        self,619        hidden_states: torch.Tensor,620        attention_mask: Optional[torch.Tensor] = None,621        position_ids: Optional[torch.LongTensor] = None,622        past_key_value: Optional[Cache] = None,623        output_attentions: bool = False,624        use_cache: bool = False,625        **kwargs,626    ):627        if "padding_mask" in kwargs:628            warnings.warn(629                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"630            )631 632            # overwrite attention_mask with padding_mask633            attention_mask = kwargs.pop("padding_mask")634        bsz, q_len, _ = hidden_states.size()635 636        query_states = self.q_proj(hidden_states)637        key_states = self.k_proj(hidden_states)638        value_states = self.v_proj(hidden_states)639 640        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)641        key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)642        value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)643 644        kv_seq_len = key_states.shape[-2]645        if past_key_value is not None:646            if self.layer_idx is None:647                raise ValueError(648                    f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "649                    "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "650                    "with a layer index."651                )652            kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)653 654        # Because the input can be padded, the absolute sequence length depends on the max position id.655        rotary_seq_len = max(kv_seq_len, position_ids[:, -1].max().item()) + 1656        cos, sin = self.rotary_emb(value_states, seq_len=rotary_seq_len)657 658        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)659 660        use_sliding_windows = (661            _flash_supports_window_size662            and getattr(self.config, "sliding_window", None) is not None663            and kv_seq_len > self.config.sliding_window664        )665 666        if not _flash_supports_window_size:667            logger.warning_once(668                "The current flash attention version does not support sliding window attention, for a more memory efficient implementation"669                " make sure to upgrade flash-attn library."670            )671 672        if past_key_value is not None:673            # Activate slicing cache only if the config has a value `sliding_windows` attribute674            cache_has_contents = past_key_value.get_seq_length(self.layer_idx) > 0675            if (676                getattr(self.config, "sliding_window", None) is not None677                and kv_seq_len > self.config.sliding_window678                and cache_has_contents679            ):680                slicing_tokens = 1 - self.config.sliding_window681 682                past_key = past_key_value[self.layer_idx][0]683                past_value = past_key_value[self.layer_idx][1]684 685                past_key = past_key[:, :, slicing_tokens:, :].contiguous()686                past_value = past_value[:, :, slicing_tokens:, :].contiguous()687 688                if past_key.shape[-2] != self.config.sliding_window - 1:689                    raise ValueError(690                        f"past key must have a shape of (`batch_size, num_heads, self.config.sliding_window-1, head_dim`), got"691                        f" {past_key.shape}"692                    )693 694                if attention_mask is not None:695                    attention_mask = attention_mask[:, slicing_tokens:]696                    attention_mask = torch.cat([attention_mask, torch.ones_like(attention_mask[:, -1:])], dim=-1)697 698            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models699            key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)700 701        # repeat k/v heads if n_kv_heads < n_heads702        key_states = repeat_kv(key_states, self.num_key_value_groups)703        value_states = repeat_kv(value_states, self.num_key_value_groups)704        dropout_rate = 0.0 if not self.training else self.attention_dropout705 706        # In PEFT, usually we cast the layer norms in float32 for training stability reasons707        # therefore the input hidden states gets silently casted in float32. Hence, we need708        # cast them back in float16 just to be sure everything works as expected.709        input_dtype = query_states.dtype710        if input_dtype == torch.float32:711            if torch.is_autocast_enabled():712                target_dtype = torch.get_autocast_gpu_dtype()713            # Handle the case where the model is quantized714            elif hasattr(self.config, "_pre_quantization_dtype"):715                target_dtype = self.config._pre_quantization_dtype716            else:717                target_dtype = self.q_proj.weight.dtype718 719            logger.warning_once(720                f"The input hidden states seems to be silently casted in float32, this might be related to"721                f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"722                f" {target_dtype}."723            )724 725            query_states = query_states.to(target_dtype)726            key_states = key_states.to(target_dtype)727            value_states = value_states.to(target_dtype)728 729        # Reashape to the expected shape for Flash Attention730        query_states = query_states.transpose(1, 2)731        key_states = key_states.transpose(1, 2)732        value_states = value_states.transpose(1, 2)733 734        attn_output = self._flash_attention_forward(735            query_states,736            key_states,737            value_states,738            attention_mask,739            q_len,740            dropout=dropout_rate,741            use_sliding_windows=use_sliding_windows,742        )743 744        attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()745        attn_output = self.o_proj(attn_output)746 747        if not output_attentions:748            attn_weights = None749 750        return attn_output, attn_weights, past_key_value751 752    def _flash_attention_forward(753        self,754        query_states,755        key_states,756        value_states,757        attention_mask,758        query_length,759        dropout=0.0,760        softmax_scale=None,761        use_sliding_windows=False,762    ):763        """764        Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token765        first unpad the input, then computes the attention scores and pad the final attention scores.766 767        Args:768            query_states (`torch.Tensor`):769                Input query states to be passed to Flash Attention API770            key_states (`torch.Tensor`):771                Input key states to be passed to Flash Attention API772            value_states (`torch.Tensor`):773                Input value states to be passed to Flash Attention API774            attention_mask (`torch.Tensor`):775                The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the776                position of padding tokens and 1 for the position of non-padding tokens.777            dropout (`float`):778                Attention dropout779            softmax_scale (`float`, *optional*):780                The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)781            use_sliding_windows (`bool`, *optional*):782                Whether to activate sliding window attention.783        """784        if not self._flash_attn_uses_top_left_mask:785            causal = self.is_causal786        else:787            # 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__.788            causal = self.is_causal and query_length != 1789 790        # Contains at least one padding token in the sequence791        if attention_mask is not None:792            batch_size = query_states.shape[0]793            query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(794                query_states, key_states, value_states, attention_mask, query_length795            )796 797            cu_seqlens_q, cu_seqlens_k = cu_seq_lens798            max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens799 800            if not use_sliding_windows:801                attn_output_unpad = flash_attn_varlen_func(802                    query_states,803                    key_states,804                    value_states,805                    cu_seqlens_q=cu_seqlens_q,806                    cu_seqlens_k=cu_seqlens_k,807                    max_seqlen_q=max_seqlen_in_batch_q,808                    max_seqlen_k=max_seqlen_in_batch_k,809                    dropout_p=dropout,810                    softmax_scale=softmax_scale,811                    causal=causal,812                )813            else:814                attn_output_unpad = flash_attn_varlen_func(815                    query_states,816                    key_states,817                    value_states,818                    cu_seqlens_q=cu_seqlens_q,819                    cu_seqlens_k=cu_seqlens_k,820                    max_seqlen_q=max_seqlen_in_batch_q,821                    max_seqlen_k=max_seqlen_in_batch_k,822                    dropout_p=dropout,823                    softmax_scale=softmax_scale,824                    causal=causal,825                    window_size=(self.config.sliding_window, self.config.sliding_window),826                )827 828            attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)829        else:830            if not use_sliding_windows:831                attn_output = flash_attn_func(832                    query_states,833                    key_states,834                    value_states,835                    dropout,836                    softmax_scale=softmax_scale,837                    causal=causal,838                )839            else:840                attn_output = flash_attn_func(841                    query_states,842                    key_states,843                    value_states,844                    dropout,845                    softmax_scale=softmax_scale,846                    causal=causal,847                    window_size=(self.config.sliding_window, self.config.sliding_window),848                )849 850        return attn_output851 852    def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):853        batch_size, kv_seq_len, num_heads, head_dim = key_layer.shape854 855        # On the first iteration we need to properly re-create the padding mask856        # by slicing it on the proper place857        if kv_seq_len != attention_mask.shape[-1]:858            attention_mask_num_tokens = attention_mask.shape[-1]859            attention_mask = attention_mask[:, attention_mask_num_tokens - kv_seq_len :]860 861        indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)862 863        key_layer = index_first_axis(key_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)864        value_layer = index_first_axis(value_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)865 866        if query_length == kv_seq_len:867            query_layer = index_first_axis(868                query_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k869            )870            cu_seqlens_q = cu_seqlens_k871            max_seqlen_in_batch_q = max_seqlen_in_batch_k872            indices_q = indices_k873        elif query_length == 1:874            max_seqlen_in_batch_q = 1875            cu_seqlens_q = torch.arange(876                batch_size + 1, dtype=torch.int32, device=query_layer.device877            )  # There is a memcpy here, that is very bad.878            indices_q = cu_seqlens_q[:-1]879            query_layer = query_layer.squeeze(1)880        else:881            # The -q_len: slice assumes left padding.882            attention_mask = attention_mask[:, -query_length:]883            query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)884 885        return (886            query_layer,887            key_layer,888            value_layer,889            indices_q,890            (cu_seqlens_q, cu_seqlens_k),891            (max_seqlen_in_batch_q, max_seqlen_in_batch_k),892        )893 894 895# copied from transformers.models.llama.modeling_llama.LlamaSdpaAttention with Llama->Mistral896# TODO @Arthur no longer copied from LLama after static cache897class MistralSdpaAttention(MistralAttention):898    """899    Mistral attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from900    `MistralAttention` as the weights of the module stays untouched. The only changes are on the forward pass to adapt to901    SDPA API.902    """903 904    # Adapted from MistralAttention.forward905    def forward(906        self,907        hidden_states: torch.Tensor,908        attention_mask: Optional[torch.Tensor] = None,909        position_ids: Optional[torch.LongTensor] = None,910        past_key_value: Optional[Cache] = None,911        output_attentions: bool = False,912        use_cache: bool = False,913    ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:914        if output_attentions:915            # TODO: Improve this warning with e.g. `model.config.attn_implementation = "manual"` once this is implemented.916            logger.warning_once(917                "MistralModel is using MistralSdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, "918                '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.'919            )920            return super().forward(921                hidden_states=hidden_states,922                attention_mask=attention_mask,923                position_ids=position_ids,924                past_key_value=past_key_value,925                output_attentions=output_attentions,926                use_cache=use_cache,927            )928 929        bsz, q_len, _ = hidden_states.size()930 931        query_states = self.q_proj(hidden_states)932        key_states = self.k_proj(hidden_states)933        value_states = self.v_proj(hidden_states)934 935        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)936        key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)937        value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)938 939        kv_seq_len = key_states.shape[-2]940        if past_key_value is not None:941            kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)942        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)943 944        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)945 946        if past_key_value is not None:947            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models948            key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)949 950        key_states = repeat_kv(key_states, self.num_key_value_groups)951        value_states = repeat_kv(value_states, self.num_key_value_groups)952 953        if attention_mask is not None:954            if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):955                raise ValueError(956                    f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"957                )958 959        # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with custom attn_mask,960        # Reference: https://github.com/pytorch/pytorch/issues/112577.961        if query_states.device.type == "cuda" and attention_mask is not None:962            query_states = query_states.contiguous()963            key_states = key_states.contiguous()964            value_states = value_states.contiguous()965 966        attn_output = torch.nn.functional.scaled_dot_product_attention(967            query_states,968            key_states,969            value_states,970            attn_mask=attention_mask,971            dropout_p=self.attention_dropout if self.training else 0.0,972            # 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.973            is_causal=self.is_causal and attention_mask is None and q_len > 1,974        )975 976        attn_output = attn_output.transpose(1, 2).contiguous()977        attn_output = attn_output.view(bsz, q_len, self.hidden_size)978 979        attn_output = self.o_proj(attn_output)980 981        return attn_output, None, past_key_value982 983 984MISTRAL_ATTENTION_CLASSES = {985    "eager": MistralAttention,986    "flash_attention_2": MistralFlashAttention2,987    "sdpa": MistralSdpaAttention,988}989 990 991class MistralDecoderLayer(nn.Module):992    def __init__(self, config: MistralConfig, layer_idx: int):993        super().__init__()994        self.hidden_size = config.hidden_size995 996        self.self_attn = MISTRAL_ATTENTION_CLASSES[config._attn_implementation](config, layer_idx)997 998        self.mlp = MistralMLP(config)999        self.input_layernorm = MistralRMSNorm(config.hidden_size, eps=config.rms_norm_eps)1000        self.post_attention_layernorm = MistralRMSNorm(config.hidden_size, eps=config.rms_norm_eps)1001 1002    def forward(1003        self,1004        hidden_states: torch.Tensor,1005        attention_mask: Optional[torch.Tensor] = None,1006        position_ids: Optional[torch.LongTensor] = None,1007        past_key_value: Optional[Tuple[torch.Tensor]] = None,1008        output_attentions: Optional[bool] = False,1009        use_cache: Optional[bool] = False,1010        **kwargs,1011    ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:1012        if "padding_mask" in kwargs:1013            warnings.warn(1014                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"1015            )1016        """1017        Args:1018            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`1019            attention_mask (`torch.FloatTensor`, *optional*): attention mask of size1020                `(batch, sequence_length)` where padding elements are indicated by 0.1021            output_attentions (`bool`, *optional*):1022                Whether or not to return the attentions tensors of all attention layers. See `attentions` under1023                returned tensors for more detail.1024            use_cache (`bool`, *optional*):1025                If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding1026                (see `past_key_values`).1027            past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states1028        """1029 1030        residual = hidden_states1031 1032        hidden_states = self.input_layernorm(hidden_states)1033 1034        # Self Attention1035        hidden_states, self_attn_weights, present_key_value = self.self_attn(1036            hidden_states=hidden_states,1037            attention_mask=attention_mask,1038            position_ids=position_ids,1039            past_key_value=past_key_value,1040            output_attentions=output_attentions,1041            use_cache=use_cache,1042        )1043        hidden_states = residual + hidden_states1044 1045        # Fully Connected1046        residual = hidden_states1047        hidden_states = self.post_attention_layernorm(hidden_states)1048        hidden_states = self.mlp(hidden_states)1049        hidden_states = residual + hidden_states1050 1051        outputs = (hidden_states,)1052 1053        if output_attentions:1054            outputs += (self_attn_weights,)1055 1056        if use_cache:1057            outputs += (present_key_value,)1058 1059        return outputs1060 1061 1062MISTRAL_START_DOCSTRING = r"""1063    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the1064    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads1065    etc.)1066 1067    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.1068    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage1069    and behavior.1070 1071    Parameters:1072        config ([`MistralConfig`]):1073            Model configuration class with all the parameters of the model. Initializing with a config file does not1074            load the weights associated with the model, only the configuration. Check out the1075            [`~PreTrainedModel.from_pretrained`] method to load the model weights.1076"""1077 1078 1079@add_start_docstrings(1080    "The bare Mistral Model outputting raw hidden-states without any specific head on top.",1081    MISTRAL_START_DOCSTRING,1082)1083class MistralPreTrainedModel(PreTrainedModel):1084    config_class = MistralConfig1085    base_model_prefix = "model"1086    supports_gradient_checkpointing = True1087    _no_split_modules = ["MistralDecoderLayer"]1088    _skip_keys_device_placement = "past_key_values"1089    _supports_flash_attn_2 = True1090    _supports_sdpa = True1091    _supports_cache_class = True1092 1093    def _init_weights(self, module):1094        std = self.config.initializer_range1095        if isinstance(module, nn.Linear):1096            module.weight.data.normal_(mean=0.0, std=std)1097            if module.bias is not None:1098                module.bias.data.zero_()1099        elif isinstance(module, nn.Embedding):1100            module.weight.data.normal_(mean=0.0, std=std)1101            if module.padding_idx is not None:1102                module.weight.data[module.padding_idx].zero_()1103 1104 1105MISTRAL_INPUTS_DOCSTRING = r"""1106    Args:1107        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):1108            Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide1109            it.1110 1111            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1112            [`PreTrainedTokenizer.__call__`] for details.1113 1114            [What are input IDs?](../glossary#input-ids)1115        attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1116            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1117 1118            - 1 for tokens that are **not masked**,1119            - 0 for tokens that are **masked**.1120 1121            [What are attention masks?](../glossary#attention-mask)1122 1123            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1124            [`PreTrainedTokenizer.__call__`] for details.1125 1126            If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see1127            `past_key_values`).1128 1129            If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]1130            and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more1131            information on the default strategy.1132 1133            - 1 indicates the head is **not masked**,1134            - 0 indicates the head is **masked**.1135        position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1136            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,1137            config.n_positions - 1]`.1138 1139            [What are position IDs?](../glossary#position-ids)1140        past_key_values (`Cache` or `tuple(tuple(torch.FloatTensor))`, *optional*):1141            Pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention1142            blocks) that can be used to speed up sequential decoding. This typically consists in the `past_key_values`1143            returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.1144 1145            Two formats are allowed:1146            - a [`~cache_utils.Cache`] instance;1147            - Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of1148            shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`). This is also known as the legacy1149            cache format.1150 1151            The model will output the same cache format that is fed as input. If no `past_key_values` are passed, the1152            legacy cache format will be returned.1153 1154            If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't1155            have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`1156            of shape `(batch_size, sequence_length)`.1157        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):1158            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This1159            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the1160            model's internal embedding lookup matrix.1161        use_cache (`bool`, *optional*):1162            If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see1163            `past_key_values`).1164        output_attentions (`bool`, *optional*):1165            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1166            tensors for more detail.1167        output_hidden_states (`bool`, *optional*):1168            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1169            more detail.1170        return_dict (`bool`, *optional*):1171            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1172"""1173 1174""" Classes to support Vision-Encoder-Text-Decoder architectures"""1175 1176VISION_ENCODER_DECODER_START_DOCSTRING = r"""1177    This class can be used to initialize an image-to-text-sequence model with any pretrained vision autoencoding model1178    as the encoder and any pretrained text autoregressive model as the decoder. The encoder is loaded via1179    [`~AutoModel.from_pretrained`] function and the decoder is loaded via [`~AutoModelForCausalLM.from_pretrained`]1180    function. Cross-attention layers are automatically added to the decoder and should be fine-tuned on a downstream1181    generative task, like image captioning.1182 1183    The effectiveness of initializing sequence-to-sequence models with pretrained checkpoints for sequence generation1184    tasks was shown in [Leveraging Pre-trained Checkpoints for Sequence Generation1185    Tasks](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn. Michael Matena, Yanqi1186    Zhou, Wei Li, Peter J. Liu.1187 1188    Additionally, in [TrOCR: Transformer-based Optical Character Recognition with Pre-trained1189    Models](https://arxiv.org/abs/2109.10282) it is shown how leveraging large pretrained vision models for optical1190    character recognition (OCR) yields a significant performance improvement.1191 1192    After such a Vision-Encoder-Text-Decoder model has been trained/fine-tuned, it can be saved/loaded just like any1193    other models (see the examples for more information).1194 1195    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the1196    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads1197    etc.)1198 1199    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.1200    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage

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