File size: 3,244 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
"""Load model-only checkpoints written in the public PXDepth format."""

from copy import deepcopy
from pathlib import Path
from typing import Any, Dict, IO, Optional, Type, TypeVar, Union

import torch
import torch.nn as nn
from huggingface_hub import hf_hub_download


ModelT = TypeVar("ModelT", bound=nn.Module)


def _merge(base: Dict[str, Any], update: Dict[str, Any]) -> Dict[str, Any]:
    """Recursively merge model-constructor overrides into a copied config.

    Args:
        base: Original nested model configuration. The dictionary is not
            modified.
        update: User overrides. Nested dictionaries update individual fields,
            while lists and scalar values replace their counterparts.

    Returns:
        A new merged dictionary suitable for constructing the model.
    """
    result = deepcopy(base)
    for key, value in update.items():
        if isinstance(result.get(key), dict) and isinstance(value, dict):
            result[key] = _merge(result[key], value)
        else:
            result[key] = deepcopy(value)
    return result


def load_pretrained(
    model_class: Type[ModelT],
    path_or_repo: Union[str, Path, IO[bytes]],
    model_kwargs: Optional[Dict[str, Any]] = None,
    strict: bool = True,
    **hf_kwargs: Any,
) -> ModelT:
    """Construct a model from a local or Hugging Face ``model.pt`` file.

    Args:
        model_class: Model class whose constructor accepts the canonical public
            configuration.
        path_or_repo: Existing local checkpoint path, binary file object, or
            Hugging Face model repository identifier.
        model_kwargs: Optional nested constructor overrides. Nested encoder or
            predictor fields are merged without discarding sibling settings.
        strict: Forwarded to :meth:`torch.nn.Module.load_state_dict`. Published
            checkpoints should normally use the default strict loading.
        **hf_kwargs: Extra arguments forwarded to ``hf_hub_download`` when
            ``path_or_repo`` is a repository identifier.

    Returns:
        An initialized model instance on CPU.
    """
    path = Path(path_or_repo) if isinstance(path_or_repo, (str, Path)) else None
    if path is not None and path.exists():
        checkpoint_path: Union[Path, IO[bytes]] = path
    elif isinstance(path_or_repo, str):
        checkpoint_path = Path(
            hf_hub_download(path_or_repo, repo_type="model", filename="model.pt", **hf_kwargs)
        )
    else:
        checkpoint_path = path_or_repo

    checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
    if "model_config" not in checkpoint or "model" not in checkpoint:
        raise ValueError("PXDepth checkpoints must contain 'model_config' and 'model'.")
    model_config = deepcopy(checkpoint["model_config"])
    model_type = model_config.pop("type", None)
    if model_type != model_class.__name__:
        raise ValueError(
            f"Expected a {model_class.__name__} checkpoint, got model type {model_type!r}."
        )
    if model_kwargs:
        model_config = _merge(model_config, model_kwargs)
    model = model_class(**model_config)
    model.load_state_dict(checkpoint["model"], strict=strict)
    return model