optimize model load

This commit is contained in:
Vladimir Mandic
2023-05-02 21:30:31 -04:00
parent ce16d18a43
commit 6f976c358f
6 changed files with 22 additions and 30 deletions
+13 -7
View File
@@ -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")
+1 -1
View File
@@ -166,7 +166,7 @@ face_restorers = []
def debug(message):
if cmd_opts.debug:
log.info(message)
log.debug(message)
class OptionInfo:
@@ -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