From bb6a84ea167b4f01bdcc0ebf620f8d983ec72fed Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Mar 2024 12:49:53 +0200 Subject: [PATCH] Use comfy VAE --- depthfm/dfm.py | 20 ++++++++------- inference.py | 65 ------------------------------------------------ nodes.py | 12 ++++----- requirements.txt | 3 --- 4 files changed, 17 insertions(+), 83 deletions(-) delete mode 100644 inference.py diff --git a/depthfm/dfm.py b/depthfm/dfm.py index 4814650..fa83a7a 100644 --- a/depthfm/dfm.py +++ b/depthfm/dfm.py @@ -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): diff --git a/inference.py b/inference.py deleted file mode 100644 index 82afec3..0000000 --- a/inference.py +++ /dev/null @@ -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) diff --git a/nodes.py b/nodes.py index 3d40414..42c3067 100644 --- a/nodes.py +++ b/nodes.py @@ -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 = { diff --git a/requirements.txt b/requirements.txt index 93b0607..1bbcc1f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,4 @@ numpy einops omegaconf -matplotlib -accelerate>=0.22.0 -diffusers>=0.20.1 torchdiffeq>=0.2.3 \ No newline at end of file