diff --git a/TODO.md b/TODO.md index 309c274fb..474717811 100644 --- a/TODO.md +++ b/TODO.md @@ -20,6 +20,8 @@ Stuff to be added... - Monitor file changes by misbehaving extensions - Kitchen theme: - Lightbox improvements +- Check duplicate extensions +- Reload browser on server restart ## Investigate diff --git a/extensions-builtin/SwinIR/scripts/swinir_model.py b/extensions-builtin/SwinIR/scripts/swinir_model.py index 86672cd9a..619f52e6d 100644 --- a/extensions-builtin/SwinIR/scripts/swinir_model.py +++ b/extensions-builtin/SwinIR/scripts/swinir_model.py @@ -4,10 +4,10 @@ import torch from PIL import Image from basicsr.utils.download_util import load_file_from_url from tqdm import tqdm -from rich import print, progress # pylint: disable=redefined-builtin +from rich import progress from modules import modelloader, devices, script_callbacks, shared -from modules.shared import cmd_opts, opts, state +from modules.shared import opts, state from swinir_model_arch import SwinIR as net from swinir_model_arch_v2 import Swin2SR as net2 from modules.upscaler import Upscaler, UpscalerData @@ -88,7 +88,7 @@ class UpscalerSwinIR(Upscaler): params = "params_ema" with progress.open(filename, 'rb', description=f'Loading weights: [cyan]{filename}', auto_refresh=True) as f: - pretrained_model = torch.load(filename) + pretrained_model = torch.load(f) if params is not None and params in pretrained_model: model.load_state_dict(pretrained_model[params], strict=True) else: @@ -151,7 +151,7 @@ def inference(img, model, tile, tile_overlap, window_size, scale): for w_idx in w_idx_list: if state.interrupted or state.skipped: break - + in_patch = img[..., h_idx: h_idx + tile, w_idx: w_idx + tile] out_patch = model(in_patch) out_patch_mask = torch.ones_like(out_patch) diff --git a/extensions-builtin/a1111-sd-webui-lycoris b/extensions-builtin/a1111-sd-webui-lycoris index b2a4e5f92..514511d72 160000 --- a/extensions-builtin/a1111-sd-webui-lycoris +++ b/extensions-builtin/a1111-sd-webui-lycoris @@ -1 +1 @@ -Subproject commit b2a4e5f9292ab0f4cb17739afc7af5d3a713eb54 +Subproject commit 514511d7260635e0eb7b67cabcbce2a484387a97 diff --git a/extensions-builtin/multidiffusion-upscaler-for-automatic1111 b/extensions-builtin/multidiffusion-upscaler-for-automatic1111 index eada8e510..0f55e98e2 160000 --- a/extensions-builtin/multidiffusion-upscaler-for-automatic1111 +++ b/extensions-builtin/multidiffusion-upscaler-for-automatic1111 @@ -1 +1 @@ -Subproject commit eada8e5101417f60850efa15f7454a826da6cf0d +Subproject commit 0f55e98e27235984a31fbd38287ba6584c4884c5 diff --git a/html/card-no-preview.png b/html/card-no-preview.png index e2beb2692..952d30580 100644 Binary files a/html/card-no-preview.png and b/html/card-no-preview.png differ diff --git a/javascript/style.css b/javascript/style.css index dee5b3704..523e2df66 100644 --- a/javascript/style.css +++ b/javascript/style.css @@ -562,7 +562,7 @@ footer { height: 9em; width: 9em; cursor: pointer; - background-image: url('./file=html/card-no-preview.png'); + background-image: url('../html/card-no-preview.png'); background-size: cover; background-position: center center; position: relative; @@ -592,11 +592,11 @@ footer { box-shadow: 0 0 5px rgba(128, 128, 128, 0.5); border-radius: 0.2em; position: relative; - background-size: auto 100%; + background-size: cover; background-position: center; overflow: hidden; cursor: pointer; - background-image: url('./file=html/card-no-preview.png') + background-image: url('../html/card-no-preview.png') } .extra-network-cards .card:hover { box-shadow: 0 0 2px 0.3em rgba(0, 128, 255, 0.35); } diff --git a/javascript/ui.js b/javascript/ui.js index 3825935aa..8e864575f 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -387,7 +387,6 @@ function reconnect_ui() { loadingStarted = Date.now(); loadingMonitor = setInterval(() => { elapsed = Date.now() - loadingStarted; - console.log('Loading', elapsed) if (elapsed > 3000 && loading) loading.style.display = 'none'; }, 5000); } diff --git a/modules/extra_networks.py b/modules/extra_networks.py index 3170bed4a..9e886219a 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -74,7 +74,7 @@ def activate(p, extra_network_data): try: extra_network.activate(p, extra_network_args) except Exception as e: - errors.display(e, f"Error activating extra network {extra_network_name} with arguments {extra_network_args}") + errors.display(e, f"activating extra network {extra_network_name} with arguments {extra_network_args}") for extra_network_name, extra_network in extra_network_registry.items(): args = extra_network_data.get(extra_network_name, None) @@ -84,7 +84,7 @@ def activate(p, extra_network_data): try: extra_network.activate(p, []) except Exception as e: - errors.display(e, f"Error activating extra network {extra_network_name}") + errors.display(e, f"activating extra network {extra_network_name}") def deactivate(p, extra_network_data): @@ -99,7 +99,7 @@ def deactivate(p, extra_network_data): try: extra_network.deactivate(p) except Exception as e: - errors.display(e, f"Error deactivating extra network {extra_network_name}") + errors.display(e, f"deactivating extra network {extra_network_name}") for extra_network_name, extra_network in extra_network_registry.items(): args = extra_network_data.get(extra_network_name, None) @@ -109,7 +109,7 @@ def deactivate(p, extra_network_data): try: extra_network.deactivate(p) except Exception as e: - errors.display(e, f"Error deactivating unmentioned extra network {extra_network_name}") + errors.display(e, f"deactivating unmentioned extra network {extra_network_name}") re_extra_net = re.compile(r"<(\w+):([^>]+)>") diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index 3f8b5fbf0..32c12d54b 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -143,8 +143,8 @@ def connect_paste_params_buttons(): binding.paste_button.click( fn=None, _js=f"switch_to_{binding.tabname}", - inputs=None, - outputs=None, + inputs=[], + outputs=[], ) @@ -343,7 +343,7 @@ def create_override_settings_dict(text_pairs): return res -def connect_paste(button, paste_fields, input_comp, override_settings_component, tabname): # pylint: disable=redefined-outer-name +def connect_paste(button, local_paste_fields, input_comp, override_settings_component, tabname): def paste_func(prompt): if 'Negative prompt' not in prompt and 'Steps' not in prompt: prompt = None @@ -357,7 +357,7 @@ def connect_paste(button, paste_fields, input_comp, override_settings_component, params = parse_generation_parameters(prompt) script_callbacks.infotext_pasted_callback(prompt, params) res = [] - for output, key in paste_fields: + for output, key in local_paste_fields: if callable(key): v = key(params) else: @@ -394,12 +394,12 @@ def connect_paste(button, paste_fields, input_comp, override_settings_component, vals[param_name] = v vals_pairs = [f"{k}: {v}" for k, v in vals.items()] return gr.Dropdown.update(value=vals_pairs, choices=vals_pairs, visible=len(vals_pairs) > 0) - paste_fields = paste_fields + [(override_settings_component, paste_settings)] + local_paste_fields = local_paste_fields + [(override_settings_component, paste_settings)] button.click( fn=paste_func, inputs=[input_comp], - outputs=[x[0] for x in paste_fields], + outputs=[x[0] for x in local_paste_fields], ) button.click( fn=None, diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index 9500c0410..d13b811d5 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -1,24 +1,21 @@ -import csv import datetime import glob import html import os -import sys +from collections import deque import inspect - -import modules.textual_inversion.dataset -import torch +from statistics import stdev, mean +from rich import progress import tqdm +import torch +from torch import einsum +from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ from einops import rearrange, repeat from ldm.util import default from modules import devices, processing, sd_models, shared, sd_samplers, hashes, sd_hijack_checkpoint, errors +import modules.textual_inversion.dataset from modules.textual_inversion import textual_inversion, logging from modules.textual_inversion.learn_schedule import LearnRateScheduler -from torch import einsum -from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ - -from collections import defaultdict, deque -from statistics import stdev, mean optimizer_dict = {optim_name : cls_obj for optim_name, cls_obj in inspect.getmembers(torch.optim, inspect.isclass) if optim_name != "Optimizer"} @@ -245,7 +242,8 @@ class Hypernetwork: if self.name is None: self.name = os.path.splitext(os.path.basename(filename))[0] - state_dict = torch.load(filename, map_location='cpu') + with progress.open(filename, 'rb', description=f'Loading hypernetwork: [cyan]{filename}', auto_refresh=True) as f: + state_dict = torch.load(f, map_location='cpu') self.layer_structure = state_dict.get('layer_structure', [1, 2, 1]) self.optional_info = state_dict.get('optional_info', None) @@ -539,7 +537,7 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi return hypernetwork, filename scheduler = LearnRateScheduler(learn_rate, steps, initial_step) - + clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else None if clip_grad: clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False) @@ -595,7 +593,7 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi scaler = torch.xpu.amp.GradScaler() else: scaler = torch.cuda.amp.GradScaler() - + batch_size = ds.batch_size gradient_step = ds.gradient_step # n steps = batch_size * gradient_step * n image processed @@ -638,7 +636,7 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi if clip_grad: clip_grad_sched.step(hypernetwork.step) - + with devices.autocast(): x = batch.latent_sample.to(devices.device, non_blocking=pin_memory) if use_weight: @@ -659,14 +657,14 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi _loss_step += loss.item() scaler.scale(loss).backward() - + # go back until we reach gradient accumulation steps if (j + 1) % gradient_step != 0: continue loss_logging.append(_loss_step) if clip_grad: clip_grad(weights, clip_grad_sched.learn_rate) - + scaler.step(optimizer) scaler.update() hypernetwork.step += 1 @@ -674,9 +672,7 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi optimizer.zero_grad(set_to_none=True) loss_step = _loss_step _loss_step = 0 - steps_done = hypernetwork.step + 1 - epoch_num = hypernetwork.step // steps_per_epoch epoch_step = hypernetwork.step % steps_per_epoch diff --git a/modules/processing.py b/modules/processing.py index 6c2877576..efba7f01e 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -492,11 +492,13 @@ def process_images(p: StableDiffusionProcessing) -> Processed: try: # if no checkpoint override or the override checkpoint can't be found, remove override entry and load opts checkpoint - if sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None: + if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None: p.override_settings.pop('sd_model_checkpoint', None) sd_models.reload_model_weights() for k, v in p.override_settings.items(): setattr(opts, k, v) + if k == 'sd_model_checkpoint': + sd_models.reload_model_weights() if k == 'sd_vae': sd_vae.reload_vae_weights() diff --git a/modules/scripts.py b/modules/scripts.py index 1f244c66c..ac3c38a65 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -260,6 +260,14 @@ class ScriptRunner: def initialize_scripts(self, is_img2img): from modules import scripts_auto_postprocessing + self.scripts.clear() + self.selectable_scripts.clear() + self.alwayson_scripts.clear() + self.titles.clear() + self.infotext_fields.clear() + self.paste_field_names.clear() + self.script_load_ctr = 0 + self.scripts.clear() self.alwayson_scripts.clear() self.selectable_scripts.clear() @@ -429,7 +437,6 @@ class ScriptRunner: self.scripts[si].args_from = args_from self.scripts[si].args_to = args_to - scripts_txt2img = ScriptRunner() scripts_img2img = ScriptRunner() scripts_postproc = scripts_postprocessing.ScriptPostprocessingRunner() diff --git a/modules/sd_models.py b/modules/sd_models.py index f422b0133..c8790b53e 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -212,18 +212,18 @@ def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unuse pl_sd = None with progress.open(checkpoint_file, 'rb', description=f'Loading weights: [cyan]{checkpoint_file}', auto_refresh=True) as f: _, extension = os.path.splitext(checkpoint_file) - if 'v1-5-pruned-emaonly.safetensors' in checkpoint_file and not shared.opts.stream_load: - if extension.lower() == ".safetensors": - pl_sd = safetensors.torch.load_file(checkpoint_file, device='cpu') - else: - pl_sd = torch.load(checkpoint_file, map_location='cpu') - else: + if shared.opts.stream_load: if extension.lower() == ".safetensors": buffer = f.read() pl_sd = safetensors.torch.load(buffer) else: buffer = io.BytesIO(f.read()) pl_sd = torch.load(buffer, map_location='cpu') + else: + if extension.lower() == ".safetensors": + pl_sd = safetensors.torch.load_file(checkpoint_file, device='cpu') + else: + pl_sd = torch.load(f, map_location='cpu') sd = get_state_dict_from_checkpoint(pl_sd) del pl_sd except Exception as e: @@ -342,10 +342,9 @@ sd2_clip_weight = 'cond_stage_model.model.transformer.resblocks.0.attn.in_proj_w def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None): - shared.debug(f'Load model: {checkpoint_info}') + shared.debug(f'Load model: {checkpoint_info} {already_loaded_state_dict}') from modules import lowvram, sd_hijack checkpoint_info = checkpoint_info or select_checkpoint() - do_inpainting_hijack() if timer is None: timer = Timer() current_checkpoint_info = None @@ -353,9 +352,10 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None) current_checkpoint_info = shared.sd_model.sd_checkpoint_info sd_hijack.model_hijack.undo_hijack(shared.sd_model) shared.sd_model = None - gc.collect() - devices.torch_gc() + gc.collect() + devices.torch_gc() shared.debug(f'Model unloaded: {memory_stats()}') + do_inpainting_hijack() if already_loaded_state_dict is not None: state_dict = already_loaded_state_dict else: @@ -380,6 +380,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None) sd_model = instantiate_from_config(sd_config.model) except Exception: sd_model = instantiate_from_config(sd_config.model) + # sd_model = instantiate_from_config(sd_config.model) sd_model.used_config = checkpoint_config timer.record("create") load_model_weights(sd_model, checkpoint_info, state_dict, timer) @@ -403,6 +404,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None) script_callbacks.model_loaded_callback(sd_model) timer.record("callbacks") shared.log.info(f"Model loaded in {timer.summary()}") + gc.collect() shared.debug(f'Model load finished: {memory_stats()}') return sd_model diff --git a/modules/shared.py b/modules/shared.py index b375dc6d4..de12c0787 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -216,7 +216,6 @@ def list_themes(): def refresh_themes(): - import requests try: req = requests.get('https://huggingface.co/datasets/freddyaboulton/gradio-theme-subdomains/resolve/main/subdomains.json', timeout=5) if req.status_code == 200: @@ -409,7 +408,7 @@ options_templates.update(options_section(('ui', "User interface"), { "font": OptionInfo("", "Font for image grids that have text"), "keyedit_precision_attention": OptionInfo(0.1, "Ctrl+up/down precision when editing (attention:1.1)", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001}), "keyedit_precision_extra": OptionInfo(0.05, "Ctrl+up/down precision when editing ", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001}), - "keyedit_delimiters": OptionInfo(".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters"), + "keyedit_delimiters": OptionInfo(".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters"), # pylint: disable=anomalous-backslash-in-string "quicksettings": OptionInfo("sd_model_checkpoint", "Quicksettings list"), "hidden_tabs": OptionInfo([], "Hidden UI tabs", ui_components.DropdownMulti, lambda: {"choices": [x for x in tab_names]}), "ui_reorder": OptionInfo(", ".join(ui_reorder_categories), "txt2img/img2img UI item order"), @@ -715,6 +714,7 @@ def restart_server(restart=True): demo.server.force_exit = True demo.close(verbose=False) demo.server.close() + demo.fns = [] except: pass if restart: diff --git a/modules/ui.py b/modules/ui.py index 1a4316c7e..8d00e7db5 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -631,8 +631,9 @@ def create_ui(): with gr.Tab(label="Resize to") as tab_scale_to: with FormRow(): with gr.Column(elem_id="img2img_column_size", scale=4): - width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512, elem_id="img2img_width") - height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512, elem_id="img2img_height") + with FormRow(): + width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512, elem_id="img2img_width") + height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512, elem_id="img2img_height") with gr.Column(elem_id="img2img_dimensions_row", scale=1, elem_classes="dimensions-tools"): res_switch_btn = ToolButton(value=switch_values_symbol, elem_id="img2img_res_switch_btn") @@ -664,9 +665,9 @@ def create_ui(): tab_scale_to.select(fn=lambda: 0, inputs=[], outputs=[selected_scale_tab]) tab_scale_by.select(fn=lambda: 1, inputs=[], outputs=[selected_scale_tab]) - with gr.Column(elem_id="img2img_column_batch"): - batch_count = gr.Slider(minimum=1, step=1, label='Batch count', value=1, elem_id="img2img_batch_count") - batch_size = gr.Slider(minimum=1, maximum=8, step=1, label='Batch size', value=1, elem_id="img2img_batch_size") + with FormRow(elem_id="img2img_column_batch"): + batch_count = gr.Slider(minimum=1, step=1, label='Batch count', value=1, elem_id="img2img_batch_count") + batch_size = gr.Slider(minimum=1, maximum=8, step=1, label='Batch size', value=1, elem_id="img2img_batch_size") elif category == "cfg": with FormGroup(): @@ -1627,7 +1628,7 @@ def html_head(): for script in modules.scripts.list_scripts("javascript", ".js"): if script.path == script_js: continue - print(script.path) + shared.log.debug(f'Loading JS script: {script.path}') head += f'\n' for script in modules.scripts.list_scripts("javascript", ".mjs"): head += f'\n' diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index 1c84e3dc7..d919acbfe 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -20,32 +20,22 @@ def check_access(): def apply_and_restart(disable_list, update_list, disable_all): check_access() - disabled = json.loads(disable_list) assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}" - update = json.loads(update_list) assert type(update) == list, f"wrong update_list data for apply_and_restart: {update_list}" - update = set(update) - for ext in extensions.extensions: if ext.name not in update: continue - try: ext.fetch_and_reset_hard() except Exception as e: errors.display(e, f'extensions apply update: {ext.name}') - shared.opts.disabled_extensions = disabled shared.opts.disable_all_extensions = disable_all shared.opts.save(shared.config_filename) - - # shared.state.interrupt() - # shared.state.need_restart = True - # shared.restart_server() - shared.log.warning('Extension list updated - please restart the server') + shared.restart_server(restart=True) def check_updates(_id_task, disable_list): @@ -313,7 +303,7 @@ def create_ui(): with gr.TabItem("Installed", id="installed"): with gr.Row(elem_id="extensions_installed_top"): - apply = gr.Button(value="Apply (restart required)", variant="primary") + apply = gr.Button(value="Apply & restart", variant="primary") check = gr.Button(value="Check for updates") extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "extra", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all") extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False).style(container=False)