style: fix indentation in estimate_depth method
This commit is contained in:
+19
-11
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user