fix(upscale): make SeedVR2 generation_step patch idempotent

UpscalerSeedVR.load_model() rebinds the module-global
generation.generation_step (called by name inside generation_loop) to
the instance's model_step wrapper, keeping the previous value to call
back into. That global was never restored, so the second pass through
load_model() saved the wrapper itself as the "original", making
model_step() call itself -> RecursionError.

The second pass is reached on any model (re)load: with upscaler_unload
enabled (self.model reset to None after each run) every subsequent run
recurses, and switching SeedVR variants (self.model_loaded != model_name)
triggers it even without unload.

Stash the pristine generation_step on the module once and have the
wrapper call that, so repeated loads never wrap the wrapper.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
CalamitousFelicitousness
2026-06-16 10:18:21 +01:00
parent ccae78ca66
commit 0769d423a7
+6 -2
View File
@@ -47,7 +47,10 @@ class UpscalerSeedVR(Upscaler):
self.model.dit.dtype = devices.dtype
self.model.vae_encode = self.vae_encode
self.model.vae_decode = self.vae_decode
self.model.model_step = generation.generation_step
# Patch generation_loop's generation_step() with our wrapper; stash the original once
# so reloads don't re-wrap the wrapper itself (infinite recursion).
if not hasattr(generation, "generation_step_original"):
generation.generation_step_original = generation.generation_step
generation.generation_step = self.model_step
self.model._internal_dict = {
'dit': self.model.dit,
@@ -119,6 +122,7 @@ class UpscalerSeedVR(Upscaler):
return samples
def model_step(self, *args, **kwargs):
from modules.seedvr.src.core import generation
from modules.seedvr.src.optimization import memory_manager
self.model.vae = self.model.vae.to(device="cpu")
self.model.dit = self.model.dit.to(device=self.device)
@@ -126,7 +130,7 @@ class UpscalerSeedVR(Upscaler):
log.debug(f'Upscaler inference: args={len(args)} kwargs={list(kwargs.keys())}')
memory_manager.preinitialize_rope_cache(self.model)
with devices.inference_context():
result = self.model.model_step(*args, **kwargs)
result = generation.generation_step_original(*args, **kwargs)
self.model.dit = self.model.dit.to(device="cpu")
devices.torch_gc()
return result