Add training code
This commit is contained in:
+15
@@ -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
|
# Byte-compiled / optimized / DLL files
|
||||||
__pycache__/
|
__pycache__/
|
||||||
*.py[cod]
|
*.py[cod]
|
||||||
|
|||||||
+48
@@ -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")
|
||||||
@@ -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)
|
||||||
@@ -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
@@ -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)
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user