Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
resnet.py338 linesDownload Raw Back to models
1import torch2from torch import Tensor3import torch.nn as nn4from typing import Type, Any, Callable, Union, List, Optional5 6try:7    from torch.hub import load_state_dict_from_url8except ImportError:9    from torch.utils.model_zoo import load_url as load_state_dict_from_url10 11 12model_urls = {13    'resnet18': 'https://download.pytorch.org/models/resnet18-f37072fd.pth',14    'resnet34': 'https://download.pytorch.org/models/resnet34-b627a593.pth',15    'resnet50': 'https://download.pytorch.org/models/resnet50-0676ba61.pth',16    'resnet101': 'https://download.pytorch.org/models/resnet101-63fe2227.pth',17    'resnet152': 'https://download.pytorch.org/models/resnet152-394f9c45.pth',18    'resnext50_32x4d': 'https://download.pytorch.org/models/resnext50_32x4d-7cdf4587.pth',19    'resnext101_32x8d': 'https://download.pytorch.org/models/resnext101_32x8d-8ba56ff5.pth',20    'wide_resnet50_2': 'https://download.pytorch.org/models/wide_resnet50_2-95faca4d.pth',21    'wide_resnet101_2': 'https://download.pytorch.org/models/wide_resnet101_2-32ee1156.pth',22}23 24 25 26 27def conv3x3(in_planes: int, out_planes: int, stride: int = 1, groups: int = 1, dilation: int = 1) -> nn.Conv2d:28    """3x3 convolution with padding"""29    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,30                     padding=dilation, groups=groups, bias=False, dilation=dilation)31 32 33def conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d:34    """1x1 convolution"""35    return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)36 37 38class BasicBlock(nn.Module):39    expansion: int = 140 41    def __init__(42        self,43        inplanes: int,44        planes: int,45        stride: int = 1,46        downsample: Optional[nn.Module] = None,47        groups: int = 1,48        base_width: int = 64,49        dilation: int = 1,50        norm_layer: Optional[Callable[..., nn.Module]] = None51    ) -> None:52        super(BasicBlock, self).__init__()53        if norm_layer is None:54            norm_layer = nn.BatchNorm2d55        if groups != 1 or base_width != 64:56            raise ValueError('BasicBlock only supports groups=1 and base_width=64')57        if dilation > 1:58            raise NotImplementedError("Dilation > 1 not supported in BasicBlock")59        # Both self.conv1 and self.downsample layers downsample the input when stride != 160        self.conv1 = conv3x3(inplanes, planes, stride)61        self.bn1 = norm_layer(planes)62        self.relu = nn.ReLU(inplace=True)63        self.conv2 = conv3x3(planes, planes)64        self.bn2 = norm_layer(planes)65        self.downsample = downsample66        self.stride = stride67 68    def forward(self, x: Tensor) -> Tensor:69        identity = x70 71        out = self.conv1(x)72        out = self.bn1(out)73        out = self.relu(out)74 75        out = self.conv2(out)76        out = self.bn2(out)77 78        if self.downsample is not None:79            identity = self.downsample(x)80 81        out += identity82        out = self.relu(out)83 84        return out85 86 87class Bottleneck(nn.Module):88    # Bottleneck in torchvision places the stride for downsampling at 3x3 convolution(self.conv2)89    # while original implementation places the stride at the first 1x1 convolution(self.conv1)90    # according to "Deep residual learning for image recognition"https://arxiv.org/abs/1512.03385.91    # This variant is also known as ResNet V1.5 and improves accuracy according to92    # https://ngc.nvidia.com/catalog/model-scripts/nvidia:resnet_50_v1_5_for_pytorch.93 94    expansion: int = 495 96    def __init__(97        self,98        inplanes: int,99        planes: int,100        stride: int = 1,101        downsample: Optional[nn.Module] = None,102        groups: int = 1,103        base_width: int = 64,104        dilation: int = 1,105        norm_layer: Optional[Callable[..., nn.Module]] = None106    ) -> None:107        super(Bottleneck, self).__init__()108        if norm_layer is None:109            norm_layer = nn.BatchNorm2d110        width = int(planes * (base_width / 64.)) * groups111        # Both self.conv2 and self.downsample layers downsample the input when stride != 1112        self.conv1 = conv1x1(inplanes, width)113        self.bn1 = norm_layer(width)114        self.conv2 = conv3x3(width, width, stride, groups, dilation)115        self.bn2 = norm_layer(width)116        self.conv3 = conv1x1(width, planes * self.expansion)117        self.bn3 = norm_layer(planes * self.expansion)118        self.relu = nn.ReLU(inplace=True)119        self.downsample = downsample120        self.stride = stride121 122    def forward(self, x: Tensor) -> Tensor:123        identity = x124 125        out = self.conv1(x)126        out = self.bn1(out)127        out = self.relu(out)128 129        out = self.conv2(out)130        out = self.bn2(out)131        out = self.relu(out)132 133        out = self.conv3(out)134        out = self.bn3(out)135 136        if self.downsample is not None:137            identity = self.downsample(x)138 139        out += identity140        out = self.relu(out)141 142        return out143 144 145class ResNet(nn.Module):146 147    def __init__(148        self,149        block: Type[Union[BasicBlock, Bottleneck]],150        layers: List[int],151        num_classes: int = 1000,152        zero_init_residual: bool = False,153        groups: int = 1,154        width_per_group: int = 64,155        replace_stride_with_dilation: Optional[List[bool]] = None,156        norm_layer: Optional[Callable[..., nn.Module]] = None157    ) -> None:158        super(ResNet, self).__init__()159        if norm_layer is None:160            norm_layer = nn.BatchNorm2d161        self._norm_layer = norm_layer162 163        self.inplanes = 64164        self.dilation = 1165        if replace_stride_with_dilation is None:166            # each element in the tuple indicates if we should replace167            # the 2x2 stride with a dilated convolution instead168            replace_stride_with_dilation = [False, False, False]169        if len(replace_stride_with_dilation) != 3:170            raise ValueError("replace_stride_with_dilation should be None "171                             "or a 3-element tuple, got {}".format(replace_stride_with_dilation))172        self.groups = groups173        self.base_width = width_per_group174        self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=2, padding=3,175                               bias=False)176        self.bn1 = norm_layer(self.inplanes)177        self.relu = nn.ReLU(inplace=True)178        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)179        self.layer1 = self._make_layer(block, 64, layers[0])180        self.layer2 = self._make_layer(block, 128, layers[1], stride=2,181                                       dilate=replace_stride_with_dilation[0])182        self.layer3 = self._make_layer(block, 256, layers[2], stride=2,183                                       dilate=replace_stride_with_dilation[1])184        self.layer4 = self._make_layer(block, 512, layers[3], stride=2,185                                       dilate=replace_stride_with_dilation[2])186        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))187        self.fc = nn.Linear(512 * block.expansion, num_classes)188 189        for m in self.modules():190            if isinstance(m, nn.Conv2d):191                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')192            elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):193                nn.init.constant_(m.weight, 1)194                nn.init.constant_(m.bias, 0)195 196        # Zero-initialize the last BN in each residual branch,197        # so that the residual branch starts with zeros, and each residual block behaves like an identity.198        # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677199        if zero_init_residual:200            for m in self.modules():201                if isinstance(m, Bottleneck):202                    nn.init.constant_(m.bn3.weight, 0)  # type: ignore[arg-type]203                elif isinstance(m, BasicBlock):204                    nn.init.constant_(m.bn2.weight, 0)  # type: ignore[arg-type]205 206    def _make_layer(self, block: Type[Union[BasicBlock, Bottleneck]], planes: int, blocks: int,207                    stride: int = 1, dilate: bool = False) -> nn.Sequential:208        norm_layer = self._norm_layer209        downsample = None210        previous_dilation = self.dilation211        if dilate:212            self.dilation *= stride213            stride = 1214        if stride != 1 or self.inplanes != planes * block.expansion:215            downsample = nn.Sequential(216                conv1x1(self.inplanes, planes * block.expansion, stride),217                norm_layer(planes * block.expansion),218            )219 220        layers = []221        layers.append(block(self.inplanes, planes, stride, downsample, self.groups,222                            self.base_width, previous_dilation, norm_layer))223        self.inplanes = planes * block.expansion224        for _ in range(1, blocks):225            layers.append(block(self.inplanes, planes, groups=self.groups,226                                base_width=self.base_width, dilation=self.dilation,227                                norm_layer=norm_layer))228 229        return nn.Sequential(*layers)230 231    def _forward_impl(self, x):232        # The comment resolution is based on input size is 224*224 imagenet233        out = {}234        x = self.conv1(x)235        x = self.bn1(x)236        x = self.relu(x)237        x = self.maxpool(x)238        out['f0'] = x # N*64*56*56 239 240        x = self.layer1(x)241        out['f1'] = x # N*64*56*56242 243        x = self.layer2(x)244        out['f2'] = x # N*128*28*28245 246        x = self.layer3(x)247        out['f3'] = x # N*256*14*14248        249        x = self.layer4(x)250        out['f4'] = x # N*512*7*7251 252        x = self.avgpool(x)253        x = torch.flatten(x, 1)254        out['penultimate'] = x # N*512255 256        x = self.fc(x)257        out['logits'] = x # N*1000258 259        # return all features 260        return out261 262        # return final classification result 263        # return x264 265    def forward(self, x):266        return self._forward_impl(x)267 268 269def _resnet(270    arch: str,271    block: Type[Union[BasicBlock, Bottleneck]],272    layers: List[int],273    pretrained: bool,274    progress: bool,275    **kwargs: Any276) -> ResNet:277    model = ResNet(block, layers, **kwargs)278    if pretrained:279        state_dict = load_state_dict_from_url(model_urls[arch], progress=progress)280        model.load_state_dict(state_dict)281    return model282 283 284def resnet18(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> ResNet:285    r"""ResNet-18 model from286    `"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_.287 288    Args:289        pretrained (bool): If True, returns a model pre-trained on ImageNet290        progress (bool): If True, displays a progress bar of the download to stderr291    """292    return _resnet('resnet18', BasicBlock, [2, 2, 2, 2], pretrained, progress, **kwargs)293 294 295def resnet34(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> ResNet:296    r"""ResNet-34 model from297    `"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_.298 299    Args:300        pretrained (bool): If True, returns a model pre-trained on ImageNet301        progress (bool): If True, displays a progress bar of the download to stderr302    """303    return _resnet('resnet34', BasicBlock, [3, 4, 6, 3], pretrained, progress, **kwargs)304 305 306def resnet50(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> ResNet:307    r"""ResNet-50 model from308    `"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_.309 310    Args:311        pretrained (bool): If True, returns a model pre-trained on ImageNet312        progress (bool): If True, displays a progress bar of the download to stderr313    """314    return _resnet('resnet50', Bottleneck, [3, 4, 6, 3], pretrained, progress, **kwargs)315 316 317def resnet101(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> ResNet:318    r"""ResNet-101 model from319    `"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_.320 321    Args:322        pretrained (bool): If True, returns a model pre-trained on ImageNet323        progress (bool): If True, displays a progress bar of the download to stderr324    """325    return _resnet('resnet101', Bottleneck, [3, 4, 23, 3], pretrained, progress, **kwargs)326 327 328def resnet152(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> ResNet:329    r"""ResNet-152 model from330    `"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_.331 332    Args:333        pretrained (bool): If True, returns a model pre-trained on ImageNet334        progress (bool): If True, displays a progress bar of the download to stderr335    """336    return _resnet('resnet152', Bottleneck, [3, 8, 36, 3], pretrained, progress, **kwargs)337 338