diff --git a/__init__.py b/__init__.py index 2ccd9ad..a062235 100644 --- a/__init__.py +++ b/__init__.py @@ -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" +] \ No newline at end of file diff --git a/depth_estimation_node.py b/depth_estimation_node.py index ba08f9a..ed0dae0 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -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" +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 3cccaa0..8741038 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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