| 21 | device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| 22 | |
| 23 | class Model(nn.Module): |
| 24 | |
| 25 | def __init__(self, opt): |
| 26 | super(Model, self).__init__() |
| 27 | self.opt = opt |
| 28 | self.stages = {'Trans': opt.Transformation, 'Feat': opt.FeatureExtraction, |
| 29 | 'Seq': opt.SequenceModeling, 'Pred': opt.Prediction} |
| 30 | |
| 31 | """ Transformation """ |
| 32 | if opt.Transformation == 'TPS17': |
| 33 | self.Transformation = TPS_SpatialTransformerNetwork( |
| 34 | F=opt.num_fiducial, I_size=(opt.imgH, opt.imgW), I_r_size=(opt.imgH, opt.imgW), I_channel_num=opt.input_channel) |
| 35 | |
| 36 | elif opt.Transformation == 'TPS19': |
| 37 | self.tps = TPSSpatialTransformer(output_image_size=[opt.imgH, opt.imgW], |
| 38 | num_control_points=opt.num_fiducial, |
| 39 | margins=[0.05, 0.05]) |
| 40 | self.stn_head = STNHead(in_planes=3, num_ctrlpoints=opt.num_fiducial, activation=None) |
| 41 | else: |
| 42 | print('No Transformation module specified') |
| 43 | |
| 44 | """ FeatureExtraction """ |
| 45 | if opt.FeatureExtraction == 'VGG': |
| 46 | self.FeatureExtraction = VGG_FeatureExtractor(opt.input_channel, opt.output_channel) |
| 47 | elif opt.FeatureExtraction == 'ResNet': |
| 48 | self.FeatureExtraction = ResNet_FeatureExtractor(opt.input_channel, opt.output_channel) |
| 49 | elif opt.FeatureExtraction == 'VIPTRv1L': |
| 50 | self.FeatureExtraction = VIPTRv1L(opt) |
| 51 | elif opt.FeatureExtraction == 'VIPTRv1T': |
| 52 | self.FeatureExtraction = VIPTRv1(opt) |
| 53 | elif opt.FeatureExtraction == 'VIPTRv1T_ch': |
| 54 | self.FeatureExtraction = VIPTRv1T_CH(opt) |
| 55 | elif opt.FeatureExtraction == 'VIPTRv2T': |
| 56 | self.FeatureExtraction = VIPTRv2(opt) |
| 57 | elif opt.FeatureExtraction == 'VIPTRv2T_ch': |
| 58 | self.FeatureExtraction = VIPTRv2T_CH(opt) |
| 59 | elif opt.FeatureExtraction == 'VIPTRv2B': |
| 60 | self.FeatureExtraction = VIPTRv2B(opt) |
| 61 | elif opt.FeatureExtraction == 'SVTR': |
| 62 | self.FeatureExtraction = SVTRNet(img_size=[32, opt.imgW], # 100 |
| 63 | in_channels=3, |
| 64 | embed_dim=[64, 128, 256], |
| 65 | depth=[3, 6, 3], |
| 66 | num_heads=[2, 4, 8], |
| 67 | mixer=['Local'] * 6 + ['Global'] * 6, # Local atten, Global atten, Conv |
| 68 | local_mixer=[[7, 11], [7, 11], [7, 11]], |
| 69 | patch_merging='Conv', # Conv, Pool, None |
| 70 | mlp_ratio=4, |
| 71 | qkv_bias=True, |
| 72 | qk_scale=None, |
| 73 | drop_rate=0., |
| 74 | last_drop=0.1, |
| 75 | attn_drop_rate=0., |
| 76 | drop_path_rate=0.1, |
| 77 | norm_layer='nn.LayerNorm', |
| 78 | sub_norm='nn.LayerNorm', |
| 79 | epsilon=1e-6, |
| 80 | out_channels=opt.output_channel, |