CoolFace
Apppublic

VisionLanguageGroup/MicroscopyMatching

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
backbone.py65 linesDownload Raw Back to enc_model
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