Add training code

This commit is contained in:
City
2023-08-15 19:50:45 +02:00
parent ca060b9eca
commit 14c6d8e75e
6 changed files with 342 additions and 0 deletions
+15
View File
@@ -1,3 +1,18 @@
raw/
images/
latent_*/
vae/
models/
other/
test.py
*.png
*.zip
*.npy
*.ckpt
*.safetensors
# default github .gitignore follows
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
+48
View File
@@ -0,0 +1,48 @@
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 = {}
def process_lines(lines):
global train_loss
global eval_loss
name = fp.split("/")[1]
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],
)
if len(vals[0]) == 3:
eval_loss[name] = (
[int(x[0]) for x in vals],
[math.log(float(x[2])) 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):
fig, ax = plt.subplots()
ax.grid()
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')
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")
+80
View File
@@ -0,0 +1,80 @@
import os
import torch
import hashlib
import argparse
import numpy as np
from torchvision import transforms
from diffusers import AutoencoderKL
from tqdm import tqdm
from PIL import Image
from vae import get_vae
def parse_args():
parser = argparse.ArgumentParser(description="Preprocess images into latents")
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("-v", "--ver", choices=["v1","xl"], default="v1", help="SD version")
parser.add_argument('--vae', help="Path to VAE (Optional)")
parser.add_argument('--src', default="raw", help="Source folder with images")
return parser.parse_args()
def encode(vae, img):
"""image [PIL Image] -> latent [np array]"""
inp = transforms.ToTensor()(img).unsqueeze(0)
inp = inp.to("cuda") # move to GPU
latent = vae.encode(inp*2.0-1.0)
latent = latent.latent_dist.sample()
return latent.cpu().detach()
def scale(path, res):
"""Crop image to the top-left corner"""
img = Image.open(path)
img = img.convert('RGB')
target = (res, res)
if min(img.height, img.width) < 256:
return
if img.width > img.height:
target = (int(img.width/img.height*res), res)
elif img.height > img.width:
target = (res, int(img.height/img.width*res))
img = img.resize(target, Image.LANCZOS)
img = img.crop([0,0,res,res])
return img
def process_folder(vae, src_dir, ver, res):
dst_dir = f"latents/{ver}_{res}px"
if not os.path.isdir(dst_dir):
os.mkdir(dst_dir)
for file in tqdm(os.listdir(src_dir)):
src = os.path.join(src_dir, file)
md5 = hashlib.md5(open(src,'rb').read()).hexdigest()
dst = os.path.join(dst_dir, f"{md5}.npy")
if os.path.isfile(dst):
continue
img = scale(src, res)
latent = encode(vae, img)
np.save(dst, latent)
def process_res(vae, src_dir, ver, res):
process_folder(vae, src_dir, ver, res)
# test image, optional
if os.path.isfile("test.png"):
if os.path.isfile(f"test_{ver}_{res}px.npy"):
return
img = scale("test.png", res)
latent = encode(vae, img)
np.save(f"test_{ver}_{res}px.npy", latent)
torch.cuda.empty_cache()
if __name__ == "__main__":
if not os.path.isdir("latents"):
os.mkdir("latents")
args = parse_args()
vae = get_vae(args.ver, args.vae)
vae.to("cuda")
## args
dst_res = int(args.res*args.fac)
process_res(vae, args.src, args.ver, args.res)
process_res(vae, args.src, args.ver, dst_res)
+126
View File
@@ -0,0 +1,126 @@
import os
import torch
import torch.nn as nn
import numpy as np
import argparse
import random
from PIL import Image
from tqdm import tqdm
from safetensors.torch import save_file, load_file
from upscaler import LatentUpscaler as Upscaler
from vae import get_vae
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("-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")
parser.add_argument("-v", "--ver", choices=["v1","xl"], default="v1", help="SD version")
parser.add_argument('--vae', help="Path to VAE (Optional)")
parser.add_argument('--resume', help="Checkpoint to resume from")
args = parser.parse_args()
try:
float(args.lr)
except:
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
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)
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)
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)
else:
sample_latent = random.choice(latents)
ss_latent = sample_latent.src.to(target_dev)
st_latent = sample_latent.dst.to(target_dev)
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)
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)
y_pred = model(src) # forward
loss = criterion(y_pred, dst) # loss
# backward
optimizer.zero_grad()
loss.backward()
optimizer.step()
# 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()
# 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")
log.close()
+25
View File
@@ -0,0 +1,25 @@
import torch
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),
nn.ReLU(),
nn.Upsample(scale_factor=fac, mode="nearest"), # bicubic was blurry
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 forward(self, x: torch.Tensor) -> torch.Tensor:
return self.sequential(x)
+48
View File
@@ -0,0 +1,48 @@
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