diff --git a/modules/sd_models.py b/modules/sd_models.py index 494c5bf1d..69abadc7e 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1573,6 +1573,8 @@ def unload_model_weights(op='model'): disable_offload(model_data.sd_model) move_model(model_data.sd_model, 'meta') model_data.sd_model = None + from modules.sdnq.common import reset_compile_caches + reset_compile_caches() # dead compiled-dequant graphs and their lifetime recompile counters otherwise accumulate across switches devices.torch_gc(force=True, reason='unload') log.debug(f'Unload {op}: {memory_stats()} fn={fn}') elif (op == 'refiner') and model_data.sd_refiner: diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 92801ce25..d82bb4823 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -368,6 +368,22 @@ else: return fn +def reset_compile_caches(): + """Drop compiled graphs and dynamo's lifetime recompile counters at model unload. + + Freed layers leave their guards invalid, so the dead graphs cannot be + reused, but the per-frame lifetime compile counter keeps climbing and a + fullgraph compile hard-fails (FailOnRecompileLimitHit) once it crosses the + accumulated limit; enough model or quant switches in one process get there. + The reset is global, so call it only when the compiled graphs' owner is + being discarded. + """ + if check_torch_compile(): + from modules.logger import log + torch._dynamo.reset() + log.debug('SDNQ compile: dynamo reset') + + common_skip_keys = ( ".time_embed", ".context_embedder",