reorg model load

This commit is contained in:
Vladimir Mandic
2024-05-10 16:53:49 -04:00
parent 1c793e8357
commit d1b0cd61b3
+15 -10
View File
@@ -1129,6 +1129,19 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba')
sd_unet.load_unet(sd_model)
timer.record("load")
if op == 'refiner':
model_data.sd_refiner = sd_model
else:
model_data.sd_model = sd_model
from modules.textual_inversion import textual_inversion
sd_model.embedding_db = textual_inversion.EmbeddingDatabase()
sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
timer.record("embeddings")
set_diffuser_options(sd_model, vae, op)
if op == 'refiner' and shared.opts.diffusers_move_refiner:
@@ -1136,6 +1149,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
move_model(sd_model, devices.cpu)
else:
move_model(sd_model, devices.device)
timer.record("move")
if shared.opts.ipex_optimize:
sd_model = sd_models_compile.ipex_optimize(sd_model)
@@ -1145,22 +1159,13 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
if (shared.opts.cuda_compile and shared.opts.cuda_compile_backend != 'none'):
sd_model = sd_models_compile.compile_diffusers(sd_model)
timer.record("compile")
except Exception as e:
shared.log.error("Failed to load diffusers model")
errors.display(e, "loading Diffusers model")
if op == 'refiner':
model_data.sd_refiner = sd_model
else:
model_data.sd_model = sd_model
timer.record("load")
from modules.textual_inversion import textual_inversion
sd_model.embedding_db = textual_inversion.EmbeddingDatabase()
sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
timer.record("embeddings")
devices.torch_gc(force=True)
if shared.cmd_opts.profile: