CoolFace
Apppublic

mfrashad/CharacterGAN

sourceHugging Facecc-by-nc-4.0updated 4y agoView on Hugging Face
10likes
app.py146 linesDownload Raw Back to root
1import nltk; nltk.download('wordnet')2 3#@title Load Model4selected_model = 'character'5 6# Load model7import torch8import PIL9import numpy as np10from PIL import Image11from models import get_instrumented_model12from decomposition import get_or_compute13from config import Config14import gradio as gr15import numpy as np16 17# Speed up computation18torch.autograd.set_grad_enabled(False)19torch.backends.cudnn.benchmark = True20 21# Specify model to use22config = Config(23  model='StyleGAN2',24  layer='style',25  output_class=selected_model,26  components=80,27  use_w=True,28  batch_size=5_000, # style layer quite small29)30device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")31 32inst = get_instrumented_model(config.model, config.output_class,33                              config.layer, torch.device(device), use_w=config.use_w)34 35path_to_components = get_or_compute(config, inst)36 37model = inst.model38 39comps = np.load(path_to_components)40lst = comps.files41latent_dirs = []42latent_stdevs = []43 44load_activations = False45 46for item in lst:47    if load_activations:48      if item == 'act_comp':49        for i in range(comps[item].shape[0]):50          latent_dirs.append(comps[item][i])51      if item == 'act_stdev':52        for i in range(comps[item].shape[0]):53          latent_stdevs.append(comps[item][i])54    else:55      if item == 'lat_comp':56        for i in range(comps[item].shape[0]):57          latent_dirs.append(comps[item][i])58      if item == 'lat_stdev':59        for i in range(comps[item].shape[0]):60          latent_stdevs.append(comps[item][i])61 62 63def display_sample_pytorch(seed, truncation, directions, distances, scale, start, end, w=None, disp=True, save=None, noise_spec=None):64    # blockPrint()65    model.truncation = truncation66    if w is None:67        w = model.sample_latent(1, seed=seed).detach().cpu().numpy()68        w = [w]*model.get_max_latents() # one per layer69    else:70        w = [np.expand_dims(x, 0) for x in w]71    72    for l in range(start, end):73      for i in range(len(directions)):74        w[l] = w[l] + directions[i] * distances[i] * scale75    76    torch.cuda.empty_cache()77    #save image and display78    out = model.sample_np(w)79    final_im = Image.fromarray((out * 255).astype(np.uint8)).resize((500,500),Image.LANCZOS)80    81    82    if save is not None:83      if disp == False:84        print(save)85      final_im.save(f'out/{seed}_{save:05}.png')86    87    return final_im88 89 90#@title Demo UI91 92 93def generate_image(seed, truncation,94                  monster, female, skimpy, light, bodysuit, bulky, human_head,95                  start_layer, end_layer):96    seed = hash(seed) % 100000000097    scale = 198    params = {'monster': monster,99          'female': female,100          'skimpy': skimpy,101          'light': light,102          'bodysuit': bodysuit,103          'bulky': bulky,104          'human_head': human_head}105 106    param_indexes = {'monster': 0,107              'female': 1,108              'skimpy': 2,109              'light': 4,110              'bodysuit': 5,111              'bulky': 6,112              'human_head': 8}113 114    directions = []115    distances = []116    for k, v in params.items():117        directions.append(latent_dirs[param_indexes[k]])118        distances.append(v)119 120    style = {'description_width': 'initial'}121    return display_sample_pytorch(int(seed), truncation, directions, distances, scale, int(start_layer), int(end_layer), disp=False)122 123truncation = gr.inputs.Slider(minimum=0, maximum=1, default=0.5, label="Truncation")124start_layer = gr.inputs.Number(default=0, label="Start Layer")125end_layer = gr.inputs.Number(default=14, label="End Layer")126seed = gr.inputs.Textbox(default="0", label="Seed")127 128slider_max_val = 20129slider_min_val = -20130slider_step = 1131 132monster = gr.inputs.Slider(label="Monsterfication", minimum=slider_min_val, maximum=slider_max_val, default=0)133female = gr.inputs.Slider(label="Gender", minimum=slider_min_val, maximum=slider_max_val, default=0)134skimpy = gr.inputs.Slider(label="Amount of Clothing", minimum=slider_min_val, maximum=slider_max_val, default=0)135light = gr.inputs.Slider(label="Brightness", minimum=slider_min_val, maximum=slider_max_val, default=0)136bodysuit = gr.inputs.Slider(label="Bodysuit", minimum=slider_min_val, maximum=slider_max_val, default=0)137bulky = gr.inputs.Slider(label="Bulkiness", minimum=slider_min_val, maximum=slider_max_val, default=0)138human_head = gr.inputs.Slider(label="Head", minimum=slider_min_val, maximum=slider_max_val, default=0)139 140 141scale = 1142 143inputs = [seed, truncation, monster, female, skimpy, light, bodysuit, bulky, human_head, start_layer, end_layer]144description = "Change the seed number to generate different character design. Made by <a href='https://www.mfrashad.com/' target='_blank'>@mfrashad</a>. For more details on how to build this, visit the <a href='https://github.com/mfrashad/gancreate-saai' target='_blank'>repo</a>. Please give a star if you find it useful :)"145 146gr.Interface(generate_image, inputs, ["image"], description=description, live=True, title="CharacterGAN").launch()