Stable Cascade previewer and fixes

This commit is contained in:
Disty0
2024-02-15 18:50:54 +03:00
parent d0ceff70a3
commit e631fd85e2
5 changed files with 121 additions and 9 deletions
+3 -3
View File
@@ -261,11 +261,11 @@ def torch_Generator(device=None):
original_torch_load = torch.load
@wraps(torch.load)
def torch_load(f, map_location=None, pickle_module=None, *, weights_only=False, mmap=None, **kwargs):
def torch_load(f, map_location=None, *args, **kwargs):
if check_device(map_location):
return original_torch_load(f, map_location=return_xpu(map_location), pickle_module=pickle_module, weights_only=weights_only, mmap=mmap, **kwargs)
return original_torch_load(f, *args, map_location=return_xpu(map_location), **kwargs)
else:
return original_torch_load(f, map_location=map_location, pickle_module=pickle_module, weights_only=weights_only, mmap=mmap, **kwargs)
return original_torch_load(f, *args, map_location=map_location, **kwargs)
# Hijack Functions:
+13 -1
View File
@@ -207,7 +207,17 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
model.scheduler.noise_sampler_seed = p.seeds[0] # some schedulers have internal noise generator and do not use pipeline generator
if 'noise_sampler_seed' in possible:
args['noise_sampler_seed'] = p.seeds[0]
if 'guidance_scale' in possible:
if hasattr(model, "decoder") and hasattr(model, "prior_prior") and 'prior_num_inference_steps' in possible:
steps = kwargs.pop("num_inference_steps", 20)
args["prior_num_inference_steps"] = steps
args["num_inference_steps"] = max(int(steps / 2), 1) # TODO: add another slider without overcrowding the UI
if hasattr(model, "decoder") and hasattr(model, "prior_prior") and 'prior_guidance_scale' in possible:
cfg_scale = kwargs.pop("guidance_scale", p.cfg_scale)
args["prior_guidance_scale"] = cfg_scale
# Using decoder_guidance_scale causes "Expected all tensors to be on the same device" errors right now
# Enabling model cpu offload fixes the error above
#args["decoder_guidance_scale"] = 0.0
elif 'guidance_scale' in possible:
args['guidance_scale'] = p.cfg_scale
if 'generator' in possible and generator is not None:
args['generator'] = generator
@@ -219,6 +229,8 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
if 'callback' in possible:
args['callback'] = diffusers_callback_legacy
elif 'callback_on_step_end_tensor_inputs' in possible:
if hasattr(model, "decoder") and hasattr(model, "prior_prior") and 'prior_guidance_scale' in possible:
args['prior_callback_on_step_end'] = diffusers_callback
args['callback_on_step_end'] = diffusers_callback
if 'prompt_embeds' in possible and 'negative_prompt_embeds' in possible and hasattr(model, '_callback_tensor_inputs'):
args['callback_on_step_end_tensor_inputs'] = model._callback_tensor_inputs # pylint: disable=protected-access
+90
View File
@@ -0,0 +1,90 @@
import os
import torch
from torch import nn
import safetensors
from modules import devices, paths
preview_model = None
# Fast Decoder for Stage C latents. E.g. 16 x 24 x 24 -> 3 x 192 x 192
# https://github.com/Stability-AI/StableCascade/blob/master/modules/previewer.py
class Previewer(nn.Module):
def __init__(self, c_in=16, c_hidden=512, c_out=3):
super().__init__()
self.blocks = nn.Sequential(
nn.Conv2d(c_in, c_hidden, kernel_size=1), # 16 channels to 512 channels
nn.GELU(),
nn.BatchNorm2d(c_hidden),
nn.Conv2d(c_hidden, c_hidden, kernel_size=3, padding=1),
nn.GELU(),
nn.BatchNorm2d(c_hidden),
nn.ConvTranspose2d(c_hidden, c_hidden // 2, kernel_size=2, stride=2), # 16 -> 32
nn.GELU(),
nn.BatchNorm2d(c_hidden // 2),
nn.Conv2d(c_hidden // 2, c_hidden // 2, kernel_size=3, padding=1),
nn.GELU(),
nn.BatchNorm2d(c_hidden // 2),
nn.ConvTranspose2d(c_hidden // 2, c_hidden // 4, kernel_size=2, stride=2), # 32 -> 64
nn.GELU(),
nn.BatchNorm2d(c_hidden // 4),
nn.Conv2d(c_hidden // 4, c_hidden // 4, kernel_size=3, padding=1),
nn.GELU(),
nn.BatchNorm2d(c_hidden // 4),
nn.ConvTranspose2d(c_hidden // 4, c_hidden // 4, kernel_size=2, stride=2), # 64 -> 128
nn.GELU(),
nn.BatchNorm2d(c_hidden // 4),
nn.Conv2d(c_hidden // 4, c_hidden // 4, kernel_size=3, padding=1),
nn.GELU(),
nn.BatchNorm2d(c_hidden // 4),
nn.Conv2d(c_hidden // 4, c_out, kernel_size=1),
)
def forward(self, x):
return self.blocks(x)
def download_model(model_path):
model_url = 'https://huggingface.co/stabilityai/stable-cascade/resolve/main/previewer.safetensors?download=true'
if not os.path.exists(model_path):
import torch
from modules.shared import log
os.makedirs(os.path.dirname(model_path), exist_ok=True)
log.info(f'Downloading Stable Cascade previewer: {model_path}')
torch.hub.download_url_to_file(model_url, model_path)
def load_model(model_path):
checkpoint = {}
with safetensors.safe_open(model_path, framework="pt", device="cpu") as f:
for key in f.keys():
checkpoint[key] = f.get_tensor(key)
return checkpoint
def decode(latents):
from modules import shared
global preview_model
if preview_model is None:
model_path = os.path.join(paths.models_path, "VAE-approx", "sd_cascade_previewer.safetensors")
download_model(model_path)
if os.path.exists(model_path):
preview_model = Previewer()
previewer_checkpoint = load_model(model_path)
preview_model.load_state_dict(previewer_checkpoint if 'state_dict' not in previewer_checkpoint else previewer_checkpoint['state_dict'])
preview_model.eval().requires_grad_(False).to(devices.device, devices.dtype_vae)
del previewer_checkpoint
shared.log.info(f"Load Stable Cascade previewer: model={model_path}")
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)
return image
except Exception as e:
shared.log.error(f'Stable Cascade previewer: {e}')
return latents
+8 -3
View File
@@ -819,9 +819,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
diffusers_load_config['variant'] = 'fp16'
if model_type in ['Stable Cascade']: # forced pipeline
try:
# set prior manually for now
prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
# TODO: Add decoder and vqgan loader
# vqgan can reuse the vae loader ui
# decoder needs a new loader
vae_file = diffusers_load_config.pop("vae", None)
prior_model = "stabilityai/stable-cascade-prior" if "models--stabilityai--stable-cascade" in checkpoint_info.path and "models--stabilityai--stable-cascade-prior" not in checkpoint_info.path else checkpoint_info.path
decoder_model = "stabilityai/stable-cascade" if vae_file is None else vae_file
prior = diffusers.StableCascadePriorPipeline.from_pretrained(prior_model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained(decoder_model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) # pylint: disable=no-member
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_scheduler=prior.scheduler, feature_extractor=prior.feature_extractor, image_encoder=prior.image_encoder) # pylint: disable=no-member
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
+7 -2
View File
@@ -3,7 +3,7 @@ from collections import namedtuple
import torch
import torchvision.transforms as T
from PIL import Image
from modules import shared, devices, processing, images, sd_vae_approx, sd_vae_taesd, sd_samplers
from modules import shared, devices, processing, images, sd_vae_approx, sd_vae_taesd, sd_cascade_previewer, sd_samplers
SamplerData = namedtuple('SamplerData', ['name', 'constructor', 'aliases', 'options'])
@@ -33,6 +33,7 @@ def setup_img2img_steps(p, steps=None):
def single_sample_to_image(sample, approximation=None):
with queue_lock:
sd_cascade = False
if approximation is None:
approximation = approximation_indexes.get(shared.opts.show_progress_type, None)
if approximation is None:
@@ -43,6 +44,8 @@ def single_sample_to_image(sample, approximation=None):
sample = sample.to(torch.float16)
if len(sample.shape) > 4: # likely unknown video latent (e.g. svd)
return Image.new(mode="RGB", size=(512, 512))
if len(sample) == 16: # sd_cascade
sd_cascade = True
if len(sample.shape) == 4 and sample.shape[0]: # likely animatediff latent
sample = sample.permute(1, 0, 2, 3)[0]
if shared.backend == shared.Backend.DIFFUSERS: # [-x,x] to [-5,5]
@@ -52,7 +55,9 @@ def single_sample_to_image(sample, approximation=None):
sample_min = torch.min(sample)
if sample_min < -5:
sample = sample * (5 / abs(sample_min))
if approximation == 0: # Simple
if sd_cascade:
x_sample = sd_cascade_previewer.decode(sample)
elif approximation == 0: # Simple
x_sample = sd_vae_approx.cheap_approximation(sample) * 0.5 + 0.5
elif approximation == 1: # Approximate
x_sample = sd_vae_approx.nn_approximation(sample) * 0.5 + 0.5