LeroyDyer/SpydazWebAI_VisionEncoderDecoderModel_Mini3b
013
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