mirror of
https://github.com/vladmandic/automatic
synced 2026-08-26 15:16:01 +02:00
Stable Cascade previewer and fixes
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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}')
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user