CoolFace
Apppublic

XaviXva/Video-LLaVA

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
modeling_depth.py1030 linesDownload Raw Back to depth
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        )