mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
add load unet override
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)}')
|
||||
@@ -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}),
|
||||
|
||||
@@ -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 [
|
||||
|
||||
Reference in New Issue
Block a user