Enhance DepthEstimationNode with type hints, better VRAM management and all models

This commit is contained in:
limbicnation
2025-02-25 12:12:33 +01:00
parent 28a4ed29ca
commit 29fac572ec
+135 -30
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,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 = {