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