functional
This commit is contained in:
@@ -1,2 +1,10 @@
|
|||||||
# Comfyui-depth-pro
|
# ComfyUI-Depth-Pro
|
||||||
https://github.com/apple/ml-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
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
from safetensors import safe_open
|
||||||
from torchvision.transforms import (
|
from torchvision.transforms import (
|
||||||
Compose,
|
Compose,
|
||||||
ConvertImageDtype,
|
ConvertImageDtype,
|
||||||
@@ -22,6 +23,7 @@ from .network.encoder import DepthProEncoder
|
|||||||
from .network.fov import FOVNetwork
|
from .network.fov import FOVNetwork
|
||||||
from .network.vit_factory import VIT_CONFIG_DICT, ViTPreset, create_vit
|
from .network.vit_factory import VIT_CONFIG_DICT, ViTPreset, create_vit
|
||||||
|
|
||||||
|
import comfy.utils
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DepthProConfig:
|
class DepthProConfig:
|
||||||
@@ -70,9 +72,10 @@ def create_backbone_model(
|
|||||||
|
|
||||||
|
|
||||||
def create_model_and_transforms(
|
def create_model_and_transforms(
|
||||||
|
model_path,
|
||||||
config: DepthProConfig = DEFAULT_MONODEPTH_CONFIG_DICT,
|
config: DepthProConfig = DEFAULT_MONODEPTH_CONFIG_DICT,
|
||||||
device: torch.device = torch.device("cpu"),
|
device: torch.device = torch.device("cpu"),
|
||||||
precision: torch.dtype = torch.float32,
|
precision: torch.dtype = torch.float16,
|
||||||
) -> Tuple[DepthPro, Compose]:
|
) -> Tuple[DepthPro, Compose]:
|
||||||
"""Create a DepthPro model and load weights from `config.checkpoint_uri`.
|
"""Create a DepthPro model and load weights from `config.checkpoint_uri`.
|
||||||
|
|
||||||
@@ -117,10 +120,7 @@ def create_model_and_transforms(
|
|||||||
last_dims=(32, 1),
|
last_dims=(32, 1),
|
||||||
use_fov_head=config.use_fov_head,
|
use_fov_head=config.use_fov_head,
|
||||||
fov_encoder=fov_encoder,
|
fov_encoder=fov_encoder,
|
||||||
).to(device)
|
).to(device, dtype=precision)
|
||||||
|
|
||||||
if precision == torch.half:
|
|
||||||
model.half()
|
|
||||||
|
|
||||||
transform = Compose(
|
transform = Compose(
|
||||||
[
|
[
|
||||||
@@ -131,22 +131,21 @@ def create_model_and_transforms(
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
if config.checkpoint_uri is not None:
|
state_dict = comfy.utils.load_torch_file(model_path)
|
||||||
state_dict = torch.load(config.checkpoint_uri, map_location="cpu")
|
missing_keys, unexpected_keys = model.load_state_dict(
|
||||||
missing_keys, unexpected_keys = model.load_state_dict(
|
state_dict=state_dict, strict=True
|
||||||
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:
|
# fc_norm is only for the classification head,
|
||||||
raise KeyError(
|
# which we would not use. We only use the encoding.
|
||||||
f"Found unexpected keys when loading monodepth: {unexpected_keys}"
|
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
|
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:
|
if len(missing_keys) != 0:
|
||||||
raise KeyError(f"Keys are missing when loading vit: {missing_keys}")
|
raise KeyError(f"Keys are missing when loading vit: {missing_keys}")
|
||||||
|
|
||||||
LOGGER.info(model)
|
# LOGGER.info(model)
|
||||||
return model.model
|
return model.model
|
||||||
|
|||||||
+2
-15
@@ -5,12 +5,7 @@ from pathlib import Path
|
|||||||
from typing import Any, Dict, List, Tuple, Union
|
from typing import Any, Dict, List, Tuple, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pillow_heif
|
|
||||||
from PIL import ExifTags, Image, TiffTags
|
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]:
|
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.
|
f_px: The optional focal length in pixels, extracting from the exif data.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
LOGGER.debug(f"Loading image {path} ...")
|
|
||||||
|
|
||||||
path = Path(path)
|
path = Path(path)
|
||||||
if path.suffix.lower() in [".heic"]:
|
img_pil = Image.open(path)
|
||||||
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_exif = extract_exif(img_pil)
|
img_exif = extract_exif(img_pil)
|
||||||
icc_profile = img_pil.info.get("icc_profile", None)
|
icc_profile = img_pil.info.get("icc_profile", None)
|
||||||
@@ -84,7 +74,7 @@ def load_rgb(
|
|||||||
elif exif_orientation == 8:
|
elif exif_orientation == 8:
|
||||||
img_pil = img_pil.transpose(Image.ROTATE_90)
|
img_pil = img_pil.transpose(Image.ROTATE_90)
|
||||||
elif exif_orientation != 1:
|
elif exif_orientation != 1:
|
||||||
LOGGER.warning(f"Ignoring image orientation {exif_orientation}.")
|
pass
|
||||||
|
|
||||||
img = np.array(img_pil)
|
img = np.array(img_pil)
|
||||||
# Convert to RGB if single channel.
|
# Convert to RGB if single channel.
|
||||||
@@ -94,8 +84,6 @@ def load_rgb(
|
|||||||
if remove_alpha:
|
if remove_alpha:
|
||||||
img = img[:, :, :3]
|
img = img[:, :, :3]
|
||||||
|
|
||||||
LOGGER.debug(f"\tHxW: {img.shape[0]}x{img.shape[1]}")
|
|
||||||
|
|
||||||
# Extract the focal length from exif data.
|
# Extract the focal length from exif data.
|
||||||
f_35mm = img_exif.get(
|
f_35mm = img_exif.get(
|
||||||
"FocalLengthIn35mmFilm",
|
"FocalLengthIn35mmFilm",
|
||||||
@@ -104,7 +92,6 @@ def load_rgb(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
if f_35mm is not None and f_35mm > 0:
|
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)
|
f_px = fpx_from_f35(img.shape[1], img.shape[0], f_35mm)
|
||||||
else:
|
else:
|
||||||
f_px = None
|
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
|
timm
|
||||||
numpy<2
|
numpy<2
|
||||||
pillow_heif
|
# pillow_heif
|
||||||
matplotlib
|
# matplotlib
|
||||||
Reference in New Issue
Block a user