CoolFace
Apppublic

kfahn/Image-to-Line-Drawings

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
app.py129 linesDownload Raw Back to root
1import numpy as np2import torch3import torch.nn as nn4import gradio as gr5from PIL import Image6import torchvision.transforms as transforms7 8norm_layer = nn.InstanceNorm2d9 10class ResidualBlock(nn.Module):11    def __init__(self, in_features):12        super(ResidualBlock, self).__init__()13 14        conv_block = [  nn.ReflectionPad2d(1),15                        nn.Conv2d(in_features, in_features, 3),16                        norm_layer(in_features),17                        nn.ReLU(inplace=True),18                        nn.ReflectionPad2d(1),19                        nn.Conv2d(in_features, in_features, 3),20                        norm_layer(in_features)21                        ]22 23        self.conv_block = nn.Sequential(*conv_block)24 25    def forward(self, x):26        return x + self.conv_block(x)27 28 29class Generator(nn.Module):30    def __init__(self, input_nc, output_nc, n_residual_blocks=9, sigmoid=True):31        super(Generator, self).__init__()32 33        # Initial convolution block34        model0 = [   nn.ReflectionPad2d(3),35                    nn.Conv2d(input_nc, 64, 7),36                    norm_layer(64),37                    nn.ReLU(inplace=True) ]38        self.model0 = nn.Sequential(*model0)39 40        # Downsampling41        model1 = []42        in_features = 6443        out_features = in_features*244        for _ in range(2):45            model1 += [  nn.Conv2d(in_features, out_features, 3, stride=2, padding=1),46                        norm_layer(out_features),47                        nn.ReLU(inplace=True) ]48            in_features = out_features49            out_features = in_features*250        self.model1 = nn.Sequential(*model1)51 52        model2 = []53        # Residual blocks54        for _ in range(n_residual_blocks):55            model2 += [ResidualBlock(in_features)]56        self.model2 = nn.Sequential(*model2)57 58        # More downsampling59        model3 = []60        out_features = in_features//261        for _ in range(2):62            model3 += [  nn.ConvTranspose2d(in_features, out_features, 3, stride=2, padding=1, output_padding=1),63                        norm_layer(out_features),64                        nn.ReLU(inplace=True) ]65            in_features = out_features66            out_features = in_features//267        self.model3 = nn.Sequential(*model3)68 69        # Output layer70        model4 = [  nn.ReflectionPad2d(3),71                        nn.Conv2d(64, output_nc, 7)]72        if sigmoid:73            model4 += [nn.Sigmoid()]74 75        self.model4 = nn.Sequential(*model4)76 77    def forward(self, x, cond=None):78        out = self.model0(x)79        out = self.model1(out)80        out = self.model2(out)81        out = self.model3(out)82        out = self.model4(out)83 84        return out85 86model1 = Generator(3, 1, 3)87model1.load_state_dict(torch.load('model.pth', map_location=torch.device('cpu')))88model1.eval()89 90model3 = Generator(3, 1, 3)91model3.load_state_dict(torch.load('model.pth', map_location=torch.device('cpu')))92model3.eval()93 94# model2 = Generator(3, 1, 3)95# model2.load_state_dict(torch.load('model2.pth', map_location=torch.device('cpu')))96# model2.eval()97 98def predict(input_img, ver):99    input_img = Image.open(input_img)100    transform = transforms.Compose([transforms.Resize(256, Image.BICUBIC), transforms.ToTensor()])101    input_img = transform(input_img)102    input_img = torch.unsqueeze(input_img, 0)103 104    drawing = 0105    with torch.no_grad():106        if ver == 'Simple Lines':107            drawing = model3(input_img)[0].detach()108        else:109            drawing = model1(input_img)[0].detach()110    111    drawing = transforms.ToPILImage()(drawing)112    return drawing113 114title="Image to Coloring Page Generator"115# examples=[116# ['01.jpeg', 'Complex Lines'], 117#]118 119# iface = gr.Interface(predict,120#     image, 121#     #gr.outputs.Image(type="pil"))122#     image)123 124 125iface = gr.Interface(predict, [gr.inputs.Image(type='filepath'),126    gr.inputs.Radio(['Complex Lines','Simple Lines'], type="value", default='Complex Lines', label='version')],127    gr.outputs.Image(type="pil"))128 129iface.launch()