diff --git a/.gitignore b/.gitignore index 8c5cb25..80720af 100644 --- a/.gitignore +++ b/.gitignore @@ -1,12 +1,15 @@ raw/ images/ -latent_*/ +latents/ +latents vae/ models/ +other/ test.py *.png *.zip *.npy +*.pth *.ckpt *.safetensors diff --git a/interposer.py b/interposer.py index 6714eaf..e8db42b 100644 --- a/interposer.py +++ b/interposer.py @@ -2,30 +2,55 @@ import torch import torch.nn as nn import numpy as np +class Block(nn.Module): + def __init__(self, size): + super().__init__() + self.join = nn.ReLU() + self.long = nn.Sequential( + nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), + nn.Dropout(0.2) + ) + def forward(self, x): + y = self.long(x) + z = self.join(y + x) + return z + class Interposer(nn.Module): def __init__(self): super().__init__() - - # it looks like a spaceship if you squint :D - module_list = [ - #############) - #############) - #||# - #||# - nn.Conv2d(4, 32, kernel_size=5, padding=2), - nn.ReLU(), - nn.Conv2d(32, 128, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(128, 32, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(32, 4, kernel_size=5, padding=2), - #||# - #||# - #############) - #############) - ] + self.chan = 4 # in/out channels + self.hid = 128 - self.sequential = nn.Sequential(*module_list) + # expand channels + self.head_join = nn.ReLU() + self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1) + self.head_long = nn.Sequential( + nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1), + ) + # not sure if this is how residuals work + self.core = nn.Sequential( + Block(self.hid), + Block(self.hid), + Block(self.hid), + ) + # reduce channels + self.tail = nn.Sequential( + nn.ReLU(), + nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1) + ) - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.sequential(x) + def forward(self, x): + y = self.head_join( + self.head_long(x)+ + self.head_short(x) + ) + z = self.core(y) + return self.tail(z) diff --git a/log_loss.py b/log_loss.py index 2c4a055..acf17ad 100644 --- a/log_loss.py +++ b/log_loss.py @@ -15,7 +15,7 @@ def process_lines(lines): [int(x[0]) for x in vals], [math.log(float(x[1])) for x in vals], ) - if len(vals[0]) == 3: + if len(vals[0]) >= 3: eval_loss[name] = ( [int(x[0]) for x in vals], [math.log(float(x[2])) for x in vals], diff --git a/preprocess_images.py b/preprocess_images.py deleted file mode 100644 index d91153b..0000000 --- a/preprocess_images.py +++ /dev/null @@ -1,65 +0,0 @@ -import os -import hashlib -import argparse -from tqdm import tqdm -from PIL import Image -from queue import Queue -from threading import Thread - -if not os.path.isdir("images"): - os.mkdir("images") - -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") - if os.path.isfile(out): - return - img = Image.open(src) - img = img.convert('RGB') - target = (resolution, resolution) - if min(img.height, img.width) < 256: - return - - if img.width > img.height: - target = (int(img.width/img.height*resolution), resolution) - elif img.height > img.width: - target = (resolution, int(img.height/img.width*resolution)) - img = img.resize(target, Image.LANCZOS) - img = img.crop([0,0,resolution,resolution]) - img.save(out) - -def thread(queue, pbar, folder, resolution): - while not queue.empty(): - fname = queue.get() - try: process(fname, folder, resolution) - except: pass - queue.task_done() - pbar.update() - -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(args.threads): - Thread( - target=thread, - args=( - queue, - pbar, - args.src, - args.res, - ), - daemon=True, - ).start() - -queue.join() diff --git a/preprocess_latents.py b/preprocess_latents.py deleted file mode 100644 index 8716f1c..0000000 --- a/preprocess_latents.py +++ /dev/null @@ -1,51 +0,0 @@ -import os -import torch -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 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 process_folder(vae, v): - if not os.path.isdir(f"latent_{v}"): - os.mkdir(f"latent_{v}") - - vae.to("cuda") - for i in tqdm(os.listdir("images")): - src = os.path.join("images", i) - img = Image.open(src) - dst = os.path.join(f"latent_{v}", f"{os.path.splitext(i)[0]}.npy") - latent = encode(vae, img) - np.save(dst, latent) - vae.to("cpu") - -def run_v1(file_path=None): - vae = get_vae("v1", file_path) - process_folder(vae, "v1") - del vae - -def run_v2(file_path=None): - vae = get_vae("v2", file_path) - process_folder(vae, "v2") - del vae - -def run_xl(file_path=None): - vae = get_vae("xl", file_path) - process_folder(vae, "xl") - del vae - -if __name__ == "__main__": - # run_v1("./vae/ft-mse-840000.ckpt") # probably doesn't reflect internal SD latent - run_v1() - # run_v2() # v2 and v1 share a latent space - run_xl("./vae/sdxl_v0.9.safetensors") # 1.0 has artifacts diff --git a/train.py b/train.py index 09ed802..832bde1 100644 --- a/train.py +++ b/train.py @@ -3,23 +3,27 @@ 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 torch.utils.data import DataLoader, Dataset from interposer import Interposer from vae import get_vae +torch.backends.cudnn.benchmark = True +torch.manual_seed(0) + 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('--bs', type=int, default=4, help="Batch size") + parser.add_argument('--lr', default="1e-4", help="Learning rate") 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") + parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR") args = parser.parse_args() if args.src == args.dst: parser.error("--src and --dst can't be the same") @@ -29,25 +33,6 @@ def parse_args(): parser.error("--lr must be a valid float eg. 0.001 or 1e-3") return args -class Latent: - def __init__(self, md5, lat_src, lat_dst, dev): - if lat_src == "v1": src = os.path.join("latent_v1", f"{md5}.npy") - if lat_src == "xl": src = os.path.join("latent_xl", f"{md5}.npy") - - if lat_dst == "v1": dst = os.path.join("latent_v1", f"{md5}.npy") - if lat_dst == "xl": dst = os.path.join("latent_xl", 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(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 @@ -66,66 +51,147 @@ def sample_decode(latent, filename, version): out = Image.fromarray(out) out.save(filename) +def get_eval_data(dataset, src_path, dst_path, target_dev): + if os.path.isfile(src_path) and os.path.isfile(dst_path): + src = LatentDataset.load_latent(None, src_path) + dst = LatentDataset.load_latent(None, dst_path) + else: + src = dataset[0][0] + dst = dataset[0][1] + src = src.float().to(target_dev).unsqueeze(0) + dst = dst.float().to(target_dev).unsqueeze(0) + return(src, dst) + +def eval_model(step, model, criterion, scheduler, src, dst): + with torch.no_grad(): + t_pred = model(src) + t_loss = criterion(t_pred, dst) + tqdm.write(f"{str(step):<10} {loss.data.item():.4e}|{t_loss.data.item():.4e} @ {float(scheduler.get_last_lr()[0]):.4e}") + log.write(f"{step},{loss.data.item()},{t_loss.data.item()},{float(scheduler.get_last_lr()[0])}\n") + log.flush() + +def save_model(step, model, optim, lat, src, dst): + with torch.no_grad(): + out = model(lat) + output_name = f"./models/{src}-to-{dst}_interposer_e{round(step/1000)}k" + sample_decode(out, f"{output_name}.png", dst) + save_file(model.state_dict(), f"{output_name}.safetensors") + torch.save(optim.state_dict(), f"{output_name}.optim.pth") + +class LatentDataset(Dataset): + class Shard: + def __init__(self, root, fname, res, src, dst): + self.fname = fname + self.src_path = f"{root}/{src}_{res}px/{fname}.npy" + self.dst_path = f"{root}/{dst}_{res}px/{fname}.npy" + + def __init__(self, res, src, dst, root="latents"): + print("Loading latents from disk") + self.latents = [] + for i in tqdm(os.listdir(f"{root}/{src}_{res}px")): + fname, ext = os.path.splitext(i) + assert ext == ".npy" + s = self.Shard(root, fname, res, src, dst) + if os.path.isfile(s.src_path) and os.path.isfile(s.dst_path): + self.latents.append(s) + + def __len__(self): + return len(self.latents) + + def __getitem__(self, index): + s = self.latents[index] + src = self.load_latent(s.src_path) + dst = self.load_latent(s.dst_path) + return (src, dst) + + def load_latent(self, path): + lat = torch.from_numpy(np.load(path)) + if lat.shape[0] == 1: + lat = torch.squeeze(lat, 0) + assert not torch.isnan(torch.sum(lat.float())) + return lat + if __name__ == "__main__": args = parse_args() target_dev = "cuda" - latent_src = args.src - latent_dst = args.dst + resolution = 768 - latents = load_latents(latent_src, latent_dst, target_dev) + dataset = LatentDataset(resolution, args.src, args.dst) + loader = DataLoader( + dataset, + batch_size=args.bs, + shuffle=True, + num_workers=0, + # num_workers=4, + # persistent_workers=True, + ) + eval_src, eval_dst = get_eval_data( + dataset, + f"latents/test_{args.src}_{resolution}px.npy", + f"latents/test_{args.dst}_{resolution}px.npy", + target_dev, + ) - if not os.path.isdir("models"): os.mkdir("models") - log = open(f"models/{latent_src}-to-{latent_dst}_interposer.csv", "w") - - if os.path.isfile(f"test_{latent_src}.npy") and os.path.isfile(f"test_{latent_dst}.npy"): - ss_latent = torch.from_numpy(np.load(f"test_{latent_src}.npy")).to(target_dev) - st_latent = torch.from_numpy(np.load(f"test_{latent_dst}.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) + os.makedirs("models", exist_ok=True) + log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w") model = Interposer() + + criterion = torch.nn.L1Loss() + optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr)) + # import bitsandbytes as bnb + # optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=float(args.lr)) + + scheduler = None + if args.cosine: + print("Using CosineAnnealingLR") + scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( + optimizer, T_max = int(args.steps/args.bs), + ) + else: + print("Using LinearLR") + scheduler = torch.optim.lr_scheduler.LinearLR( + optimizer, + start_factor = 0.1, + end_factor = 1.0, + total_iters = int(5000/args.bs), + ) + if args.resume: model.load_state_dict(load_file(args.resume)) - model.to(target_dev) + model.to(target_dev) + optimizer.load_state_dict(torch.load( + f"{os.path.splitext(args.resume)[0]}.optim.pth" + )) + else: + model.to(target_dev) - criterion = torch.nn.MSELoss(size_average=False) - optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) + progress = tqdm(total=args.steps) + while progress.n < args.steps: + for src, dst in loader: + src = src.to(target_dev) + dst = dst.to(target_dev) + with torch.cuda.amp.autocast(): + y_pred = model(src) # forward + loss = criterion(y_pred, dst) # loss - 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) + # backward + optimizer.zero_grad() + loss.backward() + optimizer.step() + scheduler.step() - y_pred = model(src) # forward - loss = criterion(y_pred, dst) # loss + # eval/save + progress.update(args.bs) + if progress.n % (1000 + 1000%args.bs) == 0: + eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst) + if progress.n % (args.save + args.save%args.bs) == 0: + save_model(progress.n, model, optimizer, eval_src, args.src, args.dst) + if progress.n >= args.steps: + break + progress.close() - # 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_src}-to-{latent_dst}_interposer_e{step/1000}k" - sample_decode(out, f"{output_name}.png", latent_dst) - save_file(model.state_dict(), f"{output_name}.safetensors") # save final output - output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{args.steps/1000}k" - sample_decode(out, f"{output_name}.png", "v1") - save_file(model.state_dict(), f"{output_name}.safetensors") + eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst) + save_model(progress.n, model, optimizer, eval_src, args.src, args.dst) log.close()