Merge V4
This commit is contained in:
@@ -4,14 +4,17 @@ A small neural network to provide interoperability between the latents generated
|
||||
I wanted to see if it was possible to pass latents generated by the new SDXL model directly into SDv1.5 models without decoding and re-encoding them using a VAE first.
|
||||
|
||||
## Installation
|
||||
To install it, simply clone this repo to your custom_nodes folder using the following command: `git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer`.
|
||||
To install it, simply clone this repo to your custom_nodes folder using the following command:
|
||||
```
|
||||
git clone https://github.com/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer
|
||||
```
|
||||
|
||||
Alternatively, you can download the [comfy_latent_interposer.py](https://github.com/city96/SD-Latent-Interposer/raw/main/comfy_latent_interposer.py) file to your `ComfyUI/custom_nodes` folder as well. You may need to install hfhub using the command `pip install huggingface-hub` inside your venv.
|
||||
|
||||
If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo.
|
||||
If you need the model weights for something else, they are [hosted on HF](https://huggingface.co/city96/SD-Latent-Interposer/tree/main) under the same Apache2 license as the rest of the repo. The current files are in the **"v4.0"** subfolder.
|
||||
|
||||
## Usage
|
||||
See the image below for an example on how to use it. xl=>v1 conversion is almost flawless, **v1=>xl seems to produce artifacts.**
|
||||
Simply place it where you would normally place a VAE decode followed by a VAE encode. Set the denoise as appropirate to hide any artifacts while keeping the composition. See image below.
|
||||
|
||||

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

|
||||
|
||||
### Local models
|
||||
The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the modules there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
|
||||
The node pulls the required files from huggingface hub by default. You can create a `models` folder and place the models there if you have a flaky connection or prefer to use it completely offline. The custom node will prefer local files over HF when available. The path should be: `ComfyUI/custom_nodes/SD-Latent-Interposer/models`
|
||||
|
||||
Alternatively, just clone the entire HF repo to it: `git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models`
|
||||
Alternatively, just clone the entire HF repo to it:
|
||||
```
|
||||
git clone https://huggingface.co/city96/SD-Latent-Interposer custom_nodes/SD-Latent-Interposer/models
|
||||
```
|
||||
|
||||
### Supported Models
|
||||
|
||||
Model names:
|
||||
|
||||
| code | name |
|
||||
| ---- | -------------------------- |
|
||||
| `v1` | SDXL |
|
||||
| `xl` | Stable Diffusion v1.x |
|
||||
| `ca` | Stable Cascade (Stage A/B) |
|
||||
|
||||
Available models:
|
||||
|
||||
| From | to `v1` | to `xl` | to `ca` |
|
||||
|:----:|:-------:|:-------:|:-------:|
|
||||
| `v1` | - | v4.0 | No |
|
||||
| `xl` | v4.0 | - | No |
|
||||
| `ca` | v4.0 | v4.0 | - |
|
||||
|
||||
## Training
|
||||
Most of the training/preprocessing code is a 1:1 mirror from my latent upscaler. The folder layout it expects is also the same.
|
||||
|
||||
The training code initializes most training parameters from the provided config file. The dataset should be a single .bin file saved with `torch.save` for each latent version. The format should be [batch, channels, height, width] with the "batch" being as large as the dataset, ie 88000.
|
||||
|
||||
### Interposer v4.0
|
||||
|
||||
The training code currently initializes two copies of the model, one in the target direction and one in the opposite. The losses are defined based on this.
|
||||
|
||||
- `p_loss` is the main criterion for the primary model.
|
||||
- `b_loss` is the main criterion for the secondary one.
|
||||
- `r_loss` is the output of the primary model back through the secondary model and checked against the source latent (basically a round trip through the two models).
|
||||
- `h_loss` is the same as `r_loss` but for the secondary model.
|
||||
|
||||
All models were trained for 50000 steps with either batch size 128 (xl/v1) or 48 (cascade).
|
||||
The training was done locally on an RTX 3080 and a Tesla V100S.
|
||||
|
||||

|
||||
|
||||
### Older versions
|
||||
|
||||
<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
|
||||
|
||||

|
||||
|
||||
### 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
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
+95
@@ -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
@@ -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
@@ -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")
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user