Aluode/PerceptionLabPortable
0
1# coding=utf-82# Copyright 2022 The BAAI Teams Authors and The HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15"""PyTorch AltCLIP model."""16 17import math18from dataclasses import dataclass19from typing import Any, Callable, Optional, Union20 21import torch22import torch.nn as nn23 24from ...activations import ACT2FN25from ...modeling_layers import GradientCheckpointingLayer26from ...modeling_outputs import (27 BaseModelOutput,28 BaseModelOutputWithPooling,29 BaseModelOutputWithPoolingAndCrossAttentions,30 BaseModelOutputWithPoolingAndProjection,31)32from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel33from ...pytorch_utils import apply_chunking_to_forward, find_pruneable_heads_and_indices, prune_linear_layer34from ...utils import ModelOutput, auto_docstring, can_return_tuple, filter_out_non_signature_kwargs, logging, torch_int35from .configuration_altclip import AltCLIPConfig, AltCLIPTextConfig, AltCLIPVisionConfig36 37 38logger = logging.get_logger(__name__)39 40 41# contrastive loss function, adapted from42# https://sachinruk.github.io/blog/pytorch/pytorch%20lightning/loss%20function/gpu/2021/03/07/CLIP.html43def contrastive_loss(logits: torch.Tensor) -> torch.Tensor:44 return nn.functional.cross_entropy(logits, torch.arange(len(logits), device=logits.device))45 46 47def clip_loss(similarity: torch.Tensor) -> torch.Tensor:48 caption_loss = contrastive_loss(similarity)49 image_loss = contrastive_loss(similarity.t())50 return (caption_loss + image_loss) / 2.051 52 53@dataclass54@auto_docstring55# Copied from transformers.models.clip.modeling_clip.CLIPOutput with CLIP->AltCLIP56class AltCLIPOutput(ModelOutput):57 r"""58 loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):59 Contrastive loss for image-text similarity.60 logits_per_image (`torch.FloatTensor` of shape `(image_batch_size, text_batch_size)`):61 The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text62 similarity scores.63 logits_per_text (`torch.FloatTensor` of shape `(text_batch_size, image_batch_size)`):64 The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image65 similarity scores.66 text_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):67 The text embeddings obtained by applying the projection layer to the pooled output of [`AltCLIPTextModel`].68 image_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):69 The image embeddings obtained by applying the projection layer to the pooled output of [`AltCLIPVisionModel`].70 text_model_output (`BaseModelOutputWithPooling`):71 The output of the [`AltCLIPTextModel`].72 vision_model_output (`BaseModelOutputWithPooling`):73 The output of the [`AltCLIPVisionModel`].74 """75 76 loss: Optional[torch.FloatTensor] = None77 logits_per_image: Optional[torch.FloatTensor] = None78 logits_per_text: Optional[torch.FloatTensor] = None79 text_embeds: Optional[torch.FloatTensor] = None80 image_embeds: Optional[torch.FloatTensor] = None81 text_model_output: BaseModelOutputWithPooling = None82 vision_model_output: BaseModelOutputWithPooling = None83 84 def to_tuple(self) -> tuple[Any]:85 return tuple(86 self[k] if k not in ["text_model_output", "vision_model_output"] else getattr(self, k).to_tuple()87 for k in self.keys()88 )89 90 91# Copied from transformers.models.roberta.modeling_roberta.RobertaEmbeddings with Roberta->AltRoberta92class AltRobertaEmbeddings(nn.Module):93 """94 Same as BertEmbeddings with a tiny tweak for positional embeddings indexing.95 """96 97 # Copied from transformers.models.bert.modeling_bert.BertEmbeddings.__init__98 def __init__(self, config):99 super().__init__()100 self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)101 self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)102 self.token_type_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size)103 104 # self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load105 # any TensorFlow checkpoint file106 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)107 self.dropout = nn.Dropout(config.hidden_dropout_prob)108 # position_ids (1, len position emb) is contiguous in memory and exported when serialized109 self.position_embedding_type = getattr(config, "position_embedding_type", "absolute")110 self.register_buffer(111 "position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False112 )113 self.register_buffer(114 "token_type_ids", torch.zeros(self.position_ids.size(), dtype=torch.long), persistent=False115 )116 117 # End copy118 self.padding_idx = config.pad_token_id119 self.position_embeddings = nn.Embedding(120 config.max_position_embeddings, config.hidden_size, padding_idx=self.padding_idx121 )122 123 def forward(124 self, input_ids=None, token_type_ids=None, position_ids=None, inputs_embeds=None, past_key_values_length=0125 ):126 if position_ids is None:127 if input_ids is not None:128 # Create the position ids from the input token ids. Any padded tokens remain padded.129 position_ids = create_position_ids_from_input_ids(input_ids, self.padding_idx, past_key_values_length)130 else:131 position_ids = self.create_position_ids_from_inputs_embeds(inputs_embeds)132 133 if input_ids is not None:134 input_shape = input_ids.size()135 else:136 input_shape = inputs_embeds.size()[:-1]137 138 seq_length = input_shape[1]139 140 # Setting the token_type_ids to the registered buffer in constructor where it is all zeros, which usually occurs141 # when its auto-generated, registered buffer helps users when tracing the model without passing token_type_ids, solves142 # issue #5664143 if token_type_ids is None:144 if hasattr(self, "token_type_ids"):145 buffered_token_type_ids = self.token_type_ids[:, :seq_length]146 buffered_token_type_ids_expanded = buffered_token_type_ids.expand(input_shape[0], seq_length)147 token_type_ids = buffered_token_type_ids_expanded148 else:149 token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=self.position_ids.device)150 151 if inputs_embeds is None:152 inputs_embeds = self.word_embeddings(input_ids)153 token_type_embeddings = self.token_type_embeddings(token_type_ids)154 155 embeddings = inputs_embeds + token_type_embeddings156 if self.position_embedding_type == "absolute":157 position_embeddings = self.position_embeddings(position_ids)158 embeddings += position_embeddings159 embeddings = self.LayerNorm(embeddings)160 embeddings = self.dropout(embeddings)161 return embeddings162 163 def create_position_ids_from_inputs_embeds(self, inputs_embeds):164 """165 We are provided embeddings directly. We cannot infer which are padded so just generate sequential position ids.166 167 Args:168 inputs_embeds: torch.Tensor169 170 Returns: torch.Tensor171 """172 input_shape = inputs_embeds.size()[:-1]173 sequence_length = input_shape[1]174 175 position_ids = torch.arange(176 self.padding_idx + 1, sequence_length + self.padding_idx + 1, dtype=torch.long, device=inputs_embeds.device177 )178 return position_ids.unsqueeze(0).expand(input_shape)179 180 181class AltRobertaSelfAttention(nn.Module):182 def __init__(self, config, position_embedding_type=None):183 super().__init__()184 if config.hidden_size % config.num_attention_heads != 0 and not hasattr(config, "embedding_size"):185 raise ValueError(186 f"The hidden size ({config.hidden_size}) is not a multiple of the number of attention "187 f"heads ({config.num_attention_heads})"188 )189 190 self.num_attention_heads = config.num_attention_heads191 self.attention_head_size = int(config.hidden_size / config.num_attention_heads)192 self.all_head_size = self.num_attention_heads * self.attention_head_size193 194 self.query = nn.Linear(config.hidden_size, self.all_head_size)195 self.key = nn.Linear(config.hidden_size, self.all_head_size)196 self.value = nn.Linear(config.hidden_size, self.all_head_size)197 198 self.dropout = nn.Dropout(config.attention_probs_dropout_prob)199 self.position_embedding_type = position_embedding_type or getattr(200 config, "position_embedding_type", "absolute"201 )202 if self.position_embedding_type == "relative_key" or self.position_embedding_type == "relative_key_query":203 self.max_position_embeddings = config.max_position_embeddings204 self.distance_embedding = nn.Embedding(2 * config.max_position_embeddings - 1, self.attention_head_size)205 206 def forward(207 self,208 hidden_states: torch.Tensor,209 attention_mask: Optional[torch.FloatTensor] = None,210 head_mask: Optional[torch.FloatTensor] = None,211 output_attentions: Optional[bool] = False,212 ) -> tuple[torch.Tensor]:213 input_shape = hidden_states.shape[:-1]214 hidden_shape = (*input_shape, -1, self.attention_head_size)215 216 query_layer = self.query(hidden_states).view(hidden_shape).transpose(1, 2)217 key_layer = self.key(hidden_states).view(hidden_shape).transpose(1, 2)218 value_layer = self.value(hidden_states).view(hidden_shape).transpose(1, 2)219 220 # Take the dot product between "query" and "key" to get the raw attention scores.221 attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))222 223 if self.position_embedding_type == "relative_key" or self.position_embedding_type == "relative_key_query":224 query_length, key_length = query_layer.shape[2], key_layer.shape[2]225 position_ids_l = torch.arange(query_length, dtype=torch.long, device=hidden_states.device).view(-1, 1)226 position_ids_r = torch.arange(key_length, dtype=torch.long, device=hidden_states.device).view(1, -1)227 distance = position_ids_l - position_ids_r228 229 positional_embedding = self.distance_embedding(distance + self.max_position_embeddings - 1)230 positional_embedding = positional_embedding.to(dtype=query_layer.dtype) # fp16 compatibility231 232 if self.position_embedding_type == "relative_key":233 relative_position_scores = torch.einsum("bhld,lrd->bhlr", query_layer, positional_embedding)234 attention_scores = attention_scores + relative_position_scores235 elif self.position_embedding_type == "relative_key_query":236 relative_position_scores_query = torch.einsum("bhld,lrd->bhlr", query_layer, positional_embedding)237 relative_position_scores_key = torch.einsum("bhrd,lrd->bhlr", key_layer, positional_embedding)238 attention_scores = attention_scores + relative_position_scores_query + relative_position_scores_key239 240 attention_scores = attention_scores / math.sqrt(self.attention_head_size)241 if attention_mask is not None:242 # Apply the attention mask is (precomputed for all layers in AltRobertaModel forward() function)243 attention_scores = attention_scores + attention_mask244 245 # Normalize the attention scores to probabilities.246 attention_probs = nn.functional.softmax(attention_scores, dim=-1)247 248 # This is actually dropping out entire tokens to attend to, which might249 # seem a bit unusual, but is taken from the original Transformer paper.250 attention_probs = self.dropout(attention_probs)251 252 # Mask heads if we want to253 if head_mask is not None:254 attention_probs = attention_probs * head_mask255 256 context_layer = torch.matmul(attention_probs, value_layer)257 258 context_layer = context_layer.permute(0, 2, 1, 3).contiguous()259 new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)260 context_layer = context_layer.view(new_context_layer_shape)261 262 outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)263 264 return outputs265 266 267# Copied from transformers.models.roberta.modeling_roberta.RobertaSelfOutput268class AltRobertaSelfOutput(nn.Module):269 def __init__(self, config):270 super().__init__()271 self.dense = nn.Linear(config.hidden_size, config.hidden_size)272 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)273 self.dropout = nn.Dropout(config.hidden_dropout_prob)274 275 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:276 hidden_states = self.dense(hidden_states)277 hidden_states = self.dropout(hidden_states)278 hidden_states = self.LayerNorm(hidden_states + input_tensor)279 return hidden_states280 281 282ALT_ROBERTA_SELF_ATTENTION_CLASSES = {283 "eager": AltRobertaSelfAttention,284}285 286 287class AltRobertaAttention(nn.Module):288 def __init__(self, config, position_embedding_type=None):289 super().__init__()290 self.self = ALT_ROBERTA_SELF_ATTENTION_CLASSES[config._attn_implementation](291 config, position_embedding_type=position_embedding_type292 )293 self.output = AltRobertaSelfOutput(config)294 self.pruned_heads = set()295 296 def prune_heads(self, heads):297 if len(heads) == 0:298 return299 heads, index = find_pruneable_heads_and_indices(300 heads, self.self.num_attention_heads, self.self.attention_head_size, self.pruned_heads301 )302 303 # Prune linear layers304 self.self.query = prune_linear_layer(self.self.query, index)305 self.self.key = prune_linear_layer(self.self.key, index)306 self.self.value = prune_linear_layer(self.self.value, index)307 self.output.dense = prune_linear_layer(self.output.dense, index, dim=1)308 309 # Update hyper params and store pruned heads310 self.self.num_attention_heads = self.self.num_attention_heads - len(heads)311 self.self.all_head_size = self.self.attention_head_size * self.self.num_attention_heads312 self.pruned_heads = self.pruned_heads.union(heads)313 314 def forward(315 self,316 hidden_states: torch.Tensor,317 attention_mask: Optional[torch.FloatTensor] = None,318 head_mask: Optional[torch.FloatTensor] = None,319 output_attentions: Optional[bool] = False,320 ) -> tuple[torch.Tensor]:321 self_outputs = self.self(322 hidden_states,323 attention_mask=attention_mask,324 head_mask=head_mask,325 output_attentions=output_attentions,326 )327 attention_output = self.output(self_outputs[0], hidden_states)328 outputs = (attention_output,) + self_outputs[1:] # add attentions if we output them329 return outputs330 331 332# Copied from transformers.models.roberta.modeling_roberta.RobertaIntermediate with Roberta->AltRoberta333class AltRobertaIntermediate(nn.Module):334 def __init__(self, config):335 super().__init__()336 self.dense = nn.Linear(config.hidden_size, config.intermediate_size)337 if isinstance(config.hidden_act, str):338 self.intermediate_act_fn = ACT2FN[config.hidden_act]339 else:340 self.intermediate_act_fn = config.hidden_act341 342 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:343 hidden_states = self.dense(hidden_states)344 hidden_states = self.intermediate_act_fn(hidden_states)345 return hidden_states346 347 348# Copied from transformers.models.roberta.modeling_roberta.RobertaOutput349class AltRobertaOutput(nn.Module):350 def __init__(self, config):351 super().__init__()352 self.dense = nn.Linear(config.intermediate_size, config.hidden_size)353 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)354 self.dropout = nn.Dropout(config.hidden_dropout_prob)355 356 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:357 hidden_states = self.dense(hidden_states)358 hidden_states = self.dropout(hidden_states)359 hidden_states = self.LayerNorm(hidden_states + input_tensor)360 return hidden_states361 362 363# Copied from transformers.models.align.modeling_align.AlignTextLayer with AlignText->AltRoberta364class AltRobertaLayer(GradientCheckpointingLayer):365 def __init__(self, config):366 super().__init__()367 self.chunk_size_feed_forward = config.chunk_size_feed_forward368 self.seq_len_dim = 1369 self.attention = AltRobertaAttention(config)370 self.intermediate = AltRobertaIntermediate(config)371 self.output = AltRobertaOutput(config)372 373 def forward(374 self,375 hidden_states: torch.Tensor,376 attention_mask: Optional[torch.FloatTensor] = None,377 head_mask: Optional[torch.FloatTensor] = None,378 output_attentions: Optional[bool] = False,379 **kwargs,380 ) -> tuple[torch.Tensor]:381 self_attention_outputs = self.attention(382 hidden_states,383 attention_mask=attention_mask,384 head_mask=head_mask,385 output_attentions=output_attentions,386 **kwargs,387 )388 attention_output = self_attention_outputs[0]389 390 outputs = self_attention_outputs[1:] # add self attentions if we output attention weights391 layer_output = apply_chunking_to_forward(392 self.feed_forward_chunk, self.chunk_size_feed_forward, self.seq_len_dim, attention_output393 )394 outputs = (layer_output,) + outputs395 396 return outputs397 398 def feed_forward_chunk(self, attention_output):399 intermediate_output = self.intermediate(attention_output)400 layer_output = self.output(intermediate_output, attention_output)401 return layer_output402 403 404# Copied from transformers.models.align.modeling_align.AlignTextEncoder with AlignText->AltRoberta405class AltRobertaEncoder(nn.Module):406 def __init__(self, config):407 super().__init__()408 self.config = config409 self.layer = nn.ModuleList([AltRobertaLayer(config) for i in range(config.num_hidden_layers)])410 self.gradient_checkpointing = False411 412 @can_return_tuple413 def forward(414 self,415 hidden_states: torch.Tensor,416 attention_mask: Optional[torch.FloatTensor] = None,417 head_mask: Optional[torch.FloatTensor] = None,418 output_attentions: Optional[bool] = False,419 output_hidden_states: Optional[bool] = False,420 return_dict: Optional[bool] = True,421 **kwargs,422 ) -> Union[tuple[torch.Tensor], BaseModelOutput]:423 all_hidden_states = () if output_hidden_states else None424 all_self_attentions = () if output_attentions else None425 426 for i, layer_module in enumerate(self.layer):427 if output_hidden_states:428 all_hidden_states = all_hidden_states + (hidden_states,)429 430 layer_head_mask = head_mask[i] if head_mask is not None else None431 432 layer_outputs = layer_module(433 hidden_states=hidden_states,434 attention_mask=attention_mask,435 head_mask=layer_head_mask,436 output_attentions=output_attentions,437 **kwargs,438 )439 440 hidden_states = layer_outputs[0]441 if output_attentions:442 all_self_attentions = all_self_attentions + (layer_outputs[1],)443 444 if output_hidden_states:445 all_hidden_states = all_hidden_states + (hidden_states,)446 447 return BaseModelOutput(448 last_hidden_state=hidden_states,449 hidden_states=all_hidden_states,450 attentions=all_self_attentions,451 )452 453 454# Copied from transformers.models.roberta.modeling_roberta.RobertaPooler455class AltRobertaPooler(nn.Module):456 def __init__(self, config):457 super().__init__()458 self.dense = nn.Linear(config.hidden_size, config.hidden_size)459 self.activation = nn.Tanh()460 461 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:462 # We "pool" the model by simply taking the hidden state corresponding463 # to the first token.464 first_token_tensor = hidden_states[:, 0]465 pooled_output = self.dense(first_token_tensor)466 pooled_output = self.activation(pooled_output)467 return pooled_output468 469 470# Copied from transformers.models.siglip.modeling_siglip.eager_attention_forward471def eager_attention_forward(472 module: nn.Module,473 query: torch.Tensor,474 key: torch.Tensor,475 value: torch.Tensor,476 attention_mask: Optional[torch.Tensor],477 scaling: float,478 dropout: float = 0.0,479 **kwargs,480):481 attn_weights = torch.matmul(query, key.transpose(-1, -2)) * scaling482 if attention_mask is not None:483 attn_weights = attn_weights + attention_mask484 485 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)486 attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)487 488 attn_output = torch.matmul(attn_weights, value)489 attn_output = attn_output.transpose(1, 2).contiguous()490 491 return attn_output, attn_weights492 493 494class AltCLIPAttention(nn.Module):495 """Multi-headed attention from 'Attention Is All You Need' paper"""496 497 def __init__(self, config):498 super().__init__()499 self.config = config500 self.embed_dim = config.hidden_size501 self.num_heads = config.num_attention_heads502 self.head_dim = self.embed_dim // self.num_heads503 if self.head_dim * self.num_heads != self.embed_dim:504 raise ValueError(505 f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"506 f" {self.num_heads})."507 )508 self.scale = self.head_dim**-0.5509 self.dropout = config.attention_dropout510 self.is_causal = False511 512 self.k_proj = nn.Linear(self.embed_dim, self.embed_dim)513 self.v_proj = nn.Linear(self.embed_dim, self.embed_dim)514 self.q_proj = nn.Linear(self.embed_dim, self.embed_dim)515 self.out_proj = nn.Linear(self.embed_dim, self.embed_dim)516 517 def forward(518 self,519 hidden_states: torch.Tensor,520 attention_mask: Optional[torch.Tensor] = None,521 causal_attention_mask: Optional[torch.Tensor] = None,522 output_attentions: Optional[bool] = False,523 ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:524 """Input shape: Batch x Time x Channel"""525 526 batch_size, seq_length, embed_dim = hidden_states.shape527 528 queries = self.q_proj(hidden_states)529 keys = self.k_proj(hidden_states)530 values = self.v_proj(hidden_states)531 532 queries = queries.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)533 keys = keys.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)534 values = values.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)535 # CLIP text model uses both `causal_attention_mask` and `attention_mask`536 # in case FA2 kernel is called, `is_causal` should be inferred from `causal_attention_mask`537 if self.config._attn_implementation != "flash_attention_2":538 if attention_mask is not None and causal_attention_mask is not None:539 attention_mask = attention_mask + causal_attention_mask540 elif causal_attention_mask is not None:541 attention_mask = causal_attention_mask542 else:543 self.is_causal = causal_attention_mask is not None544 545 attention_interface: Callable = eager_attention_forward546 if self.config._attn_implementation != "eager":547 attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]548 549 attn_output, attn_weights = attention_interface(550 self,551 queries,552 keys,553 values,554 attention_mask,555 is_causal=self.is_causal,556 scaling=self.scale,557 dropout=0.0 if not self.training else self.dropout,558 )559 560 attn_output = attn_output.reshape(batch_size, seq_length, embed_dim).contiguous()561 attn_output = self.out_proj(attn_output)562 if not output_attentions:563 attn_weights = None564 return attn_output, attn_weights565 566 567# Copied from transformers.models.clip.modeling_clip.CLIPMLP with CLIP->AltCLIP568class AltCLIPMLP(nn.Module):569 def __init__(self, config):570 super().__init__()571 self.config = config572 self.activation_fn = ACT2FN[config.hidden_act]573 self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)574 self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)575 576 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:577 hidden_states = self.fc1(hidden_states)578 hidden_states = self.activation_fn(hidden_states)579 hidden_states = self.fc2(hidden_states)580 return hidden_states581 582 583class AltCLIPEncoderLayer(GradientCheckpointingLayer):584 def __init__(self, config: AltCLIPConfig):585 super().__init__()586 self.embed_dim = config.hidden_size587 self.self_attn = AltCLIPAttention(config)588 self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)589 self.mlp = AltCLIPMLP(config)590 self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)591 592 def forward(593 self,594 hidden_states: torch.Tensor,595 attention_mask: torch.Tensor,596 causal_attention_mask: torch.Tensor,597 output_attentions: Optional[bool] = False,598 ) -> tuple[torch.FloatTensor]:599 """600 Args:601 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`602 attention_mask (`torch.FloatTensor`): attention mask of size603 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.604 `(config.encoder_attention_heads,)`.605 output_attentions (`bool`, *optional*):606 Whether or not to return the attentions tensors of all attention layers. See `attentions` under607 returned tensors for more detail.608 """609 residual = hidden_states610 611 hidden_states = self.layer_norm1(hidden_states)612 hidden_states, attn_weights = self.self_attn(613 hidden_states=hidden_states,614 attention_mask=attention_mask,615 causal_attention_mask=causal_attention_mask,616 output_attentions=output_attentions,617 )618 hidden_states = residual + hidden_states619 620 residual = hidden_states621 hidden_states = self.layer_norm2(hidden_states)622 hidden_states = self.mlp(hidden_states)623 hidden_states = residual + hidden_states624 625 outputs = (hidden_states,)626 627 if output_attentions:628 outputs += (attn_weights,)629 630 return outputs631 632 633class AltCLIPEncoder(nn.Module):634 """635 Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a636 [`AltCLIPEncoderLayer`].637 638 Args:639 config: AltCLIPConfig640 """641 642 def __init__(self, config: AltCLIPConfig):643 super().__init__()644 self.config = config645 self.layers = nn.ModuleList([AltCLIPEncoderLayer(config) for _ in range(config.num_hidden_layers)])646 self.gradient_checkpointing = False647 648 @can_return_tuple649 def forward(650 self,651 inputs_embeds,652 attention_mask: Optional[torch.Tensor] = None,653 causal_attention_mask: Optional[torch.Tensor] = None,654 output_attentions: Optional[bool] = None,655 output_hidden_states: Optional[bool] = None,656 return_dict: Optional[bool] = None,657 ) -> Union[tuple, BaseModelOutput]:658 r"""659 Args:660 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):661 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.662 This is useful if you want more control over how to convert `input_ids` indices into associated vectors663 than the model's internal embedding lookup matrix.664 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):665 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:666 667 - 1 for tokens that are **not masked**,668 - 0 for tokens that are **masked**.669 670 [What are attention masks?](../glossary#attention-mask)671 causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):672 Causal mask for the text model. Mask values selected in `[0, 1]`:673 674 - 1 for tokens that are **not masked**,675 - 0 for tokens that are **masked**.676 677 [What are attention masks?](../glossary#attention-mask)678 output_attentions (`bool`, *optional*):679 Whether or not to return the attentions tensors of all attention layers. See `attentions` under680 returned tensors for more detail.681 output_hidden_states (`bool`, *optional*):682 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors683 for more detail.684 return_dict (`bool`, *optional*):685 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.686 """687 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions688 output_hidden_states = (689 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states690 )691 return_dict = return_dict if return_dict is not None else self.config.use_return_dict692 693 encoder_states = () if output_hidden_states else None694 all_attentions = () if output_attentions else None695 696 hidden_states = inputs_embeds697 for idx, encoder_layer in enumerate(self.layers):698 if output_hidden_states:699 encoder_states = encoder_states + (hidden_states,)700 layer_outputs = encoder_layer(701 hidden_states,702 attention_mask,703 causal_attention_mask,704 output_attentions=output_attentions,705 )706 707 hidden_states = layer_outputs[0]708 709 if output_attentions:710 all_attentions = all_attentions + (layer_outputs[1],)711 712 if output_hidden_states:713 encoder_states = encoder_states + (hidden_states,)714 715 return BaseModelOutput(716 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions717 )718 719 720# Copied from transformers.models.clip.modeling_clip.CLIPVisionEmbeddings with CLIP->AltCLIP721class AltCLIPVisionEmbeddings(nn.Module):722 def __init__(self, config: AltCLIPVisionConfig):723 super().__init__()724 self.config = config725 self.embed_dim = config.hidden_size726 self.image_size = config.image_size727 self.patch_size = config.patch_size728 729 self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))730 731 self.patch_embedding = nn.Conv2d(732 in_channels=config.num_channels,733 out_channels=self.embed_dim,734 kernel_size=self.patch_size,735 stride=self.patch_size,736 bias=False,737 )738 739 self.num_patches = (self.image_size // self.patch_size) ** 2740 self.num_positions = self.num_patches + 1741 self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)742 self.register_buffer("position_ids", torch.arange(self.num_positions).expand((1, -1)), persistent=False)743 744 def interpolate_pos_encoding(self, embeddings: torch.Tensor, height: int, width: int) -> torch.Tensor:745 """746 This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution747 images. This method is also adapted to support torch.jit tracing.748 749 Adapted from:750 - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and751 - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211752 """753 754 num_patches = embeddings.shape[1] - 1755 position_embedding = self.position_embedding.weight.unsqueeze(0)756 num_positions = position_embedding.shape[1] - 1757 758 # always interpolate when tracing to ensure the exported model works for dynamic input shapes759 if not torch.jit.is_tracing() and num_patches == num_positions and height == width:760 return self.position_embedding(self.position_ids)761 762 class_pos_embed = position_embedding[:, :1]763 patch_pos_embed = position_embedding[:, 1:]764 765 dim = embeddings.shape[-1]766 767 new_height = height // self.patch_size768 new_width = width // self.patch_size769 770 sqrt_num_positions = torch_int(num_positions**0.5)771 patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)772 patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)773 774 patch_pos_embed = nn.functional.interpolate(775 patch_pos_embed,776 size=(new_height, new_width),777 mode="bicubic",778 align_corners=False,779 )780 781 patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)782 783 return torch.cat((class_pos_embed, patch_pos_embed), dim=1)784 785 def forward(self, pixel_values: torch.FloatTensor, interpolate_pos_encoding=False) -> torch.Tensor:786 batch_size, _, height, width = pixel_values.shape787 if not interpolate_pos_encoding and (height != self.image_size or width != self.image_size):788 raise ValueError(789 f"Input image size ({height}*{width}) doesn't match model ({self.image_size}*{self.image_size})."790 )791 target_dtype = self.patch_embedding.weight.dtype792 patch_embeds = self.patch_embedding(pixel_values.to(dtype=target_dtype)) # shape = [*, width, grid, grid]793 patch_embeds = patch_embeds.flatten(2).transpose(1, 2)794 795 class_embeds = self.class_embedding.expand(batch_size, 1, -1)796 embeddings = torch.cat([class_embeds, patch_embeds], dim=1)797 if interpolate_pos_encoding:798 embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)799 else:800 embeddings = embeddings + self.position_embedding(self.position_ids)801 return embeddings802 803 804@auto_docstring805class AltCLIPPreTrainedModel(PreTrainedModel):806 config: AltCLIPConfig807 base_model_prefix = "altclip"808 supports_gradient_checkpointing = True809 _no_split_module = []810 811 def _init_weights(self, module):812 """Initialize the weights"""813 factor = self.config.initializer_factor814 if isinstance(module, AltCLIPVisionEmbeddings):815 factor = self.config.initializer_factor816 nn.init.normal_(module.class_embedding, mean=0.0, std=module.embed_dim**-0.5 * factor)817 nn.init.normal_(module.patch_embedding.weight, std=module.config.initializer_range * factor)818 nn.init.normal_(module.position_embedding.weight, std=module.config.initializer_range * factor)819 elif isinstance(module, AltCLIPAttention):820 factor = self.config.initializer_factor821 in_proj_std = (module.embed_dim**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor822 out_proj_std = (module.embed_dim**-0.5) * factor823 nn.init.normal_(module.q_proj.weight, std=in_proj_std)824 nn.init.normal_(module.k_proj.weight, std=in_proj_std)825 nn.init.normal_(module.v_proj.weight, std=in_proj_std)826 nn.init.normal_(module.out_proj.weight, std=out_proj_std)827 elif isinstance(module, AltCLIPMLP):828 factor = self.config.initializer_factor829 in_proj_std = (module.config.hidden_size**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor830 fc_std = (2 * module.config.hidden_size) ** -0.5 * factor831 nn.init.normal_(module.fc1.weight, std=fc_std)832 nn.init.normal_(module.fc2.weight, std=in_proj_std)833 elif isinstance(module, AltCLIPModel):834 nn.init.normal_(835 module.text_projection.weight,836 std=module.text_embed_dim**-0.5 * self.config.initializer_factor,837 )838 module.text_projection._is_hf_initialized = True839 nn.init.normal_(840 module.visual_projection.weight,841 std=module.vision_embed_dim**-0.5 * self.config.initializer_factor,842 )843 module.visual_projection._is_hf_initialized = True844 elif isinstance(module, nn.LayerNorm):845 module.bias.data.zero_()846 module.weight.data.fill_(1.0)847 elif isinstance(module, nn.Linear):848 module.weight.data.normal_(mean=0.0, std=self.config.initializer_factor)849 if module.bias is not None:850 module.bias.data.zero_()851 elif isinstance(module, nn.Embedding):852 module.weight.data.normal_(mean=0.0, std=self.config.initializer_factor)853 if module.padding_idx is not None:854 module.weight.data[module.padding_idx].zero_()855 856 857class AltCLIPVisionTransformer(nn.Module):858 def __init__(self, config: AltCLIPVisionConfig):859 super().__init__()860 self.config = config861 embed_dim = config.hidden_size862 863 self.embeddings = AltCLIPVisionEmbeddings(config)864 self.pre_layrnorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)865 self.encoder = AltCLIPEncoder(config)866 self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)867 868 @can_return_tuple869 @auto_docstring870 def forward(871 self,872 pixel_values: Optional[torch.FloatTensor] = None,873 output_attentions: Optional[bool] = None,874 output_hidden_states: Optional[bool] = None,875 return_dict: Optional[bool] = None,876 interpolate_pos_encoding: Optional[bool] = False,877 ) -> Union[tuple, BaseModelOutputWithPooling]:878 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions879 output_hidden_states = (880 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states881 )882 return_dict = return_dict if return_dict is not None else self.config.use_return_dict883 884 if pixel_values is None:885 raise ValueError("You have to specify pixel_values")886 887 hidden_states = self.embeddings(pixel_values, interpolate_pos_encoding=interpolate_pos_encoding)888 hidden_states = self.pre_layrnorm(hidden_states)889 890 encoder_outputs = self.encoder(891 inputs_embeds=hidden_states,892 output_attentions=output_attentions,893 output_hidden_states=output_hidden_states,894 return_dict=True,895 )896 897 last_hidden_state = encoder_outputs[0]898 pooled_output = last_hidden_state[:, 0, :]899 pooled_output = self.post_layernorm(pooled_output)900 901 return BaseModelOutputWithPooling(902 last_hidden_state=last_hidden_state,903 pooler_output=pooled_output,904 hidden_states=encoder_outputs.hidden_states,905 attentions=encoder_outputs.attentions,906 )907 908 909class AltCLIPVisionModel(AltCLIPPreTrainedModel):910 config: AltCLIPVisionConfig911 main_input_name = "pixel_values"912 913 def __init__(self, config: AltCLIPVisionConfig):914 super().__init__(config)915 self.vision_model = AltCLIPVisionTransformer(config)916 # Initialize weights and apply final processing917 self.post_init()918 919 def get_input_embeddings(self) -> nn.Module:920 return self.vision_model.embeddings.patch_embedding921 922 @auto_docstring923 def forward(924 self,925 pixel_values: Optional[torch.FloatTensor] = None,926 output_attentions: Optional[bool] = None,927 output_hidden_states: Optional[bool] = None,928 interpolate_pos_encoding: bool = False,929 return_dict: Optional[bool] = None,930 ) -> Union[tuple, BaseModelOutputWithPooling]:931 r"""932 Examples:933 934 ```python935 >>> from PIL import Image936 >>> import requests937 >>> from transformers import AutoProcessor, AltCLIPVisionModel938 939 >>> model = AltCLIPVisionModel.from_pretrained("BAAI/AltCLIP")940 >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")941 942 >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"943 >>> image = Image.open(requests.get(url, stream=True).raw)944 945 >>> inputs = processor(images=image, return_tensors="pt")946 947 >>> outputs = model(**inputs)948 >>> last_hidden_state = outputs.last_hidden_state949 >>> pooled_output = outputs.pooler_output # pooled CLS states950 ```"""951 return_dict = return_dict if return_dict is not None else self.config.use_return_dict952 953 return self.vision_model(954 pixel_values=pixel_values,955 output_attentions=output_attentions,956 output_hidden_states=output_hidden_states,957 interpolate_pos_encoding=interpolate_pos_encoding,958 return_dict=return_dict,959 )960 961 962@auto_docstring(963 custom_intro="""964 The model behaves as an encoder following the architecture described in *Attention is965 all you need*_ by Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit, Llion Jones, Aidan N. Gomez, Lukasz966 Kaiser and Illia Polosukhin.967 968 .. _*Attention is all you need*: https://huggingface.co/papers/1706.03762969 """970)971class AltRobertaModel(AltCLIPPreTrainedModel):972 config: AltCLIPTextConfig973 974 # Copied from transformers.models.clap.modeling_clap.ClapTextModel.__init__ with ClapText->AltRoberta975 def __init__(self, config, add_pooling_layer=True):976 r"""977 add_pooling_layer (bool, *optional*, defaults to `True`):978 Whether to add a pooling layer979 """980 super().__init__(config)981 self.config = config982 983 self.embeddings = AltRobertaEmbeddings(config)984 self.encoder = AltRobertaEncoder(config)985 986 self.pooler = AltRobertaPooler(config) if add_pooling_layer else None987 988 # Initialize weights and apply final processing989 self.post_init()990 991 def get_input_embeddings(self):992 return self.embeddings.word_embeddings993 994 def set_input_embeddings(self, value):995 self.embeddings.word_embeddings = value996 997 def _prune_heads(self, heads_to_prune):998 """999 Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base1000 class PreTrainedModel1001 """1002 for layer, heads in heads_to_prune.items():1003 self.encoder.layer[layer].attention.prune_heads(heads)1004 1005 @auto_docstring1006 # Copied from transformers.models.clap.modeling_clap.ClapTextModel.forward1007 def forward(1008 self,1009 input_ids: Optional[torch.Tensor] = None,1010 attention_mask: Optional[torch.Tensor] = None,1011 token_type_ids: Optional[torch.Tensor] = None,1012 position_ids: Optional[torch.Tensor] = None,1013 head_mask: Optional[torch.Tensor] = None,1014 inputs_embeds: Optional[torch.Tensor] = None,1015 output_attentions: Optional[bool] = None,1016 output_hidden_states: Optional[bool] = None,1017 return_dict: Optional[bool] = None,1018 ) -> Union[tuple[torch.Tensor], BaseModelOutputWithPoolingAndCrossAttentions]:1019 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions1020 output_hidden_states = (1021 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1022 )1023 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1024 1025 if input_ids is not None and inputs_embeds is not None:1026 raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")1027 elif input_ids is not None:1028 self.warn_if_padding_and_no_attention_mask(input_ids, attention_mask)1029 input_shape = input_ids.size()1030 elif inputs_embeds is not None:1031 input_shape = inputs_embeds.size()[:-1]1032 else:1033 raise ValueError("You have to specify either input_ids or inputs_embeds")1034 1035 batch_size, seq_length = input_shape1036 device = input_ids.device if input_ids is not None else inputs_embeds.device1037 1038 if attention_mask is None:1039 attention_mask = torch.ones(((batch_size, seq_length)), device=device)1040 1041 if token_type_ids is None:1042 if hasattr(self.embeddings, "token_type_ids"):1043 buffered_token_type_ids = self.embeddings.token_type_ids[:, :seq_length]1044 buffered_token_type_ids_expanded = buffered_token_type_ids.expand(batch_size, seq_length)1045 token_type_ids = buffered_token_type_ids_expanded1046 else:1047 token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)1048 1049 # We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]1050 # ourselves in which case we just need to make it broadcastable to all heads.1051 extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(attention_mask, input_shape)1052 1053 # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]1054 head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)1055 1056 embedding_output = self.embeddings(1057 input_ids=input_ids,1058 position_ids=position_ids,1059 token_type_ids=token_type_ids,1060 inputs_embeds=inputs_embeds,1061 )1062 encoder_outputs = self.encoder(1063 embedding_output,1064 attention_mask=extended_attention_mask,1065 head_mask=head_mask,1066 output_attentions=output_attentions,1067 output_hidden_states=output_hidden_states,1068 return_dict=True,1069 )1070 sequence_output = encoder_outputs[0]1071 pooled_output = self.pooler(sequence_output) if self.pooler is not None else None1072 1073 return BaseModelOutputWithPooling(1074 last_hidden_state=sequence_output,1075 pooler_output=pooled_output,1076 hidden_states=encoder_outputs.hidden_states,1077 attentions=encoder_outputs.attentions,1078 )1079 1080 1081class AltCLIPTextModel(AltCLIPPreTrainedModel):1082 config: AltCLIPTextConfig1083 1084 def __init__(self, config):1085 super().__init__(config)1086 self.roberta = AltRobertaModel(config, add_pooling_layer=False)1087 self.transformation = nn.Linear(config.hidden_size, config.project_dim)1088 self.pre_LN = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)1089 self.post_init()1090 1091 def get_input_embeddings(self) -> nn.Module:1092 return self.roberta.embeddings.word_embeddings1093 1094 def set_input_embeddings(self, value: nn.Embedding) -> None:1095 self.roberta.embeddings.word_embeddings = value1096 1097 def resize_token_embeddings(self, new_num_tokens: Optional[int] = None) -> nn.Embedding:1098 return super().resize_token_embeddings(new_num_tokens)1099 1100 @can_return_tuple1101 @auto_docstring1102 def forward(1103 self,1104 input_ids: Optional[torch.Tensor] = None,1105 attention_mask: Optional[torch.Tensor] = None,1106 token_type_ids: Optional[torch.Tensor] = None,1107 position_ids: Optional[torch.Tensor] = None,1108 head_mask: Optional[torch.Tensor] = None,1109 inputs_embeds: Optional[torch.Tensor] = None,1110 output_attentions: Optional[bool] = None,1111 return_dict: Optional[bool] = None,1112 output_hidden_states: Optional[bool] = None,1113 ) -> Union[tuple, BaseModelOutputWithPoolingAndProjection]:1114 r"""1115 Examples:1116 1117 ```python1118 >>> from transformers import AutoProcessor, AltCLIPTextModel1119 1120 >>> model = AltCLIPTextModel.from_pretrained("BAAI/AltCLIP")1121 >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")1122 1123 >>> texts = ["it's a cat", "it's a dog"]1124 1125 >>> inputs = processor(text=texts, padding=True, return_tensors="pt")1126 1127 >>> outputs = model(**inputs)1128 >>> last_hidden_state = outputs.last_hidden_state1129 >>> pooled_output = outputs.pooler_output # pooled CLS states1130 ```"""1131 1132 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1133 1134 outputs = self.roberta(1135 input_ids=input_ids,1136 attention_mask=attention_mask,1137 token_type_ids=token_type_ids,1138 position_ids=position_ids,1139 head_mask=head_mask,1140 inputs_embeds=inputs_embeds,1141 output_attentions=output_attentions,1142 output_hidden_states=output_hidden_states,1143 return_dict=True,1144 )1145 1146 # last module outputs1147 sequence_output = outputs[0]1148 1149 # project every module1150 sequence_output = self.pre_LN(sequence_output)1151 1152 # pooler1153 projection_state = self.transformation(sequence_output)1154 pooler_output = projection_state[:, 0]1155 1156 return BaseModelOutputWithPoolingAndProjection(1157 last_hidden_state=projection_state,1158 pooler_output=pooler_output,1159 hidden_states=outputs.hidden_states,1160 attentions=outputs.attentions,1161 )1162 1163 1164class AltCLIPModel(AltCLIPPreTrainedModel):1165 config: AltCLIPConfig1166 1167 def __init__(self, config: AltCLIPConfig):1168 super().__init__(config)1169 1170 if not isinstance(config.vision_config, AltCLIPVisionConfig):1171 raise TypeError(1172 "config.vision_config is expected to be of type AltCLIPVisionConfig but is of type"1173 f" {type(config.vision_config)}."1174 )1175 if not isinstance(config.text_config, AltCLIPTextConfig):1176 raise TypeError(1177 "config.text_config is expected to be of type AltCLIPTextConfig but is of type"1178 f" {type(config.text_config)}."1179 )1180 1181 text_config = config.text_config1182 vision_config = config.vision_config1183 # The module using it is not a PreTrainedModel subclass so we need this1184 vision_config._attn_implementation = config._attn_implementation1185 1186 self.projection_dim = config.projection_dim1187 self.text_embed_dim = text_config.project_dim1188 self.vision_embed_dim = vision_config.hidden_size1189 1190 self.text_model = AltCLIPTextModel(text_config)1191 self.vision_model = AltCLIPVisionTransformer(vision_config)1192 1193 self.visual_projection = nn.Linear(self.vision_embed_dim, self.projection_dim, bias=False)1194 self.text_projection = nn.Linear(self.text_embed_dim, self.projection_dim, bias=False)1195 self.logit_scale = nn.Parameter(torch.tensor(self.config.logit_scale_init_value))1196 1197 # Initialize weights and apply final processing1198 self.post_init()1199 1200 @filter_out_non_signature_kwargs()