CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
modeling_x_clip.py1510 linesDownload Raw Back to x_clip
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()

Showing the first 1,200 of 1510 lines. Download the file for the rest.

Aluode/PerceptionLabPortable · CoolFace