optimum-intel-internal-testing/tiny-random-phi-4-multimodal
014k
1# coding=utf-82# Copyright 2024 The HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15""" Siglip model configuration"""16 17import os18from typing import Union19 20from transformers.configuration_utils import PretrainedConfig21from transformers.utils import logging22 23 24logger = logging.get_logger(__name__)25 26SIGLIP_PRETRAINED_CONFIG_ARCHIVE_MAP = {27 "google/siglip-base-patch16-224": "https://huggingface.co/google/siglip-base-patch16-224/resolve/main/config.json",28}29 30 31class SiglipTextConfig(PretrainedConfig):32 r"""33 This is the configuration class to store the configuration of a [`SiglipTextModel`]. It is used to instantiate a34 Siglip text encoder according to the specified arguments, defining the model architecture. Instantiating a35 configuration with the defaults will yield a similar configuration to that of the text encoder of the Siglip36 [google/siglip-base-patch16-224](https://huggingface.co/google/siglip-base-patch16-224) architecture.37 Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the38 documentation from [`PretrainedConfig`] for more information.39 Args:40 vocab_size (`int`, *optional*, defaults to 32000):41 Vocabulary size of the Siglip text model. Defines the number of different tokens that can be represented by42 the `inputs_ids` passed when calling [`SiglipModel`].43 hidden_size (`int`, *optional*, defaults to 768):44 Dimensionality of the encoder layers and the pooler layer.45 intermediate_size (`int`, *optional*, defaults to 3072):46 Dimensionality of the "intermediate" (i.e., feed-forward) layer in the Transformer encoder.47 num_hidden_layers (`int`, *optional*, defaults to 12):48 Number of hidden layers in the Transformer encoder.49 num_attention_heads (`int`, *optional*, defaults to 12):50 Number of attention heads for each attention layer in the Transformer encoder.51 max_position_embeddings (`int`, *optional*, defaults to 64):52 The maximum sequence length that this model might ever be used with. Typically set this to something large53 just in case (e.g., 512 or 1024 or 2048).54 hidden_act (`str` or `function`, *optional*, defaults to `"gelu_pytorch_tanh"`):55 The non-linear activation function (function or string) in the encoder and pooler. If string, `"gelu"`,56 `"relu"`, `"selu"` and `"gelu_new"` `"quick_gelu"` are supported.57 layer_norm_eps (`float`, *optional*, defaults to 1e-06):58 The epsilon used by the layer normalization layers.59 attention_dropout (`float`, *optional*, defaults to 0.0):60 The dropout ratio for the attention probabilities.61 pad_token_id (`int`, *optional*, defaults to 1):62 The id of the padding token in the vocabulary.63 bos_token_id (`int`, *optional*, defaults to 49406):64 The id of the beginning-of-sequence token in the vocabulary.65 eos_token_id (`int`, *optional*, defaults to 49407):66 The id of the end-of-sequence token in the vocabulary.67 Example:68 ```python69 >>> from transformers import SiglipTextConfig, SiglipTextModel70 >>> # Initializing a SiglipTextConfig with google/siglip-base-patch16-224 style configuration71 >>> configuration = SiglipTextConfig()72 >>> # Initializing a SiglipTextModel (with random weights) from the google/siglip-base-patch16-224 style configuration73 >>> model = SiglipTextModel(configuration)74 >>> # Accessing the model configuration75 >>> configuration = model.config76 ```"""77 78 model_type = "siglip_text_model"79 80 def __init__(81 self,82 vocab_size=32000,83 hidden_size=768,84 intermediate_size=3072,85 num_hidden_layers=12,86 num_attention_heads=12,87 max_position_embeddings=64,88 hidden_act="gelu_pytorch_tanh",89 layer_norm_eps=1e-6,90 attention_dropout=0.0,91 # This differs from `CLIPTokenizer`'s default and from openai/siglip92 # See https://github.com/huggingface/transformers/pull/24773#issuecomment-163228753893 pad_token_id=1,94 bos_token_id=49406,95 eos_token_id=49407,96 _flash_attn_2_enabled=True,97 **kwargs,98 ):99 super().__init__(pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, **kwargs)100 101 self.vocab_size = vocab_size102 self.hidden_size = hidden_size103 self.intermediate_size = intermediate_size104 self.num_hidden_layers = num_hidden_layers105 self.num_attention_heads = num_attention_heads106 self.max_position_embeddings = max_position_embeddings107 self.layer_norm_eps = layer_norm_eps108 self.hidden_act = hidden_act109 self.attention_dropout = attention_dropout110 self._flash_attn_2_enabled = _flash_attn_2_enabled111 112 @classmethod113 def from_pretrained(cls, pretrained_model_name_or_path: Union[str, os.PathLike], **kwargs) -> "PretrainedConfig":114 cls._set_token_in_kwargs(kwargs)115 116 config_dict, kwargs = cls.get_config_dict(pretrained_model_name_or_path, **kwargs)117 118 # get the text config dict if we are loading from SiglipConfig119 if config_dict.get("model_type") == "siglip":120 config_dict = config_dict["text_config"]121 122 if "model_type" in config_dict and hasattr(cls, "model_type") and config_dict["model_type"] != cls.model_type:123 logger.warning(124 f"You are using a model of type {config_dict['model_type']} to instantiate a model of type "125 f"{cls.model_type}. This is not supported for all configurations of models and can yield errors."126 )127 128 return cls.from_dict(config_dict, **kwargs)129 130 131class SiglipVisionConfig(PretrainedConfig):132 r"""133 This is the configuration class to store the configuration of a [`SiglipVisionModel`]. It is used to instantiate a134 Siglip vision encoder according to the specified arguments, defining the model architecture. Instantiating a135 configuration with the defaults will yield a similar configuration to that of the vision encoder of the Siglip136 [google/siglip-base-patch16-224](https://huggingface.co/google/siglip-base-patch16-224) architecture.137 Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the138 documentation from [`PretrainedConfig`] for more information.139 Args:140 hidden_size (`int`, *optional*, defaults to 768):141 Dimensionality of the encoder layers and the pooler layer.142 intermediate_size (`int`, *optional*, defaults to 3072):143 Dimensionality of the "intermediate" (i.e., feed-forward) layer in the Transformer encoder.144 num_hidden_layers (`int`, *optional*, defaults to 12):145 Number of hidden layers in the Transformer encoder.146 num_attention_heads (`int`, *optional*, defaults to 12):147 Number of attention heads for each attention layer in the Transformer encoder.148 num_channels (`int`, *optional*, defaults to 3):149 Number of channels in the input images.150 image_size (`int`, *optional*, defaults to 224):151 The size (resolution) of each image.152 patch_size (`int`, *optional*, defaults to 16):153 The size (resolution) of each patch.154 hidden_act (`str` or `function`, *optional*, defaults to `"gelu_pytorch_tanh"`):155 The non-linear activation function (function or string) in the encoder and pooler. If string, `"gelu"`,156 `"relu"`, `"selu"` and `"gelu_new"` ``"quick_gelu"` are supported.157 layer_norm_eps (`float`, *optional*, defaults to 1e-06):158 The epsilon used by the layer normalization layers.159 attention_dropout (`float`, *optional*, defaults to 0.0):160 The dropout ratio for the attention probabilities.161 Example:162 ```python163 >>> from transformers import SiglipVisionConfig, SiglipVisionModel164 >>> # Initializing a SiglipVisionConfig with google/siglip-base-patch16-224 style configuration165 >>> configuration = SiglipVisionConfig()166 >>> # Initializing a SiglipVisionModel (with random weights) from the google/siglip-base-patch16-224 style configuration167 >>> model = SiglipVisionModel(configuration)168 >>> # Accessing the model configuration169 >>> configuration = model.config170 ```"""171 172 model_type = "siglip_vision_model"173 174 def __init__(175 self,176 hidden_size=768,177 intermediate_size=3072,178 num_hidden_layers=12,179 num_attention_heads=12,180 num_channels=3,181 image_size=224,182 patch_size=16,183 hidden_act="gelu_pytorch_tanh",184 layer_norm_eps=1e-6,185 attention_dropout=0.0,186 _flash_attn_2_enabled=True,187 **kwargs,188 ):189 super().__init__(**kwargs)190 191 self.hidden_size = hidden_size192 self.intermediate_size = intermediate_size193 self.num_hidden_layers = num_hidden_layers194 self.num_attention_heads = num_attention_heads195 self.num_channels = num_channels196 self.patch_size = patch_size197 self.image_size = image_size198 self.attention_dropout = attention_dropout199 self.layer_norm_eps = layer_norm_eps200 self.hidden_act = hidden_act201 self._flash_attn_2_enabled = _flash_attn_2_enabled202 203 @classmethod204 def from_pretrained(cls, pretrained_model_name_or_path: Union[str, os.PathLike], **kwargs) -> "PretrainedConfig":205 cls._set_token_in_kwargs(kwargs)206 207 config_dict, kwargs = cls.get_config_dict(pretrained_model_name_or_path, **kwargs)208 209 # get the vision config dict if we are loading from SiglipConfig210 if config_dict.get("model_type") == "siglip":211 config_dict = config_dict["vision_config"]212 213 if "model_type" in config_dict and hasattr(cls, "model_type") and config_dict["model_type"] != cls.model_type:214 logger.warning(215 f"You are using a model of type {config_dict['model_type']} to instantiate a model of type "216 f"{cls.model_type}. This is not supported for all configurations of models and can yield errors."217 )218 219 return cls.from_dict(config_dict, **kwargs)220 221 222class SiglipConfig(PretrainedConfig):223 r"""224 [`SiglipConfig`] is the configuration class to store the configuration of a [`SiglipModel`]. It is used to225 instantiate a Siglip model according to the specified arguments, defining the text model and vision model configs.226 Instantiating a configuration with the defaults will yield a similar configuration to that of the Siglip227 [google/siglip-base-patch16-224](https://huggingface.co/google/siglip-base-patch16-224) architecture.228 Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the229 documentation from [`PretrainedConfig`] for more information.230 Args:231 text_config (`dict`, *optional*):232 Dictionary of configuration options used to initialize [`SiglipTextConfig`].233 vision_config (`dict`, *optional*):234 Dictionary of configuration options used to initialize [`SiglipVisionConfig`].235 kwargs (*optional*):236 Dictionary of keyword arguments.237 Example:238 ```python239 >>> from transformers import SiglipConfig, SiglipModel240 >>> # Initializing a SiglipConfig with google/siglip-base-patch16-224 style configuration241 >>> configuration = SiglipConfig()242 >>> # Initializing a SiglipModel (with random weights) from the google/siglip-base-patch16-224 style configuration243 >>> model = SiglipModel(configuration)244 >>> # Accessing the model configuration245 >>> configuration = model.config246 >>> # We can also initialize a SiglipConfig from a SiglipTextConfig and a SiglipVisionConfig247 >>> from transformers import SiglipTextConfig, SiglipVisionConfig248 >>> # Initializing a SiglipText and SiglipVision configuration249 >>> config_text = SiglipTextConfig()250 >>> config_vision = SiglipVisionConfig()251 >>> config = SiglipConfig.from_text_vision_configs(config_text, config_vision)252 ```"""253 254 model_type = "siglip"255 256 def __init__(self, text_config=None, vision_config=None, **kwargs):257 super().__init__(**kwargs)258 259 if text_config is None:260 text_config = {}261 logger.info("`text_config` is `None`. Initializing the `SiglipTextConfig` with default values.")262 263 if vision_config is None:264 vision_config = {}265 logger.info("`vision_config` is `None`. initializing the `SiglipVisionConfig` with default values.")266 267 self.text_config = SiglipTextConfig(**text_config)268 self.vision_config = SiglipVisionConfig(**vision_config)269 270 self.initializer_factor = 1.0271 272 @classmethod273 def from_text_vision_configs(cls, text_config: SiglipTextConfig, vision_config: SiglipVisionConfig, **kwargs):274 r"""275 Instantiate a [`SiglipConfig`] (or a derived class) from siglip text model configuration and siglip vision276 model configuration.277 Returns:278 [`SiglipConfig`]: An instance of a configuration object279 """280 281 return cls(text_config=text_config.to_dict(), vision_config=vision_config.to_dict(), **kwargs)282 283# coding=utf-8284# Copyright 2024 Google AI and The HuggingFace Team. All rights reserved.285#286# Licensed under the Apache License, Version 2.0 (the "License");287# you may not use this file except in compliance with the License.288# You may obtain a copy of the License at289#290# http://www.apache.org/licenses/LICENSE-2.0291#292# Unless required by applicable law or agreed to in writing, software293# distributed under the License is distributed on an "AS IS" BASIS,294# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.295# See the License for the specific language governing permissions and296# limitations under the License.297""" PyTorch Siglip model."""298 299 300import math301import warnings302from dataclasses import dataclass303from typing import Any, Optional, Tuple, Union304 305import numpy as np306import torch307import torch.nn.functional as F308import torch.utils.checkpoint309from torch import nn310from torch.nn.init import _calculate_fan_in_and_fan_out311 312from transformers.activations import ACT2FN313from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask314from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling315from transformers.modeling_utils import PreTrainedModel316from transformers.utils import (317 ModelOutput,318 add_start_docstrings,319 add_start_docstrings_to_model_forward,320 is_flash_attn_2_available,321 logging,322 replace_return_docstrings,323)324 325logger = logging.get_logger(__name__)326 327_CHECKPOINT_FOR_DOC = "google/siglip-base-patch16-224"328 329SIGLIP_PRETRAINED_MODEL_ARCHIVE_LIST = [330 "google/siglip-base-patch16-224",331 # See all SigLIP models at https://huggingface.co/models?filter=siglip332]333 334if is_flash_attn_2_available():335 from flash_attn import flash_attn_func, flash_attn_varlen_func336 from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa337 338 339# Copied from transformers.models.llama.modeling_llama._get_unpad_data340def _get_unpad_data(attention_mask):341 seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)342 indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()343 max_seqlen_in_batch = seqlens_in_batch.max().item()344 cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))345 return (346 indices,347 cu_seqlens,348 max_seqlen_in_batch,349 )350 351 352def _trunc_normal_(tensor, mean, std, a, b):353 # Cut & paste from PyTorch official master until it's in a few official releases - RW354 # Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf355 def norm_cdf(x):356 # Computes standard normal cumulative distribution function357 return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0358 359 if (mean < a - 2 * std) or (mean > b + 2 * std):360 warnings.warn(361 "mean is more than 2 std from [a, b] in nn.init.trunc_normal_. "362 "The distribution of values may be incorrect.",363 stacklevel=2,364 )365 366 # Values are generated by using a truncated uniform distribution and367 # then using the inverse CDF for the normal distribution.368 # Get upper and lower cdf values369 l = norm_cdf((a - mean) / std)370 u = norm_cdf((b - mean) / std)371 372 # Uniformly fill tensor with values from [l, u], then translate to373 # [2l-1, 2u-1].374 tensor.uniform_(2 * l - 1, 2 * u - 1)375 376 # Use inverse cdf transform for normal distribution to get truncated377 # standard normal378 if tensor.dtype in [torch.float16, torch.bfloat16]:379 # The `erfinv_` op is not (yet?) defined in float16+cpu, bfloat16+gpu380 og_dtype = tensor.dtype381 tensor = tensor.to(torch.float32)382 tensor.erfinv_()383 tensor = tensor.to(og_dtype)384 else:385 tensor.erfinv_()386 387 # Transform to proper mean, std388 tensor.mul_(std * math.sqrt(2.0))389 tensor.add_(mean)390 391 # Clamp to ensure it's in the proper range392 if tensor.dtype == torch.float16:393 # The `clamp_` op is not (yet?) defined in float16+cpu394 tensor = tensor.to(torch.float32)395 tensor.clamp_(min=a, max=b)396 tensor = tensor.to(torch.float16)397 else:398 tensor.clamp_(min=a, max=b)399 400 401def trunc_normal_tf_(402 tensor: torch.Tensor, mean: float = 0.0, std: float = 1.0, a: float = -2.0, b: float = 2.0403) -> torch.Tensor:404 """Fills the input Tensor with values drawn from a truncated405 normal distribution. The values are effectively drawn from the406 normal distribution :math:`\\mathcal{N}(\text{mean}, \text{std}^2)`407 with values outside :math:`[a, b]` redrawn until they are within408 the bounds. The method used for generating the random values works409 best when :math:`a \\leq \text{mean} \\leq b`.410 NOTE: this 'tf' variant behaves closer to Tensorflow / JAX impl where the411 bounds [a, b] are applied when sampling the normal distribution with mean=0, std=1.0412 and the result is subsquently scaled and shifted by the mean and std args.413 Args:414 tensor: an n-dimensional `torch.Tensor`415 mean: the mean of the normal distribution416 std: the standard deviation of the normal distribution417 a: the minimum cutoff value418 b: the maximum cutoff value419 """420 with torch.no_grad():421 _trunc_normal_(tensor, 0, 1.0, a, b)422 tensor.mul_(std).add_(mean)423 424 425def variance_scaling_(tensor, scale=1.0, mode="fan_in", distribution="normal"):426 fan_in, fan_out = _calculate_fan_in_and_fan_out(tensor)427 if mode == "fan_in":428 denom = fan_in429 elif mode == "fan_out":430 denom = fan_out431 elif mode == "fan_avg":432 denom = (fan_in + fan_out) / 2433 434 variance = scale / denom435 436 if distribution == "truncated_normal":437 # constant is stddev of standard normal truncated to (-2, 2)438 trunc_normal_tf_(tensor, std=math.sqrt(variance) / 0.87962566103423978)439 elif distribution == "normal":440 with torch.no_grad():441 tensor.normal_(std=math.sqrt(variance))442 elif distribution == "uniform":443 bound = math.sqrt(3 * variance)444 with torch.no_grad():445 tensor.uniform_(-bound, bound)446 else:447 raise ValueError(f"invalid distribution {distribution}")448 449 450def lecun_normal_(tensor):451 variance_scaling_(tensor, mode="fan_in", distribution="truncated_normal")452 453 454def default_flax_embed_init(tensor):455 variance_scaling_(tensor, mode="fan_in", distribution="normal")456 457 458@dataclass459# Copied from transformers.models.clip.modeling_clip.CLIPVisionModelOutput with CLIP->Siglip460class SiglipVisionModelOutput(ModelOutput):461 """462 Base class for vision model's outputs that also contains image embeddings of the pooling of the last hidden states.463 Args:464 image_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim)` *optional* returned when model is initialized with `with_projection=True`):465 The image embeddings obtained by applying the projection layer to the pooler_output.466 last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):467 Sequence of hidden-states at the output of the last layer of the model.468 hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):469 Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +470 one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.471 Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.472 attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):473 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,474 sequence_length)`.475 Attentions weights after the attention softmax, used to compute the weighted average in the self-attention476 heads.477 """478 479 image_embeds: Optional[torch.FloatTensor] = None480 last_hidden_state: torch.FloatTensor = None481 hidden_states: Optional[Tuple[torch.FloatTensor]] = None482 attentions: Optional[Tuple[torch.FloatTensor]] = None483 484 485@dataclass486# Copied from transformers.models.clip.modeling_clip.CLIPTextModelOutput with CLIP->Siglip487class SiglipTextModelOutput(ModelOutput):488 """489 Base class for text model's outputs that also contains a pooling of the last hidden states.490 Args:491 text_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim)` *optional* returned when model is initialized with `with_projection=True`):492 The text embeddings obtained by applying the projection layer to the pooler_output.493 last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):494 Sequence of hidden-states at the output of the last layer of the model.495 hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):496 Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +497 one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.498 Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.499 attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):500 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,501 sequence_length)`.502 Attentions weights after the attention softmax, used to compute the weighted average in the self-attention503 heads.504 """505 506 text_embeds: Optional[torch.FloatTensor] = None507 last_hidden_state: torch.FloatTensor = None508 hidden_states: Optional[Tuple[torch.FloatTensor]] = None509 attentions: Optional[Tuple[torch.FloatTensor]] = None510 511 512@dataclass513# Copied from transformers.models.clip.modeling_clip.CLIPOutput with CLIP->Siglip514class SiglipOutput(ModelOutput):515 """516 Args:517 loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):518 Contrastive loss for image-text similarity.519 logits_per_image:(`torch.FloatTensor` of shape `(image_batch_size, text_batch_size)`):520 The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text521 similarity scores.522 logits_per_text:(`torch.FloatTensor` of shape `(text_batch_size, image_batch_size)`):523 The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image524 similarity scores.525 text_embeds(`torch.FloatTensor` of shape `(batch_size, output_dim`):526 The text embeddings obtained by applying the projection layer to the pooled output of [`SiglipTextModel`].527 image_embeds(`torch.FloatTensor` of shape `(batch_size, output_dim`):528 The image embeddings obtained by applying the projection layer to the pooled output of [`SiglipVisionModel`].529 text_model_output(`BaseModelOutputWithPooling`):530 The output of the [`SiglipTextModel`].531 vision_model_output(`BaseModelOutputWithPooling`):532 The output of the [`SiglipVisionModel`].533 """534 535 loss: Optional[torch.FloatTensor] = None536 logits_per_image: torch.FloatTensor = None537 logits_per_text: torch.FloatTensor = None538 text_embeds: torch.FloatTensor = None539 image_embeds: torch.FloatTensor = None540 text_model_output: BaseModelOutputWithPooling = None541 vision_model_output: BaseModelOutputWithPooling = None542 543 def to_tuple(self) -> Tuple[Any]:544 return tuple(545 self[k] if k not in ["text_model_output", "vision_model_output"] else getattr(self, k).to_tuple()546 for k in self.keys()547 )548 549 550class SiglipVisionEmbeddings(nn.Module):551 def __init__(self, config: SiglipVisionConfig):552 super().__init__()553 self.config = config554 self.embed_dim = config.hidden_size555 self.image_size = config.image_size556 self.patch_size = config.patch_size557 558 self.patch_embedding = nn.Conv2d(559 in_channels=config.num_channels,560 out_channels=self.embed_dim,561 kernel_size=self.patch_size,562 stride=self.patch_size,563 padding="valid",564 )565 566 self.num_patches_per_side = self.image_size // self.patch_size567 self.num_patches = self.num_patches_per_side**2568 self.num_positions = self.num_patches569 self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)570 571 def forward(self, pixel_values: torch.FloatTensor, patch_attention_mask: torch.BoolTensor) -> torch.Tensor:572 batch_size = pixel_values.size(0)573 574 patch_embeds = self.patch_embedding(pixel_values)575 embeddings = patch_embeds.flatten(2).transpose(1, 2)576 577 max_im_h, max_im_w = pixel_values.size(2), pixel_values.size(3)578 max_nb_patches_h, max_nb_patches_w = max_im_h // self.patch_size, max_im_w // self.patch_size579 boundaries = torch.arange(1 / self.num_patches_per_side, 1.0, 1 / self.num_patches_per_side)580 position_ids = torch.full(581 size=(582 batch_size,583 max_nb_patches_h * max_nb_patches_w,584 ),585 fill_value=0,586 )587 588 for batch_idx, p_attn_mask in enumerate(patch_attention_mask):589 nb_patches_h = p_attn_mask[:, 0].sum()590 nb_patches_w = p_attn_mask[0].sum()591 592 fractional_coords_h = torch.arange(0, 1 - 1e-6, 1 / nb_patches_h)593 fractional_coords_w = torch.arange(0, 1 - 1e-6, 1 / nb_patches_w)594 595 bucket_coords_h = torch.bucketize(fractional_coords_h, boundaries, right=True)596 bucket_coords_w = torch.bucketize(fractional_coords_w, boundaries, right=True)597 598 pos_ids = (bucket_coords_h[:, None] * self.num_patches_per_side + bucket_coords_w).flatten()599 position_ids[batch_idx][p_attn_mask.view(-1).cpu()] = pos_ids600 601 position_ids = position_ids.to(self.position_embedding.weight.device)602 603 embeddings = embeddings + self.position_embedding(position_ids)604 return embeddings605 606 607# Copied from transformers.models.clip.modeling_clip.CLIPTextEmbeddings with CLIP->Siglip608class SiglipTextEmbeddings(nn.Module):609 def __init__(self, config: SiglipTextConfig):610 super().__init__()611 embed_dim = config.hidden_size612 613 self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)614 self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)615 616 # position_ids (1, len position emb) is contiguous in memory and exported when serialized617 self.register_buffer(618 "position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False619 )620 621 def forward(622 self,623 input_ids: Optional[torch.LongTensor] = None,624 position_ids: Optional[torch.LongTensor] = None,625 inputs_embeds: Optional[torch.FloatTensor] = None,626 ) -> torch.Tensor:627 seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]628 629 if position_ids is None:630 position_ids = self.position_ids[:, :seq_length]631 632 if inputs_embeds is None:633 inputs_embeds = self.token_embedding(input_ids)634 635 position_embeddings = self.position_embedding(position_ids)636 embeddings = inputs_embeds + position_embeddings637 638 return embeddings639 640 641class SiglipAttention(nn.Module):642 """Multi-headed attention from 'Attention Is All You Need' paper"""643 644 # Copied from transformers.models.clip.modeling_clip.CLIPAttention.__init__645 def __init__(self, config):646 super().__init__()647 self.config = config648 self.embed_dim = config.hidden_size649 self.num_heads = config.num_attention_heads650 self.head_dim = self.embed_dim // self.num_heads651 if self.head_dim * self.num_heads != self.embed_dim:652 raise ValueError(653 f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"654 f" {self.num_heads})."655 )656 self.scale = self.head_dim**-0.5657 self.dropout = config.attention_dropout658 659 self.k_proj = nn.Linear(self.embed_dim, self.embed_dim)660 self.v_proj = nn.Linear(self.embed_dim, self.embed_dim)661 self.q_proj = nn.Linear(self.embed_dim, self.embed_dim)662 self.out_proj = nn.Linear(self.embed_dim, self.embed_dim)663 664 def forward(665 self,666 hidden_states: torch.Tensor,667 attention_mask: Optional[torch.Tensor] = None,668 output_attentions: Optional[bool] = False,669 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:670 """Input shape: Batch x Time x Channel"""671 672 batch_size, q_len, _ = hidden_states.size()673 674 query_states = self.q_proj(hidden_states)675 key_states = self.k_proj(hidden_states)676 value_states = self.v_proj(hidden_states)677 678 query_states = query_states.view(batch_size, q_len, self.num_heads, self.head_dim).transpose(1, 2)679 key_states = key_states.view(batch_size, q_len, self.num_heads, self.head_dim).transpose(1, 2)680 value_states = value_states.view(batch_size, q_len, self.num_heads, self.head_dim).transpose(1, 2)681 682 k_v_seq_len = key_states.shape[-2]683 attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) * self.scale684 685 if attn_weights.size() != (batch_size, self.num_heads, q_len, k_v_seq_len):686 raise ValueError(687 f"Attention weights should be of size {(batch_size, self.num_heads, q_len, k_v_seq_len)}, but is"688 f" {attn_weights.size()}"689 )690 691 if attention_mask is not None:692 if attention_mask.size() != (batch_size, 1, q_len, k_v_seq_len):693 raise ValueError(694 f"Attention mask should be of size {(batch_size, 1, q_len, k_v_seq_len)}, but is {attention_mask.size()}"695 )696 attn_weights = attn_weights + attention_mask697 698 # upcast attention to fp32699 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)700 attn_weights = nn.functional.dropout(attn_weights, p=self.dropout, training=self.training)701 attn_output = torch.matmul(attn_weights, value_states)702 703 if attn_output.size() != (batch_size, self.num_heads, q_len, self.head_dim):704 raise ValueError(705 f"`attn_output` should be of size {(batch_size, self.num_heads, q_len, self.head_dim)}, but is"706 f" {attn_output.size()}"707 )708 709 attn_output = attn_output.transpose(1, 2).contiguous()710 attn_output = attn_output.reshape(batch_size, q_len, self.embed_dim)711 712 attn_output = self.out_proj(attn_output)713 714 return attn_output, attn_weights715 716 717class SiglipFlashAttention2(SiglipAttention):718 """719 Llama flash attention module. This module inherits from `LlamaAttention` as the weights of the module stays720 untouched. The only required change would be on the forward pass where it needs to correctly call the public API of721 flash attention and deal with padding tokens in case the input contains any of them.722 """723 724 def __init__(self, *args, **kwargs):725 super().__init__(*args, **kwargs)726 self.is_causal = False # Hack to make sure we don't use a causal mask727 728 def forward(729 self,730 hidden_states: torch.Tensor,731 attention_mask: Optional[torch.LongTensor] = None,732 position_ids: Optional[torch.LongTensor] = None,733 past_key_value: Optional[Tuple[torch.Tensor]] = None,734 output_attentions: bool = False,735 use_cache: bool = False,736 **kwargs,737 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:738 output_attentions = False739 740 bsz, q_len, _ = hidden_states.size()741 742 query_states = self.q_proj(hidden_states)743 key_states = self.k_proj(hidden_states)744 value_states = self.v_proj(hidden_states)745 746 # Flash attention requires the input to have the shape747 # batch_size x seq_length x head_dim x hidden_dim748 # therefore we just need to keep the original shape749 query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)750 key_states = key_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)751 value_states = value_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)752 753 kv_seq_len = key_states.shape[-2]754 if past_key_value is not None:755 kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)756 # cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)757 # query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)758 759 # if past_key_value is not None:760 # cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models761 # key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)762 763 # TODO: These transpose are quite inefficient but Flash Attention requires the layout [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache764 # to be able to avoid many of these transpose/reshape/view.765 query_states = query_states.transpose(1, 2)766 key_states = key_states.transpose(1, 2)767 value_states = value_states.transpose(1, 2)768 769 dropout_rate = self.dropout if self.training else 0.0770 771 # In PEFT, usually we cast the layer norms in float32 for training stability reasons772 # therefore the input hidden states gets silently casted in float32. Hence, we need773 # cast them back in the correct dtype just to be sure everything works as expected.774 # This might slowdown training & inference so it is recommended to not cast the LayerNorms775 # in fp32. (LlamaRMSNorm handles it correctly)776 777 input_dtype = query_states.dtype778 if input_dtype == torch.float32:779 if torch.is_autocast_enabled():780 target_dtype = torch.get_autocast_gpu_dtype()781 # Handle the case where the model is quantized782 elif hasattr(self.config, "_pre_quantization_dtype"):783 target_dtype = self.config._pre_quantization_dtype784 else:785 target_dtype = self.q_proj.weight.dtype786 787 logger.warning_once(788 "The input hidden states seems to be silently casted in float32, this might be related to the fact"789 " you have upcasted embedding or layer norm layers in float32. We will cast back the input in"790 f" {target_dtype}."791 )792 793 query_states = query_states.to(target_dtype)794 key_states = key_states.to(target_dtype)795 value_states = value_states.to(target_dtype)796 797 attn_output = self._flash_attention_forward(798 query_states, key_states, value_states, attention_mask, q_len, dropout=dropout_rate799 )800 801 attn_output = attn_output.reshape(bsz, q_len, self.embed_dim).contiguous()802 attn_output = self.out_proj(attn_output)803 804 if not output_attentions:805 attn_weights = None806 807 return attn_output, attn_weights808 809 def _flash_attention_forward(810 self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None811 ):812 """813 Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token814 first unpad the input, then computes the attention scores and pad the final attention scores.815 Args:816 query_states (`torch.Tensor`):817 Input query states to be passed to Flash Attention API818 key_states (`torch.Tensor`):819 Input key states to be passed to Flash Attention API820 value_states (`torch.Tensor`):821 Input value states to be passed to Flash Attention API822 attention_mask (`torch.Tensor`):823 The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the824 position of padding tokens and 1 for the position of non-padding tokens.825 dropout (`int`, *optional*):826 Attention dropout827 softmax_scale (`float`, *optional*):828 The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)829 """830 831 # 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__.832 causal = self.is_causal and query_length != 1833 834 # Contains at least one padding token in the sequence835 if attention_mask is not None:836 batch_size = query_states.shape[0]837 query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(838 query_states, key_states, value_states, attention_mask, query_length839 )840 841 cu_seqlens_q, cu_seqlens_k = cu_seq_lens842 max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens843 844 attn_output_unpad = flash_attn_varlen_func(845 query_states,846 key_states,847 value_states,848 cu_seqlens_q=cu_seqlens_q,849 cu_seqlens_k=cu_seqlens_k,850 max_seqlen_q=max_seqlen_in_batch_q,851 max_seqlen_k=max_seqlen_in_batch_k,852 dropout_p=dropout,853 softmax_scale=softmax_scale,854 causal=causal,855 )856 857 attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)858 else:859 attn_output = flash_attn_func(860 query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal861 )862 863 return attn_output864 865 def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):866 indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)867 batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape868 869 key_layer = index_first_axis(870 key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k871 )872 value_layer = index_first_axis(873 value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k874 )875 if query_length == kv_seq_len:876 query_layer = index_first_axis(877 query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k878 )879 cu_seqlens_q = cu_seqlens_k880 max_seqlen_in_batch_q = max_seqlen_in_batch_k881 indices_q = indices_k882 elif query_length == 1:883 max_seqlen_in_batch_q = 1884 cu_seqlens_q = torch.arange(885 batch_size + 1, dtype=torch.int32, device=query_layer.device886 ) # There is a memcpy here, that is very bad.887 indices_q = cu_seqlens_q[:-1]888 query_layer = query_layer.squeeze(1)889 else:890 # The -q_len: slice assumes left padding.891 attention_mask = attention_mask[:, -query_length:]892 query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)893 894 return (895 query_layer,896 key_layer,897 value_layer,898 indices_q,899 (cu_seqlens_q, cu_seqlens_k),900 (max_seqlen_in_batch_q, max_seqlen_in_batch_k),901 )902 903 904# Copied from transformers.models.clip.modeling_clip.CLIPMLP with CLIP->Siglip905class SiglipMLP(nn.Module):906 def __init__(self, config):907 super().__init__()908 self.config = config909 self.activation_fn = ACT2FN[config.hidden_act]910 self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)911 self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)912 913 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:914 hidden_states = self.fc1(hidden_states)915 hidden_states = self.activation_fn(hidden_states)916 hidden_states = self.fc2(hidden_states)917 return hidden_states918 919 920# Copied from transformers.models.clip.modeling_clip.CLIPEncoderLayer with CLIP->Siglip921class SiglipEncoderLayer(nn.Module):922 def __init__(self, config: SiglipConfig):923 super().__init__()924 self.embed_dim = config.hidden_size925 self.self_attn = (926 SiglipAttention(config)927 if not getattr(config, "_flash_attn_2_enabled", False)928 else SiglipFlashAttention2(config)929 )930 self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)931 self.mlp = SiglipMLP(config)932 self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)933 934 def forward(935 self,936 hidden_states: torch.Tensor,937 attention_mask: torch.Tensor,938 output_attentions: Optional[bool] = False,939 ) -> Tuple[torch.FloatTensor]:940 """941 Args:942 hidden_states (`torch.FloatTensor`):943 Input to the layer of shape `(batch, seq_len, embed_dim)`.944 attention_mask (`torch.FloatTensor`):945 Attention mask of shape `(batch, 1, q_len, k_v_seq_len)` where padding elements are indicated by very large negative values.946 output_attentions (`bool`, *optional*, defaults to `False`):947 Whether or not to return the attentions tensors of all attention layers. See `attentions` under948 returned tensors for more detail.949 """950 residual = hidden_states951 952 hidden_states = self.layer_norm1(hidden_states)953 hidden_states, attn_weights = self.self_attn(954 hidden_states=hidden_states,955 attention_mask=attention_mask,956 output_attentions=output_attentions,957 )958 hidden_states = residual + hidden_states959 960 residual = hidden_states961 hidden_states = self.layer_norm2(hidden_states)962 hidden_states = self.mlp(hidden_states)963 hidden_states = residual + hidden_states964 965 outputs = (hidden_states,)966 967 if output_attentions:968 outputs += (attn_weights,)969 970 return outputs971 972 973class SiglipPreTrainedModel(PreTrainedModel):974 """975 An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained976 models.977 """978 979 config_class = SiglipConfig980 base_model_prefix = "siglip"981 supports_gradient_checkpointing = True982 983 def _init_weights(self, module):984 """Initialize the weights"""985 986 if isinstance(module, SiglipVisionEmbeddings):987 width = (988 self.config.vision_config.hidden_size989 if isinstance(self.config, SiglipConfig)990 else self.config.hidden_size991 )992 nn.init.normal_(module.position_embedding.weight, std=1 / np.sqrt(width))993 elif isinstance(module, nn.Embedding):994 default_flax_embed_init(module.weight)995 elif isinstance(module, SiglipAttention):996 nn.init.normal_(module.q_proj.weight)997 nn.init.normal_(module.k_proj.weight)998 nn.init.normal_(module.v_proj.weight)999 nn.init.normal_(module.out_proj.weight)1000 nn.init.zeros_(module.q_proj.bias)1001 nn.init.zeros_(module.k_proj.bias)1002 nn.init.zeros_(module.v_proj.bias)1003 nn.init.zeros_(module.out_proj.bias)1004 elif isinstance(module, SiglipMLP):1005 nn.init.normal_(module.fc1.weight)1006 nn.init.normal_(module.fc2.weight)1007 nn.init.normal_(module.fc1.bias, std=1e-6)1008 nn.init.normal_(module.fc2.bias, std=1e-6)1009 elif isinstance(module, SiglipMultiheadAttentionPoolingHead):1010 nn.init.normal_(module.probe.data)1011 nn.init.normal_(module.attention.in_proj_weight.data)1012 nn.init.zeros_(module.attention.in_proj_bias.data)1013 elif isinstance(module, SiglipModel):1014 logit_scale_init = torch.tensor(0.0)1015 module.logit_scale.data.fill_(logit_scale_init)1016 module.logit_bias.data.zero_()1017 elif isinstance(module, (nn.Linear, nn.Conv2d)):1018 lecun_normal_(module.weight)1019 if module.bias is not None:1020 nn.init.zeros_(module.bias)1021 elif isinstance(module, nn.LayerNorm):1022 module.bias.data.zero_()1023 module.weight.data.fill_(1.0)1024 1025 1026SIGLIP_START_DOCSTRING = r"""1027 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the1028 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads1029 etc.)1030 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.1031 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage1032 and behavior.1033 Parameters:1034 config ([`SiglipConfig`]): Model configuration class with all the parameters of the model.1035 Initializing with a config file does not load the weights associated with the model, only the1036 configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.1037"""1038 1039SIGLIP_TEXT_INPUTS_DOCSTRING = r"""1040 Args:1041 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):1042 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide1043 it.1044 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1045 [`PreTrainedTokenizer.__call__`] for details.1046 [What are input IDs?](../glossary#input-ids)1047 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1048 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1049 - 1 for tokens that are **not masked**,1050 - 0 for tokens that are **masked**.1051 [What are attention masks?](../glossary#attention-mask)1052 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1053 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,1054 config.max_position_embeddings - 1]`.1055 [What are position IDs?](../glossary#position-ids)1056 output_attentions (`bool`, *optional*):1057 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1058 tensors for more detail.1059 output_hidden_states (`bool`, *optional*):1060 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1061 more detail.1062 return_dict (`bool`, *optional*):1063 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1064"""1065 1066SIGLIP_VISION_INPUTS_DOCSTRING = r"""1067 Args:1068 pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):1069 Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using1070 [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details.1071 output_attentions (`bool`, *optional*):1072 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1073 tensors for more detail.1074 output_hidden_states (`bool`, *optional*):1075 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1076 more detail.1077 return_dict (`bool`, *optional*):1078 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1079"""1080 1081SIGLIP_INPUTS_DOCSTRING = r"""1082 Args:1083 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):1084 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide1085 it.1086 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1087 [`PreTrainedTokenizer.__call__`] for details.1088 [What are input IDs?](../glossary#input-ids)1089 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1090 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1091 - 1 for tokens that are **not masked**,1092 - 0 for tokens that are **masked**.1093 [What are attention masks?](../glossary#attention-mask)1094 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1095 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,1096 config.max_position_embeddings - 1]`.1097 [What are position IDs?](../glossary#position-ids)1098 pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):1099 Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using1100 [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details.1101 return_loss (`bool`, *optional*):1102 Whether or not to return the contrastive loss.1103 output_attentions (`bool`, *optional*):1104 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1105 tensors for more detail.1106 output_hidden_states (`bool`, *optional*):1107 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1108 more detail.1109 return_dict (`bool`, *optional*):1110 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1111"""1112 1113 1114# Copied from transformers.models.clip.modeling_clip.CLIPEncoder with CLIP->Siglip1115class SiglipEncoder(nn.Module):1116 """1117 Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a1118 [`SiglipEncoderLayer`].1119 Args:1120 config: SiglipConfig1121 """1122 1123 def __init__(self, config: SiglipConfig):1124 super().__init__()1125 self.config = config1126 self.layers = nn.ModuleList([SiglipEncoderLayer(config) for _ in range(config.num_hidden_layers)])1127 self.gradient_checkpointing = False1128 1129 # Ignore copy1130 def forward(1131 self,1132 inputs_embeds,1133 attention_mask: Optional[torch.Tensor] = None,1134 output_attentions: Optional[bool] = None,1135 output_hidden_states: Optional[bool] = None,1136 return_dict: Optional[bool] = None,1137 ) -> Union[Tuple, BaseModelOutput]:1138 r"""1139 Args:1140 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):1141 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.1142 This is useful if you want more control over how to convert `input_ids` indices into associated vectors1143 than the model's internal embedding lookup matrix.1144 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1145 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1146 - 1 for tokens that are **not masked**,1147 - 0 for tokens that are **masked**.1148 [What are attention masks?](../glossary#attention-mask)1149 output_attentions (`bool`, *optional*):1150 Whether or not to return the attentions tensors of all attention layers. See `attentions` under1151 returned tensors for more detail.1152 output_hidden_states (`bool`, *optional*):1153 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors1154 for more detail.1155 return_dict (`bool`, *optional*):1156 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1157 """1158 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions1159 output_hidden_states = (1160 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1161 )1162 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1163 1164 encoder_states = () if output_hidden_states else None1165 all_attentions = () if output_attentions else None1166 1167 hidden_states = inputs_embeds1168 for encoder_layer in self.layers:1169 if output_hidden_states:1170 encoder_states = encoder_states + (hidden_states,)1171 if self.gradient_checkpointing and self.training:1172 layer_outputs = self._gradient_checkpointing_func(1173 encoder_layer.__call__,1174 hidden_states,1175 attention_mask,1176 output_attentions,1177 )1178 else:1179 layer_outputs = encoder_layer(1180 hidden_states,1181 attention_mask,1182 output_attentions=output_attentions,1183 )1184 1185 hidden_states = layer_outputs[0]1186 1187 if output_attentions:1188 all_attentions = all_attentions + (layer_outputs[1],)1189 1190 if output_hidden_states:1191 encoder_states = encoder_states + (hidden_states,)1192 1193 if not return_dict:1194 return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)1195 return BaseModelOutput(1196 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions1197 )1198 1199 1200class SiglipTextTransformer(nn.Module):