File size: 4,116 Bytes
d9bb75c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This software may be used and distributed in accordance with
# the terms of the DINOv3 License Agreement.

from enum import Enum
from typing import Optional, Tuple

import torch
from dinov3.eval.dense.depth.models import build_depther
from urllib.parse import urlparse
from pathlib import Path

from .utils import DINOV3_BASE_URL
from .backbones import (
    Weights as BackboneWeights,
    dinov3_vitl16,
    dinov3_vit7b16,
    convert_path_or_url_to_url,
)


class DepthWeights(Enum):
    SYNTHMIX = "SYNTHMIX"


def _get_depth_range(dataset: DepthWeights):
    depth_ranges = {
        DepthWeights.SYNTHMIX: (0.001, 100.0),
    }
    return depth_ranges[dataset]


_DPT_HEAD_CONFIG_DICT = dict(
    use_backbone_norm=True,
    use_batchnorm=True,
    use_cls_token=False,
    n_output_channels=256,
    depth_weights=DepthWeights.SYNTHMIX,
    backbone_weights=BackboneWeights.LVD1689M,
)


def _get_out_layers(backbone_name):
    if "vitl" in backbone_name:
        return [4, 11, 17, 23]
    elif "vit7b" in backbone_name:
        return [9, 19, 29, 39]
    else:
        raise ValueError(f"Unrecognized backbone name {backbone_name}")


def _get_post_process_channels(backbone_name):
    if "vitl" in backbone_name:
        return [1024, 1024, 1024, 1024]
    elif "vit7b" in backbone_name:
        return [2048, 2048, 2048, 2048]


_BACKBONE_DICT = {
    "dinov3_vit7b16": dinov3_vit7b16,
    "dinov3_vitl16": dinov3_vitl16,
}


def _make_dinov3_dpt_depther(
    *,
    backbone_name: str = "dinov3_vit7b16",
    pretrained: bool = True,
    depther_weights: DepthWeights | str = DepthWeights.SYNTHMIX,
    backbone_weights: BackboneWeights | str = BackboneWeights.LVD1689M,
    depth_range: Optional[Tuple[float, float]] = None,
    check_hash: bool = False,
    backbone_dtype: torch.dtype = torch.float32,
    **kwargs,
):
    backbone: torch.nn.Module = _BACKBONE_DICT[backbone_name](
        pretrained=pretrained,
        weights=backbone_weights,
    )
    out_index = _get_out_layers(backbone_name)
    post_process_channels = _get_post_process_channels(backbone_name)

    depth_range = depth_range or _get_depth_range(_DPT_HEAD_CONFIG_DICT["depth_weights"])
    min_depth, max_depth = depth_range

    depther = build_depther(
        backbone,
        backbone_out_layers=out_index,
        n_output_channels=_DPT_HEAD_CONFIG_DICT["n_output_channels"],
        use_backbone_norm=_DPT_HEAD_CONFIG_DICT["use_backbone_norm"],
        use_batchnorm=_DPT_HEAD_CONFIG_DICT["use_batchnorm"],
        use_cls_token=_DPT_HEAD_CONFIG_DICT["use_cls_token"],
        head_type="dpt",
        encoder_dtype=backbone_dtype,
        min_depth=min_depth,
        max_depth=max_depth,
        # DPTHead args
        channels=512,
        post_process_channels=post_process_channels,
        **kwargs,
    )

    if pretrained:
        if isinstance(depther_weights, DepthWeights):
            assert depther_weights == DepthWeights.SYNTHMIX, f"Unsupported depther weights {depther_weights}"
            weights_name = depther_weights.value.lower()
            hash = kwargs["hash"] if "hash" in kwargs else "02040be1"
            url = DINOV3_BASE_URL + f"/{backbone_name}/{backbone_name}_{weights_name}_dpt_head-{hash}.pth"
        else:
            url = convert_path_or_url_to_url(depther_weights)
        checkpoint = torch.hub.load_state_dict_from_url(url, map_location="cpu", check_hash=check_hash)
        depther[0].decoder.load_state_dict(checkpoint, strict=True)
    return depther


def dinov3_vit7b16_dd(
    *,
    pretrained: bool = True,
    weights: DepthWeights | str = DepthWeights.SYNTHMIX,
    backbone_weights: BackboneWeights | str = BackboneWeights.LVD1689M,
    check_hash: bool = False,
    backbone_dtype: torch.dtype = torch.float32,
    **kwargs,
):
    return _make_dinov3_dpt_depther(
        backbone_name="dinov3_vit7b16",
        pretrained=pretrained,
        depther_weights=weights,
        backbone_weights=backbone_weights,
        check_hash=check_hash,
        backbone_dtype=backbone_dtype,
        **kwargs,
    )