CoolFace
Modelpublic

ngocson2002/vivqa-model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes11downloads
modeling_vivqa.py211 linesDownload Raw Back to root
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        )