Stop retrying latent upscale OOMs as sampling
This commit is contained in:
+35
-20
@@ -4136,6 +4136,12 @@ def _param_value(mapping, key, default):
|
||||
return default if value is None else value
|
||||
|
||||
|
||||
def _tag_oom_stage(exc, stage):
|
||||
if _is_oom(exc):
|
||||
exc._h3_stage = stage
|
||||
return exc
|
||||
|
||||
|
||||
def _grid_1d(size, tile, ol, min_tile):
|
||||
if size <= tile:
|
||||
return [0], [size], [0]
|
||||
@@ -6645,9 +6651,7 @@ class H3LongVideos:
|
||||
# OOM retry cannot help an OOM raised here -- it just re-runs the whole
|
||||
# sampling pass and fails the same way, which on a 362-frame shot is four
|
||||
# more minutes for nothing.
|
||||
if _is_oom(e):
|
||||
e._h3_stage = "sampling"
|
||||
raise
|
||||
raise _tag_oom_stage(e, "sampling")
|
||||
refined_out = out
|
||||
latent_upscale_param = latent_upscale_param or None
|
||||
latent_upscale_mode = _latent_upscale_mode(latent_upscale_param)
|
||||
@@ -6832,9 +6836,7 @@ class H3LongVideos:
|
||||
del upscale_latent, upscaled_video, full_audio
|
||||
mm.soft_empty_cache()
|
||||
except Exception as e:
|
||||
if _is_oom(e):
|
||||
e._h3_stage = "sampling"
|
||||
raise
|
||||
raise _tag_oom_stage(e, "latent_upscale")
|
||||
# Keep a CPU copy of the sampled latent BEFORE decoding, for the `latent`
|
||||
# output. Latents are ~1000x smaller than the frames they decode to (a
|
||||
# 1344x768 124f shot is ~1.5MB against ~1.5GB), so carrying one per shot for
|
||||
@@ -7446,13 +7448,19 @@ class H3LongVideos:
|
||||
shot_refs, ref_image_size, shot_aug, shot_silent,
|
||||
latent_upscale_param=latent_upscale_param, timing_sink=shot_timing)
|
||||
break
|
||||
except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
|
||||
shot_retry_elapsed += time.perf_counter() - attempt_start
|
||||
if not _is_oom(e):
|
||||
raise
|
||||
mm.soft_empty_cache(True)
|
||||
if not tiled:
|
||||
tiled = True; backoff.append("tiled decode")
|
||||
except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
|
||||
shot_retry_elapsed += time.perf_counter() - attempt_start
|
||||
if not _is_oom(e):
|
||||
raise
|
||||
if getattr(e, "_h3_stage", "") == "latent_upscale":
|
||||
raise RuntimeError(
|
||||
f"H3 Long Videos: shot {i + 1} of {len(gens)} ran out of VRAM "
|
||||
f"during latent upscale. The latent-upscale stage is too large "
|
||||
f"for this card at {w}x{h} with the current spatial settings."
|
||||
) from e
|
||||
mm.soft_empty_cache(True)
|
||||
if not tiled:
|
||||
tiled = True; backoff.append("tiled decode")
|
||||
elif allow_res_backoff and min(w, h) > 384:
|
||||
nw, nh = res_down(w, h); backoff.append(f"res->{nw}x{nh}"); w, h = nw, nh
|
||||
else:
|
||||
@@ -7467,13 +7475,20 @@ class H3LongVideos:
|
||||
tiled, sa, shot_handoff, decode_tile_frames, decode_tile_size,
|
||||
shot_refs, ref_image_size, shot_aug, shot_silent,
|
||||
latent_upscale_param=latent_upscale_param, timing_sink=shot_timing)
|
||||
except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
|
||||
shot_retry_elapsed += time.perf_counter() - attempt_start
|
||||
if _is_oom(e) and getattr(e, "_h3_stage", "") == "sampling":
|
||||
# Retrying with tiles would re-run the whole sampling pass and
|
||||
# fail identically. Fail now, and say what actually shrinks it.
|
||||
raise RuntimeError(
|
||||
f"H3 Long Videos: shot {i + 1} of {len(gens)} ran out of VRAM "
|
||||
except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
|
||||
shot_retry_elapsed += time.perf_counter() - attempt_start
|
||||
stage = getattr(e, "_h3_stage", "")
|
||||
if _is_oom(e) and stage == "latent_upscale":
|
||||
raise RuntimeError(
|
||||
f"H3 Long Videos: shot {i + 1} of {len(gens)} ran out of VRAM "
|
||||
f"during latent upscale. The latent-upscale stage is too large "
|
||||
f"for this card at {w}x{h} with the current spatial settings."
|
||||
) from e
|
||||
if _is_oom(e) and getattr(e, "_h3_stage", "") == "sampling":
|
||||
# Retrying with tiles would re-run the whole sampling pass and
|
||||
# fail identically. Fail now, and say what actually shrinks it.
|
||||
raise RuntimeError(
|
||||
f"H3 Long Videos: shot {i + 1} of {len(gens)} ran out of VRAM "
|
||||
f"while sampling. " + sampling_oom_help(w, h, ln_i, fps, megapixels)
|
||||
) from e
|
||||
if not _is_oom(e) or tiled:
|
||||
|
||||
@@ -1154,6 +1154,12 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
||||
self.assertEqual(self.module._latent_upscale_mode({"mode": "model"}), "model")
|
||||
self.assertEqual(self.module._latent_upscale_mode({}), "off")
|
||||
|
||||
def test_tag_oom_stage_marks_oom_exceptions(self):
|
||||
exc = RuntimeError("CUDA out of memory")
|
||||
tagged = self.module._tag_oom_stage(exc, "latent_upscale")
|
||||
self.assertIs(tagged, exc)
|
||||
self.assertEqual(getattr(tagged, "_h3_stage", ""), "latent_upscale")
|
||||
|
||||
def test_compose_persistent_does_not_expand_ambiguous_plural_to_full_cast(self):
|
||||
active = self.module.parse_wardrobe(
|
||||
"Maya = she, red jacket\n"
|
||||
|
||||
Reference in New Issue
Block a user