Use comfy VAE

This commit is contained in:
kijai
2024-03-22 12:49:53 +02:00
parent 455a5bd044
commit bb6a84ea16
4 changed files with 17 additions and 83 deletions
+11 -9
View File
@@ -15,10 +15,11 @@ def exists(val):
class DepthFM(nn.Module):
def __init__(self, ckpt_path: str):
def __init__(self, vae, ckpt_path: str):
super().__init__()
vae_id = "runwayml/stable-diffusion-v1-5"
self.vae = AutoencoderKL.from_pretrained(vae_id, subfolder="vae")
#vae_id = "runwayml/stable-diffusion-v1-5"
#self.vae = AutoencoderKL.from_pretrained(vae_id, subfolder="vae")
self.vae = vae
self.scale_factor = 0.18215
# set with checkpoint
@@ -98,11 +99,12 @@ class DepthFM(nn.Module):
@torch.no_grad()
def encode(self, x: Tensor, sample_posterior: bool = True):
posterior = self.vae.encode(x)
if sample_posterior:
z = posterior.latent_dist.sample()
else:
z = posterior.latent_dist.mode()
self.vae.first_stage_model = self.vae.first_stage_model.to(x.device)
z = self.vae.first_stage_model.encode(x)
# if sample_posterior:
# z = posterior.latent_dist.sample()
# else:
# z = posterior.latent_dist.mode()
# normalize latent code
z = z * self.scale_factor
return z
@@ -110,7 +112,7 @@ class DepthFM(nn.Module):
@torch.no_grad()
def decode(self, z: Tensor):
z = 1.0 / self.scale_factor * z
return self.vae.decode(z).sample
return self.vae.first_stage_model.decode(z)
def sigmoid(x):
-65
View File
@@ -1,65 +0,0 @@
import os
import torch
import einops
import argparse
import numpy as np
from PIL import Image
from depthfm import DepthFM
import matplotlib.pyplot as plt
def load_im(fp):
assert os.path.exists(fp), f"File not found: {fp}"
im = Image.open(fp).convert('RGB')
x = np.array(im)
x = einops.rearrange(x, 'h w c -> c h w')
x = x / 127.5 - 1
x = torch.tensor(x, dtype=torch.float32)[None]
return x
def main(args):
print(f"{'Input':<10}: {args.img}")
print(f"{'Steps':<10}: {args.num_steps}")
print(f"{'Ensemble':<10}: {args.ensemble_size}")
# Load the model
model = DepthFM(args.ckpt)
model.cuda().eval()
# Load an image
im = load_im(args.img).cuda()
# Generate depth
depth = model.predict_depth(im, num_steps=args.num_steps, ensemble_size=args.ensemble_size)
depth = depth.squeeze(0).squeeze(0).cpu().numpy() # (h, w) in [0, 1]
# Convert depth to [0, 255] range
if args.no_color:
depth = (depth * 255).astype(np.uint8)
else:
depth = plt.get_cmap('magma')(depth, bytes=True)[..., :3]
# Save the depth map
depth_fp = args.img.replace('.png', '-depth.png')
Image.fromarray(depth).save(depth_fp)
print(f"==> Saved depth map to {depth_fp}")
if __name__ == "__main__":
parser = argparse.ArgumentParser("DepthFM Inference")
parser.add_argument("--img", type=str, default="assets/dog.png",
help="Path to the input image")
parser.add_argument("--ckpt", type=str, default="checkpoints/depthfm-v1.ckpt",
help="Path to the model checkpoint")
parser.add_argument("--num_steps", type=int, default=2,
help="Number of steps for ODE solver")
parser.add_argument("--ensemble_size", type=int, default=4,
help="Number of ensemble members")
parser.add_argument("--no_color", action="store_true",
help="If set, the depth map will be grayscale")
parser.add_argument("--device", type=int, default=0,
help="GPU to use")
args = parser.parse_args()
main(args)
+6 -6
View File
@@ -21,6 +21,7 @@ class Depth_fm:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("VAE",),
"depthfm_model": (folder_paths.get_filename_list("checkpoints"),),
"images": ("IMAGE",),
"steps": ("INT", {"default": 4, "min": 1, "max": 200, "step": 1}),
@@ -43,10 +44,10 @@ class Depth_fm:
FUNCTION = "process"
CATEGORY = "depth_fm"
def process(self, depthfm_model, images, ensemble_size, steps, dtype, invert, per_batch):
def process(self, depthfm_model, vae, images, ensemble_size, steps, dtype, invert, per_batch):
device = model_management.get_torch_device()
dtype = convert_dtype(dtype)
custom_config = {
"model_path": depthfm_model,
"dtype": dtype,
@@ -54,7 +55,7 @@ class Depth_fm:
if not hasattr(self, "model") or custom_config != self.current_config:
self.current_config = custom_config
DEPTHFM_MODEL_PATH = folder_paths.get_full_path("checkpoints", depthfm_model)
self.model = DepthFM(DEPTHFM_MODEL_PATH)
self.model = DepthFM(vae, DEPTHFM_MODEL_PATH)
self.model.eval().to(dtype).to(device)
images = images.permute(0, 3, 1, 2)
@@ -79,12 +80,10 @@ class Depth_fm:
for start_idx in range(0, images.shape[0], per_batch):
sub_images = self.model.predict_depth(images[start_idx:start_idx+per_batch], num_steps=steps, ensemble_size=ensemble_size)
depth_list.append(sub_images.cpu())
print(sub_images.shape[0])
batch_count = sub_images.shape[0]
pbar.update(batch_count)
depth = torch.cat(depth_list, dim=0)
#print(depth.min(), depth.max())
depth = depth.repeat(1, 3, 1, 1).permute(0, 2, 3, 1).cpu()
final_H = (orig_H // 2) * 2
@@ -95,7 +94,8 @@ class Depth_fm:
if invert:
depth = 1.0 - depth
depth = torch.clamp(depth, 0.0, 1.0)
return (depth,)
NODE_CLASS_MAPPINGS = {
-3
View File
@@ -1,7 +1,4 @@
numpy
einops
omegaconf
matplotlib
accelerate>=0.22.0
diffusers>=0.20.1
torchdiffeq>=0.2.3