Team Ai
Modelpublic

mapo80/DeQA-Doc-Sharpness

sourceHugging Faceapache-2.0updated 10mo agoView on Hugging Face
1likes57downloads
visual_encoder.py1019 linesDownload Raw Back to root
1import math2from typing import Any, Optional, Tuple, Union3 4from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling, BaseModelOutputWithPastAndCrossAttentions5from transformers.modeling_utils import PreTrainedModel6from transformers.pytorch_utils import find_pruneable_heads_and_indices, prune_linear_layer7 8import numpy as np9import torch10import torch.nn as nn11import torch.utils.checkpoint12# icecream removed for inference13 14def get_abs_pos(abs_pos, tgt_size):15    # abs_pos: L, C16    # tgt_size: M17    # return: M, C18    src_size = int(math.sqrt(abs_pos.size(0)))19    tgt_size = int(math.sqrt(tgt_size))20    dtype = abs_pos.dtype21 22    if src_size != tgt_size:23        return F.interpolate(24            abs_pos.float().reshape(1, src_size, src_size, -1).permute(0, 3, 1, 2),25            size=(tgt_size, tgt_size),26            mode="bicubic",27            align_corners=False,28        ).permute(0, 2, 3, 1).flatten(0, 2).to(dtype=dtype)29    else:30        return abs_pos31 32# https://github.com/facebookresearch/mae/blob/efb2a8062c206524e35e47d04501ed4f544c0ae8/util/pos_embed.py#L2033def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):34    """35    grid_size: int of the grid height and width36    return:37    pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)38    """39    grid_h = np.arange(grid_size, dtype=np.float32)40    grid_w = np.arange(grid_size, dtype=np.float32)41    grid = np.meshgrid(grid_w, grid_h)  # here w goes first42    grid = np.stack(grid, axis=0)43 44    grid = grid.reshape([2, 1, grid_size, grid_size])45    pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)46    if cls_token:47        pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)48    return pos_embed49 50 51def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):52    assert embed_dim % 2 == 053 54    # use half of dimensions to encode grid_h55    emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0])  # (H*W, D/2)56    emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1])  # (H*W, D/2)57 58    emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)59    return emb60 61 62def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):63    """64    embed_dim: output dimension for each position65    pos: a list of positions to be encoded: size (M,)66    out: (M, D)67    """68    assert embed_dim % 2 == 069    omega = np.arange(embed_dim // 2, dtype=np.float32)70    omega /= embed_dim / 2.71    omega = 1. / 10000**omega  # (D/2,)72 73    pos = pos.reshape(-1)  # (M,)74    out = np.einsum('m,d->md', pos, omega)  # (M, D/2), outer product75 76    emb_sin = np.sin(out) # (M, D/2)77    emb_cos = np.cos(out) # (M, D/2)78 79    emb = np.concatenate([emb_sin, emb_cos], axis=1)  # (M, D)80    return emb81 82 83 84import torch85import torch.nn as nn86import torch.nn.functional as F87 88class MplugOwlVisionEmbeddings(nn.Module):89    def __init__(self, config):90        super().__init__()91        self.config = config92        self.hidden_size = config.hidden_size93        self.image_size = config.image_size94        self.patch_size = config.patch_size95 96        self.cls_token = nn.Parameter(torch.randn(1, 1, self.hidden_size))97 98        self.patch_embed = nn.Conv2d(99            in_channels=3,100            out_channels=self.hidden_size,101            kernel_size=self.patch_size,102            stride=self.patch_size,103            bias=False,104        )105 106        # Initialize position embedding for default size (can be resized later)107        self.num_patches = (self.image_size // self.patch_size) ** 2108        self.position_embedding = nn.Parameter(torch.randn(1, self.num_patches + 1, self.hidden_size))109        self.pre_layernorm = nn.LayerNorm(self.hidden_size, eps=config.layer_norm_eps)110 111    def interpolate_pos_encoding(self, embeddings, h, w):112        """113        Interpolate position embeddings for different image sizes114        """115        npatch = embeddings.shape[1] - 1  # subtract 1 for cls token116        N = self.position_embedding.shape[1] - 1  # original number of patches117        118        if npatch == N:119            return self.position_embedding120        121        # Separate class token and patch embeddings122        class_pos_embed = self.position_embedding[:, 0:1]  # [1, 1, hidden_size]123        patch_pos_embed = self.position_embedding[:, 1:]   # [1, N, hidden_size]124        125        dim = embeddings.shape[-1]126        127        # Calculate original grid size128        w0 = h0 = int(N ** 0.5)129        130        # Reshape patch embeddings to 2D grid131        patch_pos_embed = patch_pos_embed.reshape(1, w0, h0, dim).permute(0, 3, 1, 2)132        133        # Convert to float32 for interpolation134        patch_pos_embed = patch_pos_embed.float()135        136        # Interpolate to new size137        patch_pos_embed = F.interpolate(138            patch_pos_embed,139            size=(h, w),140            mode='bicubic',141            align_corners=False,142        )143        144        # Convert back to original dtype145        patch_pos_embed = patch_pos_embed.to(dtype=embeddings.dtype)146        147        # Reshape back to sequence148        patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).reshape(1, -1, dim)149        150        # Concatenate class token and patch embeddings151        return torch.cat((class_pos_embed, patch_pos_embed), dim=1)152 153    def forward(self, pixel_values: torch.FloatTensor) -> torch.Tensor:154        batch_size = pixel_values.size(0)155        #print(f"[DEBUG] Input image shape: {pixel_values.shape}")156        157        image_embeds = self.patch_embed(pixel_values)158        #print(f"[DEBUG] After patch_embed shape: {image_embeds.shape}")159        160        # Get patch grid dimensions161        _, _, h, w = image_embeds.shape162        163        image_embeds = image_embeds.flatten(2).transpose(1, 2)164        #print(f"[DEBUG] After flatten and transpose shape: {image_embeds.shape}")165        166        class_embeds = self.cls_token.expand(batch_size, 1, -1).to(image_embeds.dtype)167        embeddings = torch.cat([class_embeds, image_embeds], dim=1)168        169        # Interpolate position embeddings to match current image size170        pos_embed = self.interpolate_pos_encoding(embeddings, h, w).to(image_embeds.dtype)171        #print(f"[DEBUG] Position embedding shape after interpolation: {pos_embed.shape}")172        173        embeddings = embeddings + pos_embed174        embeddings = self.pre_layernorm(embeddings)175        return embeddings176 177 178 179class MplugOwlVisionAttention(nn.Module):180    """Multi-headed attention from 'Attention Is All You Need' paper"""181 182    def __init__(self, config):183        super().__init__()184        self.config = config185        self.hidden_size = config.hidden_size186        self.num_heads = config.num_attention_heads187        self.head_dim = self.hidden_size // self.num_heads188        if self.head_dim * self.num_heads != self.hidden_size:189            raise ValueError(190                f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size} and `num_heads`:"191                f" {self.num_heads})."192            )193        self.scale = self.head_dim**-0.5194        self.dropout = nn.Dropout(config.attention_dropout)195 196        self.query_key_value = nn.Linear(self.hidden_size, 3 * self.hidden_size)197        self.dense = nn.Linear(self.hidden_size, self.hidden_size)198 199    def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):200        return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()201 202    def forward(203        self,204        hidden_states: torch.Tensor,205        head_mask: Optional[torch.Tensor] = None,206        output_attentions: Optional[bool] = False,207    ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:208        """Input shape: Batch x Time x Channel"""209 210        bsz, seq_len, embed_dim = hidden_states.size()211 212        mixed_qkv = self.query_key_value(hidden_states)213 214        mixed_qkv = mixed_qkv.reshape(bsz, seq_len, self.num_heads, 3, embed_dim // self.num_heads).permute(215            3, 0, 2, 1, 4216        )  # [3, b, np, sq, hn]217        query_states, key_states, value_states = (218            mixed_qkv[0],219            mixed_qkv[1],220            mixed_qkv[2],221        )222        # if self.config.use_flash_attn and flash_attn_func is not None:223        if False:224            # [b*sq, np, hn]225            query_states = query_states.permute(0, 2, 1, 3).contiguous()226            query_states = query_states.view(query_states.size(0) * query_states.size(1), query_states.size(2), -1)227 228            key_states = key_states.permute(0, 2, 1, 3).contiguous()229            key_states = key_states.view(key_states.size(0) * key_states.size(1), key_states.size(2), -1)230 231            value_states = value_states.permute(0, 2, 1, 3).contiguous()232            value_states = value_states.view(value_states.size(0) * value_states.size(1), value_states.size(2), -1)233 234            cu_seqlens = torch.arange(235                0, (bsz + 1) * seq_len, step=seq_len, dtype=torch.int32, device=query_states.device236            )237 238            context_layer = flash_attn_func(239                query_states,240                key_states,241                value_states,242                cu_seqlens,243                cu_seqlens,244                seq_len,245                seq_len,246                self.dropout if self.training else 0.0,247                softmax_scale=self.scale,248                causal=False,249                return_attn_probs=False,250            )251            # [b*sq, np, hn] => [b, sq, np, hn]252            context_layer = context_layer.view(bsz, seq_len, context_layer.size(1), context_layer.size(2))253        else:254            # Take the dot product between "query" and "key" to get the raw attention scores.255            attention_scores = torch.matmul(query_states, key_states.transpose(-1, -2))256 257            attention_scores = attention_scores * self.scale258 259            # Normalize the attention scores to probabilities.260            attention_probs = torch.softmax(attention_scores, dim=-1)261 262            # This is actually dropping out entire tokens to attend to, which might263            # seem a bit unusual, but is taken from the original Transformer paper.264            attention_probs = self.dropout(attention_probs)265 266            # Mask heads if we want to267            if head_mask is not None:268                attention_probs = attention_probs * head_mask269 270            context_layer = torch.matmul(attention_probs, value_states).permute(0, 2, 1, 3)271 272        new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size,)273        context_layer = context_layer.reshape(new_context_layer_shape)274 275        output = self.dense(context_layer)276 277        outputs = (output, attention_probs) if output_attentions else (output, None)278 279        return outputs280 281 282class QuickGELU(nn.Module):283    def forward(self, x: torch.Tensor):284        return x * torch.sigmoid(1.702 * x)285 286 287class MplugOwlMLP(nn.Module):288    def __init__(self, config):289        super().__init__()290        self.config = config291        self.activation_fn = QuickGELU()292        self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)293        self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)294 295    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:296        hidden_states = self.fc1(hidden_states)297        hidden_states = self.activation_fn(hidden_states)298        hidden_states = self.fc2(hidden_states)299        return hidden_states300 301 302class MplugOwlVisionEncoderLayer(nn.Module):303    def __init__(self, config):304        super().__init__()305        self.hidden_size = config.hidden_size306        self.self_attn = MplugOwlVisionAttention(config)307        self.input_layernorm = nn.LayerNorm(self.hidden_size, eps=config.layer_norm_eps)308        self.mlp = MplugOwlMLP(config)309        self.post_attention_layernorm = nn.LayerNorm(self.hidden_size, eps=config.layer_norm_eps)310 311    def forward(312        self,313        hidden_states: torch.Tensor,314        attention_mask: torch.Tensor,315        output_attentions: Optional[bool] = False,316    ) -> Tuple[torch.FloatTensor]:317        """318        Args:319            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`320            attention_mask (`torch.FloatTensor`): attention mask of size321                `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.322                `(config.encoder_attention_heads,)`.323            output_attentions (`bool`, *optional*):324                Whether or not to return the attentions tensors of all attention layers. See `attentions` under325                returned tensors for more detail.326        """327        residual = hidden_states328 329        hidden_states = self.input_layernorm(hidden_states)330        hidden_states, attn_weights = self.self_attn(331            hidden_states=hidden_states,332            head_mask=attention_mask,333            output_attentions=output_attentions,334        )335        hidden_states = hidden_states + residual336        residual = hidden_states337        hidden_states = self.post_attention_layernorm(hidden_states)338        hidden_states = self.mlp(hidden_states)339 340        hidden_states = hidden_states + residual341 342        outputs = (hidden_states,)343 344        if output_attentions:345            outputs += (attn_weights,)346 347        return outputs348    349    350class MplugOwlVisionEncoder(nn.Module):351    """352    Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a353    [`MplugOwlVisionEncoderLayer`].354 355    Args:356        config (`MplugOwlVisionConfig`):357            The corresponding vision configuration for the `MplugOwlEncoder`.358    """359 360    def __init__(self, config):361        super().__init__()362        self.config = config363        self.layers = nn.ModuleList([MplugOwlVisionEncoderLayer(config) for _ in range(config.num_hidden_layers)])364        self.gradient_checkpointing = True365 366    def forward(367        self,368        inputs_embeds,369        attention_mask: Optional[torch.Tensor] = None,370        output_attentions: Optional[bool] = None,371        output_hidden_states: Optional[bool] = None,372        return_dict: Optional[bool] = None,373    ) -> Union[Tuple, BaseModelOutput]:374        r"""375        Args:376            inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):377                Embedded representation of the inputs. Should be float, not int tokens.378            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):379                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:380 381                - 1 for tokens that are **not masked**,382                - 0 for tokens that are **masked**.383 384                [What are attention masks?](../glossary#attention-mask)385            output_attentions (`bool`, *optional*):386                Whether or not to return the attentions tensors of all attention layers. See `attentions` under387                returned tensors for more detail.388            output_hidden_states (`bool`, *optional*):389                Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors390                for more detail.391            return_dict (`bool`, *optional*):392                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.393        """394        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions395        output_hidden_states = (396            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states397        )398        return_dict = return_dict if return_dict is not None else self.config.use_return_dict399 400        encoder_states = () if output_hidden_states else None401        all_attentions = () if output_attentions else None402 403        hidden_states = inputs_embeds404        for idx, encoder_layer in enumerate(self.layers):405            if output_hidden_states:406                encoder_states = encoder_states + (hidden_states,)407            if self.gradient_checkpointing and self.training:408 409                def create_custom_forward(module):410                    def custom_forward(*inputs):411                        return module(*inputs, output_attentions)412 413                    return custom_forward414 415                layer_outputs = torch.utils.checkpoint.checkpoint(416                    create_custom_forward(encoder_layer),417                    hidden_states,418                    attention_mask,419                )420            else:421                layer_outputs = encoder_layer(422                    hidden_states,423                    attention_mask,424                    output_attentions=output_attentions,425                )426 427            hidden_states = layer_outputs[0]428 429            if output_attentions:430                all_attentions = all_attentions + (layer_outputs[1],)431 432        if output_hidden_states:433            encoder_states = encoder_states + (hidden_states,)434 435        if not return_dict:436            return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)437        return BaseModelOutput(438            last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions439        )440 441 442class MplugOwlVisionModel(PreTrainedModel):443    main_input_name = "pixel_values"444    _no_split_modules = ["MplugOwlVisionEncoderLayer"]445 446    def __init__(self, config):447        super().__init__(config)448        self.config = config449        self.hidden_size = config.hidden_size450 451        self.embeddings = MplugOwlVisionEmbeddings(config)452        self.encoder = MplugOwlVisionEncoder(config)453        self.post_layernorm = nn.LayerNorm(self.hidden_size, eps=config.layer_norm_eps)454 455        self.post_init()456 457 458    def forward(459        self,460        pixel_values: Optional[torch.FloatTensor] = None,461        output_attentions: Optional[bool] = None,462        output_hidden_states: Optional[bool] = None,463        return_dict: Optional[bool] = None,464    ) -> Union[Tuple, BaseModelOutputWithPooling]:465        r"""466        Returns:467 468        """469        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions470        output_hidden_states = (471            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states472        )473        return_dict = return_dict if return_dict is not None else self.config.use_return_dict474 475        if pixel_values is None:476            raise ValueError("You have to specify pixel_values")477 478        hidden_states = self.embeddings(pixel_values)479 480        encoder_outputs = self.encoder(481            inputs_embeds=hidden_states,482            output_attentions=output_attentions,483            output_hidden_states=output_hidden_states,484            return_dict=return_dict,485        )486 487        last_hidden_state = encoder_outputs[0]488        last_hidden_state = self.post_layernorm(last_hidden_state)489 490        pooled_output = last_hidden_state[:, 0, :]491        pooled_output = self.post_layernorm(pooled_output)492 493        if not return_dict:494            return (last_hidden_state, pooled_output) + encoder_outputs[1:]495 496        return BaseModelOutputWithPooling(497            last_hidden_state=last_hidden_state,498            pooler_output=pooled_output,499            hidden_states=encoder_outputs.hidden_states,500            attentions=encoder_outputs.attentions,501        )502 503    def get_input_embeddings(self):504        return self.embeddings505 506 507class MplugOwlVisualAbstractorMLP(nn.Module):508    def __init__(self, config):509        super().__init__()510        self.config = config511        in_features = config.hidden_size512        self.act = nn.SiLU()513 514        self.w1 = nn.Linear(in_features, config.intermediate_size)515        self.w2 = nn.Linear(config.intermediate_size, in_features)516        self.w3 = nn.Linear(in_features, config.intermediate_size)517        self.ffn_ln = nn.LayerNorm(config.intermediate_size, eps=config.layer_norm_eps)518 519    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:520        hidden_states = self.act(self.w1(hidden_states)) * self.w3(hidden_states)521        hidden_states = self.ffn_ln(hidden_states)522        hidden_states = self.w2(hidden_states)523        return hidden_states524 525 526class MplugOwlVisualAbstractorMultiHeadAttention(nn.Module):527    def __init__(self, config):528        super().__init__()529        self.config = config530        if config.hidden_size % config.num_attention_heads != 0:531            raise ValueError(532                "The hidden size (%d) is not a multiple of the number of attention heads (%d)"533                % (config.hidden_size, config.num_attention_heads)534            )535 536        self.num_attention_heads = config.num_attention_heads537        self.attention_head_size = int(config.hidden_size / config.num_attention_heads)538        self.all_head_size = self.num_attention_heads * self.attention_head_size539 540        self.query = nn.Linear(config.hidden_size, self.all_head_size)541        self.key = nn.Linear(config.encoder_hidden_size, self.all_head_size)542        self.value = nn.Linear(config.encoder_hidden_size, self.all_head_size)543 544        self.dropout = nn.Dropout(config.attention_probs_dropout_prob)545        self.save_attention = False546        547#         self.q_pos_embed = nn.Parameter(548#             torch.from_numpy(get_1d_sincos_pos_embed_from_grid(config.hidden_size, np.arange(config.num_learnable_queries, dtype=np.float32))).float()549#         ).requires_grad_(False)550#         grids = config.grid_size551#         self.k_pos_embed = nn.Parameter(552#             torch.from_numpy(get_2d_sincos_pos_embed(config.hidden_size, grids, cls_token=True)).float()553#         ).requires_grad_(False)554        grids = config.grid_size555        self.register_buffer(556            'q_pos_embed', 557            torch.from_numpy(get_1d_sincos_pos_embed_from_grid(config.hidden_size, np.arange(config.num_learnable_queries, dtype=np.float32))).float()558        )559        self.register_buffer(560            'k_pos_embed', 561            torch.from_numpy(get_2d_sincos_pos_embed(config.hidden_size, grids, cls_token=True)).float()562        )563        564 565    def save_attn_gradients(self, attn_gradients):566        self.attn_gradients = attn_gradients567 568    def get_attn_gradients(self):569        return self.attn_gradients570 571    def save_attention_map(self, attention_map):572        self.attention_map = attention_map573 574    def get_attention_map(self):575        return self.attention_map576 577    def transpose_for_scores(self, x):578        new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)579        x = x.view(*new_x_shape)580        return x.permute(0, 2, 1, 3)581 582    def forward(583        self,584        hidden_states,585        attention_mask=None,586        head_mask=None,587        encoder_hidden_states=None,588        encoder_attention_mask=None,589        past_key_value=None,590        output_attentions=False,591    ):592        # If this is instantiated as a cross-attention module, the keys593        # and values come from an encoder; the attention mask needs to be594        # such that the encoder's padding tokens are not attended to.595        596        # ็กฎไฟไฝ็ฝฎ็ผ–็ ็š„็ปดๅบฆไธŽ่พ“ๅ…ฅๅŒน้…597        if encoder_hidden_states is not None:598            seq_len = encoder_hidden_states.size(1)599            if seq_len != self.k_pos_embed.size(0):600                # ๅฆ‚ๆžœๅบๅˆ—้•ฟๅบฆไธๅŒน้…๏ผŒ้œ€่ฆ่ฐƒๆ•ดไฝ็ฝฎ็ผ–็ 601                # ไฝฟ็”จๆ›ด้ซ˜ๆ•ˆ็š„ๆ–นๅผ่ฐƒๆ•ดไฝ็ฝฎ็ผ–็ 602                k_pos_embed = self.k_pos_embed603                if seq_len > k_pos_embed.size(0):604                    # ๅฆ‚ๆžœ็›ฎๆ ‡ๅบๅˆ—ๆ›ด้•ฟ๏ผŒไฝฟ็”จ้‡ๅค605                    repeat_times = (seq_len + k_pos_embed.size(0) - 1) // k_pos_embed.size(0)606                    k_pos_embed = k_pos_embed.repeat(repeat_times, 1)[:seq_len]607                else:608                    # ๅฆ‚ๆžœ็›ฎๆ ‡ๅบๅˆ—ๆ›ด็Ÿญ๏ผŒไฝฟ็”จๅˆ‡็‰‡609                    k_pos_embed = k_pos_embed[:seq_len]610            else:611                k_pos_embed = self.k_pos_embed612                613            # ็กฎไฟ q_pos_embed ๅ’Œ k_pos_embed ็š„็ปดๅบฆๆญฃ็กฎ614            q_pos_embed = self.q_pos_embed.to(dtype=hidden_states.dtype)615            k_pos_embed = k_pos_embed.to(dtype=encoder_hidden_states.dtype)616            617            # ็กฎไฟ็ปดๅบฆๅŒน้…618            if q_pos_embed.size(0) + k_pos_embed.size(0) != encoder_hidden_states.size(1):619                # ๅฆ‚ๆžœ็ปดๅบฆไธๅŒน้…๏ผŒ่ฐƒๆ•ด k_pos_embed ็š„ๅคงๅฐ620                target_size = encoder_hidden_states.size(1) - q_pos_embed.size(0)621                if target_size > k_pos_embed.size(0):622                    # ๅฆ‚ๆžœ็›ฎๆ ‡ๅคงๅฐๆ›ดๅคง๏ผŒไฝฟ็”จ้‡ๅค623                    repeat_times = (target_size + k_pos_embed.size(0) - 1) // k_pos_embed.size(0)624                    k_pos_embed = k_pos_embed.repeat(repeat_times, 1)[:target_size]625                else:626                    # ๅฆ‚ๆžœ็›ฎๆ ‡ๅคงๅฐๆ›ดๅฐ๏ผŒไฝฟ็”จๅˆ‡็‰‡627                    k_pos_embed = k_pos_embed[:target_size]628            629            qk_pos_embed = torch.cat([q_pos_embed, k_pos_embed], dim=0).unsqueeze(0)630        else:631            qk_pos_embed = self.q_pos_embed.unsqueeze(0).to(dtype=hidden_states.dtype)632        633        # ็กฎไฟๆœ€็ปˆ็ปดๅบฆๅŒน้…634        assert qk_pos_embed.size(1) == encoder_hidden_states.size(1), \635            f"Position embedding size {qk_pos_embed.size(1)} does not match encoder hidden states size {encoder_hidden_states.size(1)}"636        637        key_layer = self.transpose_for_scores(self.key(encoder_hidden_states + qk_pos_embed))638        value_layer = self.transpose_for_scores(self.value(encoder_hidden_states))639        attention_mask = encoder_attention_mask640 641        mixed_query_layer = self.query(hidden_states + self.q_pos_embed.unsqueeze(0).to(dtype=hidden_states.dtype))642 643        query_layer = self.transpose_for_scores(mixed_query_layer)644 645        past_key_value = (key_layer, value_layer)646 647        # Take the dot product between "query" and "key" to get the raw attention scores.648        attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))649 650        attention_scores = attention_scores / math.sqrt(self.attention_head_size)651 652        if attention_mask is not None:653            # Apply the attention mask is (precomputed for all layers in BertModel forward() function)654            attention_scores = attention_scores + attention_mask655 656        # Normalize the attention scores to probabilities.657        attention_probs = nn.Softmax(dim=-1)(attention_scores)658 659        if self.save_attention:660            self.save_attention_map(attention_probs)661            attention_probs.register_hook(self.save_attn_gradients)662 663        # This is actually dropping out entire tokens to attend to, which might664        # seem a bit unusual, but is taken from the original Transformer paper.665        attention_probs_dropped = self.dropout(attention_probs)666 667        # Mask heads if we want to668        if head_mask is not None:669            attention_probs_dropped = attention_probs_dropped * head_mask670 671        context_layer = torch.matmul(attention_probs_dropped, value_layer)672 673        context_layer = context_layer.permute(0, 2, 1, 3).contiguous()674        new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)675        context_layer = context_layer.view(*new_context_layer_shape)676 677        outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)678 679        outputs = outputs + (past_key_value,)680        return outputs681 682 683class MplugOwlVisualAbstractorCrossOutput(nn.Module):684    def __init__(self, config):685        super().__init__()686        dim = config.hidden_size687        self.out_proj = nn.Linear(dim, dim, bias=True)688        self.norm2 = nn.LayerNorm(dim)689        self.mlp = MplugOwlVisualAbstractorMLP(config)690 691    def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:692        input_tensor = input_tensor + self.out_proj(hidden_states)693        input_tensor = input_tensor + self.mlp(self.norm2(input_tensor))694        return input_tensor695 696 697class MplugOwlVisualAbstractorAttention(nn.Module):698    def __init__(self, config):699        super().__init__()700        self.attention = MplugOwlVisualAbstractorMultiHeadAttention(config)701        self.output = MplugOwlVisualAbstractorCrossOutput(config)702        self.pruned_heads = set()703        self.norm1 = nn.LayerNorm(config.hidden_size)704        self.normk = nn.LayerNorm(config.hidden_size)705 706    def prune_heads(self, heads):707        if len(heads) == 0:708            return709        heads, index = find_pruneable_heads_and_indices(710            heads, self.attention.num_attention_heads, self.attention.attention_head_size, self.pruned_heads711        )712 713        # Prune linear layers714        self.attention.query = prune_linear_layer(self.attention.query, index)715        self.attention.key = prune_linear_layer(self.attention.key, index)716        self.attention.value = prune_linear_layer(self.attention.value, index)717        self.output.dense = prune_linear_layer(self.output.out_proj, index, dim=1)718 719        # Update hyper params and store pruned heads720        self.attention.num_attention_heads = self.attention.num_attention_heads - len(heads)721        self.attention.all_head_size = self.attention.attention_head_size * self.attention.num_attention_heads722        self.pruned_heads = self.pruned_heads.union(heads)723 724    def forward(725        self,726        hidden_states: torch.Tensor,727        attention_mask: Optional[torch.FloatTensor] = None,728        head_mask: Optional[torch.FloatTensor] = None,729        encoder_hidden_states: Optional[torch.FloatTensor] = None,730        encoder_attention_mask: Optional[torch.FloatTensor] = None,731        past_key_value: Optional[Tuple[Tuple[torch.FloatTensor]]] = None,732        output_attentions: Optional[bool] = False,733    ) -> Tuple[torch.Tensor]:734        # HACK we apply norm on q and k735        hidden_states = self.norm1(hidden_states)736        encoder_hidden_states = self.normk(encoder_hidden_states)737        encoder_hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)738        encoder_attention_mask = torch.cat([attention_mask, encoder_attention_mask], dim=-1)739        self_outputs = self.attention(740            hidden_states,741            attention_mask,742            head_mask,743            encoder_hidden_states,744            encoder_attention_mask,745            past_key_value,746            output_attentions,747        )748        attention_output = self.output(self_outputs[0], hidden_states)749        # add attentions if we output them750        outputs = (attention_output,) + self_outputs[1:]751        return outputs752 753 754class MplugOwlVisualAbstractorLayer(nn.Module):755    def __init__(self, config, layer_idx):756        super().__init__()757        self.chunk_size_feed_forward = config.chunk_size_feed_forward758        self.seq_len_dim = 1759 760        self.layer_idx = layer_idx761 762        self.crossattention = MplugOwlVisualAbstractorAttention(config)763        self.has_cross_attention = True764 765    def forward(766        self,767        hidden_states,768        attention_mask=None,769        head_mask=None,770        encoder_hidden_states=None,771        encoder_attention_mask=None,772        output_attentions=False,773    ):774        if encoder_hidden_states is None:775            raise ValueError("encoder_hidden_states must be given for cross-attention layers")776        cross_attention_outputs = self.crossattention(777            hidden_states,778            attention_mask,779            head_mask,780            encoder_hidden_states,781            encoder_attention_mask,782            output_attentions=output_attentions,783        )784        query_attention_output = cross_attention_outputs[0]785 786        outputs = (query_attention_output,)787        return outputs788 789 790class MplugOwlVisualAbstractorEncoder(nn.Module):791    def __init__(self, config):792        super().__init__()793        self.config = config794        self.layers = nn.ModuleList(795            [MplugOwlVisualAbstractorLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]796        )797        self.gradient_checkpointing = True798 799    def forward(800        self,801        hidden_states,802        attention_mask=None,803        head_mask=None,804        encoder_hidden_states=None,805        encoder_attention_mask=None,806        past_key_values=None,807        output_attentions=False,808        output_hidden_states=False,809        return_dict=True,810    ):811        all_hidden_states = () if output_hidden_states else None812 813        for i in range(self.config.num_hidden_layers):814            layer_module = self.layers[i]815            if output_hidden_states:816                all_hidden_states = all_hidden_states + (hidden_states,)817 818            layer_head_mask = head_mask[i] if head_mask is not None else None819            past_key_value = past_key_values[i] if past_key_values is not None else None820 821            if getattr(self.config, "gradient_checkpointing", False) and self.training:822 823                def create_custom_forward(module):824                    def custom_forward(*inputs):825                        return module(*inputs, past_key_value, output_attentions)826 827                    return custom_forward828 829                layer_outputs = torch.utils.checkpoint.checkpoint(830                    create_custom_forward(layer_module),831                    hidden_states,832                    attention_mask,833                    layer_head_mask,834                    encoder_hidden_states,835                    encoder_attention_mask,836                )837            else:838                layer_outputs = layer_module(839                    hidden_states,840                    attention_mask,841                    layer_head_mask,842                    encoder_hidden_states,843                    encoder_attention_mask,844                    output_attentions,845                )846 847            hidden_states = layer_outputs[0]848 849        return BaseModelOutput(850            last_hidden_state=hidden_states,851        )852 853 854class MplugOwlVisualAbstractorModel(PreTrainedModel):855    _no_split_modules = ["MplugOwlVisualAbstractorLayer"]856    def __init__(self, config, language_hidden_size):857        super().__init__(config)858        self.config = config859 860        self.encoder = MplugOwlVisualAbstractorEncoder(config)861        self.visual_fc = torch.nn.Linear(config.hidden_size, language_hidden_size)862        self.query_embeds = torch.nn.Parameter(torch.randn(1, config.num_learnable_queries, config.hidden_size))863        self.vit_eos = torch.nn.Parameter(torch.randn(1, 1, language_hidden_size))864 865        self.post_init()866 867    def _prune_heads(self, heads_to_prune):868        """869        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base870        class PreTrainedModel871        """872        for layer, heads in heads_to_prune.items():873            self.encoder.layer[layer].attention.prune_heads(heads)874 875    def get_extended_attention_mask(876        self,877        attention_mask: torch.Tensor,878        input_shape: Tuple[int],879        device: torch.device,880    ) -> torch.Tensor:881        """882        Makes broadcastable attention and causal masks so that future and masked tokens are ignored.883 884        Arguments:885            attention_mask (`torch.Tensor`):886                Mask with ones indicating tokens to attend to, zeros for tokens to ignore.887            input_shape (`Tuple[int]`):888                The shape of the input to the model.889            device: (`torch.device`):890                The device of the input to the model.891 892        Returns:893            `torch.Tensor` The extended attention mask, with a the same dtype as `attention_mask.dtype`.894        """895        # We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]896        # ourselves in which case we just need to make it broadcastable to all heads.897        if attention_mask.dim() == 3:898            extended_attention_mask = attention_mask[:, None, :, :]899        elif attention_mask.dim() == 2:900            # Provided a padding mask of dimensions [batch_size, seq_length]901            # - the model is an encoder, so make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]902            extended_attention_mask = attention_mask[:, None, None, :]903        else:904            raise ValueError(905                "Wrong shape for input_ids (shape {}) or attention_mask (shape {})".format(906                    input_shape, attention_mask.shape907                )908            )909 910        # Since attention_mask is 1.0 for positions we want to attend and 0.0 for911        # masked positions, this operation will create a tensor which is 0.0 for912        # positions we want to attend and -10000.0 for masked positions.913        # Since we are adding it to the raw scores before the softmax, this is914        # effectively the same as removing these entirely.915        extended_attention_mask = extended_attention_mask.to(dtype=self.dtype)  # fp16 compatibility916        extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0917        return extended_attention_mask918 919    def forward(920        self,921        attention_mask=None,922        head_mask=None,923        encoder_hidden_states=None,924        encoder_attention_mask=None,925        past_key_values=None,926        output_attentions=None,927        output_hidden_states=None,928        return_dict=None,929    ):930        r"""931        encoder_hidden_states  (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, `optional`):932            Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention if933            the model is configured as a decoder.934        encoder_attention_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, `optional`):935            Mask to avoid performing attention on the padding token indices of the encoder input. This mask is used in936            the cross-attention if the model is configured as a decoder. Mask values selected in `[0, 1]`:937            - 1 for tokens that are **not masked**,938            - 0 for tokens that are **masked**.939        past_key_values (`tuple(tuple(torch.FloatTensor))` of length `config.n_layers` with each tuple having 4 tensors of:940            shape `(batch_size, num_heads, sequence_length - 1, embed_size_per_head)`): Contains precomputed key and941            value hidden states of the attention blocks. Can be used to speed up decoding. If `past_key_values` are942            used, the user can optionally input only the last `decoder_input_ids` (those that don't have their past key943            value states given to this model) of shape `(batch_size, 1)` instead of all `decoder_input_ids` of shape944            `(batch_size, sequence_length)`.945        """946        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions947        output_hidden_states = (948            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states949        )950        return_dict = return_dict if return_dict is not None else self.config.use_return_dict951        952        query_embeds = self.query_embeds.repeat(encoder_hidden_states.shape[0], 1, 1)953        embedding_output = query_embeds954        input_shape = embedding_output.size()[:-1]955        batch_size, seq_length = input_shape956        device = embedding_output.device957 958        # We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]959        # ourselves in which case we just need to make it broadcastable to all heads.960        if attention_mask is None:961            attention_mask = torch.ones(962                (query_embeds.shape[0], query_embeds.shape[1]), dtype=torch.long, device=query_embeds.device963            )964        extended_attention_mask = self.get_extended_attention_mask(attention_mask, input_shape, device)965 966        # If a 2D or 3D attention mask is provided for the cross-attention967        # we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]968        if encoder_hidden_states is not None:969            if type(encoder_hidden_states) == list:970                encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states[0].size()971            else:972                (973                    encoder_batch_size,974                    encoder_sequence_length,975                    _,976                ) = encoder_hidden_states.size()977            encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)978 979            if type(encoder_attention_mask) == list:980                encoder_extended_attention_mask = [self.invert_attention_mask(mask) for mask in encoder_attention_mask]981            elif encoder_attention_mask is None:982                encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)983                encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)984            else:985                encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)986        else:987            encoder_extended_attention_mask = None988 989        # Prepare head mask if needed990        # 1.0 in head_mask indicate we keep the head991        # attention_probs has shape bsz x n_heads x N x N992        # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]993        # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]994        head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)995 996        encoder_outputs = self.encoder(997            embedding_output,998            attention_mask=extended_attention_mask,999            head_mask=head_mask,1000            encoder_hidden_states=encoder_hidden_states,1001            encoder_attention_mask=encoder_extended_attention_mask,1002            past_key_values=past_key_values,1003            output_attentions=output_attentions,1004            output_hidden_states=output_hidden_states,1005            return_dict=return_dict,1006        )1007        sequence_output = encoder_outputs[0]1008        pooled_output = sequence_output[:, 0, :]1009 1010        sequence_output = self.visual_fc(sequence_output)1011        sequence_output = torch.cat([sequence_output, self.vit_eos.repeat(sequence_output.shape[0], 1, 1)], dim=1)1012 1013        return BaseModelOutputWithPooling(1014            last_hidden_state=sequence_output,1015            pooler_output=pooled_output,1016            hidden_states=encoder_outputs.hidden_states,1017        )1018    1019