Switch to args instead of hardcoding variables
This commit is contained in:
+20
-15
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user