diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/__pycache__/style.cpython-312.pyc b/__pycache__/style.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..9929704f95046c5379e5827413718e2320df43ba Binary files /dev/null and b/__pycache__/style.cpython-312.pyc differ diff --git a/__pycache__/style.cpython-39.pyc b/__pycache__/style.cpython-39.pyc new file mode 100644 index 0000000000000000000000000000000000000000..5708e0cd8a2f1647c979fb77e20f9b77f01d9b97 Binary files /dev/null and b/__pycache__/style.cpython-39.pyc differ diff --git a/__pycache__/transformer_net.cpython-312.pyc b/__pycache__/transformer_net.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..89539de2eb5b8f8a0ecea7bae39df56724a37d8f Binary files /dev/null and b/__pycache__/transformer_net.cpython-312.pyc differ diff --git a/__pycache__/utils.cpython-312.pyc b/__pycache__/utils.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..75811c320cc91248cf10342cad65359bf3632c92 Binary files /dev/null and b/__pycache__/utils.cpython-312.pyc differ diff --git a/__pycache__/vgg.cpython-312.pyc b/__pycache__/vgg.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..35c771a4785556358eb0a7e3129ddfe769f9c87d Binary files /dev/null and b/__pycache__/vgg.cpython-312.pyc differ diff --git a/app.py b/app.py new file mode 100644 index 0000000000000000000000000000000000000000..9a378f09a634d7a4c4a8275571f4a0a6a6074c57 --- /dev/null +++ b/app.py @@ -0,0 +1,36 @@ +import streamlit as st +from PIL import Image + +import style +st.title(":rainbow[Stylize app is so kewl]:sun_with_face::sunglasses:") +#st.title(":rainbow[using _Streamlit_ is so cool] :sunglasses:") +col1, col2 = st.columns(2,gap="large") +img = st.sidebar.selectbox( + 'Select image', + ('amber.jpg','cat.jpg','owl.jpg','dipika.jpg','mayanti.jpg','model1.jpg','model2.jpg','taapse.jpg','tamanna.jpg') +) + +style_name = st.sidebar.selectbox( + 'Select style', + ('candy','mosaic','rain_princess','udnie') +) + +model = "./saved_models/" +style_name+ ".pth" +#print(model) +input_image = "./images/content-images/" + img +output_image = "./images/output-images/" + style_name + "-" + img +with col1: + st.write("### Source Image:") + image = Image.open(input_image) + st.image(image,width=300) + +clicked = st.button("Stylize",type="primary") + +if clicked: + with col2: + model = style.load_model(model) + style.stylize(model,input_image,output_image) + + st.write('### Output Image:') + image=Image.open(output_image) + st.image(image,width=300) \ No newline at end of file diff --git a/images/content-images/amber.jpg b/images/content-images/amber.jpg new file mode 100644 index 0000000000000000000000000000000000000000..22f6390f40b0ac1bdcc8a77330d40b8299c926ad Binary files /dev/null and b/images/content-images/amber.jpg differ diff --git a/images/content-images/cat.jpg b/images/content-images/cat.jpg new file mode 100644 index 0000000000000000000000000000000000000000..63ef301e78a91481f787ed603a99898c74bf2342 Binary files /dev/null and b/images/content-images/cat.jpg differ diff --git a/images/content-images/dipika.jpg b/images/content-images/dipika.jpg new file mode 100644 index 0000000000000000000000000000000000000000..3a96495a17ceae65c4a5a54bd2859ffc3665ca15 Binary files /dev/null and b/images/content-images/dipika.jpg differ diff --git a/images/content-images/mayanti.jpg b/images/content-images/mayanti.jpg new file mode 100644 index 0000000000000000000000000000000000000000..bb47295849177e7ed794cdff70851e2945a995f2 Binary files /dev/null and b/images/content-images/mayanti.jpg differ diff --git a/images/content-images/model1.jpg b/images/content-images/model1.jpg new file mode 100644 index 0000000000000000000000000000000000000000..6eb9522487881fad5103ebff68bd32e696e7b4a9 Binary files /dev/null and b/images/content-images/model1.jpg differ diff --git a/images/content-images/model2.jpg b/images/content-images/model2.jpg new file mode 100644 index 0000000000000000000000000000000000000000..76d984c87517045ff4b8f00f25104781de7e59f7 Binary files /dev/null and b/images/content-images/model2.jpg differ diff --git a/images/content-images/owl.jpg b/images/content-images/owl.jpg new file mode 100644 index 0000000000000000000000000000000000000000..9819381fb6f804808a4abe06eced333cd565fdc5 Binary files /dev/null and b/images/content-images/owl.jpg differ diff --git a/images/content-images/taapse.jpg b/images/content-images/taapse.jpg new file mode 100644 index 0000000000000000000000000000000000000000..1c8fb4c460ba95815eccbe5f0e72f8e07fb8af87 Binary files /dev/null and b/images/content-images/taapse.jpg differ diff --git a/images/content-images/tamanna.jpg b/images/content-images/tamanna.jpg new file mode 100644 index 0000000000000000000000000000000000000000..03c0539b94cb1a215cd5484556d290fbd2f6157c Binary files /dev/null and b/images/content-images/tamanna.jpg differ diff --git a/images/output-images/amber-candy.jpg b/images/output-images/amber-candy.jpg new file mode 100644 index 0000000000000000000000000000000000000000..f585fdaae1b1450bb461b497adae74f71abe4f28 Binary files /dev/null and b/images/output-images/amber-candy.jpg differ diff --git a/images/output-images/amber-mosaic.jpg b/images/output-images/amber-mosaic.jpg new file mode 100644 index 0000000000000000000000000000000000000000..5af32e759d76bda46fd3516bc92865f17fcf32ab Binary files /dev/null and b/images/output-images/amber-mosaic.jpg differ diff --git a/images/output-images/amber-rain-princess.jpg b/images/output-images/amber-rain-princess.jpg new file mode 100644 index 0000000000000000000000000000000000000000..4f9efeb2048dd0b7cb4b80eacf3bf05359fb117a Binary files /dev/null and b/images/output-images/amber-rain-princess.jpg differ diff --git a/images/output-images/amber-udnie.jpg b/images/output-images/amber-udnie.jpg new file mode 100644 index 0000000000000000000000000000000000000000..4e7261603ac16dbfd9ff09e33e071e95af9df3df Binary files /dev/null and b/images/output-images/amber-udnie.jpg differ diff --git a/images/output-images/candy-amber.jpg b/images/output-images/candy-amber.jpg new file mode 100644 index 0000000000000000000000000000000000000000..392e0f7b428580d8c63ea31f0f8d53c3cc4bd120 Binary files /dev/null and b/images/output-images/candy-amber.jpg differ diff --git a/images/output-images/candy-dipika.jpg b/images/output-images/candy-dipika.jpg new file mode 100644 index 0000000000000000000000000000000000000000..9fee6b7d0d8bb69b1fe21b5c333c938351136c49 Binary files /dev/null and b/images/output-images/candy-dipika.jpg differ diff --git a/images/output-images/candy-mayanti.jpg b/images/output-images/candy-mayanti.jpg new file mode 100644 index 0000000000000000000000000000000000000000..73cc06aa5678b34ef947f55e335db99ce6d88e67 Binary files /dev/null and b/images/output-images/candy-mayanti.jpg differ diff --git a/images/output-images/candy-taapse.jpg b/images/output-images/candy-taapse.jpg new file mode 100644 index 0000000000000000000000000000000000000000..b6da9118f043e6d4c4fe219a2003443f4e33291d Binary files /dev/null and b/images/output-images/candy-taapse.jpg differ diff --git a/images/output-images/candy-tamanna.jpg b/images/output-images/candy-tamanna.jpg new file mode 100644 index 0000000000000000000000000000000000000000..ea3fad731d40063fbf06788f4c1e49540cc1d533 Binary files /dev/null and b/images/output-images/candy-tamanna.jpg differ diff --git a/images/output-images/mosaic-amber.jpg b/images/output-images/mosaic-amber.jpg new file mode 100644 index 0000000000000000000000000000000000000000..9a336469f96f16910d0ab721ec492854a46ee31f Binary files /dev/null and b/images/output-images/mosaic-amber.jpg differ diff --git a/images/output-images/mosaic-cat.jpg b/images/output-images/mosaic-cat.jpg new file mode 100644 index 0000000000000000000000000000000000000000..2920e5e0e0afa62111b4a8070df150bb05cfd524 Binary files /dev/null and b/images/output-images/mosaic-cat.jpg differ diff --git a/images/output-images/mosaic-dipika.jpg b/images/output-images/mosaic-dipika.jpg new file mode 100644 index 0000000000000000000000000000000000000000..34723f3e00c86a5f3fa28be614417302dcb45163 Binary files /dev/null and b/images/output-images/mosaic-dipika.jpg differ diff --git a/images/output-images/mosaic-owl.jpg b/images/output-images/mosaic-owl.jpg new file mode 100644 index 0000000000000000000000000000000000000000..c652f417f162ef3b9859ec06d80a05e89ddd20d4 Binary files /dev/null and b/images/output-images/mosaic-owl.jpg differ diff --git a/images/output-images/mosaic-taapse.jpg b/images/output-images/mosaic-taapse.jpg new file mode 100644 index 0000000000000000000000000000000000000000..8841c8e69f8fa9cca59a45d8a50c1e8c9032846a Binary files /dev/null and b/images/output-images/mosaic-taapse.jpg differ diff --git a/images/output-images/mosaic-tamanna.jpg b/images/output-images/mosaic-tamanna.jpg new file mode 100644 index 0000000000000000000000000000000000000000..996cee9897534aaf715c00b281188d1c80223712 Binary files /dev/null and b/images/output-images/mosaic-tamanna.jpg differ diff --git a/images/output-images/rain_princess-amber.jpg b/images/output-images/rain_princess-amber.jpg new file mode 100644 index 0000000000000000000000000000000000000000..a713799f696177a24f8cd02f07187e2204863c18 Binary files /dev/null and b/images/output-images/rain_princess-amber.jpg differ diff --git a/images/output-images/rain_princess-cat.jpg b/images/output-images/rain_princess-cat.jpg new file mode 100644 index 0000000000000000000000000000000000000000..2b03bdbae272d7ad0d5ec754b875fa0605779173 Binary files /dev/null and b/images/output-images/rain_princess-cat.jpg differ diff --git a/images/output-images/rain_princess-dipika.jpg b/images/output-images/rain_princess-dipika.jpg new file mode 100644 index 0000000000000000000000000000000000000000..c129511b245460b8679740933c8363ab9c646c3d Binary files /dev/null and b/images/output-images/rain_princess-dipika.jpg differ diff --git a/images/output-images/rain_princess-model1.jpg b/images/output-images/rain_princess-model1.jpg new file mode 100644 index 0000000000000000000000000000000000000000..de5f255e94e3cdec1db90a970cb51dcddec6639a Binary files /dev/null and b/images/output-images/rain_princess-model1.jpg differ diff --git a/images/output-images/rain_princess-owl.jpg b/images/output-images/rain_princess-owl.jpg new file mode 100644 index 0000000000000000000000000000000000000000..f8e2166427b9b6ecfc87aa3672d2e94674161850 Binary files /dev/null and b/images/output-images/rain_princess-owl.jpg differ diff --git a/images/output-images/rain_princess-taapse.jpg b/images/output-images/rain_princess-taapse.jpg new file mode 100644 index 0000000000000000000000000000000000000000..ca576a2595c03ee763a48311e22119875def0856 Binary files /dev/null and b/images/output-images/rain_princess-taapse.jpg differ diff --git a/images/output-images/rain_princess-tamanna.jpg b/images/output-images/rain_princess-tamanna.jpg new file mode 100644 index 0000000000000000000000000000000000000000..17bda998712dd51f5a062a07caf4be9f4484d116 Binary files /dev/null and b/images/output-images/rain_princess-tamanna.jpg differ diff --git a/images/output-images/udnie-dipika.jpg b/images/output-images/udnie-dipika.jpg new file mode 100644 index 0000000000000000000000000000000000000000..0966f43ed09b8ea54fe0c9f126b6a088bc9e256e Binary files /dev/null and b/images/output-images/udnie-dipika.jpg differ diff --git a/images/output-images/udnie-mayanti.jpg b/images/output-images/udnie-mayanti.jpg new file mode 100644 index 0000000000000000000000000000000000000000..8352dee211608840c702c4378c80d489770fd078 Binary files /dev/null and b/images/output-images/udnie-mayanti.jpg differ diff --git a/images/output-images/udnie-model2.jpg b/images/output-images/udnie-model2.jpg new file mode 100644 index 0000000000000000000000000000000000000000..1076868efb1203960bb50667d4d4a9e3a1c497c3 Binary files /dev/null and b/images/output-images/udnie-model2.jpg differ diff --git a/images/output-images/udnie-tamanna.jpg b/images/output-images/udnie-tamanna.jpg new file mode 100644 index 0000000000000000000000000000000000000000..9cfa454d1bb7814a4f871f38df44edb7805ff453 Binary files /dev/null and b/images/output-images/udnie-tamanna.jpg differ diff --git a/images/style-images/candy.jpg b/images/style-images/candy.jpg new file mode 100644 index 0000000000000000000000000000000000000000..f40e5a33e93329581baaa2ba564e02ce25615cbe Binary files /dev/null and b/images/style-images/candy.jpg differ diff --git a/images/style-images/mosaic.jpg b/images/style-images/mosaic.jpg new file mode 100644 index 0000000000000000000000000000000000000000..63aa06fe4296fa6e0feb39d5c6be57fc177283d7 Binary files /dev/null and b/images/style-images/mosaic.jpg differ diff --git a/images/style-images/rain-princess-cropped.jpg b/images/style-images/rain-princess-cropped.jpg new file mode 100644 index 0000000000000000000000000000000000000000..00a83ea48e5b65385be02937d0777770aa72a3d1 Binary files /dev/null and b/images/style-images/rain-princess-cropped.jpg differ diff --git a/images/style-images/rain-princess.jpg b/images/style-images/rain-princess.jpg new file mode 100644 index 0000000000000000000000000000000000000000..520f6a22771c9130b9907381c5e0843c9e379f06 Binary files /dev/null and b/images/style-images/rain-princess.jpg differ diff --git a/images/style-images/udnie.jpg b/images/style-images/udnie.jpg new file mode 100644 index 0000000000000000000000000000000000000000..3dbb29cf89e70f9421afb3341674eec0a991cf52 Binary files /dev/null and b/images/style-images/udnie.jpg differ diff --git a/requirements.txt.txt b/requirements.txt.txt new file mode 100644 index 0000000000000000000000000000000000000000..d064f3f9d7076810930af97aac178204a285466c --- /dev/null +++ b/requirements.txt.txt @@ -0,0 +1,3 @@ +streamlit +torch +torchvision \ No newline at end of file diff --git a/saved_models/candy.pth b/saved_models/candy.pth new file mode 100644 index 0000000000000000000000000000000000000000..2f0d2bc554256d0d3ab0d59c46864b0677966879 --- /dev/null +++ b/saved_models/candy.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3da09dfa2877fe4b64154d5561ef764a666aed17f9e59bfa6e102aaa31f4966 +size 6740723 diff --git a/saved_models/mosaic.pth b/saved_models/mosaic.pth new file mode 100644 index 0000000000000000000000000000000000000000..9a618e6060bbcf5529931b223a0fe8aba76b0f96 --- /dev/null +++ b/saved_models/mosaic.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d6d92cebde0cd068b8ed508f0d2d4c54b2b2b7137ffb4032bea44f2ab87c973 +size 6740703 diff --git a/saved_models/rain_princess.pth b/saved_models/rain_princess.pth new file mode 100644 index 0000000000000000000000000000000000000000..038b21434c01795ad597622bfab99fbe74743173 --- /dev/null +++ b/saved_models/rain_princess.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b24113ab5bd2f6d84ec00b3b5d114281f4bbcf3da74dc05f9cdcd82709578fd +size 6740717 diff --git a/saved_models/udnie.pth b/saved_models/udnie.pth new file mode 100644 index 0000000000000000000000000000000000000000..65301cfa8f9b3d9bdd53bd907a96c0850a4d1c82 --- /dev/null +++ b/saved_models/udnie.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e3f2dd8dc537495ff321e61fa0b45c663682889f5b02762d29e309aed7928da +size 6740707 diff --git a/style.py b/style.py new file mode 100644 index 0000000000000000000000000000000000000000..fccb6f6758abb3a5474b5a015a8204077e1e5ffd --- /dev/null +++ b/style.py @@ -0,0 +1,50 @@ +import argparse +import os +import sys +import time +import re + +import numpy as np +import torch +from torch.optim import Adam +from torch.utils.data import DataLoader +from torchvision import datasets +from torchvision import transforms +import torch.onnx + +import utils +from transformer_net import TransformerNet +from vgg import Vgg16 +import streamlit as st + +device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + +@st.cache_resource +def load_model(model_path): + print('load model') + with torch.no_grad(): + style_model = TransformerNet() + state_dict = torch.load(model_path) + # remove saved deprecated running_* keys in InstanceNorm from the checkpoint + for k in list(state_dict.keys()): + if re.search(r'in\d+\.running_(mean|var)$', k): + del state_dict[k] + style_model.load_state_dict(state_dict) + style_model.to(device) + style_model.eval() + return style_model + +@st.cache_resource +def stylize(_style_model, content_image, output_image): + content_image = utils.load_image(content_image) + content_transform = transforms.Compose([ + transforms.ToTensor(), + transforms.Lambda(lambda x: x.mul(255)) + ]) + content_image = content_transform(content_image) + content_image = content_image.unsqueeze(0).to(device) + + with torch.no_grad(): + output = _style_model(content_image).cpu() + + utils.save_image(output_image, output[0]) \ No newline at end of file diff --git a/transformer_net.py b/transformer_net.py new file mode 100644 index 0000000000000000000000000000000000000000..2cb06971d281d616417e8d2551bb45e78e99eb2f --- /dev/null +++ b/transformer_net.py @@ -0,0 +1,99 @@ +import torch + + +class TransformerNet(torch.nn.Module): + def __init__(self): + super(TransformerNet, self).__init__() + # Initial convolution layers + self.conv1 = ConvLayer(3, 32, kernel_size=9, stride=1) + self.in1 = torch.nn.InstanceNorm2d(32, affine=True) + self.conv2 = ConvLayer(32, 64, kernel_size=3, stride=2) + self.in2 = torch.nn.InstanceNorm2d(64, affine=True) + self.conv3 = ConvLayer(64, 128, kernel_size=3, stride=2) + self.in3 = torch.nn.InstanceNorm2d(128, affine=True) + # Residual layers + self.res1 = ResidualBlock(128) + self.res2 = ResidualBlock(128) + self.res3 = ResidualBlock(128) + self.res4 = ResidualBlock(128) + self.res5 = ResidualBlock(128) + # Upsampling Layers + self.deconv1 = UpsampleConvLayer(128, 64, kernel_size=3, stride=1, upsample=2) + self.in4 = torch.nn.InstanceNorm2d(64, affine=True) + self.deconv2 = UpsampleConvLayer(64, 32, kernel_size=3, stride=1, upsample=2) + self.in5 = torch.nn.InstanceNorm2d(32, affine=True) + self.deconv3 = ConvLayer(32, 3, kernel_size=9, stride=1) + # Non-linearities + self.relu = torch.nn.ReLU() + + def forward(self, X): + y = self.relu(self.in1(self.conv1(X))) + y = self.relu(self.in2(self.conv2(y))) + y = self.relu(self.in3(self.conv3(y))) + y = self.res1(y) + y = self.res2(y) + y = self.res3(y) + y = self.res4(y) + y = self.res5(y) + y = self.relu(self.in4(self.deconv1(y))) + y = self.relu(self.in5(self.deconv2(y))) + y = self.deconv3(y) + return y + + +class ConvLayer(torch.nn.Module): + def __init__(self, in_channels, out_channels, kernel_size, stride): + super(ConvLayer, self).__init__() + reflection_padding = kernel_size // 2 + self.reflection_pad = torch.nn.ReflectionPad2d(reflection_padding) + self.conv2d = torch.nn.Conv2d(in_channels, out_channels, kernel_size, stride) + + def forward(self, x): + out = self.reflection_pad(x) + out = self.conv2d(out) + return out + + +class ResidualBlock(torch.nn.Module): + """ResidualBlock + introduced in: https://arxiv.org/abs/1512.03385 + recommended architecture: http://torch.ch/blog/2016/02/04/resnets.html + """ + + def __init__(self, channels): + super(ResidualBlock, self).__init__() + self.conv1 = ConvLayer(channels, channels, kernel_size=3, stride=1) + self.in1 = torch.nn.InstanceNorm2d(channels, affine=True) + self.conv2 = ConvLayer(channels, channels, kernel_size=3, stride=1) + self.in2 = torch.nn.InstanceNorm2d(channels, affine=True) + self.relu = torch.nn.ReLU() + + def forward(self, x): + residual = x + out = self.relu(self.in1(self.conv1(x))) + out = self.in2(self.conv2(out)) + out = out + residual + return out + + +class UpsampleConvLayer(torch.nn.Module): + """UpsampleConvLayer + Upsamples the input and then does a convolution. This method gives better results + compared to ConvTranspose2d. + ref: http://distill.pub/2016/deconv-checkerboard/ + """ + + def __init__(self, in_channels, out_channels, kernel_size, stride, upsample=None): + super(UpsampleConvLayer, self).__init__() + self.upsample = upsample + reflection_padding = kernel_size // 2 + self.reflection_pad = torch.nn.ReflectionPad2d(reflection_padding) + self.conv2d = torch.nn.Conv2d(in_channels, out_channels, kernel_size, stride) + + def forward(self, x): + x_in = x + if self.upsample: + x_in = torch.nn.functional.interpolate(x_in, mode='nearest', scale_factor=self.upsample) + out = self.reflection_pad(x_in) + out = self.conv2d(out) + return out diff --git a/utils.py b/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..ff7036b5099505b08d727d17e41d21ed135ad6a4 --- /dev/null +++ b/utils.py @@ -0,0 +1,34 @@ +import torch +from PIL import Image + + +def load_image(filename, size=None, scale=None): + img = Image.open(filename).convert('RGB') + if size is not None: + img = img.resize((size, size), Image.LANCZOS) + elif scale is not None: + img = img.resize((int(img.size[0] / scale), int(img.size[1] / scale)), Image.LANCZOS) + return img + + +def save_image(filename, data): + img = data.clone().clamp(0, 255).numpy() + img = img.transpose(1, 2, 0).astype("uint8") + img = Image.fromarray(img) + img.save(filename) + + +def gram_matrix(y): + (b, ch, h, w) = y.size() + features = y.view(b, ch, w * h) + features_t = features.transpose(1, 2) + gram = features.bmm(features_t) / (ch * h * w) + return gram + + +def normalize_batch(batch): + # normalize using imagenet mean and std + mean = batch.new_tensor([0.485, 0.456, 0.406]).view(-1, 1, 1) + std = batch.new_tensor([0.229, 0.224, 0.225]).view(-1, 1, 1) + batch = batch.div_(255.0) + return (batch - mean) / std diff --git a/vgg.py b/vgg.py new file mode 100644 index 0000000000000000000000000000000000000000..35fd25848720f5df6117c7d91d6db84defa125ad --- /dev/null +++ b/vgg.py @@ -0,0 +1,38 @@ +from collections import namedtuple + +import torch +from torchvision import models + + +class Vgg16(torch.nn.Module): + def __init__(self, requires_grad=False): + super(Vgg16, self).__init__() + vgg_pretrained_features = models.vgg16(weights=models.VGG16_Weights.IMAGENET1K_V1).features + self.slice1 = torch.nn.Sequential() + self.slice2 = torch.nn.Sequential() + self.slice3 = torch.nn.Sequential() + self.slice4 = torch.nn.Sequential() + for x in range(4): + self.slice1.add_module(str(x), vgg_pretrained_features[x]) + for x in range(4, 9): + self.slice2.add_module(str(x), vgg_pretrained_features[x]) + for x in range(9, 16): + self.slice3.add_module(str(x), vgg_pretrained_features[x]) + for x in range(16, 23): + self.slice4.add_module(str(x), vgg_pretrained_features[x]) + if not requires_grad: + for param in self.parameters(): + param.requires_grad = False + + def forward(self, X): + h = self.slice1(X) + h_relu1_2 = h + h = self.slice2(h) + h_relu2_2 = h + h = self.slice3(h) + h_relu3_3 = h + h = self.slice4(h) + h_relu4_3 = h + vgg_outputs = namedtuple("VggOutputs", ['relu1_2', 'relu2_2', 'relu3_3', 'relu4_3']) + out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3) + return out