mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
enable sdxl vae
This commit is contained in:
@@ -638,10 +638,11 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
return
|
||||
shared.log.info(f'Loading diffuser {op}: {checkpoint_info.filename}')
|
||||
|
||||
vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename)
|
||||
vae = sd_vae.load_vae_diffusers(None, vae_file, vae_source)
|
||||
if vae is not None:
|
||||
diffusers_load_config["vae"] = vae
|
||||
if op == 'model':
|
||||
vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename)
|
||||
vae = sd_vae.load_vae_diffusers(None, vae_file, vae_source)
|
||||
if vae is not None:
|
||||
diffusers_load_config["vae"] = vae
|
||||
|
||||
if not os.path.isfile(checkpoint_info.path):
|
||||
try:
|
||||
|
||||
+55
-50
@@ -45,39 +45,50 @@ def restore_base_vae(model):
|
||||
|
||||
|
||||
def get_filename(filepath):
|
||||
return os.path.basename(filepath)
|
||||
if filepath.endswith(".json"):
|
||||
return os.path.basename(os.path.dirname(filepath))
|
||||
else:
|
||||
return os.path.basename(filepath)
|
||||
|
||||
|
||||
def refresh_vae_list():
|
||||
global vae_path # pylint: disable=global-statement
|
||||
vae_path = shared.opts.vae_dir
|
||||
vae_dict.clear()
|
||||
vae_paths = [
|
||||
os.path.join(sd_models.model_path, '**/*.vae.ckpt'),
|
||||
os.path.join(sd_models.model_path, '**/*.vae.pt'),
|
||||
os.path.join(sd_models.model_path, '**/*.vae.safetensors'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.ckpt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.pt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.safetensors'),
|
||||
]
|
||||
if shared.opts.ckpt_dir is not None and os.path.isdir(shared.opts.ckpt_dir):
|
||||
vae_paths = []
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
if sd_models.model_path is not None and os.path.isdir(sd_models.model_path):
|
||||
vae_paths += [
|
||||
os.path.join(sd_models.model_path, 'VAE', '**/*.vae.ckpt'),
|
||||
os.path.join(sd_models.model_path, 'VAE', '**/*.vae.pt'),
|
||||
os.path.join(sd_models.model_path, 'VAE', '**/*.vae.safetensors'),
|
||||
]
|
||||
if shared.opts.ckpt_dir is not None and os.path.isdir(shared.opts.ckpt_dir):
|
||||
vae_paths += [
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.ckpt'),
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.pt'),
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.safetensors'),
|
||||
]
|
||||
if shared.opts.vae_dir is not None and os.path.isdir(shared.opts.vae_dir):
|
||||
vae_paths += [
|
||||
os.path.join(shared.opts.vae_dir, '**/*.ckpt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.pt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.safetensors'),
|
||||
]
|
||||
else:
|
||||
vae_paths += [
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.ckpt'),
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.pt'),
|
||||
os.path.join(shared.opts.ckpt_dir, '**/*.vae.safetensors'),
|
||||
]
|
||||
if shared.opts.vae_dir is not None and os.path.isdir(shared.opts.vae_dir):
|
||||
vae_paths += [
|
||||
os.path.join(shared.opts.vae_dir, '**/*.ckpt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.pt'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.safetensors'),
|
||||
os.path.join(sd_models.model_path, 'VAE', '**/*.json'),
|
||||
os.path.join(shared.opts.vae_dir, '**/*.json'),
|
||||
]
|
||||
candidates = []
|
||||
for path in vae_paths:
|
||||
candidates += glob.iglob(path, recursive=True)
|
||||
for filepath in candidates:
|
||||
name = get_filename(filepath)
|
||||
vae_dict[name] = filepath
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
vae_dict[name] = filepath
|
||||
else:
|
||||
vae_dict[name] = os.path.dirname(filepath)
|
||||
shared.log.info(f"Available VAEs: {vae_path} {len(vae_dict)}")
|
||||
|
||||
|
||||
@@ -86,28 +97,28 @@ def find_vae_near_checkpoint(checkpoint_file):
|
||||
for vae_location in [f"{checkpoint_path}.vae.pt", f"{checkpoint_path}.vae.ckpt", f"{checkpoint_path}.vae.safetensors"]:
|
||||
if os.path.isfile(vae_location):
|
||||
return vae_location
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def resolve_vae(checkpoint_file):
|
||||
if shared.cmd_opts.vae is not None:
|
||||
print('HERE1', checkpoint_file, shared.opts.sd_vae)
|
||||
if shared.cmd_opts.vae is not None: # 1st
|
||||
return shared.cmd_opts.vae, 'forced'
|
||||
is_automatic = shared.opts.sd_vae in {"Automatic", "auto"} # "auto" for people with old config
|
||||
vae_near_checkpoint = find_vae_near_checkpoint(checkpoint_file)
|
||||
if vae_near_checkpoint is not None:
|
||||
return vae_near_checkpoint, 'near checkpoint'
|
||||
if is_automatic:
|
||||
basename = os.path.join(vae_path, os.path.splitext(os.path.basename(checkpoint_file))[0])
|
||||
for named_vae_location in [basename + ".pt", basename + ".ckpt", basename + ".safetensors", basename + ".vae.pt", basename + ".vae.ckpt", basename + ".vae.safetensors"]:
|
||||
if os.path.isfile(named_vae_location):
|
||||
return named_vae_location, 'in VAE dir'
|
||||
if shared.opts.sd_vae == "None":
|
||||
if shared.opts.sd_vae == "None": # 2nd
|
||||
return None, None
|
||||
vae_from_options = vae_dict.get(shared.opts.sd_vae, None)
|
||||
if vae_from_options is not None:
|
||||
return vae_from_options, 'specified in settings'
|
||||
if not is_automatic:
|
||||
vae_near_checkpoint = find_vae_near_checkpoint(checkpoint_file)
|
||||
if vae_near_checkpoint is not None: # 3rd
|
||||
return vae_near_checkpoint, 'near checkpoint'
|
||||
if shared.opts.sd_vae == "Automatic": # 4th
|
||||
basename = os.path.splitext(os.path.basename(checkpoint_file))[0]
|
||||
print('HERE2', basename)
|
||||
if vae_dict.get(basename, None) is not None:
|
||||
print('HERE3', vae_dict[basename])
|
||||
return vae_dict[basename], 'in VAE dir'
|
||||
else:
|
||||
vae_from_options = vae_dict.get(shared.opts.sd_vae, None) # 5th
|
||||
if vae_from_options is not None:
|
||||
return vae_from_options, 'specified in settings'
|
||||
shared.log.warning(f"VAE not found: {shared.opts.sd_vae}")
|
||||
return None, None
|
||||
|
||||
@@ -150,18 +161,14 @@ def load_vae(model, vae_file=None, vae_source="from unknown source"):
|
||||
|
||||
|
||||
def load_vae_diffusers(_model, vae_file=None, vae_source="from unknown source"):
|
||||
global loaded_vae_file # pylint: disable=global-statement
|
||||
if loaded_vae_file == vae_file:
|
||||
return
|
||||
loaded_vae_file = None
|
||||
if vae_file is None:
|
||||
return
|
||||
if not os.path.isfile(vae_file):
|
||||
shared.log.error('VAE not found: {vae_file}')
|
||||
return
|
||||
return None
|
||||
if not os.path.exists(vae_file):
|
||||
shared.log.error(f'VAE not found: {vae_file}')
|
||||
return None
|
||||
shared.log.info(f"Loading diffusers VAE: {vae_source}: {vae_file}")
|
||||
diffusers_load_config = {
|
||||
"low_cpu_mem_usage": True,
|
||||
"low_cpu_mem_usage": False,
|
||||
"torch_dtype": devices.dtype_vae,
|
||||
"use_safetensors": True,
|
||||
}
|
||||
@@ -170,11 +177,11 @@ def load_vae_diffusers(_model, vae_file=None, vae_source="from unknown source"):
|
||||
shared.log.debug(f'Diffusers VAE load config: {diffusers_load_config}')
|
||||
try:
|
||||
import diffusers
|
||||
diffusers_vae = diffusers.AutoencoderKL.from_pretrained(vae_file, **diffusers_load_config)
|
||||
vae = diffusers.AutoencoderKL.from_pretrained(vae_file, **diffusers_load_config)
|
||||
return vae
|
||||
except Exception as e:
|
||||
shared.log.error(f"Loading diffusers VAE failed: {vae_file} {e}")
|
||||
diffusers_vae = None
|
||||
return diffusers_vae
|
||||
return None
|
||||
|
||||
|
||||
# don't call this from outside
|
||||
@@ -216,8 +223,6 @@ def reload_vae_weights(sd_model=None, vae_file=unspecified):
|
||||
load_vae(sd_model, vae_file, vae_source)
|
||||
sd_hijack.model_hijack.hijack(sd_model)
|
||||
script_callbacks.model_loaded_callback(sd_model)
|
||||
elif shared.backend == shared.Backend.DIFFUSERS:
|
||||
load_vae_diffusers(sd_model, vae_file, vae_source)
|
||||
|
||||
if not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram:
|
||||
sd_model.to(devices.device)
|
||||
|
||||
+2
-2
@@ -361,7 +361,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), {
|
||||
|
||||
options_templates.update(options_section(('diffusers', "Diffusers Settings"), {
|
||||
"diffusers_allow_safetensors": OptionInfo(True, 'Diffusers allow loading from safetensors files'),
|
||||
"diffusers_pipeline": OptionInfo(pipelines[0], 'Select diffuser pipeline when loading from safetensors', gr.Dropdown, lambda: {"choices": pipelines}),
|
||||
"diffusers_pipeline": OptionInfo(pipelines[0], 'Diffusers pipeline', gr.Dropdown, lambda: {"choices": pipelines}),
|
||||
"diffusers_refiner_latents": OptionInfo(True, "Use latents when using refiner"),
|
||||
"diffusers_move_base": OptionInfo(False, "Move base model to CPU when using refiner"),
|
||||
"diffusers_move_refiner": OptionInfo(True, "Move refiner model to CPU when not in use"),
|
||||
@@ -369,7 +369,7 @@ options_templates.update(options_section(('diffusers', "Diffusers Settings"), {
|
||||
"diffusers_generator_device": OptionInfo("default", "Generator device", gr.Radio, lambda: {"choices": ["default", "cpu"]}),
|
||||
"diffusers_seq_cpu_offload": OptionInfo(False, "Enable sequential CPU offload"),
|
||||
"diffusers_model_cpu_offload": OptionInfo(False, "Enable model CPU offload"),
|
||||
"diffusers_vae_slicing": OptionInfo(False, "Enable VAE slicing"),
|
||||
"diffusers_vae_slicing": OptionInfo(True, "Enable VAE slicing"),
|
||||
"diffusers_vae_tiling": OptionInfo(False, "Enable VAE tiling"),
|
||||
"diffusers_attention_slicing": OptionInfo(False, "Enable attention slicing"),
|
||||
# "diffusers_force_zeros": OptionInfo(False, "Force zeros for prompts when empty"),
|
||||
|
||||
@@ -5,21 +5,19 @@ def realesrgan_models_names():
|
||||
|
||||
def postprocessing_scripts():
|
||||
import modules.scripts
|
||||
|
||||
return modules.scripts.scripts_postproc.scripts
|
||||
|
||||
|
||||
def sd_vae_items():
|
||||
import modules.sd_vae
|
||||
|
||||
return ["Automatic", "None"] + list(modules.sd_vae.vae_dict)
|
||||
|
||||
|
||||
def refresh_vae_list():
|
||||
import modules.sd_vae
|
||||
|
||||
modules.sd_vae.refresh_vae_list()
|
||||
|
||||
|
||||
def list_crossattention():
|
||||
return [
|
||||
"Disable cross-attention layer optimization",
|
||||
|
||||
+1
-1
@@ -658,7 +658,7 @@ def create_ui(startup_timer = None):
|
||||
with FormRow():
|
||||
cfg_scale = gr.Slider(minimum=1.0, maximum=30.0, step=0.5, label='CFG Scale', value=6.0, elem_id="img2img_cfg_scale")
|
||||
image_cfg_scale = gr.Slider(minimum=0, maximum=30.0, step=0.05, label='Image CFG Scale', value=1.5, elem_id="img2img_image_cfg_scale")
|
||||
diffusers_guidance_rescale = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='CFG Scale', value=0.7, elem_id="txt2img_image_cfg_rescale")
|
||||
diffusers_guidance_rescale = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Guidance Rescale', value=0.7, elem_id="txt2img_image_cfg_rescale")
|
||||
with FormRow():
|
||||
denoising_strength = gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Denoising strength', value=0.75, elem_id="img2img_denoising_strength")
|
||||
refiner_denoise_start = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Denoise start', value=0.0, elem_id="txt2img_refiner_denoise_start")
|
||||
|
||||
Reference in New Issue
Block a user