Merge pull request #3136 from Disty0/dev

Stable Cascade UNet support
This commit is contained in:
Vladimir Mandic
2024-05-15 11:54:31 -04:00
committed by GitHub
5 changed files with 277 additions and 39 deletions
+64
View File
@@ -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
]
}
@@ -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
]
}
+127
View File
@@ -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
+5 -31
View File
@@ -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':
+17 -8
View File
@@ -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}')