mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
optimize model load
This commit is contained in:
+13
-7
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user