File size: 11,435 Bytes
87608ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
"""Core PXDepth architecture and stable public model API.

The module connects the Global Context Encoder to the Pixel-Space Depth
Predictor and defines raw forward computation. Checkpoint translation and
metric-scale inference live in focused helper modules, while their familiar
``from_pretrained`` and ``infer`` entry points remain methods on this class.
"""

from pathlib import Path
from typing import Any, Dict, IO, Optional, Union

import torch
import torch.nn as nn
import torch.nn.functional as F
from .Global_Context_Encoder import GlobalContextEncoder
from .Pixel_Space_Depth_Predictor import PixelSpaceDepthPredictor
from .checkpoint import load_pretrained
from .inference import infer as infer_model
from .precision import full_precision, inference_dtype, reduced_precision
from ..registry import ENCODERS, MODELS, PREDICTORS


@MODELS.register()
class PXDepth(nn.Module):
    """Complete PXDepth monocular depth model.

    A Global Context Encoder extracts semantic patch features and a Pixel-Space
    Depth Predictor estimates full-resolution normalized log-depth together
    with a finite-depth probability. ``forward`` exposes raw network outputs,
    while ``infer`` aligns them to a GT or MoGe-2 reference for metric-scale
    visualization and point-cloud reconstruction.
    """

    def __init__(
        self,
        encoder: Union[nn.Module, Dict[str, Any]],
        predictor: Union[nn.Module, Dict[str, Any]],
        remap_output: str = "linear",
        mask_threshold: float = 0.5,
    ) -> None:
        """Construct the encoder and CM-PiT pixel predictor.

        Args:
            encoder: Encoder module or registry config.
            predictor: Pixel predictor module or registry config. Its context
                patch size and channel width default to the encoder contract.
            remap_output: Output remapping applied to normalized log-depth.
                The released model uses ``'linear'``.
            mask_threshold: Probability threshold used by :meth:`infer`.

        Returns:
            ``None``. Model modules and ImageNet normalization buffers are
            registered on the instance.
        """
        super().__init__()
        if remap_output not in {"linear", "elu"}:
            raise ValueError(f"Unsupported remap_output: {remap_output}")

        self.remap_output = remap_output
        self.mask_threshold = float(mask_threshold)
        if isinstance(encoder, nn.Module):
            self.encoder = encoder
        else:
            encoder_config = dict(encoder)
            encoder_config.setdefault("type", "GlobalContextEncoder")
            self.encoder = ENCODERS.build(encoder_config)
        if not hasattr(self.encoder, "patch_size"):
            raise TypeError("The encoder must expose an integer patch_size attribute.")
        self.patch_size = self.encoder.patch_size
        self.p_enc = self.patch_size
        dim_ctx = getattr(self.encoder, "dim_out", None)
        if isinstance(predictor, nn.Module):
            self.predictor = predictor
        else:
            predictor_config = dict(predictor)
            predictor_config.setdefault("type", "PixelSpaceDepthPredictor")
            predictor_config.setdefault("in_channels", 3)
            predictor_config.setdefault("ctx_patch_size", self.patch_size)
            if dim_ctx is not None:
                predictor_config.setdefault("dim_ctx", int(dim_ctx))
            self.predictor = PREDICTORS.build(predictor_config)
        self.register_buffer("image_mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))
        self.register_buffer("image_std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))
        self._reference_model: Optional[nn.Module] = None

    @property
    def device(self) -> torch.device:
        """Return the device hosting PXDepth learnable parameters.

        No inputs are required. The value is inferred from the model's first
        parameter and is used when moving inference inputs or reference models.

        Returns:
            ``torch.device`` for the current model placement.
        """
        return next(self.parameters()).device

    @property
    def dtype(self) -> torch.dtype:
        """Return the storage dtype of PXDepth learnable parameters.

        No inputs are required. This reports parameter storage, which is
        independent from local autocast contexts used inside attention.

        Returns:
            ``torch.dtype`` of the model's first parameter.
        """
        return next(self.parameters()).dtype

    @classmethod
    def from_pretrained(
        cls,
        path_or_repo: Union[str, Path, IO[bytes]],
        model_kwargs: Optional[Dict[str, Any]] = None,
        strict: bool = True,
        **hf_kwargs: Any,
    ) -> "PXDepth":
        """Create a model from a local or Hugging Face ``model.pt`` checkpoint.

        Args:
            path_or_repo: Local checkpoint path, binary file object, or Hugging
                Face model repository identifier.
            model_kwargs: Optional constructor overrides applied after reading
                ``model_config`` from the checkpoint.
            strict: Forwarded to ``load_state_dict``. Published checkpoints
                should use the default exact matching.
            **hf_kwargs: Additional keyword arguments forwarded to
                ``huggingface_hub.hf_hub_download`` for remote repositories.

        Returns:
            Initialized :class:`PXDepth` instance on CPU.
        """
        return load_pretrained(
            cls,
            path_or_repo,
            model_kwargs=model_kwargs,
            strict=strict,
            **hf_kwargs,
        )

    def init_weights(self) -> None:
        """Initialize the Global Context Encoder from official DINOv2 weights.

        Predictor parameters retain the initialization created by their own
        module constructors.

        Returns:
            ``None``. Encoder parameters are updated in place.
        """
        self.encoder.init_weights()

    def enable_gradient_checkpointing(self) -> None:
        """Enable activation checkpointing in both encoder and predictor.

        This reduces saved activation memory during backward at the cost of
        recomputing transformer blocks.

        Returns:
            ``None``. Child module runtime behavior is updated in place.
        """
        self.encoder.enable_gradient_checkpointing()
        self.predictor.enable_gradient_checkpointing()

    def enable_pytorch_native_sdpa(self) -> None:
        """Enable the optimized SDPA attention path in the DINOv2 backbone.

        Decoder CM-PiT attention already uses PyTorch SDPA directly and is not
        modified by this method.

        Returns:
            ``None``. Encoder attention modules are wrapped in place.
        """
        self.encoder.enable_pytorch_native_sdpa()

    def _remap(self, depth: torch.Tensor) -> torch.Tensor:
        """Apply the configured output activation to raw depth predictions.

        Args:
            depth: Raw normalized-depth tensor with arbitrary batch/spatial
                shape, normally ``[B, H, W]``.

        Returns:
            Tensor with the same shape. The released ``linear`` setting returns
            the input unchanged.
        """
        return F.elu(depth) if self.remap_output == "elu" else depth

    def forward(
        self,
        image: torch.Tensor,
        use_fp16: bool = False,
        use_fp32: bool = False,
    ) -> Dict[str, torch.Tensor]:
        """Run the network without metric-scale alignment.

        Args:
            image: RGB tensor ``[B, 3, H, W]`` with values in ``[0, 1]``. ``H``
                and ``W`` must be divisible by the encoder patch size.
            use_fp16: Run attention-heavy encoder and predictor regions under
                FP16 autocast.
            use_fp32: Disable reduced-precision autocast. It is mutually
                exclusive with ``use_fp16``.

        Returns:
            Dictionary with normalized log-depth ``depth`` and finite-depth
            probability ``mask``, both FP32 tensors ``[B, H, W]``.
        """
        height, width = image.shape[-2:]
        if height % self.patch_size or width % self.patch_size:
            raise ValueError(
                f"Input resolution ({height}, {width}) must be divisible by patch size {self.patch_size}"
            )
        dtype = inference_dtype(use_fp16=use_fp16, use_fp32=use_fp32)

        with full_precision(image.device):
            image_norm = (image.float() - self.image_mean.float()) / self.image_std.float()
        with reduced_precision(image.device, dtype):
            context = self.encoder(image, height // self.patch_size, width // self.patch_size)
        context = context.flatten(2).permute(0, 2, 1).contiguous()
        depth, mask = self.predictor(image_norm, context, autocast_dtype=dtype)

        with full_precision(image.device):
            depth = self._remap(depth.float().squeeze(1))
            mask = mask.float().squeeze(1).sigmoid()
        return {"depth": depth, "mask": mask}

    def infer(
        self,
        image: torch.Tensor,
        gt_depth: Optional[torch.Tensor] = None,
        intrinsics: Optional[torch.Tensor] = None,
        fov_x: Optional[Union[float, torch.Tensor]] = None,
        ref_image: Optional[torch.Tensor] = None,
        apply_mask: bool = True,
        use_fp16: bool = True,
        use_fp32: bool = False,
    ) -> Dict[str, torch.Tensor]:
        """Recover metric-scale depth, validity, intrinsics, and 3D points.

        Raw normalized log-depth is affine-aligned in log space to ``gt_depth``
        when supplied, otherwise to a lazily loaded MoGe-2 reference. Alignment
        parameters are estimated on a 64x64 nearest-resized valid subset. The
        aligned depth is exponentiated and back-projected with normalized camera
        intrinsics.

        Args:
            image: RGB tensor ``[3,H,W]`` or batch ``[B,3,H,W]`` in ``[0,1]``.
            gt_depth: Optional reference depth ``[H,W]`` or ``[B,H,W]``. Finite
                positive pixels define log-space alignment.
            intrinsics: Optional normalized camera matrices ``[3,3]`` or
                ``[B,3,3]`` corresponding to ``gt_depth``.
            fov_x: Optional horizontal field of view in degrees, scalar or
                tensor ``[B]``, used when intrinsics are unavailable.
            ref_image: Optional original-resolution RGB tensor used only by the
                reference model; PXDepth still consumes ``image``.
            apply_mask: Replace invalid predicted depth/points with infinity.
            use_fp16: Use FP16 for attention-heavy model regions.
            use_fp32: Force those regions to FP32 and override the BF16 default.

        Returns:
            Dictionary containing aligned ``depth`` ``[B,H,W]``, boolean
            ``mask`` ``[B,H,W]``, point map ``points`` ``[B,H,W,3]``, normalized
            ``intrinsics`` ``[B,3,3]``, and horizontal ``fov_x`` ``[B]``. For an
            unbatched input, the leading batch dimension is removed.
        """
        return infer_model(
            self,
            image,
            gt_depth=gt_depth,
            intrinsics=intrinsics,
            fov_x=fov_x,
            ref_image=ref_image,
            apply_mask=apply_mask,
            use_fp16=use_fp16,
            use_fp32=use_fp32,
        )