mapo80/DeQA-Doc-Sharpness
157
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 