add load unet override

This commit is contained in:
Vladimir Mandic
2024-04-28 11:51:08 -04:00
parent 01e2ab051f
commit a26d222cc1
10 changed files with 83 additions and 18 deletions
+1
View File
@@ -100,6 +100,7 @@ def create_paths(opts):
create_path(fix_path('ckpt_dir'))
create_path(fix_path('diffusers_dir'))
create_path(fix_path('vae_dir'))
create_path(fix_path('unet_dir'))
create_path(fix_path('lora_dir'))
create_path(fix_path('embeddings_dir'))
create_path(fix_path('hypernetwork_dir'))
+10 -10
View File
@@ -20,7 +20,7 @@ from omegaconf import OmegaConf
import tomesd
from transformers import logging as transformers_logging
from ldm.util import instantiate_from_config
from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_accelerate
from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_accelerate
from modules.timer import Timer
from modules.memstats import memory_stats
from modules.modeldata import model_data
@@ -1113,6 +1113,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
if hasattr(sd_model, "set_progress_bar_config"):
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)
set_diffuser_options(sd_model, vae, op)
if op == 'refiner' and shared.opts.diffusers_move_refiner:
@@ -1134,15 +1135,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
shared.log.error("Failed to load diffusers model")
errors.display(e, "loading Diffusers model")
if sd_model is not None:
from modules.textual_inversion import textual_inversion
sd_model.embedding_db = textual_inversion.EmbeddingDatabase()
if op == 'refiner':
model_data.sd_refiner = sd_model
else:
model_data.sd_model = sd_model
sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
from modules.textual_inversion import textual_inversion
sd_model.embedding_db = textual_inversion.EmbeddingDatabase()
if op == 'refiner':
model_data.sd_refiner = sd_model
else:
model_data.sd_model = sd_model
sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
timer.record("load")
devices.torch_gc(force=True)
+42
View File
@@ -0,0 +1,42 @@
import os
from modules import shared, devices, files_cache
unet_dict = {}
def load_unet(model):
from diffusers import UNet2DConditionModel
from safetensors.torch import load_file
if shared.opts.sd_unet == 'None':
return
if shared.opts.sd_unet not in list(unet_dict):
shared.log.error(f'UNet model not found: {shared.opts.sd_unet}')
return
if not hasattr(model, 'unet') or model.unet is None:
shared.log.error('UNet not found in current model')
return
config_file = os.path.splitext(unet_dict[shared.opts.sd_unet])[0] + '.json'
if os.path.exists(config_file):
config = shared.readfile(config_file)
else:
config = None
config_file = 'default'
try:
shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" config="{config_file}"')
unet = UNet2DConditionModel.from_config(model.unet.config if config is None else config).to(devices.device, devices.dtype)
state_dict = load_file(unet_dict[shared.opts.sd_unet])
unet.load_state_dict(state_dict)
model.unet = unet.to(devices.device, devices.dtype_unet)
except Exception as e:
unet = None
shared.log.error(f'Failed to load UNet model: {e}')
return
def refresh_unet_list():
unet_dict.clear()
for file in files_cache.list_files(shared.opts.unet_dir, ext_filter=[".safetensors"]):
name = os.path.splitext(os.path.basename(file))[0]
unet_dict[name] = file
shared.log.debug(f'Available UNets: path="{shared.opts.unet_dir}" items={len(unet_dict)}')
+2
View File
@@ -386,6 +386,7 @@ options_templates.update(options_section(('sd', "Execution & Models"), {
"sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", gr.Dropdown, lambda: {"choices": list_checkpoint_tiles()}, refresh=refresh_checkpoints),
"sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints),
"sd_vae": OptionInfo("Automatic", "VAE model", gr.Dropdown, lambda: {"choices": shared_items.sd_vae_items()}, refresh=shared_items.refresh_vae_list),
"sd_unet": OptionInfo("None", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list),
"sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"),
"sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints),
"stream_load": OptionInfo(False, "Load models using stream loading method", gr.Checkbox, {"visible": backend == Backend.ORIGINAL }),
@@ -539,6 +540,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), {
"diffusers_dir": OptionInfo(os.path.join(paths.models_path, 'Diffusers'), "Folder with Huggingface models", folder=True),
"hfcache_dir": OptionInfo(os.path.join(os.path.expanduser('~'), '.cache', 'huggingface', 'hub'), "Folder for Huggingface cache", folder=True),
"vae_dir": OptionInfo(os.path.join(paths.models_path, 'VAE'), "Folder with VAE files", folder=True),
"unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True),
"sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}),
"lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True),
"lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}),
+10
View File
@@ -13,6 +13,16 @@ def refresh_vae_list():
modules.sd_vae.refresh_vae_list()
def sd_unet_items():
import modules.sd_unet
return ["None"] + list(modules.sd_unet.unet_dict)
def refresh_unet_list():
import modules.sd_unet
modules.sd_unet.refresh_unet_list()
def list_crossattention(diffusers=False):
if diffusers:
return [