sneedium/captcha_pixelplanet
1
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)