Merge pull request #9 from Limbicnation/feature/improved-depth-estimation-node

Feature/improved depth estimation node
This commit is contained in:
Gero Doll
2025-02-25 12:45:02 +01:00
committed by GitHub
+140 -40
View File
@@ -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,24 +144,54 @@ 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}")
@@ -123,31 +206,43 @@ 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) * 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
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)
# 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,)
@@ -159,11 +254,16 @@ class DepthEstimationNode:
torch.cuda.empty_cache()
gc.collect()
def gamma_correction(self, img, gamma=1.0):
def gamma_correction(self, img: Image.Image, gamma: float = 1.0) -> Image.Image:
"""Applies gamma correction to the image."""
inv_gamma = 1.0 / gamma
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])
# 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 = {