Merge pull request #15 from Limbicnation/feature/da3-integration
feat: integrate Depth Anything V3 (Small/Base) models
This commit is contained in:
@@ -13,6 +13,8 @@ A robust custom depth estimation node for ComfyUI using Depth-Anything models to
|
||||
- Depth-Anything-Large
|
||||
- Depth-Anything-V2-Small
|
||||
- Depth-Anything-V2-Base
|
||||
- Depth-Anything-V3-Small (Requires optional dependency)
|
||||
- Depth-Anything-V3-Base (Requires optional dependency)
|
||||
- Post-processing options:
|
||||
- Gaussian blur (adjustable radius)
|
||||
- Median filtering (configurable size)
|
||||
@@ -46,6 +48,11 @@ A robust custom depth estimation node for ComfyUI using Depth-Anything models to
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
4. (Optional) To enable Depth Anything V3 models:
|
||||
```bash
|
||||
pip install git+https://github.com/ByteDance-Seed/Depth-Anything-3.git
|
||||
```
|
||||
|
||||
4. Restart ComfyUI to load the new custom node.
|
||||
|
||||
> **Note**: On first use, the node will download the selected model from Hugging Face. This may take some time depending on your internet connection.
|
||||
@@ -91,6 +98,8 @@ A robust custom depth estimation node for ComfyUI using Depth-Anything models to
|
||||
| Depth-Anything-Small | Good | ~1.5 GB | Fast |
|
||||
| Depth-Anything-V2-Base | Better | ~2.5 GB | Medium |
|
||||
| Depth-Anything-Base | Better | ~2.5 GB | Medium |
|
||||
| Depth-Anything-V3-Small | Excellent | ~2.0 GB | Fast |
|
||||
| Depth-Anything-V3-Base | Superior | ~2.5 GB | Medium |
|
||||
| Depth-Anything-Large | Best | ~4.0 GB | Slow |
|
||||
|
||||
## Troubleshooting Guide
|
||||
|
||||
+36
-2
@@ -12,7 +12,7 @@ logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("DepthEstimation")
|
||||
|
||||
# Version info
|
||||
__version__ = "1.1.1"
|
||||
__version__ = "1.2.0"
|
||||
|
||||
# Node class mappings - will be populated based on dependency checks
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
@@ -31,7 +31,14 @@ required_dependencies = {
|
||||
"huggingface_hub": "0.16.0"
|
||||
}
|
||||
|
||||
# Optional dependencies for DA3 (Depth Anything V3) support
|
||||
# Node works without these but DA3 models will not be available
|
||||
optional_dependencies = {
|
||||
"depth_anything_3": "0.0.0" # For DA3 models
|
||||
}
|
||||
|
||||
missing_dependencies = []
|
||||
DA3_AVAILABLE = False
|
||||
|
||||
# Check each dependency
|
||||
for module_name, min_version in required_dependencies.items():
|
||||
@@ -49,6 +56,32 @@ for module_name, min_version in required_dependencies.items():
|
||||
missing_dependencies.append(f"{module_name}>={min_version}")
|
||||
logger.warning(f"Missing required dependency: {module_name}>={min_version}")
|
||||
|
||||
# Check optional DA3 dependencies
|
||||
for module_name, min_version in optional_dependencies.items():
|
||||
try:
|
||||
module = __import__(module_name)
|
||||
module_version = getattr(module, "__version__", "0.0.0")
|
||||
|
||||
logger.info(f"Found optional {module_name} version {module_version}")
|
||||
|
||||
except ImportError:
|
||||
logger.info(f"Optional dependency {module_name} not installed. DA3 models will not be available.")
|
||||
except Exception as e:
|
||||
logger.warning(f"Error checking version for {module_name}: {e}. DA3 models might not work correctly.")
|
||||
|
||||
# Version comparison
|
||||
if module_name == "depth_anything_3" and "module_version" in locals():
|
||||
try:
|
||||
# Compare versions. This simple check works for "X.Y.Z" formats.
|
||||
if tuple(map(int, module_version.split('.'))) >= tuple(map(int, min_version.split('.'))):
|
||||
DA3_AVAILABLE = True
|
||||
else:
|
||||
logger.warning(f"Optional dependency {module_name} version {module_version} is older than required {min_version}. DA3 models may not be available or work correctly.")
|
||||
except ValueError:
|
||||
# Fallback for non-standard version strings (e.g. from git)
|
||||
logger.info(f"Could not parse version {module_version} for {module_name}. Assuming compatible.")
|
||||
DA3_AVAILABLE = True
|
||||
|
||||
if missing_dependencies:
|
||||
# Create placeholder node with dependency error
|
||||
class DependencyErrorNode:
|
||||
@@ -123,5 +156,6 @@ __all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"__version__",
|
||||
"WEB_DIRECTORY"
|
||||
"WEB_DIRECTORY",
|
||||
"DA3_AVAILABLE"
|
||||
]
|
||||
+205
-8
@@ -25,6 +25,13 @@ except ImportError:
|
||||
TIMM_AVAILABLE = False
|
||||
print("Warning: timm not available. Direct loading of Depth Anything models may not work.")
|
||||
|
||||
# Import DA3 availability status from the package's __init__
|
||||
from . import DA3_AVAILABLE
|
||||
|
||||
# Conditionally import Depth Anything V3 if available
|
||||
if DA3_AVAILABLE:
|
||||
from depth_anything_3.api import DepthAnything3
|
||||
|
||||
# Setup logging
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("DepthEstimation")
|
||||
@@ -288,10 +295,30 @@ DEPTH_MODELS = {
|
||||
"direct_url": "https://github.com/intel-isl/MiDaS/releases/download/v2_1/midas_v21_small_256.pt"
|
||||
},
|
||||
"MiDaS-Base": {
|
||||
"path": "Intel/dpt-hybrid-midas",
|
||||
"path": "Intel/dpt-hybrid-midas",
|
||||
"vram_mb": 1200,
|
||||
"midas_type": "DPT_Hybrid",
|
||||
"direct_url": "https://github.com/intel-isl/MiDaS/releases/download/v3/dpt_hybrid-midas-501f0c75.pt"
|
||||
},
|
||||
# DA3 (Depth Anything V3) Models - Apache 2.0 Licensed (Commercial Friendly)
|
||||
# Note: These models require the depth_anything_v3 package to be installed
|
||||
"Depth-Anything-V3-Small": {
|
||||
"path": "depth-anything/DA3-Small",
|
||||
"vram_mb": 2000, # Estimated: 80M params
|
||||
"model_type": "v3",
|
||||
"encoder": "vits",
|
||||
"license": "Apache-2.0",
|
||||
"supports_batch": True, # Multi-view support
|
||||
"params": "80M"
|
||||
},
|
||||
"Depth-Anything-V3-Base": {
|
||||
"path": "depth-anything/DA3-Base",
|
||||
"vram_mb": 2500, # Estimated: 120M params
|
||||
"model_type": "v3",
|
||||
"encoder": "vitb",
|
||||
"license": "Apache-2.0",
|
||||
"supports_batch": True, # Multi-view support
|
||||
"params": "120M"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -588,6 +615,131 @@ class MiDaSWrapper:
|
||||
|
||||
return {"predicted_depth": dummy_tensor}
|
||||
|
||||
|
||||
class DA3ModelWrapper:
|
||||
"""
|
||||
Wrapper to make Depth Anything V3 API compatible with existing node interface.
|
||||
|
||||
DA3 uses a different API and output format than V1/V2, so this wrapper normalizes
|
||||
the output to be compatible with the existing post-processing pipeline.
|
||||
|
||||
Supports both single images and batches for multi-view depth estimation.
|
||||
"""
|
||||
|
||||
def __init__(self, model, device):
|
||||
"""
|
||||
Initialize the DA3 wrapper.
|
||||
|
||||
Args:
|
||||
model: The loaded DA3 model instance
|
||||
device: The device to run inference on (cuda/cpu)
|
||||
"""
|
||||
self.model = model
|
||||
self.device = device
|
||||
logger.info(f"DA3ModelWrapper initialized on device: {device}")
|
||||
|
||||
def __call__(self, image: Union[Image.Image, List[Image.Image]]) -> dict:
|
||||
"""
|
||||
Run inference and return in V1/V2 compatible format.
|
||||
|
||||
Supports both single images and batches for multi-view depth estimation.
|
||||
|
||||
Args:
|
||||
image: Single PIL Image or list of PIL Images for batch processing
|
||||
|
||||
Returns:
|
||||
Dictionary with 'predicted_depth' key containing the depth tensor
|
||||
"""
|
||||
try:
|
||||
# Handle single image or batch
|
||||
if isinstance(image, Image.Image):
|
||||
images = [image]
|
||||
is_batch = False
|
||||
else:
|
||||
images = list(image)
|
||||
is_batch = True
|
||||
|
||||
# DA3 inference
|
||||
with torch.inference_mode():
|
||||
prediction = self.model.inference(images)
|
||||
|
||||
# Extract depth maps from prediction
|
||||
# DA3 returns prediction.depth as [N, H, W] numpy array
|
||||
if hasattr(prediction, 'depth'):
|
||||
depths = prediction.depth # [N, H, W] numpy array
|
||||
else:
|
||||
# Fallback if prediction format differs
|
||||
depths = np.array(prediction)
|
||||
|
||||
# Normalize each depth map to [0, 1] range
|
||||
normalized_depths = []
|
||||
for depth in depths:
|
||||
depth_min, depth_max = depth.min(), depth.max()
|
||||
if depth_max - depth_min > 1e-6:
|
||||
depth = (depth - depth_min) / (depth_max - depth_min)
|
||||
else:
|
||||
depth = np.zeros_like(depth)
|
||||
normalized_depths.append(depth)
|
||||
|
||||
# Stack and convert to tensor
|
||||
depth_array = np.stack(normalized_depths, axis=0) # [N, H, W]
|
||||
depth_tensor = torch.from_numpy(depth_array).float().to(self.device)
|
||||
|
||||
if not is_batch:
|
||||
depth_tensor = depth_tensor.squeeze(0) # [H, W] for single image
|
||||
|
||||
logger.info(f"DA3 inference complete. Output shape: {depth_tensor.shape}")
|
||||
|
||||
return {"predicted_depth": depth_tensor}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in DA3 inference: {e}")
|
||||
logger.error(traceback.format_exc())
|
||||
|
||||
# Return placeholder depth map on error
|
||||
if isinstance(image, Image.Image):
|
||||
w, h = image.size
|
||||
elif isinstance(image, list) and len(image) > 0:
|
||||
w, h = image[0].size
|
||||
else:
|
||||
w, h = 512, 512
|
||||
|
||||
dummy_tensor = torch.ones((h, w), device=self.device, dtype=torch.float32)
|
||||
return {"predicted_depth": dummy_tensor}
|
||||
|
||||
def eval(self):
|
||||
"""Compatibility method for eval mode."""
|
||||
if hasattr(self.model, 'eval'):
|
||||
self.model.eval()
|
||||
return self
|
||||
|
||||
def to(self, device):
|
||||
"""Move model to specified device."""
|
||||
if hasattr(self.model, 'to'):
|
||||
self.model = self.model.to(device)
|
||||
self.device = device
|
||||
return self
|
||||
|
||||
|
||||
def get_available_models():
|
||||
"""
|
||||
Returns a list of available depth models based on installed dependencies.
|
||||
|
||||
DA3 models are only included if depth_anything_v3 package is installed.
|
||||
"""
|
||||
available = []
|
||||
for model_name, model_info in DEPTH_MODELS.items():
|
||||
model_type = model_info.get("model_type", "v1") if isinstance(model_info, dict) else "v1"
|
||||
|
||||
# DA3 models require the depth_anything_v3 package
|
||||
if model_type == "v3" and not DA3_AVAILABLE:
|
||||
continue
|
||||
|
||||
available.append(model_name)
|
||||
|
||||
return available
|
||||
|
||||
|
||||
class DepthEstimationNode:
|
||||
"""
|
||||
ComfyUI node for depth estimation using Depth Anything models.
|
||||
@@ -607,11 +759,14 @@ class DepthEstimationNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
|
||||
"""Define the input types for the node."""
|
||||
"""Define the input types for the node.
|
||||
|
||||
Note: DA3 models are only shown if depth_anything_v3 package is installed.
|
||||
"""
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_name": (list(DEPTH_MODELS.keys()),),
|
||||
"model_name": (get_available_models(),),
|
||||
# Ensure minimum size is enforced by the UI
|
||||
"input_size": ("INT", {"default": 1024, "min": 256, "max": 1024, "step": 1}),
|
||||
"blur_radius": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
||||
@@ -689,9 +844,15 @@ class DepthEstimationNode:
|
||||
|
||||
# Prioritized fallback selection logic:
|
||||
# 1. Try to match on similar name
|
||||
# 2. Prefer V2 models if V2 was requested
|
||||
# 3. Prefer smaller models (more reliable)
|
||||
if "v2" in model_name_lower and "small" in model_name_lower:
|
||||
# 2. Prefer V3 models if V3 was requested (and DA3 is available)
|
||||
# 3. Prefer V2 models if V2 was requested
|
||||
# 4. Prefer smaller models (more reliable)
|
||||
if "v3" in model_name_lower and DA3_AVAILABLE:
|
||||
if "small" in model_name_lower:
|
||||
fallback_model = "Depth-Anything-V3-Small"
|
||||
else:
|
||||
fallback_model = "Depth-Anything-V3-Base"
|
||||
elif "v2" in model_name_lower and "small" in model_name_lower:
|
||||
fallback_model = "Depth-Anything-V2-Small"
|
||||
elif "v2" in model_name_lower and "base" in model_name_lower:
|
||||
fallback_model = "Depth-Anything-V2-Base"
|
||||
@@ -777,8 +938,8 @@ class DepthEstimationNode:
|
||||
# Use appropriate dtype based on device and model
|
||||
# FP16 for CUDA saves VRAM but doesn't work well for all models
|
||||
if 'cuda' in str(self.device) and not force_cpu:
|
||||
# V2 models have issues with FP16 - use FP32 for them
|
||||
if model_type == "v2":
|
||||
# V2 and V3 models have issues with FP16 - use FP32 for them
|
||||
if model_type in ("v2", "v3"):
|
||||
dtype = torch.float32
|
||||
else:
|
||||
# Other models can use FP16 to save VRAM
|
||||
@@ -786,6 +947,42 @@ class DepthEstimationNode:
|
||||
else:
|
||||
# CPU always uses FP32
|
||||
dtype = torch.float32
|
||||
|
||||
# Special handling for V3 (DA3) models - use DA3 API instead of transformers pipeline
|
||||
if model_type == "v3":
|
||||
if not DA3_AVAILABLE:
|
||||
raise RuntimeError(
|
||||
f"DA3 model '{model_name}' requested but depth_anything_v3 package not installed. "
|
||||
"Please install with: pip install depth-anything-v3"
|
||||
f"DA3 model '{model_name}' requested but depth_anything_3 package not installed. "
|
||||
"Please install with: pip install depth-anything-3"
|
||||
)
|
||||
|
||||
logger.info(f"Loading DA3 model: {model_name} using Depth Anything V3 API")
|
||||
try:
|
||||
# Load DA3 model from HuggingFace
|
||||
da3_model = DepthAnything3.from_pretrained(model_path)
|
||||
|
||||
# Move to appropriate device
|
||||
target_device = self.device if not force_cpu else 'cpu'
|
||||
da3_model = da3_model.to(target_device)
|
||||
|
||||
# Set to eval mode
|
||||
da3_model.eval()
|
||||
|
||||
# Wrap in DA3ModelWrapper for compatibility
|
||||
self.depth_estimator = DA3ModelWrapper(da3_model, self.device if not force_cpu else 'cpu')
|
||||
self.current_model = model_path
|
||||
logger.info(f"Successfully loaded DA3 model: {model_name}")
|
||||
return
|
||||
|
||||
except Exception as da3_error:
|
||||
logger.error(f"Failed to load DA3 model '{model_name}': {str(da3_error)}")
|
||||
logger.error(traceback.format_exc())
|
||||
|
||||
# Fall back to V2 if DA3 loading fails
|
||||
logger.warning(f"Falling back to Depth-Anything-V2-Small due to DA3 loading failure")
|
||||
return self.ensure_model_loaded("Depth-Anything-V2-Small", True, force_cpu)
|
||||
|
||||
# Create model-specific cache directory
|
||||
# Use consistent naming to improve cache hits
|
||||
|
||||
+6
-2
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyuidepthestimation"
|
||||
description = "A robust custom depth estimation node for ComfyUI using Depth-Anything models. It integrates depth estimation with configurable post-processing options including blur, median filtering, contrast enhancement, and gamma correction."
|
||||
version = "1.1.4"
|
||||
description = "A robust custom depth estimation node for ComfyUI using Depth-Anything models (V1, V2, and V3/DA3). It integrates depth estimation with configurable post-processing options including blur, median filtering, contrast enhancement, and gamma correction."
|
||||
version = "1.2.0"
|
||||
license = { file = "LICENSE" }
|
||||
dependencies = [
|
||||
"transformers>=4.20.0",
|
||||
@@ -13,6 +13,10 @@ dependencies = [
|
||||
"wget>=3.2"
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
# DA3 (Depth Anything V3) support - Apache 2.0 licensed models
|
||||
da3 = ["depth-anything-3"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/Limbicnation/ComfyUIDepthEstimation"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
+6
-2
@@ -2,7 +2,7 @@
|
||||
# Note: These are minimum requirements. ComfyUI's environment may provide newer versions.
|
||||
|
||||
# Fix for protobuf errors with transformers
|
||||
protobuf==3.20.3
|
||||
protobuf>=3.20.3
|
||||
|
||||
# Core dependencies
|
||||
tokenizers>=0.13.3 # Pre-built version compatible with most platforms
|
||||
@@ -24,4 +24,8 @@ wget>=3.2 # For reliable model downloading
|
||||
# If you're installing this node directly, ensure torch>=2.0.0 is available
|
||||
|
||||
# Network dependencies
|
||||
requests>=2.27.0 # For model downloading
|
||||
requests>=2.27.0 # For model downloading
|
||||
|
||||
# Optional dependencies for DA3 (Depth Anything V3) support
|
||||
# Uncomment the following line to enable DA3 models (Apache 2.0 licensed)
|
||||
# depth-anything-3 @ git+https://github.com/ByteDance-Seed/Depth-Anything-3.git
|
||||
Reference in New Issue
Block a user