RabbitRUI/ruispace
0
1import torch2from torch import nn3 4assert torch.__version__ >= "1.8.1"5from torch.utils.checkpoint import checkpoint_sequential6 7__all__ = ['iresnet2060']8 9 10def conv3x3(in_planes, out_planes, stride=1, groups=1, dilation=1):11 """3x3 convolution with padding"""12 return nn.Conv2d(in_planes,13 out_planes,14 kernel_size=3,15 stride=stride,16 padding=dilation,17 groups=groups,18 bias=False,19 dilation=dilation)20 21 22def conv1x1(in_planes, out_planes, stride=1):23 """1x1 convolution"""24 return nn.Conv2d(in_planes,25 out_planes,26 kernel_size=1,27 stride=stride,28 bias=False)29 30 31class IBasicBlock(nn.Module):32 expansion = 133 34 def __init__(self, inplanes, planes, stride=1, downsample=None,35 groups=1, base_width=64, dilation=1):36 super(IBasicBlock, self).__init__()37 if groups != 1 or base_width != 64:38 raise ValueError('BasicBlock only supports groups=1 and base_width=64')39 if dilation > 1:40 raise NotImplementedError("Dilation > 1 not supported in BasicBlock")41 self.bn1 = nn.BatchNorm2d(inplanes, eps=1e-05, )42 self.conv1 = conv3x3(inplanes, planes)43 self.bn2 = nn.BatchNorm2d(planes, eps=1e-05, )44 self.prelu = nn.PReLU(planes)45 self.conv2 = conv3x3(planes, planes, stride)46 self.bn3 = nn.BatchNorm2d(planes, eps=1e-05, )47 self.downsample = downsample48 self.stride = stride49 50 def forward(self, x):51 identity = x52 out = self.bn1(x)53 out = self.conv1(out)54 out = self.bn2(out)55 out = self.prelu(out)56 out = self.conv2(out)57 out = self.bn3(out)58 if self.downsample is not None:59 identity = self.downsample(x)60 out += identity61 return out62 63 64class IResNet(nn.Module):65 fc_scale = 7 * 766 67 def __init__(self,68 block, layers, dropout=0, num_features=512, zero_init_residual=False,69 groups=1, width_per_group=64, replace_stride_with_dilation=None, fp16=False):70 super(IResNet, self).__init__()71 self.fp16 = fp1672 self.inplanes = 6473 self.dilation = 174 if replace_stride_with_dilation is None:75 replace_stride_with_dilation = [False, False, False]76 if len(replace_stride_with_dilation) != 3:77 raise ValueError("replace_stride_with_dilation should be None "78 "or a 3-element tuple, got {}".format(replace_stride_with_dilation))79 self.groups = groups80 self.base_width = width_per_group81 self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=3, stride=1, padding=1, bias=False)82 self.bn1 = nn.BatchNorm2d(self.inplanes, eps=1e-05)83 self.prelu = nn.PReLU(self.inplanes)84 self.layer1 = self._make_layer(block, 64, layers[0], stride=2)85 self.layer2 = self._make_layer(block,86 128,87 layers[1],88 stride=2,89 dilate=replace_stride_with_dilation[0])90 self.layer3 = self._make_layer(block,91 256,92 layers[2],93 stride=2,94 dilate=replace_stride_with_dilation[1])95 self.layer4 = self._make_layer(block,96 512,97 layers[3],98 stride=2,99 dilate=replace_stride_with_dilation[2])100 self.bn2 = nn.BatchNorm2d(512 * block.expansion, eps=1e-05, )101 self.dropout = nn.Dropout(p=dropout, inplace=True)102 self.fc = nn.Linear(512 * block.expansion * self.fc_scale, num_features)103 self.features = nn.BatchNorm1d(num_features, eps=1e-05)104 nn.init.constant_(self.features.weight, 1.0)105 self.features.weight.requires_grad = False106 107 for m in self.modules():108 if isinstance(m, nn.Conv2d):109 nn.init.normal_(m.weight, 0, 0.1)110 elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):111 nn.init.constant_(m.weight, 1)112 nn.init.constant_(m.bias, 0)113 114 if zero_init_residual:115 for m in self.modules():116 if isinstance(m, IBasicBlock):117 nn.init.constant_(m.bn2.weight, 0)118 119 def _make_layer(self, block, planes, blocks, stride=1, dilate=False):120 downsample = None121 previous_dilation = self.dilation122 if dilate:123 self.dilation *= stride124 stride = 1125 if stride != 1 or self.inplanes != planes * block.expansion:126 downsample = nn.Sequential(127 conv1x1(self.inplanes, planes * block.expansion, stride),128 nn.BatchNorm2d(planes * block.expansion, eps=1e-05, ),129 )130 layers = []131 layers.append(132 block(self.inplanes, planes, stride, downsample, self.groups,133 self.base_width, previous_dilation))134 self.inplanes = planes * block.expansion135 for _ in range(1, blocks):136 layers.append(137 block(self.inplanes,138 planes,139 groups=self.groups,140 base_width=self.base_width,141 dilation=self.dilation))142 143 return nn.Sequential(*layers)144 145 def checkpoint(self, func, num_seg, x):146 if self.training:147 return checkpoint_sequential(func, num_seg, x)148 else:149 return func(x)150 151 def forward(self, x):152 with torch.cuda.amp.autocast(self.fp16):153 x = self.conv1(x)154 x = self.bn1(x)155 x = self.prelu(x)156 x = self.layer1(x)157 x = self.checkpoint(self.layer2, 20, x)158 x = self.checkpoint(self.layer3, 100, x)159 x = self.layer4(x)160 x = self.bn2(x)161 x = torch.flatten(x, 1)162 x = self.dropout(x)163 x = self.fc(x.float() if self.fp16 else x)164 x = self.features(x)165 return x166 167 168def _iresnet(arch, block, layers, pretrained, progress, **kwargs):169 model = IResNet(block, layers, **kwargs)170 if pretrained:171 raise ValueError()172 return model173 174 175def iresnet2060(pretrained=False, progress=True, **kwargs):176 return _iresnet('iresnet2060', IBasicBlock, [3, 128, 1024 - 128, 3], pretrained, progress, **kwargs)177 