force apply vae config on model load

This commit is contained in:
Vladimir Mandic
2024-06-13 20:13:18 -04:00
parent c31b19cb57
commit 3f1c236ce4
3 changed files with 32 additions and 6 deletions
+4
View File
@@ -1,5 +1,9 @@
# Change Log for SD.Next
## Update for 2024-06-14
- force apply vae config on model load
## Update for 2024-06-13
### Highlights for 2024-06-13
+7 -4
View File
@@ -545,7 +545,7 @@ def change_backend():
refresh_vae_list()
def detect_pipeline(f: str, op: str = 'model', warning=True):
def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
guess = shared.opts.diffusers_pipeline
warn = shared.log.warning if warning else lambda *args, **kwargs: None
size = 0
@@ -642,7 +642,8 @@ def detect_pipeline(f: str, op: str = 'model', warning=True):
guess = 'Stable Diffusion XL Instruct'
# get actual pipeline
pipeline = shared_items.get_pipelines().get(guess, None)
shared.log.info(f'Autodetect: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
if not quiet:
shared.log.info(f'Autodetect: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
except Exception as e:
shared.log.error(f'Error detecting diffusers pipeline: model={f} {e}')
return None, None
@@ -650,7 +651,8 @@ def detect_pipeline(f: str, op: str = 'model', warning=True):
try:
size = round(os.path.getsize(f) / 1024 / 1024)
pipeline = shared_items.get_pipelines().get(guess, None)
shared.log.info(f'Diffusers: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
if not quiet:
shared.log.info(f'Diffusers: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
except Exception as e:
shared.log.error(f'Error loading diffusers pipeline: model={f} {e}')
@@ -1157,7 +1159,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
timer.record("embeddings")
set_diffuser_options(sd_model, vae, op)
if op == 'model':
sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model)
if op == 'refiner' and shared.opts.diffusers_move_refiner:
shared.log.debug('Moving refiner model to CPU')
move_model(sd_model, devices.cpu)
+21 -2
View File
@@ -155,8 +155,6 @@ def load_vae(model, vae_file=None, vae_source="unknown-source"):
except Exception as e:
shared.log.error(f"Loading VAE failed: model={vae_file} source={vae_source} {e}")
restore_base_vae(model)
# If vae used is not in dict, update it
# It will be removed on refresh though
vae_opt = get_filename(vae_file)
if vae_opt not in vae_dict:
vae_dict[vae_opt] = vae_file
@@ -165,6 +163,26 @@ def load_vae(model, vae_file=None, vae_source="unknown-source"):
loaded_vae_file = vae_file
def apply_vae_config(model_file, vae_file, sd_model):
def get_vae_config():
config_file = os.path.join(paths.sd_configs_path, os.path.splitext(os.path.basename(model_file))[0] + '_vae.json')
if config_file is not None and os.path.exists(config_file):
return shared.readfile(config_file)
config_file = os.path.join(paths.sd_configs_path, os.path.splitext(os.path.basename(vae_file))[0] + '.json') if vae_file else None
if config_file is not None and os.path.exists(config_file):
return shared.readfile(config_file)
config_file = os.path.join(paths.sd_configs_path, shared.sd_model_type, 'vae', 'config.json')
if config_file is not None and os.path.exists(config_file):
return shared.readfile(config_file)
return {}
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config'):
config = get_vae_config()
for k, v in config.items():
if k in sd_model.vae.config and not k.startswith('_'):
sd_model.vae.config[k] = v
def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"):
if vae_file is None:
return None
@@ -262,6 +280,7 @@ def reload_vae_weights(sd_model=None, vae_file=unspecified):
vae = load_vae_diffusers(shared.sd_model.sd_checkpoint_info.filename, vae_file, vae_source)
if vae is not None:
sd_models.set_diffuser_options(sd_model, vae=vae, op='vae')
apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model)
if not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram:
sd_models.move_model(sd_model, devices.device)