diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 6fcc79a86..e32bf2195 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -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: diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 96183f3f3..1f31205d3 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -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 diff --git a/modules/sd_cascade_previewer.py b/modules/sd_cascade_previewer.py new file mode 100644 index 000000000..9690eee7e --- /dev/null +++ b/modules/sd_cascade_previewer.py @@ -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 diff --git a/modules/sd_models.py b/modules/sd_models.py index 4935796fa..a8a666b48 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -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}') diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 2a46a8a28..6fadfdcfd 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -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