From 92f9b64e8bc0b314c4f0503721ba7f67934e8338 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Sat, 11 Nov 2023 22:19:42 +0100 Subject: [PATCH 1/4] Snapshot --- dataset.py | 89 +++++++++++++++++++++++++ interposer.py | 75 ++++++++++----------- log_loss.py | 54 ++++++++++----- train.py | 178 ++++++++++++++------------------------------------ utils.py | 123 ++++++++++++++++++++++++++++++++++ vae.py | 48 -------------- 6 files changed, 337 insertions(+), 230 deletions(-) create mode 100644 dataset.py create mode 100644 utils.py delete mode 100644 vae.py 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 From 8ccc8208b24e0936dbff7b3c90075f2e0873f743 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 18 Mar 2024 23:29:15 +0100 Subject: [PATCH 2/4] Version 4 --- comfy_latent_interposer.py | 189 ++++++++++++++---------- config/ca-to-v1.yaml | 40 +++++ config/ca-to-xl.yaml | 40 +++++ config/v1-to-xl.yaml | 40 +++++ config/xl-to-v1.yaml | 40 +++++ dataset.py | 94 ++++++------ interposer.py | 22 +-- log_loss.py | 72 --------- train.py | 292 +++++++++++++++++++++++++++---------- vae.py | 114 +++++++++++++++ 10 files changed, 666 insertions(+), 277 deletions(-) create mode 100644 config/ca-to-v1.yaml create mode 100644 config/ca-to-xl.yaml create mode 100644 config/v1-to-xl.yaml create mode 100644 config/xl-to-v1.yaml delete mode 100644 log_loss.py create mode 100644 vae.py diff --git a/comfy_latent_interposer.py b/comfy_latent_interposer.py index 6306187..ecd7451 100644 --- a/comfy_latent_interposer.py +++ b/comfy_latent_interposer.py @@ -4,111 +4,154 @@ import torch.nn as nn from safetensors.torch import load_file from huggingface_hub import hf_hub_download +# v1 = Stable Diffusion 1.x +# xl = Stable Diffusion Extra Large (SDXL) +# cc = Stable Cascade (Stage C) [not used] +# ca = Stable Cascade (Stage A/B) +config = { + "v1-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, + "xl-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, + "ca-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12}, + "ca-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12}, +} -class Interposer(nn.Module): - """ - Basic NN layout, ported from: - https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py - """ - version = 3.1 # network revision - def __init__(self): +class ResBlock(nn.Module): + """Block with residuals""" + def __init__(self, ch): super().__init__() - self.chan = 4 - self.hid = 128 + self.join = nn.ReLU() + self.norm = nn.BatchNorm2d(ch) + self.long = nn.Sequential( + nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), + nn.SiLU(), + nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), + nn.SiLU(), + nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), + nn.Dropout(0.1) + ) + def forward(self, x): + x = self.norm(x) + return self.join(self.long(x) + x) - 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), +class ExtractBlock(nn.Module): + """Increase no. of channels by [out/in]""" + def __init__(self, ch_in, ch_out): + super().__init__() + 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.SiLU(), + nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1), + nn.SiLU(), + 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)) + +class InterposerModel(nn.Module): + """ + NN layout, ported from: + https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py + """ + def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, blocks=12): + super().__init__() + self.ch_in = ch_in + self.ch_out = ch_out + self.ch_mid = ch_mid + self.blocks = blocks + self.scale = scale + + self.head = ExtractBlock(self.ch_in, self.ch_mid) self.core = nn.Sequential( - Block(self.hid), - Block(self.hid), - Block(self.hid), - ) - 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) for _ in range(blocks)], + nn.BatchNorm2d(self.ch_mid), + nn.SiLU(), ) + 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) -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), - ) - def forward(self, x): - y = self.long(x) - z = self.join(y + x) - return z - - -class LatentInterposer: +class ComfyLatentInterposer: + """Custom node""" def __init__(self): - pass + self.version = 4.0 # network revision + self.loaded = None # current model name + self.model = None # current model @classmethod def INPUT_TYPES(s): return { "required": { "samples": ("LATENT", ), - "latent_src": (["v1", "xl"],), - "latent_dst": (["v1", "xl"],), + "latent_src": (["v1", "xl", "ca"],), + "latent_dst": (["v1", "xl", "ca"],), } } RETURN_TYPES = ("LATENT",) - FUNCTION = "convert" - CATEGORY = "latent" + FUNCTION = "convert" + CATEGORY = "latent" + TITLE = "Latent Interposer" + + def get_model_path(self, model_name): + fname = f"{model_name}_interposer-v{self.version}.safetensors" + path = os.path.join(os.path.dirname(os.path.realpath(__file__)),"models") + + # local path: [models/xl-to-v1_interposer-v4.2.safetensors] + if os.path.isfile(os.path.join(path, fname)): + print("LatentInterposer: Using local model") + return os.path.join(path, fname) + + # local path: [models/v4.2/xl-to-v1_interposer-v4.2.safetensors] + if os.path.isfile(os.path.join(path, os.path.join(f"v{self.version}", fname))): + print("LatentInterposer: Using local model") + return os.path.join(path, os.path.join(f"v{self.version}", fname)) + + # huggingface hub fallback + print("LatentInterposer: Using HF Hub model") + return str(hf_hub_download( + repo_id = "city96/SD-Latent-Interposer", + subfolder = f"v{self.version}", + filename = fname, + )) def convert(self, samples, latent_src, latent_dst): + samples = samples.copy() if latent_src == latent_dst: return (samples,) - model = Interposer() - model.eval() - filename = f"{latent_src}-to-{latent_dst}_interposer-v{model.version}.safetensors" - local = os.path.join( - os.path.join(os.path.dirname(os.path.realpath(__file__)),"models"), - filename - ) - if os.path.isfile(local): - print("LatentInterposer: Using local model") - weights = local - else: - print("LatentInterposer: Using HF Hub model") - weights = str(hf_hub_download( - repo_id="city96/SD-Latent-Interposer", - filename=filename) - ) + model_name = f"{latent_src}-to-{latent_dst}" + if model_name not in config: + raise ValueError(f"No model exists for this conversion! ({model_name})") + + # only reload if changed + if self.loaded != model_name or self.model is None: + # load/init model + path = self.get_model_path(model_name) + model = InterposerModel(**config[model_name]) + model.eval() + model.load_state_dict(load_file(path)) + # keep for later runs + self.model = model + self.loaded = model_name - model.load_state_dict(load_file(weights)) lt = samples["samples"] - lt = model(lt) - del model - return ({"samples": lt},) + with torch.no_grad(): + # force FP32, always run on CPU + lt = self.model(lt.cpu().float()).to(lt.device).to(lt.dtype) + samples["samples"] = lt + return (samples,) NODE_CLASS_MAPPINGS = { - "LatentInterposer": LatentInterposer, + "LatentInterposer": ComfyLatentInterposer, } NODE_DISPLAY_NAME_MAPPINGS = { - "LatentInterposer": "Latent Interposer" + "LatentInterposer": ComfyLatentInterposer.TITLE, } diff --git a/config/ca-to-v1.yaml b/config/ca-to-v1.yaml new file mode 100644 index 0000000..24ba67c --- /dev/null +++ b/config/ca-to-v1.yaml @@ -0,0 +1,40 @@ +steps: 20000 +batch: 48 +fconst: 0 +cosine: False +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 1.4 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: ca # Stable Cascade Stage A + dst: v1 # Stable Diffusion 1.x + rev: "v4.0-rc16" + args: + scale: 0.5 + ch_in: 4 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/ca_256px_combined.bin" + dst: "./latents/v1_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_ca_768px.npy" + dst: "./latents/test_eru/test_v1_768px.npy" + aux: + src: "./latents/test_bga/test_ca_768px.npy" + dst: "./latents/test_bga/test_v1_768px.npy" diff --git a/config/ca-to-xl.yaml b/config/ca-to-xl.yaml new file mode 100644 index 0000000..6fe47ad --- /dev/null +++ b/config/ca-to-xl.yaml @@ -0,0 +1,40 @@ +steps: 20000 +batch: 48 +fconst: 0 +cosine: False +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 1.4 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: ca # Stable Cascade Stage A + dst: xl # Stable Diffusion Extra Large + rev: "v4.0-rc16" + args: + scale: 0.5 + ch_in: 4 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/ca_256px_combined.bin" + dst: "./latents/xl_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_ca_768px.npy" + dst: "./latents/test_eru/test_xl_768px.npy" + aux: + src: "./latents/test_bga/test_ca_768px.npy" + dst: "./latents/test_bga/test_xl_768px.npy" diff --git a/config/v1-to-xl.yaml b/config/v1-to-xl.yaml new file mode 100644 index 0000000..61ad3f4 --- /dev/null +++ b/config/v1-to-xl.yaml @@ -0,0 +1,40 @@ +steps: 50000 +batch: 128 +fconst: 35000 +cosine: True +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 1.4 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: v1 # Stable Diffusion 1.x + dst: xl # Stable Diffusion Extra Large + rev: "v4.0-rc15" + args: + scale: 1.0 + ch_in: 4 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/v1_256px_combined.bin" + dst: "./latents/xl_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_v1_768px.npy" + dst: "./latents/test_eru/test_xl_768px.npy" + aux: + src: "./latents/test_bga/test_v1_768px.npy" + dst: "./latents/test_bga/test_xl_768px.npy" diff --git a/config/xl-to-v1.yaml b/config/xl-to-v1.yaml new file mode 100644 index 0000000..0d8800c --- /dev/null +++ b/config/xl-to-v1.yaml @@ -0,0 +1,40 @@ +steps: 50000 +batch: 128 +fconst: 30000 +cosine: True +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 0.0 +b_loss_weight: 0.0 +h_loss_weight: 0.0 +save_image: 1000 +eval_model: 10 + +model: + src: xl # Stable Diffusion Extra Large + dst: v1 # Stable Diffusion 1.x + rev: "v4.0-rc16" + args: + scale: 1.0 + ch_in: 4 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/xl_256px_combined.bin" + dst: "./latents/v1_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_xl_768px.npy" + dst: "./latents/test_eru/test_v1_768px.npy" + aux: + src: "./latents/test_bga/test_xl_768px.npy" + dst: "./latents/test_bga/test_v1_768px.npy" diff --git a/dataset.py b/dataset.py index cb4f71d..3cb2456 100644 --- a/dataset.py +++ b/dataset.py @@ -1,45 +1,37 @@ -# 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 FileLatentDataset(Dataset): + def __init__(self, src_file, dst_file, device="cpu", dtype=torch.float16): + assert os.path.isfile(src_file), f"src bin missing! ({src_file})" + assert os.path.isfile(dst_file), f"dst bin missing! ({dst_file})" + self.src_data = torch.load(src_file).to(dtype).to(device) + self.dst_data = torch.load(dst_file).to(dtype).to(device) + assert self.src_data.shape[0] == self.dst_data.shape[0], "Data size mismatch!" + + def __len__(self): + return self.src_data.shape[0] + + def __getitem__(self, index): + return { + "src": self.src_data[index].float(), + "dst": self.dst_data[index].float(), + } 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]) + return all([os.path.isfile(x) for x in self.paths.values()]) def get_data(self): if self.data is not None: return self.data - return tuple([self.load_latent(x) for x in self.paths]) + return {k:self.load_latent(v) for k,v in self.paths.items()} def load_latent(self, path): lat = torch.from_numpy(np.load(path)) @@ -52,29 +44,36 @@ class Shard: 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 - """ + def __init__(self, src_root, dst_root, preload=True): + assert os.path.isdir(src_root), f"Source folder missing! ({src_root})" + assert os.path.isdir(dst_root), f"Destination folder missing! ({dst_root})" + print("Dataset: Parsing data from disk") - self.specs = specs - self.res = res - self.root = root + fnames = list( + set(os.listdir(src_root)).intersection( + set(os.listdir(dst_root))) + ) + assert len(fnames) > 0, "Source/destination have no overlapping files" + self.shards = [] - for fname in tqdm(os.listdir(f"{root}/{specs[0]}_{res}px")): + for fname in tqdm(fnames): + src_path = os.path.join(src_root, fname) + dst_path = os.path.join(dst_root, fname) 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 ext not in [".npy"]: + continue + shard = Shard({ + "src": src_path, + "dst": dst_path, + }) if shard.exists(): self.shards.append(shard) + assert len(self.shards) > 0, "No valid files found." 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): @@ -83,7 +82,14 @@ class LatentDataset(Dataset): 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]) +def load_evals(evals): + data = {} + for name, paths in evals.items(): + shard = Shard(paths) + assert shard.exists(), f"Eval data missing ({name})" + data[name] = {} + for k, v in shard.get_data().items(): + if len(v.shape) == 3: + v = v.unsqueeze(0) + data[name][k] = v.float() + return data diff --git a/interposer.py b/interposer.py index 1c97a53..a7288f0 100644 --- a/interposer.py +++ b/interposer.py @@ -6,14 +6,17 @@ class ResBlock(nn.Module): def __init__(self, ch): super().__init__() self.join = nn.ReLU() + self.norm = nn.BatchNorm2d(ch) self.long = nn.Sequential( nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), - nn.LeakyReLU(0.1), + nn.SiLU(), nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), - nn.LeakyReLU(0.1), + nn.SiLU(), nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), + nn.Dropout(0.1) ) def forward(self, x): + x = self.norm(x) return self.join(self.long(x) + x) class ExtractBlock(nn.Module): @@ -24,9 +27,9 @@ class ExtractBlock(nn.Module): 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.SiLU(), nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1), - nn.LeakyReLU(0.1), + nn.SiLU(), nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1), nn.Dropout(0.1) ) @@ -35,19 +38,20 @@ class ExtractBlock(nn.Module): class InterposerModel(nn.Module): """Main neural network""" - def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0): + def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, blocks=12): super().__init__() - self.scale = scale self.ch_in = ch_in self.ch_out = ch_out self.ch_mid = ch_mid + self.blocks = blocks + self.scale = scale self.head = ExtractBlock(self.ch_in, self.ch_mid) self.core = nn.Sequential( 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), + *[ResBlock(self.ch_mid) for _ in range(blocks)], + nn.BatchNorm2d(self.ch_mid), + nn.SiLU(), ) self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1) diff --git a/log_loss.py b/log_loss.py deleted file mode 100644 index 9fd0da9..0000000 --- a/log_loss.py +++ /dev/null @@ -1,72 +0,0 @@ -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 = {} -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] - 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])+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])+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 -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, title=None, smw=0.9): - fig, ax = plt.subplots() - plt.tight_layout() - ax.grid() - dmax = 0 - for name, val in data.items(): - 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') - -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 afb8541..f2e6e5a 100644 --- a/train.py +++ b/train.py @@ -1,115 +1,249 @@ import os +import yaml import torch import argparse from tqdm import tqdm from torch.utils.data import DataLoader -from safetensors.torch import load_file +from safetensors.torch import save_file, load_file -from interposer import InterposerModel as Model -from dataset import LatentDataset -from utils import ModelWrapper +from interposer import InterposerModel +from dataset import LatentDataset, FileLatentDataset, load_evals +from vae import load_vae 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("-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("--config", help="Config for training") args = parser.parse_args() - if args.src == args.dst: - parser.error("--src and --dst can't be the same") - try: - float(args.lr) - except: - parser.error("--lr must be a valid float eg. 0.001 or 1e-3") - return args + with open(args.config) as f: + conf = yaml.safe_load(f) + args.dataset = argparse.Namespace(**conf.pop("dataset")) + args.model = argparse.Namespace(**conf.pop("model")) + return argparse.Namespace(**vars(args), **conf) + +def eval_images(model, vae, evals): + preds = eval_model(model, evals, loss=False) + out = {} + for name, pred in preds.items(): + images = vae.decode(pred).cpu().float() + # for image in images: # eval isn't batched + out[f"eval/{name}"] = images[0] + return out + +def eval_model(model, evals, loss=True): + model.eval() + preds = {} + losses = [] + for name, data in evals.items(): + src = data["src"].to(args.device) + dst = data["dst"].to(args.device) + with torch.no_grad(): + pred = model(src) + if loss: + loss = torch.nn.functional.l1_loss(dst, pred) + losses.append(loss) + else: + preds[name] = pred + model.train() + if loss: + return (sum(losses) / len(losses)).data.item() + else: + return preds + +# from pytorch GAN tutorial +def weights_init(m): + classname = m.__class__.__name__ + if classname.find('Conv') != -1: + torch.nn.init.normal_(m.weight.data, 0.0, 0.02) + elif classname.find('BatchNorm') != -1: + torch.nn.init.normal_(m.weight.data, 1.0, 0.02) + torch.nn.init.constant_(m.bias.data, 0) if __name__ == "__main__": args = parse_args() + base_name = f"models/{args.model.src}-to-{args.model.dst}_interposer-{args.model.rev}" - dataset = LatentDataset([args.src, args.dst]) + # dataset + if os.path.isfile(args.dataset.src): + dataset = FileLatentDataset( + args.dataset.src, + args.dataset.dst, + ) + elif os.path.isdir(args.dataset.src): + dataset = LatentDataset( + args.dataset.src, + args.dataset.dst, + args.dataset.preload + ) + else: + raise OSError(f"Missing dataset source {args.dataset.src}") loader = DataLoader( dataset, batch_size = args.batch, shuffle = True, drop_last = True, pin_memory = False, - # num_workers = 0, - num_workers = 4, - persistent_workers=True, + num_workers = 0, + # num_workers = 6, + # persistent_workers=True, ) - 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)) + + # evals + try: + evals = load_evals(args.dataset.evals) + except: + print(f"No evals, fallback to dataset.") + evals = dataset[0] + + # defaults + crit = torch.nn.L1Loss() + optim_args = { + "lr": args.optim["lr"], + "betas": (args.optim["beta1"], args.optim["beta2"]) + } + + # model + model = InterposerModel(**args.model.args) + model.apply(weights_init) + model.to(args.device) + optim = torch.optim.AdamW(model.parameters(), **optim_args) + + # aux model for reverse pass + model_back = InterposerModel( + ch_in = args.model.args["ch_out"], + ch_mid = args.model.args["ch_mid"], + ch_out = args.model.args["ch_in"], + scale = 1.0 / args.model.args["scale"], + blocks = args.model.args["blocks"], + ) + model_back.apply(weights_init) + model_back.to(args.device) + optim_back = torch.optim.AdamW(model_back.parameters(), **optim_args) + + # scheduler scheduler = None if args.cosine: - print("Using CosineAnnealingLR") scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( - optimizer, T_max = int(args.steps/args.batch), - ) - else: - print("Using LinearLR") - scheduler = torch.optim.lr_scheduler.LinearLR( - optimizer, - start_factor = 0.1, - end_factor = 1.0, - total_iters = int(5000/args.batch), + optim, + T_max = (args.steps - args.fconst), + eta_min = 1e-8, ) - if args.resume: - model.load_state_dict(load_file(args.resume)) - 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) + # vae + vae = None + if args.save_image: + vae = load_vae(args.model.dst, device=args.device, dtype=torch.float16, dec_only=True) - 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, - ) + # main loop + import time + from torch.utils.tensorboard import SummaryWriter + writer = SummaryWriter(log_dir=f"{base_name}_{int(time.time())}") - 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) + pbar = tqdm(total=args.steps) + while pbar.n < args.steps: + for batch in loader: + # get training data + src = batch.get("src").to(args.device) + dst = batch.get("dst").to(args.device) + + ### Train main model ### + optim.zero_grad() + logs = {} + loss = [] with torch.cuda.amp.autocast(): - y_pred = model(src) # forward - loss = criterion(y_pred, dst) # loss + # pass first model + pred = model(src) - # backward - optimizer.zero_grad() + p_loss = crit(pred, dst) * args.p_loss_weight + loss.append(p_loss) + logs["p_loss"] = p_loss.data.item() + + # pass second model + if args.r_loss_weight: + pred_back = model_back(pred) + + r_loss = crit(pred_back, src) * args.r_loss_weight + loss.append(r_loss) + logs["r_loss"] = r_loss.data.item() + + # loss logic + loss = sum(loss) + logs["main"] = loss.data.item() loss.backward() - optimizer.step() - if progress.n >= args.lrskip: scheduler.step() + optim.step() - # eval/save - 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: + # logging + for name, value in logs.items(): + writer.add_scalar(f"loss/{name}", value, pbar.n) + + ### Train backwards model ### + if args.r_loss_weight: + optim_back.zero_grad() + logs = {} + loss = [] + with torch.cuda.amp.autocast(): + # pass second model + pred = model_back(dst) + + p_loss = crit(pred, src) * args.b_loss_weight + loss.append(p_loss) + logs["p_loss"] = p_loss.data.item() + + # pass first model + if args.h_loss_weight: # better w/o this? + pred_back = model(pred) + + r_loss = crit(pred_back, dst) * args.h_loss_weight + loss.append(r_loss) + logs["r_loss"] = r_loss.data.item() + + # loss logic + loss = sum(loss) + logs["main"] = loss.data.item() + loss.backward() + optim_back.step() + + # logging + for name, value in logs.items(): + writer.add_scalar(f"loss_aux/{name}", value, pbar.n) + + # run eval/save eval image + if args.eval_model and pbar.n % args.eval_model == 0: + writer.add_scalar("loss/eval_loss", eval_model(model, evals), pbar.n) + if args.save_image and pbar.n % args.save_image == 0: + for name, image in eval_images(model, vae, evals).items(): + writer.add_image(name, image, pbar.n) + + # scheduler logic main + if scheduler is not None and pbar.n >= args.fconst: + lr = scheduler.get_last_lr()[0] + scheduler.step() + else: + lr = args.optim["lr"] + writer.add_scalar("lr/model", lr, pbar.n) + + # aux model doesn't have a scheduler + writer.add_scalar("lr/model_aux", args.optim["lr"], pbar.n) + + # step + pbar.update() + if pbar.n > args.steps: break - progress.close() - wrapper.save_model(epoch="") # final save - wrapper.close() + + # hacky workaround when the colors are off. + # Save the last n versions and just pick the best one later. + # if pbar.n > (args.steps-2500) and pbar.n%500==0: + # from torchvision.utils import save_image + # save_file(model.state_dict(), f"{base_name}_{pbar.n:07}.safetensors") + # for name, image in eval_images(model, vae, evals).items(): + # name = f"models/{name.replace('/', '_')}_{pbar.n:07}.png" + # save_image(image, name) + + # final save/cleanup + pbar.close() + writer.close() + + save_file(model.state_dict(), f"{base_name}.safetensors") + torch.save(optim.state_dict(), f"{base_name}.optim.pth") diff --git a/vae.py b/vae.py new file mode 100644 index 0000000..e4fa097 --- /dev/null +++ b/vae.py @@ -0,0 +1,114 @@ +import torch +from diffusers import AutoencoderKL + +DTYPE = torch.float16 +DEVICE = "cuda:0" + +class SDv1_VAE: + scale = 1/8 + channels = 4 + def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False): + self.device = device + self.dtype = dtype + self.model = AutoencoderKL.from_pretrained( + "stabilityai/sd-vae-ft-mse" + ) + self.model.eval().to(self.dtype).to(self.device) + if dec_only: + del self.model.encoder + + def encode(self, image): + image = image.to(self.dtype).to(self.device) + image = (image * 2.0) - 1.0 # assuming input is [0;1] + with torch.no_grad(): + latent = self.model.encode(image).latent_dist.sample() + return latent.to(image.dtype).to(image.device) + + def decode(self, latent, grad=False): + latent = latent.to(self.dtype).to(self.device) + if grad: + out = self.model.decode(latent)[0] + else: + with torch.no_grad(): + out = self.model.decode(latent).sample + out = torch.clamp(out, min=-1.0, max=1.0) + out = (out + 1.0) / 2.0 + return out.to(latent.dtype).to(latent.device) + +class SDXL_VAE(SDv1_VAE): + scale = 1/8 + channels = 4 + def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False): + self.device = device + self.dtype = dtype + self.model = AutoencoderKL.from_pretrained( + "madebyollin/sdxl-vae-fp16-fix" + ) + self.model.eval().to(self.dtype).to(self.device) + if dec_only: + del self.model.encoder + +class CascadeC_VAE(SDv1_VAE): + scale = 1/32 + channels = 16 + def __init__(self, device=DEVICE, dtype=DTYPE, **kwargs): + self.device = device + self.dtype = dtype + + #For now this is just piggybacking off of koyha-ss/sd-scripts + from library import stable_cascade as sc + from safetensors.torch import load_file + from huggingface_hub import hf_hub_download + + self.model = sc.EfficientNetEncoder() + self.model.load_state_dict(load_file( + str(hf_hub_download( + repo_id = "stabilityai/stable-cascade", + filename = "effnet_encoder.safetensors", + )) + )) + self.model.eval().to(self.dtype).to(self.device) + +class CascadeA_VAE(): + scale = 1/4 + channels = 4 + def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False): + self.device = device + self.dtype = dtype + + # not sure if this will change in the future? + from diffusers.pipelines.wuerstchen.modeling_paella_vq_model import PaellaVQModel + self.model = PaellaVQModel.from_pretrained( + "stabilityai/stable-cascade", + subfolder="vqgan" + ) + self.model.eval().to(self.dtype).to(self.device) + if dec_only: + del self.model.encoder + + def encode(self, image): + image = image.to(self.dtype).to(self.device) + with torch.no_grad(): + latent = self.model.encode(image).latents + return latent.to(image.dtype).to(image.device) + + def decode(self, latent, grad=False): + latent = latent.to(self.dtype).to(self.device) + if grad: + out = self.model.decode(latent)[0] + else: + with torch.no_grad(): + out = self.model.decode(latent).sample + out = torch.clamp(out, min=0.0, max=1.0) + return out.to(latent.dtype).to(latent.device) + +def load_vae(ver, *args, **kwargs): + if ver == "v1": + VAE = SDv1_VAE + elif ver == "xl": + VAE = SDXL_VAE + elif ver == "cc": + VAE = CascadeC_VAE + elif ver == "ca": + VAE = CascadeA_VAE + return VAE(*args, **kwargs) From 56bfde4ac8624d3903ff11c68e0d95ee5e0c98e2 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 18 Mar 2024 23:51:50 +0100 Subject: [PATCH 3/4] Update README.md --- README.md | 57 ++++++++++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 50 insertions(+), 7 deletions(-) diff --git a/README.md b/README.md index 18a3d8a..86a0dbd 100644 --- a/README.md +++ b/README.md @@ -4,14 +4,17 @@ A small neural network to provide interoperability between the latents generated I wanted to see if it was possible to pass latents generated by the new SDXL model directly into SDv1.5 models without decoding and re-encoding them using a VAE first. ## Installation -To install it, simply clone this repo to your custom_nodes folder using the following command: `git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer`. +To install it, simply clone this repo to your custom_nodes folder using the following command: +``` +git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer +``` Alternatively, you can download the [comfy_latent_interposer.py](https://github.com/city96/SD-Latent-Interposer/raw/main/comfy_latent_interposer.py) file to your `ComfyUI/custom_nodes` folder as well. You may need to install hfhub using the command `pip install huggingface-hub` inside your venv. -If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo. +If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo. The current files are in the **"v4.0"** subfolder. ## Usage -See the image below for an example on how to use it. xl=>v1 conversion is almost flawless, **v1=>xl seems to produce artifacts.** +Simply place it where you would normally place a VAE decode followed by a VAE encode. Set the denoise as appropirate to hide any artifacts while keeping the composition. See image below. ![LATENT_INTERPOSER_V3 1_TEST](https://github.com/city96/SD-Latent-Interposer/assets/125218114/849574b4-2565-4090-85d3-ae63ab425ee2) @@ -20,14 +23,53 @@ Without the interposer, the two latent spaces are incompatible: ![LATENT_INTERPOSER_V3 1](https://github.com/city96/SD-Latent-Interposer/assets/125218114/13e2c01f-580e-4ecb-af1f-b6b21699127b) ### Local models -The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the modules there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models` +The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the models there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models` -Alternatively, just clone the entire HF repo to it: `git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models` +Alternatively, just clone the entire HF repo to it: +``` +git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models +``` + +### Supported Models + +Model names: + +| code | name | +| ---- | -------------------------- | +| `v1` | SDXL | +| `xl` | Stable Diffusion v1.x | +| `ca` | Stable Cascade (Stage A/B) | + +Available models: + +| From | to `v1` | to `xl` | to `ca` | +|:----:|:-------:|:-------:|:-------:| +| `v1` | - | v4.0 | No | +| `xl` | v4.0 | - | No | +| `ca` | v4.0 | v4.0 | - | ## Training -Most of the training/preprocessing code is a 1:1 mirror from my latent upscaler. The folder layout it expects is also the same. + +The training code initializes most training parameters from the provided config file. The dataset should be a single .bin file saved with `torch.save` for each latent version. The format should be [batch, channels, height, width] with the "batch" being as large as the dataset, ie 88000. + +### Interposer v4.0 + +The training code currently initializes two copies of the model, one in the target direction and one in the opposite. The losses are defined based on this. + +- `p_loss` is the main criterion for the primary model. +- `b_loss` is the main criterion for the secondary one. +- `r_loss` is the output of the primary model back through the secondary model and checked against the source latent (basically a round trip through the two models). +- `h_loss` is the same as `r_loss` but for the secondary model. + +All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade). +The training was done locally on an RTX 3080 and a Tesla V100S. + +### Older versions + +
Interposer v3.1 ### Interposer v3.1 + This is basically a complete rewrite. Replaced the mediocre bunch of conv2d layers with something that looks more like a proper neural network. No VGG loss because I still don't have a better GPU. Training was done on combined Flickr2K + DIV2K, with each image being processed into 6 1024x1024 segments. Padded with some of my random images for a total of 22,000 source images in the dataset. @@ -38,7 +80,7 @@ v3.0 was 500k steps at a constant LR of 1e-4, v3.1 was 1M steps using a CosineAn ![INTERPOSER_V3 1](https://github.com/city96/SD-Latent-Interposer/assets/125218114/daff0ae2-4739-4cef-ba54-ac1d156d3388) -### Older versions +
Interposer v1.1 @@ -50,6 +92,7 @@ Overall, it seems to perform a lot better, especially for real life photos. I al
+
Interposer v1.0 ### Interposer v1.0 From a91698a3b84df69e4a12c3054e86c01e7b08d0b8 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Wed, 20 Mar 2024 22:48:47 +0100 Subject: [PATCH 4/4] Update README.md --- README.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/README.md b/README.md index 86a0dbd..8ba2454 100644 --- a/README.md +++ b/README.md @@ -64,6 +64,8 @@ The training code currently initializes two copies of the model, one in the targ All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade). The training was done locally on an RTX 3080 and a Tesla V100S. +![LATENT_INTERPOSER_V4_LOSS](https://github.com/city96/SD-Latent-Interposer/assets/125218114/3a0d8920-ed48-42f0-96c9-897263525efb) + ### Older versions
Interposer v3.1