From 8c3e6123f3d9bd7ee9aca7ff213327c6d4fd0741 Mon Sep 17 00:00:00 2001 From: limbicnation Date: Fri, 29 Nov 2024 20:51:23 +0100 Subject: [PATCH] style: fix indentation in estimate_depth method --- depth_estimation_node.py | 30 +++++++++++++++++++----------- 1 file changed, 19 insertions(+), 11 deletions(-) diff --git a/depth_estimation_node.py b/depth_estimation_node.py index 6e26f15..a1dd9ba 100644 --- a/depth_estimation_node.py +++ b/depth_estimation_node.py @@ -15,7 +15,7 @@ DEPTH_MODELS = { } class DepthEstimationNode: - MEDIAN_SIZES = ["3", "5", "7", "9", "11"] # Valid median sizes + MEDIAN_SIZES = ["3", "5", "7", "9", "11"] def __init__(self): self.device = get_torch_device() @@ -40,7 +40,6 @@ class DepthEstimationNode: CATEGORY = "image/depth" def ensure_model_loaded(self, model_name): - """Ensure the depth estimation model is loaded.""" model_path = DEPTH_MODELS[model_name] if self.depth_estimator is None or self.current_model != model_path: try: @@ -55,9 +54,6 @@ class DepthEstimationNode: def estimate_depth(self, image, model_name, blur_radius=2.0, median_size="5", apply_auto_contrast=True, apply_gamma=True): - """ - Estimate depth from input image with optional post-processing. - """ try: # Validate median_size if median_size not in self.MEDIAN_SIZES: @@ -87,14 +83,24 @@ class DepthEstimationNode: # Get depth map depth_result = self.depth_estimator(pil_image) - depth_map = depth_result["predicted_depth"] - # Convert tensor to numpy if necessary + # Convert tensor to numpy and ensure correct dimensions if torch.is_tensor(depth_map): - depth_map = depth_map.cpu().numpy() + depth_map = depth_map.squeeze().cpu().numpy() + + # Ensure depth_map is 2D + if len(depth_map.shape) > 2: + depth_map = depth_map.squeeze() # Normalize depth values to 0-255 range - depth_map = ((depth_map - depth_map.min()) * (255 / (depth_map.max() - depth_map.min()))).astype(np.uint8) + depth_min = depth_map.min() + depth_max = depth_map.max() + if depth_max > depth_min: + depth_map = ((depth_map - depth_min) * (255.0 / (depth_max - depth_min))).astype(np.uint8) + else: + depth_map = np.zeros_like(depth_map, dtype=np.uint8) + + # Convert to PIL Image depth_map = Image.fromarray(depth_map) # Apply post-processing @@ -116,7 +122,10 @@ class DepthEstimationNode: # Convert back to tensor format depth_array = np.array(depth_map).astype(np.float32) / 255.0 - depth_tensor = depth_array[None, ..., None] # Add batch and channel dims + depth_tensor = torch.from_numpy(depth_array)[None, ..., None] # Convert to tensor and add batch and channel dims + + # Move tensor to the correct device + depth_tensor = depth_tensor.to(self.device) return (depth_tensor,) @@ -124,7 +133,6 @@ class DepthEstimationNode: raise RuntimeError(f"Depth estimation failed: {str(e)}") def gamma_correction(self, img, gamma=1.0): - """Apply gamma correction to the image.""" inv_gamma = 1.0 / gamma table = [((i / 255.0) ** inv_gamma) * 255 for i in range(256)] table = np.array(table, np.uint8)