feat: enhance depth node with robust error handling and resource management

This comprehensive update improves the depth estimation node with:
- Robust error handling that continues workflow execution instead of crashing
- Visual error reporting with informative messages displayed on error images
- Intelligent resource management with VRAM usage tracking and requirements
- Automatic fallback to CPU when insufficient VRAM is detected
- Multiple fallback strategies for model loading issues
- Better handling of problematic inputs like NaN values
- Detailed logging for easier troubleshooting

These changes make the node much more stable and user-friendly in
complex ComfyUI setups, preventing workflow-breaking errors.
This commit is contained in:
limbicnation
2025-05-03 18:29:07 +02:00
parent e8269297e9
commit d3b1095d43
+366 -88
View File
@@ -1,8 +1,10 @@
import os
import numpy as np
import torch
import traceback
import time
from transformers import pipeline
from PIL import Image, ImageFilter, ImageOps
from PIL import Image, ImageFilter, ImageOps, ImageDraw, ImageFont
import folder_paths
from comfy.model_management import get_torch_device, get_free_memory
import gc
@@ -28,13 +30,28 @@ 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
# Define all models mentioned in the README with memory requirements
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",
"Depth-Anything-Small": {
"path": "LiheYoung/depth-anything-small",
"vram_mb": 1500
},
"Depth-Anything-Base": {
"path": "LiheYoung/depth-anything-base",
"vram_mb": 2500
},
"Depth-Anything-Large": {
"path": "LiheYoung/depth-anything-large",
"vram_mb": 4000
},
"Depth-Anything-V2-Small": {
"path": "LiheYoung/depth-anything-small-hf",
"vram_mb": 1500
},
"Depth-Anything-V2-Base": {
"path": "LiheYoung/depth-anything-base-hf",
"vram_mb": 2500
},
}
class DepthEstimationNode:
@@ -65,6 +82,10 @@ class DepthEstimationNode:
"median_size": (cls.MEDIAN_SIZES, {"default": "5"}),
"apply_auto_contrast": ("BOOLEAN", {"default": True}),
"apply_gamma": ("BOOLEAN", {"default": True})
},
"optional": {
"force_reload": ("BOOLEAN", {"default": False}),
"force_cpu": ("BOOLEAN", {"default": False})
}
}
@@ -76,6 +97,10 @@ class DepthEstimationNode:
"""Clean up resources and free VRAM."""
try:
if self.depth_estimator is not None:
# Save model name before deletion for logging
model_name = self.current_model
# Delete the estimator
del self.depth_estimator
self.depth_estimator = None
self.current_model = None
@@ -85,41 +110,68 @@ class DepthEstimationNode:
torch.cuda.empty_cache()
gc.collect()
logger.info("Cleaned up model resources")
logger.info(f"Cleaned up model resources for {model_name}")
# Log available memory after cleanup if CUDA is available
if torch.cuda.is_available():
free_mem, total_mem = get_free_memory(get_torch_device())
logger.info(f"Available VRAM after cleanup: {free_mem/1024:.2f}MB of {total_mem/1024:.2f}MB")
except Exception as e:
logger.warning(f"Error during cleanup: {e}")
logger.debug(traceback.format_exc())
def ensure_model_loaded(self, model_name: str) -> None:
def ensure_model_loaded(self, model_name: str, force_reload: bool = False, force_cpu: bool = False) -> None:
"""
Ensures the correct model is loaded with proper VRAM management and fallback options.
Args:
model_name: The name of the model to load
force_reload: If True, reload the model even if it's already loaded
force_cpu: If True, force loading on CPU regardless of GPU availability
Raises:
RuntimeError: If the model fails to load after all fallback attempts
"""
try:
if model_name not in DEPTH_MODELS:
raise ValueError(f"Unknown model: {model_name}. Available models: {list(DEPTH_MODELS.keys())}")
available_models = list(DEPTH_MODELS.keys())
if len(available_models) > 0:
fallback_model = available_models[0]
logger.warning(f"Unknown model: {model_name}. Falling back to {fallback_model}")
model_name = fallback_model
else:
raise ValueError(f"No depth models available. Please check your installation.")
model_path = DEPTH_MODELS[model_name]
model_info = DEPTH_MODELS[model_name]
model_path = model_info["path"]
# Only reload if needed
if self.depth_estimator is None or self.current_model != model_path:
# Only reload if needed or forced
if force_reload or self.depth_estimator is None or self.current_model != model_path:
self.cleanup()
# Set up device
if self.device is None:
self.device = get_torch_device()
logger.info(f"Loading depth model: {model_name} on device {self.device}")
logger.info(f"Loading depth model: {model_name} on {'CPU' if force_cpu else self.device}")
# Check available memory if using CUDA
if torch.cuda.is_available() and not force_cpu:
free_mem, total_mem = get_free_memory(self.device)
required_mem = model_info.get("vram_mb", 2000) * 1024 # Convert to KB
logger.info(f"Available VRAM: {free_mem/1024:.2f}MB, Required: {required_mem/1024:.2f}MB")
# If not enough memory, fall back to CPU
if free_mem < required_mem:
logger.warning(f"Insufficient VRAM for {model_name} ({required_mem/1024:.1f}MB required, {free_mem/1024:.1f}MB available). Falling back to CPU.")
force_cpu = True
# Determine device type for pipeline
device_type = 'cuda' if torch.cuda.is_available() else 'cpu'
device_type = 'cpu' if force_cpu else ('cuda' if torch.cuda.is_available() else 'cpu')
# Use FP16 for CUDA devices to save VRAM
dtype = torch.float16 if 'cuda' in str(self.device) else torch.float32
dtype = torch.float16 if 'cuda' in str(self.device) and not force_cpu else torch.float32
# Create a dedicated cache directory for this model
cache_dir = os.path.join(MODELS_DIR, model_name.replace("-", "_").lower())
@@ -136,6 +188,8 @@ class DepthEstimationNode:
success = False
last_error = None
logger.info(f"Loading model with device={device_type}, dtype={dtype}")
for path in model_paths_to_try:
try:
logger.info(f"Attempting to load from: {path}")
@@ -180,21 +234,31 @@ class DepthEstimationNode:
continue
if not success:
# If all attempts failed, show helpful message with instructions
# If all attempts failed, try a different model
if model_name != "Depth-Anything-V2-Small" and "Depth-Anything-V2-Small" in DEPTH_MODELS:
logger.warning(f"Failed to load {model_name}, trying Depth-Anything-V2-Small as fallback")
try:
# Increase chances of success with CPU
return self.ensure_model_loaded("Depth-Anything-V2-Small", True, True)
except Exception as fallback_error:
logger.error(f"Fallback model also failed: {str(fallback_error)}")
# If still failing, show helpful message with instructions
error_msg = f"""
Failed to load model {model_name} after trying multiple sources.
Last error: {str(last_error)}
Try these solutions:
1. Run 'huggingface-cli login' in your terminal to authenticate
2. Check your internet connection
3. Try a different model version (e.g. Depth-Anything-V2-Small instead of Depth-Anything-Small)
"""
Failed to load model {model_name} after trying multiple sources.
Last error: {str(last_error)}
Try these solutions:
1. Run 'huggingface-cli login' in your terminal to authenticate
2. Check your internet connection
3. Try a different model version (e.g. Depth-Anything-V2-Small instead of Depth-Anything-Small)
4. Ensure you have enough VRAM available or use force_cpu=True
"""
logger.error(error_msg)
raise RuntimeError(error_msg)
# Ensure model is on the correct device
if hasattr(self.depth_estimator, 'model'):
if not force_cpu and hasattr(self.depth_estimator, 'model'):
self.depth_estimator.model = self.depth_estimator.model.to(self.device)
self.current_model = model_path
@@ -203,6 +267,7 @@ class DepthEstimationNode:
self.cleanup()
error_msg = f"Failed to load model {model_name}: {str(e)}"
logger.error(error_msg)
logger.debug(traceback.format_exc())
raise RuntimeError(error_msg)
def process_image(self, image: Union[torch.Tensor, np.ndarray]) -> Image.Image:
@@ -215,18 +280,134 @@ class DepthEstimationNode:
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)
try:
if torch.is_tensor(image):
# Check for NaN values in tensor
if torch.isnan(image).any():
logger.warning("Input tensor contains NaN values. Replacing with zeros.")
image = torch.nan_to_num(image, nan=0.0)
image_np = (image.cpu().numpy()[0] * 255).astype(np.uint8)
else:
# Check for NaN values in numpy array
if np.isnan(image).any():
logger.warning("Input array contains NaN values. Replacing with zeros.")
image = np.nan_to_num(image, nan=0.0)
image_np = (image * 255).astype(np.uint8)
if len(image_np.shape) == 3:
if image_np.shape[-1] == 4: # Handle RGBA images
image_np = image_np[..., :3]
elif len(image_np.shape) == 2: # Handle grayscale images
image_np = np.stack([image_np] * 3, axis=-1)
return Image.fromarray(image_np)
except Exception as e:
logger.error(f"Error processing image: {str(e)}")
logger.debug(traceback.format_exc())
# Return a placeholder image on error
return Image.new('RGB', (512, 512), (128, 128, 128))
def _create_error_image(self, input_image=None):
"""Create an error image placeholder based on input image if possible."""
try:
if input_image is not None and isinstance(input_image, torch.Tensor) and input_image.shape[0] > 0:
# Create gray error image with same dimensions as input
h, w = input_image.shape[2], input_image.shape[3]
# Gray background with slight red tint to indicate error
placeholder = torch.ones((1, h, w, 3), dtype=torch.float32) * torch.tensor([0.5, 0.4, 0.4])
if self.device is not None:
placeholder = placeholder.to(self.device)
return placeholder
else:
return self._create_basic_error_image()
except Exception:
return self._create_basic_error_image()
if len(image_np.shape) == 3:
if image_np.shape[-1] == 4: # Handle RGBA images
image_np = image_np[..., :3]
elif len(image_np.shape) == 2: # Handle grayscale images
image_np = np.stack([image_np] * 3, axis=-1)
def _create_basic_error_image(self):
"""Create a basic error image when no input dimensions are available."""
# Standard size error image (512x512)
h, w = 512, 512
# Gray background with slight red tint to indicate error
placeholder = torch.ones((1, h, w, 3), dtype=torch.float32) * torch.tensor([0.5, 0.4, 0.4])
return Image.fromarray(image_np)
if self.device is not None:
placeholder = placeholder.to(self.device)
return placeholder
def _add_error_text_to_image(self, image_tensor, error_text):
"""Add error text to the image tensor for visual feedback."""
try:
# Convert tensor to PIL for text rendering
if image_tensor is None:
return
temp_img = self._tensor_to_pil(image_tensor)
# Draw error text
draw = ImageDraw.Draw(temp_img)
# Try to get a font, fall back to default if needed
try:
font = ImageFont.truetype("arial.ttf", 20)
except:
font = ImageFont.load_default()
# Split text into multiple lines if too long
lines = []
words = error_text.split()
current_line = words[0] if words else "Error"
for word in words[1:]:
if len(current_line + " " + word) < 50:
current_line += " " + word
else:
lines.append(current_line)
current_line = word
lines.append(current_line)
# Draw title
draw.text((10, 10), "Depth Estimation Error", fill=(255, 50, 50), font=font)
# Draw error message
y_position = 40
for line in lines:
draw.text((10, y_position), line, fill=(255, 255, 255), font=font)
y_position += 25
# Convert back to tensor
result = self._pil_to_tensor(temp_img)
# Copy to original tensor if shapes match
if image_tensor.shape == result.shape:
image_tensor.copy_(result)
return image_tensor
except Exception as e:
logger.error(f"Error adding text to error image: {e}")
return image_tensor
def _tensor_to_pil(self, tensor):
"""Convert a tensor to PIL Image."""
if tensor.shape[0] == 1: # Batch size 1
img_np = (tensor[0].cpu().numpy() * 255).astype(np.uint8)
return Image.fromarray(img_np)
return Image.new('RGB', (512, 512), color=(128, 100, 100))
def _pil_to_tensor(self, pil_img):
"""Convert PIL Image back to tensor."""
img_np = np.array(pil_img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np).unsqueeze(0)
if self.device is not None:
tensor = tensor.to(self.device)
return tensor
def estimate_depth(self,
image: torch.Tensor,
@@ -234,7 +415,9 @@ class DepthEstimationNode:
blur_radius: float = 2.0,
median_size: str = "5",
apply_auto_contrast: bool = True,
apply_gamma: bool = True) -> Tuple[torch.Tensor]:
apply_gamma: bool = True,
force_reload: bool = False,
force_cpu: bool = False) -> Tuple[torch.Tensor]:
"""
Estimates depth from input image with error handling and cleanup.
@@ -245,73 +428,168 @@ class DepthEstimationNode:
median_size: Size of median filter for noise reduction
apply_auto_contrast: Whether to enhance contrast automatically
apply_gamma: Whether to apply gamma correction
force_reload: Whether to force reload the model
force_cpu: Whether to force using CPU for inference
Returns:
Tuple containing depth map tensor
Raises:
RuntimeError: If depth estimation fails
ValueError: If invalid parameters are provided
"""
error_image = None
start_time = time.time()
try:
# Validate inputs
if image is None or image.numel() == 0:
raise ValueError("Empty or null input image")
if image.ndim != 4:
raise ValueError(f"Expected 4D tensor for image, got {image.ndim}D.")
# Create error image placeholder based on input dimensions
error_image = self._create_error_image(image)
if torch.isnan(image).any():
logger.warning("Input image contains NaN values. These will be replaced.")
image = torch.nan_to_num(image, nan=0.0)
if median_size not in self.MEDIAN_SIZES:
raise ValueError(f"Invalid median_size. Must be one of {self.MEDIAN_SIZES}")
logger.warning(f"Invalid median_size: {median_size}. Defaulting to 5")
median_size = "5"
self.ensure_model_loaded(model_name)
pil_image = self.process_image(image)
# Load model with fallback strategy - wrapped in try-except
try:
self.ensure_model_loaded(model_name, force_reload, force_cpu)
except Exception as model_error:
# Special handling for model loading errors - common issue
error_msg = f"Failed to load model '{model_name}': {str(model_error)}"
logger.error(error_msg)
# Add error text to error image
self._add_error_text_to_image(error_image, f"Model Error: {str(model_error)[:100]}...")
return (error_image,)
with torch.inference_mode():
depth_result = self.depth_estimator(pil_image)
depth_map = depth_result["predicted_depth"].squeeze().cpu().numpy()
# Process input image
try:
pil_image = self.process_image(image)
except Exception as img_error:
logger.error(f"Image processing error: {str(img_error)}")
self._add_error_text_to_image(error_image, f"Image Error: {str(img_error)[:100]}...")
return (error_image,)
# 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) * 255.0)
depth_map = depth_map.astype(np.uint8)
# Perform depth estimation with error catching
try:
with torch.inference_mode():
depth_result = self.depth_estimator(pil_image)
depth_map = depth_result["predicted_depth"].squeeze().cpu().numpy()
except RuntimeError as rt_error:
# Check specifically for CUDA out-of-memory errors
if "CUDA out of memory" in str(rt_error):
error_msg = (
f"CUDA out of memory while processing depth map. "
f"Try using a smaller model or reducing image size."
)
logger.error(error_msg)
# Try to fall back to CPU if we hit OOM
if not force_cpu:
logger.info("Attempting to fall back to CPU due to CUDA OOM error")
try:
return self.estimate_depth(
image, model_name, blur_radius, median_size,
apply_auto_contrast, apply_gamma, True, True
)
except Exception as cpu_fallback_error:
logger.error(f"CPU fallback also failed: {str(cpu_fallback_error)}")
self._add_error_text_to_image(error_image, "CUDA Out of Memory. Try a smaller model.")
return (error_image,)
else:
# Other runtime errors
error_msg = f"Runtime error during depth estimation: {str(rt_error)}"
logger.error(error_msg)
logger.debug(traceback.format_exc())
self._add_error_text_to_image(error_image, f"Runtime Error: {str(rt_error)[:100]}...")
return (error_image,)
except Exception as e:
# General exceptions
error_msg = f"Depth estimation failed: {str(e)}"
logger.error(error_msg)
logger.debug(traceback.format_exc())
self._add_error_text_to_image(error_image, f"Error: {str(e)[:100]}...")
return (error_image,)
# Create PIL image explicitly with L mode (grayscale)
depth_pil = Image.fromarray(depth_map, mode='L')
# Check for NaN values in depth map
if np.isnan(depth_map).any():
logger.warning("Depth map contains NaN values. Replacing with zeros.")
depth_map = np.nan_to_num(depth_map, nan=0.0)
# Apply post-processing
if blur_radius > 0:
depth_pil = depth_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius))
if int(median_size) > 0:
depth_pil = depth_pil.filter(ImageFilter.MedianFilter(size=int(median_size)))
if apply_auto_contrast:
depth_pil = ImageOps.autocontrast(depth_pil)
if apply_gamma:
# Continue with the normal depth map processing
try:
# 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) * 255.0)
depth_map = depth_map.astype(np.uint8)
# Create PIL image explicitly with L mode (grayscale)
depth_pil = Image.fromarray(depth_map, mode='L')
# Apply post-processing
if blur_radius > 0:
depth_pil = depth_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius))
if int(median_size) > 0:
depth_pil = depth_pil.filter(ImageFilter.MedianFilter(size=int(median_size)))
if apply_auto_contrast:
depth_pil = ImageOps.autocontrast(depth_pil)
if apply_gamma:
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)
# 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 - explicitly handle as grayscale
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)
# 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 - 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)
if self.device is not None:
depth_tensor = depth_tensor.to(self.device)
return (depth_tensor,)
# Make it compatible with ComfyUI by creating a 3-channel image
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)
if self.device is not None:
depth_tensor = depth_tensor.to(self.device)
processing_time = time.time() - start_time
logger.info(f"Depth processing completed in {processing_time:.2f} seconds")
return (depth_tensor,)
except Exception as post_error:
error_msg = f"Error during depth map post-processing: {str(post_error)}"
logger.error(error_msg)
logger.debug(traceback.format_exc())
self._add_error_text_to_image(error_image, f"Post-processing Error: {str(post_error)[:100]}...")
return (error_image,)
except Exception as e:
# Catch-all for any other exceptions
error_msg = f"Depth estimation failed: {str(e)}"
logger.error(error_msg)
raise RuntimeError(error_msg)
logger.debug(traceback.format_exc())
# If error_image hasn't been created yet, create a basic one
if error_image is None:
error_image = self._create_basic_error_image()
self._add_error_text_to_image(error_image, f"Unexpected Error: {str(e)[:100]}...")
return (error_image,)
finally:
# Always clean up resources
torch.cuda.empty_cache()
gc.collect()