mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
Wuerstchen V3 fixes and custom model support
This commit is contained in:
@@ -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
@@ -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}')
|
||||
|
||||
@@ -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}')
|
||||
|
||||
Reference in New Issue
Block a user