mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
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:
@@ -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'
|
||||
|
||||
@@ -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"]):
|
||||
|
||||
@@ -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']
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user