diff --git a/modules/control/proc/depth_pro/__init__.py b/modules/control/proc/depth_pro/__init__.py index 43c2458a0..ac2075632 100644 --- a/modules/control/proc/depth_pro/__init__.py +++ b/modules/control/proc/depth_pro/__init__.py @@ -1,6 +1,7 @@ import cv2 -import numpy as np import torch +import torch.nn.functional as F +import numpy as np from PIL import Image from modules import devices, masking @@ -8,88 +9,54 @@ from modules.shared import opts class DepthProDetector: - """Wrapper around Apple's DepthPro depth estimation model.""" + """Apple DepthPro detector (aligned with Depth Anything style).""" def __init__(self, model, processor): self.model = model self.processor = processor @classmethod - def from_pretrained(cls, pretrained_model_or_path: str, cache_dir: str, use_fast_processor: bool = False, **kwargs): + def from_pretrained(cls, pretrained_model_or_path: str = "apple/DepthPro-hf", cache_dir: str | None = None) -> "DepthProDetector": from transformers import AutoImageProcessor, DepthProForDepthEstimation - processor_kwargs = {"cache_dir": cache_dir} - processor_kwargs.update(kwargs) - if use_fast_processor: - from transformers.models.depth_pro.image_processing_depth_pro_fast import DepthProImageProcessorFast - - processor = DepthProImageProcessorFast.from_pretrained( - pretrained_model_or_path, - **processor_kwargs, - ) - else: - processor = AutoImageProcessor.from_pretrained( - pretrained_model_or_path, - **processor_kwargs, - ) - + processor = AutoImageProcessor.from_pretrained(pretrained_model_or_path, cache_dir=cache_dir) model = DepthProForDepthEstimation.from_pretrained( pretrained_model_or_path, cache_dir=cache_dir, - ) - model = model.to(device=devices.device).eval() + ).to(devices.device).eval() return cls(model, processor) - def _prepare_inputs(self, image: Image.Image) -> dict: - inputs = self.processor(images=image, return_tensors="pt") - tensor_inputs = {} - for key, value in inputs.items(): - if isinstance(value, torch.Tensor): - tensor_inputs[key] = value.to(device=devices.device) - else: - tensor_inputs[key] = value - return tensor_inputs + def __call__(self, image, color_map: str = "none", output_type: str = "pil"): + self.model.to(devices.device) + if isinstance(image, Image.Image): + image = np.array(image) + h, w = image.shape[:2] + image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) + pil_image = Image.fromarray(image_rgb) - def __call__( - self, - image, - color_map: str = "inferno", - output_type: str = "pil", - ): - if isinstance(image, list): - image = image[0] - if image is None: - return image - if not isinstance(image, Image.Image): - image = Image.fromarray(np.array(image)) + inputs = self.processor(images=pil_image, return_tensors="pt") + inputs = {k: v.to(devices.device) if isinstance(v, torch.Tensor) else v for k, v in inputs.items()} - original_size = (image.height, image.width) - inputs = self._prepare_inputs(image) with devices.inference_context(): outputs = self.model(**inputs) - results = self.processor.post_process_depth_estimation(outputs, target_sizes=[original_size]) - depth_tensor = results[0]["predicted_depth"].to(torch.float32) + results = self.processor.post_process_depth_estimation(outputs, target_sizes=[(h, w)]) + depth_tensor = results[0]["predicted_depth"].to(devices.device, dtype=torch.float32) + if opts.control_move_processor: self.model.to("cpu") - # Invert to align with other depth processors that render near as bright + depth_tensor = F.interpolate(depth_tensor[None, None], size=(h, w), mode="bilinear", align_corners=False)[0, 0] depth_tensor = 1.0 / torch.clamp(depth_tensor, min=1e-6) depth_tensor -= depth_tensor.min() - max_val = depth_tensor.max() - if max_val > 0: - depth_tensor /= max_val - depth_tensor = (depth_tensor * 255.0).clamp(0, 255).to(torch.uint8) - depth = depth_tensor.cpu().numpy() - - if color_map and color_map.lower() != "none": - color = color_map.lower() - if color not in masking.COLORMAP: - color = "inferno" - processed = cv2.applyColorMap(depth, masking.COLORMAP.index(color))[:, :, ::-1] - else: - processed = depth + depth_max = depth_tensor.max() + if depth_max > 0: + depth_tensor /= depth_max + depth = (depth_tensor * 255.0).clamp(0, 255).to(torch.uint8).cpu().numpy() + if color_map != "none": + colormap_key = color_map if color_map in masking.COLORMAP else "inferno" + depth = cv2.applyColorMap(depth, masking.COLORMAP.index(colormap_key))[:, :, ::-1] if output_type == "pil": - mode = "RGB" if processed.ndim == 3 else "L" - processed = Image.fromarray(processed, mode=mode) - return processed + mode = "RGB" if depth.ndim == 3 else "L" + depth = Image.fromarray(depth, mode=mode) + return depth