Art generation with VQGAN + CLIP in Docker, with a simple web UI for anyone with a GPU. A simplified and expanded take on Kevin Costa's work.
aiai-artclipdockerdocker-composegpuhandcodedimagenetpythonpytorchtorchtorchvisionvqganvqgan-clip
1import argparse
2import json
3import os
4import uuid
5
6import numpy as np
7import torch
8import torch.nn.functional as F
9import torchvision.transforms.functional as TF
10from clip import clip
11from PIL import Image
12from tqdm import tqdm
13
14from core.schemas import Config
15from core.utils import (
16 MakeCutouts,
17 Normalize,
18 get_optimizer,
19 get_scheduler,
20 global_seed,
21 load_vqgan_model,
22 resize_image,
23)
24from core.utils.gradients import ClampWithGrad, vector_quantize
25from core.utils.noises import (
26 random_fractal_image,
27 random_gradient_image,
28 random_noise_image,
29)
30from core.utils.prompt import Prompt, parse_prompt
31
32from .check import check_files_and_folders
33
34PARAMS: Config = None
35DEVICE = torch.device(
36 os.environ.get("DEVICE", "cuda" if torch.cuda.is_available() else "cpu")
37)
38NORMALIZE = Normalize(
39 mean=[0.48145466, 0.4578275, 0.40821073],
40 std=[0.26862954, 0.26130258, 0.27577711],
41 device=DEVICE,
42)
43UUID = str(uuid.uuid4())
44
45
46def parse_args():
47 parser = argparse.ArgumentParser()
48 parser.add_argument(
49 "-c", "--config", type=str, required=True, help="Path to configuration file."
50 )
51 return parser.parse_args()
52
53
54def initialize_image(model):
55 f = 2 ** (model.decoder.num_resolutions - 1)
56 toksX, toksY = PARAMS.size[0] // f, PARAMS.size[1] // f
57 sideX, sideY = toksX * f, toksY * f
58
59 def encode(img):
60 pil_image = img.convert("RGB").resize((sideX, sideY), Image.Resampling.LANCZOS)
61 pil_tensor = TF.to_tensor(pil_image)
62 z, *_ = model.encode(pil_tensor.to(DEVICE).unsqueeze(0) * 2 - 1)
63 return z
64
65 if PARAMS.init_image and os.path.exists(PARAMS.init_image):
66 z = encode(Image.open(PARAMS.init_image))
67 elif PARAMS.init_noise == "pixels":
68 z = encode(random_noise_image(PARAMS.size[0], PARAMS.size[1]))
69 elif PARAMS.init_noise == "fractal":
70 z = encode(random_fractal_image(PARAMS.size[0], PARAMS.size[1]))
71 elif PARAMS.init_noise == "gradient":
72 z = encode(random_gradient_image(PARAMS.size[0], PARAMS.size[1]))
73 else:
74 e_dim = model.quantize.e_dim
75 n_toks = model.quantize.n_e
76
77 one_hot = F.one_hot(
78 torch.randint(n_toks, [toksY * toksX], device=DEVICE), n_toks
79 ).float()
80 z = one_hot @ model.quantize.embedding.weight
81 z = z.view([-1, toksY, toksX, e_dim]).permute(0, 3, 1, 2)
82
83 return z
84
85
86def tokenize(model, perceptor, make_cutouts):
87 f = 2 ** (model.decoder.num_resolutions - 1)
88 toksX, toksY = PARAMS.size[0] // f, PARAMS.size[1] // f
89 sideX, sideY = toksX * f, toksY * f
90
91 prompts = []
92 for prompt in PARAMS.prompts:
93 txt, weight, stop = parse_prompt(prompt)
94 embed = perceptor.encode_text(clip.tokenize(txt).to(DEVICE)).float()
95 prompts.append(Prompt(embed, weight, stop).to(DEVICE))
96
97 for prompt in PARAMS.image_prompts:
98 path, weight, stop = parse_prompt(prompt)
99 img = Image.open(path)
100 pil_image = img.convert("RGB")
101 img = resize_image(pil_image, (sideX, sideY))
102 batch = make_cutouts(TF.to_tensor(img).unsqueeze(0).to(DEVICE))
103 embed = perceptor.encode_image(NORMALIZE(batch)).float()
104 prompts.append(Prompt(embed, weight, stop).to(DEVICE))
105
106 for seed, weight in zip(PARAMS.noise_prompt_seeds, PARAMS.noise_prompt_weights):
107 gen = torch.Generator().manual_seed(seed)
108 embed = torch.empty([1, perceptor.visual.output_dim]).normal_(generator=gen)
109 prompts.append(Prompt(embed, weight).to(DEVICE))
110
111 return prompts
112
113
114def synth(z, *, model):
115 z_q = vector_quantize(z.movedim(1, 3), model.quantize.embedding.weight).movedim(
116 3, 1
117 )
118 z_q = ClampWithGrad.apply(model.decode(z_q).add(1).div(2), 0, 1)
119
120 if PARAMS.pixelart:
121 z_q = F.avg_pool2d(
122 z_q, tuple(np.ceil(np.divide(PARAMS.size, PARAMS.pixelart)).astype("uint8"))
123 )
124
125 return z_q
126
127
128@torch.no_grad()
129def checkin(z, losses, **kwargs):
130 losses_str = ", ".join(f"{loss.item():g}" for loss in losses)
131 tqdm.write(
132 f"step: {kwargs['step']}, loss: {sum(losses).item():g}, losses: {losses_str}"
133 )
134 out = synth(z, model=kwargs["model"])
135
136 filename = "output"
137 if len(PARAMS.prompts):
138 filename = "-".join(PARAMS.prompts).replace(" ", "-")
139 filename = f"{UUID}--{filename}"
140
141 path = f"{PARAMS.output_dir}/{filename}.png"
142 TF.to_pil_image(out[0].cpu()).save(path)
143
144
145def ascend_txt(z, **kwargs):
146 out = synth(z, model=kwargs["model"])
147 cutouts = kwargs["make_cutouts"](out)
148 iii = kwargs["perceptor"].encode_image(NORMALIZE(cutouts)).float()
149
150 step = kwargs["step"]
151 result = []
152 if PARAMS.init_weight:
153 mse_weight = kwargs["mse_weight"]
154 result.append(F.mse_loss(z, kwargs["z_orig"]) * mse_weight / 2)
155
156 mse_decay = PARAMS.init_weight / (PARAMS.max_iterations / PARAMS.mse_decay_rate)
157 with torch.no_grad():
158 if step > 0 and step % PARAMS.mse_decay_rate == 0:
159 kwargs["mse_weight"] = max(mse_weight - mse_decay, 0)
160
161 for prompt in kwargs["prompts"]:
162 result.append(prompt(iii))
163
164 TF.to_pil_image(out[0].cpu()).save(f"{PARAMS.output_dir}/steps/{step}.png")
165 return result
166
167
168def train(z, **kwargs):
169 kwargs["optimizer"].zero_grad(set_to_none=True)
170 lossAll = ascend_txt(z, **kwargs)
171
172 if (
173 kwargs["step"] % PARAMS.save_freq == 0
174 or kwargs["step"] == PARAMS.max_iterations
175 ):
176 checkin(z, lossAll, **kwargs)
177
178 loss = sum(lossAll)
179 loss.backward()
180 kwargs["optimizer"].step()
181
182 if kwargs["scheduler"] is not None:
183 kwargs["scheduler"].step()
184
185 with torch.no_grad():
186 z.copy_(z.maximum(kwargs["z_min"]).minimum(kwargs["z_max"]))
187
188
189def main():
190 model = load_vqgan_model(
191 PARAMS.vqgan_config, PARAMS.vqgan_checkpoint, PARAMS.models_dir
192 ).to(DEVICE)
193 perceptor = (
194 clip.load(PARAMS.clip_model, device=DEVICE, download_root=PARAMS.models_dir)[0]
195 .eval()
196 .requires_grad_(False)
197 .to(DEVICE)
198 )
199
200 cut_size = perceptor.visual.input_resolution
201 make_cutouts = MakeCutouts(
202 PARAMS.augments, cut_size, PARAMS.cutn, cut_pow=PARAMS.cut_pow
203 )
204
205 z_min = model.quantize.embedding.weight.min(dim=0).values[None, :, None, None]
206 z_max = model.quantize.embedding.weight.max(dim=0).values[None, :, None, None]
207 z = initialize_image(model)
208 z_orig = torch.zeros_like(z)
209 z.requires_grad_(True)
210
211 prompts = tokenize(model, perceptor, make_cutouts)
212 optimizer = get_optimizer(z, PARAMS.optimizer, PARAMS.step_size)
213 scheduler = get_scheduler(optimizer, PARAMS.max_iterations, PARAMS.nwarm_restarts)
214
215 kwargs = {
216 "model": model,
217 "perceptor": perceptor,
218 "optimizer": optimizer,
219 "scheduler": scheduler,
220 "prompts": prompts,
221 "make_cutouts": make_cutouts,
222 "z_orig": z_orig,
223 "z_min": z_min,
224 "z_max": z_max,
225 "mse_weight": PARAMS.init_weight,
226 }
227 try:
228 for step in tqdm(range(PARAMS.max_iterations)):
229 kwargs["step"] = step + 1
230 train(z, **kwargs)
231 except KeyboardInterrupt:
232 pass
233
234
235if __name__ == "__main__":
236 check_files_and_folders()
237
238 args = parse_args()
239
240 if not os.path.exists(args.config):
241 exit(f"ERROR: {args.config} not found.")
242
243 print(f"Loading configuration from '{args.config}'")
244 with open(args.config, "r") as f:
245 PARAMS = Config(**json.load(f))
246
247 print(f"Running on {DEVICE}.")
248 print(PARAMS)
249
250 global_seed(PARAMS.seed)
251
252 main()