Wuerstchen V3 fixes and custom model support

This commit is contained in:
Disty0
2024-04-25 23:13:08 +03:00
parent 0137331696
commit 5b444c39d7
3 changed files with 29 additions and 22 deletions
+9 -2
View File
@@ -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'):
+19 -19
View File
@@ -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}')
+1 -1
View File
@@ -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}')