CoolFace
Apppublic

CodingTeading/NeonGAN_Demo

sourceHugging Facemitupdated 4y agoView on Hugging Face
0likes
app.py164 linesDownload Raw Back to root
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