diff --git a/configs/stable-cascade/prior/config.json b/configs/stable-cascade/prior/config.json new file mode 100644 index 000000000..0c46a693c --- /dev/null +++ b/configs/stable-cascade/prior/config.json @@ -0,0 +1,64 @@ +{ + "_class_name": "StableCascadeUNet", + "_diffusers_version": "0.27.0.dev0", + "block_out_channels": [ + 2048, + 2048 + ], + "block_types_per_layer": [ + [ + "SDCascadeResBlock", + "SDCascadeTimestepBlock", + "SDCascadeAttnBlock" + ], + [ + "SDCascadeResBlock", + "SDCascadeTimestepBlock", + "SDCascadeAttnBlock" + ] + ], + "clip_image_in_channels": 768, + "clip_seq": 4, + "clip_text_in_channels": 1280, + "clip_text_pooled_in_channels": 1280, + "conditioning_dim": 2048, + "down_blocks_repeat_mappers": [ + 1, + 1 + ], + "down_num_layers_per_block": [ + 8, + 24 + ], + "dropout": [ + 0.1, + 0.1 + ], + "effnet_in_channels": null, + "in_channels": 16, + "kernel_size": 3, + "num_attention_heads": [ + 32, + 32 + ], + "out_channels": 16, + "patch_size": 1, + "pixel_mapper_in_channels": null, + "self_attn": true, + "switch_level": [ + false + ], + "timestep_conditioning_type": [ + "sca", + "crp" + ], + "timestep_ratio_embedding_dim": 64, + "up_blocks_repeat_mappers": [ + 1, + 1 + ], + "up_num_layers_per_block": [ + 24, + 8 + ] +} diff --git a/configs/stable-cascade/prior_lite/config.json b/configs/stable-cascade/prior_lite/config.json new file mode 100644 index 000000000..7b9fc4fa2 --- /dev/null +++ b/configs/stable-cascade/prior_lite/config.json @@ -0,0 +1,64 @@ +{ + "_class_name": "StableCascadeUNet", + "_diffusers_version": "0.27.0.dev0", + "block_out_channels": [ + 1536, + 1536 + ], + "block_types_per_layer": [ + [ + "SDCascadeResBlock", + "SDCascadeTimestepBlock", + "SDCascadeAttnBlock" + ], + [ + "SDCascadeResBlock", + "SDCascadeTimestepBlock", + "SDCascadeAttnBlock" + ] + ], + "clip_image_in_channels": 768, + "clip_seq": 4, + "clip_text_in_channels": 1280, + "clip_text_pooled_in_channels": 1280, + "conditioning_dim": 1536, + "down_blocks_repeat_mappers": [ + 1, + 1 + ], + "down_num_layers_per_block": [ + 4, + 12 + ], + "dropout": [ + 0.1, + 0.1 + ], + "effnet_in_channels": null, + "in_channels": 16, + "kernel_size": 3, + "num_attention_heads": [ + 24, + 24 + ], + "out_channels": 16, + "patch_size": 1, + "pixel_mapper_in_channels": null, + "self_attn": true, + "switch_level": [ + false + ], + "timestep_conditioning_type": [ + "sca", + "crp" + ], + "timestep_ratio_embedding_dim": 64, + "up_blocks_repeat_mappers": [ + 1, + 1 + ], + "up_num_layers_per_block": [ + 12, + 4 + ] +} diff --git a/modules/sd_cascade.py b/modules/sd_cascade.py new file mode 100644 index 000000000..bccc750d8 --- /dev/null +++ b/modules/sd_cascade.py @@ -0,0 +1,127 @@ +import os +from modules import shared, devices + +def load_text_encoder(path): + from transformers import CLIPTextConfig, CLIPTextModelWithProjection + from accelerate.utils.modeling import set_module_tensor_to_device + from accelerate import init_empty_weights + from safetensors.torch import load_file + + try: + config = CLIPTextConfig( + architectures=["CLIPTextModelWithProjection"], + attention_dropout=0.0, + bos_token_id=49406, + dropout=0.0, + eos_token_id=49407, + hidden_act="gelu", + hidden_size=1280, + initializer_factor=1.0, + initializer_range=0.02, + intermediate_size=5120, + layer_norm_eps=1e-05, + max_position_embeddings=77, + model_type="clip_text_model", + num_attention_heads=20, + num_hidden_layers=32, + pad_token_id=1, + projection_dim=1280, + vocab_size=49408 + ) + + shared.log.info(f'Loading Text Encoder: name="{os.path.basename(os.path.splitext(path)[0])}" file="{path}"') + + with init_empty_weights(): + text_encoder = CLIPTextModelWithProjection(config) + + state_dict = load_file(path) + + for key in list(state_dict.keys()): + set_module_tensor_to_device(text_encoder, key, devices.device, value=state_dict.pop(key), dtype=devices.dtype) + + return text_encoder + + except Exception as e: + text_encoder = None + shared.log.error(f'Failed to load Text Encoder model: {e}') + return None + + +def load_prior(path, config_file="default"): + from diffusers import StableCascadeUNet + prior_text_encoder = None + + if config_file == "default": + config_file = os.path.splitext(path)[0] + '.json' + if not os.path.exists(config_file): + if round(os.path.getsize(path) / 1024 / 1024 / 1024) < 5: # diffusers fails to find the configs from huggingface + config_file = "configs/stable-cascade/prior_lite/config.json" + else: + config_file = "configs/stable-cascade/prior/config.json" + + shared.log.info(f'Loading UNet: name="{os.path.basename(os.path.splitext(path)[0])}" file="{path}" config="{config_file}"') + prior_unet = StableCascadeUNet.from_single_file(path, config=config_file, torch_dtype=devices.dtype_unet, cache_dir=shared.opts.diffusers_dir) + + if os.path.isfile(os.path.splitext(path)[0] + "_text_encoder.safetensors"): # OneTrainer + prior_text_encoder = load_text_encoder(os.path.splitext(path)[0] + "_text_encoder.safetensors") + elif os.path.isfile(os.path.splitext(path)[0] + "_text_model.safetensors"): # KohyaSS + prior_text_encoder = load_text_encoder(os.path.splitext(path)[0] + "_text_model.safetensors") + + 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 diff --git a/modules/sd_models.py b/modules/sd_models.py index 840f68af0..1dd6a6abb 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -939,36 +939,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 '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' - 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) - shared.log.debug(f'StableCascade {decoder_folder}: scale={decoder.latent_dim_scale}') - prior_unet = diffusers.models.StableCascadeUNet.from_pretrained("stabilityai/stable-cascade-prior", subfolder=prior_folder, 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 {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) - 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: @@ -1128,7 +1101,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No if hasattr(sd_model, "set_progress_bar_config"): sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba') - sd_unet.load_unet(sd_model) + if "StableCascade" not in sd_model.__class__.__name__: + sd_unet.load_unet(sd_model) timer.record("load") if op == 'refiner': diff --git a/modules/sd_unet.py b/modules/sd_unet.py index 792f80a8d..9f5be9655 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -6,14 +6,12 @@ unet_dict = {} def load_unet(model): - from diffusers import UNet2DConditionModel - from safetensors.torch import load_file if shared.opts.sd_unet == 'None': return if shared.opts.sd_unet not in list(unet_dict): shared.log.error(f'UNet model not found: {shared.opts.sd_unet}') return - if not hasattr(model, 'unet') or model.unet is None: + if (not hasattr(model, 'unet') or model.unet is None) and (not hasattr(model, 'prior_prior') or model.prior_prior is None): shared.log.error('UNet not found in current model') return config_file = os.path.splitext(unet_dict[shared.opts.sd_unet])[0] + '.json' @@ -23,11 +21,22 @@ def load_unet(model): config = None config_file = 'default' try: - shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" config="{config_file}"') - unet = UNet2DConditionModel.from_config(model.unet.config if config is None else config).to(devices.device, devices.dtype) - state_dict = load_file(unet_dict[shared.opts.sd_unet]) - unet.load_state_dict(state_dict) - model.unet = unet.to(devices.device, devices.dtype_unet) + 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 = 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 = 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 + from safetensors.torch import load_file + unet = UNet2DConditionModel.from_config(model.unet.config if config is None else config).to(devices.device, devices.dtype) + state_dict = load_file(unet_dict[shared.opts.sd_unet]) + unet.load_state_dict(state_dict) + model.unet = unet.to(devices.device, devices.dtype_unet) except Exception as e: unet = None shared.log.error(f'Failed to load UNet model: {e}')