Add V2 model and improved training code

This commit is contained in:
City
2023-08-17 20:31:30 +02:00
parent 002b326f1e
commit 7cf1669be2
4 changed files with 156 additions and 85 deletions
+32 -15
View File
@@ -9,24 +9,41 @@ class Upscaler(nn.Module):
Basic NN layout, ported from: Basic NN layout, ported from:
https://github.com/city96/SD-Latent-Upscaler/blob/main/upscaler.py https://github.com/city96/SD-Latent-Upscaler/blob/main/upscaler.py
""" """
version = 1.0 # network revision version = 2.0 # network revision
def __init__(self, fac): def head(self):
super().__init__() return [
nn.Conv2d(self.chan, self.size, kernel_size=self.krn, padding=self.pad),
module_list = [
nn.Conv2d(4, 64, kernel_size=5, padding=2),
nn.ReLU(), nn.ReLU(),
nn.Upsample(scale_factor=fac, mode="nearest"), nn.Upsample(scale_factor=self.fac, mode="nearest"),
nn.ReLU(), 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: def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.sequential(x) return self.sequential(x)
+1 -1
View File
@@ -15,7 +15,7 @@ def process_lines(lines):
[int(x[0]) for x in vals], [int(x[0]) for x in vals],
[math.log(float(x[1])) 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] = ( eval_loss[name] = (
[int(x[0]) for x in vals], [int(x[0]) for x in vals],
[math.log(float(x[2])) for x in vals], [math.log(float(x[2])) for x in vals],
+86 -49
View File
@@ -7,15 +7,18 @@ import random
from PIL import Image from PIL import Image
from tqdm import tqdm from tqdm import tqdm
from safetensors.torch import save_file, load_file from safetensors.torch import save_file, load_file
from torch.utils.data import DataLoader, Dataset
from upscaler import LatentUpscaler as Upscaler from upscaler import LatentUpscaler as Upscaler
from vae import get_vae from vae import get_vae
torch.backends.cudnn.benchmark = True
def parse_args(): def parse_args():
parser = argparse.ArgumentParser(description="Train latent interposer model") parser = argparse.ArgumentParser(description="Train latent interposer model")
parser.add_argument("--steps", type=int, default=500000, help="No. of training steps") 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('--bs', type=int, default=4, help="Batch size")
parser.add_argument('--lr', default="1e-8", help="Learning rate") 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("-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("-r", "--res", type=int, default=512, help="Source resolution")
parser.add_argument("-f", "--fac", type=float, default=1.5, help="Upscale factor") 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") parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
return args 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 vae = None
def sample_decode(latent, filename, version): def sample_decode(latent, filename, version):
global vae global vae
@@ -62,39 +50,94 @@ def sample_decode(latent, filename, version):
out = Image.fromarray(out) out = Image.fromarray(out)
out.save(filename) 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__": if __name__ == "__main__":
args = parse_args() args = parse_args()
target_dev = "cuda" target_dev = "cuda"
dst_res = int(args.res*args.fac) 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") if not os.path.isdir("models"): os.mkdir("models")
log = open(f"models/latent-upscaler_SD{args.ver}-x{args.fac}.csv", "w") 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"): 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) eval_src = 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_dst = torch.from_numpy(np.load(f"test_{args.ver}_{dst_res}px.npy")).to(target_dev)
else: else:
sample_latent = random.choice(latents) eval_src = dataset[0][0]
ss_latent = sample_latent.src.to(target_dev) eval_dst = dataset[0][1]
st_latent = sample_latent.dst.to(target_dev)
model = Upscaler(args.fac) model = Upscaler(args.fac)
if args.resume: if args.resume:
model.load_state_dict(load_file(args.resume)) model.load_state_dict(load_file(args.resume))
model.to(target_dev) model.to(target_dev)
criterion = torch.nn.MSELoss(size_average=False) # criterion = torch.nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) criterion = torch.nn.L1Loss()
for t in tqdm(range(int(args.steps/args.bs)), unit_scale=args.bs): # optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs)
step = t*args.bs optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr)/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)
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)
while progress.n < args.steps:
for src, dst in loader:
with torch.cuda.amp.autocast():
y_pred = model(src) # forward y_pred = model(src) # forward
loss = criterion(y_pred, dst) # loss loss = criterion(y_pred, dst) # loss
@@ -102,25 +145,19 @@ if __name__ == "__main__":
optimizer.zero_grad() optimizer.zero_grad()
loss.backward() loss.backward()
optimizer.step() optimizer.step()
scheduler.step()
# print loss # eval/save
if step%1000 == 0: progress.update(args.bs)
# test loss if progress.n % (1000 + 1000%args.bs) == 0:
with torch.no_grad(): eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
t_pred = model(ss_latent) if progress.n % (args.save + args.save%args.bs) == 0:
t_loss = criterion(t_pred, st_latent) save_model(progress.n, model, args.ver, args.fac, eval_src)
tqdm.write(f"{step} - {loss.data.item()/args.bs:.2f}|{t_loss.data.item()/args.bs:.2f}") if progress.n >= args.steps:
log.write(f"{step},{loss.data.item()/args.bs:.2f},{t_loss.data.item()/args.bs:.2f}\n") break
log.flush() 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 # save final output
output_name = f"./models/latent-upscaler_SD{args.ver}-x{args.fac}_e{step/1000}k" eval_model(args.steps, model, criterion, scheduler, eval_src, eval_dst)
sample_decode(out, f"{output_name}.png", args.ver) save_model(args.steps, model, args.ver, args.fac, eval_src)
save_file(model.state_dict(), f"{output_name}.safetensors")
log.close() log.close()
+31 -14
View File
@@ -3,23 +3,40 @@ import torch.nn as nn
import numpy as np import numpy as np
class LatentUpscaler(nn.Module): class LatentUpscaler(nn.Module):
def __init__(self, fac): def head(self):
super().__init__() return [
nn.Conv2d(self.chan, self.size, kernel_size=self.krn, padding=self.pad),
module_list = [
nn.Conv2d(4, 64, kernel_size=5, padding=2),
nn.ReLU(), nn.ReLU(),
nn.Upsample(scale_factor=fac, mode="nearest"), # bicubic was blurry nn.Upsample(scale_factor=self.fac, mode="nearest"),
nn.ReLU(), 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: def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.sequential(x) return self.sequential(x)