repos
/ ai-art master

ai-art

mirror archived upstream

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

7.4 KB · 253 lines · Python Raw History
  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()