Version 4

This commit is contained in:
City
2024-03-18 23:29:15 +01:00
parent 92f9b64e8b
commit 8ccc8208b2
10 changed files with 666 additions and 277 deletions
+112 -69
View File
@@ -4,111 +4,154 @@ import torch.nn as nn
from safetensors.torch import load_file from safetensors.torch import load_file
from huggingface_hub import hf_hub_download 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): class ResBlock(nn.Module):
"""Block with residuals"""
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.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)
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):
""" """
Basic NN layout, ported from: NN layout, ported from:
https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py https://github.com/city96/SD-Latent-Interposer/blob/main/interposer.py
""" """
version = 3.1 # network revision def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, blocks=12):
def __init__(self):
super().__init__() super().__init__()
self.chan = 4 self.ch_in = ch_in
self.hid = 128 self.ch_out = ch_out
self.ch_mid = ch_mid
self.blocks = blocks
self.scale = scale
self.head_join = nn.ReLU() self.head = ExtractBlock(self.ch_in, self.ch_mid)
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),
)
self.core = nn.Sequential( self.core = nn.Sequential(
Block(self.hid), nn.Upsample(scale_factor=self.scale, mode="nearest"),
Block(self.hid), *[ResBlock(self.ch_mid) for _ in range(blocks)],
Block(self.hid), nn.BatchNorm2d(self.ch_mid),
) nn.SiLU(),
self.tail = nn.Sequential(
nn.ReLU(),
nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
) )
self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
def forward(self, x): def forward(self, x):
y = self.head_join( y = self.head(x)
self.head_long(x)+
self.head_short(x)
)
z = self.core(y) z = self.core(y)
return self.tail(z) return self.tail(z)
class Block(nn.Module): class ComfyLatentInterposer:
def __init__(self, size): """Custom node"""
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:
def __init__(self): def __init__(self):
pass self.version = 4.0 # network revision
self.loaded = None # current model name
self.model = None # current model
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"samples": ("LATENT", ), "samples": ("LATENT", ),
"latent_src": (["v1", "xl"],), "latent_src": (["v1", "xl", "ca"],),
"latent_dst": (["v1", "xl"],), "latent_dst": (["v1", "xl", "ca"],),
} }
} }
RETURN_TYPES = ("LATENT",) RETURN_TYPES = ("LATENT",)
FUNCTION = "convert" FUNCTION = "convert"
CATEGORY = "latent" 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): def convert(self, samples, latent_src, latent_dst):
samples = samples.copy()
if latent_src == latent_dst: if latent_src == latent_dst:
return (samples,) return (samples,)
model = Interposer()
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.eval()
filename = f"{latent_src}-to-{latent_dst}_interposer-v{model.version}.safetensors" model.load_state_dict(load_file(path))
local = os.path.join( # keep for later runs
os.path.join(os.path.dirname(os.path.realpath(__file__)),"models"), self.model = model
filename self.loaded = model_name
)
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.load_state_dict(load_file(weights))
lt = samples["samples"] lt = samples["samples"]
lt = model(lt) with torch.no_grad():
del model # force FP32, always run on CPU
return ({"samples": lt},) lt = self.model(lt.cpu().float()).to(lt.device).to(lt.dtype)
samples["samples"] = lt
return (samples,)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"LatentInterposer": LatentInterposer, "LatentInterposer": ComfyLatentInterposer,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"LatentInterposer": "Latent Interposer" "LatentInterposer": ComfyLatentInterposer.TITLE,
} }
+40
View File
@@ -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"
+40
View File
@@ -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"
+40
View File
@@ -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"
+40
View File
@@ -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"
+50 -44
View File
@@ -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 os
import torch import torch
import numpy as np import numpy as np
from tqdm import tqdm from tqdm import tqdm
from torch.utils.data import Dataset from torch.utils.data import Dataset
DEFAULT_ROOT = "latents" class FileLatentDataset(Dataset):
ALLOWED_EXTS = [".npy"] 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: class Shard:
"""
Shard to store groups of latents in
paths: List containing paths to latent encoded images
"""
def __init__(self, paths): def __init__(self, paths):
self.paths = paths self.paths = paths
self.data = None self.data = None
def exists(self): 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): def get_data(self):
if self.data is not None: return self.data 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): def load_latent(self, path):
lat = torch.from_numpy(np.load(path)) lat = torch.from_numpy(np.load(path))
@@ -52,29 +44,36 @@ class Shard:
self.data = self.get_data() self.data = self.get_data()
class LatentDataset(Dataset): class LatentDataset(Dataset):
def __init__(self, specs, res=768, root=DEFAULT_ROOT, preload=False): def __init__(self, src_root, dst_root, preload=True):
""" assert os.path.isdir(src_root), f"Source folder missing! ({src_root})"
Main dataset that returns list of requested images as (C, H, W) latents assert os.path.isdir(dst_root), f"Destination folder missing! ({dst_root})"
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") print("Dataset: Parsing data from disk")
self.specs = specs fnames = list(
self.res = res set(os.listdir(src_root)).intersection(
self.root = root set(os.listdir(dst_root)))
)
assert len(fnames) > 0, "Source/destination have no overlapping files"
self.shards = [] 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) name, ext = os.path.splitext(fname)
if ext not in ALLOWED_EXTS: continue if ext not in [".npy"]:
shard = Shard([f"{root}/{x}_{res}px/{name}{ext}" for x in specs]) continue
shard = Shard({
"src": src_path,
"dst": dst_path,
})
if shard.exists(): if shard.exists():
self.shards.append(shard) self.shards.append(shard)
assert len(self.shards) > 0, "No valid files found."
if preload: # cache to RAM if preload: # cache to RAM
print("Dataset: Preloading data to system RAM") print("Dataset: Preloading data to system RAM")
[x.preload() for x in tqdm(self.shards)] [x.preload() for x in tqdm(self.shards)]
print(f"Dataset: OK, {len(self)} items") print(f"Dataset: OK, {len(self)} items")
def __len__(self): def __len__(self):
@@ -83,7 +82,14 @@ class LatentDataset(Dataset):
def __getitem__(self, index): def __getitem__(self, index):
return self.shards[index].get_data() return self.shards[index].get_data()
def get_eval(self): def load_evals(evals):
shard = Shard([f"{self.root}/test_{x}_{self.res}px.npy" for x in self.specs]) data = {}
data = shard.get_data() if shard.exists() else self[0] for name, paths in evals.items():
return tuple([x.unsqueeze(0).to(torch.float32) for x in data]) 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
+13 -9
View File
@@ -6,14 +6,17 @@ class ResBlock(nn.Module):
def __init__(self, ch): def __init__(self, ch):
super().__init__() super().__init__()
self.join = nn.ReLU() self.join = nn.ReLU()
self.norm = nn.BatchNorm2d(ch)
self.long = nn.Sequential( self.long = nn.Sequential(
nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1), 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.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.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.1)
) )
def forward(self, x): def forward(self, x):
x = self.norm(x)
return self.join(self.long(x) + x) return self.join(self.long(x) + x)
class ExtractBlock(nn.Module): 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.short = nn.Conv2d(ch_in, ch_out, kernel_size=3, stride=1, padding=1)
self.long = nn.Sequential( self.long = nn.Sequential(
nn.Conv2d( ch_in, ch_out, kernel_size=3, stride=1, padding=1), 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.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.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.1) nn.Dropout(0.1)
) )
@@ -35,19 +38,20 @@ class ExtractBlock(nn.Module):
class InterposerModel(nn.Module): class InterposerModel(nn.Module):
"""Main neural network""" """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__() super().__init__()
self.scale = scale
self.ch_in = ch_in self.ch_in = ch_in
self.ch_out = ch_out self.ch_out = ch_out
self.ch_mid = ch_mid self.ch_mid = ch_mid
self.blocks = blocks
self.scale = scale
self.head = ExtractBlock(self.ch_in, self.ch_mid) self.head = ExtractBlock(self.ch_in, self.ch_mid)
self.core = nn.Sequential( self.core = nn.Sequential(
nn.Upsample(scale_factor=self.scale, mode="nearest"), 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) for _ in range(blocks)],
ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), nn.BatchNorm2d(self.ch_mid),
ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), nn.SiLU(),
) )
self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1) self.tail = nn.Conv2d(self.ch_mid, self.ch_out, kernel_size=3, stride=1, padding=1)
-72
View File
@@ -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)
+213 -79
View File
@@ -1,115 +1,249 @@
import os import os
import yaml
import torch import torch
import argparse import argparse
from tqdm import tqdm from tqdm import tqdm
from torch.utils.data import DataLoader 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 interposer import InterposerModel
from dataset import LatentDataset from dataset import LatentDataset, FileLatentDataset, load_evals
from utils import ModelWrapper from vae import load_vae
torch.backends.cudnn.benchmark = True torch.backends.cudnn.benchmark = True
torch.manual_seed(0) torch.manual_seed(0)
TARGET_DEV = "cuda"
def parse_args(): def parse_args():
parser = argparse.ArgumentParser(description="Train latent interposer model") parser = argparse.ArgumentParser(description="Train latent interposer model")
parser.add_argument("-s", "--steps", type=int, default=500000, help="No. of training steps") parser.add_argument("--config", help="Config for training")
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")
args = parser.parse_args() args = parser.parse_args()
if args.src == args.dst: with open(args.config) as f:
parser.error("--src and --dst can't be the same") conf = yaml.safe_load(f)
try: args.dataset = argparse.Namespace(**conf.pop("dataset"))
float(args.lr) args.model = argparse.Namespace(**conf.pop("model"))
except: return argparse.Namespace(**vars(args), **conf)
parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
return args 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__": if __name__ == "__main__":
args = parse_args() 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( loader = DataLoader(
dataset, dataset,
batch_size = args.batch, batch_size = args.batch,
shuffle = True, shuffle = True,
drop_last = True, drop_last = True,
pin_memory = False, pin_memory = False,
# num_workers = 0, num_workers = 0,
num_workers = 4, # num_workers = 6,
persistent_workers=True, # persistent_workers=True,
) )
model = Model() # TODO: handle scale factor/channels for non-sd VAEs
criterion = torch.nn.L1Loss() # evals
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr)) 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 scheduler = None
if args.cosine: if args.cosine:
print("Using CosineAnnealingLR")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max = int(args.steps/args.batch), optim,
) T_max = (args.steps - args.fconst),
else: eta_min = 1e-8,
print("Using LinearLR")
scheduler = torch.optim.lr_scheduler.LinearLR(
optimizer,
start_factor = 0.1,
end_factor = 1.0,
total_iters = int(5000/args.batch),
) )
if args.resume: # vae
model.load_state_dict(load_file(args.resume)) vae = None
model.to(TARGET_DEV) if args.save_image:
optimizer.load_state_dict(torch.load( vae = load_vae(args.model.dst, device=args.device, dtype=torch.float16, dec_only=True)
f"{os.path.splitext(args.resume)[0]}.optim.pth"
))
optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0]
else:
model.to(TARGET_DEV)
wrapper = ModelWrapper( # model wrapper for saving/eval/etc # main loop
name = f"{args.src}-to-{args.dst}_interposer-{args.rev}", import time
specs = [args.src, args.dst], from torch.utils.tensorboard import SummaryWriter
model = model, writer = SummaryWriter(log_dir=f"{base_name}_{int(time.time())}")
evals = dataset.get_eval(),
device = TARGET_DEV,
criterion = criterion,
optimizer = optimizer,
scheduler = scheduler,
)
progress = tqdm(total=args.steps) pbar = tqdm(total=args.steps)
while progress.n < args.steps: while pbar.n < args.steps:
for src, dst in loader: for batch in loader:
src = src.to(TARGET_DEV) # get training data
dst = dst.to(TARGET_DEV) 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(): with torch.cuda.amp.autocast():
y_pred = model(src) # forward # pass first model
loss = criterion(y_pred, dst) # loss pred = model(src)
# backward p_loss = crit(pred, dst) * args.p_loss_weight
optimizer.zero_grad() 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() loss.backward()
optimizer.step() optim.step()
if progress.n >= args.lrskip: scheduler.step()
# eval/save # logging
progress.update(args.batch) for name, value in logs.items():
wrapper.log_step(loss.data.item(), progress.n) writer.add_scalar(f"loss/{name}", value, pbar.n)
if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0:
wrapper.save_model(step=progress.n) ### Train backwards model ###
if progress.n >= args.steps: 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 break
progress.close()
wrapper.save_model(epoch="") # final save # hacky workaround when the colors are off.
wrapper.close() # 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")
+114
View File
@@ -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)