depth_pro / app.py
ashwmurt's picture
Upload depth_pro recipe (v1)
2b156ca verified
Raw
History Blame Contribute Delete
4.8 kB
# ---------------------------------------------------------------------
# Copyright (c) 2026 Qualcomm Technologies, Inc. and/or its subsidiaries.
# SPDX-License-Identifier: BSD-3-Clause
# ---------------------------------------------------------------------
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, cast
import matplotlib.pyplot as plt
import numpy as np
import numpy.typing as npt
import torch
from PIL import Image
from torchvision import transforms
from qai_hub_models.utils.image_processing import pil_resize_pad, undo_resize_pad
@dataclass
class DepthProPrediction:
"""Structured output of :class:`DepthProApp`.
``depth`` is metric depth in scene units (aligned via ``focal_length_px``);
``heatmap`` is a plasma-colored visualization of inverse depth for display;
``field_of_view`` is horizontal FoV in degrees; ``focal_length_px`` is
the pixel focal length derived from FoV and the original image width.
"""
depth: npt.NDArray[np.float32]
heatmap: Image.Image
field_of_view: float
focal_length_px: float
class DepthProApp:
"""End-to-end app for Apple DepthPro depth estimation.
Wraps a callable returning ``(predicted_depth, field_of_view)`` — either
the torch model or an on-device runner. Preprocessing resizes with
aspect-preserving padding to the network's 1536x1536 input; post-
processing mirrors HuggingFace's
``DepthProImageProcessorFast.post_process_depth_estimation`` (metric
scaling by ``width / focal_length_px``, then inversion of the canonical
inverse depth).
"""
def __init__(
self,
model: Callable[
[torch.Tensor], tuple[torch.Tensor, torch.Tensor]
],
input_height: int | None = None,
input_width: int | None = None,
) -> None:
self.model = model
if input_height is None or input_width is None:
get_input_spec = getattr(model, "get_input_spec", None)
if get_input_spec is None:
raise TypeError(
"DepthProApp needs input_height and input_width when the "
"provided model is not a BaseModel (has no get_input_spec)."
)
_, _, h, w = get_input_spec()["image"][0]
input_height = input_height if input_height is not None else h
input_width = input_width if input_width is not None else w
self.input_height = input_height
self.input_width = input_width
def predict(self, *args: Any, **kwargs: Any) -> DepthProPrediction:
return self.estimate_depth(*args, **kwargs)
def estimate_depth(self, image: Image.Image) -> DepthProPrediction:
"""Estimate depth, FoV, and focal length for a single image.
Parameters
----------
image
PIL image in any resolution / aspect ratio.
"""
resized_image, scale, padding = pil_resize_pad(
image, (self.input_height, self.input_width)
)
image_tensor = transforms.ToTensor()(resized_image).unsqueeze(0)
predicted_depth, field_of_view = self.model(image_tensor)
# Horizontal FoV -> focal length in pixels of the *original* image.
# Matches HF's post_process_depth_estimation.
orig_width = float(image.size[0])
fov_deg = field_of_view.detach().float().view(-1)
focal_length_px = 0.5 * orig_width / torch.tan(
0.5 * torch.deg2rad(fov_deg)
)
# Metric scaling of canonical inverse depth (again, from HF).
depth_scaled = predicted_depth * (
orig_width / focal_length_px
).view(-1, 1, 1)
# (B, 1, H, W) shape is what undo_resize_pad expects.
depth_map = undo_resize_pad(
depth_scaled.unsqueeze(1), image.size, scale, padding
)
# Canonical inverse depth -> metric depth. Clamp mirrors HF.
depth_map = 1.0 / torch.clamp(depth_map, min=1e-4, max=1e4)
depth_np = cast(
npt.NDArray[np.float32],
depth_map.squeeze().detach().cpu().numpy().astype(np.float32),
)
# Visualize inverse depth so closer objects appear brighter, matching
# the shared depth-estimation demo convention.
inv = 1.0 / np.maximum(depth_np, 1e-6)
inv_norm = inv / max(inv.max(), 1e-6)
heatmap = plt.cm.get_cmap("plasma")(inv_norm)[..., :3]
heatmap_image = Image.fromarray((heatmap * 255).astype(np.uint8))
return DepthProPrediction(
depth=depth_np,
heatmap=heatmap_image,
field_of_view=float(fov_deg.squeeze().item()),
focal_length_px=float(focal_length_px.squeeze().item()),
)