Aluode/PerceptionLabPortable
0
1# coding=utf-82# Copyright 2022 Microsoft Research and The HuggingFace 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 X-CLIP model."""16 17import copy18from dataclasses import dataclass19from typing import Any, Callable, Optional, Union20 21import torch22from torch import nn23 24from ...activations import ACT2FN25from ...modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask26from ...modeling_layers import GradientCheckpointingLayer27from ...modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling28from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel29from ...utils import (30 ModelOutput,31 auto_docstring,32 can_return_tuple,33 filter_out_non_signature_kwargs,34 logging,35 torch_int,36)37from .configuration_x_clip import XCLIPConfig, XCLIPTextConfig, XCLIPVisionConfig38 39 40logger = logging.get_logger(__name__)41 42 43# contrastive loss function, adapted from44# https://sachinruk.github.io/blog/pytorch/pytorch%20lightning/loss%20function/gpu/2021/03/07/CLIP.html45def contrastive_loss(logits: torch.Tensor) -> torch.Tensor:46 return nn.functional.cross_entropy(logits, torch.arange(len(logits), device=logits.device))47 48 49# Copied from transformers.models.clip.modeling_clip.clip_loss with clip->x_clip50def x_clip_loss(similarity: torch.Tensor) -> torch.Tensor:51 caption_loss = contrastive_loss(similarity)52 image_loss = contrastive_loss(similarity.t())53 return (caption_loss + image_loss) / 2.054 55 56@dataclass57@auto_docstring58class XCLIPOutput(ModelOutput):59 r"""60 loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):61 Contrastive loss for video-text similarity.62 logits_per_video (`torch.FloatTensor` of shape `(video_batch_size, text_batch_size)`):63 The scaled dot product scores between `video_embeds` and `text_embeds`. This represents the video-text64 similarity scores.65 logits_per_text (`torch.FloatTensor` of shape `(text_batch_size, video_batch_size)`):66 The scaled dot product scores between `text_embeds` and `video_embeds`. This represents the text-video67 similarity scores.68 text_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):69 The text embeddings obtained by applying the projection layer to the pooled output of [`XCLIPTextModel`].70 video_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):71 The video embeddings obtained by applying the projection layer to the pooled output of72 [`XCLIPVisionModel`].73 text_model_output (`BaseModelOutputWithPooling`):74 The output of the [`XCLIPTextModel`].75 vision_model_output (`BaseModelOutputWithPooling`):76 The output of the [`XCLIPVisionModel`].77 mit_output (`BaseModelOutputWithPooling`):78 The output of `XCLIPMultiframeIntegrationTransformer` (MIT for short).79 """80 81 loss: Optional[torch.FloatTensor] = None82 logits_per_video: Optional[torch.FloatTensor] = None83 logits_per_text: Optional[torch.FloatTensor] = None84 text_embeds: Optional[torch.FloatTensor] = None85 video_embeds: Optional[torch.FloatTensor] = None86 text_model_output: BaseModelOutputWithPooling = None87 vision_model_output: BaseModelOutputWithPooling = None88 mit_output: BaseModelOutputWithPooling = None89 90 def to_tuple(self) -> tuple[Any]:91 return tuple(92 self[k]93 if k not in ["text_model_output", "vision_model_output", "mit_output"]94 else getattr(self, k).to_tuple()95 for k in self.keys()96 )97 98 99# Copied from transformers.models.clip.modeling_clip.CLIPVisionEmbeddings with CLIP->XCLIP100class XCLIPVisionEmbeddings(nn.Module):101 def __init__(self, config: XCLIPVisionConfig):102 super().__init__()103 self.config = config104 self.embed_dim = config.hidden_size105 self.image_size = config.image_size106 self.patch_size = config.patch_size107 108 self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))109 110 self.patch_embedding = nn.Conv2d(111 in_channels=config.num_channels,112 out_channels=self.embed_dim,113 kernel_size=self.patch_size,114 stride=self.patch_size,115 bias=False,116 )117 118 self.num_patches = (self.image_size // self.patch_size) ** 2119 self.num_positions = self.num_patches + 1120 self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)121 self.register_buffer("position_ids", torch.arange(self.num_positions).expand((1, -1)), persistent=False)122 123 def interpolate_pos_encoding(self, embeddings: torch.Tensor, height: int, width: int) -> torch.Tensor:124 """125 This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution126 images. This method is also adapted to support torch.jit tracing.127 128 Adapted from:129 - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and130 - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211131 """132 133 num_patches = embeddings.shape[1] - 1134 position_embedding = self.position_embedding.weight.unsqueeze(0)135 num_positions = position_embedding.shape[1] - 1136 137 # always interpolate when tracing to ensure the exported model works for dynamic input shapes138 if not torch.jit.is_tracing() and num_patches == num_positions and height == width:139 return self.position_embedding(self.position_ids)140 141 class_pos_embed = position_embedding[:, :1]142 patch_pos_embed = position_embedding[:, 1:]143 144 dim = embeddings.shape[-1]145 146 new_height = height // self.patch_size147 new_width = width // self.patch_size148 149 sqrt_num_positions = torch_int(num_positions**0.5)150 patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)151 patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)152 153 patch_pos_embed = nn.functional.interpolate(154 patch_pos_embed,155 size=(new_height, new_width),156 mode="bicubic",157 align_corners=False,158 )159 160 patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)161 162 return torch.cat((class_pos_embed, patch_pos_embed), dim=1)163 164 def forward(self, pixel_values: torch.FloatTensor, interpolate_pos_encoding=False) -> torch.Tensor:165 batch_size, _, height, width = pixel_values.shape166 if not interpolate_pos_encoding and (height != self.image_size or width != self.image_size):167 raise ValueError(168 f"Input image size ({height}*{width}) doesn't match model ({self.image_size}*{self.image_size})."169 )170 target_dtype = self.patch_embedding.weight.dtype171 patch_embeds = self.patch_embedding(pixel_values.to(dtype=target_dtype)) # shape = [*, width, grid, grid]172 patch_embeds = patch_embeds.flatten(2).transpose(1, 2)173 174 class_embeds = self.class_embedding.expand(batch_size, 1, -1)175 embeddings = torch.cat([class_embeds, patch_embeds], dim=1)176 if interpolate_pos_encoding:177 embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)178 else:179 embeddings = embeddings + self.position_embedding(self.position_ids)180 return embeddings181 182 183# Copied from transformers.models.clip.modeling_clip.CLIPTextEmbeddings with CLIP->XCLIP184class XCLIPTextEmbeddings(nn.Module):185 def __init__(self, config: XCLIPTextConfig):186 super().__init__()187 embed_dim = config.hidden_size188 189 self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)190 self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)191 192 # position_ids (1, len position emb) is contiguous in memory and exported when serialized193 self.register_buffer(194 "position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False195 )196 197 def forward(198 self,199 input_ids: Optional[torch.LongTensor] = None,200 position_ids: Optional[torch.LongTensor] = None,201 inputs_embeds: Optional[torch.FloatTensor] = None,202 ) -> torch.Tensor:203 seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]204 max_position_embedding = self.position_embedding.weight.shape[0]205 206 if seq_length > max_position_embedding:207 raise ValueError(208 f"Sequence length must be less than max_position_embeddings (got `sequence length`: "209 f"{seq_length} and max_position_embeddings: {max_position_embedding}"210 )211 212 if position_ids is None:213 position_ids = self.position_ids[:, :seq_length]214 215 if inputs_embeds is None:216 inputs_embeds = self.token_embedding(input_ids)217 218 position_embeddings = self.position_embedding(position_ids)219 embeddings = inputs_embeds + position_embeddings220 221 return embeddings222 223 224# Copied from transformers.models.siglip.modeling_siglip.eager_attention_forward225def eager_attention_forward(226 module: nn.Module,227 query: torch.Tensor,228 key: torch.Tensor,229 value: torch.Tensor,230 attention_mask: Optional[torch.Tensor],231 scaling: float,232 dropout: float = 0.0,233 **kwargs,234):235 attn_weights = torch.matmul(query, key.transpose(-1, -2)) * scaling236 if attention_mask is not None:237 attn_weights = attn_weights + attention_mask238 239 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)240 attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)241 242 attn_output = torch.matmul(attn_weights, value)243 attn_output = attn_output.transpose(1, 2).contiguous()244 245 return attn_output, attn_weights246 247 248class XCLIPAttention(nn.Module):249 """Multi-headed attention from 'Attention Is All You Need' paper"""250 251 def __init__(self, config):252 super().__init__()253 self.config = config254 self.embed_dim = config.hidden_size255 self.num_heads = config.num_attention_heads256 self.head_dim = self.embed_dim // self.num_heads257 if self.head_dim * self.num_heads != self.embed_dim:258 raise ValueError(259 f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"260 f" {self.num_heads})."261 )262 self.scale = self.head_dim**-0.5263 self.dropout = config.attention_dropout264 self.is_causal = False265 266 self.k_proj = nn.Linear(self.embed_dim, self.embed_dim)267 self.v_proj = nn.Linear(self.embed_dim, self.embed_dim)268 self.q_proj = nn.Linear(self.embed_dim, self.embed_dim)269 self.out_proj = nn.Linear(self.embed_dim, self.embed_dim)270 271 def forward(272 self,273 hidden_states: torch.Tensor,274 attention_mask: Optional[torch.Tensor] = None,275 causal_attention_mask: Optional[torch.Tensor] = None,276 output_attentions: Optional[bool] = False,277 ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:278 """Input shape: Batch x Time x Channel"""279 280 batch_size, seq_length, embed_dim = hidden_states.shape281 282 queries = self.q_proj(hidden_states)283 keys = self.k_proj(hidden_states)284 values = self.v_proj(hidden_states)285 286 queries = queries.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)287 keys = keys.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)288 values = values.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)289 # CLIP text model uses both `causal_attention_mask` and `attention_mask`290 # in case FA2 kernel is called, `is_causal` should be inferred from `causal_attention_mask`291 if self.config._attn_implementation != "flash_attention_2":292 if attention_mask is not None and causal_attention_mask is not None:293 attention_mask = attention_mask + causal_attention_mask294 elif causal_attention_mask is not None:295 attention_mask = causal_attention_mask296 else:297 self.is_causal = causal_attention_mask is not None298 299 attention_interface: Callable = eager_attention_forward300 if self.config._attn_implementation != "eager":301 attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]302 303 attn_output, attn_weights = attention_interface(304 self,305 queries,306 keys,307 values,308 attention_mask,309 is_causal=self.is_causal,310 scaling=self.scale,311 dropout=0.0 if not self.training else self.dropout,312 )313 314 attn_output = attn_output.reshape(batch_size, seq_length, embed_dim).contiguous()315 attn_output = self.out_proj(attn_output)316 if not output_attentions:317 attn_weights = None318 319 return attn_output, attn_weights320 321 322class XCLIPMLP(nn.Module):323 def __init__(self, config):324 super().__init__()325 self.config = config326 self.activation_fn = ACT2FN[config.hidden_act]327 self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)328 self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)329 330 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:331 hidden_states = self.fc1(hidden_states)332 hidden_states = self.activation_fn(hidden_states)333 hidden_states = self.fc2(hidden_states)334 return hidden_states335 336 337# Copied from transformers.models.altclip.modeling_altclip.AltCLIPEncoderLayer with AltCLIP->XCLIP338class XCLIPEncoderLayer(GradientCheckpointingLayer):339 def __init__(self, config: XCLIPConfig):340 super().__init__()341 self.embed_dim = config.hidden_size342 self.self_attn = XCLIPAttention(config)343 self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)344 self.mlp = XCLIPMLP(config)345 self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)346 347 def forward(348 self,349 hidden_states: torch.Tensor,350 attention_mask: torch.Tensor,351 causal_attention_mask: torch.Tensor,352 output_attentions: Optional[bool] = False,353 ) -> tuple[torch.FloatTensor]:354 """355 Args:356 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`357 attention_mask (`torch.FloatTensor`): attention mask of size358 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.359 `(config.encoder_attention_heads,)`.360 output_attentions (`bool`, *optional*):361 Whether or not to return the attentions tensors of all attention layers. See `attentions` under362 returned tensors for more detail.363 """364 residual = hidden_states365 366 hidden_states = self.layer_norm1(hidden_states)367 hidden_states, attn_weights = self.self_attn(368 hidden_states=hidden_states,369 attention_mask=attention_mask,370 causal_attention_mask=causal_attention_mask,371 output_attentions=output_attentions,372 )373 hidden_states = residual + hidden_states374 375 residual = hidden_states376 hidden_states = self.layer_norm2(hidden_states)377 hidden_states = self.mlp(hidden_states)378 hidden_states = residual + hidden_states379 380 outputs = (hidden_states,)381 382 if output_attentions:383 outputs += (attn_weights,)384 385 return outputs386 387 388# Copied from transformers.models.beit.modeling_beit.drop_path389def drop_path(input: torch.Tensor, drop_prob: float = 0.0, training: bool = False) -> torch.Tensor:390 """391 Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).392 393 Comment by Ross Wightman: This is the same as the DropConnect impl I created for EfficientNet, etc networks,394 however, the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...395 See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for changing the396 layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use 'survival rate' as the397 argument.398 """399 if drop_prob == 0.0 or not training:400 return input401 keep_prob = 1 - drop_prob402 shape = (input.shape[0],) + (1,) * (input.ndim - 1) # work with diff dim tensors, not just 2D ConvNets403 random_tensor = keep_prob + torch.rand(shape, dtype=input.dtype, device=input.device)404 random_tensor.floor_() # binarize405 output = input.div(keep_prob) * random_tensor406 return output407 408 409# Copied from transformers.models.beit.modeling_beit.BeitDropPath with Beit->XCLIP410class XCLIPDropPath(nn.Module):411 """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks)."""412 413 def __init__(self, drop_prob: Optional[float] = None) -> None:414 super().__init__()415 self.drop_prob = drop_prob416 417 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:418 return drop_path(hidden_states, self.drop_prob, self.training)419 420 def extra_repr(self) -> str:421 return f"p={self.drop_prob}"422 423 424class XCLIPVisionEncoderLayer(GradientCheckpointingLayer):425 """426 This corresponds to the `CrossFramelAttentionBlock` class in the original implementation.427 """428 429 def __init__(self, config: XCLIPConfig):430 super().__init__()431 self.num_frames = config.num_frames432 self.embed_dim = config.hidden_size433 434 self.message_fc = nn.Linear(self.embed_dim, self.embed_dim)435 self.message_ln = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)436 self.message_attn = XCLIPAttention(config)437 438 self.drop_path = XCLIPDropPath(config.drop_path_rate) if config.drop_path_rate > 0.0 else nn.Identity()439 440 self.self_attn = XCLIPAttention(config)441 self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)442 self.mlp = XCLIPMLP(config)443 self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)444 445 def forward(446 self,447 hidden_states: torch.Tensor,448 attention_mask: torch.Tensor,449 causal_attention_mask: torch.Tensor,450 output_attentions: Optional[bool] = False,451 ) -> tuple[torch.FloatTensor]:452 """453 Args:454 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`455 attention_mask (`torch.FloatTensor`): attention mask of size456 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.457 `(config.encoder_attention_heads,)`.458 causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):459 Causal mask for the text model. Mask values selected in `[0, 1]`:460 - 1 for tokens that are **not masked**,461 - 0 for tokens that are **masked**.462 [What are attention masks?](../glossary#attention-mask)463 output_attentions (`bool`, *optional*):464 Whether or not to return the attentions tensors of all attention layers. See `attentions` under465 returned tensors for more detail.466 """467 batch_time, seq_length, hidden_size = hidden_states.size()468 batch_size = batch_time // self.num_frames469 msg_token = self.message_fc(hidden_states[:, 0, :])470 msg_token = msg_token.view(batch_size, self.num_frames, hidden_size)471 472 msg_token = msg_token + self.drop_path(self.message_attn(self.message_ln(msg_token))[0])473 # add dummy sequence dimension474 msg_token = msg_token.view(-1, 1, hidden_size)475 476 hidden_states = torch.cat([hidden_states, msg_token], dim=1)477 478 residual = hidden_states479 480 hidden_states = self.layer_norm1(hidden_states)481 hidden_states, attn_weights = self.self_attn(482 hidden_states=hidden_states,483 attention_mask=attention_mask,484 causal_attention_mask=causal_attention_mask,485 output_attentions=output_attentions,486 )487 hidden_states = residual + hidden_states488 489 hidden_states = hidden_states[:, :seq_length, :]490 491 residual = hidden_states492 hidden_states = self.layer_norm2(hidden_states)493 hidden_states = self.mlp(hidden_states)494 hidden_states = residual + hidden_states495 496 outputs = (hidden_states,)497 498 if output_attentions:499 outputs += (attn_weights,)500 501 return outputs502 503 504@auto_docstring505class XCLIPPreTrainedModel(PreTrainedModel):506 config: XCLIPConfig507 base_model_prefix = "x_clip"508 supports_gradient_checkpointing = True509 510 def _init_weights(self, module):511 """Initialize the weights"""512 factor = self.config.initializer_factor513 if isinstance(module, XCLIPTextEmbeddings):514 module.token_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)515 module.position_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)516 elif isinstance(module, XCLIPVisionEmbeddings):517 factor = self.config.initializer_factor518 nn.init.normal_(module.class_embedding, mean=0.0, std=module.embed_dim**-0.5 * factor)519 nn.init.normal_(module.patch_embedding.weight, std=module.config.initializer_range * factor)520 nn.init.normal_(module.position_embedding.weight, std=module.config.initializer_range * factor)521 elif isinstance(module, XCLIPAttention):522 factor = self.config.initializer_factor523 in_proj_std = (module.embed_dim**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor524 out_proj_std = (module.embed_dim**-0.5) * factor525 nn.init.normal_(module.q_proj.weight, std=in_proj_std)526 nn.init.normal_(module.k_proj.weight, std=in_proj_std)527 nn.init.normal_(module.v_proj.weight, std=in_proj_std)528 nn.init.normal_(module.out_proj.weight, std=out_proj_std)529 elif isinstance(module, XCLIPMLP):530 factor = self.config.initializer_factor531 in_proj_std = (module.config.hidden_size**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor532 fc_std = (2 * module.config.hidden_size) ** -0.5 * factor533 nn.init.normal_(module.fc1.weight, std=fc_std)534 nn.init.normal_(module.fc2.weight, std=in_proj_std)535 elif isinstance(module, XCLIPModel):536 factor = self.config.initializer_factor537 nn.init.normal_(538 module.text_projection.weight,539 std=module.text_embed_dim**-0.5 * factor,540 )541 nn.init.normal_(542 module.visual_projection.weight,543 std=module.vision_embed_dim**-0.5 * factor,544 )545 nn.init.normal_(module.prompts_visual_projection, mean=0.0, std=module.vision_embed_dim**-0.5 * factor)546 elif isinstance(module, XCLIPMultiframeIntegrationTransformer):547 nn.init.normal_(module.position_embedding, std=self.config.initializer_factor)548 549 if isinstance(module, nn.LayerNorm):550 module.bias.data.zero_()551 module.weight.data.fill_(1.0)552 if isinstance(module, nn.Linear):553 module.weight.data.normal_(mean=0.0, std=self.config.initializer_factor)554 if module.bias is not None:555 module.bias.data.zero_()556 557 558# Copied from transformers.models.altclip.modeling_altclip.AltCLIPEncoder with AltCLIP->XCLIP559class XCLIPEncoder(nn.Module):560 """561 Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a562 [`XCLIPEncoderLayer`].563 564 Args:565 config: XCLIPConfig566 """567 568 def __init__(self, config: XCLIPConfig):569 super().__init__()570 self.config = config571 self.layers = nn.ModuleList([XCLIPEncoderLayer(config) for _ in range(config.num_hidden_layers)])572 self.gradient_checkpointing = False573 574 @can_return_tuple575 def forward(576 self,577 inputs_embeds,578 attention_mask: Optional[torch.Tensor] = None,579 causal_attention_mask: Optional[torch.Tensor] = None,580 output_attentions: Optional[bool] = None,581 output_hidden_states: Optional[bool] = None,582 return_dict: Optional[bool] = None,583 ) -> Union[tuple, BaseModelOutput]:584 r"""585 Args:586 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):587 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.588 This is useful if you want more control over how to convert `input_ids` indices into associated vectors589 than the model's internal embedding lookup matrix.590 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):591 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:592 593 - 1 for tokens that are **not masked**,594 - 0 for tokens that are **masked**.595 596 [What are attention masks?](../glossary#attention-mask)597 causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):598 Causal mask for the text model. Mask values selected in `[0, 1]`:599 600 - 1 for tokens that are **not masked**,601 - 0 for tokens that are **masked**.602 603 [What are attention masks?](../glossary#attention-mask)604 output_attentions (`bool`, *optional*):605 Whether or not to return the attentions tensors of all attention layers. See `attentions` under606 returned tensors for more detail.607 output_hidden_states (`bool`, *optional*):608 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors609 for more detail.610 return_dict (`bool`, *optional*):611 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.612 """613 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions614 output_hidden_states = (615 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states616 )617 return_dict = return_dict if return_dict is not None else self.config.use_return_dict618 619 encoder_states = () if output_hidden_states else None620 all_attentions = () if output_attentions else None621 622 hidden_states = inputs_embeds623 for idx, encoder_layer in enumerate(self.layers):624 if output_hidden_states:625 encoder_states = encoder_states + (hidden_states,)626 layer_outputs = encoder_layer(627 hidden_states,628 attention_mask,629 causal_attention_mask,630 output_attentions=output_attentions,631 )632 633 hidden_states = layer_outputs[0]634 635 if output_attentions:636 all_attentions = all_attentions + (layer_outputs[1],)637 638 if output_hidden_states:639 encoder_states = encoder_states + (hidden_states,)640 641 return BaseModelOutput(642 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions643 )644 645 646class XCLIPTextTransformer(nn.Module):647 def __init__(self, config: XCLIPTextConfig):648 super().__init__()649 self.config = config650 embed_dim = config.hidden_size651 self.embeddings = XCLIPTextEmbeddings(config)652 self.encoder = XCLIPEncoder(config)653 self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)654 655 @auto_docstring656 def forward(657 self,658 input_ids: Optional[torch.Tensor] = None,659 attention_mask: Optional[torch.Tensor] = None,660 position_ids: Optional[torch.Tensor] = None,661 output_attentions: Optional[bool] = None,662 output_hidden_states: Optional[bool] = None,663 return_dict: Optional[bool] = None,664 ) -> Union[tuple, BaseModelOutputWithPooling]:665 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions666 output_hidden_states = (667 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states668 )669 return_dict = return_dict if return_dict is not None else self.config.use_return_dict670 671 if input_ids is None:672 raise ValueError("You have to specify either input_ids")673 674 input_shape = input_ids.size()675 input_ids = input_ids.view(-1, input_shape[-1])676 677 hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)678 679 # X_CLIP's text model uses causal mask, prepare it here.680 # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324681 causal_attention_mask = _create_4d_causal_attention_mask(682 input_shape, hidden_states.dtype, device=hidden_states.device683 )684 # expand attention_mask685 if attention_mask is not None:686 # [batch_size, seq_len] -> [batch_size, 1, tgt_seq_len, src_seq_len]687 attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)688 689 encoder_outputs = self.encoder(690 inputs_embeds=hidden_states,691 attention_mask=attention_mask,692 causal_attention_mask=causal_attention_mask,693 output_attentions=output_attentions,694 output_hidden_states=output_hidden_states,695 return_dict=return_dict,696 )697 698 last_hidden_state = encoder_outputs[0]699 last_hidden_state = self.final_layer_norm(last_hidden_state)700 701 # text_embeds.shape = [batch_size, sequence_length, transformer.width]702 # take features from the eot embedding (eot_token is the highest number in each sequence)703 pooled_output = last_hidden_state[torch.arange(last_hidden_state.shape[0]), input_ids.argmax(dim=-1)]704 705 if not return_dict:706 return (last_hidden_state, pooled_output) + encoder_outputs[1:]707 708 return BaseModelOutputWithPooling(709 last_hidden_state=last_hidden_state,710 pooler_output=pooled_output,711 hidden_states=encoder_outputs.hidden_states,712 attentions=encoder_outputs.attentions,713 )714 715 716class XCLIPTextModel(XCLIPPreTrainedModel):717 config: XCLIPTextConfig718 719 def __init__(self, config: XCLIPTextConfig):720 super().__init__(config)721 self.text_model = XCLIPTextTransformer(config)722 # Initialize weights and apply final processing723 self.post_init()724 725 def get_input_embeddings(self) -> nn.Module:726 return self.text_model.embeddings.token_embedding727 728 def set_input_embeddings(self, value):729 self.text_model.embeddings.token_embedding = value730 731 @auto_docstring732 def forward(733 self,734 input_ids: Optional[torch.Tensor] = None,735 attention_mask: Optional[torch.Tensor] = None,736 position_ids: Optional[torch.Tensor] = None,737 output_attentions: Optional[bool] = None,738 output_hidden_states: Optional[bool] = None,739 return_dict: Optional[bool] = None,740 ) -> Union[tuple, BaseModelOutputWithPooling]:741 r"""742 Examples:743 744 ```python745 >>> from transformers import AutoTokenizer, XCLIPTextModel746 747 >>> model = XCLIPTextModel.from_pretrained("microsoft/xclip-base-patch32")748 >>> tokenizer = AutoTokenizer.from_pretrained("microsoft/xclip-base-patch32")749 750 >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")751 752 >>> outputs = model(**inputs)753 >>> last_hidden_state = outputs.last_hidden_state754 >>> pooled_output = outputs.pooler_output # pooled (EOS token) states755 ```"""756 return self.text_model(757 input_ids=input_ids,758 attention_mask=attention_mask,759 position_ids=position_ids,760 output_attentions=output_attentions,761 output_hidden_states=output_hidden_states,762 return_dict=return_dict,763 )764 765 766class XCLIPVisionEncoder(nn.Module):767 """768 Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a769 [`XCLIPVisionEncoderLayer`].770 771 Args:772 config: XCLIPConfig773 """774 775 def __init__(self, config: XCLIPConfig):776 super().__init__()777 self.config = config778 self.layers = nn.ModuleList([XCLIPVisionEncoderLayer(config) for _ in range(config.num_hidden_layers)])779 self.gradient_checkpointing = False780 781 def forward(782 self,783 inputs_embeds,784 attention_mask: Optional[torch.Tensor] = None,785 causal_attention_mask: Optional[torch.Tensor] = None,786 output_attentions: Optional[bool] = None,787 output_hidden_states: Optional[bool] = None,788 return_dict: Optional[bool] = None,789 ) -> Union[tuple, BaseModelOutput]:790 r"""791 Args:792 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):793 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.794 This is useful if you want more control over how to convert `input_ids` indices into associated vectors795 than the model's internal embedding lookup matrix.796 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):797 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:798 799 - 1 for tokens that are **not masked**,800 - 0 for tokens that are **masked**.801 802 [What are attention masks?](../glossary#attention-mask)803 causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):804 Causal mask for the text model. Mask values selected in `[0, 1]`:805 806 - 1 for tokens that are **not masked**,807 - 0 for tokens that are **masked**.808 809 [What are attention masks?](../glossary#attention-mask)810 output_attentions (`bool`, *optional*):811 Whether or not to return the attentions tensors of all attention layers. See `attentions` under812 returned tensors for more detail.813 output_hidden_states (`bool`, *optional*):814 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors815 for more detail.816 return_dict (`bool`, *optional*):817 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.818 """819 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions820 output_hidden_states = (821 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states822 )823 return_dict = return_dict if return_dict is not None else self.config.use_return_dict824 825 encoder_states = () if output_hidden_states else None826 all_attentions = () if output_attentions else None827 828 hidden_states = inputs_embeds829 for idx, encoder_layer in enumerate(self.layers):830 if output_hidden_states:831 encoder_states = encoder_states + (hidden_states,)832 layer_outputs = encoder_layer(833 hidden_states,834 attention_mask,835 causal_attention_mask,836 output_attentions=output_attentions,837 )838 839 hidden_states = layer_outputs[0]840 841 if output_attentions:842 all_attentions = all_attentions + (layer_outputs[1],)843 844 if output_hidden_states:845 encoder_states = encoder_states + (hidden_states,)846 847 if not return_dict:848 return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)849 return BaseModelOutput(850 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions851 )852 853 854class XCLIPVisionTransformer(nn.Module):855 """856 This corresponds to the `CrossFrameCommunicationTransformer` class in the original implementation.857 """858 859 def __init__(self, config: XCLIPVisionConfig):860 super().__init__()861 self.config = config862 embed_dim = config.hidden_size863 864 self.embeddings = XCLIPVisionEmbeddings(config)865 self.pre_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)866 self.encoder = XCLIPVisionEncoder(config)867 self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)868 869 @auto_docstring870 def forward(871 self,872 pixel_values: torch.FloatTensor,873 output_attentions: Optional[bool] = None,874 output_hidden_states: Optional[bool] = None,875 interpolate_pos_encoding: bool = False,876 return_dict: Optional[bool] = None,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 hidden_states = self.embeddings(pixel_values, interpolate_pos_encoding=interpolate_pos_encoding)885 hidden_states = self.pre_layernorm(hidden_states)886 887 encoder_outputs = self.encoder(888 inputs_embeds=hidden_states,889 output_attentions=output_attentions,890 output_hidden_states=output_hidden_states,891 return_dict=return_dict,892 )893 894 last_hidden_state = encoder_outputs[0]895 pooled_output = last_hidden_state[:, 0, :]896 pooled_output = self.post_layernorm(pooled_output)897 898 if not return_dict:899 return (last_hidden_state, pooled_output) + encoder_outputs[1:]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 XCLIPVisionModel(XCLIPPreTrainedModel):910 config: XCLIPVisionConfig911 main_input_name = "pixel_values"912 913 def __init__(self, config: XCLIPVisionConfig):914 super().__init__(config)915 self.vision_model = XCLIPVisionTransformer(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 return_dict: Optional[bool] = None,929 ) -> Union[tuple, BaseModelOutputWithPooling]:930 r"""931 Examples:932 933 ```python934 >>> import av935 >>> import torch936 >>> import numpy as np937 938 >>> from transformers import AutoProcessor, XCLIPVisionModel939 >>> from huggingface_hub import hf_hub_download940 941 >>> np.random.seed(0)942 943 944 >>> def read_video_pyav(container, indices):945 ... '''946 ... Decode the video with PyAV decoder.947 ... Args:948 ... container (`av.container.input.InputContainer`): PyAV container.949 ... indices (`list[int]`): List of frame indices to decode.950 ... Returns:951 ... result (np.ndarray): np array of decoded frames of shape (num_frames, height, width, 3).952 ... '''953 ... frames = []954 ... container.seek(0)955 ... start_index = indices[0]956 ... end_index = indices[-1]957 ... for i, frame in enumerate(container.decode(video=0)):958 ... if i > end_index:959 ... break960 ... if i >= start_index and i in indices:961 ... frames.append(frame)962 ... return np.stack([x.to_ndarray(format="rgb24") for x in frames])963 964 965 >>> def sample_frame_indices(clip_len, frame_sample_rate, seg_len):966 ... '''967 ... Sample a given number of frame indices from the video.968 ... Args:969 ... clip_len (`int`): Total number of frames to sample.970 ... frame_sample_rate (`int`): Sample every n-th frame.971 ... seg_len (`int`): Maximum allowed index of sample's last frame.972 ... Returns:973 ... indices (`list[int]`): List of sampled frame indices974 ... '''975 ... converted_len = int(clip_len * frame_sample_rate)976 ... end_idx = np.random.randint(converted_len, seg_len)977 ... start_idx = end_idx - converted_len978 ... indices = np.linspace(start_idx, end_idx, num=clip_len)979 ... indices = np.clip(indices, start_idx, end_idx - 1).astype(np.int64)980 ... return indices981 982 983 >>> # video clip consists of 300 frames (10 seconds at 30 FPS)984 >>> file_path = hf_hub_download(985 ... repo_id="nielsr/video-demo", filename="eating_spaghetti.mp4", repo_type="dataset"986 ... )987 >>> container = av.open(file_path)988 989 >>> # sample 16 frames990 >>> indices = sample_frame_indices(clip_len=8, frame_sample_rate=1, seg_len=container.streams.video[0].frames)991 >>> video = read_video_pyav(container, indices)992 993 >>> processor = AutoProcessor.from_pretrained("microsoft/xclip-base-patch32")994 >>> model = XCLIPVisionModel.from_pretrained("microsoft/xclip-base-patch32")995 996 >>> pixel_values = processor(videos=list(video), return_tensors="pt").pixel_values997 998 >>> batch_size, num_frames, num_channels, height, width = pixel_values.shape999 >>> pixel_values = pixel_values.reshape(-1, num_channels, height, width)1000 1001 >>> outputs = model(pixel_values)1002 >>> last_hidden_state = outputs.last_hidden_state1003 ```"""1004 return self.vision_model(1005 pixel_values=pixel_values,1006 output_attentions=output_attentions,1007 output_hidden_states=output_hidden_states,1008 return_dict=return_dict,1009 )1010 1011 1012class XCLIPMultiframeIntegrationTransformer(nn.Module):1013 """1014 This corresponds to the `MultiframeIntegrationTransformer` class in the original implementation.1015 """1016 1017 def __init__(self, config: XCLIPVisionConfig):1018 super().__init__()1019 1020 self.position_embedding = nn.Parameter(torch.empty(1, config.num_frames, config.hidden_size))1021 self.encoder = XCLIPEncoder(config)1022 1023 def forward(1024 self,1025 hidden_states,1026 output_attentions: Optional[bool] = None,1027 output_hidden_states: Optional[bool] = None,1028 return_dict: Optional[bool] = None,1029 ) -> Union[tuple, BaseModelOutput]:1030 residual = hidden_states1031 1032 # add position embeddings1033 hidden_states = hidden_states + self.position_embedding1034 1035 encoder_outputs = self.encoder(1036 inputs_embeds=hidden_states,1037 output_attentions=output_attentions,1038 output_hidden_states=output_hidden_states,1039 return_dict=return_dict,1040 )1041 last_hidden_state = encoder_outputs[0]1042 1043 last_hidden_state = last_hidden_state.type(hidden_states.dtype) + residual1044 1045 pooled_output = last_hidden_state.mean(dim=1, keepdim=False)1046 1047 if not return_dict:1048 return (last_hidden_state, pooled_output) + encoder_outputs[1:]1049 1050 return BaseModelOutputWithPooling(1051 last_hidden_state=last_hidden_state,1052 pooler_output=pooled_output,1053 hidden_states=encoder_outputs.hidden_states,1054 attentions=encoder_outputs.attentions,1055 )1056 1057 1058class XCLIPCrossAttention(nn.Module):1059 """Multi-headed attention from 'Attention Is All You Need' paper"""1060 1061 def __init__(self, config):1062 super().__init__()1063 self.num_heads = config.prompt_num_attention_heads1064 1065 dim = config.projection_dim1066 head_dim = dim // self.num_heads1067 self.scale = head_dim**-0.51068 1069 self.q_proj = nn.Linear(dim, dim, bias=False)1070 self.k_proj = nn.Linear(dim, dim, bias=False)1071 self.v_proj = nn.Linear(dim, dim, bias=False)1072 1073 self.attn_drop = nn.Dropout(config.prompt_attention_dropout)1074 self.proj = nn.Linear(dim, dim)1075 self.proj_drop = nn.Dropout(config.prompt_projection_dropout)1076 1077 def _shape(self, tensor: torch.Tensor, seq_len: int, batch_size: int):1078 return tensor.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()1079 1080 def forward(self, queries, keys, values):1081 """Input shape: Batch x Time x Channel"""1082 batch_size, query_seq_len, hidden_size = queries.shape1083 batch_size, key_seq_len, hidden_size = keys.shape1084 queries = (1085 self.q_proj(queries)1086 .reshape(batch_size, query_seq_len, self.num_heads, hidden_size // self.num_heads)1087 .permute(0, 2, 1, 3)1088 )1089 keys = (1090 self.k_proj(keys)1091 .reshape(batch_size, key_seq_len, self.num_heads, hidden_size // self.num_heads)1092 .permute(0, 2, 1, 3)1093 )1094 values = (1095 self.v_proj(values)1096 .reshape(batch_size, key_seq_len, self.num_heads, hidden_size // self.num_heads)1097 .permute(0, 2, 1, 3)1098 )1099 1100 attn = (queries @ keys.transpose(-2, -1)) * self.scale1101 attn = attn.softmax(dim=-1)1102 attn = self.attn_drop(attn)1103 1104 x = (attn @ values).transpose(1, 2).reshape(batch_size, query_seq_len, hidden_size)1105 x = self.proj(x)1106 x = self.proj_drop(x)1107 return x1108 1109 1110class PromptGeneratorLayer(nn.Module):1111 def __init__(self, config):1112 super().__init__()1113 1114 embed_dim = config.projection_dim1115 self.cross_attn = XCLIPCrossAttention(config)1116 self.norm1 = nn.LayerNorm(embed_dim, eps=config.text_config.layer_norm_eps)1117 self.norm3 = nn.LayerNorm(embed_dim, eps=config.text_config.layer_norm_eps)1118 self.mlp = nn.Sequential(1119 nn.Linear(embed_dim, embed_dim * 4),1120 ACT2FN[config.prompt_hidden_act],1121 nn.Dropout(config.prompt_attention_dropout),1122 nn.Linear(embed_dim * 4, embed_dim),1123 )1124 1125 def forward(self, x, visual):1126 x = x + self.cross_attn(self.norm1(x), visual, visual)1127 x = x + self.mlp(self.norm3(x))1128 return x1129 1130 1131class XCLIPPromptGenerator(nn.Module):1132 """This corresponds to the `VideoSpecificPrompt` class in the original implementation."""1133 1134 def __init__(self, config):1135 super().__init__()1136 embed_dim = config.projection_dim1137 self.layernorm = nn.LayerNorm(embed_dim, eps=config.vision_config.layer_norm_eps)1138 self.decoder = nn.ModuleList([PromptGeneratorLayer(config) for _ in range(config.prompt_layers)])1139 self.alpha = nn.Parameter(torch.ones(embed_dim) * config.prompt_alpha)1140 1141 def forward(self, text, visual):1142 visual = self.layernorm(visual)1143 for layer in self.decoder:1144 text = layer(text, visual)1145 1146 return self.alpha * text1147 1148 1149@auto_docstring1150class XCLIPModel(XCLIPPreTrainedModel):1151 config: XCLIPConfig1152 1153 def __init__(self, config: XCLIPConfig):1154 super().__init__(config)1155 1156 if not isinstance(config.text_config, XCLIPTextConfig):1157 raise TypeError(1158 "config.text_config is expected to be of type XCLIPTextConfig but is of type"1159 f" {type(config.text_config)}."1160 )1161 1162 if not isinstance(config.vision_config, XCLIPVisionConfig):1163 raise TypeError(1164 "config.vision_config is expected to be of type XCLIPVisionConfig but is of type"1165 f" {type(config.vision_config)}."1166 )1167 1168 text_config = config.text_config1169 vision_config = config.vision_config1170 # The module using it is not a PreTrainedModel subclass so we need this1171 text_config._attn_implementation = config._attn_implementation1172 # The module using it is not a PreTrainedModel subclass so we need this1173 vision_config._attn_implementation = config._attn_implementation1174 1175 self.projection_dim = config.projection_dim1176 self.text_embed_dim = text_config.hidden_size1177 self.vision_embed_dim = vision_config.hidden_size1178 1179 self.text_model = XCLIPTextTransformer(text_config)1180 self.vision_model = XCLIPVisionTransformer(vision_config)1181 1182 self.visual_projection = nn.Linear(self.vision_embed_dim, self.projection_dim, bias=False)1183 self.text_projection = nn.Linear(self.text_embed_dim, self.projection_dim, bias=False)1184 self.logit_scale = nn.Parameter(torch.tensor(self.config.logit_scale_init_value))1185 1186 self.prompts_visual_layernorm = nn.LayerNorm(self.vision_embed_dim, eps=config.vision_config.layer_norm_eps)1187 self.prompts_visual_projection = nn.Parameter(torch.randn(self.vision_embed_dim, self.projection_dim))1188 mit_config = copy.copy(vision_config)1189 mit_config.hidden_size = vision_config.mit_hidden_size1190 mit_config.intermediate_size = vision_config.mit_intermediate_size1191 mit_config.num_hidden_layers = vision_config.mit_num_hidden_layers1192 mit_config.num_attention_heads = vision_config.mit_num_attention_heads1193 self.mit = XCLIPMultiframeIntegrationTransformer(mit_config)1194 1195 self.prompts_generator = XCLIPPromptGenerator(config)1196 1197 # Initialize weights and apply final processing1198 self.post_init()1199 1200 @filter_out_non_signature_kwargs()