gem-x-motion-capture / gem /utils /hf_utils.py
cs686's picture
Deploy GEM-X ZeroGPU motion capture
49d36c0 verified
Raw
History Blame Contribute Delete
4.17 kB
# 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