Harden latent upscale GPU cleanup

This commit is contained in:
2026-09-04 10:05:24 +00:00
parent cbabcf8208
commit 618a48e4d9
2 changed files with 45 additions and 16 deletions
+25 -3
View File
@@ -6085,6 +6085,29 @@ def _evict_all_but(keep_model):
pass pass
def _evict_for_latent_upscale(model):
"""Clear the sampler model before loading the auxiliary latent upscaler."""
try:
unload_clones = getattr(mm, "unload_model_and_clones", None)
if callable(unload_clones):
try:
unload_clones(model, unload_additional_models=False)
mm.soft_empty_cache()
return
except Exception:
pass
mm.unload_all_models()
except Exception:
pass
try:
mm.soft_empty_cache(True)
except Exception:
try:
mm.soft_empty_cache()
except Exception:
pass
@@ -6662,9 +6685,8 @@ class H3LongVideos:
parts = _nested_tensor_parts(out_samples) parts = _nested_tensor_parts(out_samples)
if not getattr(out_samples, "is_nested", False) or len(parts) < 2: if not getattr(out_samples, "is_nested", False) or len(parts) < 2:
raise RuntimeError("latent upscale expects a nested AV latent") 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"): if latent_upscale_mode == "model" and str(latent_upscale_param.get("device", "cuda")) == "cuda":
mm.unload_model_and_clones(model, unload_additional_models=False) _evict_for_latent_upscale(model)
mm.soft_empty_cache()
upscaled_video, up_h, up_w = _upscale_latent_video(parts[0], latent_upscale_param) upscaled_video, up_h, up_w = _upscale_latent_video(parts[0], latent_upscale_param)
full_audio = parts[1] full_audio = parts[1]
# Drop the first-pass sampling state before we start the refinement # Drop the first-pass sampling state before we start the refinement
+10 -3
View File
@@ -374,6 +374,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
original_upscale = self.module._upscale_latent_video original_upscale = self.module._upscale_latent_video
original_copy_sample = self.module._copy_sample_latent original_copy_sample = self.module._copy_sample_latent
original_unload = getattr(self.module.mm, "unload_model_and_clones", None) original_unload = getattr(self.module.mm, "unload_model_and_clones", None)
original_unload_all = self.module.mm.unload_all_models
original_nested = getattr(self.module.comfy.nested_tensor, "NestedTensor", None) original_nested = getattr(self.module.comfy.nested_tensor, "NestedTensor", None)
try: try:
self.module.comfy.nested_tensor.NestedTensor = FakeNestedTensor self.module.comfy.nested_tensor.NestedTensor = FakeNestedTensor
@@ -397,7 +398,11 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
order.append("cleanup") order.append("cleanup")
def unload_model_and_clones(*_args, **_kwargs): def unload_model_and_clones(*_args, **_kwargs):
order.append("unload_h3") order.append("unload_h3_failed")
raise RuntimeError("model wrapper does not expose clone metadata")
def unload_all_models(*_args, **_kwargs):
order.append("unload_all")
def upscale_latent_video(video, param): def upscale_latent_video(video, param):
order.append("upscale") order.append("upscale")
@@ -410,6 +415,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
) )
self.module._evict_all_but = lambda *_args, **_kwargs: None self.module._evict_all_but = lambda *_args, **_kwargs: None
self.module.mm.unload_model_and_clones = unload_model_and_clones self.module.mm.unload_model_and_clones = unload_model_and_clones
self.module.mm.unload_all_models = unload_all_models
self.module._upscale_latent_video = upscale_latent_video self.module._upscale_latent_video = upscale_latent_video
self.module._copy_sample_latent = lambda sampled: sampled["samples"].unbind() self.module._copy_sample_latent = lambda sampled: sampled["samples"].unbind()
self.module._decode_video = decode_video self.module._decode_video = decode_video
@@ -417,7 +423,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
self.module._deep_cleanup = cleanup self.module._deep_cleanup = cleanup
self.module.H3LongVideos()._render( self.module.H3LongVideos()._render(
model=types.SimpleNamespace(clone_base_uuid="h3"), model=object(),
clip=types.SimpleNamespace( clip=types.SimpleNamespace(
tokenize=lambda text, **kwargs: text, tokenize=lambda text, **kwargs: text,
encode_from_tokens_scheduled=lambda tokens: tokens, encode_from_tokens_scheduled=lambda tokens: tokens,
@@ -447,7 +453,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
) )
self.assertEqual(order[0], "sample") self.assertEqual(order[0], "sample")
self.assertLess(order.index("unload_h3"), order.index("upscale")) self.assertLess(order.index("unload_all"), order.index("upscale"))
self.assertLess(order.index("upscale"), order.index("latent_upscale_sample")) self.assertLess(order.index("upscale"), order.index("latent_upscale_sample"))
self.assertLess(order.index("audio"), order.index("video")) self.assertLess(order.index("audio"), order.index("video"))
self.assertEqual(order[-1], "cleanup") self.assertEqual(order[-1], "cleanup")
@@ -464,6 +470,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
delattr(self.module.mm, "unload_model_and_clones") delattr(self.module.mm, "unload_model_and_clones")
else: else:
self.module.mm.unload_model_and_clones = original_unload self.module.mm.unload_model_and_clones = original_unload
self.module.mm.unload_all_models = original_unload_all
if original_nested is None: if original_nested is None:
delattr(self.module.comfy.nested_tensor, "NestedTensor") delattr(self.module.comfy.nested_tensor, "NestedTensor")
else: else: