mfrashad/CharacterGAN
10
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()