CLYang617/RemoteSensingChangeDetection-RSCD.HA2F
0
1import torch.nn as nn2import math3import torch4import torch.utils.model_zoo as model_zoo5import torch.nn.functional as F6from einops import rearrange7 8 9__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101',10 'resnet152']11 12 13model_urls = {14 'resnet18': 'https://download.pytorch.org/models/resnet18-5c106cde.pth',15 'resnet34': 'https://download.pytorch.org/models/resnet34-333f7ec4.pth',16 'resnet50': 'https://download.pytorch.org/models/resnet50-19c8e357.pth',17 'resnet101': 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth',18 'resnet152': 'https://download.pytorch.org/models/resnet152-b121ed2d.pth',19}20 21 22def conv3x3(in_planes, out_planes, stride=1):23 """3x3 convolution with padding"""24 return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,25 padding=1, bias=False)26 27 28 29 30class BasicBlock(nn.Module):31 expansion = 132 33 def __init__(self, inplanes, planes, stride=1, downsample=None):34 super(BasicBlock, self).__init__()35 self.conv1 = conv3x3(inplanes, planes, stride)36 self.bn1 = nn.BatchNorm2d(planes)37 self.relu = nn.ReLU(inplace=True)38 self.conv2 = conv3x3(planes, planes)39 self.bn2 = nn.BatchNorm2d(planes)40 self.downsample = downsample41 self.stride = stride42 43 def forward(self, x):44 residual = x45 46 out = self.conv1(x)47 out = self.bn1(out)48 out = self.relu(out)49 50 out = self.conv2(out)51 out = self.bn2(out)52 53 if self.downsample is not None:54 residual = self.downsample(x)55 56 out += residual57 out = self.relu(out)58 59 return out60 61 62class Bottleneck(nn.Module):63 expansion = 464 65 def __init__(self, inplanes, planes, stride=1, downsample=None):66 super(Bottleneck, self).__init__()67 self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)68 self.bn1 = nn.BatchNorm2d(planes)69 self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride,70 padding=1, bias=False)71 self.bn2 = nn.BatchNorm2d(planes)72 self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)73 self.bn3 = nn.BatchNorm2d(planes * 4)74 self.relu = nn.ReLU(inplace=True)75 self.downsample = downsample76 self.stride = stride77 78 def forward(self, x):79 residual = x80 81 out = self.conv1(x)82 out = self.bn1(out)83 out = self.relu(out)84 85 out = self.conv2(out)86 out = self.bn2(out)87 out = self.relu(out)88 89 out = self.conv3(out)90 out = self.bn3(out)91 92 if self.downsample is not None:93 residual = self.downsample(x)94 95 out += residual96 out = self.relu(out)97 98 return out99 100 101class ResNet(nn.Module):102 103 def __init__(self, block, layers, num_classes=1000):104 self.inplanes = 64105 super(ResNet, self).__init__()106 self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,107 bias=False)108 self.bn1 = nn.BatchNorm2d(64)109 self.relu = nn.ReLU(inplace=True)110 self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)111 self.layer1 = self._make_layer(block, 64, layers[0])112 self.layer2 = self._make_layer(block, 128, layers[1], stride=2)113 self.layer3 = self._make_layer(block, 256, layers[2], stride=2)114 self.layer4 = self._make_layer(block, 512, layers[3], stride=2)115 self.avgpool = nn.AvgPool2d(7, stride=1)116 self.fc = nn.Linear(512 * block.expansion, num_classes)117 118 def _make_layer(self, block, planes, blocks, stride=1):119 downsample = None120 if stride != 1 or self.inplanes != planes * block.expansion:121 downsample = nn.Sequential(122 nn.Conv2d(self.inplanes, planes * block.expansion,123 kernel_size=1, stride=stride, bias=False),124 nn.BatchNorm2d(planes * block.expansion),125 )126 127 layers = []128 layers.append(block(self.inplanes, planes, stride, downsample))129 self.inplanes = planes * block.expansion130 for i in range(1, blocks):131 layers.append(block(self.inplanes, planes))132 133 return nn.Sequential(*layers)134 135 def forward(self, x):136 x = self.conv1(x)137 x = self.bn1(x)138 x = self.relu(x)139 x = self.maxpool(x)140 141 x = self.layer1(x)142 x = self.layer2(x)143 x = self.layer3(x)144 x = self.layer4(x)145 146 x = self.avgpool(x)147 x = x.view(x.size(0), -1)148 x = self.fc(x)149 150 return x151 152 153def resnet18(pretrained=False, **kwargs):154 """Constructs a ResNet-18 model.155 Args:156 pretrained (bool): If True, returns a model pre-trained on ImageNet157 """158 model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)159 if pretrained:160 model.load_state_dict(model_zoo.load_url(model_urls['resnet18']), strict=False)161 return model162 163 164def resnet34(pretrained=False, **kwargs):165 """Constructs a ResNet-34 model.166 Args:167 pretrained (bool): If True, returns a model pre-trained on ImageNet168 """169 model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)170 if pretrained:171 model.load_state_dict(model_zoo.load_url(model_urls['resnet34']))172 return model173 174 175def resnet50(pretrained=False, **kwargs):176 """Constructs a ResNet-50 model.177 Args:178 pretrained (bool): If True, returns a model pre-trained on ImageNet179 """180 model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)181 if pretrained:182 model.load_state_dict(model_zoo.load_url(model_urls['resnet50']))183 return model184 185 186def resnet101(pretrained=False, **kwargs):187 """Constructs a ResNet-101 model.188 Args:189 pretrained (bool): If True, returns a model pre-trained on ImageNet190 """191 model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)192 if pretrained:193 model.load_state_dict(model_zoo.load_url(model_urls['resnet101']))194 return model195 196 197def resnet152(pretrained=False, **kwargs):198 """Constructs a ResNet-152 model.199 Args:200 pretrained (bool): If True, returns a model pre-trained on ImageNet201 """202 model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)203 if pretrained:204 model.load_state_dict(model_zoo.load_url(model_urls['resnet152']))205 return model206 207 208if __name__ == '__main__':209 m = resnet18(pretrained=True, vit_dim=768)210 x = torch.rand(1, 3, 256, 256)211 vit = [torch.rand(1, 256, 768), torch.rand(1, 256, 768), torch.rand(1, 256, 768)]212 x2, x3, x4 = m(x, vit)213 print(x2.shape, x3.shape, x4.shape)