diff --git a/dumas_h3_latent_upscale.py b/dumas_h3_latent_upscale.py index 38f51ab..eb669e8 100644 --- a/dumas_h3_latent_upscale.py +++ b/dumas_h3_latent_upscale.py @@ -72,6 +72,10 @@ def _is_oom_error(exc): return "out of memory" in text or "exhausted its gpu spatial fallbacks" in text +def _should_retry_temporal_before_spatial(param): + return _uses_cuda_model_upscale(param) and int(param.get("chunk_length", 0) or 0) > CUDA_MODEL_CHUNK_LENGTH + + def _retry_with_smaller_temporal(video, param, upscaler, exc): if not _is_oom_error(exc): raise exc @@ -665,6 +669,11 @@ def upscale_video_model(video, param): except RuntimeError as exc: if "out of memory" not in str(exc).lower(): raise + if _should_retry_temporal_before_spatial(param): + raise RuntimeError( + "out of memory: retry H3 latent upscale with a smaller temporal chunk " + "before spatial fallback" + ) from exc smaller = _shrink_model_tile_param(param) if smaller is None: raise RuntimeError( diff --git a/tests/test_dumas_h3_longvideos.py b/tests/test_dumas_h3_longvideos.py index ac2a1ad..cd75bca 100644 --- a/tests/test_dumas_h3_longvideos.py +++ b/tests/test_dumas_h3_longvideos.py @@ -1290,6 +1290,40 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): self.assertEqual(chunk_length, 85) self.assertEqual(temporal_overlap, 17) + def test_cuda_model_oom_retries_temporal_before_spatial_fallback(self): + latent = importlib.import_module("dumas_h3_latent_upscale") + + calls = [] + original_tiled = latent._upscale_video_model_tiled + original_shrink_model = latent._shrink_model_tile_param + try: + def tiled(_video, param): + calls.append(("tiled", param.get("chunk_length"), param.get("tile_size_mode"))) + raise RuntimeError("out of memory") + + latent._upscale_video_model_tiled = tiled + latent._shrink_model_tile_param = ( + lambda param: calls.append(("shrink_spatial", param.get("tile_size_mode"))) or None + ) + + with self.assertRaisesRegex(RuntimeError, "smaller temporal chunk"): + latent.upscale_video_model( + "video", + { + "mode": "model", + "device": "cuda", + "model_name": "upscale.safetensors", + "chunk_length": 85, + "temporal_overlap": 17, + }, + ) + + self.assertEqual(calls[0], ("tiled", 85, None)) + self.assertNotIn(("shrink_spatial", None), calls) + finally: + latent._upscale_video_model_tiled = original_tiled + latent._shrink_model_tile_param = original_shrink_model + def test_interp_temporal_params_preserve_upstream_defaults(self): latent = importlib.import_module("dumas_h3_latent_upscale") chunk_length, temporal_overlap = latent._effective_temporal_params({