diff --git a/javascript/sdnext.css b/javascript/sdnext.css
index 8b52b1a47..f79e2898f 100644
--- a/javascript/sdnext.css
+++ b/javascript/sdnext.css
@@ -291,7 +291,8 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
#pnginfo_html_info .gradio-html > div { margin: 0.5em; }
#models_image, #models_image > div { min-height: 0; }
#models_error { font-family: monospace; color: var(--body-text-color-subdued) }
-
+#model_loader_md li { margin-left: -2em; color: var(--body-text-color-subdued); }
+#model_loader_df button { display: none !important; }
/* log monitor */
.log-monitor { display: none; justify-content: unset !important; overflow: hidden; padding: 0; margin-top: auto; font-family: monospace; font-size: var(--text-xxs); }
diff --git a/modules/processing_args.py b/modules/processing_args.py
index a839c9995..884f189ed 100644
--- a/modules/processing_args.py
+++ b/modules/processing_args.py
@@ -14,7 +14,7 @@ from modules.api import helpers
debug_enabled = os.environ.get('SD_DIFFUSERS_DEBUG', None)
-debug_log = shared.log.trace if os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None else lambda *args, **kwargs: None
+debug_log = shared.log.trace if debug_enabled else lambda *args, **kwargs: None
disable_pbar = os.environ.get('SD_DISABLE_PBAR', None) is not None
diff --git a/modules/shared_items.py b/modules/shared_items.py
index c951e201f..c7c7d7dcf 100644
--- a/modules/shared_items.py
+++ b/modules/shared_items.py
@@ -34,10 +34,10 @@ pipelines = {
'Amused': getattr(diffusers, 'AmusedPipeline', None),
# dynamically imported and redefined later
- 'Meissonic': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
- 'OmniGenPipeline': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
- 'InstaFlow': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
- 'SegMoE': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
+ 'Meissonic': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
+ 'OmniGenPipeline': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
+ 'InstaFlow': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
+ 'SegMoE': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
}
onnx_pipelines = {
'ONNX Stable Diffusion': getattr(diffusers, 'OnnxStableDiffusionPipeline', None),
@@ -116,3 +116,16 @@ def get_pipelines():
if k != 'Autodetect' and v is None:
log.error(f'Not available: pipeline={k} diffusers={diffusers.__version__} path={diffusers.__file__}')
return pipelines
+
+
+def get_repo(model):
+ if model == 'StableDiffusionPipeline' or model == 'Stable Diffusion 1.5':
+ return 'stable-diffusion-v1-5/stable-diffusion-v1-5'
+ elif model == 'StableDiffusionXLPipeline' or model == 'Stable Diffusion XL':
+ return 'stabilityai/stable-diffusion-xl-base-1.0'
+ elif model == 'StableDiffusion3Pipeline' or model == 'Stable Diffusion 3.x':
+ return 'stabilityai/stable-diffusion-3.5-medium'
+ elif model == 'FluxPipeline' or model == 'FLUX':
+ return 'black-forest-labs/FLUX.1-dev'
+ else:
+ return None
diff --git a/modules/ui_models.py b/modules/ui_models.py
index 59d8a9196..807f6611a 100644
--- a/modules/ui_models.py
+++ b/modules/ui_models.py
@@ -44,13 +44,17 @@ def create_ui():
model_desc = gr.HTML(value="", elem_id="model_desc")
with gr.Row():
module_headers = ['Module', 'Class', 'Device', 'DType', 'Params', 'Modules', 'Config']
- model_types = ['str', 'str', 'str', 'str', 'number', 'number', 'str']
- model_modules = gr.DataFrame(value=None, label=None, show_label=False, interactive=False, wrap=True, headers=module_headers, datatype=model_types, type='array')
+ module_types = ['str', 'str', 'str', 'str', 'number', 'number', 'str']
+ model_modules = gr.DataFrame(value=None, label=None, show_label=False, interactive=False, wrap=True, headers=module_headers, datatype=module_types, type='array')
with gr.Row():
model_meta = gr.JSON(label="Metadata", value={}, elem_id="model_meta")
model_analyze.click(fn=analyze, inputs=[], outputs=[model_desc, model_modules, model_meta])
+ with gr.Tab(label="Load custom"):
+ from modules import ui_models_load
+ ui_models_load.create_ui()
+
with gr.Tab(label="Merge"):
def sd_model_choices():
return ['None'] + sd_models.checkpoint_titles()
diff --git a/modules/ui_models_load.py b/modules/ui_models_load.py
new file mode 100644
index 000000000..c91392900
--- /dev/null
+++ b/modules/ui_models_load.py
@@ -0,0 +1,109 @@
+import os
+import inspect
+import gradio as gr
+from modules import shared, shared_items
+
+
+debug_enabled = os.environ.get('SD_LOAD_DEBUG', None)
+debug_log = shared.log.trace if debug_enabled else lambda *args, **kwargs: None
+
+
+class Component():
+ def __init__(self, param):
+ self.name = param.name
+ self.cls = param.annotation
+ self.str = str(param.annotation)
+ self.val = param.default if param.default is not inspect.Parameter.empty else None
+ self.enum = None
+ if self.cls in [str, int, float, bool]:
+ self.type = 'variable'
+ elif 'enum' in self.str:
+ self.type = 'enum'
+ self.enum = [v.name for v in self.cls]
+ elif inspect.isclass(param.annotation):
+ self.type = 'class'
+ elif inspect.ismodule(param.annotation):
+ self.type = 'module'
+ elif inspect.isfunction(param.annotation):
+ self.type = 'function'
+ elif 'typing.Optional' in self.str:
+ self.type = 'optional'
+ self.cls = param.annotation.__args__[0]
+ self.str = str(self.cls)
+ self.val = None
+ else:
+ self.type = 'unknown'
+ self.loadable = self.type = 'class' and hasattr(self.cls, 'from_pretrained')
+
+ def __str__(self):
+ return f'name="{self.name}" type={self.type} cls={self.cls} str={self.str} val={self.val} loadable={self.loadable} enum={self.enum}'
+
+ def dataframe(self):
+ return [self.name, self.loadable, self.val, self.str, self.enum]
+
+
+def create_ui():
+ def get_components(cls):
+ if cls is None:
+ return []
+ signature = inspect.signature(cls.__init__, follow_wrapped=True)
+ components = []
+ for param in signature.parameters.values():
+ if param.name == 'self':
+ continue
+ component = Component(param)
+ debug_log(f'Component: {str(component)}')
+ components.append(component.dataframe())
+ return components
+
+ def get_model(model):
+ cls = shared_items.pipelines.get(model, None)
+ name = cls.__name__
+ repo = shared_items.get_repo(name) or shared_items.get_repo(model)
+ link = f'Link
{repo}' if repo else ''
+ components = get_components(cls)
+ shared.log.debug(f'Model select: name="{model}" cls={name} repo={repo} link={link} components={len(components)}')
+ return [name, repo, link, components]
+
+ with gr.Row():
+ model = gr.Dropdown(label="Model type", choices=['Select'] + list(shared_items.pipelines), value='Select')
+ cls = gr.Textbox(label="Model class", placeholder="Class name", interactive=False)
+ with gr.Row():
+ repo = gr.Textbox(label="Model repo", placeholder="Repo name", interactive=True)
+ link = gr.HTML(value="", interactive=False)
+ with gr.Row():
+ headers = ['Name', 'Loadable', 'Value', 'Class', 'Choices']
+ datatype = ['str', 'bool', 'str', 'str', 'str']
+ components = gr.DataFrame(
+ value=None,
+ label=None,
+ show_label=False,
+ interactive=True,
+ wrap=True,
+ headers=headers,
+ datatype=datatype,
+ max_rows=None,
+ max_cols=None,
+ # row_count=(20, 'fixed'),
+ # col_count=(5, 'fixed'),
+ type='array',
+ elem_id="model_loader_df",
+ )
+
+ model.change(get_model, inputs=[model], outputs=[cls, repo, link, components])
+
+ gr.Markdown("""
+ - Model repo is required to access base model config
+ - Default model repo is provided for common models
+ - Model repo can be overriden to any valid repo on huggingface
+ - Any loadable model component without set value will be loaded from default repo
+ - Any loadable model component with set value will be loaded from that value
+ - Value can be local path to safetensors file or path on huggingface
+ """, elem_id="model_loader_md")
+
+ with gr.Row():
+ btn_load_receipe = gr.Button(value="Load receipe") # pylint: disable=unused-variable
+ btn_save_receipe = gr.Button(value="Save receipe") # pylint: disable=unused-variable
+ with gr.Row():
+ btn_load_model = gr.Button(value="Load model") # pylint: disable=unused-variable
+ btn_unload_model = gr.Button(value="Unload model") # pylint: disable=unused-variable