mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user