From 7cf1669be28195b4f491aac3c8736478dbf0761d Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 17 Aug 2023 20:31:30 +0200 Subject: [PATCH] Add V2 model and improved training code --- comfy_latent_upscaler.py | 47 +++++++++---- log_loss.py | 2 +- train.py | 147 ++++++++++++++++++++++++--------------- upscaler.py | 45 ++++++++---- 4 files changed, 156 insertions(+), 85 deletions(-) diff --git a/comfy_latent_upscaler.py b/comfy_latent_upscaler.py index d1f1aa8..605d812 100644 --- a/comfy_latent_upscaler.py +++ b/comfy_latent_upscaler.py @@ -9,24 +9,41 @@ class Upscaler(nn.Module): Basic NN layout, ported from: https://github.com/city96/SD-Latent-Upscaler/blob/main/upscaler.py """ - version = 1.0 # network revision - def __init__(self, fac): - super().__init__() - - module_list = [ - nn.Conv2d(4, 64, kernel_size=5, padding=2), + version = 2.0 # network revision + def head(self): + return [ + nn.Conv2d(self.chan, self.size, kernel_size=self.krn, padding=self.pad), nn.ReLU(), - nn.Upsample(scale_factor=fac, mode="nearest"), + nn.Upsample(scale_factor=self.fac, mode="nearest"), nn.ReLU(), - nn.Conv2d(64, 64, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(64, 64, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(64, 32, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(32, 4, kernel_size=5, padding=2), ] - self.sequential = nn.Sequential(*module_list) + def core(self): + layers = [] + for _ in range(self.depth): + layers += [ + nn.Conv2d(self.size, self.size, kernel_size=self.krn, padding=self.pad), + nn.ReLU(), + ] + return layers + def tail(self): + return [ + nn.Conv2d(self.size, self.chan, kernel_size=self.krn, padding=self.pad), + ] + + def __init__(self, fac, depth=16): + super().__init__() + self.size = 64 # Conv2d size + self.chan = 4 # in/out channels + self.depth = depth # no. of layers + self.fac = fac # scale factor + self.krn = 3 # kernel size + self.pad = 1 # padding + + self.sequential = nn.Sequential( + *self.head(), + *self.core(), + *self.tail(), + ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.sequential(x) diff --git a/log_loss.py b/log_loss.py index ac769eb..0a35ba7 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/train.py b/train.py index 19ead4d..8381f36 100644 --- a/train.py +++ b/train.py @@ -7,15 +7,18 @@ 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 upscaler import LatentUpscaler as Upscaler from vae import get_vae +torch.backends.cudnn.benchmark = True + 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="5e-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("-r", "--res", type=int, default=512, help="Source resolution") parser.add_argument("-f", "--fac", type=float, default=1.5, help="Upscale factor") @@ -29,21 +32,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, ver, src_res, dst_res, dev): - src = os.path.join(f"latents/{ver}_{src_res}px", f"{md5}.npy") - dst = os.path.join(f"latents/{ver}_{dst_res}px", 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(ver, src_res, dst_res, dev): - print("Loading latents from disk") - latents = [] - for i in tqdm(os.listdir(f"latents/{ver}_{src_res}px")): - md5 = os.path.splitext(i)[0] - latents.append(Latent(md5, ver, src_res, dst_res, dev)) - return latents - vae = None def sample_decode(latent, filename, version): global vae @@ -62,65 +50,114 @@ def sample_decode(latent, filename, version): out = Image.fromarray(out) out.save(filename) +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, ver, fac, src): + out = model(src) + output_name = f"./models/latent-upscaler_SD{ver}-x{fac}_e{round(step/1000)}k" + sample_decode(out, f"{output_name}.png", ver) + save_file(model.state_dict(), f"{output_name}.safetensors") + +class Latent: + def __init__(self, md5, ver, src_res, dst_res): + src = os.path.join(f"latents/{ver}_{src_res}px", f"{md5}.npy") + dst = os.path.join(f"latents/{ver}_{dst_res}px", f"{md5}.npy") + self.src = torch.from_numpy(np.load(src)).to("cuda") + self.dst = torch.from_numpy(np.load(dst)).to("cuda") + self.src = torch.squeeze(self.src, 0) + self.dst = torch.squeeze(self.dst, 0) + +class LatentDataset(Dataset): + def __init__(self, ver, src_res, dst_res): + print("Loading latents from disk") + self.latents = [] + for i in tqdm(os.listdir(f"latents/{ver}_{src_res}px")): + md5 = os.path.splitext(i)[0] + self.latents.append( + Latent(md5, ver, src_res, dst_res) + ) + + def __len__(self): + return len(self.latents) + + def __getitem__(self, index): + return ( + self.latents[index].src, + self.latents[index].dst, + ) + if __name__ == "__main__": args = parse_args() target_dev = "cuda" dst_res = int(args.res*args.fac) - latents = load_latents(args.ver, args.res, dst_res, target_dev) + dataset = LatentDataset(args.ver, args.res, dst_res) + loader = DataLoader( + dataset, + batch_size=args.bs, + shuffle=True, + num_workers=0, + ) if not os.path.isdir("models"): os.mkdir("models") log = open(f"models/latent-upscaler_SD{args.ver}-x{args.fac}.csv", "w") if os.path.isfile(f"test_{args.ver}_{args.res}px.npy") and os.path.isfile(f"test_{args.ver}_{dst_res}px.npy"): - ss_latent = torch.from_numpy(np.load(f"test_{args.ver}_{args.res}px.npy")).to(target_dev) - st_latent = torch.from_numpy(np.load(f"test_{args.ver}_{dst_res}px.npy")).to(target_dev) + eval_src = torch.from_numpy(np.load(f"test_{args.ver}_{args.res}px.npy")).to(target_dev) + eval_dst = torch.from_numpy(np.load(f"test_{args.ver}_{dst_res}px.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) + eval_src = dataset[0][0] + eval_dst = dataset[0][1] model = Upscaler(args.fac) if args.resume: model.load_state_dict(load_file(args.resume)) model.to(target_dev) - criterion = torch.nn.MSELoss(size_average=False) - optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) + # criterion = torch.nn.MSELoss() + criterion = torch.nn.L1Loss() - 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) + # optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) + optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr)/args.bs) - y_pred = model(src) # forward - loss = criterion(y_pred, dst) # loss + scheduler = torch.optim.lr_scheduler.OneCycleLR( + optimizer, + total_steps=int(args.steps/args.bs), + max_lr=float(args.lr)/args.bs, + pct_start=0.015, + ) + # scaler = torch.cuda.amp.GradScaler() + progress = tqdm(total=args.steps) - # backward - optimizer.zero_grad() - loss.backward() - optimizer.step() + while progress.n < args.steps: + for src, dst in loader: + with torch.cuda.amp.autocast(): + y_pred = model(src) # forward + loss = criterion(y_pred, dst) # loss - # 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() + # backward + optimizer.zero_grad() + loss.backward() + optimizer.step() + 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, args.ver, args.fac, eval_src) + if progress.n >= args.steps: + break + progress.close() - # sample/save - if step%args.save == 0: - out = model(ss_latent) - output_name = f"./models/latent-upscaler_SD{args.ver}-x{args.fac}_e{step/1000}k" - sample_decode(out, f"{output_name}.png", args.ver) - save_file(model.state_dict(), f"{output_name}.safetensors") # save final output - output_name = f"./models/latent-upscaler_SD{args.ver}-x{args.fac}_e{step/1000}k" - sample_decode(out, f"{output_name}.png", args.ver) - save_file(model.state_dict(), f"{output_name}.safetensors") + eval_model(args.steps, model, criterion, scheduler, eval_src, eval_dst) + save_model(args.steps, model, args.ver, args.fac, eval_src) log.close() diff --git a/upscaler.py b/upscaler.py index 094e50c..a40d477 100644 --- a/upscaler.py +++ b/upscaler.py @@ -3,23 +3,40 @@ import torch.nn as nn import numpy as np class LatentUpscaler(nn.Module): - def __init__(self, fac): - super().__init__() - - module_list = [ - nn.Conv2d(4, 64, kernel_size=5, padding=2), + def head(self): + return [ + nn.Conv2d(self.chan, self.size, kernel_size=self.krn, padding=self.pad), nn.ReLU(), - nn.Upsample(scale_factor=fac, mode="nearest"), # bicubic was blurry + nn.Upsample(scale_factor=self.fac, mode="nearest"), nn.ReLU(), - nn.Conv2d(64, 64, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(64, 64, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(64, 32, kernel_size=7, padding=3), - nn.ReLU(), - nn.Conv2d(32, 4, kernel_size=5, padding=2), ] - self.sequential = nn.Sequential(*module_list) + def core(self): + layers = [] + for _ in range(self.depth): + layers += [ + nn.Conv2d(self.size, self.size, kernel_size=self.krn, padding=self.pad), + nn.ReLU(), + ] + return layers + def tail(self): + return [ + nn.Conv2d(self.size, self.chan, kernel_size=self.krn, padding=self.pad), + ] + + def __init__(self, fac, depth=16): + super().__init__() + self.size = 64 # Conv2d size + self.chan = 4 # in/out channels + self.depth = depth # no. of layers + self.fac = fac # scale factor + self.krn = 3 # kernel size + self.pad = 1 # padding + + self.sequential = nn.Sequential( + *self.head(), + *self.core(), + *self.tail(), + ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.sequential(x)