cezar-hapiko commited on
Commit
d8da5e8
·
1 Parent(s): 328801b

Fix stage-2 OOM: offload Lyra-2 net to CPU during DA3 GS recon

Browse files
Files changed (2) hide show
  1. app.py +23 -6
  2. entrypoint.sh +5 -0
app.py CHANGED
@@ -219,16 +219,33 @@ def _stage1_video(
219
  def _stage2_splat(
220
  run_dir: Path,
221
  video_path: Path,
 
222
  res2: Stage2Resources,
223
  progress: gr.Progress,
224
  ) -> Path:
225
  progress(0.75, desc="Stage 2/2 — reconstructing 3D Gaussian splat (~3 min)")
226
  t0 = time.monotonic()
227
- ply_path = run_stage2_single(
228
- res2,
229
- video_path=video_path,
230
- output_dir=run_dir / "recon",
231
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
232
  log.info("[stage2] done in %.1fs -> %s", time.monotonic() - t0, ply_path)
233
  if not ply_path.exists():
234
  raise gr.Error(f"Stage 2 did not produce {ply_path}.")
@@ -277,7 +294,7 @@ def generate(
277
 
278
  try:
279
  video_path = _stage1_video(run_dir, input_image, prompt.strip(), preset_params, res1, progress)
280
- ply_path = _stage2_splat(run_dir, video_path, res2, progress)
281
  except gr.Error:
282
  _CONSECUTIVE_FAILURES += 1
283
  if torch.cuda.is_available():
 
219
  def _stage2_splat(
220
  run_dir: Path,
221
  video_path: Path,
222
+ res1: Stage1Resources,
223
  res2: Stage2Resources,
224
  progress: gr.Progress,
225
  ) -> Path:
226
  progress(0.75, desc="Stage 2/2 — reconstructing 3D Gaussian splat (~3 min)")
227
  t0 = time.monotonic()
228
+
229
+ # Offload Lyra-2 diffusion net to CPU so DA3 GS recon has VRAM headroom.
230
+ # First-run measurement on A100-80GB: with stage-1 resident, DA3's DINOv2
231
+ # backbone OOMs trying to allocate 1.47 GiB (78.98/79.25 GiB already used).
232
+ # DA3 stays on GPU since it's reused across both stages.
233
+ res1.model.net.cpu()
234
+ if torch.cuda.is_available():
235
+ torch.cuda.empty_cache()
236
+
237
+ try:
238
+ ply_path = run_stage2_single(
239
+ res2,
240
+ video_path=video_path,
241
+ output_dir=run_dir / "recon",
242
+ )
243
+ finally:
244
+ # Restore the diffusion net to GPU for the next stage-1 request.
245
+ res1.model.net.to(device=res1.desired_device, dtype=res1.desired_dtype)
246
+ if torch.cuda.is_available():
247
+ torch.cuda.empty_cache()
248
+
249
  log.info("[stage2] done in %.1fs -> %s", time.monotonic() - t0, ply_path)
250
  if not ply_path.exists():
251
  raise gr.Error(f"Stage 2 did not produce {ply_path}.")
 
294
 
295
  try:
296
  video_path = _stage1_video(run_dir, input_image, prompt.strip(), preset_params, res1, progress)
297
+ ply_path = _stage2_splat(run_dir, video_path, res1, res2, progress)
298
  except gr.Error:
299
  _CONSECUTIVE_FAILURES += 1
300
  if torch.cuda.is_available():
entrypoint.sh CHANGED
@@ -4,6 +4,11 @@
4
  # 2. Launch the Gradio app.
5
  set -euo pipefail
6
 
 
 
 
 
 
7
  cd /home/user/app
8
 
9
  # The checkpoint downloader is idempotent — on warm starts with a populated
 
4
  # 2. Launch the Gradio app.
5
  set -euo pipefail
6
 
7
+ # expandable_segments reduces fragmentation when the resident Lyra-2 diffusion
8
+ # net is swapped between GPU/CPU to make room for DA3's stage-2 GS recon pass.
9
+ # Without this, we can still OOM even though aggregate free memory is enough.
10
+ export PYTORCH_CUDA_ALLOC_CONF=${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}
11
+
12
  cd /home/user/app
13
 
14
  # The checkpoint downloader is idempotent — on warm starts with a populated