Version 3 / rewrite

This commit is contained in:
City
2023-10-11 06:11:27 +02:00
parent 451cb196b7
commit 31c3eb5a82
6 changed files with 188 additions and 210 deletions
+4 -1
View File
@@ -1,12 +1,15 @@
raw/ raw/
images/ images/
latent_*/ latents/
latents
vae/ vae/
models/ models/
other/
test.py test.py
*.png *.png
*.zip *.zip
*.npy *.npy
*.pth
*.ckpt *.ckpt
*.safetensors *.safetensors
+46 -21
View File
@@ -2,30 +2,55 @@ import torch
import torch.nn as nn import torch.nn as nn
import numpy as np import numpy as np
class Block(nn.Module):
def __init__(self, size):
super().__init__()
self.join = nn.ReLU()
self.long = nn.Sequential(
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.2)
)
def forward(self, x):
y = self.long(x)
z = self.join(y + x)
return z
class Interposer(nn.Module): class Interposer(nn.Module):
def __init__(self): def __init__(self):
super().__init__() super().__init__()
self.chan = 4 # in/out channels
self.hid = 128
# it looks like a spaceship if you squint :D # expand channels
module_list = [ self.head_join = nn.ReLU()
#############) self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1)
#############) self.head_long = nn.Sequential(
#||# nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1),
#||# nn.LeakyReLU(0.1),
nn.Conv2d(4, 32, kernel_size=5, padding=2), nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
)
# not sure if this is how residuals work
self.core = nn.Sequential(
Block(self.hid),
Block(self.hid),
Block(self.hid),
)
# reduce channels
self.tail = nn.Sequential(
nn.ReLU(), nn.ReLU(),
nn.Conv2d(32, 128, kernel_size=7, padding=3), nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
nn.ReLU(), )
nn.Conv2d(128, 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):
y = self.head_join(
def forward(self, x: torch.Tensor) -> torch.Tensor: self.head_long(x)+
return self.sequential(x) self.head_short(x)
)
z = self.core(y)
return self.tail(z)
+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],
-65
View File
@@ -1,65 +0,0 @@
import os
import hashlib
import argparse
from tqdm import tqdm
from PIL import Image
from queue import Queue
from threading import Thread
if not os.path.isdir("images"):
os.mkdir("images")
def parse_args():
parser = argparse.ArgumentParser(description="Preprocess images")
parser.add_argument("-r", "--res", type=int, default=768, help="Target resolution")
parser.add_argument("-t", "--threads", type=int, default=4, help="No. of CPU threads to use")
parser.add_argument('--src', default="raw", help="Source folder with images")
return parser.parse_args()
def process(fname, folder, resolution):
src = os.path.join(folder, fname)
md5 = hashlib.md5(open(src,'rb').read()).hexdigest()
out = os.path.join("images", f"{md5}.png")
if os.path.isfile(out):
return
img = Image.open(src)
img = img.convert('RGB')
target = (resolution, resolution)
if min(img.height, img.width) < 256:
return
if img.width > img.height:
target = (int(img.width/img.height*resolution), resolution)
elif img.height > img.width:
target = (resolution, int(img.height/img.width*resolution))
img = img.resize(target, Image.LANCZOS)
img = img.crop([0,0,resolution,resolution])
img.save(out)
def thread(queue, pbar, folder, resolution):
while not queue.empty():
fname = queue.get()
try: process(fname, folder, resolution)
except: pass
queue.task_done()
pbar.update()
args = parse_args()
files = os.listdir(args.src)
pbar = tqdm(total=len(files),unit="img")
queue = Queue()
[queue.put(x) for x in files]
for _ in range(args.threads):
Thread(
target=thread,
args=(
queue,
pbar,
args.src,
args.res,
),
daemon=True,
).start()
queue.join()
-51
View File
@@ -1,51 +0,0 @@
import os
import torch
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 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 process_folder(vae, v):
if not os.path.isdir(f"latent_{v}"):
os.mkdir(f"latent_{v}")
vae.to("cuda")
for i in tqdm(os.listdir("images")):
src = os.path.join("images", i)
img = Image.open(src)
dst = os.path.join(f"latent_{v}", f"{os.path.splitext(i)[0]}.npy")
latent = encode(vae, img)
np.save(dst, latent)
vae.to("cpu")
def run_v1(file_path=None):
vae = get_vae("v1", file_path)
process_folder(vae, "v1")
del vae
def run_v2(file_path=None):
vae = get_vae("v2", file_path)
process_folder(vae, "v2")
del vae
def run_xl(file_path=None):
vae = get_vae("xl", file_path)
process_folder(vae, "xl")
del vae
if __name__ == "__main__":
# run_v1("./vae/ft-mse-840000.ckpt") # probably doesn't reflect internal SD latent
run_v1()
# run_v2() # v2 and v1 share a latent space
run_xl("./vae/sdxl_v0.9.safetensors") # 1.0 has artifacts
+129 -63
View File
@@ -3,23 +3,27 @@ import torch
import torch.nn as nn import torch.nn as nn
import numpy as np import numpy as np
import argparse import argparse
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 interposer import Interposer from interposer import Interposer
from vae import get_vae from vae import get_vae
torch.backends.cudnn.benchmark = True
torch.manual_seed(0)
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="1e-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('--src', choices=["v1","xl"], required=True, help="Source latent format") parser.add_argument('--src', choices=["v1","xl"], required=True, help="Source latent format")
parser.add_argument('--dst', choices=["v1","xl"], required=True, help="Destination latent format") parser.add_argument('--dst', choices=["v1","xl"], required=True, help="Destination latent format")
parser.add_argument('--resume', help="Checkpoint to resume from") parser.add_argument('--resume', help="Checkpoint to resume from")
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR")
args = parser.parse_args() args = parser.parse_args()
if args.src == args.dst: if args.src == args.dst:
parser.error("--src and --dst can't be the same") parser.error("--src and --dst can't be the same")
@@ -29,25 +33,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, lat_src, lat_dst, dev):
if lat_src == "v1": src = os.path.join("latent_v1", f"{md5}.npy")
if lat_src == "xl": src = os.path.join("latent_xl", f"{md5}.npy")
if lat_dst == "v1": dst = os.path.join("latent_v1", f"{md5}.npy")
if lat_dst == "xl": dst = os.path.join("latent_xl", 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(src, dst, dev):
print("Loading latents from disk")
latents = []
for i in tqdm(os.listdir("images")):
md5 = os.path.splitext(i)[0]
latents.append(Latent(md5, src, dst, dev))
return latents
vae = None vae = None
def sample_decode(latent, filename, version): def sample_decode(latent, filename, version):
global vae global vae
@@ -66,40 +51,127 @@ def sample_decode(latent, filename, version):
out = Image.fromarray(out) out = Image.fromarray(out)
out.save(filename) out.save(filename)
def get_eval_data(dataset, src_path, dst_path, target_dev):
if os.path.isfile(src_path) and os.path.isfile(dst_path):
src = LatentDataset.load_latent(None, src_path)
dst = LatentDataset.load_latent(None, dst_path)
else:
src = dataset[0][0]
dst = dataset[0][1]
src = src.float().to(target_dev).unsqueeze(0)
dst = dst.float().to(target_dev).unsqueeze(0)
return(src, dst)
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, optim, lat, src, dst):
with torch.no_grad():
out = model(lat)
output_name = f"./models/{src}-to-{dst}_interposer_e{round(step/1000)}k"
sample_decode(out, f"{output_name}.png", dst)
save_file(model.state_dict(), f"{output_name}.safetensors")
torch.save(optim.state_dict(), f"{output_name}.optim.pth")
class LatentDataset(Dataset):
class Shard:
def __init__(self, root, fname, res, src, dst):
self.fname = fname
self.src_path = f"{root}/{src}_{res}px/{fname}.npy"
self.dst_path = f"{root}/{dst}_{res}px/{fname}.npy"
def __init__(self, res, src, dst, root="latents"):
print("Loading latents from disk")
self.latents = []
for i in tqdm(os.listdir(f"{root}/{src}_{res}px")):
fname, ext = os.path.splitext(i)
assert ext == ".npy"
s = self.Shard(root, fname, res, src, dst)
if os.path.isfile(s.src_path) and os.path.isfile(s.dst_path):
self.latents.append(s)
def __len__(self):
return len(self.latents)
def __getitem__(self, index):
s = self.latents[index]
src = self.load_latent(s.src_path)
dst = self.load_latent(s.dst_path)
return (src, dst)
def load_latent(self, path):
lat = torch.from_numpy(np.load(path))
if lat.shape[0] == 1:
lat = torch.squeeze(lat, 0)
assert not torch.isnan(torch.sum(lat.float()))
return lat
if __name__ == "__main__": if __name__ == "__main__":
args = parse_args() args = parse_args()
target_dev = "cuda" target_dev = "cuda"
latent_src = args.src resolution = 768
latent_dst = args.dst
latents = load_latents(latent_src, latent_dst, target_dev) dataset = LatentDataset(resolution, args.src, args.dst)
loader = DataLoader(
dataset,
batch_size=args.bs,
shuffle=True,
num_workers=0,
# num_workers=4,
# persistent_workers=True,
)
eval_src, eval_dst = get_eval_data(
dataset,
f"latents/test_{args.src}_{resolution}px.npy",
f"latents/test_{args.dst}_{resolution}px.npy",
target_dev,
)
if not os.path.isdir("models"): os.mkdir("models") os.makedirs("models", exist_ok=True)
log = open(f"models/{latent_src}-to-{latent_dst}_interposer.csv", "w") log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
if os.path.isfile(f"test_{latent_src}.npy") and os.path.isfile(f"test_{latent_dst}.npy"):
ss_latent = torch.from_numpy(np.load(f"test_{latent_src}.npy")).to(target_dev)
st_latent = torch.from_numpy(np.load(f"test_{latent_dst}.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 = Interposer() model = Interposer()
criterion = torch.nn.L1Loss()
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr))
# import bitsandbytes as bnb
# optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=float(args.lr))
scheduler = None
if args.cosine:
print("Using CosineAnnealingLR")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max = int(args.steps/args.bs),
)
else:
print("Using LinearLR")
scheduler = torch.optim.lr_scheduler.LinearLR(
optimizer,
start_factor = 0.1,
end_factor = 1.0,
total_iters = int(5000/args.bs),
)
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)
optimizer.load_state_dict(torch.load(
f"{os.path.splitext(args.resume)[0]}.optim.pth"
))
else:
model.to(target_dev)
criterion = torch.nn.MSELoss(size_average=False) progress = tqdm(total=args.steps)
optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs) while progress.n < args.steps:
for src, dst in loader:
for t in tqdm(range(int(args.steps/args.bs)), unit_scale=args.bs): src = src.to(target_dev)
step = t*args.bs dst = dst.to(target_dev)
# input batch with torch.cuda.amp.autocast():
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 y_pred = model(src) # forward
loss = criterion(y_pred, dst) # loss loss = criterion(y_pred, dst) # loss
@@ -107,25 +179,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, optimizer, eval_src, args.src, args.dst)
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_src}-to-{latent_dst}_interposer_e{step/1000}k"
sample_decode(out, f"{output_name}.png", latent_dst)
save_file(model.state_dict(), f"{output_name}.safetensors")
# save final output # save final output
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{args.steps/1000}k" eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
sample_decode(out, f"{output_name}.png", "v1") save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
save_file(model.state_dict(), f"{output_name}.safetensors")
log.close() log.close()