functional
This commit is contained in:
@@ -1,2 +1,10 @@
|
||||
# Comfyui-depth-pro
|
||||
https://github.com/apple/ml-depth-pro
|
||||
# ComfyUI-Depth-Pro
|
||||
|
||||
Based on https://github.com/apple/ml-depth-pro
|
||||
|
||||
## License
|
||||
|
||||
All code that is unique to this repository is covered by the Apache-2.0 license. Any
|
||||
code and models that are redistributed without modification from the original codebase
|
||||
may be subject to the original license from https://github.com/apple/ml-depth-pro/blob/main/LICENSE
|
||||
if applicable. This project is not affiliated in any way with Apple Inc.
|
||||
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -0,0 +1,47 @@
|
||||
Copyright (C) 2024 Apple Inc. All Rights Reserved.
|
||||
|
||||
Disclaimer: IMPORTANT: This Apple software is supplied to you by Apple
|
||||
Inc. ("Apple") in consideration of your agreement to the following
|
||||
terms, and your use, installation, modification or redistribution of
|
||||
this Apple software constitutes acceptance of these terms. If you do
|
||||
not agree with these terms, please do not use, install, modify or
|
||||
redistribute this Apple software.
|
||||
|
||||
In consideration of your agreement to abide by the following terms, and
|
||||
subject to these terms, Apple grants you a personal, non-exclusive
|
||||
license, under Apple's copyrights in this original Apple software (the
|
||||
"Apple Software"), to use, reproduce, modify and redistribute the Apple
|
||||
Software, with or without modifications, in source and/or binary forms;
|
||||
provided that if you redistribute the Apple Software in its entirety and
|
||||
without modifications, you must retain this notice and the following
|
||||
text and disclaimers in all such redistributions of the Apple Software.
|
||||
Neither the name, trademarks, service marks or logos of Apple Inc. may
|
||||
be used to endorse or promote products derived from the Apple Software
|
||||
without specific prior written permission from Apple. Except as
|
||||
expressly stated in this notice, no other rights or licenses, express or
|
||||
implied, are granted by Apple herein, including but not limited to any
|
||||
patent rights that may be infringed by your derivative works or by other
|
||||
works in which the Apple Software may be incorporated.
|
||||
|
||||
The Apple Software is provided by Apple on an "AS IS" basis. APPLE
|
||||
MAKES NO WARRANTIES, EXPRESS OR IMPLIED, INCLUDING WITHOUT LIMITATION
|
||||
THE IMPLIED WARRANTIES OF NON-INFRINGEMENT, MERCHANTABILITY AND FITNESS
|
||||
FOR A PARTICULAR PURPOSE, REGARDING THE APPLE SOFTWARE OR ITS USE AND
|
||||
OPERATION ALONE OR IN COMBINATION WITH YOUR PRODUCTS.
|
||||
|
||||
IN NO EVENT SHALL APPLE BE LIABLE FOR ANY SPECIAL, INDIRECT, INCIDENTAL
|
||||
OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
|
||||
SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
|
||||
INTERRUPTION) ARISING IN ANY WAY OUT OF THE USE, REPRODUCTION,
|
||||
MODIFICATION AND/OR DISTRIBUTION OF THE APPLE SOFTWARE, HOWEVER CAUSED
|
||||
AND WHETHER UNDER THEORY OF CONTRACT, TORT (INCLUDING NEGLIGENCE),
|
||||
STRICT LIABILITY OR OTHERWISE, EVEN IF APPLE HAS BEEN ADVISED OF THE
|
||||
POSSIBILITY OF SUCH DAMAGE.
|
||||
|
||||
|
||||
-------------------------------------------------------------------------------
|
||||
SOFTWARE DISTRIBUTED IN THIS REPOSITORY:
|
||||
|
||||
This software includes a number of subcomponents with separate
|
||||
copyright notices and license terms - please see the file ACKNOWLEDGEMENTS.
|
||||
-------------------------------------------------------------------------------
|
||||
+18
-19
@@ -9,6 +9,7 @@ from typing import Mapping, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from safetensors import safe_open
|
||||
from torchvision.transforms import (
|
||||
Compose,
|
||||
ConvertImageDtype,
|
||||
@@ -22,6 +23,7 @@ from .network.encoder import DepthProEncoder
|
||||
from .network.fov import FOVNetwork
|
||||
from .network.vit_factory import VIT_CONFIG_DICT, ViTPreset, create_vit
|
||||
|
||||
import comfy.utils
|
||||
|
||||
@dataclass
|
||||
class DepthProConfig:
|
||||
@@ -70,9 +72,10 @@ def create_backbone_model(
|
||||
|
||||
|
||||
def create_model_and_transforms(
|
||||
model_path,
|
||||
config: DepthProConfig = DEFAULT_MONODEPTH_CONFIG_DICT,
|
||||
device: torch.device = torch.device("cpu"),
|
||||
precision: torch.dtype = torch.float32,
|
||||
precision: torch.dtype = torch.float16,
|
||||
) -> Tuple[DepthPro, Compose]:
|
||||
"""Create a DepthPro model and load weights from `config.checkpoint_uri`.
|
||||
|
||||
@@ -117,10 +120,7 @@ def create_model_and_transforms(
|
||||
last_dims=(32, 1),
|
||||
use_fov_head=config.use_fov_head,
|
||||
fov_encoder=fov_encoder,
|
||||
).to(device)
|
||||
|
||||
if precision == torch.half:
|
||||
model.half()
|
||||
).to(device, dtype=precision)
|
||||
|
||||
transform = Compose(
|
||||
[
|
||||
@@ -131,22 +131,21 @@ def create_model_and_transforms(
|
||||
]
|
||||
)
|
||||
|
||||
if config.checkpoint_uri is not None:
|
||||
state_dict = torch.load(config.checkpoint_uri, map_location="cpu")
|
||||
missing_keys, unexpected_keys = model.load_state_dict(
|
||||
state_dict=state_dict, strict=True
|
||||
state_dict = comfy.utils.load_torch_file(model_path)
|
||||
missing_keys, unexpected_keys = model.load_state_dict(
|
||||
state_dict=state_dict, strict=True
|
||||
)
|
||||
|
||||
if len(unexpected_keys) != 0:
|
||||
raise KeyError(
|
||||
f"Found unexpected keys when loading monodepth: {unexpected_keys}"
|
||||
)
|
||||
|
||||
if len(unexpected_keys) != 0:
|
||||
raise KeyError(
|
||||
f"Found unexpected keys when loading monodepth: {unexpected_keys}"
|
||||
)
|
||||
|
||||
# fc_norm is only for the classification head,
|
||||
# which we would not use. We only use the encoding.
|
||||
missing_keys = [key for key in missing_keys if "fc_norm" not in key]
|
||||
if len(missing_keys) != 0:
|
||||
raise KeyError(f"Keys are missing when loading monodepth: {missing_keys}")
|
||||
# fc_norm is only for the classification head,
|
||||
# which we would not use. We only use the encoding.
|
||||
missing_keys = [key for key in missing_keys if "fc_norm" not in key]
|
||||
if len(missing_keys) != 0:
|
||||
raise KeyError(f"Keys are missing when loading monodepth: {missing_keys}")
|
||||
|
||||
return model, transform
|
||||
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.9 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 1.2 MiB |
@@ -120,5 +120,5 @@ def create_vit(
|
||||
if len(missing_keys) != 0:
|
||||
raise KeyError(f"Keys are missing when loading vit: {missing_keys}")
|
||||
|
||||
LOGGER.info(model)
|
||||
# LOGGER.info(model)
|
||||
return model.model
|
||||
|
||||
+2
-15
@@ -5,12 +5,7 @@ from pathlib import Path
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import pillow_heif
|
||||
from PIL import ExifTags, Image, TiffTags
|
||||
from pillow_heif import register_heif_opener
|
||||
|
||||
register_heif_opener()
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def extract_exif(img_pil: Image) -> Dict[str, Any]:
|
||||
@@ -62,14 +57,9 @@ def load_rgb(
|
||||
f_px: The optional focal length in pixels, extracting from the exif data.
|
||||
|
||||
"""
|
||||
LOGGER.debug(f"Loading image {path} ...")
|
||||
|
||||
path = Path(path)
|
||||
if path.suffix.lower() in [".heic"]:
|
||||
heif_file = pillow_heif.open_heif(path, convert_hdr_to_8bit=True)
|
||||
img_pil = heif_file.to_pillow()
|
||||
else:
|
||||
img_pil = Image.open(path)
|
||||
img_pil = Image.open(path)
|
||||
|
||||
img_exif = extract_exif(img_pil)
|
||||
icc_profile = img_pil.info.get("icc_profile", None)
|
||||
@@ -84,7 +74,7 @@ def load_rgb(
|
||||
elif exif_orientation == 8:
|
||||
img_pil = img_pil.transpose(Image.ROTATE_90)
|
||||
elif exif_orientation != 1:
|
||||
LOGGER.warning(f"Ignoring image orientation {exif_orientation}.")
|
||||
pass
|
||||
|
||||
img = np.array(img_pil)
|
||||
# Convert to RGB if single channel.
|
||||
@@ -94,8 +84,6 @@ def load_rgb(
|
||||
if remove_alpha:
|
||||
img = img[:, :, :3]
|
||||
|
||||
LOGGER.debug(f"\tHxW: {img.shape[0]}x{img.shape[1]}")
|
||||
|
||||
# Extract the focal length from exif data.
|
||||
f_35mm = img_exif.get(
|
||||
"FocalLengthIn35mmFilm",
|
||||
@@ -104,7 +92,6 @@ def load_rgb(
|
||||
),
|
||||
)
|
||||
if f_35mm is not None and f_35mm > 0:
|
||||
LOGGER.debug(f"\tfocal length @ 35mm film: {f_35mm}mm")
|
||||
f_px = fpx_from_f35(img.shape[1], img.shape[0], f_35mm)
|
||||
else:
|
||||
f_px = None
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
import os
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import trange
|
||||
from torchvision.transforms import Normalize
|
||||
|
||||
import comfy.utils
|
||||
import model_management
|
||||
import folder_paths
|
||||
|
||||
from . import depth_pro
|
||||
|
||||
|
||||
class LoadDepthPro:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"precision": (["fp16", "fp32"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("DEPTH_PRO_MODEL",)
|
||||
RETURN_NAMES = ("depth_pro_model",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "Depth-Pro"
|
||||
|
||||
def load_model(self, precision):
|
||||
device = model_management.get_torch_device()
|
||||
dtype = torch.float16 if precision == "fp16" else torch.float32
|
||||
|
||||
depth_model_path = os.path.join(folder_paths.models_dir, "depth", "ml-depth-pro")
|
||||
if not os.path.exists(depth_model_path):
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(
|
||||
repo_id="spacepxl/ml-depth-pro",
|
||||
local_dir=depth_model_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
|
||||
depth_model_path = os.path.join(depth_model_path, "depth_pro.fp16.safetensors")
|
||||
model, transform = depth_pro.create_model_and_transforms(depth_model_path, device=device, precision=dtype)
|
||||
model.eval()
|
||||
|
||||
model_dict = {
|
||||
"model": model,
|
||||
"device": device,
|
||||
"dtype": dtype,
|
||||
}
|
||||
|
||||
return (model_dict,)
|
||||
|
||||
|
||||
class DepthPro:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"depth_pro_model": ("DEPTH_PRO_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "LIST", "FLOAT",)
|
||||
RETURN_NAMES = ("metric_depth", "focal_list", "focal_avg",)
|
||||
FUNCTION = "estimate_depth"
|
||||
CATEGORY = "Depth-Pro"
|
||||
|
||||
def estimate_depth(self, depth_pro_model, image):
|
||||
model = depth_pro_model["model"]
|
||||
device = depth_pro_model["device"]
|
||||
dtype = depth_pro_model["dtype"]
|
||||
|
||||
transform = Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
|
||||
|
||||
rgb = image.unsqueeze(0) if len(image.shape) < 4 else image
|
||||
rgb = rgb.movedim(-1, 1) # BCHW
|
||||
|
||||
depth = []
|
||||
focal_px = []
|
||||
|
||||
# add comfyui progress bar
|
||||
for i in trange(rgb.size(0)):
|
||||
rgb_image = rgb[i, :3].unsqueeze(0).to(device, dtype=dtype)
|
||||
rgb_image = transform(rgb_image)
|
||||
|
||||
prediction = model.infer(rgb_image)
|
||||
depth.append(prediction["depth"].unsqueeze(-1))
|
||||
focal_px.append(prediction["focallength_px"].item())
|
||||
|
||||
depth = torch.stack(depth, dim=0).repeat(1,1,1,3)
|
||||
focal_list = focal_px
|
||||
focal_avg = np.mean(focal_px)
|
||||
|
||||
return (depth, focal_list, focal_avg)
|
||||
|
||||
|
||||
class MetricDepthToRelative:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"depth": ("IMAGE",),
|
||||
"per_image": ("BOOLEAN", {"default": False,}),
|
||||
"invert": ("BOOLEAN", {"default": True,}),
|
||||
"std_dev": ("FLOAT", {"default": 5.0, "min": 0.1, "max": 100.0, "step": 0.1, "round": 0.1}),
|
||||
"gamma": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("depth",)
|
||||
FUNCTION = "convert_depth"
|
||||
CATEGORY = "Depth-Pro"
|
||||
|
||||
def convert_depth(self, depth, per_image, invert, std_dev, gamma):
|
||||
relative_depth = depth.detach().clone()
|
||||
|
||||
if per_image:
|
||||
for i in range(relative_depth.size(0)):
|
||||
std, mean = torch.std_mean(relative_depth[i], dim=None)
|
||||
relative_depth[i] = torch.clamp(relative_depth[i], min = 0, max = std * std_dev + mean)
|
||||
relative_depth[i] = relative_depth[i] - relative_depth[i].min()
|
||||
relative_depth[i] = relative_depth[i] / relative_depth[i].max()
|
||||
else:
|
||||
std, mean = torch.std_mean(relative_depth, dim=None)
|
||||
relative_depth = torch.clamp(relative_depth, min = 0, max = std * std_dev + mean)
|
||||
relative_depth = relative_depth - relative_depth.min()
|
||||
relative_depth = relative_depth / relative_depth.max()
|
||||
|
||||
if invert:
|
||||
relative_depth = 1 - relative_depth
|
||||
|
||||
if gamma != 1:
|
||||
relative_depth = relative_depth ** (1 / gamma)
|
||||
|
||||
return (relative_depth,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LoadDepthPro": LoadDepthPro,
|
||||
"DepthPro": DepthPro,
|
||||
"MetricDepthToRelative": MetricDepthToRelative,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LoadDepthPro": "(Down)Load Depth Pro model",
|
||||
"DepthPro": "Depth Pro",
|
||||
"MetricDepthToRelative": "Metric Depth to Relative",
|
||||
}
|
||||
+2
-2
@@ -1,4 +1,4 @@
|
||||
timm
|
||||
numpy<2
|
||||
pillow_heif
|
||||
matplotlib
|
||||
# pillow_heif
|
||||
# matplotlib
|
||||
Reference in New Issue
Block a user