Add V2 model and improved training code
This commit is contained in:
+32
-15
@@ -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)
|
||||
|
||||
+1
-1
@@ -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],
|
||||
|
||||
@@ -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()
|
||||
|
||||
+31
-14
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user