This commit is contained in:
City
2024-03-20 22:55:03 +01:00
committed by GitHub
12 changed files with 908 additions and 367 deletions
+52 -7
View File
@@ -4,14 +4,17 @@ A small neural network to provide interoperability between the latents generated
I wanted to see if it was possible to pass latents generated by the new SDXL model directly into SDv1.5 models without decoding and re-encoding them using a VAE first.
## Installation
To install it, simply clone this repo to your custom_nodes folder using the following command: `git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer`.
To install it, simply clone this repo to your custom_nodes folder using the following command:
```
git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer
```
Alternatively, you can download the [comfy_latent_interposer.py](https://github.com/city96/SD-Latent-Interposer/raw/main/comfy_latent_interposer.py) file to your `ComfyUI/custom_nodes` folder as well. You may need to install hfhub using the command `pip install huggingface-hub` inside your venv.
If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo.
If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo. The current files are in the **"v4.0"** subfolder.
## Usage
See the image below for an example on how to use it. xl=>v1 conversion is almost flawless, **v1=>xl seems to produce artifacts.**
Simply place it where you would normally place a VAE decode followed by a VAE encode. Set the denoise as appropirate to hide any artifacts while keeping the composition. See image below.
![LATENT_INTERPOSER_V3 1_TEST](https://github.com/city96/SD-Latent-Interposer/assets/125218114/849574b4-2565-4090-85d3-ae63ab425ee2)
@@ -20,14 +23,55 @@ Without the interposer, the two latent spaces are incompatible:
![LATENT_INTERPOSER_V3 1](https://github.com/city96/SD-Latent-Interposer/assets/125218114/13e2c01f-580e-4ecb-af1f-b6b21699127b)
### Local models
The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the modules there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the models there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
Alternatively, just clone the entire HF repo to it: `git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models`
Alternatively, just clone the entire HF repo to it:
```
git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models
```
### Supported Models
Model names:
| code | name |
| ---- | -------------------------- |
| `v1` | SDXL |
| `xl` | Stable Diffusion v1.x |
| `ca` | Stable Cascade (Stage A/B) |
Available models:
| From | to `v1` | to `xl` | to `ca` |
|:----:|:-------:|:-------:|:-------:|
| `v1` | - | v4.0 | No |
| `xl` | v4.0 | - | No |
| `ca` | v4.0 | v4.0 | - |
## Training
Most of the training/preprocessing code is a 1:1 mirror from my latent upscaler. The folder layout it expects is also the same.
The training code initializes most training parameters from the provided config file. The dataset should be a single .bin file saved with `torch.save` for each latent version. The format should be [batch, channels, height, width] with the "batch" being as large as the dataset, ie 88000.
### Interposer v4.0
The training code currently initializes two copies of the model, one in the target direction and one in the opposite. The losses are defined based on this.
- `p_loss` is the main criterion for the primary model.
- `b_loss` is the main criterion for the secondary one.
- `r_loss` is the output of the primary model back through the secondary model and checked against the source latent (basically a round trip through the two models).
- `h_loss` is the same as `r_loss` but for the secondary model.
All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade).
The training was done locally on an RTX 3080 and a Tesla V100S.
![LATENT_INTERPOSER_V4_LOSS](https://github.com/city96/SD-Latent-Interposer/assets/125218114/3a0d8920-ed48-42f0-96c9-897263525efb)
### Older versions
<details><summary>Interposer v3.1</summary>
### Interposer v3.1
This is basically a complete rewrite. Replaced the mediocre bunch of conv2d layers with something that looks more like a proper neural network. No VGG loss because I still don't have a better GPU.
Training was done on combined Flickr2K + DIV2K, with each image being processed into 6 1024x1024 segments. Padded with some of my random images for a total of 22,000 source images in the dataset.
@@ -38,7 +82,7 @@ v3.0 was 500k steps at a constant LR of 1e-4, v3.1 was 1M steps using a CosineAn
![INTERPOSER_V3 1](https://github.com/city96/SD-Latent-Interposer/assets/125218114/daff0ae2-4739-4cef-ba54-ac1d156d3388)
### Older versions
</details>
<details><summary>Interposer v1.1</summary>
@@ -50,6 +94,7 @@ Overall, it seems to perform a lot better, especially for real life photos. I al
</details>
<details><summary>Interposer v1.0</summary>
### Interposer v1.0
+116 -73
View File
@@ -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,
}
+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"
+95
View File
@@ -0,0 +1,95 @@
import os
import torch
import numpy as np
from tqdm import tqdm
from torch.utils.data import Dataset
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:
def __init__(self, paths):
self.paths = paths
self.data = None
def exists(self):
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 {k:self.load_latent(v) for k,v in self.paths.items()}
def load_latent(self, path):
lat = torch.from_numpy(np.load(path))
if lat.shape[0] == 1:
lat = torch.squeeze(lat, 0)
assert not torch.isnan(torch.sum(lat.float()))
return lat
def preload(self):
self.data = self.get_data()
class LatentDataset(Dataset):
def __init__(self, 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")
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(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 [".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):
return len(self.shards)
def __getitem__(self, index):
return self.shards[index].get_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
+44 -39
View File
@@ -1,56 +1,61 @@
import torch
import torch.nn as nn
import numpy as np
class Block(nn.Module):
def __init__(self, size):
class ResBlock(nn.Module):
"""Block with residuals"""
def __init__(self, ch):
super().__init__()
self.join = nn.ReLU()
self.norm = nn.BatchNorm2d(ch)
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),
nn.Dropout(0.2)
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):
y = self.long(x)
z = self.join(y + x)
return z
x = self.norm(x)
return self.join(self.long(x) + x)
class Interposer(nn.Module):
def __init__(self):
class ExtractBlock(nn.Module):
"""Increase no. of channels by [out/in]"""
def __init__(self, ch_in, ch_out):
super().__init__()
self.chan = 4 # in/out channels
self.hid = 128
self.join = nn.ReLU()
self.short = nn.Conv2d(ch_in, ch_out, kernel_size=3, stride=1, padding=1)
self.long = nn.Sequential(
nn.Conv2d( ch_in, ch_out, kernel_size=3, stride=1, padding=1),
nn.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))
# expand channels
self.head_join = nn.ReLU()
self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1)
self.head_long = nn.Sequential(
nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
)
# not sure if this is how residuals work
class InterposerModel(nn.Module):
"""Main neural network"""
def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0, 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),
)
# reduce channels
self.tail = nn.Sequential(
nn.ReLU(),
nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
nn.Upsample(scale_factor=self.scale, mode="nearest"),
*[ResBlock(self.ch_mid) 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)
-48
View File
@@ -1,48 +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 = {}
def process_lines(lines):
global train_loss
global eval_loss
name = fp.split("/")[1].split("_")[0]
vals = [x.split(",") for x in lines]
train_loss[name] = (
[int(x[0]) for x in vals],
[math.log(float(x[1])) for x in vals],
)
if len(vals[0]) >= 3:
eval_loss[name] = (
[int(x[0]) for x in vals],
[math.log(float(x[2])) 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):
fig, ax = plt.subplots()
ax.grid()
for name, val in data.items():
ax.plot(val[0], smooth(val[1], 0.9), label=name)
plt.legend(loc="upper right")
plt.savefig(fname, dpi=300, bbox_inches='tight')
for fp in files:
with open(fp) as f:
lines = f.readlines()
process_lines(lines)
plot(train_loss, "loss.png")
plot(eval_loss, "loss-eval.png")
+210 -158
View File
@@ -1,197 +1,249 @@
import os
import yaml
import torch
import torch.nn as nn
import numpy as np
import argparse
from PIL import Image
from tqdm import tqdm
from torch.utils.data import DataLoader
from safetensors.torch import save_file, load_file
from torch.utils.data import DataLoader, Dataset
from interposer import Interposer
from vae import get_vae
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)
def parse_args():
parser = argparse.ArgumentParser(description="Train latent interposer model")
parser.add_argument("--steps", type=int, default=500000, help="No. of training steps")
parser.add_argument('--bs', type=int, default=4, help="Batch size")
parser.add_argument('--lr', default="1e-4", help="Learning rate")
parser.add_argument("-n", "--save_every_n", type=int, dest="save", default=50000, help="Save model/sample periodically")
parser.add_argument('--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('--resume', help="Checkpoint to resume from")
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR")
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)
vae = None
def sample_decode(latent, filename, version):
global vae
if not vae:
vae = get_vae(version, fp16=True)
vae.to("cuda")
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
latent = latent.half().to("cuda")
out = vae.decode(latent).sample
out = out.cpu().detach().numpy()
out = np.squeeze(out, 0)
out = out.transpose((1, 2, 0))
out = np.clip(out, -1.0, 1.0)
out = (out+1)/2 * 255
out = out.astype(np.uint8)
out = Image.fromarray(out)
out.save(filename)
def get_eval_data(dataset, src_path, dst_path, target_dev):
if os.path.isfile(src_path) and os.path.isfile(dst_path):
src = LatentDataset.load_latent(None, src_path)
dst = LatentDataset.load_latent(None, dst_path)
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:
src = dataset[0][0]
dst = dataset[0][1]
src = src.float().to(target_dev).unsqueeze(0)
dst = dst.float().to(target_dev).unsqueeze(0)
return(src, dst)
return preds
def eval_model(step, model, criterion, scheduler, src, dst):
with torch.no_grad():
t_pred = model(src)
t_loss = criterion(t_pred, dst)
tqdm.write(f"{str(step):<10} {loss.data.item():.4e}|{t_loss.data.item():.4e} @ {float(scheduler.get_last_lr()[0]):.4e}")
log.write(f"{step},{loss.data.item()},{t_loss.data.item()},{float(scheduler.get_last_lr()[0])}\n")
log.flush()
def save_model(step, model, optim, lat, src, dst):
with torch.no_grad():
out = model(lat)
output_name = f"./models/{src}-to-{dst}_interposer_e{round(step/1000)}k"
sample_decode(out, f"{output_name}.png", dst)
save_file(model.state_dict(), f"{output_name}.safetensors")
torch.save(optim.state_dict(), f"{output_name}.optim.pth")
class LatentDataset(Dataset):
class Shard:
def __init__(self, root, fname, res, src, dst):
self.fname = fname
self.src_path = f"{root}/{src}_{res}px/{fname}.npy"
self.dst_path = f"{root}/{dst}_{res}px/{fname}.npy"
def __init__(self, res, src, dst, root="latents"):
print("Loading latents from disk")
self.latents = []
for i in tqdm(os.listdir(f"{root}/{src}_{res}px")):
fname, ext = os.path.splitext(i)
assert ext == ".npy"
s = self.Shard(root, fname, res, src, dst)
if os.path.isfile(s.src_path) and os.path.isfile(s.dst_path):
self.latents.append(s)
def __len__(self):
return len(self.latents)
def __getitem__(self, index):
s = self.latents[index]
src = self.load_latent(s.src_path)
dst = self.load_latent(s.dst_path)
return (src, dst)
def load_latent(self, path):
lat = torch.from_numpy(np.load(path))
if lat.shape[0] == 1:
lat = torch.squeeze(lat, 0)
assert not torch.isnan(torch.sum(lat.float()))
return lat
# 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()
target_dev = "cuda"
resolution = 768
base_name = f"models/{args.model.src}-to-{args.model.dst}_interposer-{args.model.rev}"
dataset = LatentDataset(resolution, 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.bs,
shuffle=True,
num_workers=0,
# num_workers=4,
batch_size = args.batch,
shuffle = True,
drop_last = True,
pin_memory = False,
num_workers = 0,
# num_workers = 6,
# persistent_workers=True,
)
eval_src, eval_dst = get_eval_data(
dataset,
f"latents/test_{args.src}_{resolution}px.npy",
f"latents/test_{args.dst}_{resolution}px.npy",
target_dev,
# 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)
os.makedirs("models", exist_ok=True)
log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
model = Interposer()
criterion = torch.nn.L1Loss()
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr))
# import bitsandbytes as bnb
# optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=float(args.lr))
# scheduler
scheduler = None
if args.cosine:
print("Using CosineAnnealingLR")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max = int(args.steps/args.bs),
)
else:
print("Using LinearLR")
scheduler = torch.optim.lr_scheduler.LinearLR(
optimizer,
start_factor = 0.1,
end_factor = 1.0,
total_iters = int(5000/args.bs),
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"
))
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)
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)
# main loop
import time
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter(log_dir=f"{base_name}_{int(time.time())}")
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()
scheduler.step()
optim.step()
# eval/save
progress.update(args.bs)
if progress.n % (1000 + 1000%args.bs) == 0:
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
if progress.n % (args.save + args.save%args.bs) == 0:
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
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()
# save final output
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
log.close()
# 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")
+123
View File
@@ -0,0 +1,123 @@
#
# This file just has all the random saving/logging/eval related code
#
import os
import torch
from tqdm import tqdm
from diffusers import AutoencoderKL
from safetensors.torch import save_file
from torchvision.utils import save_image
LOSS_MEMORY = 500
LOG_EVERY_N = 500
SAVE_FOLDER = "models"
class ModelWrapper:
def __init__(self, name, specs, model, optimizer, criterion, scheduler, device="cpu", evals=[None,None], stdout=True):
self.name = name
self.specs = specs
self.losses = []
self.model = model
self.optimizer = optimizer
self.criterion = criterion
self.scheduler = scheduler
self.device = device
self.vae = self.get_vae(self.specs[1], fp16=True)
self.eval_src = evals[0]
self.eval_dst = evals[1]
os.makedirs(SAVE_FOLDER, exist_ok=True)
self.csvlog = open(f"{SAVE_FOLDER}/{self.name}.csv", "w")
self.stdout = stdout
def log_step(self, loss, step=None):
self.losses.append(loss)
step = step if step else len(self.losses)
if step % LOG_EVERY_N == 0:
self.log_main(step)
def log_main(self, step=None):
lr = float(self.scheduler.get_last_lr()[0])
avg = sum(self.losses[-LOSS_MEMORY:])/LOSS_MEMORY
evl = self.eval_model()[0]
if self.stdout:
tqdm.write(f"{str(step):<10} {avg:.4e}|{evl:.4e} @ {lr:.4e}")
if self.csvlog:
self.csvlog.write(f"{step},{avg},{evl},{lr}\n")
self.csvlog.flush()
def eval_model(self):
with torch.no_grad():
pred = self.model(self.eval_src.to(self.device))
loss = self.criterion(pred, self.eval_dst.to(self.device))
return loss, pred
def save_model(self, step=None, epoch=None):
step = step if step else len(self.losses)
if epoch is None and step >= 10**6:
epoch = f"_e{round(step/10**6,2)}M"
elif epoch is None:
epoch = f"_e{round(step/10**3)}K"
output_name = f"./{SAVE_FOLDER}/{self.name}{epoch}"
if self.vae:
out = self.eval_model()[1]
img = self.vae_decode(out).detach()
save_image(img, f"{output_name}.png")
torch.cuda.empty_cache()
save_file(self.model.state_dict(), f"{output_name}.safetensors")
torch.save(self.optimizer.state_dict(), f"{output_name}.optim.pth")
def close(self):
del self.vae
self.csvlog.close()
def vae_decode(self, latent):
latent = latent.to(torch.float16).to("cuda")
out = self.vae.decode(latent).sample
out = out.float().to(latent.device)
out = torch.clamp(out, min=-1.0, max=1.0)
return ((out + 1.0) / 2.0)
def get_vae(self, version, file_path=None, fp16=False):
"""Load VAE from file or default hf repo. fp16 only works from hf"""
vae = None
dtype = torch.float16 if fp16 else torch.float32
if version == "v1" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=512,
)
elif version == "v1":
vae = AutoencoderKL.from_pretrained(
"runwayml/stable-diffusion-v1-5",
subfolder="vae",
torch_dtype=dtype,
)
elif version == "xl" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=1024
)
elif version == "xl" and fp16:
vae = AutoencoderKL.from_pretrained(
"madebyollin/sdxl-vae-fp16-fix",
torch_dtype=torch.float16,
)
elif version == "xl":
vae = AutoencoderKL.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
subfolder="vae"
)
else:
raise NotImplementedError(f"Unknown VAE version '{version}'")
# save VRAM
vae.to(dtype).to("cuda")
vae.decoder.eval()
vae.set_use_memory_efficient_attention_xformers(True)
vae.enable_xformers_memory_efficient_attention()
vae.enable_gradient_checkpointing()
del vae.encoder
return vae
+108 -42
View File
@@ -1,48 +1,114 @@
import torch
from diffusers import AutoencoderKL
def get_vae(version, file_path=None, fp16=False):
"""Load VAE from file or default hf repo. fp16 only works from hf"""
vae = None
dtype = torch.float16 if fp16 else torch.float32
if version == "v1" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=512,
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"
)
elif version == "v1":
vae = AutoencoderKL.from_pretrained(
"runwayml/stable-diffusion-v1-5",
subfolder="vae",
torch_dtype=dtype,
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"
)
elif version == "v2" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=768,
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"
)
elif version == "v2":
vae = AutoencoderKL.from_pretrained(
"stabilityai/stable-diffusion-2-1",
subfolder="vae",
torch_dtype=dtype,
)
elif version == "xl" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=1024
)
elif version == "xl" and fp16:
vae = AutoencoderKL.from_pretrained(
"madebyollin/sdxl-vae-fp16-fix",
torch_dtype=torch.float16,
)
elif version == "xl":
vae = AutoencoderKL.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
subfolder="vae"
)
else:
input("Invalid VAE version. Press any key to exit")
exit(1)
return vae
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)