David310/Detect_AI-generated_Image
4
1import torch2import torch.nn as nn3from typing import Union, List, Dict, Any, cast4import torchvision5import torch.nn.functional as F6 7 8 9 10 11class VGG(torch.nn.Module):12 def __init__(self, arch_type, pretrained, progress):13 super().__init__()14 15 self.layer1 = torch.nn.Sequential()16 self.layer2 = torch.nn.Sequential()17 self.layer3 = torch.nn.Sequential()18 self.layer4 = torch.nn.Sequential()19 self.layer5 = torch.nn.Sequential()20 21 if arch_type == 'vgg11':22 official_vgg = torchvision.models.vgg11(pretrained=pretrained, progress=progress)23 blocks = [ [0,2], [2,5], [5,10], [10,15], [15,20] ]24 last_idx = 2025 elif arch_type == 'vgg19':26 official_vgg = torchvision.models.vgg19(pretrained=pretrained, progress=progress)27 blocks = [ [0,4], [4,9], [9,18], [18,27], [27,36] ]28 last_idx = 3629 else:30 raise NotImplementedError31 32 33 for x in range( *blocks[0] ):34 self.layer1.add_module(str(x), official_vgg.features[x])35 for x in range( *blocks[1] ):36 self.layer2.add_module(str(x), official_vgg.features[x])37 for x in range( *blocks[2] ):38 self.layer3.add_module(str(x), official_vgg.features[x])39 for x in range( *blocks[3] ):40 self.layer4.add_module(str(x), official_vgg.features[x])41 for x in range( *blocks[4] ):42 self.layer5.add_module(str(x), official_vgg.features[x])43 44 self.max_pool = official_vgg.features[last_idx]45 self.avgpool = nn.AdaptiveAvgPool2d((7, 7))46 47 self.fc1 = official_vgg.classifier[0]48 self.fc2 = official_vgg.classifier[3]49 self.fc3 = official_vgg.classifier[6]50 self.dropout = nn.Dropout()51 52 53 def forward(self, x):54 out = {}55 56 x = self.layer1(x)57 out['f0'] = x58 59 x = self.layer2(x)60 out['f1'] = x61 62 x = self.layer3(x)63 out['f2'] = x64 65 x = self.layer4(x)66 out['f3'] = x67 68 x = self.layer5(x)69 out['f4'] = x70 71 x = self.max_pool(x)72 x = self.avgpool(x)73 x = x.view(-1,512*7*7) 74 75 x = self.fc1(x)76 x = F.relu(x)77 x = self.dropout(x) 78 x = self.fc2(x)79 x = F.relu(x)80 out['penultimate'] = x 81 x = self.dropout(x) 82 x = self.fc3(x)83 out['logits'] = x 84 85 return out86 87 88 89 90 91 92 93 94 95 96def vgg11(pretrained=False, progress=True):97 r"""VGG 11-layer model (configuration "A") from98 `"Very Deep Convolutional Networks For Large-Scale Image Recognition" <https://arxiv.org/pdf/1409.1556.pdf>`_.99 100 Args:101 pretrained (bool): If True, returns a model pre-trained on ImageNet102 progress (bool): If True, displays a progress bar of the download to stderr103 """104 return VGG('vgg11', pretrained, progress)105 106 107 108def vgg19(pretrained=False, progress=True):109 r"""VGG 19-layer model (configuration "E")110 `"Very Deep Convolutional Networks For Large-Scale Image Recognition" <https://arxiv.org/pdf/1409.1556.pdf>`_.111 112 Args:113 pretrained (bool): If True, returns a model pre-trained on ImageNet114 progress (bool): If True, displays a progress bar of the download to stderr115 """116 return VGG('vgg19', pretrained, progress)117 118 119 120 121 