From d171f5801871a8ac7ec192d945e05fe1becdd1e4 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Wed, 15 May 2024 09:18:18 +0300 Subject: [PATCH] Move Cascade from sd_models --- modules/sd_cascade.py | 56 +++++++++++++++++++++++++++++++++++++++++++ modules/sd_models.py | 47 +++--------------------------------- modules/sd_unet.py | 6 +++-- 3 files changed, 63 insertions(+), 46 deletions(-) diff --git a/modules/sd_cascade.py b/modules/sd_cascade.py index 80a27effd..bccc750d8 100644 --- a/modules/sd_cascade.py +++ b/modules/sd_cascade.py @@ -69,3 +69,59 @@ def load_prior(path, config_file="default"): return prior_unet, prior_text_encoder + +def load_cascade_combined(checkpoint_info, diffusers_load_config): + from diffusers import StableCascadeUNet, StableCascadeDecoderPipeline, StableCascadePriorPipeline, StableCascadeCombinedPipeline + from modules.sd_unet import unet_dict + + diffusers_load_config.pop("vae", None) + if 'stabilityai' in checkpoint_info.name: + diffusers_load_config["variant"] = 'bf16' + + if shared.opts.sd_unet != "None" or 'stabilityai' in checkpoint_info.name: + if 'stabilityai' in checkpoint_info.name and ('lite' in checkpoint_info.name or (checkpoint_info.hash is not None and 'abc818bb0d' in checkpoint_info.hash)): + decoder_folder = 'decoder_lite' + prior_folder = 'prior_lite' + else: + decoder_folder = 'decoder' + prior_folder = 'prior' + + if 'stabilityai' in checkpoint_info.name: + decoder_unet = StableCascadeUNet.from_pretrained("stabilityai/stable-cascade", subfolder=decoder_folder, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + decoder = StableCascadeDecoderPipeline.from_pretrained("stabilityai/stable-cascade", cache_dir=shared.opts.diffusers_dir, decoder=decoder_unet, **diffusers_load_config) + else: + decoder = StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + + shared.log.debug(f'StableCascade {decoder_folder}: scale={decoder.latent_dim_scale}') + + prior_text_encoder = None + if shared.opts.sd_unet != "None": + prior_unet, prior_text_encoder = load_prior(unet_dict[shared.opts.sd_unet]) + else: + prior_unet = StableCascadeUNet.from_pretrained("stabilityai/stable-cascade-prior", subfolder=prior_folder, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + + if prior_text_encoder is not None: + prior = StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, prior=prior_unet, text_encoder=prior_text_encoder, **diffusers_load_config) + else: + prior = StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, prior=prior_unet, **diffusers_load_config) + + shared.log.debug(f'StableCascade {prior_folder}: scale={prior.resolution_multiple}') + + sd_model = StableCascadeCombinedPipeline( + tokenizer=decoder.tokenizer, + text_encoder=decoder.text_encoder, + decoder=decoder.decoder, + scheduler=decoder.scheduler, + vqgan=decoder.vqgan, + prior_prior=prior.prior, + prior_text_encoder=prior.text_encoder, + prior_tokenizer=prior.tokenizer, + prior_scheduler=prior.scheduler, + prior_feature_extractor=prior.feature_extractor, + prior_image_encoder=prior.image_encoder) + else: + sd_model = StableCascadeCombinedPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + + shared.log.debug(f'StableCascade combined: {sd_model.__class__.__name__}') + + return sd_model diff --git a/modules/sd_models.py b/modules/sd_models.py index 6ea785c30..e0d31a97c 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -942,50 +942,9 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No if 'variant' not in diffusers_load_config and any('diffusion_pytorch_model.fp16' in f for f in files): # deal with diffusers lack of variant fallback when loading diffusers_load_config['variant'] = 'fp16' if model_type in ['Stable Cascade']: # forced pipeline - try: # this is horrible special-case handling for stable-cascade multi-stage pipeline with variants and non-standard revision - shared.opts.data['diffusers_model_cpu_offload'] = True # override - diffusers_load_config.pop("vae", None) - if 'stabilityai' in checkpoint_info.name: - diffusers_load_config["variant"] = 'bf16' - if shared.opts.sd_unet != "None" or 'stabilityai' in checkpoint_info.name: - if 'stabilityai' in checkpoint_info.name and ('lite' in checkpoint_info.name or (checkpoint_info.hash is not None and 'abc818bb0d' in checkpoint_info.hash)): - decoder_folder = 'decoder_lite' - prior_folder = 'prior_lite' - else: - decoder_folder = 'decoder' - prior_folder = 'prior' - if 'stabilityai' in checkpoint_info.name: - decoder_unet = diffusers.models.StableCascadeUNet.from_pretrained("stabilityai/stable-cascade", subfolder=decoder_folder, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained("stabilityai/stable-cascade", cache_dir=shared.opts.diffusers_dir, decoder=decoder_unet, **diffusers_load_config) - else: - decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - shared.log.debug(f'StableCascade {decoder_folder}: scale={decoder.latent_dim_scale}') - prior_text_encoder = None - if shared.opts.sd_unet != "None": - from modules.sd_cascade import load_prior - prior_unet, prior_text_encoder = load_prior(sd_unet.unet_dict[shared.opts.sd_unet]) - else: - prior_unet = diffusers.models.StableCascadeUNet.from_pretrained("stabilityai/stable-cascade-prior", subfolder=prior_folder, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - if prior_text_encoder is not None: - prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, prior=prior_unet, text_encoder=prior_text_encoder, **diffusers_load_config) - else: - prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, prior=prior_unet, **diffusers_load_config) - shared.log.debug(f'StableCascade {prior_folder}: scale={prior.resolution_multiple}') - sd_model = diffusers.StableCascadeCombinedPipeline( - tokenizer=decoder.tokenizer, - text_encoder=decoder.text_encoder, - decoder=decoder.decoder, - scheduler=decoder.scheduler, - vqgan=decoder.vqgan, - prior_prior=prior.prior, - prior_text_encoder=prior.text_encoder, - prior_tokenizer=prior.tokenizer, - prior_scheduler=prior.scheduler, - prior_feature_extractor=prior.feature_extractor, - prior_image_encoder=prior.image_encoder) - else: - sd_model = diffusers.StableCascadeCombinedPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - shared.log.debug(f'StableCascade combined: {sd_model.__class__.__name__}') + try: + from modules.sd_cascade import load_cascade_combined + sd_model = load_cascade_combined(checkpoint_info, diffusers_load_config) except Exception as e: shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}') if debug_load: diff --git a/modules/sd_unet.py b/modules/sd_unet.py index 2983b16f9..9f5be9655 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -24,9 +24,11 @@ def load_unet(model): if "StableCascade" in model.__class__.__name__: from modules.sd_cascade import load_prior prior_unet, prior_text_encoder = load_prior(unet_dict[shared.opts.sd_unet], config_file=config_file) - model.prior_pipe.prior = prior_unet.to(devices.device, dtype=devices.dtype_unet) + model.prior_pipe.prior = model.prior_prior = None # Prevent OOM + model.prior_pipe.prior = model.prior_prior = prior_unet.to(devices.device, dtype=devices.dtype_unet) if prior_text_encoder is not None: - model.prior_pipe.text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) + model.prior_pipe.text_encoder = model.prior_text_encoder = None # Prevent OOM + model.prior_pipe.text_encoder = model.prior_text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) else: shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" config="{config_file}"') from diffusers import UNet2DConditionModel