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

1.8 KB · 67 lines · Python Raw History
 1import json
 2import random
 3
 4import numpy as np
 5
 6import torch
 7import torch.optim as optim
 8from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts
 9
10from PIL import Image
11
12from core.taming.models import vqgan
13
14
15def resize_image(image, out_size):
16    ratio = image.size[0] / image.size[1]
17    area = min(image.size[0] * image.size[1], out_size[0] * out_size[1])
18    size = round((area * ratio) ** 0.5), round((area / ratio) ** 0.5)
19    return image.resize(size, Image.LANCZOS)
20
21
22def get_optimizer(z, optimizer="Adam", step_size=0.1):
23    if optimizer == "Adam":
24        opt = optim.Adam([z], lr=step_size)  # LR=0.1 (Default)
25    elif optimizer == "AdamW":
26        opt = optim.AdamW([z], lr=step_size)  # LR=0.2
27    elif optimizer == "Adagrad":
28        opt = optim.Adagrad([z], lr=step_size)  # LR=0.5+
29    elif optimizer == "Adamax":
30        opt = optim.Adamax([z], lr=step_size)  # LR=0.5+?
31    return opt
32
33
34def get_scheduler(optimizer, max_iterations, nwarm_restarts=-1):
35    if nwarm_restarts == -1:
36        return None
37
38    T_0 = max_iterations
39    if nwarm_restarts > 0:
40        T_0 = int(np.ceil(max_iterations / nwarm_restarts))
41
42    return CosineAnnealingWarmRestarts(optimizer, T_0=T_0)
43
44
45def load_vqgan_model(config_path, checkpoint_path, model_dir=None):
46    with open(config_path, "r") as f:
47        config = json.load(f)
48
49    model = vqgan.VQModel(model_dir=model_dir, **config["params"])
50    model.eval().requires_grad_(False)
51    model.init_from_ckpt(checkpoint_path)
52
53    del model.loss
54    return model
55
56
57def global_seed(seed: int):
58    seed = seed if seed != -1 else torch.seed()
59    if seed > 2**32 - 1:
60        seed = seed >> 32
61
62    random.seed(seed)
63    np.random.seed(seed)
64    torch.manual_seed(seed)
65    torch.cuda.manual_seed_all(seed)
66    print(f"Global seed set to {seed}.")