fix(model): skip redundant model reload after unet override load

load_transformer consumes the sd_unet dropdown selection during a full
model load but never marked it as loaded, so the queued sd_unet
onchange callback always forced a second full reload. Sync
sd_unet.loaded_unet once the override is successfully consumed; the
incompatible-override fallback keeps its reset to Default.
This commit is contained in:
CalamitousFelicitousness
2026-07-10 02:39:24 +01:00
parent 2a7d4b4037
commit cbe3c148e0
2 changed files with 44 additions and 0 deletions
+7
View File
@@ -69,6 +69,7 @@ def load_transformer(
)
local_file = None
override_name = None
fallback = True
from modules import sd_unet
@@ -77,6 +78,7 @@ def load_transformer(
log.error(f'Load module: type=transformer file="{shared.opts.sd_unet}" not found')
elif os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]):
local_file = sd_unet.unet_dict[shared.opts.sd_unet]
override_name = shared.opts.sd_unet
if repo_id.startswith(shared.opts.ckpt_dir) and os.path.exists(repo_id):
log.error(f'Load model: transformer="{repo_id}" is incorrectly placed in the checkpoints folder')
@@ -142,6 +144,11 @@ def load_transformer(
else:
transformer = load_from_repo()
# mark the dropdown selection as loaded so the sd_unet onchange callback
# does not force a redundant full reload for an already-consumed override
if transformer is not None and override_name is not None and shared.opts.sd_unet == override_name:
sd_unet.loaded_unet = override_name
sd_models.allow_post_quant = False # we already handled it
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
sd_models.move_model(transformer, devices.cpu)
+37
View File
@@ -1202,6 +1202,42 @@ def test_load_comfy_marker_for_unknown_module_raises_mismatch():
os.unlink(path)
def test_load_transformer_syncs_loaded_unet():
"""A full model load that consumes the UNET dropdown override must mark it
as loaded, so the queued sd_unet onchange callback does not trigger a
second, redundant full model reload."""
from modules import sd_unet, shared
from pipelines import generic_transformer as gt
fd, path = tempfile.mkstemp(suffix='.safetensors')
dim = 8
raw = {
'model.diffusion_model.in_proj.weight': torch.randn(dim, dim),
'model.diffusion_model.in_proj.bias': torch.zeros(dim),
'model.diffusion_model.out_proj.weight': torch.randn(dim, dim),
'model.diffusion_model.out_proj.bias': torch.zeros(dim),
}
write_fixture(raw, fd, path)
orig_unet_opt = shared.opts.sd_unet
orig_loaded = sd_unet.loaded_unet
sd_unet.unet_dict['mock-unet'] = path
shared.opts.data['sd_unet'] = 'mock-unet'
sd_unet.loaded_unet = None
try:
with ComfyTestEnv(dim):
spec = nt.TransformerSpec(cls=MockMiniTransformer)
transformer = gt.load_transformer('fake/repo', cls_name=MockMiniTransformer, native_spec=spec)
assert transformer is not None
assert sd_unet.loaded_unet == 'mock-unet', f'override consumed but loaded_unet={sd_unet.loaded_unet}'
finally:
sd_unet.unet_dict.pop('mock-unet', None)
shared.opts.data['sd_unet'] = orig_unet_opt
sd_unet.loaded_unet = orig_loaded
if os.path.exists(path):
os.unlink(path)
def test_build_component_comfy_preempts_sdnq_fresh_quant():
"""When SDNQ on-load quant settings are active (quant_type=SDNQConfig), a
comfy_quant file must still take the pre-quantized path: fresh quant of
@@ -1358,6 +1394,7 @@ def run_all():
test_load_comfy_marker_dtype_mismatch_raises,
test_load_comfy_unsupported_format_raises_mismatch,
test_load_comfy_marker_for_unknown_module_raises_mismatch,
test_load_transformer_syncs_loaded_unet,
test_build_component_comfy_preempts_sdnq_fresh_quant,
]:
run_test(cat, fn)