Merge pull request #11 from Limbicnation/fix/depth-model-loading-fallback
fix: add robust fallback for depth model loading
This commit is contained in:
+74
-13
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user