VisionLanguageGroup/MicroscopyMatching
0
1import torch2from torch import nn3from torch.nn import functional as F4from torchvision import models5from torchvision.ops.misc import FrozenBatchNorm2d6 7 8class Backbone(nn.Module):9 10 def __init__(11 self,12 name: str,13 pretrained: bool,14 dilation: bool,15 reduction: int,16 swav: bool,17 requires_grad: bool18 ):19 20 super(Backbone, self).__init__()21 22 resnet = getattr(models, name)(23 replace_stride_with_dilation=[False, False, dilation],24 pretrained=pretrained, norm_layer=FrozenBatchNorm2d25 )26 27 self.backbone = resnet28 self.reduction = reduction29 30 if name == 'resnet50' and swav:31 checkpoint = torch.hub.load_state_dict_from_url(32 'https://dl.fbaipublicfiles.com/deepcluster/swav_800ep_pretrain.pth.tar',33 map_location="cpu"34 )35 state_dict = {k.replace("module.", ""): v for k, v in checkpoint.items()}36 self.backbone.load_state_dict(state_dict, strict=False)37 38 # concatenation of layers 2, 3 and 439 self.num_channels = 896 if name in ['resnet18', 'resnet34'] else 358440 41 for n, param in self.backbone.named_parameters():42 if 'layer2' not in n and 'layer3' not in n and 'layer4' not in n:43 param.requires_grad_(False)44 else:45 param.requires_grad_(requires_grad)46 47 def forward(self, x):48 size = x.size(-2) // self.reduction, x.size(-1) // self.reduction49 x = self.backbone.conv1(x)50 x = self.backbone.bn1(x)51 x = self.backbone.relu(x)52 x = self.backbone.maxpool(x)53 54 x = self.backbone.layer1(x)55 x = layer2 = self.backbone.layer2(x)56 x = layer3 = self.backbone.layer3(x)57 x = layer4 = self.backbone.layer4(x)58 59 x = torch.cat([60 F.interpolate(f, size=size, mode='bilinear', align_corners=True)61 for f in [layer2, layer3, layer4]62 ], dim=1)63 64 return x65 