From b9a35b42e3c665a3fa66b0079a0a9d016b1b83d6 Mon Sep 17 00:00:00 2001 From: limbicnation Date: Thu, 27 Feb 2025 12:11:34 +0100 Subject: [PATCH] fix: add robust fallback for depth model loading This change implements multiple fallback paths for depth model loading, handles both online and offline scenarios, and provides clear error messages for troubleshooting. It maintains backward compatibility while addressing the 'model not found' error. --- depth_estimation_node.py | 87 ++++++++++++++++++++++++++++++++++------ 1 file changed, 74 insertions(+), 13 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index ea3c027..745c4f1 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -91,13 +91,13 @@ class DepthEstimationNode: def ensure_model_loaded(self, model_name: str) -> None: """ - Ensures the correct model is loaded with proper VRAM management. + Ensures the correct model is loaded with proper VRAM management and fallback options. Args: model_name: The name of the model to load Raises: - RuntimeError: If the model fails to load + RuntimeError: If the model fails to load after all fallback attempts """ try: if model_name not in DEPTH_MODELS: @@ -105,9 +105,11 @@ class DepthEstimationNode: model_path = DEPTH_MODELS[model_name] + # Only reload if needed if 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() @@ -119,24 +121,83 @@ class DepthEstimationNode: # Use FP16 for CUDA devices to save VRAM 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") + # Create a dedicated cache directory for this model + cache_dir = os.path.join(MODELS_DIR, model_name.replace("-", "_").lower()) + os.makedirs(cache_dir, exist_ok=True) - self.depth_estimator = pipeline( - "depth-estimation", - model=model_path, - device_map=device_type, - torch_dtype=dtype - ) + # List of model paths to try (original and fallback) + model_paths_to_try = [ + model_path, # Original path + model_path + "-hf", # Try with -hf suffix + model_path.replace("depth-anything", "depth-anything-hf") # Alternative format + ] + + # Try each model path + success = False + last_error = None + + for path in model_paths_to_try: + try: + logger.info(f"Attempting to load from: {path}") + + # Try with online mode first + try: + self.depth_estimator = pipeline( + "depth-estimation", + model=path, + cache_dir=cache_dir, + local_files_only=False, # Try online first + device_map=device_type, + torch_dtype=dtype + ) + success = True + logger.info(f"Successfully loaded model from {path}") + break + except Exception as online_error: + logger.warning(f"Online loading failed for {path}: {str(online_error)}") + + # Try with local_files_only if online fails + try: + self.depth_estimator = pipeline( + "depth-estimation", + model=path, + cache_dir=cache_dir, + local_files_only=True, # Try local only as fallback + device_map=device_type, + torch_dtype=dtype + ) + success = True + logger.info(f"Successfully loaded model from local cache: {path}") + break + except Exception as local_error: + last_error = local_error + logger.warning(f"Local loading failed for {path}: {str(local_error)}") + continue + + except Exception as path_error: + last_error = path_error + logger.warning(f"Failed to load model from {path}: {str(path_error)}") + continue + + if not success: + # If all attempts failed, 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) + """ + logger.error(error_msg) + raise RuntimeError(error_msg) # 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}") except Exception as e: self.cleanup()