CodingTeading/NeonGAN_Demo
0
1import torch2from torchvision.utils import make_grid3from torchvision import transforms4import torchvision.transforms.functional as TF5from torch import nn, optim6from torch.optim.lr_scheduler import CosineAnnealingLR7from torch.utils.data import DataLoader, Dataset8from huggingface_hub import hf_hub_download9import requests10import gradio as gr11import numpy as np12from PIL import Image13 14class Upsample(nn.Module):15 def __init__(self, in_channels, out_channels, kernel_size=4, stride=2, padding=1, dropout=True):16 super(Upsample, self).__init__()17 self.dropout = dropout18 self.block = nn.Sequential(19 nn.ConvTranspose2d(in_channels, out_channels, kernel_size, stride, padding, bias=nn.InstanceNorm2d),20 nn.InstanceNorm2d(out_channels),21 nn.ReLU(inplace=True)22 )23 self.dropout_layer = nn.Dropout2d(0.5)24 25 def forward(self, x, shortcut=None):26 x = self.block(x)27 if self.dropout:28 x = self.dropout_layer(x)29 30 if shortcut is not None:31 x = torch.cat([x, shortcut], dim=1)32 33 return x34 35 36class Downsample(nn.Module):37 def __init__(self, in_channels, out_channels, kernel_size=4, stride=2, padding=1, apply_instancenorm=True):38 super(Downsample, self).__init__()39 self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, bias=nn.InstanceNorm2d)40 self.norm = nn.InstanceNorm2d(out_channels)41 self.relu = nn.LeakyReLU(0.2, inplace=True)42 self.apply_norm = apply_instancenorm43 44 def forward(self, x):45 x = self.conv(x)46 if self.apply_norm:47 x = self.norm(x)48 x = self.relu(x)49 50 return x51 52 53class CycleGAN_Unet_Generator(nn.Module):54 def __init__(self, filter=64):55 super(CycleGAN_Unet_Generator, self).__init__()56 self.downsamples = nn.ModuleList([57 Downsample(3, filter, kernel_size=4, apply_instancenorm=False), # (b, filter, 128, 128)58 Downsample(filter, filter * 2), # (b, filter * 2, 64, 64)59 Downsample(filter * 2, filter * 4), # (b, filter * 4, 32, 32)60 Downsample(filter * 4, filter * 8), # (b, filter * 8, 16, 16)61 Downsample(filter * 8, filter * 8), # (b, filter * 8, 8, 8)62 Downsample(filter * 8, filter * 8), # (b, filter * 8, 4, 4)63 Downsample(filter * 8, filter * 8), # (b, filter * 8, 2, 2)64 ])65 66 self.upsamples = nn.ModuleList([67 Upsample(filter * 8, filter * 8),68 Upsample(filter * 16, filter * 8),69 Upsample(filter * 16, filter * 8),70 Upsample(filter * 16, filter * 4, dropout=False),71 Upsample(filter * 8, filter * 2, dropout=False),72 Upsample(filter * 4, filter, dropout=False)73 ])74 75 self.last = nn.Sequential(76 nn.ConvTranspose2d(filter * 2, 3, kernel_size=4, stride=2, padding=1),77 nn.Tanh()78 )79 80 def forward(self, x):81 skips = []82 for l in self.downsamples:83 x = l(x)84 skips.append(x)85 86 skips = reversed(skips[:-1])87 for l, s in zip(self.upsamples, skips):88 x = l(x, s)89 90 out = self.last(x)91 92 return out93 94class ImageTransform:95 def __init__(self, img_size=256):96 self.transform = {97 'train': transforms.Compose([98 transforms.Resize((img_size, img_size)),99 transforms.RandomHorizontalFlip(),100 transforms.RandomVerticalFlip(),101 transforms.ToTensor(),102 transforms.Normalize(mean=[0.5], std=[0.5])103 ]),104 'test': transforms.Compose([105 transforms.Resize((img_size, img_size)),106 transforms.ToTensor(),107 transforms.Normalize(mean=[0.5], std=[0.5])108 ])}109 def __call__(self, img, phase='train'):110 img = self.transform[phase](img)111 return img112 113 114title = "Generate Futuristic Images with NeonGAN"115 116path = hf_hub_download('huggan/NeonGAN', 'model.bin')117model_gen_n = torch.load(path, map_location=torch.device('cpu'))118 119transform = ImageTransform(img_size=256)120 121inputs = [122 gr.inputs.Image(type="pil", label="Original Image")123]124 125outputs = [126 gr.outputs.Image(type="pil", label="Neon Image")127]128 129examples = [['img_1.jpg'],['img_2.jpg']]130 131def get_output_image(img):132 133 img = transform(img, phase='test')134 gen_img = model_gen_n(img.unsqueeze(0))[0]135 136 # Reverse Normalization137 gen_img = gen_img * 0.5 + 0.5138 gen_img = gen_img * 255139 gen_img = gen_img.detach().cpu().numpy().astype(np.uint8)140 141 gen_img = np.transpose(gen_img, [1,2,0])142 143 gen_img = Image.fromarray(gen_img)144 print(gen_img)145 146 return gen_img147 148gr.Interface(149 get_output_image,150 inputs,151 outputs,152 examples = examples,153 title=title,154 theme="huggingface",155).launch(enable_queue=True)156 157 158 159 160 161 162 163 164 