From a347f3960601a3e796b0fb8e0ace78718af084ab Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 31 Jul 2023 18:15:13 +0200 Subject: [PATCH] Switch to args instead of hardcoding variables --- preprocess_images.py | 35 +++++++++++++++++++--------------- train.py | 45 +++++++++++++++++++++++++++++++------------- 2 files changed, 52 insertions(+), 28 deletions(-) diff --git a/preprocess_images.py b/preprocess_images.py index b82fab1..8cb2498 100644 --- a/preprocess_images.py +++ b/preprocess_images.py @@ -1,23 +1,22 @@ import os import hashlib +import argparse from tqdm import tqdm from PIL import Image from queue import Queue from threading import Thread -# target resolution [latent res * 8] -resolution = 768 -# threads used for resizing -threads = 4 -# source folder with images -folder = "raw" - if not os.path.isdir("images"): os.mkdir("images") -def process(fname): - global folder - global resolution +def parse_args(): + parser = argparse.ArgumentParser(description="Preprocess images") + parser.add_argument("-r", "--res", type=int, default=768, help="Target resolution") + parser.add_argument("-t", "--threads", type=int, default=4, help="No. of CPU threads to use") + parser.add_argument('--src', default="raw", help="Source folder with images") + return parser.parse_args() + +def process(fname, folder, resolution): src = os.path.join(folder, fname) md5 = hashlib.md5(open(src,'rb').read()).hexdigest() out = os.path.join("images", f"{md5}.png") @@ -28,22 +27,28 @@ def process(fname): img = img.resize((resolution,resolution), Image.LANCZOS) img.save(out) -def thread(queue, pbar): +def thread(queue, pbar, folder, resolution): while not queue.empty(): fname = queue.get() - process(fname) + process(fname, folder, resolution) queue.task_done() pbar.update() -files = os.listdir(folder) +args = parse_args() +files = os.listdir(args.src) pbar = tqdm(total=len(files),unit="img") queue = Queue() [queue.put(x) for x in files] -for _ in range(threads): +for _ in range(args.threads): Thread( target=thread, - args=(queue,pbar), + args=( + queue, + pbar, + args.src, + args.res, + ), daemon=True, ).start() diff --git a/train.py b/train.py index 4d2cd1f..a9de84f 100644 --- a/train.py +++ b/train.py @@ -2,20 +2,26 @@ 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 +from safetensors.torch import save_file, load_file from interposer import Interposer from vae import get_vae -# options -target_dev = "cuda" -target_steps = 500000 -save_every_n = 50000 -latent_src = "v1" -latent_dst = "xl" +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("-n", "--save_every_n", type=int, dest="save", default=50000, help="Save model/sample periodically") + parser.add_argument('--src', choices=["v1","xl"], required=True, help="Source latent format") + parser.add_argument('--dst', choices=["v1","xl"], required=True, help="Destination latent format") + parser.add_argument('--resume', help="Checkpoint to resume from") + args = parser.parse_args() + if args.src == args.dst: + parser.error("--src and --dst can't be the same") + return args class Latent: def __init__(self, md5, lat_src, lat_dst, dev): @@ -28,6 +34,14 @@ class Latent: self.src = torch.from_numpy(np.load(src)).to(dev) self.dst = torch.from_numpy(np.load(dst)).to(dev) +def load_latents(src, dst, dev): + print("Loading latents from disk") + latents = [] + for i in tqdm(os.listdir("images")): + md5 = os.path.splitext(i)[0] + latents.append(Latent(md5, src, dst, dev)) + return latents + vae = None def sample_decode(latent, filename, version): global vae @@ -47,21 +61,26 @@ def sample_decode(latent, filename, version): out.save(filename) if __name__ == "__main__": + args = parse_args() + target_dev = "cuda" + target_steps = args.steps + save_every_n = args.save + latent_src = args.src + latent_dst = args.dst + + latents = load_latents(latent_src, latent_dst, target_dev) + if not os.path.isdir("models"): os.mkdir("models") log = open(f"models/{latent_src}-to-{latent_dst}_interposer.csv", "w") - print("Loading latents from disk") - latents = [] - for i in tqdm(os.listdir("images")): - md5 = os.path.splitext(i)[0] - latents.append(Latent(md5, latent_src, latent_dst, target_dev)) - if os.path.isfile(f"test_{latent_src}.npy"): sample_latent = torch.from_numpy(np.load(f"test_{latent_src}.npy")).to(target_dev) else: sample_latent = random.choice(latents).src model = Interposer() + if args.resume: + model.load_state_dict(load_file(args.resume)) model.to(target_dev) criterion = torch.nn.MSELoss(size_average=False)