Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
vgg.py91 linesDownload Raw Back to models
1from collections import namedtuple
2
3import torch
4from torchvision import models
5
6
7class Vgg16(torch.nn.Module):
8    def __init__(self, requires_grad=False):
9        super(Vgg16, self).__init__()
10        vgg_pretrained_features = models.vgg16(pretrained=True).features
11        self.slice1 = torch.nn.Sequential()
12        self.slice2 = torch.nn.Sequential()
13        self.slice3 = torch.nn.Sequential()
14        self.slice4 = torch.nn.Sequential()
15        for x in range(4):
16            self.slice1.add_module(str(x), vgg_pretrained_features[x])
17        for x in range(4, 9):
18            self.slice2.add_module(str(x), vgg_pretrained_features[x])
19        for x in range(9, 16):
20            self.slice3.add_module(str(x), vgg_pretrained_features[x])
21        for x in range(16, 23):
22            self.slice4.add_module(str(x), vgg_pretrained_features[x])
23        if not requires_grad:
24            for param in self.parameters():
25                param.requires_grad = False
26
27    def forward(self, X):
28        h = self.slice1(X)
29        h_relu1_2 = h
30        h = self.slice2(h)
31        h_relu2_2 = h
32        h = self.slice3(h)
33        h_relu3_3 = h
34        h = self.slice4(h)
35        h_relu4_3 = h
36        vgg_outputs = namedtuple("VggOutputs", ['relu1_2', 'relu2_2', 'relu3_3', 'relu4_3'])
37        out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3)
38        return out
39
40
41class Vgg19(torch.nn.Module):
42    def __init__(self, requires_grad=False):
43        super(Vgg19, self).__init__()
44        # vgg_pretrained_features = models.vgg19(pretrained=True).features
45        self.vgg_pretrained_features = models.vgg19(pretrained=True).features
46        # self.slice1 = torch.nn.Sequential()
47        # self.slice2 = torch.nn.Sequential()
48        # self.slice3 = torch.nn.Sequential()
49        # self.slice4 = torch.nn.Sequential()
50        # self.slice5 = torch.nn.Sequential()
51        # for x in range(2):
52        #     self.slice1.add_module(str(x), vgg_pretrained_features[x])
53        # for x in range(2, 7):
54        #     self.slice2.add_module(str(x), vgg_pretrained_features[x])
55        # for x in range(7, 12):
56        #     self.slice3.add_module(str(x), vgg_pretrained_features[x])
57        # for x in range(12, 21):
58        #     self.slice4.add_module(str(x), vgg_pretrained_features[x])
59        # for x in range(21, 30):
60        #     self.slice5.add_module(str(x), vgg_pretrained_features[x])
61        if not requires_grad:
62            for param in self.parameters():
63                param.requires_grad = False
64
65    def forward(self, X, indices=None):
66        if indices is None:
67            indices = [2, 7, 12, 21, 30]
68        out = []
69        # indices = sorted(indices)
70        for i in range(indices[-1]):
71            X = self.vgg_pretrained_features[i](X)
72            if (i + 1) in indices:
73                out.append(X)
74
75        return out
76
77        # h_relu1 = self.slice1(X)
78        # h_relu2 = self.slice2(h_relu1)
79        # h_relu3 = self.slice3(h_relu2)
80        # h_relu4 = self.slice4(h_relu3)
81        # h_relu5 = self.slice5(h_relu4)
82        # out = [h_relu1, h_relu2, h_relu3, h_relu4, h_relu5]
83        # return out
84
85
86if __name__ == '__main__':
87    vgg = Vgg19()
88    import ipdb
89
90    ipdb.set_trace()
91