lyra2-explorable-scene / warm_model_test.py
cezar-hapiko's picture
Refactor: load Lyra-2 model once at startup, reuse across requests
328801b
Raw History Blame Contribute Delete
4.61 kB
"""Ad-hoc Modal-side harness for verifying resident_inference.py end-to-end.
Runs the full load-once → run-twice flow against Lyra-2's bundled samples 04 and
05. The Modal kernel must already have the Lyra-2 env resident and checkpoints
downloaded under ``checkpoints/`` in the current working directory (the kernel
ships this layout out of the box).
Usage (inside the Modal kernel):
cd /mnt/lyra-data/lyra/Lyra-2
python /mnt/lyra-data/warm_model_test.py
Pass condition: sample-05 wall-clock ≤ 15 min (target ~11 min). If it's not, the
resident loader is missing a hot path — grep the log for "load_da3_model" /
"load_model_from_checkpoint" / "Loading" to spot the culprit.
"""
from __future__ import annotations
import sys
import time
from pathlib import Path
# Assume this file is copied alongside Lyra-2 at /mnt/lyra-data/. Make the
# sibling Lyra-2 tree + the hf_space tools dir importable.
HERE = Path(__file__).resolve().parent
LYRA2_DIR = Path("Lyra-2").resolve() if Path("Lyra-2").is_dir() else (HERE / "lyra" / "Lyra-2").resolve()
if str(LYRA2_DIR) not in sys.path:
sys.path.insert(0, str(LYRA2_DIR))
if str(HERE) not in sys.path:
sys.path.insert(0, str(HERE))
from resident_inference import ( # noqa: E402
build_stage1_args,
load_stage1_resources,
load_stage2_resources,
run_stage1_single,
run_stage2_single,
)
PRESET = {
"num_frames_zoom_in": 81,
"num_frames_zoom_out": 241,
"zoom_in_strength": 0.5,
"zoom_out_strength": 1.5,
}
SAMPLES = ["04", "05"]
def _caption(samples_dir: Path, stem: str) -> str:
txt = samples_dir / f"{stem}.txt"
if not txt.is_file():
raise FileNotFoundError(f"Missing caption file: {txt}")
return txt.read_text(encoding="utf-8").strip()
def _run_one(res1, res2, samples_dir: Path, stem: str, out_root: Path) -> float:
img = samples_dir / f"{stem}.png"
if not img.is_file():
raise FileNotFoundError(f"Missing sample image: {img}")
caption = _caption(samples_dir, stem)
run_dir = out_root / stem
run_dir.mkdir(parents=True, exist_ok=True)
t0 = time.monotonic()
print(f"\n[warm_test] ====== sample {stem} ======", flush=True)
print(f"[warm_test] image={img}", flush=True)
print(f"[warm_test] caption={caption[:120]}...", flush=True)
video_path = run_stage1_single(
res1,
image_path=img,
prompt=caption,
preset_params=PRESET,
output_path=run_dir / "zoomgs",
)
t_stage1 = time.monotonic() - t0
ply_path = run_stage2_single(res2, video_path=video_path, output_dir=run_dir / "recon")
t_total = time.monotonic() - t0
print(
f"[warm_test] sample {stem} done: stage1={t_stage1:.1f}s "
f"total={t_total:.1f}s ({t_total/60:.1f} min)",
flush=True,
)
print(f"[warm_test] video={video_path}", flush=True)
print(f"[warm_test] ply={ply_path}", flush=True)
return t_total
def main() -> int:
samples_dir = LYRA2_DIR / "assets" / "samples"
if not samples_dir.is_dir():
print(f"[warm_test] FATAL: samples dir not found at {samples_dir}", flush=True)
return 2
out_root = Path("/tmp/warm_model_test")
out_root.mkdir(parents=True, exist_ok=True)
t_setup = time.monotonic()
args = build_stage1_args(checkpoint_dir="checkpoints/model", experiment="lyra2", use_dmd=True)
print("[warm_test] loading stage-1 resources (cold, ~30 min on A100)...", flush=True)
res1 = load_stage1_resources(args)
print("[warm_test] loading stage-2 resources (reusing DA3 from stage-1)...", flush=True)
res2 = load_stage2_resources(da3_from_stage1=res1.da3_model)
print(
f"[warm_test] setup done in {time.monotonic() - t_setup:.1f}s "
f"({(time.monotonic() - t_setup)/60:.1f} min)",
flush=True,
)
timings = {}
for stem in SAMPLES:
timings[stem] = _run_one(res1, res2, samples_dir, stem, out_root)
print("\n[warm_test] ===== summary =====", flush=True)
for stem in SAMPLES:
print(f" sample {stem}: {timings[stem]/60:.1f} min", flush=True)
second_min = timings[SAMPLES[1]] / 60.0
if second_min <= 15.0:
print(
f"[warm_test] PASS: second request ({second_min:.1f} min) <= 15 min target",
flush=True,
)
return 0
print(
f"[warm_test] FAIL: second request ({second_min:.1f} min) > 15 min target. "
"A hot path is still reloading per request — check log for load_* calls.",
flush=True,
)
return 1
if __name__ == "__main__":
raise SystemExit(main())