From 5b444c39d7284dbb561fcbfa052dd9a236d9b365 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Thu, 25 Apr 2024 23:13:08 +0300 Subject: [PATCH] Wuerstchen V3 fixes and custom model support --- modules/processing_diffusers.py | 11 ++++++++-- modules/sd_models.py | 38 ++++++++++++++++----------------- modules/sd_vae_stablecascade.py | 2 +- 3 files changed, 29 insertions(+), 22 deletions(-) diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 9fb4054c8..3e57dab48 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -251,9 +251,16 @@ def process_diffusers(p: processing.StableDiffusionProcessing): args["decoder_guidance_scale"] = p.image_cfg_scale # set callbacks - if 'callback_steps' in possible: + if 'prior_callback_steps' in possible: # Wuerstchen / Cascade + args['prior_callback_steps'] = 1 + elif 'callback_steps' in possible: args['callback_steps'] = 1 - if 'callback_on_step_end' in possible: + + if 'prior_callback_on_step_end' in possible: # Wuerstchen / Cascade + args['prior_callback_on_step_end'] = diffusers_callback + if 'prior_callback_on_step_end_tensor_inputs' in possible: + args['prior_callback_on_step_end_tensor_inputs'] = ['latents'] + elif 'callback_on_step_end' in possible: args['callback_on_step_end'] = diffusers_callback if 'callback_on_step_end_tensor_inputs' in possible: if 'prompt_embeds' in possible and 'negative_prompt_embeds' in possible and hasattr(model, '_callback_tensor_inputs'): diff --git a/modules/sd_models.py b/modules/sd_models.py index 267f8c407..b4dc870df 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -606,7 +606,7 @@ def detect_pipeline(f: str, op: str = 'model', warning=True): if shared.backend == shared.Backend.ORIGINAL: warn(f'Model detected as PixArt Alpha model, but attempting to load using backend=original: {op}={f} size={size} MB') guess = 'PixArt-Alpha' - if 'stable-cascade' in f.lower() or 'stablecascade' in f.lower(): + if 'stable-cascade' in f.lower() or 'stablecascade' in f.lower() or 'wuerstchen3' in f.lower(): if shared.backend == shared.Backend.ORIGINAL: warn(f'Model detected as Stable Cascade model, but attempting to load using backend=original: {op}={f} size={size} MB') guess = 'Stable Cascade' @@ -788,6 +788,8 @@ def move_model(model, device=None, force=False): return try: model.to(device) + if hasattr(model, "prior_pipe"): + model.prior_pipe.to(device) except Exception as e: shared.log.error(f'Model move: device={device} {e}') devices.torch_gc() @@ -930,31 +932,29 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No 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 diffusers_load_config.pop("vae", None) - diffusers_load_config["variant"] = 'bf16' - if 'lite' in checkpoint_info.name or 'abc818bb0d' in checkpoint_info.hash: + if 'stabilityai' in checkpoint_info.name: + diffusers_load_config["variant"] = 'bf16' + 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_unet = diffusers.models.StableCascadeUNet.from_pretrained("stabilityai/stable-cascade", subfolder="decoder_lite", 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) shared.log.debug(f'StableCascade lite decoder: scale={decoder.latent_dim_scale}') prior_unet = diffusers.models.StableCascadeUNet.from_pretrained("stabilityai/stable-cascade-prior", subfolder="prior_lite", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) 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 lite prior: 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: - decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained("stabilityai/stable-cascade", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - shared.log.debug(f'StableCascade full decoder: scale={decoder.latent_dim_scale}') - prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - shared.log.debug(f'StableCascade full prior: 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) + 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__}') except Exception as e: shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}') diff --git a/modules/sd_vae_stablecascade.py b/modules/sd_vae_stablecascade.py index 44f620a0a..00c04fc9d 100644 --- a/modules/sd_vae_stablecascade.py +++ b/modules/sd_vae_stablecascade.py @@ -82,7 +82,7 @@ def decode(latents): try: with devices.inference_context(): latents = latents.detach().clone().unsqueeze(0).to(devices.device, devices.dtype_vae) - image = preview_model(latents)[0].clamp(0, 1) + image = preview_model(latents)[0].clamp(0, 1).float() return image except Exception as e: shared.log.error(f'Stable Cascade previewer: {e}')