David310/Detect_AI-generated_Image
4
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 