ngocson2002/vivqa-model
011
1from timm.models.layers import trunc_normal_ as __call_trunc_normal_2from torchscale.component.multiway_network import MutliwayEmbedding3from torchscale.component.embedding import PositionalEmbedding4from torchscale.architecture.encoder import Encoder5from transformers import PreTrainedModel6import torch.nn as nn7import torch.nn.functional as F8import torch9import math10from transformers import AutoModel11from transformers.utils.generic import ModelOutput12from dataclasses import dataclass13from typing import Optional14from efficientnet_pytorch import EfficientNet15from lavis.common.registry import registry16from .configuration_vivqa import ViVQAConfig17 18class BartPhoExtractor(nn.Module):19 def __init__(self):20 super(BartPhoExtractor, self).__init__()21 self.bartpho_word = AutoModel.from_pretrained("vinai/bartpho-word")22 23 def forward(self, input_ids, attention_mask):24 last_hidden_states = self.bartpho_word(input_ids, attention_mask)25 features = last_hidden_states[0]26 return features27 28class Blip2EfficientExtractor(nn.Module):29 def __init__(self):30 super(Blip2EfficientExtractor, self).__init__()31 self.device = "cuda" if torch.cuda.is_available() else "cpu"32 33 # BLIP-234 self.model_blip2 = registry.get_model_class(name="blip2_feature_extractor").from_pretrained(model_type="pretrain").to(self.device)35 if self.device == "cpu" or self.device == torch.device("cpu"):36 self.model_blip2 = self.model_blip2.float()37 self.model_blip2.eval()38 39 # Efficientnet40 self.model_efficientnet = EfficientNet.from_pretrained('efficientnet-b7', advprop=True).to(self.device)41 self.model_efficientnet.eval()42 self.pooling1 = nn.AdaptiveAvgPool2d((1, 32))43 self.pooling2 = nn.AdaptiveAvgPool2d((1, 768))44 45 def forward(self, images):46 47 global_features = self.model_blip2.extract_features(samples={"image": images}, mode="image").image_embeds48 49 local_features = self.model_efficientnet.extract_features(images)50 local_features = self.pooling1(local_features)51 local_features = local_features.permute(0, 3, 2, 1)52 local_features = self.pooling2(local_features)53 batch_size = images.shape[0]54 local_features = local_features.reshape(batch_size, local_features.shape[1], -1)55 56 v = torch.cat([global_features, local_features], dim=1)57 return v58 59@dataclass60class ViVQAOutput(ModelOutput):61 loss: Optional[torch.FloatTensor] = None62 logits: torch.FloatTensor = None63 64def trunc_normal_(tensor, mean=0., std=1.):65 __call_trunc_normal_(tensor, mean=mean, std=std, a=-std, b=std)66 67class Pooler(nn.Module):68 def __init__(self, input_features, output_features, norm_layer):69 super().__init__()70 self.norm = norm_layer(input_features)71 self.dense = nn.Linear(input_features, output_features)72 self.activation = nn.Tanh()73 74 def forward(self, x):75 cls_rep = x[:, 0, :]76 cls_rep = self.norm(cls_rep)77 pooled_output = self.dense(cls_rep)78 pooled_output = self.activation(pooled_output)79 return pooled_output80 81class ViVQABEiT3(PreTrainedModel):82 def __init__(self, args):83 super().__init__(args)84 assert args.multiway85 assert not args.share_encoder_input_output_embed86 87 self.text_embed = BartPhoExtractor()88 89 self.vision_embed = Blip2EfficientExtractor()90 for param in self.vision_embed.parameters():91 param.requires_grad = False92 93 94 self.linear = nn.Linear(1024, 768)95 96 # being consistent with Fairseq, which starts from 2 for position embedding97 num_position_embeddings = 6498 embed_positions = MutliwayEmbedding(99 modules=[100 PositionalEmbedding(num_position_embeddings + 2, args.encoder_embed_dim),101 PositionalEmbedding(args.max_source_positions, args.encoder_embed_dim),102 ],103 dim=1,104 )105 self.encoder = Encoder(106 args,107 embed_tokens=None,108 embed_positions=embed_positions,109 output_projection=None,110 is_encoder_decoder=False,111 )112 113 def forward(self, textual_tokens, visual_tokens, text_padding_position):114 x1 = self.vision_embed(visual_tokens)115 multiway_split_position = x1.size(1)116 117 x2 = self.text_embed(textual_tokens, 1-text_padding_position)118 x2 = self.linear(x2)119 120 x = torch.cat([x1, x2], dim=1)121 122 encoder_padding_mask = torch.cat(123 [124 torch.zeros(x1.shape[:-1]).to(x1.device).bool(),125 text_padding_position,126 ],127 dim=1,128 )129 130 encoder_out = self.encoder(131 src_tokens=None,132 encoder_padding_mask=encoder_padding_mask,133 token_embeddings=x,134 multiway_split_position=multiway_split_position135 )136 encoder_out["multiway_split_position"] = multiway_split_position137 return encoder_out138 139class BEiT3Wrapper(PreTrainedModel):140 def __init__(self, args, **kwargs):141 super().__init__(args)142 self.beit3 = ViVQABEiT3(args)143 # self.apply(self._init_weights)144 145 def fix_init_weight(self):146 def rescale(param, layer_id):147 param.div_(math.sqrt(2.0 * layer_id))148 149 for layer_id, layer in enumerate(self.blocks):150 rescale(layer.attn.proj.weight.data, layer_id + 1)151 rescale(layer.mlp.fc2.weight.data, layer_id + 1)152 153 def get_num_layers(self):154 return self.beit3.encoder.num_layers155 156 @torch.jit.ignore157 def no_weight_decay(self):158 return {'pos_embed', 'cls_token', 'beit3.encoder.embed_positions.A.weight', 'beit3.vision_embed.cls_token', 'logit_scale'}159 160 def _init_weights(self, m):161 if isinstance(m, nn.Linear):162 trunc_normal_(m.weight, std=.02)163 if isinstance(m, nn.Linear) and m.bias is not None:164 nn.init.constant_(m.bias, 0)165 elif isinstance(m, nn.LayerNorm):166 nn.init.constant_(m.bias, 0)167 nn.init.constant_(m.weight, 1.0)168 169 170class BEiT3ForVietnameseVisualQuestionAnswering(BEiT3Wrapper):171 config_class = ViVQAConfig172 def __init__(173 self, 174 args, 175 num_classes=353, 176 **kwargs177 ):178 super(BEiT3ForVietnameseVisualQuestionAnswering, self).__init__(args=args)179 embed_dim = args.encoder_embed_dim180 self.pooler = Pooler(181 input_features=embed_dim, 182 output_features=embed_dim, 183 norm_layer=nn.LayerNorm,184 )185 self.pooler.apply(self._init_weights)186 self.head = nn.Sequential(187 nn.Linear(embed_dim, embed_dim * 2),188 nn.LayerNorm(embed_dim * 2), 189 nn.GELU(),190 nn.Linear(embed_dim * 2, num_classes), 191 )192 self.head.apply(self._init_weights)193 194 def forward(self, image, question, padding_mask, labels=None, **kwargs): 195 outputs = self.beit3(196 textual_tokens=question, 197 visual_tokens=image, 198 text_padding_position=padding_mask, 199 )200 x = outputs["encoder_out"]201 cls_rep = self.pooler(x)202 logits = self.head(cls_rep)203 204 loss = None205 if labels is not None:206 loss = F.cross_entropy(logits, labels)207 208 return ViVQAOutput(209 loss=loss,210 logits=logits,211 )