Merge pull request #6 from Limbicnation/fix/depth-estimation-structure
Fix/depth estimation structure
This commit is contained in:
+23
-4
@@ -1,11 +1,30 @@
|
||||
from .your_node_file import ComfyUIDepthEstimationNode
|
||||
"""
|
||||
ComfyUI Depth Estimation Node
|
||||
A custom node for depth map estimation using Depth-Anything-V2-Small model.
|
||||
"""
|
||||
|
||||
from .depth_estimation_node import DepthEstimationNode
|
||||
|
||||
# Version info
|
||||
__version__ = "1.0.0"
|
||||
|
||||
# Node class mapping for ComfyUI registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ComfyUIDepthEstimationNode": ComfyUIDepthEstimationNode
|
||||
"DepthEstimationNode": DepthEstimationNode
|
||||
}
|
||||
|
||||
# Display names for UI
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ComfyUIDepthEstimationNode": "Depth Estimation Node"
|
||||
"DepthEstimationNode": "Depth Estimation"
|
||||
}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
# Web extension info for ComfyUI
|
||||
WEB_DIRECTORY = "./js"
|
||||
|
||||
# Module exports
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"__version__",
|
||||
"WEB_DIRECTORY"
|
||||
]
|
||||
+148
-98
@@ -3,106 +3,156 @@ import numpy as np
|
||||
import torch
|
||||
from transformers import pipeline
|
||||
from PIL import Image, ImageFilter, ImageOps
|
||||
from comfy.nodes import Node, register_node
|
||||
import folder_paths
|
||||
from comfy.model_management import get_torch_device
|
||||
|
||||
def ensure_odd(value):
|
||||
"""Ensure the value is an odd integer."""
|
||||
value = int(value)
|
||||
return value if value % 2 == 1 else value + 1
|
||||
DEPTH_MODELS = {
|
||||
"Depth-Anything-Small": "LiheYoung/depth-anything-small",
|
||||
"Depth-Anything-Base": "LiheYoung/depth-anything-base",
|
||||
"Depth-Anything-Large": "LiheYoung/depth-anything-large",
|
||||
"Depth-Anything-V2-Small": "LiheYoung/depth-anything-small-hf",
|
||||
"Depth-Anything-V2-Base": "LiheYoung/depth-anything-base-hf",
|
||||
}
|
||||
|
||||
def convert_path(path):
|
||||
"""Convert path for compatibility between Windows and WSL."""
|
||||
if os.name == 'nt': # If running on Windows
|
||||
return path.replace('\\', '/')
|
||||
return path
|
||||
class DepthEstimationNode:
|
||||
MEDIAN_SIZES = ["3", "5", "7", "9", "11"]
|
||||
|
||||
def gamma_correction(img, gamma=1.0):
|
||||
"""Apply gamma correction to the image."""
|
||||
inv_gamma = 1.0 / gamma
|
||||
table = [((i / 255.0) ** inv_gamma) * 255 for i in range(256)]
|
||||
table = np.array(table, np.uint8)
|
||||
return Image.fromarray(np.array(img).astype(np.uint8)).point(lambda i: table[i])
|
||||
|
||||
def auto_gamma_correction(image):
|
||||
"""Automatically adjust gamma correction for the image."""
|
||||
image_array = np.array(image).astype(np.float32) / 255.0
|
||||
mean_luminance = np.mean(image_array)
|
||||
gamma = np.log(0.5) / np.log(mean_luminance)
|
||||
return gamma_correction(image, gamma=gamma)
|
||||
|
||||
def auto_contrast(image):
|
||||
"""Apply automatic contrast adjustment to the image."""
|
||||
return ImageOps.autocontrast(image)
|
||||
|
||||
class DepthEstimationNode(Node):
|
||||
def __init__(self, blur_radius=2.0, median_size=5, device="cpu"):
|
||||
super().__init__()
|
||||
self.blur_radius = blur_radius
|
||||
self.median_size = ensure_odd(median_size)
|
||||
self.device = 0 if device == "gpu" and torch.cuda.is_available() else -1
|
||||
self.pipe = pipeline(task="depth-estimation", model="LiheYoung/depth-anything-large-hf", device=self.device)
|
||||
|
||||
def process_image(self, image):
|
||||
if self.device == 0:
|
||||
image = image.convert("RGB") # Ensure image is in RGB format
|
||||
inputs = self.pipe.feature_extractor(images=image, return_tensors="pt").to(self.device)
|
||||
with torch.no_grad():
|
||||
outputs = self.pipe.model(**inputs)
|
||||
result = self.pipe.post_process(outputs, (image.height, image.width))
|
||||
else:
|
||||
result = self.pipe(image)
|
||||
|
||||
# Convert depth data to a NumPy array if not already one
|
||||
depth_data = np.array(result["depth"])
|
||||
|
||||
# Normalize and convert to uint8
|
||||
depth_normalized = (depth_data - depth_data.min()) / (depth_data.max() - depth_data.min() + 1e-8) # Avoid zero division
|
||||
depth_uint8 = (255 * depth_normalized).astype(np.uint8)
|
||||
|
||||
# Create an image from the processed depth data
|
||||
depth_image = Image.fromarray(depth_uint8)
|
||||
|
||||
# Apply a median filter to reduce noise
|
||||
depth_image = depth_image.filter(ImageFilter.MedianFilter(size=self.median_size))
|
||||
|
||||
# Enhanced edge detection with more feathering
|
||||
edges = depth_image.filter(ImageFilter.FIND_EDGES)
|
||||
edges = edges.filter(ImageFilter.GaussianBlur(radius=2 * self.blur_radius))
|
||||
edges = edges.point(lambda x: 255 if x > 20 else 0) # Adjusted threshold
|
||||
|
||||
# Create a mask from the edges
|
||||
mask = edges.convert("L")
|
||||
|
||||
# Blur only the edges using the mask
|
||||
blurred_edges = depth_image.filter(ImageFilter.GaussianBlur(radius=self.blur_radius * 2))
|
||||
|
||||
# Combine the blurred edges with the original depth image using the mask
|
||||
combined_image = Image.composite(blurred_edges, depth_image, mask)
|
||||
|
||||
# Apply auto gamma correction with a lower gamma to darken the image
|
||||
gamma_corrected_image = gamma_correction(combined_image, gamma=0.7)
|
||||
|
||||
# Apply auto contrast
|
||||
final_image = auto_contrast(gamma_corrected_image)
|
||||
|
||||
# Additional post-processing: Sharpen the final image
|
||||
final_image = final_image.filter(ImageFilter.SHARPEN)
|
||||
|
||||
return final_image
|
||||
|
||||
def forward(self, image_path: str):
|
||||
if not os.path.exists(image_path):
|
||||
raise FileNotFoundError(f"The input image path does not exist: {image_path}")
|
||||
|
||||
image = Image.open(image_path)
|
||||
return self.process_image(image)
|
||||
|
||||
@register_node
|
||||
class ComfyUIDepthEstimationNode(DepthEstimationNode):
|
||||
def __init__(self):
|
||||
super().__init__(blur_radius=2.0, median_size=5, device="cpu")
|
||||
self.device = get_torch_device()
|
||||
self.depth_estimator = None
|
||||
self.current_model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_name": (list(DEPTH_MODELS.keys()),),
|
||||
"blur_radius": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
||||
"median_size": (cls.MEDIAN_SIZES, {"default": "5"}),
|
||||
"apply_auto_contrast": ("BOOLEAN", {"default": True}),
|
||||
"apply_gamma": ("BOOLEAN", {"default": True})
|
||||
}
|
||||
}
|
||||
|
||||
def execute(self, image_path: str, output_path: str):
|
||||
final_image = self.forward(image_path)
|
||||
final_image.save(output_path)
|
||||
print(f"Processed and saved: {output_path}")
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "estimate_depth"
|
||||
CATEGORY = "image/depth"
|
||||
|
||||
def ensure_model_loaded(self, model_name):
|
||||
model_path = DEPTH_MODELS[model_name]
|
||||
if self.depth_estimator is None or self.current_model != model_path:
|
||||
try:
|
||||
self.depth_estimator = pipeline(
|
||||
"depth-estimation",
|
||||
model=model_path,
|
||||
device=self.device
|
||||
)
|
||||
self.current_model = model_path
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load model {model_name}: {str(e)}")
|
||||
|
||||
def estimate_depth(self, image, model_name, blur_radius=2.0, median_size="5",
|
||||
apply_auto_contrast=True, apply_gamma=True):
|
||||
try:
|
||||
# Validate median_size
|
||||
if median_size not in self.MEDIAN_SIZES:
|
||||
raise ValueError(f"Invalid median_size. Must be one of {self.MEDIAN_SIZES}")
|
||||
|
||||
median_size_int = int(median_size)
|
||||
self.ensure_model_loaded(model_name)
|
||||
|
||||
# Handle tensor conversion
|
||||
if torch.is_tensor(image):
|
||||
image_np = image.cpu().numpy()[0] # Remove batch dimension
|
||||
else:
|
||||
image_np = image
|
||||
|
||||
# Ensure proper RGB format and scaling
|
||||
if image_np.max() <= 1.0:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
else:
|
||||
image_np = image_np.astype(np.uint8)
|
||||
|
||||
# Convert to RGB if necessary
|
||||
if len(image_np.shape) == 3 and image_np.shape[-1] == 4:
|
||||
image_np = image_np[..., :3]
|
||||
elif len(image_np.shape) == 2:
|
||||
# Convert grayscale to RGB
|
||||
image_np = np.stack([image_np] * 3, axis=-1)
|
||||
|
||||
# Convert to PIL for processing
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Get depth map
|
||||
depth_result = self.depth_estimator(pil_image)
|
||||
depth_map = depth_result["predicted_depth"]
|
||||
|
||||
# Convert tensor to numpy and ensure correct dimensions
|
||||
if torch.is_tensor(depth_map):
|
||||
depth_map = depth_map.squeeze().cpu().numpy()
|
||||
|
||||
# Ensure depth_map is 2D
|
||||
depth_map = depth_map.squeeze()
|
||||
|
||||
# Normalize depth values to 0-255 range
|
||||
depth_min = depth_map.min()
|
||||
depth_max = depth_map.max()
|
||||
if depth_max > depth_min:
|
||||
depth_map = ((depth_map - depth_min) * (255.0 / (depth_max - depth_min)))
|
||||
else:
|
||||
depth_map = np.zeros_like(depth_map)
|
||||
|
||||
depth_map = depth_map.astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
depth_map = Image.fromarray(depth_map, mode='L') # Convert as grayscale
|
||||
|
||||
# Apply post-processing
|
||||
if blur_radius > 0:
|
||||
depth_map = depth_map.filter(ImageFilter.GaussianBlur(radius=blur_radius))
|
||||
|
||||
if median_size_int > 0:
|
||||
depth_map = depth_map.filter(ImageFilter.MedianFilter(size=median_size_int))
|
||||
|
||||
if apply_auto_contrast:
|
||||
depth_map = ImageOps.autocontrast(depth_map)
|
||||
|
||||
if apply_gamma:
|
||||
depth_array = np.array(depth_map).astype(np.float32) / 255.0
|
||||
mean_luminance = np.mean(depth_array)
|
||||
if mean_luminance > 0:
|
||||
gamma = np.log(0.5) / np.log(mean_luminance)
|
||||
depth_map = self.gamma_correction(depth_map, gamma)
|
||||
|
||||
# Convert back to tensor format
|
||||
depth_array = np.array(depth_map).astype(np.float32) / 255.0
|
||||
|
||||
# Convert single channel to 3 channels
|
||||
depth_array = np.stack([depth_array] * 3, axis=-1)
|
||||
|
||||
# Add batch dimension
|
||||
depth_tensor = torch.from_numpy(depth_array).unsqueeze(0)
|
||||
|
||||
# Move tensor to the correct device
|
||||
depth_tensor = depth_tensor.to(self.device)
|
||||
|
||||
return (depth_tensor,)
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Depth estimation failed: {str(e)}")
|
||||
|
||||
def gamma_correction(self, img, gamma=1.0):
|
||||
inv_gamma = 1.0 / gamma
|
||||
table = [((i / 255.0) ** inv_gamma) * 255 for i in range(256)]
|
||||
table = np.array(table, np.uint8)
|
||||
return Image.fromarray(np.array(img).astype(np.uint8)).point(lambda i: table[i])
|
||||
|
||||
# Node registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DepthEstimationNode": DepthEstimationNode
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DepthEstimationNode": "Depth Estimation"
|
||||
}
|
||||
+8
-4
@@ -1,4 +1,8 @@
|
||||
transformers==4.12.3
|
||||
Pillow==8.4.0
|
||||
numpy==1.21.2
|
||||
# Add other necessary dependencies without specifying torch again
|
||||
# Core dependencies matching ComfyUI
|
||||
torch>=2.0.0
|
||||
transformers>=4.28.1
|
||||
Pillow>=9.0.0
|
||||
numpy>=1.21.2
|
||||
|
||||
# Additional dependencies specific to depth estimation node
|
||||
timm>=0.6.12 # Required for depth estimation models
|
||||
|
||||
Reference in New Issue
Block a user