From 29fac572ec0b84e7acde8e85dd55b70aa0fc808b Mon Sep 17 00:00:00 2001 From: limbicnation Date: Tue, 25 Feb 2025 12:12:33 +0100 Subject: [PATCH 1/4] Enhance DepthEstimationNode with type hints, better VRAM management and all models --- depth_estimation_node.py | 165 ++++++++++++++++++++++++++++++++------- 1 file changed, 135 insertions(+), 30 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index 7955d86..6b3303a 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -7,23 +7,44 @@ import folder_paths from comfy.model_management import get_torch_device, get_free_memory import gc import logging +from typing import Tuple, List, Dict, Any, Optional, Union # Setup logging logging.basicConfig(level=logging.INFO) logger = logging.getLogger("DepthEstimation") -# Configure model caching directory -MODELS_DIR = os.path.join(folder_paths.get_folder_paths("models")[0], "depth_anything") +# Configure model paths +if not hasattr(folder_paths, "models_dir"): + folder_paths.models_dir = os.path.join(folder_paths.base_path, "models") + +# Register depth models path +DEPTH_DIR = "depth_anything" +folder_paths.folder_names_and_paths[DEPTH_DIR] = ([ + os.path.join(folder_paths.models_dir, DEPTH_DIR) +], folder_paths.supported_pt_extensions) + +# Set models directory +MODELS_DIR = folder_paths.folder_names_and_paths[DEPTH_DIR][0][0] os.makedirs(MODELS_DIR, exist_ok=True) os.environ["TRANSFORMERS_CACHE"] = MODELS_DIR +# Define all models mentioned in the README 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", } class DepthEstimationNode: - """ComfyUI node for depth estimation using Depth Anything models.""" + """ + ComfyUI node for depth estimation using Depth Anything models. + + This node provides depth map generation from images using various Depth Anything models + with configurable post-processing options like blur, median filtering, contrast enhancement, + and gamma correction. + """ MEDIAN_SIZES = ["3", "5", "7", "9", "11"] @@ -34,7 +55,8 @@ class DepthEstimationNode: logger.info("Initialized DepthEstimationNode") @classmethod - def INPUT_TYPES(cls): + def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]: + """Define the input types for the node.""" return { "required": { "image": ("IMAGE",), @@ -50,19 +72,37 @@ class DepthEstimationNode: FUNCTION = "estimate_depth" CATEGORY = "depth" - def cleanup(self): - """Clean up resources and VRAM.""" - if self.depth_estimator is not None: - del self.depth_estimator - self.depth_estimator = None - self.current_model = None - torch.cuda.empty_cache() - gc.collect() - logger.info("Cleaned up model resources") - - def ensure_model_loaded(self, model_name): - """Ensures the correct model is loaded with proper VRAM management.""" + def cleanup(self) -> None: + """Clean up resources and free VRAM.""" try: + if self.depth_estimator is not None: + del self.depth_estimator + self.depth_estimator = None + self.current_model = None + + # Force CUDA cache clearing + if torch.cuda.is_available(): + torch.cuda.empty_cache() + gc.collect() + + logger.info("Cleaned up model resources") + except Exception as e: + logger.warning(f"Error during cleanup: {e}") + + def ensure_model_loaded(self, model_name: str) -> None: + """ + Ensures the correct model is loaded with proper VRAM management. + + Args: + model_name: The name of the model to load + + Raises: + RuntimeError: If the model fails to load + """ + try: + if model_name not in DEPTH_MODELS: + raise ValueError(f"Unknown model: {model_name}. Available models: {list(DEPTH_MODELS.keys())}") + model_path = DEPTH_MODELS[model_name] if self.depth_estimator is None or self.current_model != model_path: @@ -73,15 +113,28 @@ class DepthEstimationNode: logger.info(f"Loading depth model: {model_name} on device {self.device}") + # Determine device type for pipeline + device_type = 'cuda' if torch.cuda.is_available() else 'cpu' + # Use FP16 for CUDA devices to save VRAM - dtype = torch.float16 if 'cuda' in self.device else torch.float32 + dtype = torch.float16 if 'cuda' in str(self.device) else torch.float32 + + # Check available VRAM before loading + if torch.cuda.is_available(): + free_vram = get_free_memory(self.device) + logger.info(f"Available VRAM before loading: {free_vram / (1024**3):.2f} GB") self.depth_estimator = pipeline( "depth-estimation", model=model_path, - device=self.device, + device_map=device_type, torch_dtype=dtype ) + + # Ensure model is on the correct device + if hasattr(self.depth_estimator, 'model'): + self.depth_estimator.model = self.depth_estimator.model.to(self.device) + self.current_model = model_path logger.info(f"Successfully loaded {model_name}") @@ -91,32 +144,67 @@ class DepthEstimationNode: logger.error(error_msg) raise RuntimeError(error_msg) - def process_image(self, image): - """Converts input image to proper format for depth estimation.""" + def process_image(self, image: Union[torch.Tensor, np.ndarray]) -> Image.Image: + """ + Converts input image to proper format for depth estimation. + + Args: + image: Input image as tensor or numpy array + + Returns: + PIL Image ready for depth estimation + """ if torch.is_tensor(image): image_np = (image.cpu().numpy()[0] * 255).astype(np.uint8) else: image_np = (image * 255).astype(np.uint8) if len(image_np.shape) == 3: - if image_np.shape[-1] == 4: + if image_np.shape[-1] == 4: # Handle RGBA images image_np = image_np[..., :3] - elif len(image_np.shape) == 2: + elif len(image_np.shape) == 2: # Handle grayscale images image_np = np.stack([image_np] * 3, axis=-1) return Image.fromarray(image_np) - def estimate_depth(self, image, model_name, blur_radius=2.0, median_size="5", - apply_auto_contrast=True, apply_gamma=True): - """Estimates depth from input image with error handling and cleanup.""" + def estimate_depth(self, + image: torch.Tensor, + model_name: str, + blur_radius: float = 2.0, + median_size: str = "5", + apply_auto_contrast: bool = True, + apply_gamma: bool = True) -> Tuple[torch.Tensor]: + """ + Estimates depth from input image with error handling and cleanup. + + Args: + image: Input image tensor + model_name: Name of the depth model to use + blur_radius: Gaussian blur radius for smoothing + median_size: Size of median filter for noise reduction + apply_auto_contrast: Whether to enhance contrast automatically + apply_gamma: Whether to apply gamma correction + + Returns: + Tuple containing depth map tensor + + Raises: + RuntimeError: If depth estimation fails + ValueError: If invalid parameters are provided + """ try: if median_size not in self.MEDIAN_SIZES: raise ValueError(f"Invalid median_size. Must be one of {self.MEDIAN_SIZES}") + # Load model if needed self.ensure_model_loaded(model_name) + + # Process image to PIL format pil_image = self.process_image(image) + # Run inference with torch.inference_mode(): + logger.info(f"Running depth estimation on image size {pil_image.size}") depth_result = self.depth_estimator(pil_image) depth_map = depth_result["predicted_depth"].squeeze().cpu().numpy() @@ -147,8 +235,13 @@ class DepthEstimationNode: # Convert to tensor depth_array = np.array(depth_map).astype(np.float32) / 255.0 depth_array = np.stack([depth_array] * 3, axis=-1) - depth_tensor = torch.from_numpy(depth_array).unsqueeze(0).to(self.device) + depth_tensor = torch.from_numpy(depth_array).unsqueeze(0) + # Move tensor to the correct device if needed + if self.device is not None: + depth_tensor = depth_tensor.to(self.device) + + logger.info(f"Depth estimation completed successfully") return (depth_tensor,) except Exception as e: @@ -156,14 +249,26 @@ class DepthEstimationNode: logger.error(error_msg) raise RuntimeError(error_msg) finally: - torch.cuda.empty_cache() + # Ensure proper cleanup + if torch.cuda.is_available(): + torch.cuda.empty_cache() gc.collect() - def gamma_correction(self, img, gamma=1.0): - """Applies gamma correction to the image.""" + def gamma_correction(self, img: Image.Image, gamma: float = 1.0) -> Image.Image: + """ + Applies gamma correction to the image. + + Args: + img: Input PIL image + gamma: Gamma value for correction + + Returns: + Gamma-corrected PIL image + """ inv_gamma = 1.0 / gamma + # Create lookup table for faster processing table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in range(256)], np.uint8) - return Image.fromarray(np.array(img)).point(lambda x: table[x]) + return ImageOps.gamma(img, gamma) # Using built-in PIL gamma for better performance # Node registration NODE_CLASS_MAPPINGS = { From 196315e3313b6d390dc6a44649f31a035a44979c Mon Sep 17 00:00:00 2001 From: limbicnation Date: Tue, 25 Feb 2025 12:23:02 +0100 Subject: [PATCH 2/4] Enhance DepthEstimationNode with type hints, better VRAM management and all models --- depth_estimation_node.py | 34 ++++++++++++++++++++-------------- 1 file changed, 20 insertions(+), 14 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index 6b3303a..8d2869a 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -168,12 +168,12 @@ class DepthEstimationNode: return Image.fromarray(image_np) def estimate_depth(self, - image: torch.Tensor, - model_name: str, - blur_radius: float = 2.0, - median_size: str = "5", - apply_auto_contrast: bool = True, - apply_gamma: bool = True) -> Tuple[torch.Tensor]: + image: torch.Tensor, + model_name: str, + blur_radius: float = 2.0, + median_size: str = "5", + apply_auto_contrast: bool = True, + apply_gamma: bool = True) -> Tuple[torch.Tensor]: """ Estimates depth from input image with error handling and cleanup. @@ -211,7 +211,7 @@ class DepthEstimationNode: # Normalize depth values depth_min, depth_max = depth_map.min(), depth_map.max() if depth_max > depth_min: - depth_map = ((depth_map - depth_min) * (255.0 / (depth_max - depth_min))) + depth_map = ((depth_map - depth_min) / (depth_max - depth_min + 1e-8) * 255) depth_map = depth_map.astype(np.uint8) depth_map = Image.fromarray(depth_map, mode='L') @@ -232,10 +232,10 @@ class DepthEstimationNode: gamma = np.log(0.5) / np.log(mean_luminance) depth_map = self.gamma_correction(depth_map, gamma) - # Convert to tensor + # Convert to tensor - give option for single-channel output to save memory depth_array = np.array(depth_map).astype(np.float32) / 255.0 - depth_array = np.stack([depth_array] * 3, axis=-1) - depth_tensor = torch.from_numpy(depth_array).unsqueeze(0) + # Return single-channel depth map to save memory (ComfyUI can handle both) + depth_tensor = torch.from_numpy(depth_array).unsqueeze(0).unsqueeze(0) # Move tensor to the correct device if needed if self.device is not None: @@ -265,10 +265,16 @@ class DepthEstimationNode: Returns: Gamma-corrected PIL image """ - inv_gamma = 1.0 / gamma - # Create lookup table for faster processing - table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in range(256)], np.uint8) - return ImageOps.gamma(img, gamma) # Using built-in PIL gamma for better performance + # Use built-in PIL gamma correction with error handling + try: + return ImageOps.autocontrast(ImageOps.gamma(img, gamma)) + except Exception as e: + logger.warning(f"Built-in gamma correction failed: {e}, using manual implementation") + # Fallback to manual implementation + inv_gamma = 1.0 / gamma + # Create lookup table for faster processing + table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in range(256)], np.uint8) + return Image.fromarray(np.array(img)).point(lambda x: table[x]) # Node registration NODE_CLASS_MAPPINGS = { From 8b7fc49a1659c958e738c60dea3685597f69f15a Mon Sep 17 00:00:00 2001 From: limbicnation Date: Tue, 25 Feb 2025 12:27:53 +0100 Subject: [PATCH 3/4] fix: improve depth estimation with robust gamma correction and error handling --- depth_estimation_node.py | 38 +++++++++++++++++++++----------------- 1 file changed, 21 insertions(+), 17 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index 8d2869a..5f7a8fb 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -255,26 +255,30 @@ class DepthEstimationNode: gc.collect() def gamma_correction(self, img: Image.Image, gamma: float = 1.0) -> Image.Image: - """ - Applies gamma correction to the image. - - Args: - img: Input PIL image - gamma: Gamma value for correction - - Returns: - Gamma-corrected PIL image - """ - # Use built-in PIL gamma correction with error handling + """Applies gamma correction to the image with proper error handling.""" try: - return ImageOps.autocontrast(ImageOps.gamma(img, gamma)) - except Exception as e: - logger.warning(f"Built-in gamma correction failed: {e}, using manual implementation") - # Fallback to manual implementation + # Convert PIL Image to numpy array + img_array = np.array(img) + + # Apply gamma correction inv_gamma = 1.0 / gamma - # Create lookup table for faster processing + # Create a lookup table for gamma correction table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in range(256)], np.uint8) - return Image.fromarray(np.array(img)).point(lambda x: table[x]) + + # Apply the lookup table (avoiding PIL's point method which can cause the error) + corrected_array = table[img_array] + + # Ensure the array has the right shape for PIL + # If it's a single-channel grayscale image, it should be 2D for PIL + if len(corrected_array.shape) == 3 and corrected_array.shape[2] == 1: + corrected_array = corrected_array.squeeze(2) + + # Convert back to PIL Image with explicit mode + return Image.fromarray(corrected_array, mode='L') + except Exception as e: + logger.error(f"Gamma correction failed: {e}") + # Return the original image if correction fails + return img # Node registration NODE_CLASS_MAPPINGS = { From b184437560fa4e3c677c0b5b8fd97bb93647953d Mon Sep 17 00:00:00 2001 From: limbicnation Date: Tue, 25 Feb 2025 12:39:08 +0100 Subject: [PATCH 4/4] Fix PIL Image conversion error in depth estimation pipeline --- depth_estimation_node.py | 75 ++++++++++++++++------------------------ 1 file changed, 30 insertions(+), 45 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index 5f7a8fb..ea3c027 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -196,52 +196,54 @@ class DepthEstimationNode: if median_size not in self.MEDIAN_SIZES: raise ValueError(f"Invalid median_size. Must be one of {self.MEDIAN_SIZES}") - # Load model if needed self.ensure_model_loaded(model_name) - - # Process image to PIL format pil_image = self.process_image(image) - # Run inference with torch.inference_mode(): - logger.info(f"Running depth estimation on image size {pil_image.size}") depth_result = self.depth_estimator(pil_image) depth_map = depth_result["predicted_depth"].squeeze().cpu().numpy() # Normalize depth values depth_min, depth_max = depth_map.min(), depth_map.max() if depth_max > depth_min: - depth_map = ((depth_map - depth_min) / (depth_max - depth_min + 1e-8) * 255) + depth_map = ((depth_map - depth_min) / (depth_max - depth_min) * 255.0) depth_map = depth_map.astype(np.uint8) - depth_map = Image.fromarray(depth_map, mode='L') + + # Create PIL image explicitly with L mode (grayscale) + depth_pil = Image.fromarray(depth_map, mode='L') # Apply post-processing if blur_radius > 0: - depth_map = depth_map.filter(ImageFilter.GaussianBlur(radius=blur_radius)) + depth_pil = depth_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius)) if int(median_size) > 0: - depth_map = depth_map.filter(ImageFilter.MedianFilter(size=int(median_size))) + depth_pil = depth_pil.filter(ImageFilter.MedianFilter(size=int(median_size))) if apply_auto_contrast: - depth_map = ImageOps.autocontrast(depth_map) + depth_pil = ImageOps.autocontrast(depth_pil) if apply_gamma: - depth_array = np.array(depth_map).astype(np.float32) / 255.0 + depth_array = np.array(depth_pil).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) + # Use direct numpy operations for gamma correction + corrected = np.power(depth_array, 1.0/gamma) * 255.0 + depth_pil = Image.fromarray(corrected.astype(np.uint8), mode='L') - # Convert to tensor - give option for single-channel output to save memory - depth_array = np.array(depth_map).astype(np.float32) / 255.0 - # Return single-channel depth map to save memory (ComfyUI can handle both) - depth_tensor = torch.from_numpy(depth_array).unsqueeze(0).unsqueeze(0) + # Convert to tensor - explicitly handle as grayscale + depth_array = np.array(depth_pil).astype(np.float32) / 255.0 + + # Make it compatible with ComfyUI by creating a 3-channel image + # Use proper reshaping to avoid dimension issues + h, w = depth_array.shape + depth_rgb = np.stack([depth_array] * 3, axis=-1) # Create proper 3D array with shape (h, w, 3) + + depth_tensor = torch.from_numpy(depth_rgb).unsqueeze(0) - # Move tensor to the correct device if needed if self.device is not None: depth_tensor = depth_tensor.to(self.device) - logger.info(f"Depth estimation completed successfully") return (depth_tensor,) except Exception as e: @@ -249,36 +251,19 @@ class DepthEstimationNode: logger.error(error_msg) raise RuntimeError(error_msg) finally: - # Ensure proper cleanup - if torch.cuda.is_available(): - torch.cuda.empty_cache() + torch.cuda.empty_cache() gc.collect() def gamma_correction(self, img: Image.Image, gamma: float = 1.0) -> Image.Image: - """Applies gamma correction to the image with proper error handling.""" - try: - # Convert PIL Image to numpy array - img_array = np.array(img) - - # Apply gamma correction - inv_gamma = 1.0 / gamma - # Create a lookup table for gamma correction - table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in range(256)], np.uint8) - - # Apply the lookup table (avoiding PIL's point method which can cause the error) - corrected_array = table[img_array] - - # Ensure the array has the right shape for PIL - # If it's a single-channel grayscale image, it should be 2D for PIL - if len(corrected_array.shape) == 3 and corrected_array.shape[2] == 1: - corrected_array = corrected_array.squeeze(2) - - # Convert back to PIL Image with explicit mode - return Image.fromarray(corrected_array, mode='L') - except Exception as e: - logger.error(f"Gamma correction failed: {e}") - # Return the original image if correction fails - return img + """Applies gamma correction to the image.""" + # Convert to numpy array + img_array = np.array(img) + + # Apply gamma correction directly with numpy + corrected = np.power(img_array.astype(np.float32) / 255.0, 1.0/gamma) * 255.0 + + # Ensure uint8 type and create image with explicit mode + return Image.fromarray(corrected.astype(np.uint8), mode='L') # Node registration NODE_CLASS_MAPPINGS = {