XaviXva/Video-LLaVA
0
1import math2from typing import Optional, Tuple, Union3 4import torch5from einops import rearrange6from peft import LoraConfig, get_peft_model7from torch import nn8from torch.nn import functional as F9from transformers import PreTrainedModel, add_start_docstrings10from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling11from transformers.models.clip.modeling_clip import CLIPMLP, CLIPAttention, CLIPTextEmbeddings, CLIPVisionEmbeddings, \12 CLIPVisionModelWithProjection, CLIPTextModelWithProjection, _expand_mask, CLIPOutput, clip_loss13from transformers.utils import add_start_docstrings_to_model_forward, replace_return_docstrings14 15from .configuration_depth import LanguageBindDepthConfig, CLIPVisionConfig, CLIPTextConfig16 17 18 19class PatchDropout(nn.Module):20 """21 https://arxiv.org/abs/2212.0079422 """23 24 def __init__(self, prob, exclude_first_token=True):25 super().__init__()26 assert 0 <= prob < 1.27 self.prob = prob28 self.exclude_first_token = exclude_first_token # exclude CLS token29 30 def forward(self, x, B, T):31 if not self.training or self.prob == 0.:32 return x33 34 if self.exclude_first_token:35 cls_tokens, x = x[:, :1], x[:, 1:]36 else:37 cls_tokens = torch.jit.annotate(torch.Tensor, x[:, :1])38 39 batch = x.size()[0]40 num_tokens = x.size()[1]41 42 batch_indices = torch.arange(batch)43 batch_indices = batch_indices[..., None]44 45 keep_prob = 1 - self.prob46 num_patches_keep = max(1, int(num_tokens * keep_prob))47 48 if T == 1:49 rand = torch.randn(batch, num_tokens)50 patch_indices_keep = rand.topk(num_patches_keep, dim=-1).indices51 else:52 rand = torch.randn(B, num_tokens)53 patch_indices_keep = rand.topk(num_patches_keep, dim=-1).indices54 patch_indices_keep = patch_indices_keep.unsqueeze(1).repeat(1, T, 1)55 patch_indices_keep = rearrange(patch_indices_keep, 'b t n -> (b t) n')56 57 58 x = x[batch_indices, patch_indices_keep]59 60 if self.exclude_first_token:61 x = torch.cat((cls_tokens, x), dim=1)62 63 return x64 65class CLIPEncoderLayer(nn.Module):66 def __init__(self, config: LanguageBindDepthConfig):67 super().__init__()68 self.embed_dim = config.hidden_size69 self.self_attn = CLIPAttention(config)70 self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)71 self.mlp = CLIPMLP(config)72 self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)73 74 self.add_time_attn = config.add_time_attn75 if self.add_time_attn:76 self.t = config.num_frames77 self.temporal_embedding = nn.Parameter(torch.zeros(1, config.num_frames, config.hidden_size))78 nn.init.normal_(self.temporal_embedding, std=config.hidden_size ** -0.5)79 80 self.embed_dim = config.hidden_size81 self.temporal_attn = CLIPAttention(config)82 self.temporal_layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)83 self.temporal_mlp = CLIPMLP(config)84 self.temporal_layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)85 86 def forward(87 self,88 hidden_states: torch.Tensor,89 attention_mask: torch.Tensor,90 causal_attention_mask: torch.Tensor,91 output_attentions: Optional[bool] = False,92 ) -> Tuple[torch.FloatTensor]:93 """94 Args:95 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`96 attention_mask (`torch.FloatTensor`): attention mask of size97 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.98 `(config.encoder_attention_heads,)`.99 output_attentions (`bool`, *optional*):100 Whether or not to return the attentions tensors of all attention layers. See `attentions` under101 returned tensors for more detail.102 """103 104 105 if self.add_time_attn:106 bt, n, d = hidden_states.shape107 t = self.t108 109 # time embed110 if t != 1:111 n = hidden_states.shape[1]112 hidden_states = rearrange(hidden_states, '(b t) n d -> (b n) t d', t=t)113 hidden_states = hidden_states + self.temporal_embedding[:, :t, :]114 hidden_states = rearrange(hidden_states, '(b n) t d -> (b t) n d', n=n)115 116 # time attn117 residual = hidden_states118 hidden_states = rearrange(hidden_states, '(b t) n d -> (b n) t d', t=t)119 # hidden_states = self.layer_norm1(hidden_states) # share layernorm120 hidden_states = self.temporal_layer_norm1(hidden_states)121 hidden_states, attn_weights = self.temporal_attn(122 hidden_states=hidden_states,123 attention_mask=attention_mask,124 causal_attention_mask=causal_attention_mask,125 output_attentions=output_attentions,126 )127 hidden_states = residual + rearrange(hidden_states, '(b n) t d -> (b t) n d', n=n)128 129 residual = hidden_states130 hidden_states = rearrange(hidden_states, '(b t) n d -> (b n) t d', t=t)131 # hidden_states = self.layer_norm2(hidden_states) # share layernorm132 hidden_states = self.temporal_layer_norm2(hidden_states)133 hidden_states = self.temporal_mlp(hidden_states)134 hidden_states = residual + rearrange(hidden_states, '(b n) t d -> (b t) n d', n=n)135 136 # spatial attn137 residual = hidden_states138 139 hidden_states = self.layer_norm1(hidden_states)140 hidden_states, attn_weights = self.self_attn(141 hidden_states=hidden_states,142 attention_mask=attention_mask,143 causal_attention_mask=causal_attention_mask,144 output_attentions=output_attentions,145 )146 hidden_states = residual + hidden_states147 148 residual = hidden_states149 hidden_states = self.layer_norm2(hidden_states)150 hidden_states = self.mlp(hidden_states)151 hidden_states = residual + hidden_states152 153 outputs = (hidden_states,)154 155 if output_attentions:156 outputs += (attn_weights,)157 158 return outputs159 160 161 162 163 164 165 166 167 168class CLIPPreTrainedModel(PreTrainedModel):169 """170 An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained171 models.172 """173 174 config_class = LanguageBindDepthConfig175 base_model_prefix = "clip"176 supports_gradient_checkpointing = True177 _keys_to_ignore_on_load_missing = [r"position_ids"]178 179 def _init_weights(self, module):180 """Initialize the weights"""181 factor = self.config.initializer_factor182 if isinstance(module, CLIPTextEmbeddings):183 module.token_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)184 module.position_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)185 elif isinstance(module, CLIPVisionEmbeddings):186 factor = self.config.initializer_factor187 nn.init.normal_(module.class_embedding, mean=0.0, std=module.embed_dim**-0.5 * factor)188 nn.init.normal_(module.patch_embedding.weight, std=module.config.initializer_range * factor)189 nn.init.normal_(module.position_embedding.weight, std=module.config.initializer_range * factor)190 elif isinstance(module, CLIPAttention):191 factor = self.config.initializer_factor192 in_proj_std = (module.embed_dim**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor193 out_proj_std = (module.embed_dim**-0.5) * factor194 nn.init.normal_(module.q_proj.weight, std=in_proj_std)195 nn.init.normal_(module.k_proj.weight, std=in_proj_std)196 nn.init.normal_(module.v_proj.weight, std=in_proj_std)197 nn.init.normal_(module.out_proj.weight, std=out_proj_std)198 elif isinstance(module, CLIPMLP):199 factor = self.config.initializer_factor200 in_proj_std = (201 (module.config.hidden_size**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor202 )203 fc_std = (2 * module.config.hidden_size) ** -0.5 * factor204 nn.init.normal_(module.fc1.weight, std=fc_std)205 nn.init.normal_(module.fc2.weight, std=in_proj_std)206 elif isinstance(module, LanguageBindDepth):207 nn.init.normal_(208 module.text_projection.weight,209 std=module.text_embed_dim**-0.5 * self.config.initializer_factor,210 )211 nn.init.normal_(212 module.visual_projection.weight,213 std=module.vision_embed_dim**-0.5 * self.config.initializer_factor,214 )215 elif isinstance(module, CLIPVisionModelWithProjection):216 nn.init.normal_(217 module.visual_projection.weight,218 std=self.config.hidden_size**-0.5 * self.config.initializer_factor,219 )220 elif isinstance(module, CLIPTextModelWithProjection):221 nn.init.normal_(222 module.text_projection.weight,223 std=self.config.hidden_size**-0.5 * self.config.initializer_factor,224 )225 226 if isinstance(module, nn.LayerNorm):227 module.bias.data.zero_()228 module.weight.data.fill_(1.0)229 if isinstance(module, nn.Linear) and module.bias is not None:230 module.bias.data.zero_()231 232 def _set_gradient_checkpointing(self, module, value=False):233 if isinstance(module, CLIPEncoder):234 module.gradient_checkpointing = value235 236 237CLIP_START_DOCSTRING = r"""238 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the239 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads240 etc.)241 242 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.243 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage244 and behavior.245 246 Parameters:247 config ([`CLIPConfig`]): Model configuration class with all the parameters of the model.248 Initializing with a config file does not load the weights associated with the model, only the249 configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.250"""251 252CLIP_TEXT_INPUTS_DOCSTRING = r"""253 Args:254 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):255 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide256 it.257 258 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and259 [`PreTrainedTokenizer.__call__`] for details.260 261 [What are input IDs?](../glossary#input-ids)262 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):263 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:264 265 - 1 for tokens that are **not masked**,266 - 0 for tokens that are **masked**.267 268 [What are attention masks?](../glossary#attention-mask)269 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):270 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,271 config.max_position_embeddings - 1]`.272 273 [What are position IDs?](../glossary#position-ids)274 output_attentions (`bool`, *optional*):275 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned276 tensors for more detail.277 output_hidden_states (`bool`, *optional*):278 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for279 more detail.280 return_dict (`bool`, *optional*):281 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.282"""283 284CLIP_VISION_INPUTS_DOCSTRING = r"""285 Args:286 pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):287 Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using288 [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details.289 output_attentions (`bool`, *optional*):290 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned291 tensors for more detail.292 output_hidden_states (`bool`, *optional*):293 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for294 more detail.295 return_dict (`bool`, *optional*):296 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.297"""298 299CLIP_INPUTS_DOCSTRING = r"""300 Args:301 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):302 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide303 it.304 305 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and306 [`PreTrainedTokenizer.__call__`] for details.307 308 [What are input IDs?](../glossary#input-ids)309 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):310 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:311 312 - 1 for tokens that are **not masked**,313 - 0 for tokens that are **masked**.314 315 [What are attention masks?](../glossary#attention-mask)316 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):317 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,318 config.max_position_embeddings - 1]`.319 320 [What are position IDs?](../glossary#position-ids)321 pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):322 Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using323 [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details.324 return_loss (`bool`, *optional*):325 Whether or not to return the contrastive loss.326 output_attentions (`bool`, *optional*):327 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned328 tensors for more detail.329 output_hidden_states (`bool`, *optional*):330 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for331 more detail.332 return_dict (`bool`, *optional*):333 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.334"""335 336 337class CLIPEncoder(nn.Module):338 """339 Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a340 [`CLIPEncoderLayer`].341 342 Args:343 config: CLIPConfig344 """345 346 def __init__(self, config: LanguageBindDepthConfig):347 super().__init__()348 self.config = config349 self.layers = nn.ModuleList([CLIPEncoderLayer(config) for _ in range(config.num_hidden_layers)])350 self.gradient_checkpointing = False351 352 def forward(353 self,354 inputs_embeds,355 attention_mask: Optional[torch.Tensor] = None,356 causal_attention_mask: Optional[torch.Tensor] = None,357 output_attentions: Optional[bool] = None,358 output_hidden_states: Optional[bool] = None,359 return_dict: Optional[bool] = None,360 ) -> Union[Tuple, BaseModelOutput]:361 r"""362 Args:363 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):364 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.365 This is useful if you want more control over how to convert `input_ids` indices into associated vectors366 than the model's internal embedding lookup matrix.367 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):368 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:369 370 - 1 for tokens that are **not masked**,371 - 0 for tokens that are **masked**.372 373 [What are attention masks?](../glossary#attention-mask)374 causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):375 Causal mask for the text model. Mask values selected in `[0, 1]`:376 377 - 1 for tokens that are **not masked**,378 - 0 for tokens that are **masked**.379 380 [What are attention masks?](../glossary#attention-mask)381 output_attentions (`bool`, *optional*):382 Whether or not to return the attentions tensors of all attention layers. See `attentions` under383 returned tensors for more detail.384 output_hidden_states (`bool`, *optional*):385 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors386 for more detail.387 return_dict (`bool`, *optional*):388 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.389 """390 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions391 output_hidden_states = (392 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states393 )394 return_dict = return_dict if return_dict is not None else self.config.use_return_dict395 396 encoder_states = () if output_hidden_states else None397 all_attentions = () if output_attentions else None398 399 hidden_states = inputs_embeds400 for idx, encoder_layer in enumerate(self.layers):401 if output_hidden_states:402 encoder_states = encoder_states + (hidden_states,)403 if self.gradient_checkpointing and self.training:404 405 def create_custom_forward(module):406 def custom_forward(*inputs):407 return module(*inputs, output_attentions)408 409 return custom_forward410 411 layer_outputs = torch.utils.checkpoint.checkpoint(412 create_custom_forward(encoder_layer),413 hidden_states,414 attention_mask,415 causal_attention_mask,416 )417 else:418 layer_outputs = encoder_layer(419 hidden_states,420 attention_mask,421 causal_attention_mask,422 output_attentions=output_attentions,423 )424 425 hidden_states = layer_outputs[0]426 427 if output_attentions:428 all_attentions = all_attentions + (layer_outputs[1],)429 430 if output_hidden_states:431 encoder_states = encoder_states + (hidden_states,)432 433 if not return_dict:434 return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)435 return BaseModelOutput(436 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions437 )438 439 440# Copied from transformers.models.bart.modeling_bart._make_causal_mask441def _make_causal_mask(442 input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0443):444 """445 Make causal mask used for bi-directional self-attention.446 """447 bsz, tgt_len = input_ids_shape448 mask = torch.full((tgt_len, tgt_len), torch.finfo(dtype).min, device=device)449 mask_cond = torch.arange(mask.size(-1), device=device)450 mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)451 mask = mask.to(dtype)452 453 if past_key_values_length > 0:454 mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)455 return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)456 457 458class CLIPTextTransformer(nn.Module):459 def __init__(self, config: CLIPTextConfig):460 super().__init__()461 self.config = config462 embed_dim = config.hidden_size463 self.embeddings = CLIPTextEmbeddings(config)464 self.encoder = CLIPEncoder(config)465 self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)466 467 @add_start_docstrings_to_model_forward(CLIP_TEXT_INPUTS_DOCSTRING)468 @replace_return_docstrings(output_type=BaseModelOutputWithPooling, config_class=CLIPTextConfig)469 def forward(470 self,471 input_ids: Optional[torch.Tensor] = None,472 attention_mask: Optional[torch.Tensor] = None,473 position_ids: Optional[torch.Tensor] = None,474 output_attentions: Optional[bool] = None,475 output_hidden_states: Optional[bool] = None,476 return_dict: Optional[bool] = None,477 ) -> Union[Tuple, BaseModelOutputWithPooling]:478 r"""479 Returns:480 481 """482 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions483 output_hidden_states = (484 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states485 )486 return_dict = return_dict if return_dict is not None else self.config.use_return_dict487 488 if input_ids is None:489 raise ValueError("You have to specify input_ids")490 491 input_shape = input_ids.size()492 input_ids = input_ids.view(-1, input_shape[-1])493 494 hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)495 496 # CLIP's text model uses causal mask, prepare it here.497 # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324498 causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)499 # expand attention_mask500 if attention_mask is not None:501 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]502 attention_mask = _expand_mask(attention_mask, hidden_states.dtype)503 504 encoder_outputs = self.encoder(505 inputs_embeds=hidden_states,506 attention_mask=attention_mask,507 causal_attention_mask=causal_attention_mask,508 output_attentions=output_attentions,509 output_hidden_states=output_hidden_states,510 return_dict=return_dict,511 )512 513 last_hidden_state = encoder_outputs[0]514 last_hidden_state = self.final_layer_norm(last_hidden_state)515 516 # text_embeds.shape = [batch_size, sequence_length, transformer.width]517 # take features from the eot embedding (eot_token is the highest number in each sequence)518 # casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14519 pooled_output = last_hidden_state[520 torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),521 input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),522 ]523 524 if not return_dict:525 return (last_hidden_state, pooled_output) + encoder_outputs[1:]526 527 return BaseModelOutputWithPooling(528 last_hidden_state=last_hidden_state,529 pooler_output=pooled_output,530 hidden_states=encoder_outputs.hidden_states,531 attentions=encoder_outputs.attentions,532 )533 534 535@add_start_docstrings(536 """The text model from CLIP without any head or projection on top.""",537 CLIP_START_DOCSTRING,538)539class CLIPTextModel(CLIPPreTrainedModel):540 config_class = CLIPTextConfig541 542 _no_split_modules = ["CLIPEncoderLayer"]543 544 def __init__(self, config: CLIPTextConfig):545 super().__init__(config)546 self.text_model = CLIPTextTransformer(config)547 # Initialize weights and apply final processing548 self.post_init()549 550 def get_input_embeddings(self) -> nn.Module:551 return self.text_model.embeddings.token_embedding552 553 def set_input_embeddings(self, value):554 self.text_model.embeddings.token_embedding = value555 556 @add_start_docstrings_to_model_forward(CLIP_TEXT_INPUTS_DOCSTRING)557 @replace_return_docstrings(output_type=BaseModelOutputWithPooling, config_class=CLIPTextConfig)558 def forward(559 self,560 input_ids: Optional[torch.Tensor] = None,561 attention_mask: Optional[torch.Tensor] = None,562 position_ids: Optional[torch.Tensor] = None,563 output_attentions: Optional[bool] = None,564 output_hidden_states: Optional[bool] = None,565 return_dict: Optional[bool] = None,566 ) -> Union[Tuple, BaseModelOutputWithPooling]:567 r"""568 Returns:569 570 Examples:571 572 ```python573 >>> from transformers import AutoTokenizer, CLIPTextModel574 575 >>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")576 >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")577 578 >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")579 580 >>> outputs = model(**inputs)581 >>> last_hidden_state = outputs.last_hidden_state582 >>> pooled_output = outputs.pooler_output # pooled (EOS token) states583 ```"""584 return_dict = return_dict if return_dict is not None else self.config.use_return_dict585 586 return self.text_model(587 input_ids=input_ids,588 attention_mask=attention_mask,589 position_ids=position_ids,590 output_attentions=output_attentions,591 output_hidden_states=output_hidden_states,592 return_dict=return_dict,593 )594 595 596class CLIPVisionTransformer(nn.Module):597 def __init__(self, config: CLIPVisionConfig):598 super().__init__()599 self.config = config600 embed_dim = config.hidden_size601 602 self.embeddings = CLIPVisionEmbeddings(config)603 self.patch_dropout = PatchDropout(config.force_patch_dropout)604 self.pre_layrnorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)605 self.encoder = CLIPEncoder(config)606 self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)607 608 @add_start_docstrings_to_model_forward(CLIP_VISION_INPUTS_DOCSTRING)609 @replace_return_docstrings(output_type=BaseModelOutputWithPooling, config_class=CLIPVisionConfig)610 def forward(611 self,612 pixel_values: Optional[torch.FloatTensor] = None,613 output_attentions: Optional[bool] = None,614 output_hidden_states: Optional[bool] = None,615 return_dict: Optional[bool] = None,616 ) -> Union[Tuple, BaseModelOutputWithPooling]:617 r"""618 Returns:619 620 """621 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions622 output_hidden_states = (623 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states624 )625 return_dict = return_dict if return_dict is not None else self.config.use_return_dict626 627 if pixel_values is None:628 raise ValueError("You have to specify pixel_values")629 ######################################630 if len(pixel_values.shape) == 7:631 b_new, pair_new, T, bs_new, channel_new, h_new, w_new = pixel_values.shape632 # print(pixel_values.shape)633 B = b_new * pair_new * bs_new634 pixel_values = pixel_values.reshape(B*T, channel_new, h_new, w_new)635 636 elif len(pixel_values.shape) == 5:637 B, _, T, _, _ = pixel_values.shape638 # print(pixel_values.shape)639 pixel_values = rearrange(pixel_values, 'b c t h w -> (b t) c h w')640 else:641 # print(pixel_values.shape)642 B, _, _, _ = pixel_values.shape643 T = 1644 ###########################645 hidden_states = self.embeddings(pixel_values)646 647 hidden_states = self.patch_dropout(hidden_states, B, T) ##############################################648 649 hidden_states = self.pre_layrnorm(hidden_states)650 651 encoder_outputs = self.encoder(652 inputs_embeds=hidden_states,653 output_attentions=output_attentions,654 output_hidden_states=output_hidden_states,655 return_dict=return_dict,656 )657 658 last_hidden_state = encoder_outputs[0]659 pooled_output = last_hidden_state[:, 0, :]660 pooled_output = self.post_layernorm(pooled_output)661 662 pooled_output = pooled_output.reshape(B, T, -1).mean(1) ################################663 664 if not return_dict:665 return (last_hidden_state, pooled_output) + encoder_outputs[1:]666 667 return BaseModelOutputWithPooling(668 last_hidden_state=last_hidden_state,669 pooler_output=pooled_output,670 hidden_states=encoder_outputs.hidden_states,671 attentions=encoder_outputs.attentions,672 )673 674 675@add_start_docstrings(676 """The vision model from CLIP without any head or projection on top.""",677 CLIP_START_DOCSTRING,678)679class CLIPVisionModel(CLIPPreTrainedModel):680 config_class = CLIPVisionConfig681 main_input_name = "pixel_values"682 683 def __init__(self, config: CLIPVisionConfig):684 super().__init__(config)685 self.vision_model = CLIPVisionTransformer(config)686 # Initialize weights and apply final processing687 self.post_init()688 689 def get_input_embeddings(self) -> nn.Module:690 return self.vision_model.embeddings.patch_embedding691 692 @add_start_docstrings_to_model_forward(CLIP_VISION_INPUTS_DOCSTRING)693 @replace_return_docstrings(output_type=BaseModelOutputWithPooling, config_class=CLIPVisionConfig)694 def forward(695 self,696 pixel_values: Optional[torch.FloatTensor] = None,697 output_attentions: Optional[bool] = None,698 output_hidden_states: Optional[bool] = None,699 return_dict: Optional[bool] = None,700 ) -> Union[Tuple, BaseModelOutputWithPooling]:701 r"""702 Returns:703 704 Examples:705 706 ```python707 >>> from PIL import Image708 >>> import requests709 >>> from transformers import AutoProcessor, CLIPVisionModel710 711 >>> model = CLIPVisionModel.from_pretrained("openai/clip-vit-base-patch32")712 >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")713 714 >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"715 >>> image = Image.open(requests.get(url, stream=True).raw)716 717 >>> inputs = processor(images=image, return_tensors="pt")718 719 >>> outputs = model(**inputs)720 >>> last_hidden_state = outputs.last_hidden_state721 >>> pooled_output = outputs.pooler_output # pooled CLS states722 ```"""723 return_dict = return_dict if return_dict is not None else self.config.use_return_dict724 725 return self.vision_model(726 pixel_values=pixel_values,727 output_attentions=output_attentions,728 output_hidden_states=output_hidden_states,729 return_dict=return_dict,730 )731 732 733@add_start_docstrings(CLIP_START_DOCSTRING)734class LanguageBindDepth(CLIPPreTrainedModel):735 config_class = LanguageBindDepthConfig736 737 def __init__(self, config: LanguageBindDepthConfig):738 super().__init__(config)739 740 if not isinstance(config.text_config, CLIPTextConfig):741 raise ValueError(742 "config.text_config is expected to be of type CLIPTextConfig but is of type"743 f" {type(config.text_config)}."744 )745 746 if not isinstance(config.vision_config, CLIPVisionConfig):747 raise ValueError(748 "config.vision_config is expected to be of type CLIPVisionConfig but is of type"749 f" {type(config.vision_config)}."750 )751 752 text_config = config.text_config753 vision_config = config.vision_config754 self.add_time_attn = vision_config.add_time_attn755 self.lora_r = vision_config.lora_r756 self.lora_alpha = vision_config.lora_alpha757 self.lora_dropout = vision_config.lora_dropout758 759 self.projection_dim = config.projection_dim760 self.text_embed_dim = text_config.hidden_size761 self.vision_embed_dim = vision_config.hidden_size762 763 self.text_model = CLIPTextTransformer(text_config)764 self.vision_model = CLIPVisionTransformer(vision_config)765 766 self.visual_projection = nn.Linear(self.vision_embed_dim, self.projection_dim, bias=False)767 self.text_projection = nn.Linear(self.text_embed_dim, self.projection_dim, bias=False)768 self.logit_scale = nn.Parameter(torch.tensor(self.config.logit_scale_init_value))769 770 # Initialize weights and apply final processing771 self.post_init()772 self.convert_to_lora()773 self.resize_pos(self.vision_model.embeddings, vision_config)774 775 def convert_to_lora(self):776 if self.lora_r == 0:777 return778 if self.add_time_attn:779 target_modules = ["temporal_attn.k_proj", "temporal_attn.v_proj",780 "temporal_attn.q_proj", "temporal_attn.out_proj",781 "temporal_mlp.fc1", "temporal_mlp.fc2"]782 else:783 target_modules = ["k_proj", "v_proj", "q_proj", "out_proj"]784 config = LoraConfig(785 r=self.lora_r, # 16786 lora_alpha=self.lora_alpha, # 16787 target_modules=target_modules, # self_attn.out_proj788 lora_dropout=self.lora_dropout, # 0.1789 bias="none",790 modules_to_save=[],791 )792 self.vision_model.encoder.is_gradient_checkpointing = False793 self.vision_model.encoder = get_peft_model(self.vision_model.encoder, config)794 795 def resize_pos(self, m, vision_config):796 # convert embedding797 if vision_config.num_mel_bins!=0 and vision_config.target_length!=0:798 m.image_size = [vision_config.num_mel_bins, vision_config.target_length]799 m.config.image_size = [m.image_size, m.image_size] if isinstance(m.image_size, int) else m.image_size800 # pos resize801 old_pos_embed_state_dict = m.position_embedding.state_dict()802 old_pos_embed = old_pos_embed_state_dict['weight']803 dtype = old_pos_embed.dtype804 grid_size = [m.config.image_size[0] // m.patch_size, m.config.image_size[1] // m.patch_size]805 extra_tokens = 1 # FIXME detect different token configs (ie no class token, or more)806 new_seq_len = grid_size[0] * grid_size[1] + extra_tokens807 if new_seq_len == old_pos_embed.shape[0]:808 # m.to(args.device)809 return810 811 m.num_patches = grid_size[0] * grid_size[1]812 m.num_positions = m.num_patches + 1813 m.register_buffer("position_ids", torch.arange(m.num_positions).expand((1, -1)))814 new_position_embedding = nn.Embedding(m.num_positions, m.embed_dim)815 816 if extra_tokens:817 pos_emb_tok, pos_emb_img = old_pos_embed[:extra_tokens], old_pos_embed[extra_tokens:]818 else:819 pos_emb_tok, pos_emb_img = None, old_pos_embed820 old_grid_size = [int(math.sqrt(len(pos_emb_img)))] * 2821 822 # if is_master(args):823 # logging.info('Resizing position embedding grid-size from %s to %s', old_grid_size, grid_size)824 pos_emb_img = pos_emb_img.reshape(1, old_grid_size[0], old_grid_size[1], -1).permute(0, 3, 1, 2)825 pos_emb_img = F.interpolate(826 pos_emb_img,827 size=grid_size,828 mode='bicubic',829 antialias=True,830 align_corners=False,831 )832 pos_emb_img = pos_emb_img.permute(0, 2, 3, 1).reshape(1, grid_size[0] * grid_size[1], -1)[0]833 if pos_emb_tok is not None:834 new_pos_embed = torch.cat([pos_emb_tok, pos_emb_img], dim=0)835 else:836 new_pos_embed = pos_emb_img837 old_pos_embed_state_dict['weight'] = new_pos_embed.to(dtype)838 m.position_embedding = new_position_embedding839 m.position_embedding.load_state_dict(old_pos_embed_state_dict)840 841 # m.to(args.device)842 843 @add_start_docstrings_to_model_forward(CLIP_TEXT_INPUTS_DOCSTRING)844 def get_text_features(845 self,846 input_ids: Optional[torch.Tensor] = None,847 attention_mask: Optional[torch.Tensor] = None,848 position_ids: Optional[torch.Tensor] = None,849 output_attentions: Optional[bool] = None,850 output_hidden_states: Optional[bool] = None,851 return_dict: Optional[bool] = None,852 ) -> torch.FloatTensor:853 r"""854 Returns:855 text_features (`torch.FloatTensor` of shape `(batch_size, output_dim`): The text embeddings obtained by856 applying the projection layer to the pooled output of [`CLIPTextModel`].857 858 Examples:859 860 ```python861 >>> from transformers import AutoTokenizer, CLIPModel862 863 >>> model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")864 >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")865 866 >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")867 >>> text_features = model.get_text_features(**inputs)868 ```"""869 # Use CLIP model's config for some fields (if specified) instead of those of vision & text components.870 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions871 output_hidden_states = (872 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states873 )874 return_dict = return_dict if return_dict is not None else self.config.use_return_dict875 876 text_outputs = self.text_model(877 input_ids=input_ids,878 attention_mask=attention_mask,879 position_ids=position_ids,880 output_attentions=output_attentions,881 output_hidden_states=output_hidden_states,882 return_dict=return_dict,883 )884 885 pooled_output = text_outputs[1]886 text_features = self.text_projection(pooled_output)887 888 return text_features889 890 @add_start_docstrings_to_model_forward(CLIP_VISION_INPUTS_DOCSTRING)891 def get_image_features(892 self,893 pixel_values: Optional[torch.FloatTensor] = None,894 output_attentions: Optional[bool] = None,895 output_hidden_states: Optional[bool] = None,896 return_dict: Optional[bool] = None,897 ) -> torch.FloatTensor:898 r"""899 Returns:900 image_features (`torch.FloatTensor` of shape `(batch_size, output_dim`): The image embeddings obtained by901 applying the projection layer to the pooled output of [`CLIPVisionModel`].902 903 Examples:904 905 ```python906 >>> from PIL import Image907 >>> import requests908 >>> from transformers import AutoProcessor, CLIPModel909 910 >>> model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")911 >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")912 913 >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"914 >>> image = Image.open(requests.get(url, stream=True).raw)915 916 >>> inputs = processor(images=image, return_tensors="pt")917 918 >>> image_features = model.get_image_features(**inputs)919 ```"""920 # Use CLIP model's config for some fields (if specified) instead of those of vision & text components.921 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions922 output_hidden_states = (923 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states924 )925 return_dict = return_dict if return_dict is not None else self.config.use_return_dict926 927 vision_outputs = self.vision_model(928 pixel_values=pixel_values,929 output_attentions=output_attentions,930 output_hidden_states=output_hidden_states,931 return_dict=return_dict,932 )933 934 pooled_output = vision_outputs[1] # pooled_output935 image_features = self.visual_projection(pooled_output)936 937 return image_features938 939 @add_start_docstrings_to_model_forward(CLIP_INPUTS_DOCSTRING)940 @replace_return_docstrings(output_type=CLIPOutput, config_class=LanguageBindDepthConfig)941 def forward(942 self,943 input_ids: Optional[torch.LongTensor] = None,944 pixel_values: Optional[torch.FloatTensor] = None,945 attention_mask: Optional[torch.Tensor] = None,946 position_ids: Optional[torch.LongTensor] = None,947 return_loss: Optional[bool] = None,948 output_attentions: Optional[bool] = None,949 output_hidden_states: Optional[bool] = None,950 return_dict: Optional[bool] = None,951 ) -> Union[Tuple, CLIPOutput]:952 r"""953 Returns:954 955 Examples:956 957 ```python958 >>> from PIL import Image959 >>> import requests960 >>> from transformers import AutoProcessor, CLIPModel961 962 >>> model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")963 >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")964 965 >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"966 >>> image = Image.open(requests.get(url, stream=True).raw)967 968 >>> inputs = processor(969 ... text=["a photo of a cat", "a photo of a dog"], images=image, return_tensors="pt", padding=True970 ... )971 972 >>> outputs = model(**inputs)973 >>> logits_per_image = outputs.logits_per_image # this is the image-text similarity score974 >>> probs = logits_per_image.softmax(dim=1) # we can take the softmax to get the label probabilities975 ```"""976 # Use CLIP model's config for some fields (if specified) instead of those of vision & text components.977 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions978 output_hidden_states = (979 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states980 )981 return_dict = return_dict if return_dict is not None else self.config.use_return_dict982 983 vision_outputs = self.vision_model(984 pixel_values=pixel_values,985 output_attentions=output_attentions,986 output_hidden_states=output_hidden_states,987 return_dict=return_dict,988 )989 990 text_outputs = self.text_model(991 input_ids=input_ids,992 attention_mask=attention_mask,993 position_ids=position_ids,994 output_attentions=output_attentions,995 output_hidden_states=output_hidden_states,996 return_dict=return_dict,997 )998 999 image_embeds = vision_outputs[1]1000 image_embeds = self.visual_projection(image_embeds)1001 1002 text_embeds = text_outputs[1]1003 text_embeds = self.text_projection(text_embeds)1004 1005 # normalized features1006 image_embeds = image_embeds / image_embeds.norm(p=2, dim=-1, keepdim=True)1007 text_embeds = text_embeds / text_embeds.norm(p=2, dim=-1, keepdim=True)1008 1009 # cosine similarity as logits1010 logit_scale = self.logit_scale.exp()1011 logits_per_text = torch.matmul(text_embeds, image_embeds.t()) * logit_scale1012 logits_per_image = logits_per_text.t()1013 1014 loss = None1015 if return_loss:1016 loss = clip_loss(logits_per_text)1017 1018 if not return_dict:1019 output = (logits_per_image, logits_per_text, text_embeds, image_embeds, text_outputs, vision_outputs)1020 return ((loss,) + output) if loss is not None else output1021 1022 return CLIPOutput(1023 loss=loss,1024 logits_per_image=logits_per_image,1025 logits_per_text=logits_per_text,1026 text_embeds=text_embeds,1027 image_embeds=image_embeds,1028 text_model_output=text_outputs,1029 vision_model_output=vision_outputs,1030 )