diff --git a/dataset.py b/dataset.py new file mode 100644 index 0000000..cb4f71d --- /dev/null +++ b/dataset.py @@ -0,0 +1,89 @@ +# Custom dataset to load encoded latents from disk. +# Files should contain latents as (1, C, H, W) or (C, H, W) +# Latents should be in their original format without scaling + + ######### Folder Layout ######### + # latents # + # |- test_v1_768px.npy <=eval # + # |- test_xl_768px.npy <=^ # + # |- v1_768px <= ver/res # + # | |- 000001.npy # + # | |- 000002.npy # + # | | ... # + # | |- 000999.npy # + # | \- 001000.npy # + # |- xl_768px # + # ... # + ################################# + +import os +import torch +import numpy as np +from tqdm import tqdm +from torch.utils.data import Dataset + +DEFAULT_ROOT = "latents" +ALLOWED_EXTS = [".npy"] + +class Shard: + """ + Shard to store groups of latents in + paths: List containing paths to latent encoded images + """ + def __init__(self, paths): + self.paths = paths + self.data = None + + def exists(self): + return all([os.path.isfile(x) for x in self.paths]) + + def get_data(self): + if self.data is not None: return self.data + return tuple([self.load_latent(x) for x in self.paths]) + + 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 + + def preload(self): + self.data = self.get_data() + +class LatentDataset(Dataset): + def __init__(self, specs, res=768, root=DEFAULT_ROOT, preload=False): + """ + Main dataset that returns list of requested images as (C, H, W) latents + specs: List of latent versions in the other to return them in + res: Native resolution of images (before latent encoding) + root: Path to folder with sorted files + preload: Load all files into memory on initialization + """ + print("Dataset: Parsing data from disk") + self.specs = specs + self.res = res + self.root = root + self.shards = [] + for fname in tqdm(os.listdir(f"{root}/{specs[0]}_{res}px")): + name, ext = os.path.splitext(fname) + if ext not in ALLOWED_EXTS: continue + shard = Shard([f"{root}/{x}_{res}px/{name}{ext}" for x in specs]) + if shard.exists(): + self.shards.append(shard) + + if preload: # cache to RAM + print("Dataset: Preloading data to system RAM") + [x.preload() for x in tqdm(self.shards)] + print(f"Dataset: OK, {len(self)} items") + + def __len__(self): + return len(self.shards) + + def __getitem__(self, index): + return self.shards[index].get_data() + + def get_eval(self): + shard = Shard([f"{self.root}/test_{x}_{self.res}px.npy" for x in self.specs]) + data = shard.get_data() if shard.exists() else self[0] + return tuple([x.unsqueeze(0).to(torch.float32) for x in data]) diff --git a/interposer.py b/interposer.py index e8db42b..1c97a53 100644 --- a/interposer.py +++ b/interposer.py @@ -1,56 +1,57 @@ import torch import torch.nn as nn -import numpy as np -class Block(nn.Module): - def __init__(self, size): +class ResBlock(nn.Module): + """Block with residuals""" + def __init__(self, ch): super().__init__() self.join = nn.ReLU() self.long = nn.Sequential( - nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), + nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(0.1), - nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), + nn.Conv2d(ch, ch, 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) + nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), ) def forward(self, x): - y = self.long(x) - z = self.join(y + x) - return z + return self.join(self.long(x) + x) -class Interposer(nn.Module): - def __init__(self): +class ExtractBlock(nn.Module): + """Increase no. of channels by [out/in]""" + def __init__(self, ch_in, ch_out): super().__init__() - self.chan = 4 # in/out channels - self.hid = 128 + self.join = nn.ReLU() + self.short = nn.Conv2d(ch_in, ch_out, kernel_size=3, stride=1, padding=1) + self.long = nn.Sequential( + nn.Conv2d( ch_in, ch_out, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1), + nn.LeakyReLU(0.1), + nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1), + nn.Dropout(0.1) + ) + def forward(self, x): + return self.join(self.long(x) + self.short(x)) - # 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 +class InterposerModel(nn.Module): + """Main neural network""" + def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0): + super().__init__() + self.scale = scale + self.ch_in = ch_in + self.ch_out = ch_out + self.ch_mid = ch_mid + + self.head = ExtractBlock(self.ch_in, self.ch_mid) 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) + nn.Upsample(scale_factor=self.scale, mode="nearest"), + ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), + ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), + ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ) + self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1) def forward(self, x): - y = self.head_join( - self.head_long(x)+ - self.head_short(x) - ) + y = self.head(x) z = self.core(y) return self.tail(z) diff --git a/log_loss.py b/log_loss.py index acf17ad..9fd0da9 100644 --- a/log_loss.py +++ b/log_loss.py @@ -5,20 +5,36 @@ import matplotlib.pyplot as plt files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")] train_loss = {} eval_loss = {} - +lr_vals = {} +fskip = 0 + +offsets = { # offset to display resumed training runs +} +sep = ".csv" +rep = "_interposer" +model = "Latent Interposer" + def process_lines(lines): global train_loss global eval_loss - name = fp.split("/")[1].split("_")[0] + name = fp.split("/")[1] + print(name) + if sep: name = name.split(sep)[0] + if rep: name = name.replace(rep,"") 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], + [math.log(float(x[1])+1e-10) 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], + [math.log(float(x[2])+1e-10) for x in vals], + ) + if len(vals[0]) >= 4: + lr_vals[name] = ( + [int(x[0]) for x in vals], + [float(x[3]) for x in vals], ) # https://stackoverflow.com/a/49357445 @@ -31,18 +47,26 @@ def smooth(scalars, weight): last = smoothed_val return smoothed -def plot(data, fname): +def plot(data, fname, title=None, smw=0.9): fig, ax = plt.subplots() + plt.tight_layout() ax.grid() + dmax = 0 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') + data = [x + offsets[name] for x in val[0]] if name in offsets.keys() else val[0] + dmax = max(dmax, round(data[-1],10000)) + sval = val[1][:fskip] + smooth(val[1][fskip:], smw) # skip first N + ax.plot(data, sval, label=name) + ax.set_xticks([dmax//10*x for x in range(10)]) + plt.legend(loc="lower left", bbox_to_anchor=(0.00, -0.20), ncol=5) + if title: plt.title(title) + plt.savefig(fname, 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") +if __name__ == "__main__": + for fp in files: + with open(fp) as f: + lines = f.readlines() + process_lines(lines) + plot(train_loss, "loss.png", f"{model} Training loss", 0.2) + plot(eval_loss, "loss-eval.png", f"{model} Eval. loss", 0.7) + plot(lr_vals, "loss-lr.png", f"{model} Learning rate", 0.0) diff --git a/train.py b/train.py index 832bde1..afb8541 100644 --- a/train.py +++ b/train.py @@ -1,29 +1,31 @@ import os import torch -import torch.nn as nn -import numpy as np import argparse -from PIL import Image from tqdm import tqdm -from safetensors.torch import save_file, load_file -from torch.utils.data import DataLoader, Dataset +from torch.utils.data import DataLoader +from safetensors.torch import load_file -from interposer import Interposer -from vae import get_vae +from interposer import InterposerModel as Model +from dataset import LatentDataset +from utils import ModelWrapper torch.backends.cudnn.benchmark = True torch.manual_seed(0) +TARGET_DEV = "cuda" + 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=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("-s", "--steps", type=int, default=500000, help="No. of training steps") + parser.add_argument("-b", "--batch", type=int, default= 1, help="Batch size") + parser.add_argument("-n", "--nsave", type=int, default= 50000, help="Save model/sample periodically") + parser.add_argument('--rev', default="v4.0-rc1", help="Revision/log ID") 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('--lr', default="1e-4", help="Learning rate") + parser.add_argument('--lrskip', type=int, default=0, help="Constant lr for first N steps") + parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler") 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") @@ -33,120 +35,28 @@ def parse_args(): parser.error("--lr must be a valid float eg. 0.001 or 1e-3") return args -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) - -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" - resolution = 768 - dataset = LatentDataset(resolution, args.src, args.dst) + dataset = LatentDataset([args.src, args.dst]) loader = DataLoader( dataset, - batch_size=args.bs, - shuffle=True, - num_workers=0, - # num_workers=4, - # persistent_workers=True, + batch_size = args.batch, + shuffle = True, + drop_last = True, + pin_memory = False, + # 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, - ) - - os.makedirs("models", exist_ok=True) - log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w") - - model = Interposer() - + model = Model() # TODO: handle scale factor/channels for non-sd VAEs 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), + optimizer, T_max = int(args.steps/args.batch), ) else: print("Using LinearLR") @@ -154,23 +64,35 @@ if __name__ == "__main__": optimizer, start_factor = 0.1, end_factor = 1.0, - total_iters = int(5000/args.bs), + total_iters = int(5000/args.batch), ) 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" )) + optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0] else: - model.to(target_dev) + model.to(TARGET_DEV) + + wrapper = ModelWrapper( # model wrapper for saving/eval/etc + name = f"{args.src}-to-{args.dst}_interposer-{args.rev}", + specs = [args.src, args.dst], + model = model, + evals = dataset.get_eval(), + device = TARGET_DEV, + criterion = criterion, + optimizer = optimizer, + scheduler = scheduler, + ) 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) + 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 @@ -179,19 +101,15 @@ if __name__ == "__main__": optimizer.zero_grad() loss.backward() optimizer.step() - scheduler.step() + if progress.n >= args.lrskip: scheduler.step() # 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) + progress.update(args.batch) + wrapper.log_step(loss.data.item(), progress.n) + if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0: + wrapper.save_model(step=progress.n) if progress.n >= args.steps: break progress.close() - - # save final output - 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() + wrapper.save_model(epoch="") # final save + wrapper.close() diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..ed550aa --- /dev/null +++ b/utils.py @@ -0,0 +1,123 @@ +# +# This file just has all the random saving/logging/eval related code +# +import os +import torch +from tqdm import tqdm +from diffusers import AutoencoderKL +from safetensors.torch import save_file +from torchvision.utils import save_image + +LOSS_MEMORY = 500 +LOG_EVERY_N = 500 +SAVE_FOLDER = "models" + +class ModelWrapper: + def __init__(self, name, specs, model, optimizer, criterion, scheduler, device="cpu", evals=[None,None], stdout=True): + self.name = name + self.specs = specs + self.losses = [] + + self.model = model + self.optimizer = optimizer + self.criterion = criterion + self.scheduler = scheduler + + self.device = device + self.vae = self.get_vae(self.specs[1], fp16=True) + self.eval_src = evals[0] + self.eval_dst = evals[1] + + os.makedirs(SAVE_FOLDER, exist_ok=True) + self.csvlog = open(f"{SAVE_FOLDER}/{self.name}.csv", "w") + self.stdout = stdout + + def log_step(self, loss, step=None): + self.losses.append(loss) + step = step if step else len(self.losses) + if step % LOG_EVERY_N == 0: + self.log_main(step) + + def log_main(self, step=None): + lr = float(self.scheduler.get_last_lr()[0]) + avg = sum(self.losses[-LOSS_MEMORY:])/LOSS_MEMORY + evl = self.eval_model()[0] + if self.stdout: + tqdm.write(f"{str(step):<10} {avg:.4e}|{evl:.4e} @ {lr:.4e}") + if self.csvlog: + self.csvlog.write(f"{step},{avg},{evl},{lr}\n") + self.csvlog.flush() + + def eval_model(self): + with torch.no_grad(): + pred = self.model(self.eval_src.to(self.device)) + loss = self.criterion(pred, self.eval_dst.to(self.device)) + return loss, pred + + def save_model(self, step=None, epoch=None): + step = step if step else len(self.losses) + if epoch is None and step >= 10**6: + epoch = f"_e{round(step/10**6,2)}M" + elif epoch is None: + epoch = f"_e{round(step/10**3)}K" + output_name = f"./{SAVE_FOLDER}/{self.name}{epoch}" + if self.vae: + out = self.eval_model()[1] + img = self.vae_decode(out).detach() + save_image(img, f"{output_name}.png") + torch.cuda.empty_cache() + save_file(self.model.state_dict(), f"{output_name}.safetensors") + torch.save(self.optimizer.state_dict(), f"{output_name}.optim.pth") + + def close(self): + del self.vae + self.csvlog.close() + + def vae_decode(self, latent): + latent = latent.to(torch.float16).to("cuda") + out = self.vae.decode(latent).sample + out = out.float().to(latent.device) + out = torch.clamp(out, min=-1.0, max=1.0) + return ((out + 1.0) / 2.0) + + def get_vae(self, 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 == "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: + raise NotImplementedError(f"Unknown VAE version '{version}'") + + # save VRAM + vae.to(dtype).to("cuda") + vae.decoder.eval() + vae.set_use_memory_efficient_attention_xformers(True) + vae.enable_xformers_memory_efficient_attention() + vae.enable_gradient_checkpointing() + del vae.encoder + return vae diff --git a/vae.py b/vae.py deleted file mode 100644 index 0bd4f9e..0000000 --- a/vae.py +++ /dev/null @@ -1,48 +0,0 @@ -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