diff --git a/.gitignore b/.gitignore index 68bc17f..83bae75 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,18 @@ +raw/ +images/ +latent_*/ +vae/ +models/ +other/ +test.py +*.png +*.zip +*.npy +*.ckpt +*.safetensors + +# default github .gitignore follows + # Byte-compiled / optimized / DLL files __pycache__/ *.py[cod] diff --git a/log_loss.py b/log_loss.py new file mode 100644 index 0000000..ac769eb --- /dev/null +++ b/log_loss.py @@ -0,0 +1,48 @@ +import os +import math +import matplotlib.pyplot as plt + +files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")] +train_loss = {} +eval_loss = {} + +def process_lines(lines): + global train_loss + global eval_loss + name = fp.split("/")[1] + vals = [x.split(",") for x in lines] + train_loss[name] = ( + [int(x[0]) for x in vals], + [math.log(float(x[1])) for x in vals], + ) + if len(vals[0]) == 3: + eval_loss[name] = ( + [int(x[0]) for x in vals], + [math.log(float(x[2])) for x in vals], + ) + +# https://stackoverflow.com/a/49357445 +def smooth(scalars, weight): + last = scalars[0] + smoothed = list() + for point in scalars: + smoothed_val = last * weight + (1 - weight) * point + smoothed.append(smoothed_val) + last = smoothed_val + return smoothed + +def plot(data, fname): + fig, ax = plt.subplots() + ax.grid() + for name, val in data.items(): + ax.plot(val[0], smooth(val[1], 0.9), label=name) + plt.legend(loc="upper right") + plt.savefig(fname, dpi=300, bbox_inches='tight') + +for fp in files: + with open(fp) as f: + lines = f.readlines() + process_lines(lines) + +plot(train_loss, "loss.png") +plot(eval_loss, "loss-eval.png") diff --git a/preprocess_latents.py b/preprocess_latents.py new file mode 100644 index 0000000..8275332 --- /dev/null +++ b/preprocess_latents.py @@ -0,0 +1,80 @@ +import os +import torch +import hashlib +import argparse +import numpy as np +from torchvision import transforms +from diffusers import AutoencoderKL +from tqdm import tqdm +from PIL import Image + +from vae import get_vae + +def parse_args(): + parser = argparse.ArgumentParser(description="Preprocess images into latents") + parser.add_argument("-r", "--res", type=int, default=512, help="Source resolution") + parser.add_argument("-f", "--fac", type=float, default=1.5, help="Upscale factor") + parser.add_argument("-v", "--ver", choices=["v1","xl"], default="v1", help="SD version") + parser.add_argument('--vae', help="Path to VAE (Optional)") + parser.add_argument('--src', default="raw", help="Source folder with images") + return parser.parse_args() + +def encode(vae, img): + """image [PIL Image] -> latent [np array]""" + inp = transforms.ToTensor()(img).unsqueeze(0) + inp = inp.to("cuda") # move to GPU + latent = vae.encode(inp*2.0-1.0) + latent = latent.latent_dist.sample() + return latent.cpu().detach() + +def scale(path, res): + """Crop image to the top-left corner""" + img = Image.open(path) + img = img.convert('RGB') + target = (res, res) + if min(img.height, img.width) < 256: + return + if img.width > img.height: + target = (int(img.width/img.height*res), res) + elif img.height > img.width: + target = (res, int(img.height/img.width*res)) + img = img.resize(target, Image.LANCZOS) + img = img.crop([0,0,res,res]) + return img + +def process_folder(vae, src_dir, ver, res): + dst_dir = f"latents/{ver}_{res}px" + if not os.path.isdir(dst_dir): + os.mkdir(dst_dir) + + for file in tqdm(os.listdir(src_dir)): + src = os.path.join(src_dir, file) + md5 = hashlib.md5(open(src,'rb').read()).hexdigest() + dst = os.path.join(dst_dir, f"{md5}.npy") + if os.path.isfile(dst): + continue + img = scale(src, res) + latent = encode(vae, img) + np.save(dst, latent) + +def process_res(vae, src_dir, ver, res): + process_folder(vae, src_dir, ver, res) + # test image, optional + if os.path.isfile("test.png"): + if os.path.isfile(f"test_{ver}_{res}px.npy"): + return + img = scale("test.png", res) + latent = encode(vae, img) + np.save(f"test_{ver}_{res}px.npy", latent) + torch.cuda.empty_cache() + +if __name__ == "__main__": + if not os.path.isdir("latents"): + os.mkdir("latents") + args = parse_args() + vae = get_vae(args.ver, args.vae) + vae.to("cuda") + ## args + dst_res = int(args.res*args.fac) + process_res(vae, args.src, args.ver, args.res) + process_res(vae, args.src, args.ver, dst_res) diff --git a/train.py b/train.py new file mode 100644 index 0000000..19ead4d --- /dev/null +++ b/train.py @@ -0,0 +1,126 @@ +import os +import torch +import torch.nn as nn +import numpy as np +import argparse +import random +from PIL import Image +from tqdm import tqdm +from safetensors.torch import save_file, load_file + +from upscaler import LatentUpscaler as Upscaler +from vae import get_vae + +def parse_args(): + parser = argparse.ArgumentParser(description="Train latent interposer model") + parser.add_argument("--steps", type=int, default=500000, help="No. of training steps") + parser.add_argument('--bs', type=int, default=1, help="Batch size") + parser.add_argument('--lr', default="1e-8", help="Learning rate") + parser.add_argument("-n", "--save_every_n", type=int, dest="save", default=50000, help="Save model/sample periodically") + parser.add_argument("-r", "--res", type=int, default=512, help="Source resolution") + parser.add_argument("-f", "--fac", type=float, default=1.5, help="Upscale factor") + parser.add_argument("-v", "--ver", choices=["v1","xl"], default="v1", help="SD version") + parser.add_argument('--vae', help="Path to VAE (Optional)") + parser.add_argument('--resume', help="Checkpoint to resume from") + args = parser.parse_args() + try: + float(args.lr) + except: + parser.error("--lr must be a valid float eg. 0.001 or 1e-3") + return args + +class Latent: + def __init__(self, md5, ver, src_res, dst_res, dev): + src = os.path.join(f"latents/{ver}_{src_res}px", f"{md5}.npy") + dst = os.path.join(f"latents/{ver}_{dst_res}px", f"{md5}.npy") + self.src = torch.from_numpy(np.load(src)).to(dev) + self.dst = torch.from_numpy(np.load(dst)).to(dev) + +def load_latents(ver, src_res, dst_res, dev): + print("Loading latents from disk") + latents = [] + for i in tqdm(os.listdir(f"latents/{ver}_{src_res}px")): + md5 = os.path.splitext(i)[0] + latents.append(Latent(md5, ver, src_res, dst_res, dev)) + return latents + +vae = None +def sample_decode(latent, filename, version): + global vae + if not vae: + vae = get_vae(version, fp16=True) + vae.to("cuda") + + latent = latent.half().to("cuda") + out = vae.decode(latent).sample + out = out.cpu().detach().numpy() + out = np.squeeze(out, 0) + out = out.transpose((1, 2, 0)) + out = np.clip(out, -1.0, 1.0) + out = (out+1)/2 * 255 + out = out.astype(np.uint8) + out = Image.fromarray(out) + out.save(filename) + +if __name__ == "__main__": + args = parse_args() + target_dev = "cuda" + dst_res = int(args.res*args.fac) + + latents = load_latents(args.ver, args.res, dst_res, target_dev) + + if not os.path.isdir("models"): os.mkdir("models") + log = open(f"models/latent-upscaler_SD{args.ver}-x{args.fac}.csv", "w") + + if os.path.isfile(f"test_{args.ver}_{args.res}px.npy") and os.path.isfile(f"test_{args.ver}_{dst_res}px.npy"): + ss_latent = torch.from_numpy(np.load(f"test_{args.ver}_{args.res}px.npy")).to(target_dev) + st_latent = torch.from_numpy(np.load(f"test_{args.ver}_{dst_res}px.npy")).to(target_dev) + else: + sample_latent = random.choice(latents) + ss_latent = sample_latent.src.to(target_dev) + st_latent = sample_latent.dst.to(target_dev) + + model = Upscaler(args.fac) + if args.resume: + model.load_state_dict(load_file(args.resume)) + model.to(target_dev) + + criterion = torch.nn.MSELoss(size_average=False) + optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) + + for t in tqdm(range(int(args.steps/args.bs)), unit_scale=args.bs): + step = t*args.bs + # input batch + lts = [random.choice(latents) for _ in range(args.bs)] + src = torch.cat([x.src for x in lts],0) + dst = torch.cat([x.dst for x in lts],0) + + y_pred = model(src) # forward + loss = criterion(y_pred, dst) # loss + + # backward + optimizer.zero_grad() + loss.backward() + optimizer.step() + + # print loss + if step%1000 == 0: + # test loss + with torch.no_grad(): + t_pred = model(ss_latent) + t_loss = criterion(t_pred, st_latent) + tqdm.write(f"{step} - {loss.data.item()/args.bs:.2f}|{t_loss.data.item()/args.bs:.2f}") + log.write(f"{step},{loss.data.item()/args.bs:.2f},{t_loss.data.item()/args.bs:.2f}\n") + log.flush() + + # sample/save + if step%args.save == 0: + out = model(ss_latent) + output_name = f"./models/latent-upscaler_SD{args.ver}-x{args.fac}_e{step/1000}k" + sample_decode(out, f"{output_name}.png", args.ver) + save_file(model.state_dict(), f"{output_name}.safetensors") + # save final output + output_name = f"./models/latent-upscaler_SD{args.ver}-x{args.fac}_e{step/1000}k" + sample_decode(out, f"{output_name}.png", args.ver) + save_file(model.state_dict(), f"{output_name}.safetensors") + log.close() diff --git a/upscaler.py b/upscaler.py new file mode 100644 index 0000000..094e50c --- /dev/null +++ b/upscaler.py @@ -0,0 +1,25 @@ +import torch +import torch.nn as nn +import numpy as np + +class LatentUpscaler(nn.Module): + def __init__(self, fac): + super().__init__() + + module_list = [ + nn.Conv2d(4, 64, kernel_size=5, padding=2), + nn.ReLU(), + nn.Upsample(scale_factor=fac, mode="nearest"), # bicubic was blurry + nn.ReLU(), + nn.Conv2d(64, 64, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(64, 64, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(64, 32, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(32, 4, kernel_size=5, padding=2), + ] + self.sequential = nn.Sequential(*module_list) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.sequential(x) diff --git a/vae.py b/vae.py new file mode 100644 index 0000000..0bd4f9e --- /dev/null +++ b/vae.py @@ -0,0 +1,48 @@ +import torch +from diffusers import AutoencoderKL + +def get_vae(version, file_path=None, fp16=False): + """Load VAE from file or default hf repo. fp16 only works from hf""" + vae = None + dtype = torch.float16 if fp16 else torch.float32 + if version == "v1" and file_path: + vae = AutoencoderKL.from_single_file( + file_path, + image_size=512, + ) + elif version == "v1": + vae = AutoencoderKL.from_pretrained( + "runwayml/stable-diffusion-v1-5", + subfolder="vae", + torch_dtype=dtype, + ) + elif version == "v2" and file_path: + vae = AutoencoderKL.from_single_file( + file_path, + image_size=768, + ) + elif version == "v2": + vae = AutoencoderKL.from_pretrained( + "stabilityai/stable-diffusion-2-1", + subfolder="vae", + torch_dtype=dtype, + ) + elif version == "xl" and file_path: + vae = AutoencoderKL.from_single_file( + file_path, + image_size=1024 + ) + elif version == "xl" and fp16: + vae = AutoencoderKL.from_pretrained( + "madebyollin/sdxl-vae-fp16-fix", + torch_dtype=torch.float16, + ) + elif version == "xl": + vae = AutoencoderKL.from_pretrained( + "stabilityai/stable-diffusion-xl-base-1.0", + subfolder="vae" + ) + else: + input("Invalid VAE version. Press any key to exit") + exit(1) + return vae