diff --git a/extensions-builtin/Lora/lora.py b/extensions-builtin/Lora/lora.py index 32f55eafe..ee4d91974 100644 --- a/extensions-builtin/Lora/lora.py +++ b/extensions-builtin/Lora/lora.py @@ -269,32 +269,19 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu return current_names = getattr(self, "lora_current_names", ()) - lora_prev_names = getattr(self, "lora_prev_names", ()) wanted_names = tuple((x.name, x.multiplier) for x in loaded_loras) weights_backup = getattr(self, "lora_weights_backup", None) - if weights_backup is None and len(loaded_loras): + if weights_backup is None: if isinstance(self, torch.nn.MultiheadAttention): weights_backup = (self.in_proj_weight.to(devices.cpu, copy=True), self.out_proj.weight.to(devices.cpu, copy=True)) else: weights_backup = self.weight.to(devices.cpu, copy=True) self.lora_weights_backup = weights_backup - elif lora_prev_names != current_names: - self.lora_weights_backup = None - weights_backup = None - elif len(loaded_loras) == 0: - self.lora_weights_backup = None - if current_names != wanted_names or current_names != lora_prev_names: - if weights_backup is not None and current_names != lora_prev_names: - if isinstance(self, torch.nn.MultiheadAttention): - self.in_proj_weight.copy_(weights_backup[0]) - self.out_proj.weight.copy_(weights_backup[1]) - else: - self.weight.copy_(weights_backup) - elif weights_backup is not None and current_names == (): - # print('lora restore weight') + if current_names != wanted_names: + if weights_backup is not None: if isinstance(self, torch.nn.MultiheadAttention): self.in_proj_weight.copy_(weights_backup[0]) self.out_proj.weight.copy_(weights_backup[1]) @@ -327,11 +314,9 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu print(f'failed to calculate lora weights for layer {lora_layer_name}') - setattr(self, "lora_prev_names", current_names) setattr(self, "lora_current_names", wanted_names) - def lora_reset_cached_weight(self: Union[torch.nn.Conv2d, torch.nn.Linear]): setattr(self, "lora_current_names", ()) setattr(self, "lora_weights_backup", None) diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 815b93021..a482867ee 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 815b930217873dc3bd72f7d9b518b017a6404ad8 +Subproject commit a482867ee5e82b08b221c53662ff0c70c2f18d09 diff --git a/modules/sd_models.py b/modules/sd_models.py index cfc2cde78..812356c94 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1,6 +1,5 @@ import collections import os.path -import sys import gc import re import io @@ -57,7 +56,7 @@ class CheckpointInfo: self.metadata = read_metadata_from_safetensors(filename) except Exception as e: errors.display(e, f"reading checkpoint metadata: {filename}") - + def register(self): checkpoints_list[self.title] = self for i in self.ids: @@ -109,7 +108,7 @@ def list_models(): checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt) checkpoint_info.register() shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title - elif shared.cmd_opts.ckpt != shared.default_sd_model_file: + elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None: shared.log.warning(f"Checkpoint not found: {shared.cmd_opts.ckpt}") for filename in sorted(model_list, key=str.lower): checkpoint_info = CheckpointInfo(filename) @@ -347,6 +346,7 @@ 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}') from modules import lowvram, sd_hijack checkpoint_info = checkpoint_info or select_checkpoint() do_inpainting_hijack() @@ -407,16 +407,23 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None) shared.debug(f'Model load finished: {memory_stats()}') return sd_model +skip_next_load = False def reload_model_weights(sd_model=None, info=None): + global skip_next_load # pylint: disable=global-statement + if skip_next_load: + shared.debug('Reload model weights skip') + skip_next_load = False + return + shared.debug(f'Reload model weights: {sd_model} {info}') from modules import lowvram, sd_hijack checkpoint_info = info or select_checkpoint() if not sd_model: sd_model = shared.sd_model - if not shared.opts.model_reuse_dict and sd_model is not None: - sd_model = None - else: + if shared.opts.model_reuse_dict and sd_model is not None: shared.log.info('Reusing previous model dictionary') + else: + sd_model = None if sd_model is None: # previous model load failed current_checkpoint_info = None else: @@ -442,7 +449,6 @@ def reload_model_weights(sd_model=None, info=None): except Exception: shared.log.error("Failed to load checkpoint, restoring previous") load_model_weights(sd_model, current_checkpoint_info, None, timer) - raise finally: sd_hijack.model_hijack.hijack(sd_model) timer.record("hijack") diff --git a/modules/shared.py b/modules/shared.py index 748e4f9f6..499b11b83 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -166,7 +166,7 @@ face_restorers = [] def debug(message): if cmd_opts.debug: - log.info(message) + log.debug(message) class OptionInfo: diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 95ebfb4f0..fc4507ac3 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -239,7 +239,7 @@ class EmbeddingDatabase: displayed_embeddings = (tuple(self.word_embeddings.keys()), tuple(self.skipped_embeddings.keys())) if self.previously_displayed_embeddings != displayed_embeddings: self.previously_displayed_embeddings = displayed_embeddings - shared.log.info(f"Embeddings loaded: {', '.join(self.word_embeddings.keys())} ({len(self.word_embeddings)})") + shared.log.info(f"Embeddings loaded: {len(self.word_embeddings)} {[k for k in self.word_embeddings.keys()]}") if len(self.skipped_embeddings) > 0: shared.log.info(f"Textual inversion embeddings skipped({len(self.skipped_embeddings)}): {', '.join(self.skipped_embeddings.keys())}") @@ -351,7 +351,7 @@ def validate_train_inputs(model_name, learn_rate, batch_size, gradient_step, dat assert log_directory, "Log directory is empty" -def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused_argument +def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument save_embedding_every = save_embedding_every or 0 create_image_every = create_image_every or 0 diff --git a/webui.py b/webui.py index 60b75915e..5b2157f0d 100644 --- a/webui.py +++ b/webui.py @@ -60,7 +60,7 @@ import modules.hypernetworks.hypernetwork from modules.middleware import setup_middleware startup_timer.record("libraries") - +log.info('Libraries loaded') log.setLevel(logging.DEBUG if cmd_opts.debug else logging.INFO) logging.disable(logging.NOTSET if cmd_opts.debug else logging.DEBUG) if cmd_opts.server_name: @@ -150,6 +150,7 @@ def load_model(): shared.state.job = 'load model' try: modules.sd_models.load_model() + modules.sd_models.skip_next_load = True except Exception as e: errors.display(e, "loading stable diffusion model") log.error("Stable diffusion model failed to load")