CoolFace
Apppublic

sneedium/captcha_pixelplanet

sourceHugging Facebsdupdated 4y agoView on Hugging Face
1likes
model_vision.py151 linesDownload Raw Back to modules
1import logging2import torch.nn as nn3from fastai.vision import *4 5from modules.attention import *6from modules.backbone import ResTranformer7from modules.model import Model8from modules.resnet import resnet459 10 11class BaseVision(Model):12    def __init__(self, config):13        super().__init__(config)14        self.loss_weight = ifnone(config.model_vision_loss_weight, 1.0)15        self.out_channels = ifnone(config.model_vision_d_model, 512)16 17        if config.model_vision_backbone == 'transformer':18            self.backbone = ResTranformer(config)19        else: self.backbone = resnet45()20        21        if config.model_vision_attention == 'position':22            mode = ifnone(config.model_vision_attention_mode, 'nearest')23            self.attention = PositionAttention(24                in_channels=self.out_channels,25                max_length=config.dataset_max_length + 1,  # additional stop token26                mode=mode,27            )28        elif config.model_vision_attention == 'attention':29            self.attention = Attention(30                in_channels=self.out_channels,31                max_length=config.dataset_max_length + 1,  # additional stop token32                n_feature=8*32,33            )34        else:35            raise Exception(f'{config.model_vision_attention} is not valid.')36        self.cls = nn.Linear(self.out_channels, self.charset.num_classes)37 38        if config.model_vision_checkpoint is not None:39            logging.info(f'Read vision model from {config.model_vision_checkpoint}.')40            self.load(config.model_vision_checkpoint)41 42    def _forward(self, b_features):43        attn_vecs, attn_scores = self.attention(b_features)  # (N, T, E), (N, T, H, W)44        logits = self.cls(attn_vecs) # (N, T, C)45        pt_lengths = self._get_length(logits)46 47        return {'feature': attn_vecs, 'logits': logits, 'pt_lengths': pt_lengths,48                'attn_scores': attn_scores, 'loss_weight':self.loss_weight, 'name': 'vision', 'b_features':b_features}49 50    def forward(self, images, *args, **kwargs):51        features = self.backbone(images, **kwargs)  # (N, E, H, W)52        return self._forward(features)53        54 55class BaseIterVision(BaseVision):56    def __init__(self, config):57        super().__init__(config)58        assert config.model_vision_backbone == 'transformer'59        self.iter_size = ifnone(config.model_vision_iter_size, 1)60        self.share_weights = ifnone(config.model_vision_share_weights, False)61        self.share_cnns = ifnone(config.model_vision_share_cnns, False)62        self.add_transformer = ifnone(config.model_vision_add_transformer, False)63        self.simple_trans = ifnone(config.model_vision_simple_trans, False)64        self.deep_supervision = ifnone(config.model_vision_deep_supervision, True)65        self.backbones = nn.ModuleList()66        self.trans = nn.ModuleList()67        for i in range(self.iter_size-1):68            B = None if self.share_weights else ResTranformer(config)69            if self.share_cnns:70                del B.resnet71            self.backbones.append(B)72            output_channel = self.out_channels73            if self.add_transformer:74                self.split_sizes = [output_channel]75            elif self.simple_trans:76                # self.split_sizes=[output_channel//16] + [0] * 577                # self.split_sizes= [output_channel//16, output_channel//16, output_channel//8, output_channel//4, output_channel//2] + [0]78                self.split_sizes= [output_channel//16, output_channel//16, 0, output_channel//4, output_channel//2, output_channel]79            else:80                self.split_sizes=[output_channel//16, output_channel//16, output_channel//8, output_channel//4, output_channel//2, output_channel]81            self.trans.append(nn.Conv2d(output_channel, sum(self.split_sizes), 1))82            torch.nn.init.zeros_(self.trans[-1].weight)83        84        if config.model_vision_checkpoint is not None:85            logging.info(f'Read vision model from {config.model_vision_checkpoint}.')86            self.load(config.model_vision_checkpoint)87        cb_init = ifnone(config.model_vision_cb_init, True)88        if cb_init:89            self.cb_init()90 91    def load(self, source, device=None, strict=False):92        state = torch.load(source, map_location=device)93        msg = self.load_state_dict(state['model'], strict=strict)94        print(msg)95 96    def cb_init(self):97        model_state_dict = self.backbone.state_dict()98 99        for m in self.backbones:100            if m:101                print('cb_init')102                msg = m.load_state_dict(model_state_dict, strict=False)103                print(msg)104 105    def forward_test(self, images, *args):106        l_feats = self.backbone.resnet(images)107        b_feats = self.backbone.forward_transformer(l_feats)108        cnt = len(self.backbones)109        if cnt == 0:110            v_res = super()._forward(b_feats)111        for B,T in zip(self.backbones, self.trans):112            cnt -= 1113            extra_feats = T(b_feats).split(self.split_sizes, dim=1)114            if self.share_weights:115                v_res = super().forward(images, extra_feats=extra_feats)116            else:117                if self.add_transformer:118                    if not self.share_cnns:119                        l_feats = B.resnet(images)120                    b_feats = B.forward_transformer(extra_feats[-1] + l_feats)121                else:122                    b_feats = B(images, extra_feats=extra_feats)123                v_res = super()._forward(b_feats) if cnt==0 else None124        return v_res125 126    def forward_train(self, images, *args):127        l_feats = self.backbone.resnet(images)128        b_feats = self.backbone.forward_transformer(l_feats)129        v_res = super()._forward(b_feats)130        # v_res = super().forward(images)131        all_v_res = [v_res]132        for B,T in zip(self.backbones, self.trans):133            extra_feats = T(v_res['b_features']).split(self.split_sizes, dim=1)134            if self.share_weights:135                v_res = super().forward(images, extra_feats=extra_feats)136            else:137                if self.add_transformer:138                    if not self.share_cnns:139                        l_feats = B.resnet(images)140                    b_feats = B.forward_transformer(extra_feats[-1] + l_feats)141                else:142                    b_feats = B(images, extra_feats=extra_feats)143                v_res = super()._forward(b_feats)144            all_v_res.append(v_res)145        return all_v_res146 147    def forward(self, images, *args):148        if self.training and self.deep_supervision:149            return self.forward_train(images, *args)150        else:151            return self.forward_test(images, *args)