diff --git a/pipelines/generic_transformer.py b/pipelines/generic_transformer.py index 142b5ce45..99b77a716 100644 --- a/pipelines/generic_transformer.py +++ b/pipelines/generic_transformer.py @@ -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) diff --git a/test/test-native-transformer.py b/test/test-native-transformer.py index 0d3ab10af..616fc8c5a 100644 --- a/test/test-native-transformer.py +++ b/test/test-native-transformer.py @@ -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)