Make VAE options not require model reload

This commit is contained in:
Disty0
2025-06-10 15:56:19 +03:00
parent 78f99abec8
commit c81b712ddb
3 changed files with 14 additions and 2 deletions
+2 -1
View File
@@ -9,7 +9,8 @@
- Increase the medvram mode threshold from 8GB to 12GB
- Set CPU backend to use FP32 by default
- Relax Python version checks for Zluda
- don't override user set gfx version with ROCm
- Don't override user set gfx version with ROCm
- Make VAE options not require model reload
- **Torch**
- Set default to `torch==2.7.1`
+2
View File
@@ -117,6 +117,7 @@ def full_vae_decode(latents, model):
elif shared.opts.diffusers_offload_mode != "sequential":
sd_models.move_model(model.vae, devices.device)
sd_models.set_vae_options(model, vae=None, op='decode')
upcast = (model.vae.dtype == torch.float16) and (getattr(model.vae.config, 'force_upcast', False) or shared.opts.no_half_vae)
if upcast:
if hasattr(model, 'upcast_vae'): # this is done by diffusers automatically if output_type != 'latent'
@@ -193,6 +194,7 @@ def full_vae_encode(image, model):
vae_name = sd_vae.loaded_vae_file if sd_vae.loaded_vae_file is not None else "default"
log_debug(f'Encode vae="{vae_name}" dtype={model.vae.dtype} upcast={model.vae.config.get("force_upcast", None)}')
sd_models.set_vae_options(model, vae=None, op='encode')
upcast = (model.vae.dtype == torch.float16) and (getattr(model.vae.config, 'force_upcast', False) or shared.opts.no_half_vae)
if upcast:
if hasattr(model, 'upcast_vae'): # this is done by diffusers automatically if output_type != 'latent'
+10 -1
View File
@@ -92,7 +92,7 @@ def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False):
if shared.opts.diffusers_vae_upcast != 'default':
sd_model.vae.config.force_upcast = True if shared.opts.diffusers_vae_upcast == 'true' else False
shared.log.quiet(quiet, f'Setting {op}: component=VAE upcast={sd_model.vae.config.force_upcast}')
if shared.opts.no_half_vae:
if shared.opts.no_half_vae and op not in {'decode', 'encode'}:
devices.dtype_vae = torch.float32
sd_model.vae.to(devices.dtype_vae)
shared.log.quiet(quiet, f'Setting {op}: component=VAE no-half=True')
@@ -105,11 +105,20 @@ def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False):
if hasattr(sd_model, "enable_vae_tiling"):
if shared.opts.diffusers_vae_tiling:
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int):
if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None:
sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size
sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size
sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor
if shared.opts.diffusers_vae_tile_size > 0:
sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size)
sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1)))
else:
sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size)
sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size)
if shared.opts.diffusers_vae_tile_overlap != 0.25:
sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap)
else:
sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor)
shared.log.quiet(quiet, f'Setting {op}: component=VAE tiling=True tile={sd_model.vae.tile_sample_min_size} overlap={sd_model.vae.tile_overlap_factor}')
else:
shared.log.quiet(quiet, f'Setting {op}: component=VAE tiling=True')