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.
This commit is contained in:
limbicnation
2025-02-27 12:11:34 +01:00
parent dc3b1572ec
commit b9a35b42e3
+74 -13
View File
@@ -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()