cleanup stable cascade

This commit is contained in:
Vladimir Mandic
2024-02-14 09:43:20 -05:00
parent ae10ae6997
commit 0e91c46a68
2 changed files with 6 additions and 5 deletions
+3 -4
View File
@@ -819,11 +819,10 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
diffusers_load_config['variant'] = 'fp16'
if model_type in ['Stable Cascade']: # forced pipeline
try:
from diffusers import StableCascadeDecoderPipeline, StableCascadePriorPipeline, StableCascadeCombinedPipeline
# set prior manually for now
prior = StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
decoder = StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
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_scheduler=prior.scheduler, feature_extractor=prior.feature_extractor, image_encoder=prior.image_encoder)
prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
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_scheduler=prior.scheduler, feature_extractor=prior.feature_extractor, image_encoder=prior.image_encoder) # pylint: disable=no-member
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
return
+3 -1
View File
@@ -54,7 +54,6 @@ def get_pipelines():
'PixArt Alpha': getattr(diffusers, 'PixArtAlphaPipeline', None),
'UniDiffuser': getattr(diffusers, 'UniDiffuserPipeline', None),
'Wuerstchen': getattr(diffusers, 'WuerstchenCombinedPipeline', None),
'Stable Cascade': getattr(diffusers, 'StableCascadeCombinedPipeline', None),
'Kandinsky 2.1': getattr(diffusers, 'KandinskyPipeline', None),
'Kandinsky 2.2': getattr(diffusers, 'KandinskyV22Pipeline', None),
'Kandinsky 3': getattr(diffusers, 'Kandinsky3Pipeline', None),
@@ -71,6 +70,9 @@ def get_pipelines():
# Segmind SSD-1B, Segmind Tiny
}
if hasattr(diffusers, 'StableCascadeCombinedPipeline'):
pipelines['Stable Cascade'] = getattr(diffusers, 'StableCascadeCombinedPipeline', None)
for k, v in pipelines.items():
if k != 'Autodetect' and v is None:
log.error(f'Not available: pipeline={k} diffusers={diffusers.__version__} path={diffusers.__file__}')