From cbabcf8208d5b7caf932256bb159578eba546221 Mon Sep 17 00:00:00 2001 From: Chris Dumas Date: Fri, 4 Sep 2026 09:49:34 +0000 Subject: [PATCH] Offload H3 before latent upscale --- dumas_h3_longvideos.py | 6 +++--- tests/test_dumas_h3_longvideos.py | 22 ++++++++++++++++++---- 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/dumas_h3_longvideos.py b/dumas_h3_longvideos.py index eb6bfcb..28e88a6 100644 --- a/dumas_h3_longvideos.py +++ b/dumas_h3_longvideos.py @@ -6662,6 +6662,9 @@ class H3LongVideos: parts = _nested_tensor_parts(out_samples) if not getattr(out_samples, "is_nested", False) or len(parts) < 2: raise RuntimeError("latent upscale expects a nested AV latent") + if latent_upscale_mode == "model" and str(latent_upscale_param.get("device", "cuda")) == "cuda" and hasattr(model, "clone_base_uuid"): + mm.unload_model_and_clones(model, unload_additional_models=False) + mm.soft_empty_cache() upscaled_video, up_h, up_w = _upscale_latent_video(parts[0], latent_upscale_param) full_audio = parts[1] # Drop the first-pass sampling state before we start the refinement @@ -6674,9 +6677,6 @@ class H3LongVideos: target_h = int(up_h) * 16 if target_w <= 0 or target_h <= 0: raise RuntimeError("latent upscale target size must be positive") - if latent_upscale_mode == "model" and str(latent_upscale_param.get("device", "cuda")) == "cuda" and hasattr(model, "clone_base_uuid"): - mm.unload_model_and_clones(model, unload_additional_models=False) - mm.soft_empty_cache() upscale_cond, upscale_latent = _build_shot_conditioning( clip, vae, prompt, target_w, target_h, ln, fps, handoff, ref_images=refs, ref_image_size=ref_image_size, diff --git a/tests/test_dumas_h3_longvideos.py b/tests/test_dumas_h3_longvideos.py index 165b252..4539299 100644 --- a/tests/test_dumas_h3_longvideos.py +++ b/tests/test_dumas_h3_longvideos.py @@ -373,6 +373,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): original_cleanup = self.module._deep_cleanup original_upscale = self.module._upscale_latent_video original_copy_sample = self.module._copy_sample_latent + original_unload = getattr(self.module.mm, "unload_model_and_clones", None) original_nested = getattr(self.module.comfy.nested_tensor, "NestedTensor", None) try: self.module.comfy.nested_tensor.NestedTensor = FakeNestedTensor @@ -395,20 +396,28 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): def cleanup(): order.append("cleanup") + def unload_model_and_clones(*_args, **_kwargs): + order.append("unload_h3") + + def upscale_latent_video(video, param): + order.append("upscale") + return FakeTensor("upv"), 8, 16 + self.module.nodes.common_ksampler = common_ksampler self.module._build_shot_conditioning = lambda *_args, **_kwargs: ( "cond", {"samples": FakeNestedTensor((FakeTensor("basev"), FakeTensor("basea")))}, ) self.module._evict_all_but = lambda *_args, **_kwargs: None - self.module._upscale_latent_video = lambda video, param: (FakeTensor("upv"), 8, 16) + self.module.mm.unload_model_and_clones = unload_model_and_clones + self.module._upscale_latent_video = upscale_latent_video self.module._copy_sample_latent = lambda sampled: sampled["samples"].unbind() self.module._decode_video = decode_video self.module._decode_audio = decode_audio self.module._deep_cleanup = cleanup self.module.H3LongVideos()._render( - model=object(), + model=types.SimpleNamespace(clone_base_uuid="h3"), clip=types.SimpleNamespace( tokenize=lambda text, **kwargs: text, encode_from_tokens_scheduled=lambda tokens: tokens, @@ -427,7 +436,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): latent_upscale_param={ "mode": "model", "model_name": "upscale.safetensors", - "device": "cpu", + "device": "cuda", "precision": "fp16", "sampler_name": "euler_ancestral", "scheduler": "simple", @@ -438,7 +447,8 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): ) self.assertEqual(order[0], "sample") - self.assertEqual(order[1], "latent_upscale_sample") + self.assertLess(order.index("unload_h3"), order.index("upscale")) + self.assertLess(order.index("upscale"), order.index("latent_upscale_sample")) self.assertLess(order.index("audio"), order.index("video")) self.assertEqual(order[-1], "cleanup") finally: @@ -450,6 +460,10 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): self.module._deep_cleanup = original_cleanup self.module._upscale_latent_video = original_upscale self.module._copy_sample_latent = original_copy_sample + if original_unload is None: + delattr(self.module.mm, "unload_model_and_clones") + else: + self.module.mm.unload_model_and_clones = original_unload if original_nested is None: delattr(self.module.comfy.nested_tensor, "NestedTensor") else: