feat(model): secondary unet override slot and ideogram4 native loading

One UNET override cannot serve dual-transformer arches: ideogram4
conditional/unconditional and wan combined-stage experts need separate
files, and previously a single override landed on both experts.

- sd_unet_secondary option with per-slot tracking, consumed-state sync,
  arch-change reset, and incompatible-override fallback
- dropdown renders beside the primary, follows it into quicksettings,
  and is visible only for dual-transformer model types
- ideogram4 native single-file spec with a quant-aware fused-qkv
  converter; such converters run before comfy_quant detection via
  TransformerSpec.converter_handles_quant
- quicksettings render in configured order (sort keyed on the option
  object and always fell back to alphabetical)
- post-load dtype warning skips quantized transformers
This commit is contained in:
CalamitousFelicitousness
2026-07-11 05:52:10 +01:00
parent 3831e4563f
commit 8c884c1e02
13 changed files with 551 additions and 72 deletions
+6 -1
View File
@@ -1475,8 +1475,9 @@ def reload_model_weights(sd_model=None, info: CheckpointInfo | None = None, op='
loaded_ckpt = getattr(sd_model, 'sd_checkpoint_info', None) if sd_model is not None else None
changed_checkpoint = loaded_ckpt is None or checkpoint_info is None or loaded_ckpt.filename != checkpoint_info.filename
reset_unet = shared.opts.sd_unet not in (None, 'Default', 'None')
reset_unet_secondary = shared.opts.sd_unet_secondary not in (None, 'Default', 'None')
reset_te = shared.opts.sd_text_encoder not in (None, 'Default', 'None')
if op == 'model' and sd_model is not None and changed_checkpoint and (reset_unet or reset_te):
if op == 'model' and sd_model is not None and changed_checkpoint and (reset_unet or reset_unet_secondary or reset_te):
# compare detected model type, not pipeline class: custom-loader arches (e.g. Krea2) load as a
# concrete class but detect as generic DiffusionPipeline, so a class compare would falsely reset
# across same-arch checkpoints (Base vs Turbo). detect both sides so the comparison is symmetric.
@@ -1490,6 +1491,10 @@ def reload_model_weights(sd_model=None, info: CheckpointInfo | None = None, op='
log.info(f'Load model: type="{old_type}" changed="{new_type}" unet="{shared.opts.sd_unet}" set to default')
shared.opts.data["sd_unet"] = 'Default'
sd_unet.loaded_unet = None
if reset_unet_secondary:
log.info(f'Load model: type="{old_type}" changed="{new_type}" unet_secondary="{shared.opts.sd_unet_secondary}" set to default')
shared.opts.data["sd_unet_secondary"] = 'Default'
sd_unet.loaded_unet_secondary = None
if reset_te:
log.info(f'Load model: type="{old_type}" changed="{new_type}" te="{shared.opts.sd_text_encoder}" set to default')
shared.opts.data["sd_text_encoder"] = 'Default'
+32
View File
@@ -5,11 +5,14 @@ from modules.logger import log
unet_dict = {}
loaded_unet = None
loaded_unet_secondary = None
failed_unet = []
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
dit_models = ['Flux', 'StableDiffusion3', 'HiDream', 'Lumina2', 'Chroma', 'Wan', 'Qwen', 'Anima']
# model types (shared.sd_model_type keyspace) the secondary UNET override applies to
DUAL_TRANSFORMER_TYPES = ('ideogram4', 'wanai')
def load_unet_sdxl_nunchaku(repo_id):
@@ -101,6 +104,35 @@ def load_unet(model, repo_id: str | None = None):
devices.torch_gc()
def load_unet_secondary(model): # pylint: disable=unused-argument
"""Onchange handler for the secondary UNET override: a change means a
full reload; for single-transformer models the selection is stored and
applies on the next dual-transformer load.
"""
global loaded_unet_secondary # pylint: disable=global-statement
selected = shared.opts.sd_unet_secondary
if selected is None or selected in ('Default', 'None'):
if loaded_unet_secondary in (None, 'Default', 'None'):
return
log.info(f'Load module: type=UNet slot=secondary name="Default" (was="{loaded_unet_secondary}") reverting to base transformer')
loaded_unet_secondary = selected
sd_models.reload_model_weights(force=True)
return
if selected not in list(unet_dict):
log.error(f'Load module: type=UNet slot=secondary not found: {selected}')
return
if selected == loaded_unet_secondary or selected in failed_unet:
return
if shared.sd_model_type not in DUAL_TRANSFORMER_TYPES:
log.warning(f'Load module: type=UNet slot=secondary name="{selected}" stored: model type={shared.sd_model_type} has a single transformer, applies on next dual-transformer load')
return
loaded_unet_secondary = selected
sd_models.reload_model_weights(force=True)
devices.torch_gc()
def refresh_unet_list():
unet_dict.clear()
for file in files_cache.list_files(shared.opts.unet_dir, ext_filter=[".safetensors", ".gguf", ".pth"]):
+6
View File
@@ -116,6 +116,12 @@ def refresh_unet_list():
modules.sd_unet.refresh_unet_list()
def sd_unet_secondary_visible():
import modules.sd_unet
from modules import shared
return shared.sd_model_type in modules.sd_unet.DUAL_TRANSFORMER_TYPES
def sd_te_items():
import modules.model_te
predefined = ['Default']
+1
View File
@@ -94,6 +94,7 @@ def create_settings(cmd_opts):
"sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", DropdownEditable, lambda: {"choices": list_checkpoint_titles()}, refresh=refresh_checkpoints),
"sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints),
"sd_unet": OptionInfo("Default", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list),
"sd_unet_secondary": OptionInfo("Default", "UNET model secondary", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items(), "visible": shared_items.sd_unet_secondary_visible()}, refresh=shared_items.refresh_unet_list),
"latent_history": OptionInfo(20, "Latent history size", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}),
"advanced_sep": OptionInfo("<h2>Advanced Options</h2>", "", gr.HTML),
+34 -8
View File
@@ -10,9 +10,20 @@ ui_system_tabs = None # required for system-info
dummy_component = gr.Textbox(visible=False, value='dummy')
loadsave = ui_loadsave.UiLoadsave(shared.cmd_opts.ui_config)
quicksettings_names = {x: i for i, x in enumerate(shared.opts.quicksettings_list) if x != 'quicksettings'}
# companion settings render beside their parent wherever it lives:
# quicksettings the parent and the companion follows, keeping its own visibility
companion_settings = (('sd_unet', 'sd_unet_secondary'),)
for parent_key, companion_key in companion_settings:
if parent_key in quicksettings_names and companion_key not in quicksettings_names:
quicksettings_names[companion_key] = quicksettings_names[parent_key] + 0.5
quicksettings_list = []
hidden_list = []
components = []
# settings with model-dependent visibility: their wrapper group (settings page) or
# refresh button (quicksettings) is registered here so visibility pushes toggle the
# whole control, not just the inner dropdown
dynamic_visibility_keys = ('sd_unet_secondary',)
dynamic_visibility: dict = {}
def apply_setting(key, value):
@@ -76,11 +87,15 @@ def create_setting_component(key, is_quicksettings=False):
dirtyable_setting = gr.Group(elem_classes="dirtyable", visible=args.get("visible", True))
dirtyable_setting.__enter__()
dirty_indicator = gr.Button("", elem_classes="modification-indicator", elem_id=f"modification_indicator_{key}")
if key in dynamic_visibility_keys:
dynamic_visibility.setdefault(key, []).append(dirtyable_setting)
if info.refresh is not None:
if is_quicksettings:
res = comp(label=info.label, value=fun(), elem_id=elem_id, **args)
ui_common.create_refresh_button(res, info.refresh, info.component_args, f"settings_{key}_refresh")
refresh_button = ui_common.create_refresh_button(res, info.refresh, info.component_args, f"settings_{key}_refresh", visible=args.get("visible", True))
if key in dynamic_visibility_keys:
dynamic_visibility.setdefault(key, []).append(refresh_button)
else:
with gr.Row():
res = comp(label=info.label, value=fun(), elem_id=elem_id, **args)
@@ -202,7 +217,7 @@ def run_settings_single(value, key, progress=False, force=False):
shared.opts.save(silent=True)
if key == 'sd_text_encoder':
sd_models.reload_text_encoder() # apply the change now; reloads the model for encoders with no in-place swap
if key not in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_vae', 'sd_te', 'sd_unet'] or force:
if key not in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_vae', 'sd_te', 'sd_unet', 'sd_unet_secondary'] or force:
log.debug(f'Setting changed: {key}="{value}" progress={progress} force={force}')
return get_value_for_setting(key), shared.opts.dumpjson()
@@ -267,6 +282,7 @@ def create_ui(disabled_tabs=None):
with gr.Tabs(elem_id="settings"):
quicksettings_list.clear()
dynamic_visibility.clear()
for (section_id, section_text) in sections:
items = [item for item in shared.opts.data_labels.items() if item[1].section[0] == section_id] # find all items in this section
hidden = section_id is None or 'hidden' in section_id.lower() or 'hidden' in section_text.lower()
@@ -364,7 +380,7 @@ def create_quicksettings(interfaces):
with gr.Row(elem_id="quicksettings", variant="compact"):
quicksetting_components = []
quicksetting_keys = []
for k, _item in sorted(quicksettings_list, key=lambda x: quicksettings_names.get(x[1], x[0])):
for k, _item in sorted(quicksettings_list, key=lambda x: quicksettings_names.get(x[0], 0)):
component = create_setting_component(k, is_quicksettings=True)
quicksetting_components.append(component)
quicksetting_keys.append(k)
@@ -394,10 +410,20 @@ def create_quicksettings(interfaces):
gr.Audio(interactive=False, value=os.path.join(paths.script_path, shared.opts.notification_audio_path), elem_id="audio_notification", visible=False)
def sync_checkpoint_components(value, progress=False, force=False):
# a checkpoint change can reset sd_unet / sd_text_encoder to Default (arch changed);
# push both back so the dropdowns reflect it, not just the stored option
# a checkpoint change can reset sd_unet / sd_unet_secondary / sd_text_encoder
# (arch changed); push the current values back to the dropdowns. The secondary
# dropdown and its dynamic_visibility companions also toggle visibility here,
# since get_value_for_setting strips 'visible'.
from modules import sd_unet, shared_items
checkpoint_update, settings_text = run_settings_single(value, key='sd_model_checkpoint', progress=progress, force=force)
return checkpoint_update, get_value_for_setting('sd_unet'), get_value_for_setting('sd_text_encoder'), settings_text
secondary_visible = shared.sd_model_type in sd_unet.DUAL_TRANSFORMER_TYPES
secondary_update = gr.update(
value=shared.opts.sd_unet_secondary,
choices=shared_items.sd_unet_items(),
visible=secondary_visible,
)
companion_updates = [gr.update(visible=secondary_visible) for _ in dynamic_visibility.get('sd_unet_secondary', [])]
return checkpoint_update, get_value_for_setting('sd_unet'), secondary_update, *companion_updates, get_value_for_setting('sd_text_encoder'), settings_text
for k, _item in quicksettings_list:
component = shared.settings_components[k]
@@ -417,7 +443,7 @@ def create_quicksettings(interfaces):
if k == 'sd_model_checkpoint':
def fn(value, progress=progress_flag):
return sync_checkpoint_components(value, progress=progress)
outputs = [component, shared.settings_components['sd_unet'], shared.settings_components['sd_text_encoder'], text_settings]
outputs = [component, shared.settings_components['sd_unet'], shared.settings_components['sd_unet_secondary'], *dynamic_visibility.get('sd_unet_secondary', []), shared.settings_components['sd_text_encoder'], text_settings]
else:
def fn(value, k=k, progress=progress_flag):
return run_settings_single(value, key=k, progress=progress)
@@ -438,7 +464,7 @@ def create_quicksettings(interfaces):
fn=sync_checkpoint_components_forced,
_js="consumeDesiredCheckpointName",
inputs=[shared.settings_components['sd_model_checkpoint'], dummy_component],
outputs=[shared.settings_components['sd_model_checkpoint'], shared.settings_components['sd_unet'], shared.settings_components['sd_text_encoder'], text_settings],
outputs=[shared.settings_components['sd_model_checkpoint'], shared.settings_components['sd_unet'], shared.settings_components['sd_unet_secondary'], *dynamic_visibility.get('sd_unet_secondary', []), shared.settings_components['sd_text_encoder'], text_settings],
)
button_set_refiner = gr.Button('Change refiner', elem_id='change_refiner', visible=False)
button_set_refiner.click(