Cascade fix get_timestep_ratio_conditioning

This commit is contained in:
Disty0
2024-08-07 19:39:09 +03:00
parent a6b6d16273
commit db0f6c73dd
2 changed files with 13 additions and 1 deletions
+1 -1
View File
@@ -582,7 +582,7 @@ def install_ipex(torch_command):
def install_openvino(torch_command):
check_python(supported_minors=[10,11], reason='IPEX backend requires Python 3.10 or 3.11')
check_python(supported_minors=[9, 10, 11], reason='OpenVINO backend requires Python 3.9, 3.10 or 3.11')
log.info('Using OpenVINO')
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.2.0 torchvision==0.17.0 --index-url https://download.pytorch.org/whl/cpu')
install(os.environ.get('OPENVINO_PACKAGE', 'openvino==2023.3.0'), 'openvino')
+12
View File
@@ -1,7 +1,18 @@
import os
import copy
import torch
from modules import shared, devices
def get_timestep_ratio_conditioning(t, alphas_cumprod):
s = torch.tensor([0.008]) # diffusers uses 0.003 while the original is 0.008
clamp_range = [0, 1]
min_var = torch.cos(s / (1 + s) * torch.pi * 0.5) ** 2
var = alphas_cumprod[t]
var = var.clamp(*clamp_range)
s, min_var = s.to(var.device), min_var.to(var.device)
ratio = (((var * min_var) ** 0.5).acos() / (torch.pi * 0.5)) * (1 + s) - s
return ratio
def load_text_encoder(path):
from transformers import CLIPTextConfig, CLIPTextModelWithProjection
from accelerate.utils.modeling import set_module_tensor_to_device
@@ -125,6 +136,7 @@ def load_cascade_combined(checkpoint_info, diffusers_load_config):
def cascade_post_load(sd_model):
sd_model.prior_pipe.scheduler.config.clip_sample = False
sd_model.default_scheduler = copy.deepcopy(sd_model.prior_pipe.scheduler)
sd_model.prior_pipe.get_timestep_ratio_conditioning = get_timestep_ratio_conditioning
sd_model.decoder_pipe.text_encoder = sd_model.text_encoder = None # Nothing uses the decoder's text encoder
sd_model.prior_pipe.image_encoder = sd_model.prior_image_encoder = None # No img2img is implemented yet
sd_model.prior_pipe.feature_extractor = sd_model.prior_feature_extractor = None # No img2img is implemented yet