"""Download model checkpoints from yagmurakarken/mmdiff at Space startup.""" import os from huggingface_hub import hf_hub_download CHECKPOINT_DIR = "/tmp/checkpoints" MODEL_REPO = "yagmurakarken/mmdiff" CHECKPOINT_FILES = ["duts_saliency.ckpt", "nyu_depth.ckpt", "pascal_segmentation.ckpt"] def download_all(): os.makedirs(CHECKPOINT_DIR, exist_ok=True) token = os.environ.get("HF_TOKEN") for fname in CHECKPOINT_FILES: path = os.path.join(CHECKPOINT_DIR, fname) if not os.path.exists(path): print(f"[DOWNLOAD] {fname} ...") downloaded = hf_hub_download( repo_id=MODEL_REPO, filename=fname, repo_type="model", local_dir=CHECKPOINT_DIR, token=token, ) print(f"[DOWNLOAD] {fname} -> {downloaded}") else: print(f"[DOWNLOAD] {fname} already exists at {path}") if __name__ == "__main__": download_all()