diff --git a/README.md b/README.md index 8b0c23c..b3feef6 100644 --- a/README.md +++ b/README.md @@ -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. \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/depth_pro/LICENSE b/depth_pro/LICENSE new file mode 100644 index 0000000..02fa0ad --- /dev/null +++ b/depth_pro/LICENSE @@ -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. +------------------------------------------------------------------------------- diff --git a/depth_pro/depth_pro.py b/depth_pro/depth_pro.py index bff2a68..ef3fb95 100644 --- a/depth_pro/depth_pro.py +++ b/depth_pro/depth_pro.py @@ -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 diff --git a/depth_pro/example/pexels-rob-brennecke-1709780672-28506544.jpg b/depth_pro/example/pexels-rob-brennecke-1709780672-28506544.jpg new file mode 100644 index 0000000..3240e87 Binary files /dev/null and b/depth_pro/example/pexels-rob-brennecke-1709780672-28506544.jpg differ diff --git a/depth_pro/example/workflow.png b/depth_pro/example/workflow.png new file mode 100644 index 0000000..3096a45 Binary files /dev/null and b/depth_pro/example/workflow.png differ diff --git a/depth_pro/network/vit_factory.py b/depth_pro/network/vit_factory.py index 2cd899f..30eb003 100644 --- a/depth_pro/network/vit_factory.py +++ b/depth_pro/network/vit_factory.py @@ -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 diff --git a/depth_pro/utils.py b/depth_pro/utils.py index 0a401de..9302805 100644 --- a/depth_pro/utils.py +++ b/depth_pro/utils.py @@ -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 diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..4c2e3b3 --- /dev/null +++ b/nodes.py @@ -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", + } diff --git a/requirements.txt b/requirements.txt index 3c1c2e4..0f81b24 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ timm numpy<2 -pillow_heif -matplotlib \ No newline at end of file +# pillow_heif +# matplotlib \ No newline at end of file