File size: 4,171 Bytes
49d36c0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from pathlib import Path

HF_REPO_ID = "nvidia/GEM-X"

# GEM-SOMA checkpoint
DEFAULT_CKPT_FILENAME = "gem_soma.ckpt"
DEFAULT_CKPT_DIR = "inputs/pretrained"

# ViTPose checkpoint
VITPOSE_CKPT_FILENAME = "vitpose.pth"
VITPOSE_CKPT_DIR = "inputs/checkpoints/vitpose"

# SAM-3D-Body checkpoint
SAM3D_CKPT_FILENAME = "sam3d_body.ckpt"
SAM3D_CONFIG_FILENAME = "model_config.yaml"
SAM3D_CKPT_DIR = "inputs/checkpoints/sam-3d-body-dinov3"

# MHR model
MHR_MODEL_FILENAME = "mhr_model.pt"
MHR_MODEL_DIR = "inputs/mhr_data"

# SOMA scale data
SOMA_SCALE_MEAN_FILENAME = "scale_mean.pth"
SOMA_SCALE_COMPS_FILENAME = "scale_comps.pth"
SOMA_DATA_DIR = "inputs/soma_data"


def _download_hf_file(repo_id, filename, local_dir):
    """Download a single file from HuggingFace Hub if not already present."""
    local_path = Path(local_dir) / filename
    if local_path.exists():
        return str(local_path)
    from huggingface_hub import hf_hub_download

    return hf_hub_download(repo_id=repo_id, filename=filename, local_dir=local_dir)


def download_checkpoint(
    repo_id=HF_REPO_ID, filename=DEFAULT_CKPT_FILENAME, local_dir=DEFAULT_CKPT_DIR
):
    """Download GEM-SOMA checkpoint from HuggingFace Hub if not already cached."""
    return _download_hf_file(repo_id, filename, local_dir)


def download_vitpose_checkpoint(
    repo_id=HF_REPO_ID, filename=VITPOSE_CKPT_FILENAME, local_dir=VITPOSE_CKPT_DIR
):
    """Download ViTPose checkpoint from HuggingFace Hub if not already cached."""
    return _download_hf_file(repo_id, filename, local_dir)


def download_sam3d_checkpoint(repo_id=HF_REPO_ID, local_dir=SAM3D_CKPT_DIR):
    """Download SAM-3D-Body checkpoint and config from HuggingFace Hub.

    Returns the checkpoint path. The config (model_config.yaml) is downloaded
    alongside it so load_sam_3d_body() can find it in the same directory.
    """
    _download_hf_file(repo_id, SAM3D_CONFIG_FILENAME, local_dir)
    return _download_hf_file(repo_id, SAM3D_CKPT_FILENAME, local_dir)


def download_mhr_model(repo_id=HF_REPO_ID, filename=MHR_MODEL_FILENAME, local_dir=MHR_MODEL_DIR):
    """Download MHR model from HuggingFace Hub if not already cached."""
    return _download_hf_file(repo_id, filename, local_dir)


def download_soma_data(repo_id=HF_REPO_ID, local_dir=SOMA_DATA_DIR):
    """Download SOMA scale data (scale_mean.pth, scale_comps.pth) from HuggingFace Hub.

    Returns the directory path containing the downloaded files.
    """
    _download_hf_file(repo_id, SOMA_SCALE_MEAN_FILENAME, local_dir)
    _download_hf_file(repo_id, SOMA_SCALE_COMPS_FILENAME, local_dir)
    return local_dir


# ONNX models for fast demo
ONNX_DIR = "inputs/onnx"
ONNX_MODELS = {
    "vitpose": ("vitpose.onnx", "vitpose.onnx.data"),
    "gem_denoiser": ("gem_denoiser.onnx", "gem_denoiser.onnx.data"),
    "gem_denoiser_no_imgfeat": (
        "gem_denoiser_no_imgfeat.onnx",
        "gem_denoiser_no_imgfeat.onnx.data",
    ),
    "sam3db_backbone": ("sam3db_backbone.onnx", "sam3db_backbone.onnx.data"),
}


def download_onnx_model(name, repo_id=HF_REPO_ID, local_dir=ONNX_DIR):
    """Download an ONNX model (and its .data file) from HuggingFace Hub.

    Args:
        name: Model name key, one of: "vitpose", "gem_denoiser",
              "gem_denoiser_no_imgfeat", "sam3db_backbone".

    Returns:
        Path to the main .onnx file.
    """
    if name not in ONNX_MODELS:
        raise ValueError(f"Unknown ONNX model '{name}'. Available: {list(ONNX_MODELS)}")
    files = ONNX_MODELS[name]
    # Files live under onnx/ in the HF repo; download into the parent of
    # local_dir so that onnx/<file> maps to <local_dir>/<file>.
    parent = str(Path(local_dir).parent)
    for f in files:
        _download_hf_file(repo_id, f"onnx/{f}", parent)
    return str(Path(local_dir) / files[0])


def download_all_onnx(repo_id=HF_REPO_ID, local_dir=ONNX_DIR):
    """Download all ONNX models for the fast demo pipeline."""
    for name in ONNX_MODELS:
        download_onnx_model(name, repo_id, local_dir)
    return local_dir