Move Cascade from sd_models

This commit is contained in:
Disty0
2024-05-15 09:18:18 +03:00
parent aca870ac4f
commit d171f58018
3 changed files with 63 additions and 46 deletions
+56
View File
@@ -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
+3 -44
View File
@@ -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:
+4 -2
View File
@@ -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