First working version
This commit is contained in:
+13
@@ -1,3 +1,16 @@
|
||||
raw/
|
||||
images/
|
||||
latent_*/
|
||||
vae/
|
||||
models/
|
||||
test.py
|
||||
test.png
|
||||
*.npy
|
||||
*.ckpt
|
||||
*.safetensors
|
||||
|
||||
# default github .gitignore follows
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
|
||||
class Interposer(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
# it looks like a spaceship if you squint :P
|
||||
module_list = [
|
||||
#############)
|
||||
#############)
|
||||
#||#
|
||||
#||#
|
||||
nn.Conv2d(4, 32, kernel_size=5, padding=2),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(32, 128, kernel_size=7, padding=3),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(128, 32, kernel_size=7, padding=3),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(32, 4, kernel_size=5, padding=2),
|
||||
#||#
|
||||
#||#
|
||||
#############)
|
||||
#############)
|
||||
]
|
||||
|
||||
self.sequential = nn.Sequential(*module_list)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.sequential(x)
|
||||
@@ -0,0 +1,21 @@
|
||||
import os
|
||||
import hashlib
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
|
||||
# target resolution [latent res * 8]
|
||||
resolution = 768
|
||||
|
||||
if not os.path.isdir("images"):
|
||||
os.mkdir("images")
|
||||
|
||||
for i in tqdm(os.listdir("raw")):
|
||||
src = os.path.join("raw", i)
|
||||
md5 = hashlib.md5(open(src,'rb').read()).hexdigest()
|
||||
out = os.path.join("images", f"{md5}.png")
|
||||
if os.path.isfile(out):
|
||||
continue
|
||||
img = Image.open(src)
|
||||
img = img.convert('RGB')
|
||||
img = img.resize((resolution,resolution), Image.LANCZOS)
|
||||
img.save(out)
|
||||
@@ -0,0 +1,51 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from torchvision import transforms
|
||||
from diffusers import AutoencoderKL
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
|
||||
from vae import get_vae
|
||||
|
||||
def encode(vae, img):
|
||||
"""image [PIL Image] -> latent [np array]"""
|
||||
inp = transforms.ToTensor()(img).unsqueeze(0)
|
||||
inp = inp.to("cuda") # move to GPU
|
||||
latent = vae.encode(inp*2.0-1.0)
|
||||
latent = latent.latent_dist.sample()
|
||||
return latent.cpu().detach()
|
||||
|
||||
def process_folder(vae, v):
|
||||
if not os.path.isdir(f"latent_{v}"):
|
||||
os.mkdir(f"latent_{v}")
|
||||
|
||||
vae.to("cuda")
|
||||
for i in tqdm(os.listdir("images")):
|
||||
src = os.path.join("images", i)
|
||||
img = Image.open(src)
|
||||
dst = os.path.join(f"latent_{v}", f"{os.path.splitext(i)[0]}.npy")
|
||||
latent = encode(vae, img)
|
||||
np.save(dst, latent)
|
||||
vae.to("cpu")
|
||||
|
||||
def run_v1(file_path=None):
|
||||
vae = get_vae("v1", file_path)
|
||||
process_folder(vae, "v1")
|
||||
del vae
|
||||
|
||||
def run_v2(file_path=None):
|
||||
vae = get_vae("v2", file_path)
|
||||
process_folder(vae, "v2")
|
||||
del vae
|
||||
|
||||
def run_xl(file_path=None):
|
||||
vae = get_vae("xl", file_path)
|
||||
process_folder(vae, "xl")
|
||||
del vae
|
||||
|
||||
if __name__ == "__main__":
|
||||
# run_v1("./vae/ft-mse-840000.ckpt") # probably doesn't reflect internal SD latent
|
||||
run_v1()
|
||||
run_v2()
|
||||
run_xl("./vae/sdxl_v0.9.safetensors") # 1.0 has artifacts
|
||||
@@ -0,0 +1,95 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
import random
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from interposer import Interposer
|
||||
from vae import get_vae
|
||||
|
||||
# options
|
||||
target_dev = "cuda"
|
||||
target_steps = 500000
|
||||
save_every_n = 50000
|
||||
latent_src = "v1"
|
||||
latent_dst = "xl"
|
||||
|
||||
class Latent:
|
||||
def __init__(self, md5, lat_src, lat_dst, dev):
|
||||
if lat_src == "v1": src = os.path.join("latent_v1", f"{md5}.npy")
|
||||
if lat_src == "v2": src = os.path.join("latent_v2", f"{md5}.npy")
|
||||
if lat_src == "xl": src = os.path.join("latent_xl", f"{md5}.npy")
|
||||
|
||||
if lat_dst == "v1": dst = os.path.join("latent_v1", f"{md5}.npy")
|
||||
if lat_dst == "v2": dst = os.path.join("latent_v2", f"{md5}.npy")
|
||||
if lat_dst == "xl": dst = os.path.join("latent_xl", f"{md5}.npy")
|
||||
|
||||
self.src = torch.from_numpy(np.load(src)).to(dev)
|
||||
self.dst = torch.from_numpy(np.load(dst)).to(dev)
|
||||
|
||||
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)
|
||||
|
||||
if not os.path.isdir("models"): os.mkdir("models")
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Loading latents from disk")
|
||||
latents = []
|
||||
for i in tqdm(os.listdir("images")):
|
||||
md5 = os.path.splitext(i)[0]
|
||||
latents.append(Latent(md5, latent_src, latent_dst, target_dev))
|
||||
|
||||
if os.path.isfile(f"test_{latent_src}.npy"):
|
||||
sample_latent = torch.from_numpy(np.load(f"test_{latent_src}.npy")).to(target_dev)
|
||||
else:
|
||||
sample_latent = random.choice(latents).src
|
||||
|
||||
model = Interposer()
|
||||
model.to(target_dev)
|
||||
|
||||
criterion = torch.nn.MSELoss(size_average=False)
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=1e-8)
|
||||
|
||||
for t in tqdm(range(target_steps)):
|
||||
# io = latents[t%len(latents)]
|
||||
io = random.choice(latents)
|
||||
|
||||
y_pred = model(io.src) # forward
|
||||
loss = criterion(y_pred, io.dst) # loss
|
||||
|
||||
# backward
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
# print loss
|
||||
if t%1000 == 0: tqdm.write(f"{t} - {loss.data.item():.2f}")
|
||||
|
||||
# sample/save
|
||||
if t%save_every_n == 0:
|
||||
out = model(sample_latent)
|
||||
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{t/1000}k"
|
||||
sample_decode(out, f"{output_name}.png", latent_dst)
|
||||
save_file(model.state_dict(), f"{output_name}.safetensors")
|
||||
# save final output
|
||||
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{target_steps/1000}k"
|
||||
sample_decode(out, f"{output_name}.png", "v1")
|
||||
save_file(model.state_dict(), f"{output_name}.safetensors")
|
||||
@@ -0,0 +1,48 @@
|
||||
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
|
||||
Reference in New Issue
Block a user