This commit is contained in:
City
2023-11-11 22:19:42 +01:00
parent bf5cec6eb4
commit 92f9b64e8b
6 changed files with 337 additions and 230 deletions
+89
View File
@@ -0,0 +1,89 @@
# Custom dataset to load encoded latents from disk.
# Files should contain latents as (1, C, H, W) or (C, H, W)
# Latents should be in their original format without scaling
######### Folder Layout #########
# latents #
# |- test_v1_768px.npy <=eval #
# |- test_xl_768px.npy <=^ #
# |- v1_768px <= ver/res #
# | |- 000001.npy #
# | |- 000002.npy #
# | | ... #
# | |- 000999.npy #
# | \- 001000.npy #
# |- xl_768px #
# ... #
#################################
import os
import torch
import numpy as np
from tqdm import tqdm
from torch.utils.data import Dataset
DEFAULT_ROOT = "latents"
ALLOWED_EXTS = [".npy"]
class Shard:
"""
Shard to store groups of latents in
paths: List containing paths to latent encoded images
"""
def __init__(self, paths):
self.paths = paths
self.data = None
def exists(self):
return all([os.path.isfile(x) for x in self.paths])
def get_data(self):
if self.data is not None: return self.data
return tuple([self.load_latent(x) for x in self.paths])
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, specs, res=768, root=DEFAULT_ROOT, preload=False):
"""
Main dataset that returns list of requested images as (C, H, W) latents
specs: List of latent versions in the other to return them in
res: Native resolution of images (before latent encoding)
root: Path to folder with sorted files
preload: Load all files into memory on initialization
"""
print("Dataset: Parsing data from disk")
self.specs = specs
self.res = res
self.root = root
self.shards = []
for fname in tqdm(os.listdir(f"{root}/{specs[0]}_{res}px")):
name, ext = os.path.splitext(fname)
if ext not in ALLOWED_EXTS: continue
shard = Shard([f"{root}/{x}_{res}px/{name}{ext}" for x in specs])
if shard.exists():
self.shards.append(shard)
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 get_eval(self):
shard = Shard([f"{self.root}/test_{x}_{self.res}px.npy" for x in self.specs])
data = shard.get_data() if shard.exists() else self[0]
return tuple([x.unsqueeze(0).to(torch.float32) for x in data])
+38 -37
View File
@@ -1,56 +1,57 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
import numpy as np
class Block(nn.Module): class ResBlock(nn.Module):
def __init__(self, size): """Block with residuals"""
def __init__(self, ch):
super().__init__() super().__init__()
self.join = nn.ReLU() self.join = nn.ReLU()
self.long = nn.Sequential( self.long = nn.Sequential(
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1), nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1), nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1), nn.Conv2d(ch, ch, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.2)
) )
def forward(self, x): def forward(self, x):
y = self.long(x) return self.join(self.long(x) + x)
z = self.join(y + x)
return z
class Interposer(nn.Module): class ExtractBlock(nn.Module):
def __init__(self): """Increase no. of channels by [out/in]"""
def __init__(self, ch_in, ch_out):
super().__init__() super().__init__()
self.chan = 4 # in/out channels self.join = nn.ReLU()
self.hid = 128 self.short = nn.Conv2d(ch_in, ch_out, kernel_size=3, stride=1, padding=1)
self.long = nn.Sequential(
nn.Conv2d( ch_in, ch_out, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(ch_out, ch_out, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
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 class InterposerModel(nn.Module):
self.head_join = nn.ReLU() """Main neural network"""
self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1) def __init__(self, ch_in=4, ch_out=4, ch_mid=64, scale=1.0):
self.head_long = nn.Sequential( super().__init__()
nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1), self.scale = scale
nn.LeakyReLU(0.1), self.ch_in = ch_in
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1), self.ch_out = ch_out
nn.LeakyReLU(0.1), self.ch_mid = ch_mid
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
) self.head = ExtractBlock(self.ch_in, self.ch_mid)
# not sure if this is how residuals work
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), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
Block(self.hid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
) ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid), ResBlock(self.ch_mid),
# reduce channels
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)
+39 -15
View File
@@ -5,20 +5,36 @@ import matplotlib.pyplot as plt
files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")] files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")]
train_loss = {} train_loss = {}
eval_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): def process_lines(lines):
global train_loss global train_loss
global eval_loss global eval_loss
name = fp.split("/")[1].split("_")[0] 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] vals = [x.split(",") for x in lines]
train_loss[name] = ( train_loss[name] = (
[int(x[0]) for x in vals], [int(x[0]) for x in vals],
[math.log(float(x[1])) for x in vals], [math.log(float(x[1])+1e-10) for x in vals],
) )
if len(vals[0]) >= 3: if len(vals[0]) >= 3:
eval_loss[name] = ( eval_loss[name] = (
[int(x[0]) for x in vals], [int(x[0]) for x in vals],
[math.log(float(x[2])) 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 # https://stackoverflow.com/a/49357445
@@ -31,18 +47,26 @@ def smooth(scalars, weight):
last = smoothed_val last = smoothed_val
return smoothed return smoothed
def plot(data, fname): def plot(data, fname, title=None, smw=0.9):
fig, ax = plt.subplots() fig, ax = plt.subplots()
plt.tight_layout()
ax.grid() ax.grid()
dmax = 0
for name, val in data.items(): for name, val in data.items():
ax.plot(val[0], smooth(val[1], 0.9), label=name) data = [x + offsets[name] for x in val[0]] if name in offsets.keys() else val[0]
plt.legend(loc="upper right") dmax = max(dmax, round(data[-1],10000))
plt.savefig(fname, dpi=300, bbox_inches='tight') 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')
for fp in files: if __name__ == "__main__":
with open(fp) as f: for fp in files:
lines = f.readlines() with open(fp) as f:
process_lines(lines) lines = f.readlines()
process_lines(lines)
plot(train_loss, "loss.png") plot(train_loss, "loss.png", f"{model} Training loss", 0.2)
plot(eval_loss, "loss-eval.png") plot(eval_loss, "loss-eval.png", f"{model} Eval. loss", 0.7)
plot(lr_vals, "loss-lr.png", f"{model} Learning rate", 0.0)
+48 -130
View File
@@ -1,29 +1,31 @@
import os import os
import torch import torch
import torch.nn as nn
import numpy as np
import argparse import argparse
from PIL import Image
from tqdm import tqdm from tqdm import tqdm
from safetensors.torch import save_file, load_file from torch.utils.data import DataLoader
from torch.utils.data import DataLoader, Dataset from safetensors.torch import load_file
from interposer import Interposer from interposer import InterposerModel as Model
from vae import get_vae from dataset import LatentDataset
from utils import ModelWrapper
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("--steps", type=int, default=500000, help="No. of training steps") parser.add_argument("-s", "--steps", type=int, default=500000, help="No. of training steps")
parser.add_argument('--bs', type=int, default=4, help="Batch size") parser.add_argument("-b", "--batch", type=int, default= 1, help="Batch size")
parser.add_argument('--lr', default="1e-4", help="Learning rate") parser.add_argument("-n", "--nsave", type=int, default= 50000, help="Save model/sample periodically")
parser.add_argument("-n", "--save_every_n", type=int, dest="save", 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('--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('--dst', choices=["v1","xl"], required=True, help="Destination latent format")
parser.add_argument('--lr', default="1e-4", help="Learning rate")
parser.add_argument('--lrskip', type=int, default=0, help="Constant lr for first N steps")
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler")
parser.add_argument('--resume', help="Checkpoint to resume from") parser.add_argument('--resume', help="Checkpoint to resume from")
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR")
args = parser.parse_args() args = parser.parse_args()
if args.src == args.dst: if args.src == args.dst:
parser.error("--src and --dst can't be the same") parser.error("--src and --dst can't be the same")
@@ -33,120 +35,28 @@ def parse_args():
parser.error("--lr must be a valid float eg. 0.001 or 1e-3") parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
return args return args
vae = None
def sample_decode(latent, filename, version):
global vae
if not vae:
vae = get_vae(version, fp16=True)
vae.to("cuda")
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)
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)
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
if __name__ == "__main__": if __name__ == "__main__":
args = parse_args() args = parse_args()
target_dev = "cuda"
resolution = 768
dataset = LatentDataset(resolution, args.src, args.dst) dataset = LatentDataset([args.src, args.dst])
loader = DataLoader( loader = DataLoader(
dataset, dataset,
batch_size=args.bs, batch_size = args.batch,
shuffle=True, shuffle = True,
num_workers=0, drop_last = True,
# num_workers=4, pin_memory = False,
# persistent_workers=True, # num_workers = 0,
num_workers = 4,
persistent_workers=True,
) )
eval_src, eval_dst = get_eval_data( model = Model() # TODO: handle scale factor/channels for non-sd VAEs
dataset,
f"latents/test_{args.src}_{resolution}px.npy",
f"latents/test_{args.dst}_{resolution}px.npy",
target_dev,
)
os.makedirs("models", exist_ok=True)
log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
model = Interposer()
criterion = torch.nn.L1Loss() criterion = torch.nn.L1Loss()
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr)) 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 = None scheduler = None
if args.cosine: if args.cosine:
print("Using CosineAnnealingLR") print("Using CosineAnnealingLR")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max = int(args.steps/args.bs), optimizer, T_max = int(args.steps/args.batch),
) )
else: else:
print("Using LinearLR") print("Using LinearLR")
@@ -154,23 +64,35 @@ if __name__ == "__main__":
optimizer, optimizer,
start_factor = 0.1, start_factor = 0.1,
end_factor = 1.0, end_factor = 1.0,
total_iters = int(5000/args.bs), total_iters = int(5000/args.batch),
) )
if args.resume: if args.resume:
model.load_state_dict(load_file(args.resume)) model.load_state_dict(load_file(args.resume))
model.to(target_dev) model.to(TARGET_DEV)
optimizer.load_state_dict(torch.load( optimizer.load_state_dict(torch.load(
f"{os.path.splitext(args.resume)[0]}.optim.pth" f"{os.path.splitext(args.resume)[0]}.optim.pth"
)) ))
optimizer.param_groups[0]['lr'] = scheduler.base_lrs[0]
else: else:
model.to(target_dev) model.to(TARGET_DEV)
wrapper = ModelWrapper( # model wrapper for saving/eval/etc
name = f"{args.src}-to-{args.dst}_interposer-{args.rev}",
specs = [args.src, args.dst],
model = model,
evals = dataset.get_eval(),
device = TARGET_DEV,
criterion = criterion,
optimizer = optimizer,
scheduler = scheduler,
)
progress = tqdm(total=args.steps) progress = tqdm(total=args.steps)
while progress.n < args.steps: while progress.n < args.steps:
for src, dst in loader: for src, dst in loader:
src = src.to(target_dev) src = src.to(TARGET_DEV)
dst = dst.to(target_dev) dst = dst.to(TARGET_DEV)
with torch.cuda.amp.autocast(): with torch.cuda.amp.autocast():
y_pred = model(src) # forward y_pred = model(src) # forward
loss = criterion(y_pred, dst) # loss loss = criterion(y_pred, dst) # loss
@@ -179,19 +101,15 @@ if __name__ == "__main__":
optimizer.zero_grad() optimizer.zero_grad()
loss.backward() loss.backward()
optimizer.step() optimizer.step()
scheduler.step() if progress.n >= args.lrskip: scheduler.step()
# eval/save # eval/save
progress.update(args.bs) progress.update(args.batch)
if progress.n % (1000 + 1000%args.bs) == 0: wrapper.log_step(loss.data.item(), progress.n)
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst) if args.nsave > 0 and progress.n % (args.nsave + args.nsave%args.batch) == 0:
if progress.n % (args.save + args.save%args.bs) == 0: wrapper.save_model(step=progress.n)
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
if progress.n >= args.steps: if progress.n >= args.steps:
break break
progress.close() progress.close()
wrapper.save_model(epoch="") # final save
# save final output wrapper.close()
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()
+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
-48
View File
@@ -1,48 +0,0 @@
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,
)
elif version == "v1":
vae = AutoencoderKL.from_pretrained(
"runwayml/stable-diffusion-v1-5",
subfolder="vae",
torch_dtype=dtype,
)
elif version == "v2" and file_path:
vae = AutoencoderKL.from_single_file(
file_path,
image_size=768,
)
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