CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
iresnet2060.py177 linesDownload Raw Back to backbones
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