cezar-hapiko commited on
Commit
01055ed
·
1 Parent(s): 3cbe3bf

Smarter download_checkpoints: skip download if /data/Lyra-2 already populated (direct model mount)

Browse files
Files changed (1) hide show
  1. download_checkpoints.py +65 -54
download_checkpoints.py CHANGED
@@ -1,9 +1,10 @@
1
- """Download Lyra-2 checkpoints from HuggingFace to the persistent /data volume.
2
 
3
- Run at container boot from entrypoint.sh. Idempotent: if the expected directory
4
- already exists and contains the model checkpoint, skip the download entirely.
5
- The ~50-80 GB download only happens on the first cold start after the paid
6
- Space's /data volume is provisioned — subsequent starts reuse the cached copy.
 
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
- APP_LYRA2 = APP_ROOT / "Lyra-2"
32
- APP_CHECKPOINTS = APP_LYRA2 / "checkpoints"
 
 
 
 
 
33
 
34
 
35
- def _ensure_writable_data_dir() -> Path:
36
  try:
37
- DATA_ROOT.mkdir(parents=True, exist_ok=True)
38
- probe = DATA_ROOT / ".writable_probe"
39
  probe.write_text("ok")
40
  probe.unlink()
41
- return DATA_ROOT
42
- except OSError as exc:
43
- # /data isn't writable (local dev, wrong hardware tier, etc.). Fall back
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 main() -> int:
53
- data_root = _ensure_writable_data_dir()
54
- lyra2_root = data_root / "Lyra-2"
55
- lyra2_root.mkdir(parents=True, exist_ok=True)
 
 
 
 
 
 
 
 
 
 
 
56
 
57
- model_marker = lyra2_root / "checkpoints" / "model"
58
- if model_marker.exists() and any(model_marker.iterdir()):
59
- log.info("Checkpoints already present at %s; skipping download.", model_marker)
60
- else:
61
- log.info("Downloading %s checkpoints/* to %s — this is ~50-80 GB and takes 20-40 min on cold start.",
 
 
 
 
 
 
 
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
- # The Lyra-2 inference modules expect `checkpoints/` at Lyra-2/ repo root.
74
- # Symlink the persistent-volume copy into the app's Lyra-2 tree so every
75
- # existing `--checkpoint_dir checkpoints/model` CLI stays unchanged.
76
- src = lyra2_root / "checkpoints"
77
- dst = APP_CHECKPOINTS
78
- if dst.is_symlink() or dst.exists():
79
- if dst.is_symlink() and dst.resolve() == src.resolve():
80
- log.info("Checkpoints symlink already points to %s", src)
81
- return 0
82
- log.info("Replacing stale %s with symlink to %s", dst, src)
83
- if dst.is_symlink() or dst.is_file():
84
- dst.unlink()
85
- else:
86
- shutil.rmtree(dst)
87
- dst.parent.mkdir(parents=True, exist_ok=True)
88
- dst.symlink_to(src, target_is_directory=True)
89
- log.info("Linked %s -> %s", dst, src)
 
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