| 943 | |
| 944 | |
| 945 | class VIPTRNet(nn.Module): |
| 946 | def __init__(self, in_chans=3, out_dim=192, |
| 947 | embed_dims=[96, 192, 384, 768], depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], |
| 948 | init_values=[1, 1, 1, 1], heads_ranges=[3, 3, 3, 3], mlp_ratios=[3, 3, 3, 3], split_sizes=[1, 2, 2, 4], |
| 949 | sr_ratios=[8, 4, 2, 1], drop_path_rate=0.1, norm_layer=nn.LayerNorm, |
| 950 | patch_norm=True, use_checkpoints=[False, False, False, False], |
| 951 | mixer_types=['Local1', 'LG1', 'Global2'], chunkwise_recurrents=[True, True, False, False], |
| 952 | layerscales=[False, False, False, False], layer_init_values=1e-6): |
| 953 | super().__init__() |
| 954 | |
| 955 | self.out_dim = out_dim |
| 956 | self.num_layers = len(depths) |
| 957 | self.embed_dim = embed_dims[0] |
| 958 | self.patch_norm = patch_norm |
| 959 | self.num_features = embed_dims[-1] |
| 960 | self.mlp_ratios = mlp_ratios |
| 961 | |
| 962 | # split image into non-overlapping patches |
| 963 | self.patch_embed = PatchEmbed(in_chans=in_chans, embed_dim=embed_dims[0], |
| 964 | norm_layer=norm_layer if self.patch_norm else None) |
| 965 | |
| 966 | # stochastic depth |
| 967 | dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule |
| 968 | |
| 969 | # build layers |
| 970 | self.layers = nn.ModuleList() |
| 971 | for i_layer in range(self.num_layers): |
| 972 | layer = BasicLayer( |
| 973 | embed_dim=embed_dims[i_layer], |
| 974 | out_dim=embed_dims[i_layer + 1] if (i_layer < self.num_layers - 1) else None, |
| 975 | depth=depths[i_layer], |
| 976 | num_heads=num_heads[i_layer], |
| 977 | init_value=init_values[i_layer], |
| 978 | heads_range=heads_ranges[i_layer], |
| 979 | mlp_ratio=mlp_ratios[i_layer], |
| 980 | split_size=split_sizes[i_layer], |
| 981 | sr_ratio=sr_ratios[i_layer], |
| 982 | qkv_bias=True, |
| 983 | qk_scale=None, |
| 984 | drop_rate=0., |
| 985 | attn_drop=0.0, |
| 986 | drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])], |
| 987 | # norm_layer=norm_layer, |
| 988 | chunkwise_recurrent=chunkwise_recurrents[i_layer], |
| 989 | downsample=PatchMerging if (i_layer in [0, 1]) else None, # PatchMerging if (i_layer < self.num_layers - 1) else None, |
| 990 | use_checkpoint=use_checkpoints[i_layer], |
| 991 | mixer_type=mixer_types[i_layer], |
| 992 | layerscale=layerscales[i_layer], |
| 993 | layer_init_values=layer_init_values |
| 994 | ) |
| 995 | self.layers.append(layer) |
| 996 | |
| 997 | self.pooling = nn.AdaptiveAvgPool2d((embed_dims[self.num_layers - 1], 1)) |
| 998 | self.mlp_head = nn.Sequential( |
| 999 | nn.Linear(embed_dims[self.num_layers - 1], out_dim, bias=False), |
| 1000 | nn.Hardswish(), |
| 1001 | nn.Dropout(p=0.1) |
| 1002 | ) |