From 29fac572ec0b84e7acde8e85dd55b70aa0fc808b Mon Sep 17 00:00:00 2001 From: limbicnation Date: Tue, 25 Feb 2025 12:12:33 +0100 Subject: [PATCH] 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 = {