Enhance DepthEstimationNode with type hints, better VRAM management and all models
This commit is contained in:
+135
-30
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user