functional

This commit is contained in:
spacepxl
2024-10-04 01:02:30 -04:00
parent 36c3be8b80
commit b312b2aabe
10 changed files with 233 additions and 39 deletions
+10 -2
View File
@@ -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.
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+47
View File
@@ -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
View File
@@ -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

+1 -1
View File
@@ -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
View File
@@ -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
+150
View File
@@ -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
View File
@@ -1,4 +1,4 @@
timm
numpy<2
pillow_heif
matplotlib
# pillow_heif
# matplotlib