MCPcopy Create free account
hub / github.com/buaacxf/VIPTR / Model

Class Model

model.py:23–139  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

21device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
22
23class 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,

Callers 3

testFunction · 0.90
trainFunction · 0.90
model.pyFile · 0.85

Calls

no outgoing calls

Tested by 1

testFunction · 0.72