Tl;DR: How could I access the pytorch pre-trained model for Swin-Transformer so that I could extract features from it to train it on segmentation task using DeepLabv3+ head on a custom data set with image sizes of 512
I am testing SwinTransformer backbone with Deeplabv3+ as head for semantic segmentation.
I already have the code for Head and it get the features from backbone and then process the features in different way. It is working fine for ResNet and Xception. The thing is that I want to work with pre-trained SWIN. Now there are many ways but each of those has it's own set of problems.
1. Using timm library for pre-trainned models`
! pip install timm
import timm
import torch
all_swins = timm.list_models('*swin*')
print(all_swins)
model = timm.create_model('swin_large_patch4_window12_384_in22k', in_chans = 3, pretrained = True,)
print(model.default_cfg)
dummy_image = torch.randn(1,3,512,512) # create an image of size 512
model.forward_features(dummy_image) # extract features : Will result in error
new_model = timm.create_model('swin_large_patch4_window12_384_in22k', in_chans = 3, pretrained = True, img_size = 512) # This won't work I'll explain below
The problem is that None of the above code work. Reason being that the model accepts here an image of size 384 but my images are of 512. So you could change the argument img_size for other CNNs but as the author has clarified in this question
It should work with the vit, vit_deit, vit_deit_distilled. Has not been implemented for pit, swin, and tnt yet.
2. Using MMcv / MMSeg library:
Please open this colab notebook. I have commented and documented the part
Problem: The pre-trained weights are for only for a specific method which produced SOTA results i.e ADE dataset using UperNet backbone. I can not use it with DeepLabv3+ on a custom dataset.
3. Segmentations Models Pytorch Library which uses timm encoders
Problem: Again, as it uses timm, so the image resolutions can't be changed.
It has Swin transformer but Deeplabv3+ works only with Resnet50 and 101
Last Resort: In the end, I pulled up the official code from microsoft where I found couple of useful things:
- configuration
ymlfile - Code which they use for model building
def build_model() - Code inside
def main()which parses the arguments and builds the whole model
I don't really know how to build the model changing the image size configurations. If anyone has any idea, please help.