Snapshot
This commit is contained in:
+89
@@ -0,0 +1,89 @@
|
|||||||
|
# Custom dataset to load encoded latents from disk.
|
||||||
|
# Files should contain latents as (1, C, H, W) or (C, H, W)
|
||||||
|
# Latents should be in their original format without scaling
|
||||||
|
|
||||||
|
######### Folder Layout #########
|
||||||
|
# latents #
|
||||||
|
# |- test_v1_768px.npy <=eval #
|
||||||
|
# |- test_xl_768px.npy <=^ #
|
||||||
|
# |- v1_768px <= ver/res #
|
||||||
|
# | |- 000001.npy #
|
||||||
|
# | |- 000002.npy #
|
||||||
|
# | | ... #
|
||||||
|
# | |- 000999.npy #
|
||||||
|
# | \- 001000.npy #
|
||||||
|
# |- xl_768px #
|
||||||
|
# ... #
|
||||||
|
#################################
|
||||||
|
|
||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
from tqdm import tqdm
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
DEFAULT_ROOT = "latents"
|
||||||
|
ALLOWED_EXTS = [".npy"]
|
||||||
|
|
||||||
|
class Shard:
|
||||||
|
"""
|
||||||
|
Shard to store groups of latents in
|
||||||
|
paths: List containing paths to latent encoded images
|
||||||
|
"""
|
||||||
|
def __init__(self, paths):
|
||||||
|
self.paths = paths
|
||||||
|
self.data = None
|
||||||
|
|
||||||
|
def exists(self):
|
||||||
|
return all([os.path.isfile(x) for x in self.paths])
|
||||||
|
|
||||||
|
def get_data(self):
|
||||||
|
if self.data is not None: return self.data
|
||||||
|
return tuple([self.load_latent(x) for x in self.paths])
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
def preload(self):
|
||||||
|
self.data = self.get_data()
|
||||||
|
|
||||||
|
class LatentDataset(Dataset):
|
||||||
|
def __init__(self, specs, res=768, root=DEFAULT_ROOT, preload=False):
|
||||||
|
"""
|
||||||
|
Main dataset that returns list of requested images as (C, H, W) latents
|
||||||
|
specs: List of latent versions in the other to return them in
|
||||||
|
res: Native resolution of images (before latent encoding)
|
||||||
|
root: Path to folder with sorted files
|
||||||
|
preload: Load all files into memory on initialization
|
||||||
|
"""
|
||||||
|
print("Dataset: Parsing data from disk")
|
||||||
|
self.specs = specs
|
||||||
|
self.res = res
|
||||||
|
self.root = root
|
||||||
|
self.shards = []
|
||||||
|
for fname in tqdm(os.listdir(f"{root}/{specs[0]}_{res}px")):
|
||||||
|
name, ext = os.path.splitext(fname)
|
||||||
|
if ext not in ALLOWED_EXTS: continue
|
||||||
|
shard = Shard([f"{root}/{x}_{res}px/{name}{ext}" for x in specs])
|
||||||
|
if shard.exists():
|
||||||
|
self.shards.append(shard)
|
||||||
|
|
||||||
|
if preload: # cache to RAM
|
||||||
|
print("Dataset: Preloading data to system RAM")
|
||||||
|
[x.preload() for x in tqdm(self.shards)]
|
||||||
|
print(f"Dataset: OK, {len(self)} items")
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.shards)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
return self.shards[index].get_data()
|
||||||
|
|
||||||
|
def get_eval(self):
|
||||||
|
shard = Shard([f"{self.root}/test_{x}_{self.res}px.npy" for x in self.specs])
|
||||||
|
data = shard.get_data() if shard.exists() else self[0]
|
||||||
|
return tuple([x.unsqueeze(0).to(torch.float32) for x in data])
|
||||||
+38
-37
@@ -1,56 +1,57 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
class Block(nn.Module):
|
class ResBlock(nn.Module):
|
||||||
def __init__(self, size):
|
"""Block with residuals"""
|
||||||
|
def __init__(self, ch):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.join = nn.ReLU()
|
self.join = nn.ReLU()
|
||||||
self.long = nn.Sequential(
|
self.long = nn.Sequential(
|
||||||
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
|
||||||
nn.LeakyReLU(0.1),
|
nn.LeakyReLU(0.1),
|
||||||
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
|
||||||
nn.LeakyReLU(0.1),
|
nn.LeakyReLU(0.1),
|
||||||
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
|
||||||
nn.Dropout(0.2)
|
|
||||||
)
|
)
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
y = self.long(x)
|
return self.join(self.long(x) + x)
|
||||||
z = self.join(y + x)
|
|
||||||
return z
|
|
||||||
|
|
||||||
class Interposer(nn.Module):
|
class ExtractBlock(nn.Module):
|
||||||
def __init__(self):
|
"""Increase no. of channels by [out/in]"""
|
||||||
|
def __init__(self, ch_in, ch_out):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.chan = 4 # in/out channels
|
self.join = nn.ReLU()
|
||||||
self.hid = 128
|
self.short = nn.Conv2d(ch_in, ch_out, kernel_size=3, stride=1, padding=1)
|
||||||
|
self.long = nn.Sequential(
|
||||||
|
nn.Conv2d( ch_in, ch_out, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.Dropout(0.1)
|
||||||
|
)
|
||||||
|
def forward(self, x):
|
||||||
|
return self.join(self.long(x) + self.short(x))
|
||||||
|
|
||||||
# expand channels
|
class InterposerModel(nn.Module):
|
||||||
self.head_join = nn.ReLU()
|
"""Main neural network"""
|
||||||
self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1)
|
def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0):
|
||||||
self.head_long = nn.Sequential(
|
super().__init__()
|
||||||
nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1),
|
self.scale = scale
|
||||||
nn.LeakyReLU(0.1),
|
self.ch_in = ch_in
|
||||||
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
|
self.ch_out = ch_out
|
||||||
nn.LeakyReLU(0.1),
|
self.ch_mid = ch_mid
|
||||||
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
|
|
||||||
)
|
self.head = ExtractBlock(self.ch_in, self.ch_mid)
|
||||||
# not sure if this is how residuals work
|
|
||||||
self.core = nn.Sequential(
|
self.core = nn.Sequential(
|
||||||
Block(self.hid),
|
nn.Upsample(scale_factor=self.scale, mode="nearest"),
|
||||||
Block(self.hid),
|
ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
|
||||||
Block(self.hid),
|
ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
|
||||||
)
|
ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
|
||||||
# reduce channels
|
|
||||||
self.tail = nn.Sequential(
|
|
||||||
nn.ReLU(),
|
|
||||||
nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
|
|
||||||
)
|
)
|
||||||
|
self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
y = self.head_join(
|
y = self.head(x)
|
||||||
self.head_long(x)+
|
|
||||||
self.head_short(x)
|
|
||||||
)
|
|
||||||
z = self.core(y)
|
z = self.core(y)
|
||||||
return self.tail(z)
|
return self.tail(z)
|
||||||
|
|||||||
+39
-15
@@ -5,20 +5,36 @@ import matplotlib.pyplot as plt
|
|||||||
files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")]
|
files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")]
|
||||||
train_loss = {}
|
train_loss = {}
|
||||||
eval_loss = {}
|
eval_loss = {}
|
||||||
|
lr_vals = {}
|
||||||
|
fskip = 0
|
||||||
|
|
||||||
|
offsets = { # offset to display resumed training runs
|
||||||
|
}
|
||||||
|
sep = ".csv"
|
||||||
|
rep = "_interposer"
|
||||||
|
model = "Latent Interposer"
|
||||||
|
|
||||||
def process_lines(lines):
|
def process_lines(lines):
|
||||||
global train_loss
|
global train_loss
|
||||||
global eval_loss
|
global eval_loss
|
||||||
name = fp.split("/")[1].split("_")[0]
|
name = fp.split("/")[1]
|
||||||
|
print(name)
|
||||||
|
if sep: name = name.split(sep)[0]
|
||||||
|
if rep: name = name.replace(rep,"")
|
||||||
vals = [x.split(",") for x in lines]
|
vals = [x.split(",") for x in lines]
|
||||||
train_loss[name] = (
|
train_loss[name] = (
|
||||||
[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])+1e-10) 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])+1e-10) for x in vals],
|
||||||
|
)
|
||||||
|
if len(vals[0]) >= 4:
|
||||||
|
lr_vals[name] = (
|
||||||
|
[int(x[0]) for x in vals],
|
||||||
|
[float(x[3]) for x in vals],
|
||||||
)
|
)
|
||||||
|
|
||||||
# https://stackoverflow.com/a/49357445
|
# https://stackoverflow.com/a/49357445
|
||||||
@@ -31,18 +47,26 @@ def smooth(scalars, weight):
|
|||||||
last = smoothed_val
|
last = smoothed_val
|
||||||
return smoothed
|
return smoothed
|
||||||
|
|
||||||
def plot(data, fname):
|
def plot(data, fname, title=None, smw=0.9):
|
||||||
fig, ax = plt.subplots()
|
fig, ax = plt.subplots()
|
||||||
|
plt.tight_layout()
|
||||||
ax.grid()
|
ax.grid()
|
||||||
|
dmax = 0
|
||||||
for name, val in data.items():
|
for name, val in data.items():
|
||||||
ax.plot(val[0], smooth(val[1], 0.9), label=name)
|
data = [x + offsets[name] for x in val[0]] if name in offsets.keys() else val[0]
|
||||||
plt.legend(loc="upper right")
|
dmax = max(dmax, round(data[-1],10000))
|
||||||
plt.savefig(fname, dpi=300, bbox_inches='tight')
|
sval = val[1][:fskip] + smooth(val[1][fskip:], smw) # skip first N
|
||||||
|
ax.plot(data, sval, label=name)
|
||||||
|
ax.set_xticks([dmax//10*x for x in range(10)])
|
||||||
|
plt.legend(loc="lower left", bbox_to_anchor=(0.00, -0.20), ncol=5)
|
||||||
|
if title: plt.title(title)
|
||||||
|
plt.savefig(fname, bbox_inches='tight')
|
||||||
|
|
||||||
for fp in files:
|
if __name__ == "__main__":
|
||||||
with open(fp) as f:
|
for fp in files:
|
||||||
lines = f.readlines()
|
with open(fp) as f:
|
||||||
process_lines(lines)
|
lines = f.readlines()
|
||||||
|
process_lines(lines)
|
||||||
plot(train_loss, "loss.png")
|
plot(train_loss, "loss.png", f"{model} Training loss", 0.2)
|
||||||
plot(eval_loss, "loss-eval.png")
|
plot(eval_loss, "loss-eval.png", f"{model} Eval. loss", 0.7)
|
||||||
|
plot(lr_vals, "loss-lr.png", f"{model} Learning rate", 0.0)
|
||||||
|
|||||||
@@ -1,29 +1,31 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
|
||||||
import numpy as np
|
|
||||||
import argparse
|
import argparse
|
||||||
from PIL import Image
|
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from safetensors.torch import save_file, load_file
|
from torch.utils.data import DataLoader
|
||||||
from torch.utils.data import DataLoader, Dataset
|
from safetensors.torch import load_file
|
||||||
|
|
||||||
from interposer import Interposer
|
from interposer import InterposerModel as Model
|
||||||
from vae import get_vae
|
from dataset import LatentDataset
|
||||||
|
from utils import ModelWrapper
|
||||||
|
|
||||||
torch.backends.cudnn.benchmark = True
|
torch.backends.cudnn.benchmark = True
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
|
|
||||||
|
TARGET_DEV = "cuda"
|
||||||
|
|
||||||
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("-s", "--steps", type=int, default=500000, help="No. of training steps")
|
||||||
parser.add_argument('--bs', type=int, default=4, help="Batch size")
|
parser.add_argument("-b", "--batch", type=int, default= 1, help="Batch size")
|
||||||
parser.add_argument('--lr', default="1e-4", help="Learning rate")
|
parser.add_argument("-n", "--nsave", type=int, 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('--rev', default="v4.0-rc1", help="Revision/log ID")
|
||||||
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('--lr', default="1e-4", help="Learning rate")
|
||||||
|
parser.add_argument('--lrskip', type=int, default=0, help="Constant lr for first N steps")
|
||||||
|
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler")
|
||||||
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")
|
||||||
@@ -33,120 +35,28 @@ 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
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
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"
|
|
||||||
resolution = 768
|
|
||||||
|
|
||||||
dataset = LatentDataset(resolution, args.src, args.dst)
|
dataset = LatentDataset([args.src, args.dst])
|
||||||
loader = DataLoader(
|
loader = DataLoader(
|
||||||
dataset,
|
dataset,
|
||||||
batch_size=args.bs,
|
batch_size = args.batch,
|
||||||
shuffle=True,
|
shuffle = True,
|
||||||
num_workers=0,
|
drop_last = True,
|
||||||
# num_workers=4,
|
pin_memory = False,
|
||||||
# persistent_workers=True,
|
# num_workers = 0,
|
||||||
|
num_workers = 4,
|
||||||
|
persistent_workers=True,
|
||||||
)
|
)
|
||||||
eval_src, eval_dst = get_eval_data(
|
model = Model() # TODO: handle scale factor/channels for non-sd VAEs
|
||||||
dataset,
|
|
||||||
f"latents/test_{args.src}_{resolution}px.npy",
|
|
||||||
f"latents/test_{args.dst}_{resolution}px.npy",
|
|
||||||
target_dev,
|
|
||||||
)
|
|
||||||
|
|
||||||
os.makedirs("models", exist_ok=True)
|
|
||||||
log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
|
|
||||||
|
|
||||||
model = Interposer()
|
|
||||||
|
|
||||||
criterion = torch.nn.L1Loss()
|
criterion = torch.nn.L1Loss()
|
||||||
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr))
|
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
|
scheduler = None
|
||||||
if args.cosine:
|
if args.cosine:
|
||||||
print("Using CosineAnnealingLR")
|
print("Using CosineAnnealingLR")
|
||||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||||
optimizer, T_max = int(args.steps/args.bs),
|
optimizer, T_max = int(args.steps/args.batch),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
print("Using LinearLR")
|
print("Using LinearLR")
|
||||||
@@ -154,23 +64,35 @@ if __name__ == "__main__":
|
|||||||
optimizer,
|
optimizer,
|
||||||
start_factor = 0.1,
|
start_factor = 0.1,
|
||||||
end_factor = 1.0,
|
end_factor = 1.0,
|
||||||
total_iters = int(5000/args.bs),
|
total_iters = int(5000/args.batch),
|
||||||
)
|
)
|
||||||
|
|
||||||
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(
|
optimizer.load_state_dict(torch.load(
|
||||||
f"{os.path.splitext(args.resume)[0]}.optim.pth"
|
f"{os.path.splitext(args.resume)[0]}.optim.pth"
|
||||||
))
|
))
|
||||||
|
optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0]
|
||||||
else:
|
else:
|
||||||
model.to(target_dev)
|
model.to(TARGET_DEV)
|
||||||
|
|
||||||
|
wrapper = ModelWrapper( # model wrapper for saving/eval/etc
|
||||||
|
name = f"{args.src}-to-{args.dst}_interposer-{args.rev}",
|
||||||
|
specs = [args.src, args.dst],
|
||||||
|
model = model,
|
||||||
|
evals = dataset.get_eval(),
|
||||||
|
device = TARGET_DEV,
|
||||||
|
criterion = criterion,
|
||||||
|
optimizer = optimizer,
|
||||||
|
scheduler = scheduler,
|
||||||
|
)
|
||||||
|
|
||||||
progress = tqdm(total=args.steps)
|
progress = tqdm(total=args.steps)
|
||||||
while progress.n < args.steps:
|
while progress.n < args.steps:
|
||||||
for src, dst in loader:
|
for src, dst in loader:
|
||||||
src = src.to(target_dev)
|
src = src.to(TARGET_DEV)
|
||||||
dst = dst.to(target_dev)
|
dst = dst.to(TARGET_DEV)
|
||||||
with torch.cuda.amp.autocast():
|
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
|
||||||
@@ -179,19 +101,15 @@ if __name__ == "__main__":
|
|||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
loss.backward()
|
loss.backward()
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
scheduler.step()
|
if progress.n >= args.lrskip: scheduler.step()
|
||||||
|
|
||||||
# eval/save
|
# eval/save
|
||||||
progress.update(args.bs)
|
progress.update(args.batch)
|
||||||
if progress.n % (1000 + 1000%args.bs) == 0:
|
wrapper.log_step(loss.data.item(), progress.n)
|
||||||
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
|
if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0:
|
||||||
if progress.n % (args.save + args.save%args.bs) == 0:
|
wrapper.save_model(step=progress.n)
|
||||||
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
|
|
||||||
if progress.n >= args.steps:
|
if progress.n >= args.steps:
|
||||||
break
|
break
|
||||||
progress.close()
|
progress.close()
|
||||||
|
wrapper.save_model(epoch="") # final save
|
||||||
# save final output
|
wrapper.close()
|
||||||
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
|
|
||||||
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
|
|
||||||
log.close()
|
|
||||||
|
|||||||
@@ -0,0 +1,123 @@
|
|||||||
|
#
|
||||||
|
# This file just has all the random saving/logging/eval related code
|
||||||
|
#
|
||||||
|
import os
|
||||||
|
import torch
|
||||||
|
from tqdm import tqdm
|
||||||
|
from diffusers import AutoencoderKL
|
||||||
|
from safetensors.torch import save_file
|
||||||
|
from torchvision.utils import save_image
|
||||||
|
|
||||||
|
LOSS_MEMORY = 500
|
||||||
|
LOG_EVERY_N = 500
|
||||||
|
SAVE_FOLDER = "models"
|
||||||
|
|
||||||
|
class ModelWrapper:
|
||||||
|
def __init__(self, name, specs, model, optimizer, criterion, scheduler, device="cpu", evals=[None,None], stdout=True):
|
||||||
|
self.name = name
|
||||||
|
self.specs = specs
|
||||||
|
self.losses = []
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
self.optimizer = optimizer
|
||||||
|
self.criterion = criterion
|
||||||
|
self.scheduler = scheduler
|
||||||
|
|
||||||
|
self.device = device
|
||||||
|
self.vae = self.get_vae(self.specs[1], fp16=True)
|
||||||
|
self.eval_src = evals[0]
|
||||||
|
self.eval_dst = evals[1]
|
||||||
|
|
||||||
|
os.makedirs(SAVE_FOLDER, exist_ok=True)
|
||||||
|
self.csvlog = open(f"{SAVE_FOLDER}/{self.name}.csv", "w")
|
||||||
|
self.stdout = stdout
|
||||||
|
|
||||||
|
def log_step(self, loss, step=None):
|
||||||
|
self.losses.append(loss)
|
||||||
|
step = step if step else len(self.losses)
|
||||||
|
if step % LOG_EVERY_N == 0:
|
||||||
|
self.log_main(step)
|
||||||
|
|
||||||
|
def log_main(self, step=None):
|
||||||
|
lr = float(self.scheduler.get_last_lr()[0])
|
||||||
|
avg = sum(self.losses[-LOSS_MEMORY:])/LOSS_MEMORY
|
||||||
|
evl = self.eval_model()[0]
|
||||||
|
if self.stdout:
|
||||||
|
tqdm.write(f"{str(step):<10} {avg:.4e}|{evl:.4e} @ {lr:.4e}")
|
||||||
|
if self.csvlog:
|
||||||
|
self.csvlog.write(f"{step},{avg},{evl},{lr}\n")
|
||||||
|
self.csvlog.flush()
|
||||||
|
|
||||||
|
def eval_model(self):
|
||||||
|
with torch.no_grad():
|
||||||
|
pred = self.model(self.eval_src.to(self.device))
|
||||||
|
loss = self.criterion(pred, self.eval_dst.to(self.device))
|
||||||
|
return loss, pred
|
||||||
|
|
||||||
|
def save_model(self, step=None, epoch=None):
|
||||||
|
step = step if step else len(self.losses)
|
||||||
|
if epoch is None and step >= 10**6:
|
||||||
|
epoch = f"_e{round(step/10**6,2)}M"
|
||||||
|
elif epoch is None:
|
||||||
|
epoch = f"_e{round(step/10**3)}K"
|
||||||
|
output_name = f"./{SAVE_FOLDER}/{self.name}{epoch}"
|
||||||
|
if self.vae:
|
||||||
|
out = self.eval_model()[1]
|
||||||
|
img = self.vae_decode(out).detach()
|
||||||
|
save_image(img, f"{output_name}.png")
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
save_file(self.model.state_dict(), f"{output_name}.safetensors")
|
||||||
|
torch.save(self.optimizer.state_dict(), f"{output_name}.optim.pth")
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
del self.vae
|
||||||
|
self.csvlog.close()
|
||||||
|
|
||||||
|
def vae_decode(self, latent):
|
||||||
|
latent = latent.to(torch.float16).to("cuda")
|
||||||
|
out = self.vae.decode(latent).sample
|
||||||
|
out = out.float().to(latent.device)
|
||||||
|
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||||
|
return ((out + 1.0) / 2.0)
|
||||||
|
|
||||||
|
def get_vae(self, 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 == "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:
|
||||||
|
raise NotImplementedError(f"Unknown VAE version '{version}'")
|
||||||
|
|
||||||
|
# save VRAM
|
||||||
|
vae.to(dtype).to("cuda")
|
||||||
|
vae.decoder.eval()
|
||||||
|
vae.set_use_memory_efficient_attention_xformers(True)
|
||||||
|
vae.enable_xformers_memory_efficient_attention()
|
||||||
|
vae.enable_gradient_checkpointing()
|
||||||
|
del vae.encoder
|
||||||
|
return vae
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
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