Use comfy VAE
This commit is contained in:
+11
-9
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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,7 +44,7 @@ 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)
|
||||
|
||||
@@ -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
|
||||
@@ -96,6 +95,7 @@ class Depth_fm:
|
||||
if invert:
|
||||
depth = 1.0 - depth
|
||||
|
||||
depth = torch.clamp(depth, 0.0, 1.0)
|
||||
return (depth,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
@@ -1,7 +1,4 @@
|
||||
numpy
|
||||
einops
|
||||
omegaconf
|
||||
matplotlib
|
||||
accelerate>=0.22.0
|
||||
diffusers>=0.20.1
|
||||
torchdiffeq>=0.2.3
|
||||
Reference in New Issue
Block a user