Version 4
This commit is contained in:
+112
-69
@@ -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,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -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"
|
||||||
@@ -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"
|
||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
|
||||||
@@ -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")
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user