diff --git a/.gitignore b/.gitignore index 68bc17f..0cda4e1 100644 --- a/.gitignore +++ b/.gitignore @@ -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] diff --git a/interposer.py b/interposer.py new file mode 100644 index 0000000..f35229d --- /dev/null +++ b/interposer.py @@ -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) diff --git a/preprocess_images.py b/preprocess_images.py new file mode 100644 index 0000000..1752625 --- /dev/null +++ b/preprocess_images.py @@ -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) diff --git a/preprocess_latents.py b/preprocess_latents.py new file mode 100644 index 0000000..43b32c7 --- /dev/null +++ b/preprocess_latents.py @@ -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 diff --git a/train.py b/train.py new file mode 100644 index 0000000..bdf5091 --- /dev/null +++ b/train.py @@ -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") diff --git a/vae.py b/vae.py new file mode 100644 index 0000000..0bd4f9e --- /dev/null +++ b/vae.py @@ -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