From 92f9b64e8bc0b314c4f0503721ba7f67934e8338 Mon Sep 17 00:00:00 2001
From: City <125218114+city96@users.noreply.github.com>
Date: Sat, 11 Nov 2023 22:19:42 +0100
Subject: [PATCH 1/4] Snapshot
---
dataset.py | 89 +++++++++++++++++++++++++
interposer.py | 75 ++++++++++-----------
log_loss.py | 54 ++++++++++-----
train.py | 178 ++++++++++++++------------------------------------
utils.py | 123 ++++++++++++++++++++++++++++++++++
vae.py | 48 --------------
6 files changed, 337 insertions(+), 230 deletions(-)
create mode 100644 dataset.py
create mode 100644 utils.py
delete mode 100644 vae.py
diff --git a/dataset.py b/dataset.py
new file mode 100644
index 0000000..cb4f71d
--- /dev/null
+++ b/dataset.py
@@ -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])
diff --git a/interposer.py b/interposer.py
index e8db42b..1c97a53 100644
--- a/interposer.py
+++ b/interposer.py
@@ -1,56 +1,57 @@
import torch
import torch.nn as nn
-import numpy as np
-class Block(nn.Module):
- def __init__(self, size):
+class ResBlock(nn.Module):
+ """Block with residuals"""
+ def __init__(self, ch):
super().__init__()
self.join = nn.ReLU()
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.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.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
- nn.Dropout(0.2)
+ nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
)
def forward(self, x):
- y = self.long(x)
- z = self.join(y + x)
- return z
+ return self.join(self.long(x) + x)
-class Interposer(nn.Module):
- def __init__(self):
+class ExtractBlock(nn.Module):
+ """Increase no. of channels by [out/in]"""
+ def __init__(self, ch_in, ch_out):
super().__init__()
- self.chan = 4 # in/out channels
- self.hid = 128
+ self.join = nn.ReLU()
+ 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
- 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(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
+class InterposerModel(nn.Module):
+ """Main neural network"""
+ def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0):
+ super().__init__()
+ self.scale = scale
+ self.ch_in = ch_in
+ self.ch_out = ch_out
+ self.ch_mid = ch_mid
+
+ self.head = ExtractBlock(self.ch_in, self.ch_mid)
self.core = nn.Sequential(
- Block(self.hid),
- Block(self.hid),
- Block(self.hid),
- )
- # reduce channels
- self.tail = nn.Sequential(
- nn.ReLU(),
- nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
+ nn.Upsample(scale_factor=self.scale, mode="nearest"),
+ 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),
+ ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
)
+ self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
def forward(self, x):
- y = self.head_join(
- self.head_long(x)+
- self.head_short(x)
- )
+ y = self.head(x)
z = self.core(y)
return self.tail(z)
diff --git a/log_loss.py b/log_loss.py
index acf17ad..9fd0da9 100644
--- a/log_loss.py
+++ b/log_loss.py
@@ -5,20 +5,36 @@ import matplotlib.pyplot as plt
files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")]
train_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):
global train_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]
train_loss[name] = (
[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:
eval_loss[name] = (
[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
@@ -31,18 +47,26 @@ def smooth(scalars, weight):
last = smoothed_val
return smoothed
-def plot(data, fname):
+def plot(data, fname, title=None, smw=0.9):
fig, ax = plt.subplots()
+ plt.tight_layout()
ax.grid()
+ dmax = 0
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')
+ data = [x + offsets[name] for x in val[0]] if name in offsets.keys() else val[0]
+ dmax = max(dmax, round(data[-1],10000))
+ 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:
- with open(fp) as f:
- lines = f.readlines()
- process_lines(lines)
-
-plot(train_loss, "loss.png")
-plot(eval_loss, "loss-eval.png")
+if __name__ == "__main__":
+ for fp in files:
+ with open(fp) as f:
+ lines = f.readlines()
+ process_lines(lines)
+ plot(train_loss, "loss.png", f"{model} Training loss", 0.2)
+ plot(eval_loss, "loss-eval.png", f"{model} Eval. loss", 0.7)
+ plot(lr_vals, "loss-lr.png", f"{model} Learning rate", 0.0)
diff --git a/train.py b/train.py
index 832bde1..afb8541 100644
--- a/train.py
+++ b/train.py
@@ -1,29 +1,31 @@
import os
import torch
-import torch.nn as nn
-import numpy as np
import argparse
-from PIL import Image
from tqdm import tqdm
-from safetensors.torch import save_file, load_file
-from torch.utils.data import DataLoader, Dataset
+from torch.utils.data import DataLoader
+from safetensors.torch import load_file
-from interposer import Interposer
-from vae import get_vae
+from interposer import InterposerModel as Model
+from dataset import LatentDataset
+from utils import ModelWrapper
torch.backends.cudnn.benchmark = True
torch.manual_seed(0)
+TARGET_DEV = "cuda"
+
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=4, help="Batch size")
- 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("-s", "--steps", type=int, default=500000, help="No. of training steps")
+ parser.add_argument("-b", "--batch", type=int, default= 1, help="Batch size")
+ parser.add_argument("-n", "--nsave", type=int, 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('--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('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR")
args = parser.parse_args()
if args.src == args.dst:
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")
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__":
args = parse_args()
- target_dev = "cuda"
- resolution = 768
- dataset = LatentDataset(resolution, args.src, args.dst)
+ dataset = LatentDataset([args.src, args.dst])
loader = DataLoader(
dataset,
- batch_size=args.bs,
- shuffle=True,
- num_workers=0,
- # num_workers=4,
- # persistent_workers=True,
+ batch_size = args.batch,
+ shuffle = True,
+ drop_last = True,
+ pin_memory = False,
+ # 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,
- )
-
- os.makedirs("models", exist_ok=True)
- log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
-
- model = Interposer()
-
+ model = Model() # TODO: handle scale factor/channels for non-sd VAEs
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),
+ optimizer, T_max = int(args.steps/args.batch),
)
else:
print("Using LinearLR")
@@ -154,23 +64,35 @@ if __name__ == "__main__":
optimizer,
start_factor = 0.1,
end_factor = 1.0,
- total_iters = int(5000/args.bs),
+ total_iters = int(5000/args.batch),
)
if 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"
))
+ optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0]
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)
while progress.n < args.steps:
for src, dst in loader:
- src = src.to(target_dev)
- dst = dst.to(target_dev)
+ src = src.to(TARGET_DEV)
+ dst = dst.to(TARGET_DEV)
with torch.cuda.amp.autocast():
y_pred = model(src) # forward
loss = criterion(y_pred, dst) # loss
@@ -179,19 +101,15 @@ if __name__ == "__main__":
optimizer.zero_grad()
loss.backward()
optimizer.step()
- scheduler.step()
+ if progress.n >= args.lrskip: 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, optimizer, eval_src, args.src, args.dst)
+ progress.update(args.batch)
+ wrapper.log_step(loss.data.item(), progress.n)
+ if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0:
+ wrapper.save_model(step=progress.n)
if progress.n >= args.steps:
break
progress.close()
-
- # save final output
- 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()
+ wrapper.save_model(epoch="") # final save
+ wrapper.close()
diff --git a/utils.py b/utils.py
new file mode 100644
index 0000000..ed550aa
--- /dev/null
+++ b/utils.py
@@ -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
diff --git a/vae.py b/vae.py
deleted file mode 100644
index 0bd4f9e..0000000
--- a/vae.py
+++ /dev/null
@@ -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
From 8ccc8208b24e0936dbff7b3c90075f2e0873f743 Mon Sep 17 00:00:00 2001
From: City <125218114+city96@users.noreply.github.com>
Date: Mon, 18 Mar 2024 23:29:15 +0100
Subject: [PATCH 2/4] Version 4
---
comfy_latent_interposer.py | 189 ++++++++++++++----------
config/ca-to-v1.yaml | 40 +++++
config/ca-to-xl.yaml | 40 +++++
config/v1-to-xl.yaml | 40 +++++
config/xl-to-v1.yaml | 40 +++++
dataset.py | 94 ++++++------
interposer.py | 22 +--
log_loss.py | 72 ---------
train.py | 292 +++++++++++++++++++++++++++----------
vae.py | 114 +++++++++++++++
10 files changed, 666 insertions(+), 277 deletions(-)
create mode 100644 config/ca-to-v1.yaml
create mode 100644 config/ca-to-xl.yaml
create mode 100644 config/v1-to-xl.yaml
create mode 100644 config/xl-to-v1.yaml
delete mode 100644 log_loss.py
create mode 100644 vae.py
diff --git a/comfy_latent_interposer.py b/comfy_latent_interposer.py
index 6306187..ecd7451 100644
--- a/comfy_latent_interposer.py
+++ b/comfy_latent_interposer.py
@@ -4,111 +4,154 @@ import torch.nn as nn
from safetensors.torch import load_file
from huggingface_hub import hf_hub_download
+# v1 = Stable Diffusion 1.x
+# xl = Stable Diffusion Extra Large (SDXL)
+# cc = Stable Cascade (Stage C) [not used]
+# ca = Stable Cascade (Stage A/B)
+config = {
+ "v1-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
+ "xl-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
+ "ca-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
+ "ca-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
+}
-class Interposer(nn.Module):
- """
- Basic NN layout, ported from:
- https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py
- """
- version = 3.1 # network revision
- def __init__(self):
+class ResBlock(nn.Module):
+ """Block with residuals"""
+ def __init__(self, ch):
super().__init__()
- self.chan = 4
- self.hid = 128
+ self.join = nn.ReLU()
+ self.norm = nn.BatchNorm2d(ch)
+ self.long = nn.Sequential(
+ nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
+ nn.SiLU(),
+ nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
+ nn.SiLU(),
+ nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
+ nn.Dropout(0.1)
+ )
+ def forward(self, x):
+ x = self.norm(x)
+ return self.join(self.long(x) + x)
- 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(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),
+class ExtractBlock(nn.Module):
+ """Increase no. of channels by [out/in]"""
+ def __init__(self, ch_in, ch_out):
+ super().__init__()
+ self.join = nn.ReLU()
+ 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.SiLU(),
+ nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
+ nn.SiLU(),
+ 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))
+
+class InterposerModel(nn.Module):
+ """
+ NN layout, ported from:
+ https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py
+ """
+ def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, blocks=12):
+ super().__init__()
+ self.ch_in = ch_in
+ self.ch_out = ch_out
+ self.ch_mid = ch_mid
+ self.blocks = blocks
+ self.scale = scale
+
+ self.head = ExtractBlock(self.ch_in, self.ch_mid)
self.core = nn.Sequential(
- Block(self.hid),
- Block(self.hid),
- Block(self.hid),
- )
- self.tail = nn.Sequential(
- nn.ReLU(),
- nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
+ nn.Upsample(scale_factor=self.scale, mode="nearest"),
+ *[ResBlock(self.ch_mid) for _ in range(blocks)],
+ nn.BatchNorm2d(self.ch_mid),
+ nn.SiLU(),
)
+ self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
def forward(self, x):
- y = self.head_join(
- self.head_long(x)+
- self.head_short(x)
- )
+ y = self.head(x)
z = self.core(y)
return self.tail(z)
-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),
- )
- def forward(self, x):
- y = self.long(x)
- z = self.join(y + x)
- return z
-
-
-class LatentInterposer:
+class ComfyLatentInterposer:
+ """Custom node"""
def __init__(self):
- pass
+ self.version = 4.0 # network revision
+ self.loaded = None # current model name
+ self.model = None # current model
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"samples": ("LATENT", ),
- "latent_src": (["v1", "xl"],),
- "latent_dst": (["v1", "xl"],),
+ "latent_src": (["v1", "xl", "ca"],),
+ "latent_dst": (["v1", "xl", "ca"],),
}
}
RETURN_TYPES = ("LATENT",)
- FUNCTION = "convert"
- CATEGORY = "latent"
+ FUNCTION = "convert"
+ CATEGORY = "latent"
+ TITLE = "Latent Interposer"
+
+ def get_model_path(self, model_name):
+ fname = f"{model_name}_interposer-v{self.version}.safetensors"
+ path = os.path.join(os.path.dirname(os.path.realpath(__file__)),"models")
+
+ # local path: [models/xl-to-v1_interposer-v4.2.safetensors]
+ if os.path.isfile(os.path.join(path, fname)):
+ print("LatentInterposer: Using local model")
+ return os.path.join(path, fname)
+
+ # local path: [models/v4.2/xl-to-v1_interposer-v4.2.safetensors]
+ if os.path.isfile(os.path.join(path, os.path.join(f"v{self.version}", fname))):
+ print("LatentInterposer: Using local model")
+ return os.path.join(path, os.path.join(f"v{self.version}", fname))
+
+ # huggingface hub fallback
+ print("LatentInterposer: Using HF Hub model")
+ return str(hf_hub_download(
+ repo_id = "city96/SD-Latent-Interposer",
+ subfolder = f"v{self.version}",
+ filename = fname,
+ ))
def convert(self, samples, latent_src, latent_dst):
+ samples = samples.copy()
if latent_src == latent_dst:
return (samples,)
- model = Interposer()
- model.eval()
- filename = f"{latent_src}-to-{latent_dst}_interposer-v{model.version}.safetensors"
- local = os.path.join(
- os.path.join(os.path.dirname(os.path.realpath(__file__)),"models"),
- filename
- )
- if os.path.isfile(local):
- print("LatentInterposer: Using local model")
- weights = local
- else:
- print("LatentInterposer: Using HF Hub model")
- weights = str(hf_hub_download(
- repo_id="city96/SD-Latent-Interposer",
- filename=filename)
- )
+ model_name = f"{latent_src}-to-{latent_dst}"
+ if model_name not in config:
+ raise ValueError(f"No model exists for this conversion! ({model_name})")
+
+ # only reload if changed
+ if self.loaded != model_name or self.model is None:
+ # load/init model
+ path = self.get_model_path(model_name)
+ model = InterposerModel(**config[model_name])
+ model.eval()
+ model.load_state_dict(load_file(path))
+ # keep for later runs
+ self.model = model
+ self.loaded = model_name
- model.load_state_dict(load_file(weights))
lt = samples["samples"]
- lt = model(lt)
- del model
- return ({"samples": lt},)
+ with torch.no_grad():
+ # force FP32, always run on CPU
+ lt = self.model(lt.cpu().float()).to(lt.device).to(lt.dtype)
+ samples["samples"] = lt
+ return (samples,)
NODE_CLASS_MAPPINGS = {
- "LatentInterposer": LatentInterposer,
+ "LatentInterposer": ComfyLatentInterposer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
- "LatentInterposer": "Latent Interposer"
+ "LatentInterposer": ComfyLatentInterposer.TITLE,
}
diff --git a/config/ca-to-v1.yaml b/config/ca-to-v1.yaml
new file mode 100644
index 0000000..24ba67c
--- /dev/null
+++ b/config/ca-to-v1.yaml
@@ -0,0 +1,40 @@
+steps: 20000
+batch: 48
+fconst: 0
+cosine: False
+resume: False
+device: "cuda"
+p_loss_weight: 1.0
+r_loss_weight: 1.4
+b_loss_weight: 1.0
+h_loss_weight: 0.0
+save_image: 100
+eval_model: 10
+
+model:
+ src: ca # Stable Cascade Stage A
+ dst: v1 # Stable Diffusion 1.x
+ rev: "v4.0-rc16"
+ args:
+ scale: 0.5
+ ch_in: 4
+ ch_out: 4
+ ch_mid: 64
+ blocks: 12
+
+optim:
+ lr: 5.0e-4
+ beta1: 0.5
+ beta2: 0.95
+
+dataset:
+ src: "./latents/ca_256px_combined.bin"
+ dst: "./latents/v1_256px_combined.bin"
+ preload: False
+ evals:
+ main:
+ src: "./latents/test_eru/test_ca_768px.npy"
+ dst: "./latents/test_eru/test_v1_768px.npy"
+ aux:
+ src: "./latents/test_bga/test_ca_768px.npy"
+ dst: "./latents/test_bga/test_v1_768px.npy"
diff --git a/config/ca-to-xl.yaml b/config/ca-to-xl.yaml
new file mode 100644
index 0000000..6fe47ad
--- /dev/null
+++ b/config/ca-to-xl.yaml
@@ -0,0 +1,40 @@
+steps: 20000
+batch: 48
+fconst: 0
+cosine: False
+resume: False
+device: "cuda"
+p_loss_weight: 1.0
+r_loss_weight: 1.4
+b_loss_weight: 1.0
+h_loss_weight: 0.0
+save_image: 100
+eval_model: 10
+
+model:
+ src: ca # Stable Cascade Stage A
+ dst: xl # Stable Diffusion Extra Large
+ rev: "v4.0-rc16"
+ args:
+ scale: 0.5
+ ch_in: 4
+ ch_out: 4
+ ch_mid: 64
+ blocks: 12
+
+optim:
+ lr: 5.0e-4
+ beta1: 0.5
+ beta2: 0.95
+
+dataset:
+ src: "./latents/ca_256px_combined.bin"
+ dst: "./latents/xl_256px_combined.bin"
+ preload: False
+ evals:
+ main:
+ src: "./latents/test_eru/test_ca_768px.npy"
+ dst: "./latents/test_eru/test_xl_768px.npy"
+ aux:
+ src: "./latents/test_bga/test_ca_768px.npy"
+ dst: "./latents/test_bga/test_xl_768px.npy"
diff --git a/config/v1-to-xl.yaml b/config/v1-to-xl.yaml
new file mode 100644
index 0000000..61ad3f4
--- /dev/null
+++ b/config/v1-to-xl.yaml
@@ -0,0 +1,40 @@
+steps: 50000
+batch: 128
+fconst: 35000
+cosine: True
+resume: False
+device: "cuda"
+p_loss_weight: 1.0
+r_loss_weight: 1.4
+b_loss_weight: 1.0
+h_loss_weight: 0.0
+save_image: 100
+eval_model: 10
+
+model:
+ src: v1 # Stable Diffusion 1.x
+ dst: xl # Stable Diffusion Extra Large
+ rev: "v4.0-rc15"
+ args:
+ scale: 1.0
+ ch_in: 4
+ ch_out: 4
+ ch_mid: 64
+ blocks: 12
+
+optim:
+ lr: 5.0e-4
+ beta1: 0.5
+ beta2: 0.95
+
+dataset:
+ src: "./latents/v1_256px_combined.bin"
+ dst: "./latents/xl_256px_combined.bin"
+ preload: False
+ evals:
+ main:
+ src: "./latents/test_eru/test_v1_768px.npy"
+ dst: "./latents/test_eru/test_xl_768px.npy"
+ aux:
+ src: "./latents/test_bga/test_v1_768px.npy"
+ dst: "./latents/test_bga/test_xl_768px.npy"
diff --git a/config/xl-to-v1.yaml b/config/xl-to-v1.yaml
new file mode 100644
index 0000000..0d8800c
--- /dev/null
+++ b/config/xl-to-v1.yaml
@@ -0,0 +1,40 @@
+steps: 50000
+batch: 128
+fconst: 30000
+cosine: True
+resume: False
+device: "cuda"
+p_loss_weight: 1.0
+r_loss_weight: 0.0
+b_loss_weight: 0.0
+h_loss_weight: 0.0
+save_image: 1000
+eval_model: 10
+
+model:
+ src: xl # Stable Diffusion Extra Large
+ dst: v1 # Stable Diffusion 1.x
+ rev: "v4.0-rc16"
+ args:
+ scale: 1.0
+ ch_in: 4
+ ch_out: 4
+ ch_mid: 64
+ blocks: 12
+
+optim:
+ lr: 5.0e-4
+ beta1: 0.5
+ beta2: 0.95
+
+dataset:
+ src: "./latents/xl_256px_combined.bin"
+ dst: "./latents/v1_256px_combined.bin"
+ preload: False
+ evals:
+ main:
+ src: "./latents/test_eru/test_xl_768px.npy"
+ dst: "./latents/test_eru/test_v1_768px.npy"
+ aux:
+ src: "./latents/test_bga/test_xl_768px.npy"
+ dst: "./latents/test_bga/test_v1_768px.npy"
diff --git a/dataset.py b/dataset.py
index cb4f71d..3cb2456 100644
--- a/dataset.py
+++ b/dataset.py
@@ -1,45 +1,37 @@
-# 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 FileLatentDataset(Dataset):
+ def __init__(self, src_file, dst_file, device="cpu", dtype=torch.float16):
+ assert os.path.isfile(src_file), f"src bin missing! ({src_file})"
+ assert os.path.isfile(dst_file), f"dst bin missing! ({dst_file})"
+ self.src_data = torch.load(src_file).to(dtype).to(device)
+ self.dst_data = torch.load(dst_file).to(dtype).to(device)
+ assert self.src_data.shape[0] == self.dst_data.shape[0], "Data size mismatch!"
+
+ def __len__(self):
+ return self.src_data.shape[0]
+
+ def __getitem__(self, index):
+ return {
+ "src": self.src_data[index].float(),
+ "dst": self.dst_data[index].float(),
+ }
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])
+ return all([os.path.isfile(x) for x in self.paths.values()])
def get_data(self):
if self.data is not None: return self.data
- return tuple([self.load_latent(x) for x in self.paths])
+ return {k:self.load_latent(v) for k,v in self.paths.items()}
def load_latent(self, path):
lat = torch.from_numpy(np.load(path))
@@ -52,29 +44,36 @@ class Shard:
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
- """
+ def __init__(self, src_root, dst_root, preload=True):
+ assert os.path.isdir(src_root), f"Source folder missing! ({src_root})"
+ assert os.path.isdir(dst_root), f"Destination folder missing! ({dst_root})"
+
print("Dataset: Parsing data from disk")
- self.specs = specs
- self.res = res
- self.root = root
+ fnames = list(
+ set(os.listdir(src_root)).intersection(
+ set(os.listdir(dst_root)))
+ )
+ assert len(fnames) > 0, "Source/destination have no overlapping files"
+
self.shards = []
- for fname in tqdm(os.listdir(f"{root}/{specs[0]}_{res}px")):
+ for fname in tqdm(fnames):
+ src_path = os.path.join(src_root, fname)
+ dst_path = os.path.join(dst_root, fname)
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 ext not in [".npy"]:
+ continue
+ shard = Shard({
+ "src": src_path,
+ "dst": dst_path,
+ })
if shard.exists():
self.shards.append(shard)
+ assert len(self.shards) > 0, "No valid files found."
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):
@@ -83,7 +82,14 @@ class LatentDataset(Dataset):
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])
+def load_evals(evals):
+ data = {}
+ for name, paths in evals.items():
+ shard = Shard(paths)
+ assert shard.exists(), f"Eval data missing ({name})"
+ data[name] = {}
+ for k, v in shard.get_data().items():
+ if len(v.shape) == 3:
+ v = v.unsqueeze(0)
+ data[name][k] = v.float()
+ return data
diff --git a/interposer.py b/interposer.py
index 1c97a53..a7288f0 100644
--- a/interposer.py
+++ b/interposer.py
@@ -6,14 +6,17 @@ class ResBlock(nn.Module):
def __init__(self, ch):
super().__init__()
self.join = nn.ReLU()
+ self.norm = nn.BatchNorm2d(ch)
self.long = nn.Sequential(
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
- nn.LeakyReLU(0.1),
+ nn.SiLU(),
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
- nn.LeakyReLU(0.1),
+ nn.SiLU(),
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
+ nn.Dropout(0.1)
)
def forward(self, x):
+ x = self.norm(x)
return self.join(self.long(x) + x)
class ExtractBlock(nn.Module):
@@ -24,9 +27,9 @@ class ExtractBlock(nn.Module):
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.SiLU(),
nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
- nn.LeakyReLU(0.1),
+ nn.SiLU(),
nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.1)
)
@@ -35,19 +38,20 @@ class ExtractBlock(nn.Module):
class InterposerModel(nn.Module):
"""Main neural network"""
- def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0):
+ def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, blocks=12):
super().__init__()
- self.scale = scale
self.ch_in = ch_in
self.ch_out = ch_out
self.ch_mid = ch_mid
+ self.blocks = blocks
+ self.scale = scale
self.head = ExtractBlock(self.ch_in, self.ch_mid)
self.core = nn.Sequential(
nn.Upsample(scale_factor=self.scale, mode="nearest"),
- 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),
- ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
+ *[ResBlock(self.ch_mid) for _ in range(blocks)],
+ nn.BatchNorm2d(self.ch_mid),
+ nn.SiLU(),
)
self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
diff --git a/log_loss.py b/log_loss.py
deleted file mode 100644
index 9fd0da9..0000000
--- a/log_loss.py
+++ /dev/null
@@ -1,72 +0,0 @@
-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 = {}
-lr_vals = {}
-fskip = 0
-
-offsets = { # offset to display resumed training runs
-}
-sep = ".csv"
-rep = "_interposer"
-model = "Latent Interposer"
-
-def process_lines(lines):
- global train_loss
- global eval_loss
- 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]
- train_loss[name] = (
- [int(x[0]) for x in vals],
- [math.log(float(x[1])+1e-10) for x in vals],
- )
- if len(vals[0]) >= 3:
- eval_loss[name] = (
- [int(x[0]) 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
-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, title=None, smw=0.9):
- fig, ax = plt.subplots()
- plt.tight_layout()
- ax.grid()
- dmax = 0
- for name, val in data.items():
- data = [x + offsets[name] for x in val[0]] if name in offsets.keys() else val[0]
- dmax = max(dmax, round(data[-1],10000))
- 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')
-
-if __name__ == "__main__":
- for fp in files:
- with open(fp) as f:
- lines = f.readlines()
- process_lines(lines)
- plot(train_loss, "loss.png", f"{model} Training loss", 0.2)
- plot(eval_loss, "loss-eval.png", f"{model} Eval. loss", 0.7)
- plot(lr_vals, "loss-lr.png", f"{model} Learning rate", 0.0)
diff --git a/train.py b/train.py
index afb8541..f2e6e5a 100644
--- a/train.py
+++ b/train.py
@@ -1,115 +1,249 @@
import os
+import yaml
import torch
import argparse
from tqdm import tqdm
from torch.utils.data import DataLoader
-from safetensors.torch import load_file
+from safetensors.torch import save_file, load_file
-from interposer import InterposerModel as Model
-from dataset import LatentDataset
-from utils import ModelWrapper
+from interposer import InterposerModel
+from dataset import LatentDataset, FileLatentDataset, load_evals
+from vae import load_vae
torch.backends.cudnn.benchmark = True
torch.manual_seed(0)
-TARGET_DEV = "cuda"
-
def parse_args():
parser = argparse.ArgumentParser(description="Train latent interposer model")
- parser.add_argument("-s", "--steps", type=int, default=500000, help="No. of training steps")
- parser.add_argument("-b", "--batch", type=int, default= 1, help="Batch size")
- parser.add_argument("-n", "--nsave", type=int, 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('--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("--config", help="Config for training")
args = parser.parse_args()
- if args.src == args.dst:
- parser.error("--src and --dst can't be the same")
- try:
- float(args.lr)
- except:
- parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
- return args
+ with open(args.config) as f:
+ conf = yaml.safe_load(f)
+ args.dataset = argparse.Namespace(**conf.pop("dataset"))
+ args.model = argparse.Namespace(**conf.pop("model"))
+ return argparse.Namespace(**vars(args), **conf)
+
+def eval_images(model, vae, evals):
+ preds = eval_model(model, evals, loss=False)
+ out = {}
+ for name, pred in preds.items():
+ images = vae.decode(pred).cpu().float()
+ # for image in images: # eval isn't batched
+ out[f"eval/{name}"] = images[0]
+ return out
+
+def eval_model(model, evals, loss=True):
+ model.eval()
+ preds = {}
+ losses = []
+ for name, data in evals.items():
+ src = data["src"].to(args.device)
+ dst = data["dst"].to(args.device)
+ with torch.no_grad():
+ pred = model(src)
+ if loss:
+ loss = torch.nn.functional.l1_loss(dst, pred)
+ losses.append(loss)
+ else:
+ preds[name] = pred
+ model.train()
+ if loss:
+ return (sum(losses) / len(losses)).data.item()
+ else:
+ return preds
+
+# from pytorch GAN tutorial
+def weights_init(m):
+ classname = m.__class__.__name__
+ if classname.find('Conv') != -1:
+ torch.nn.init.normal_(m.weight.data, 0.0, 0.02)
+ elif classname.find('BatchNorm') != -1:
+ torch.nn.init.normal_(m.weight.data, 1.0, 0.02)
+ torch.nn.init.constant_(m.bias.data, 0)
if __name__ == "__main__":
args = parse_args()
+ base_name = f"models/{args.model.src}-to-{args.model.dst}_interposer-{args.model.rev}"
- dataset = LatentDataset([args.src, args.dst])
+ # dataset
+ if os.path.isfile(args.dataset.src):
+ dataset = FileLatentDataset(
+ args.dataset.src,
+ args.dataset.dst,
+ )
+ elif os.path.isdir(args.dataset.src):
+ dataset = LatentDataset(
+ args.dataset.src,
+ args.dataset.dst,
+ args.dataset.preload
+ )
+ else:
+ raise OSError(f"Missing dataset source {args.dataset.src}")
loader = DataLoader(
dataset,
batch_size = args.batch,
shuffle = True,
drop_last = True,
pin_memory = False,
- # num_workers = 0,
- num_workers = 4,
- persistent_workers=True,
+ num_workers = 0,
+ # num_workers = 6,
+ # persistent_workers=True,
)
- model = Model() # TODO: handle scale factor/channels for non-sd VAEs
- criterion = torch.nn.L1Loss()
- optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr))
+
+ # evals
+ try:
+ evals = load_evals(args.dataset.evals)
+ except:
+ print(f"No evals, fallback to dataset.")
+ evals = dataset[0]
+
+ # defaults
+ crit = torch.nn.L1Loss()
+ optim_args = {
+ "lr": args.optim["lr"],
+ "betas": (args.optim["beta1"], args.optim["beta2"])
+ }
+
+ # model
+ model = InterposerModel(**args.model.args)
+ model.apply(weights_init)
+ model.to(args.device)
+ optim = torch.optim.AdamW(model.parameters(), **optim_args)
+
+ # aux model for reverse pass
+ model_back = InterposerModel(
+ ch_in = args.model.args["ch_out"],
+ ch_mid = args.model.args["ch_mid"],
+ ch_out = args.model.args["ch_in"],
+ scale = 1.0 / args.model.args["scale"],
+ blocks = args.model.args["blocks"],
+ )
+ model_back.apply(weights_init)
+ model_back.to(args.device)
+ optim_back = torch.optim.AdamW(model_back.parameters(), **optim_args)
+
+ # scheduler
scheduler = None
if args.cosine:
- print("Using CosineAnnealingLR")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
- optimizer, T_max = int(args.steps/args.batch),
- )
- else:
- print("Using LinearLR")
- scheduler = torch.optim.lr_scheduler.LinearLR(
- optimizer,
- start_factor = 0.1,
- end_factor = 1.0,
- total_iters = int(5000/args.batch),
+ optim,
+ T_max = (args.steps - args.fconst),
+ eta_min = 1e-8,
)
- if args.resume:
- model.load_state_dict(load_file(args.resume))
- model.to(TARGET_DEV)
- optimizer.load_state_dict(torch.load(
- f"{os.path.splitext(args.resume)[0]}.optim.pth"
- ))
- optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0]
- else:
- model.to(TARGET_DEV)
+ # vae
+ vae = None
+ if args.save_image:
+ vae = load_vae(args.model.dst, device=args.device, dtype=torch.float16, dec_only=True)
- 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,
- )
+ # main loop
+ import time
+ from torch.utils.tensorboard import SummaryWriter
+ writer = SummaryWriter(log_dir=f"{base_name}_{int(time.time())}")
- progress = tqdm(total=args.steps)
- while progress.n < args.steps:
- for src, dst in loader:
- src = src.to(TARGET_DEV)
- dst = dst.to(TARGET_DEV)
+ pbar = tqdm(total=args.steps)
+ while pbar.n < args.steps:
+ for batch in loader:
+ # get training data
+ src = batch.get("src").to(args.device)
+ dst = batch.get("dst").to(args.device)
+
+ ### Train main model ###
+ optim.zero_grad()
+ logs = {}
+ loss = []
with torch.cuda.amp.autocast():
- y_pred = model(src) # forward
- loss = criterion(y_pred, dst) # loss
+ # pass first model
+ pred = model(src)
- # backward
- optimizer.zero_grad()
+ p_loss = crit(pred, dst) * args.p_loss_weight
+ loss.append(p_loss)
+ logs["p_loss"] = p_loss.data.item()
+
+ # pass second model
+ if args.r_loss_weight:
+ pred_back = model_back(pred)
+
+ r_loss = crit(pred_back, src) * args.r_loss_weight
+ loss.append(r_loss)
+ logs["r_loss"] = r_loss.data.item()
+
+ # loss logic
+ loss = sum(loss)
+ logs["main"] = loss.data.item()
loss.backward()
- optimizer.step()
- if progress.n >= args.lrskip: scheduler.step()
+ optim.step()
- # eval/save
- progress.update(args.batch)
- wrapper.log_step(loss.data.item(), progress.n)
- if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0:
- wrapper.save_model(step=progress.n)
- if progress.n >= args.steps:
+ # logging
+ for name, value in logs.items():
+ writer.add_scalar(f"loss/{name}", value, pbar.n)
+
+ ### Train backwards model ###
+ if args.r_loss_weight:
+ optim_back.zero_grad()
+ logs = {}
+ loss = []
+ with torch.cuda.amp.autocast():
+ # pass second model
+ pred = model_back(dst)
+
+ p_loss = crit(pred, src) * args.b_loss_weight
+ loss.append(p_loss)
+ logs["p_loss"] = p_loss.data.item()
+
+ # pass first model
+ if args.h_loss_weight: # better w/o this?
+ pred_back = model(pred)
+
+ r_loss = crit(pred_back, dst) * args.h_loss_weight
+ loss.append(r_loss)
+ logs["r_loss"] = r_loss.data.item()
+
+ # loss logic
+ loss = sum(loss)
+ logs["main"] = loss.data.item()
+ loss.backward()
+ optim_back.step()
+
+ # logging
+ for name, value in logs.items():
+ writer.add_scalar(f"loss_aux/{name}", value, pbar.n)
+
+ # run eval/save eval image
+ if args.eval_model and pbar.n % args.eval_model == 0:
+ writer.add_scalar("loss/eval_loss", eval_model(model, evals), pbar.n)
+ if args.save_image and pbar.n % args.save_image == 0:
+ for name, image in eval_images(model, vae, evals).items():
+ writer.add_image(name, image, pbar.n)
+
+ # scheduler logic main
+ if scheduler is not None and pbar.n >= args.fconst:
+ lr = scheduler.get_last_lr()[0]
+ scheduler.step()
+ else:
+ lr = args.optim["lr"]
+ writer.add_scalar("lr/model", lr, pbar.n)
+
+ # aux model doesn't have a scheduler
+ writer.add_scalar("lr/model_aux", args.optim["lr"], pbar.n)
+
+ # step
+ pbar.update()
+ if pbar.n > args.steps:
break
- progress.close()
- wrapper.save_model(epoch="") # final save
- wrapper.close()
+
+ # hacky workaround when the colors are off.
+ # Save the last n versions and just pick the best one later.
+ # if pbar.n > (args.steps-2500) and pbar.n%500==0:
+ # from torchvision.utils import save_image
+ # save_file(model.state_dict(), f"{base_name}_{pbar.n:07}.safetensors")
+ # for name, image in eval_images(model, vae, evals).items():
+ # name = f"models/{name.replace('/', '_')}_{pbar.n:07}.png"
+ # save_image(image, name)
+
+ # final save/cleanup
+ pbar.close()
+ writer.close()
+
+ save_file(model.state_dict(), f"{base_name}.safetensors")
+ torch.save(optim.state_dict(), f"{base_name}.optim.pth")
diff --git a/vae.py b/vae.py
new file mode 100644
index 0000000..e4fa097
--- /dev/null
+++ b/vae.py
@@ -0,0 +1,114 @@
+import torch
+from diffusers import AutoencoderKL
+
+DTYPE = torch.float16
+DEVICE = "cuda:0"
+
+class SDv1_VAE:
+ scale = 1/8
+ channels = 4
+ def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
+ self.device = device
+ self.dtype = dtype
+ self.model = AutoencoderKL.from_pretrained(
+ "stabilityai/sd-vae-ft-mse"
+ )
+ self.model.eval().to(self.dtype).to(self.device)
+ if dec_only:
+ del self.model.encoder
+
+ def encode(self, image):
+ image = image.to(self.dtype).to(self.device)
+ image = (image * 2.0) - 1.0 # assuming input is [0;1]
+ with torch.no_grad():
+ latent = self.model.encode(image).latent_dist.sample()
+ return latent.to(image.dtype).to(image.device)
+
+ def decode(self, latent, grad=False):
+ latent = latent.to(self.dtype).to(self.device)
+ if grad:
+ out = self.model.decode(latent)[0]
+ else:
+ with torch.no_grad():
+ out = self.model.decode(latent).sample
+ out = torch.clamp(out, min=-1.0, max=1.0)
+ out = (out + 1.0) / 2.0
+ return out.to(latent.dtype).to(latent.device)
+
+class SDXL_VAE(SDv1_VAE):
+ scale = 1/8
+ channels = 4
+ def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
+ self.device = device
+ self.dtype = dtype
+ self.model = AutoencoderKL.from_pretrained(
+ "madebyollin/sdxl-vae-fp16-fix"
+ )
+ self.model.eval().to(self.dtype).to(self.device)
+ if dec_only:
+ del self.model.encoder
+
+class CascadeC_VAE(SDv1_VAE):
+ scale = 1/32
+ channels = 16
+ def __init__(self, device=DEVICE, dtype=DTYPE, **kwargs):
+ self.device = device
+ self.dtype = dtype
+
+ #For now this is just piggybacking off of koyha-ss/sd-scripts
+ from library import stable_cascade as sc
+ from safetensors.torch import load_file
+ from huggingface_hub import hf_hub_download
+
+ self.model = sc.EfficientNetEncoder()
+ self.model.load_state_dict(load_file(
+ str(hf_hub_download(
+ repo_id = "stabilityai/stable-cascade",
+ filename = "effnet_encoder.safetensors",
+ ))
+ ))
+ self.model.eval().to(self.dtype).to(self.device)
+
+class CascadeA_VAE():
+ scale = 1/4
+ channels = 4
+ def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
+ self.device = device
+ self.dtype = dtype
+
+ # not sure if this will change in the future?
+ from diffusers.pipelines.wuerstchen.modeling_paella_vq_model import PaellaVQModel
+ self.model = PaellaVQModel.from_pretrained(
+ "stabilityai/stable-cascade",
+ subfolder="vqgan"
+ )
+ self.model.eval().to(self.dtype).to(self.device)
+ if dec_only:
+ del self.model.encoder
+
+ def encode(self, image):
+ image = image.to(self.dtype).to(self.device)
+ with torch.no_grad():
+ latent = self.model.encode(image).latents
+ return latent.to(image.dtype).to(image.device)
+
+ def decode(self, latent, grad=False):
+ latent = latent.to(self.dtype).to(self.device)
+ if grad:
+ out = self.model.decode(latent)[0]
+ else:
+ with torch.no_grad():
+ out = self.model.decode(latent).sample
+ out = torch.clamp(out, min=0.0, max=1.0)
+ return out.to(latent.dtype).to(latent.device)
+
+def load_vae(ver, *args, **kwargs):
+ if ver == "v1":
+ VAE = SDv1_VAE
+ elif ver == "xl":
+ VAE = SDXL_VAE
+ elif ver == "cc":
+ VAE = CascadeC_VAE
+ elif ver == "ca":
+ VAE = CascadeA_VAE
+ return VAE(*args, **kwargs)
From 56bfde4ac8624d3903ff11c68e0d95ee5e0c98e2 Mon Sep 17 00:00:00 2001
From: City <125218114+city96@users.noreply.github.com>
Date: Mon, 18 Mar 2024 23:51:50 +0100
Subject: [PATCH 3/4] Update README.md
---
README.md | 57 ++++++++++++++++++++++++++++++++++++++++++++++++-------
1 file changed, 50 insertions(+), 7 deletions(-)
diff --git a/README.md b/README.md
index 18a3d8a..86a0dbd 100644
--- a/README.md
+++ b/README.md
@@ -4,14 +4,17 @@ A small neural network to provide interoperability between the latents generated
I wanted to see if it was possible to pass latents generated by the new SDXL model directly into SDv1.5 models without decoding and re-encoding them using a VAE first.
## Installation
-To install it, simply clone this repo to your custom_nodes folder using the following command: `git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer`.
+To install it, simply clone this repo to your custom_nodes folder using the following command:
+```
+git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer
+```
Alternatively, you can download the [comfy_latent_interposer.py](https://github.com/city96/SD-Latent-Interposer/raw/main/comfy_latent_interposer.py) file to your `ComfyUI/custom_nodes` folder as well. You may need to install hfhub using the command `pip install huggingface-hub` inside your venv.
-If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo.
+If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo. The current files are in the **"v4.0"** subfolder.
## Usage
-See the image below for an example on how to use it. xl=>v1 conversion is almost flawless, **v1=>xl seems to produce artifacts.**
+Simply place it where you would normally place a VAE decode followed by a VAE encode. Set the denoise as appropirate to hide any artifacts while keeping the composition. See image below.

@@ -20,14 +23,53 @@ Without the interposer, the two latent spaces are incompatible:

### Local models
-The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the modules there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
+The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the models there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
-Alternatively, just clone the entire HF repo to it: `git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models`
+Alternatively, just clone the entire HF repo to it:
+```
+git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models
+```
+
+### Supported Models
+
+Model names:
+
+| code | name |
+| ---- | -------------------------- |
+| `v1` | SDXL |
+| `xl` | Stable Diffusion v1.x |
+| `ca` | Stable Cascade (Stage A/B) |
+
+Available models:
+
+| From | to `v1` | to `xl` | to `ca` |
+|:----:|:-------:|:-------:|:-------:|
+| `v1` | - | v4.0 | No |
+| `xl` | v4.0 | - | No |
+| `ca` | v4.0 | v4.0 | - |
## Training
-Most of the training/preprocessing code is a 1:1 mirror from my latent upscaler. The folder layout it expects is also the same.
+
+The training code initializes most training parameters from the provided config file. The dataset should be a single .bin file saved with `torch.save` for each latent version. The format should be [batch, channels, height, width] with the "batch" being as large as the dataset, ie 88000.
+
+### Interposer v4.0
+
+The training code currently initializes two copies of the model, one in the target direction and one in the opposite. The losses are defined based on this.
+
+- `p_loss` is the main criterion for the primary model.
+- `b_loss` is the main criterion for the secondary one.
+- `r_loss` is the output of the primary model back through the secondary model and checked against the source latent (basically a round trip through the two models).
+- `h_loss` is the same as `r_loss` but for the secondary model.
+
+All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade).
+The training was done locally on an RTX 3080 and a Tesla V100S.
+
+### Older versions
+
+Interposer v3.1
### Interposer v3.1
+
This is basically a complete rewrite. Replaced the mediocre bunch of conv2d layers with something that looks more like a proper neural network. No VGG loss because I still don't have a better GPU.
Training was done on combined Flickr2K + DIV2K, with each image being processed into 6 1024x1024 segments. Padded with some of my random images for a total of 22,000 source images in the dataset.
@@ -38,7 +80,7 @@ v3.0 was 500k steps at a constant LR of 1e-4, v3.1 was 1M steps using a CosineAn

-### Older versions
+
Interposer v1.1
@@ -50,6 +92,7 @@ Overall, it seems to perform a lot better, especially for real life photos. I al
+
Interposer v1.0
### Interposer v1.0
From a91698a3b84df69e4a12c3054e86c01e7b08d0b8 Mon Sep 17 00:00:00 2001
From: City <125218114+city96@users.noreply.github.com>
Date: Wed, 20 Mar 2024 22:48:47 +0100
Subject: [PATCH 4/4] Update README.md
---
README.md | 2 ++
1 file changed, 2 insertions(+)
diff --git a/README.md b/README.md
index 86a0dbd..8ba2454 100644
--- a/README.md
+++ b/README.md
@@ -64,6 +64,8 @@ The training code currently initializes two copies of the model, one in the targ
All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade).
The training was done locally on an RTX 3080 and a Tesla V100S.
+
+
### Older versions
Interposer v3.1