From ae35209d04b6c577530082e29b98871d0fd0a23a Mon Sep 17 00:00:00 2001 From: gero Date: Thu, 30 Jan 2025 18:03:34 +0100 Subject: [PATCH] feat(depth): enhance depth estimation node - Add proper VRAM management and cleanup - Improve model caching and error handling - Add FP16 support for CUDA devices - Optimize image processing pipeline - Add detailed logging - Fix tensor device handling - Add robust error recovery --- depth_estimation_node.py | 167 +++++++++++++++++++++------------------ 1 file changed, 92 insertions(+), 75 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index ed0dae0..7955d86 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -4,24 +4,35 @@ import torch from transformers import pipeline from PIL import Image, ImageFilter, ImageOps import folder_paths -from comfy.model_management import get_torch_device +from comfy.model_management import get_torch_device, get_free_memory +import gc +import logging + +# 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") +os.makedirs(MODELS_DIR, exist_ok=True) +os.environ["TRANSFORMERS_CACHE"] = MODELS_DIR 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.""" + MEDIAN_SIZES = ["3", "5", "7", "9", "11"] - + def __init__(self): - self.device = get_torch_device() + self.device = None self.depth_estimator = None self.current_model = None - + logger.info("Initialized DepthEstimationNode") + @classmethod def INPUT_TYPES(cls): return { @@ -37,84 +48,92 @@ class DepthEstimationNode: RETURN_TYPES = ("IMAGE",) FUNCTION = "estimate_depth" - CATEGORY = "image/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): - model_path = DEPTH_MODELS[model_name] - if self.depth_estimator is None or self.current_model != model_path: - try: + """Ensures the correct model is loaded with proper VRAM management.""" + try: + model_path = DEPTH_MODELS[model_name] + + if self.depth_estimator is None or self.current_model != model_path: + self.cleanup() + + if self.device is None: + self.device = get_torch_device() + + logger.info(f"Loading depth model: {model_name} on device {self.device}") + + # Use FP16 for CUDA devices to save VRAM + dtype = torch.float16 if 'cuda' in self.device else torch.float32 + self.depth_estimator = pipeline( "depth-estimation", model=model_path, - device=self.device + device=self.device, + torch_dtype=dtype ) self.current_model = model_path - except Exception as e: - raise RuntimeError(f"Failed to load model {model_name}: {str(e)}") + logger.info(f"Successfully loaded {model_name}") + + except Exception as e: + self.cleanup() + error_msg = f"Failed to load model {model_name}: {str(e)}" + logger.error(error_msg) + raise RuntimeError(error_msg) + + def process_image(self, image): + """Converts input image to proper format 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: + image_np = image_np[..., :3] + elif len(image_np.shape) == 2: + 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.""" 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) + pil_image = self.process_image(image) - # 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) + with torch.inference_mode(): + depth_result = self.depth_estimator(pil_image) + depth_map = depth_result["predicted_depth"].squeeze().cpu().numpy() - # 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() + # 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))) - 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 + depth_map = Image.fromarray(depth_map, mode='L') # 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 int(median_size) > 0: + depth_map = depth_map.filter(ImageFilter.MedianFilter(size=int(median_size))) + if apply_auto_contrast: depth_map = ImageOps.autocontrast(depth_map) @@ -125,28 +144,26 @@ class DepthEstimationNode: gamma = np.log(0.5) / np.log(mean_luminance) depth_map = self.gamma_correction(depth_map, gamma) - # Convert back to tensor format + # Convert to tensor 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) + depth_tensor = torch.from_numpy(depth_array).unsqueeze(0).to(self.device) return (depth_tensor,) except Exception as e: - raise RuntimeError(f"Depth estimation failed: {str(e)}") + error_msg = f"Depth estimation failed: {str(e)}" + logger.error(error_msg) + raise RuntimeError(error_msg) + finally: + torch.cuda.empty_cache() + gc.collect() def gamma_correction(self, img, gamma=1.0): + """Applies 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]) + 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 = { @@ -154,5 +171,5 @@ NODE_CLASS_MAPPINGS = { } NODE_DISPLAY_NAME_MAPPINGS = { - "DepthEstimationNode": "Depth Estimation" + "DepthEstimationNode": "Depth Estimation (V2)" } \ No newline at end of file