Commit ·
01055ed
1
Parent(s): 3cbe3bf
Smarter download_checkpoints: skip download if /data/Lyra-2 already populated (direct model mount)
Browse files- download_checkpoints.py +65 -54
download_checkpoints.py
CHANGED
|
@@ -1,9 +1,10 @@
|
|
| 1 |
-
"""
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
|
|
|
| 7 |
"""
|
| 8 |
|
| 9 |
from __future__ import annotations
|
|
@@ -20,73 +21,83 @@ log = logging.getLogger("lyra2-space.checkpoints")
|
|
| 20 |
logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(message)s")
|
| 21 |
|
| 22 |
REPO_ID = "nvidia/Lyra-2.0"
|
| 23 |
-
|
| 24 |
-
# /data is the HF Spaces persistent volume mount point on paid GPU Spaces.
|
| 25 |
-
# Fall back to a local directory if /data is unavailable (e.g. local dev runs).
|
| 26 |
DATA_ROOT = Path(os.environ.get("LYRA2_DATA_ROOT", "/data"))
|
| 27 |
-
CHECKPOINTS_ROOT = DATA_ROOT / "Lyra-2"
|
| 28 |
-
MODEL_MARKER = CHECKPOINTS_ROOT / "checkpoints" / "model"
|
| 29 |
-
|
| 30 |
APP_ROOT = Path("/home/user/app")
|
| 31 |
-
|
| 32 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
|
| 34 |
|
| 35 |
-
def
|
| 36 |
try:
|
| 37 |
-
|
| 38 |
-
probe =
|
| 39 |
probe.write_text("ok")
|
| 40 |
probe.unlink()
|
| 41 |
-
return
|
| 42 |
-
except OSError
|
| 43 |
-
|
| 44 |
-
# to an in-container path; note this means checkpoints re-download on
|
| 45 |
-
# every cold start. Good for local dev, bad for prod.
|
| 46 |
-
fallback = Path.home() / "app" / "_checkpoints_cache"
|
| 47 |
-
log.warning("Cannot use %s (%s); falling back to %s", DATA_ROOT, exc, fallback)
|
| 48 |
-
fallback.mkdir(parents=True, exist_ok=True)
|
| 49 |
-
return fallback
|
| 50 |
|
| 51 |
|
| 52 |
-
def
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
REPO_ID, lyra2_root)
|
| 63 |
snapshot_download(
|
| 64 |
repo_id=REPO_ID,
|
| 65 |
allow_patterns=["checkpoints/*"],
|
| 66 |
local_dir=str(lyra2_root),
|
| 67 |
-
local_dir_use_symlinks=False,
|
| 68 |
-
# Resume partial downloads across container restarts.
|
| 69 |
max_workers=8,
|
| 70 |
)
|
| 71 |
log.info("Checkpoint download complete.")
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
#
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
|
|
|
| 90 |
return 0
|
| 91 |
|
| 92 |
|
|
|
|
| 1 |
+
"""Verify or download Lyra-2.0 checkpoints.
|
| 2 |
|
| 3 |
+
Three cases handled, in priority order:
|
| 4 |
+
1. /data/Lyra-2/checkpoints/model is already populated (read-only direct-mount
|
| 5 |
+
OR warm bucket): skip download, just symlink.
|
| 6 |
+
2. /data is writable + empty: download from HF Hub.
|
| 7 |
+
3. /data is unwritable + empty: fall back to ephemeral cache.
|
| 8 |
"""
|
| 9 |
|
| 10 |
from __future__ import annotations
|
|
|
|
| 21 |
logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(message)s")
|
| 22 |
|
| 23 |
REPO_ID = "nvidia/Lyra-2.0"
|
|
|
|
|
|
|
|
|
|
| 24 |
DATA_ROOT = Path(os.environ.get("LYRA2_DATA_ROOT", "/data"))
|
|
|
|
|
|
|
|
|
|
| 25 |
APP_ROOT = Path("/home/user/app")
|
| 26 |
+
APP_CHECKPOINTS = APP_ROOT / "Lyra-2" / "checkpoints"
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def _already_populated(root: Path) -> bool:
|
| 30 |
+
"""True if root/Lyra-2/checkpoints/model has content."""
|
| 31 |
+
marker = root / "Lyra-2" / "checkpoints" / "model"
|
| 32 |
+
return marker.is_dir() and any(marker.iterdir())
|
| 33 |
|
| 34 |
|
| 35 |
+
def _is_writable(root: Path) -> bool:
|
| 36 |
try:
|
| 37 |
+
root.mkdir(parents=True, exist_ok=True)
|
| 38 |
+
probe = root / ".writable_probe"
|
| 39 |
probe.write_text("ok")
|
| 40 |
probe.unlink()
|
| 41 |
+
return True
|
| 42 |
+
except OSError:
|
| 43 |
+
return False
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 44 |
|
| 45 |
|
| 46 |
+
def _symlink_app_checkpoints(lyra2_root: Path) -> None:
|
| 47 |
+
"""Point /home/user/app/Lyra-2/checkpoints at <lyra2_root>/checkpoints."""
|
| 48 |
+
src = lyra2_root / "checkpoints"
|
| 49 |
+
dst = APP_CHECKPOINTS
|
| 50 |
+
if dst.is_symlink() and dst.resolve() == src.resolve():
|
| 51 |
+
log.info("Checkpoints symlink already points to %s", src)
|
| 52 |
+
return
|
| 53 |
+
if dst.is_symlink() or dst.is_file():
|
| 54 |
+
dst.unlink()
|
| 55 |
+
elif dst.exists():
|
| 56 |
+
shutil.rmtree(dst)
|
| 57 |
+
dst.parent.mkdir(parents=True, exist_ok=True)
|
| 58 |
+
dst.symlink_to(src, target_is_directory=True)
|
| 59 |
+
log.info("Linked %s -> %s", dst, src)
|
| 60 |
+
|
| 61 |
|
| 62 |
+
def main() -> int:
|
| 63 |
+
# Case 1: /data already has the checkpoints (direct-mount OR warm bucket)
|
| 64 |
+
if DATA_ROOT.is_dir() and _already_populated(DATA_ROOT):
|
| 65 |
+
log.info("Checkpoints present at %s/Lyra-2 — skipping download.", DATA_ROOT)
|
| 66 |
+
_symlink_app_checkpoints(DATA_ROOT / "Lyra-2")
|
| 67 |
+
return 0
|
| 68 |
+
|
| 69 |
+
# Case 2: /data writable but empty — download
|
| 70 |
+
if _is_writable(DATA_ROOT):
|
| 71 |
+
lyra2_root = DATA_ROOT / "Lyra-2"
|
| 72 |
+
lyra2_root.mkdir(parents=True, exist_ok=True)
|
| 73 |
+
log.info("Downloading %s checkpoints/* to %s (first cold start only).",
|
| 74 |
REPO_ID, lyra2_root)
|
| 75 |
snapshot_download(
|
| 76 |
repo_id=REPO_ID,
|
| 77 |
allow_patterns=["checkpoints/*"],
|
| 78 |
local_dir=str(lyra2_root),
|
|
|
|
|
|
|
| 79 |
max_workers=8,
|
| 80 |
)
|
| 81 |
log.info("Checkpoint download complete.")
|
| 82 |
+
_symlink_app_checkpoints(lyra2_root)
|
| 83 |
+
return 0
|
| 84 |
+
|
| 85 |
+
# Case 3: /data unwritable + empty — ephemeral fallback
|
| 86 |
+
fallback = APP_ROOT / "_checkpoints_cache"
|
| 87 |
+
log.warning("Cannot use %s (read-only AND empty). Falling back to ephemeral %s",
|
| 88 |
+
DATA_ROOT, fallback)
|
| 89 |
+
fallback.mkdir(parents=True, exist_ok=True)
|
| 90 |
+
lyra2_root = fallback / "Lyra-2"
|
| 91 |
+
lyra2_root.mkdir(parents=True, exist_ok=True)
|
| 92 |
+
if not _already_populated(fallback):
|
| 93 |
+
log.info("Downloading %s checkpoints/* to %s.", REPO_ID, lyra2_root)
|
| 94 |
+
snapshot_download(
|
| 95 |
+
repo_id=REPO_ID,
|
| 96 |
+
allow_patterns=["checkpoints/*"],
|
| 97 |
+
local_dir=str(lyra2_root),
|
| 98 |
+
max_workers=8,
|
| 99 |
+
)
|
| 100 |
+
_symlink_app_checkpoints(lyra2_root)
|
| 101 |
return 0
|
| 102 |
|
| 103 |
|