Keep latent upscaler loaded during pass

This commit is contained in:
2026-09-04 11:04:26 +00:00
parent e8599055a1
commit 0b189dcf7a
2 changed files with 107 additions and 2 deletions
+79
View File
@@ -1278,6 +1278,85 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
self.assertEqual(chunk_length, 85)
self.assertEqual(temporal_overlap, 17)
def test_unload_upscale_model_defers_while_held(self):
latent = importlib.import_module("dumas_h3_latent_upscale")
class FakeParam:
device = "cuda"
class FakeModel:
def __init__(self):
self.moves = []
def parameters(self):
return iter((FakeParam(),))
def to(self, device):
self.moves.append(device)
return self
cache_key = "upscale.safetensors::cuda::fp16"
original_cache_value = latent._MODEL_CACHE.get(cache_key)
original_hold_depth = latent._MODEL_HOLD_DEPTH
fake_model = FakeModel()
try:
latent._MODEL_CACHE[cache_key] = fake_model
latent._MODEL_HOLD_DEPTH = 0
with latent._hold_upscale_model_loaded():
latent.unload_upscale_model("upscale.safetensors", "cuda", "fp16")
self.assertEqual(fake_model.moves, [])
latent.unload_upscale_model("upscale.safetensors", "cuda", "fp16")
self.assertEqual(fake_model.moves, ["cpu"])
finally:
latent._MODEL_HOLD_DEPTH = original_hold_depth
if original_cache_value is None:
latent._MODEL_CACHE.pop(cache_key, None)
else:
latent._MODEL_CACHE[cache_key] = original_cache_value
def test_model_upscale_releases_cached_model_after_pass(self):
latent = importlib.import_module("dumas_h3_latent_upscale")
calls = []
original_temporal = latent._upscale_video_temporal_chunks
original_unload_now = latent._unload_upscale_model_now
original_cuda = latent.torch.cuda
original_device = getattr(latent.torch, "device", None)
try:
latent.torch.cuda = types.SimpleNamespace(is_available=lambda: True)
latent.torch.device = lambda value: value
def temporal(video, param, upscaler):
calls.append(("temporal", latent._MODEL_HOLD_DEPTH))
return "video", 8, 16
def unload_now(name, device, precision):
calls.append(("unload", name, device, precision, latent._MODEL_HOLD_DEPTH))
latent._upscale_video_temporal_chunks = temporal
latent._unload_upscale_model_now = unload_now
result = latent.upscale_latent_video("source", {
"mode": "model",
"model_name": "upscale.safetensors",
"device": "cuda",
"precision": "fp16",
})
self.assertEqual(result, ("video", 8, 16))
self.assertEqual(calls[0], ("temporal", 1))
self.assertEqual(calls[1], ("unload", "upscale.safetensors", "cuda", "fp16", 1))
self.assertEqual(latent._MODEL_HOLD_DEPTH, 0)
finally:
latent._upscale_video_temporal_chunks = original_temporal
latent._unload_upscale_model_now = original_unload_now
latent.torch.cuda = original_cuda
if original_device is None:
delattr(latent.torch, "device")
else:
latent.torch.device = original_device
def test_upscale_video_model_raises_when_gpu_cannot_shrink(self):
latent = importlib.import_module("dumas_h3_latent_upscale")