vae exception handling

This commit is contained in:
Vladimir Mandic
2024-08-29 17:46:32 -04:00
parent ec85ab408f
commit 2c8cb5cd67
4 changed files with 47 additions and 33 deletions
+4 -4
View File
@@ -448,9 +448,9 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
SD3 = hasattr(pipe, 'text_encoder_3')
prompt, prompt_2, _prompt_3 = split_prompts(prompt, SD3)
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(neg_prompt, SD3)
is_sd3 = hasattr(pipe, 'text_encoder_3')
prompt, prompt_2, _prompt_3 = split_prompts(prompt, is_sd3)
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(neg_prompt, is_sd3)
try:
prompt = pipe.maybe_convert_prompt(prompt, pipe.tokenizer)
neg_prompt = pipe.maybe_convert_prompt(neg_prompt, pipe.tokenizer)
@@ -471,7 +471,7 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
te3_device = pipe.text_encoder_3.device
pipe.text_encoder_3 = pipe.text_encoder_3.to(devices.device)
if SD3:
if is_sd3:
prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sd3(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, use_t5_encoder=bool(pipe.text_encoder_3))
elif 'Flux' in pipe.__class__.__name__:
prompt_embed, positive_pooled = get_weighted_text_embeddings_flux1(pipe=pipe, prompt=prompt, prompt2=prompt_2, device=devices.device)
+36 -27
View File
@@ -734,37 +734,46 @@ def set_diffuser_offload(sd_model, op: str = 'model'):
sd_model.has_accelerate = False
if hasattr(sd_model, "enable_model_cpu_offload"):
if shared.opts.diffusers_offload_mode == "model":
shared.log.debug(f'Setting {op}: enable model CPU offload')
if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner:
shared.opts.diffusers_move_base = False
shared.opts.diffusers_move_unet = False
shared.opts.diffusers_move_refiner = False
shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled')
if not hasattr(sd_model, "_all_hooks") or len(sd_model._all_hooks) == 0: # pylint: disable=protected-access
sd_model.enable_model_cpu_offload(device=devices.device)
else:
sd_model.maybe_free_model_hooks()
sd_model.has_accelerate = True
try:
shared.log.debug(f'Setting {op}: enable model CPU offload')
if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner:
shared.opts.diffusers_move_base = False
shared.opts.diffusers_move_unet = False
shared.opts.diffusers_move_refiner = False
shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled')
if not hasattr(sd_model, "_all_hooks") or len(sd_model._all_hooks) == 0: # pylint: disable=protected-access
sd_model.enable_model_cpu_offload(device=devices.device)
else:
sd_model.maybe_free_model_hooks()
sd_model.has_accelerate = True
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
if hasattr(sd_model, "enable_sequential_cpu_offload"):
if shared.opts.diffusers_offload_mode == "sequential":
shared.log.debug(f'Setting {op}: enable sequential CPU offload')
if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner:
shared.opts.diffusers_move_base = False
shared.opts.diffusers_move_unet = False
shared.opts.diffusers_move_refiner = False
shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled')
if sd_model.has_accelerate:
if op == "vae": # reapply sequential offload to vae
from accelerate import cpu_offload
sd_model.vae.to("cpu")
cpu_offload(sd_model.vae, devices.device, offload_buffers=len(sd_model.vae._parameters) > 0) # pylint: disable=protected-access
try:
shared.log.debug(f'Setting {op}: enable sequential CPU offload')
if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner:
shared.opts.diffusers_move_base = False
shared.opts.diffusers_move_unet = False
shared.opts.diffusers_move_refiner = False
shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled')
if sd_model.has_accelerate:
if op == "vae": # reapply sequential offload to vae
from accelerate import cpu_offload
sd_model.vae.to("cpu")
cpu_offload(sd_model.vae, devices.device, offload_buffers=len(sd_model.vae._parameters) > 0) # pylint: disable=protected-access
else:
pass # do nothing if offload is already applied
else:
pass # do nothing if offload is already applied
else:
sd_model.enable_sequential_cpu_offload(device=devices.device)
sd_model.has_accelerate = True
sd_model.enable_sequential_cpu_offload(device=devices.device)
sd_model.has_accelerate = True
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
if shared.opts.diffusers_offload_mode == "balanced":
sd_model = apply_balanced_offload(sd_model)
try:
sd_model = apply_balanced_offload(sd_model)
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
def apply_balanced_offload(sd_model):
+6 -1
View File
@@ -2,7 +2,7 @@ import os
import glob
from copy import deepcopy
import torch
from modules import shared, paths, devices, script_callbacks, sd_models
from modules import shared, errors, paths, devices, script_callbacks, sd_models
vae_ignore_keys = {"model_ema.decay", "model_ema.num_updates"}
@@ -11,6 +11,7 @@ base_vae = None
loaded_vae_file = None
checkpoint_info = None
vae_path = os.path.abspath(os.path.join(paths.models_path, 'VAE'))
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
def get_base_vae(model):
@@ -154,6 +155,8 @@ def load_vae(model, vae_file=None, vae_source="unknown-source"):
_load_vae_dict(model, vae_dict_1)
except Exception as e:
shared.log.error(f"Loading VAE failed: model={vae_file} source={vae_source} {e}")
if debug:
errors.display(e, 'VAE')
restore_base_vae(model)
vae_opt = get_filename(vae_file)
if vae_opt not in vae_dict:
@@ -229,6 +232,8 @@ def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"):
return vae
except Exception as e:
shared.log.error(f"Loading VAE failed: model={vae_file} {e}")
if debug:
errors.display(e, 'VAE')
return None
+1 -1
Submodule wiki updated: 426ad49241...93959071e1