Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
vgg.py121 linesDownload Raw Back to models
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