mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
add sdxl support
This commit is contained in:
+136
-45
@@ -23,6 +23,7 @@ from modules.timer import Timer
|
||||
from modules.memstats import memory_stats
|
||||
from modules.paths_internal import models_path
|
||||
|
||||
|
||||
transformers_logging.set_verbosity_error()
|
||||
model_dir = "Stable-diffusion"
|
||||
model_path = os.path.abspath(os.path.join(paths.models_path, model_dir))
|
||||
@@ -127,7 +128,9 @@ def list_models():
|
||||
checkpoints_list.clear()
|
||||
checkpoint_aliases.clear()
|
||||
ext_filter=[".safetensors"] if shared.opts.sd_disable_ckpt else [".ckpt", ".safetensors"]
|
||||
model_list = modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])
|
||||
model_list = []
|
||||
if shared.backend == shared.Backend.ORIGINAL or shared.opts.diffusers_pipeline == shared.pipelines[0]:
|
||||
model_list += modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])
|
||||
if shared.backend == shared.Backend.DIFFUSERS:
|
||||
model_list += modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir)
|
||||
|
||||
@@ -210,11 +213,18 @@ def model_hash(filename):
|
||||
return 'NOHASH'
|
||||
|
||||
|
||||
def select_checkpoint(model=True):
|
||||
model_checkpoint = shared.opts.sd_model_checkpoint if model else shared.opts.sd_model_dict
|
||||
def select_checkpoint(op='model'):
|
||||
if op == 'model':
|
||||
model_checkpoint = shared.opts.sd_model_checkpoint
|
||||
elif op == 'dict':
|
||||
model_checkpoint = shared.opts.sd_model_dict
|
||||
elif op == 'refiner':
|
||||
model_checkpoint = shared.opts.data['sd_model_refiner']
|
||||
if model_checkpoint is None or model_checkpoint == 'None':
|
||||
return None
|
||||
checkpoint_info = get_closet_checkpoint_match(model_checkpoint)
|
||||
if checkpoint_info is not None:
|
||||
shared.log.debug(f'Select checkpoint: {checkpoint_info.title if checkpoint_info is not None else None}')
|
||||
shared.log.debug(f'Select checkpoint: {op} {checkpoint_info.title if checkpoint_info is not None else None}')
|
||||
return checkpoint_info
|
||||
if len(checkpoints_list) == 0:
|
||||
shared.log.error("Cannot run without a checkpoint")
|
||||
@@ -458,9 +468,10 @@ sd1_clip_weight = 'cond_stage_model.transformer.text_model.embeddings.token_embe
|
||||
sd2_clip_weight = 'cond_stage_model.model.transformer.resblocks.0.attn.in_proj_weight'
|
||||
|
||||
|
||||
class SdModelData:
|
||||
class ModelData:
|
||||
def __init__(self):
|
||||
self.sd_model = None
|
||||
self.sd_refiner = None
|
||||
self.sd_dict = 'None'
|
||||
self.initial = True
|
||||
self.lock = threading.Lock()
|
||||
@@ -470,9 +481,9 @@ class SdModelData:
|
||||
with self.lock:
|
||||
try:
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
reload_model_weights()
|
||||
reload_model_weights(op='model')
|
||||
elif shared.backend == shared.Backend.DIFFUSERS:
|
||||
load_diffuser()
|
||||
load_diffuser(op='model')
|
||||
else:
|
||||
shared.log.error(f"Unknown Stable Diffusion backend: {shared.backend}")
|
||||
self.initial = False
|
||||
@@ -483,11 +494,31 @@ class SdModelData:
|
||||
return self.sd_model
|
||||
|
||||
def set_sd_model(self, v):
|
||||
shared.log.debug(f"Class model: {v}")
|
||||
self.sd_model = v
|
||||
|
||||
def get_sd_refiner(self):
|
||||
if self.sd_model is None:
|
||||
with self.lock:
|
||||
try:
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
reload_model_weights(op='refiner')
|
||||
elif shared.backend == shared.Backend.DIFFUSERS:
|
||||
load_diffuser(op='refiner')
|
||||
else:
|
||||
shared.log.error(f"Unknown Stable Diffusion backend: {shared.backend}")
|
||||
self.initial = False
|
||||
except Exception as e:
|
||||
shared.log.error("Failed to load stable diffusion model")
|
||||
errors.display(e, "loading stable diffusion model")
|
||||
self.sd_refiner = None
|
||||
return self.sd_refiner
|
||||
|
||||
model_data = SdModelData()
|
||||
def set_sd_refiner(self, v):
|
||||
shared.log.debug(f"Class refiner: {v}")
|
||||
self.sd_refiner = v
|
||||
|
||||
model_data = ModelData()
|
||||
|
||||
class PriorPipeline:
|
||||
def __init__(self, prior, main):
|
||||
@@ -531,7 +562,7 @@ class PriorPipeline:
|
||||
return result
|
||||
|
||||
|
||||
def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None): # pylint: disable=unused-argument
|
||||
def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): # pylint: disable=unused-argument
|
||||
if timer is None:
|
||||
timer = Timer()
|
||||
import logging
|
||||
@@ -548,27 +579,61 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
|
||||
if shared.opts.data['sd_model_checkpoint'] == 'model.ckpt':
|
||||
shared.opts.data['sd_model_checkpoint'] = "runwayml/stable-diffusion-v1-5"
|
||||
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
else:
|
||||
if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
|
||||
sd_model = None
|
||||
try:
|
||||
devices.set_cuda_params() # todo
|
||||
devices.set_cuda_params()
|
||||
if shared.cmd_opts.ckpt is not None and model_data.initial: # initial load
|
||||
model_name = modelloader.find_diffuser(shared.cmd_opts.ckpt)
|
||||
if model_name is not None:
|
||||
shared.log.info(f'Loading diffuser model: {model_name}')
|
||||
shared.log.info(f'Loading diffuser {op}: {model_name}')
|
||||
model_file = modelloader.download_diffusers_model(hub_id=model_name)
|
||||
sd_model = diffusers.DiffusionPipeline.from_pretrained(model_file, **diffusers_load_config)
|
||||
list_models() # rescan for downloaded model
|
||||
checkpoint_info = CheckpointInfo(model_name)
|
||||
|
||||
if sd_model is None:
|
||||
checkpoint_info = checkpoint_info or select_checkpoint()
|
||||
shared.log.info(f'Loading diffuser model: {checkpoint_info.filename}')
|
||||
checkpoint_info = checkpoint_info or select_checkpoint(op=op)
|
||||
if checkpoint_info is None:
|
||||
unload_model_weights(op=op)
|
||||
return
|
||||
shared.log.info(f'Loading diffuser {op}: {checkpoint_info.filename}')
|
||||
if not os.path.isfile(checkpoint_info.path):
|
||||
sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, **diffusers_load_config)
|
||||
else:
|
||||
diffusers_load_config["local_files_only "] = True
|
||||
diffusers_load_config["extract_ema"] = shared.opts.diffusers_extract_ema
|
||||
sd_model = diffusers.StableDiffusionPipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config)
|
||||
try:
|
||||
# pipelines = ['Stable Diffusion', 'Stable Diffusion XL', 'Kandinsky V1', 'Kandinsky V2', 'DeepFloyd IF', 'Shap-E']
|
||||
if shared.opts.diffusers_pipeline == shared.pipelines[0]:
|
||||
pipeline = diffusers.StableDiffusionPipeline
|
||||
elif shared.opts.diffusers_pipeline == shared.pipelines[1]:
|
||||
pipeline = diffusers.StableDiffusionXLPipeline
|
||||
elif shared.opts.diffusers_pipeline == shared.pipelines[2]:
|
||||
pipeline = diffusers.KandinskyPipeline
|
||||
elif shared.opts.diffusers_pipeline == shared.pipelines[3]:
|
||||
pipeline = diffusers.KandinskyV22Pipeline
|
||||
elif shared.opts.diffusers_pipeline == shared.pipelines[4]:
|
||||
pipeline = diffusers.IFPipeline
|
||||
elif shared.opts.diffusers_pipeline == shared.pipelines[5]:
|
||||
pipeline = diffusers.ShapEPipeline
|
||||
else:
|
||||
shared.log.error(f'Diffusers unknown pipeline: {shared.opts.diffusers_pipeline}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'Diffusers failed initializing pipeline: {shared.opts.diffusers_pipeline} {e}')
|
||||
return
|
||||
try:
|
||||
sd_model = pipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config)
|
||||
except Exception as e:
|
||||
shared.log.error(f'Diffusers failed loading model using pipeline: {checkpoint_info.path} {shared.opts.diffusers_pipeline} {e}')
|
||||
return
|
||||
|
||||
if "StableDiffusion" in sd_model.__class__.__name__:
|
||||
pass # scheduler is created on first use
|
||||
@@ -615,7 +680,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
sd_model.unet.to(memory_format=torch.channels_last)
|
||||
if shared.opts.cuda_compile and torch.cuda.is_available():
|
||||
sd_model.to(devices.device)
|
||||
import torch._dynamo # pylint: disable=unused-import
|
||||
import torch._dynamo # pylint: disable=unused-import,redefined-outer-name
|
||||
log_level = logging.WARNING if shared.opts.cuda_compile_verbose else logging.CRITICAL # pylint: disable=protected-access
|
||||
torch._logging.set_logs(dynamo=log_level, aot=log_level, inductor=log_level) # pylint: disable=protected-access
|
||||
torch._dynamo.config.verbose = shared.opts.cuda_compile_verbose # pylint: disable=protected-access
|
||||
@@ -632,7 +697,11 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
except Exception as e:
|
||||
shared.log.error("Failed to load diffusers model")
|
||||
errors.display(e, "loading Diffusers model")
|
||||
shared.sd_model = sd_model
|
||||
|
||||
if op == 'refiner':
|
||||
model_data.sd_refiner = sd_model
|
||||
else:
|
||||
model_data.sd_model = sd_model
|
||||
|
||||
from modules.textual_inversion import textual_inversion
|
||||
embedding_db = textual_inversion.EmbeddingDatabase()
|
||||
@@ -685,7 +754,7 @@ def set_diffuser_pipe(pipe, new_pipe_type):
|
||||
new_pipe.sd_model_checkpoint = sd_model_checkpoint
|
||||
new_pipe.sd_model_hash = sd_model_hash
|
||||
|
||||
shared.sd_model = new_pipe
|
||||
model_data.sd_model = new_pipe
|
||||
shared.log.info(f"Pipeline class changed from {pipe.__class__.__name__} to {new_pipe_cls.__name__}")
|
||||
|
||||
|
||||
@@ -700,21 +769,32 @@ def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType:
|
||||
return DiffusersTaskType.TEXT_2_IMAGE
|
||||
|
||||
|
||||
def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None):
|
||||
def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'):
|
||||
from modules import lowvram, sd_hijack
|
||||
checkpoint_info = checkpoint_info or select_checkpoint()
|
||||
checkpoint_info = checkpoint_info or select_checkpoint(op=op)
|
||||
if checkpoint_info is None:
|
||||
return
|
||||
if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
shared.log.debug(f'Load model: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}')
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
else:
|
||||
if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
shared.log.debug(f'Load {op}: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}')
|
||||
if timer is None:
|
||||
timer = Timer()
|
||||
current_checkpoint_info = None
|
||||
if model_data.sd_model is not None:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
current_checkpoint_info = model_data.sd_model.sd_checkpoint_info
|
||||
unload_model_weights()
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model is not None:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
current_checkpoint_info = model_data.sd_model.sd_checkpoint_info
|
||||
unload_model_weights(op=op)
|
||||
else:
|
||||
if model_data.sd_refiner is not None:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner)
|
||||
current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info
|
||||
unload_model_weights(op=op)
|
||||
|
||||
do_inpainting_hijack()
|
||||
devices.set_cuda_params()
|
||||
if already_loaded_state_dict is not None:
|
||||
@@ -760,7 +840,10 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None)
|
||||
sd_model = torch.xpu.optimize(sd_model, dtype=devices.dtype, auto_kernel_selection=True, optimize_lstm=True,
|
||||
graph_mode=True if shared.opts.cuda_compile and shared.opts.cuda_compile_mode == 'ipex' else False)
|
||||
shared.log.info("Applied IPEX Optimize")
|
||||
model_data.sd_model = sd_model
|
||||
if op == 'refiner':
|
||||
model_data.sd_refiner = sd_model
|
||||
else:
|
||||
model_data.sd_model = sd_model
|
||||
sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True) # Reload embeddings after model load as they may or may not fit the model
|
||||
timer.record("embeddings")
|
||||
script_callbacks.model_loaded_callback(sd_model)
|
||||
@@ -771,7 +854,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None)
|
||||
shared.log.info(f'Model load finished: {memory_stats()} cached={len(checkpoints_loaded.keys())}')
|
||||
|
||||
|
||||
def reload_model_weights(sd_model=None, info=None, reuse_dict=False):
|
||||
def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model'):
|
||||
load_dict = shared.opts.sd_model_dict != model_data.sd_dict
|
||||
global skip_next_load # pylint: disable=global-statement
|
||||
if skip_next_load:
|
||||
@@ -779,15 +862,18 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False):
|
||||
skip_next_load = False
|
||||
return
|
||||
from modules import lowvram, sd_hijack
|
||||
checkpoint_info = info or select_checkpoint(model=not load_dict) # are we selecting model or dictionary
|
||||
next_checkpoint_info = info or select_checkpoint(model=load_dict) if load_dict else None
|
||||
checkpoint_info = info or select_checkpoint(op=op) # are we selecting model or dictionary
|
||||
next_checkpoint_info = info or select_checkpoint(op='dict' if load_dict else 'model') if load_dict else None
|
||||
if checkpoint_info is None:
|
||||
unload_model_weights(op=op)
|
||||
return
|
||||
if load_dict:
|
||||
shared.log.debug(f'Model dict: existing={sd_model is not None} target={checkpoint_info.filename} info={info}')
|
||||
else:
|
||||
model_data.sd_dict = 'None'
|
||||
shared.log.debug(f'Load model weights: existing={sd_model is not None} target={checkpoint_info.filename} info={info}')
|
||||
if not sd_model:
|
||||
sd_model = model_data.sd_model
|
||||
sd_model = model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner
|
||||
if sd_model is None: # previous model load failed
|
||||
current_checkpoint_info = None
|
||||
else:
|
||||
@@ -802,7 +888,7 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False):
|
||||
shared.log.info('Reusing previous model dictionary')
|
||||
sd_hijack.model_hijack.undo_hijack(sd_model)
|
||||
else:
|
||||
unload_model_weights()
|
||||
unload_model_weights(op=op)
|
||||
sd_model = None
|
||||
timer = Timer()
|
||||
state_dict = get_checkpoint_state_dict(checkpoint_info, timer)
|
||||
@@ -811,14 +897,14 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False):
|
||||
if sd_model is None or checkpoint_config != sd_model.used_config:
|
||||
del sd_model
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer)
|
||||
load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op)
|
||||
else:
|
||||
load_diffuser(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer)
|
||||
load_diffuser(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op)
|
||||
if load_dict and next_checkpoint_info is not None:
|
||||
model_data.sd_dict = shared.opts.sd_model_dict
|
||||
shared.opts.data["sd_model_checkpoint"] = next_checkpoint_info.title
|
||||
reload_model_weights(reuse_dict=True) # ok we loaded dict now lets redo and load model on top of it
|
||||
return model_data.sd_model
|
||||
return model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner
|
||||
try:
|
||||
load_model_weights(sd_model, checkpoint_info, state_dict, timer)
|
||||
except Exception:
|
||||
@@ -835,17 +921,22 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False):
|
||||
shared.log.info(f"Weights loaded in {timer.summary()}")
|
||||
|
||||
|
||||
def unload_model_weights(sd_model=None, _info=None):
|
||||
def unload_model_weights(op='model'):
|
||||
from modules import sd_hijack
|
||||
if model_data.sd_model:
|
||||
model_data.sd_model.to(devices.cpu)
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
sd_model = None
|
||||
model_data.sd_model = None
|
||||
devices.torch_gc(force=True)
|
||||
shared.log.debug(f'Model weights unloaded: {memory_stats()}')
|
||||
return sd_model
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model:
|
||||
model_data.sd_model.to(devices.cpu)
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
model_data.sd_model = None
|
||||
else:
|
||||
if model_data.sd_refiner:
|
||||
model_data.sd_refiner.to(devices.cpu)
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner)
|
||||
model_data.sd_refiner = None
|
||||
shared.log.debug(f'Weights unloaded {op}: {memory_stats()}')
|
||||
devices.torch_gc(force=True)
|
||||
|
||||
|
||||
def apply_token_merging(sd_model, token_merging_ratio):
|
||||
|
||||
Reference in New Issue
Block a user