mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
custom vae loader
Co-authored-by: Copilot <copilot@github.com> Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -9,6 +9,8 @@
|
||||
- option *inputs -> skip processing* to force images to passed to model as-is without any pre-processing
|
||||
examples of models that support multi-inputs: *qwen-image-edit, flux.2, google-gemini*
|
||||
- **Anima** support for *img2img* and *inpaint* workflows
|
||||
- custom **VAE** loader for all pipelines
|
||||
*note*: vae still needs to be compatible with the model
|
||||
- **UI**
|
||||
- add button to manually reorient input/output panels
|
||||
- all ui panels can be minimized/maximized by clicking on their header
|
||||
|
||||
+26
-13
@@ -169,25 +169,38 @@ def load_vae(model_file, vae_file=None, vae_source="unknown-source"):
|
||||
vae_config = sd_detect.get_load_config(model_file, model_type, config_type='json')
|
||||
if vae_config is not None:
|
||||
diffusers_load_config['config'] = os.path.join(vae_config, 'vae')
|
||||
log.info(f'Load module: type=VAE model="{vae_file}" source={vae_source} config={diffusers_load_config}')
|
||||
try:
|
||||
import diffusers
|
||||
if os.path.isfile(vae_file):
|
||||
vae_class = None
|
||||
vae_loader = None
|
||||
if shared.sd_loaded and getattr(shared.sd_model, 'vae', None) is not None:
|
||||
vae_class = shared.sd_model.vae.__class__
|
||||
vae_loader = vae_class.from_single_file if os.path.isfile(vae_file) else vae_class.from_pretrained
|
||||
elif os.path.isfile(vae_file):
|
||||
if os.path.getsize(vae_file) > 1310944880: # 1.3GB
|
||||
vae = diffusers.ConsistencyDecoderVAE.from_pretrained('openai/consistency-decoder', **diffusers_load_config) # consistency decoder does not have from single file, so we'll just download it once more
|
||||
vae_class = diffusers.ConsistencyDecoderVAE
|
||||
vae_loader = vae_class.from_pretrained
|
||||
vae_file = 'openai/consistency-decoder'
|
||||
elif os.path.getsize(vae_file) < 10000000: # 10MB
|
||||
vae = diffusers.AutoencoderTiny.from_single_file(vae_file, **diffusers_load_config)
|
||||
else:
|
||||
vae = diffusers.AutoencoderKL.from_single_file(vae_file, **diffusers_load_config)
|
||||
if getattr(vae.config, 'scaling_factor', 0) == 0.18125 and shared.sd_model_type == 'sdxl':
|
||||
vae.config.scaling_factor = 0.13025
|
||||
log.debug('Setting model: component=VAE fix scaling factor')
|
||||
vae = vae.to(devices.dtype_vae)
|
||||
vae_class = diffusers.AutoencoderTiny
|
||||
vae_loader = vae_class.from_single_file
|
||||
else: # fallback
|
||||
vae_class = diffusers.AutoencoderKL
|
||||
# if getattr(vae.config, 'scaling_factor', 0) == 0.18125 and shared.sd_model_type == 'sdxl':
|
||||
# vae.config.scaling_factor = 0.13025
|
||||
# log.debug('Setting model: component=VAE fix scaling factor')
|
||||
vae_loader = vae_class.from_single_file
|
||||
else:
|
||||
if 'consistency-decoder' in vae_file:
|
||||
vae = diffusers.ConsistencyDecoderVAE.from_pretrained(vae_file, **diffusers_load_config)
|
||||
else:
|
||||
vae = diffusers.AutoencoderKL.from_pretrained(vae_file, **diffusers_load_config)
|
||||
vae_class = diffusers.ConsistencyDecoderVAE
|
||||
else: # fallback
|
||||
vae_class = diffusers.AutoencoderKL
|
||||
vae_loader = vae_class.from_pretrained
|
||||
if vae_loader is not None:
|
||||
log.info(f'Load module: type=VAE model="{vae_file}" source={vae_source} cls={vae_class.__name__} config={diffusers_load_config}')
|
||||
vae = vae_loader(vae_file, **diffusers_load_config)
|
||||
vae = vae.to(devices.dtype_vae)
|
||||
|
||||
global loaded_vae_file # pylint: disable=global-statement
|
||||
loaded_vae_file = os.path.basename(vae_file)
|
||||
# log.debug(f'Diffusers VAE config: {vae.config}')
|
||||
|
||||
@@ -250,3 +250,33 @@ def load_text_encoder(repo_id, cls_name, load_config=None, subfolder="text_encod
|
||||
devices.torch_gc()
|
||||
shared.state.end(jobid)
|
||||
return text_encoder
|
||||
|
||||
|
||||
def load_vae_override(pipe, load_config=None, override_cls=None, override_args={}):
|
||||
if (shared.opts.sd_vae in [None, 'None', 'Default', 'Automatic']):
|
||||
return
|
||||
if (pipe is None) or (getattr(pipe, 'vae', None) is None):
|
||||
return
|
||||
if load_config is None:
|
||||
load_config = {}
|
||||
|
||||
cls = override_cls or pipe.vae.__class__
|
||||
if not hasattr(cls, 'from_single_file'):
|
||||
log.error(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} safetensors=unsupported')
|
||||
return
|
||||
load_args, quant_args = model_quant.get_dit_args(load_config, module='VAE')
|
||||
log.info(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} args={load_args} quant={quant_args}')
|
||||
try:
|
||||
fn = os.path.join(shared.opts.vae_dir, shared.opts.sd_vae)
|
||||
vae = cls.from_single_file(
|
||||
fn,
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
**override_args,
|
||||
**load_args,
|
||||
**quant_args,
|
||||
)
|
||||
if vae is not None:
|
||||
pipe.vae = vae
|
||||
except Exception as e:
|
||||
log.error(f'Load model: vae="{shared.opts.sd_vae}" cls={cls.__name__} {e}')
|
||||
# errors.display(e, 'Load')
|
||||
|
||||
@@ -125,6 +125,9 @@ def load_anima(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
# generic.load_vae_override(pipe, diffusers_load_config, override_cls=diffusers.AutoencoderKLQwenImage, override_args={'low_cpu_mem_usage': False, 'ignore_mismatched_sizes': True})
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
del llm_adapter
|
||||
|
||||
@@ -25,6 +25,8 @@ def load_auraflow(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -28,6 +28,9 @@ def load_chroma(checkpoint_info, diffusers_load_config=None):
|
||||
diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaPipeline
|
||||
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaImg2ImgPipeline
|
||||
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["chroma"] = diffusers.ChromaInpaintPipeline
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -40,6 +40,8 @@ def load_ernie_image(checkpoint_info, diffusers_load_config=None):
|
||||
'use_pe': shared.opts.model_ernie_enable_pe,
|
||||
}
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del transformer
|
||||
del text_encoder
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -60,6 +60,8 @@ def load_flux(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
if os.environ.get('SD_REMOTE_T5', None) is not None:
|
||||
from modules import sd_te_remote
|
||||
log.warning('Remote-TE: applying patch')
|
||||
|
||||
@@ -31,6 +31,8 @@ def load_flux2(checkpoint_info, diffusers_load_config=None):
|
||||
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flux2"] = diffusers.Flux2Pipeline
|
||||
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flux2"] = diffusers.Flux2Pipeline
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
from pipelines.flux import flux2_lora
|
||||
flux2_lora.apply_patch()
|
||||
|
||||
|
||||
@@ -34,6 +34,8 @@ def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
|
||||
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline
|
||||
diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
from pipelines.flux import flux2_lora
|
||||
flux2_lora.apply_patch()
|
||||
|
||||
|
||||
@@ -18,6 +18,9 @@ def load_lumina(checkpoint_info, diffusers_load_config=None):
|
||||
cache_dir = shared.opts.diffusers_dir,
|
||||
**load_config,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
devices.torch_gc(force=True, reason='load')
|
||||
return pipe
|
||||
@@ -47,6 +50,8 @@ def load_lumina2(checkpoint_info, diffusers_load_config=None):
|
||||
**load_config,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del transformer
|
||||
del text_encoder
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -42,6 +42,8 @@ def load_nucleus(checkpoint_info, diffusers_load_config=None):
|
||||
'output_type': 'np',
|
||||
}
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del transformer
|
||||
del text_encoder
|
||||
del processor
|
||||
|
||||
@@ -35,6 +35,8 @@ def load_pixart(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -88,6 +88,8 @@ def load_qwen(checkpoint_info, diffusers_load_config=None):
|
||||
pipe.task_args['layers'] = shared.opts.model_qwen_layers
|
||||
pipe.task_args['resolution'] = 640
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -32,6 +32,8 @@ def load_sd3(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder_3
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
@@ -52,6 +52,8 @@ def load_z_image(checkpoint_info, diffusers_load_config=None):
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del transformer
|
||||
del text_encoder
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
Reference in New Issue
Block a user