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)