Tiled vae, batch support, progress bar
Breaks old nodes as they are, should be remade.
This commit is contained in:
+1
-2
@@ -1,4 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
Binary file not shown.
@@ -28,6 +28,7 @@ from ....ldm.modules.diffusionmodules.util import make_beta_schedule, extract_in
|
||||
from ....ldm.models.diffusion.ddim import DDIMSampler
|
||||
from ....model.mixins import ImageLoggerMixin
|
||||
|
||||
from ....utils.tilevae import VAEHook
|
||||
|
||||
__conditioning_keys__ = {'concat': 'c_concat',
|
||||
'crossattn': 'c_crossattn',
|
||||
@@ -115,7 +116,8 @@ class DDPM(pl.LightningModule, ImageLoggerMixin):
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet)
|
||||
if reset_ema:
|
||||
assert self.use_ema
|
||||
print(f"Resetting ema to pure model weights. This is useful when restoring from an ema-only checkpoint.")
|
||||
print(
|
||||
f"Resetting ema to pure model weights. This is useful when restoring from an ema-only checkpoint.")
|
||||
self.model_ema = LitEma(self.model)
|
||||
if reset_num_ema_updates:
|
||||
print(" +++++++++++ WARNING: RESETTING NUM_EMA UPDATES TO ZERO +++++++++++ ")
|
||||
@@ -461,11 +463,11 @@ class DDPM(pl.LightningModule, ImageLoggerMixin):
|
||||
|
||||
@torch.no_grad()
|
||||
def validation_step(self, batch, batch_idx):
|
||||
_, loss_dict_no_ema_0 = self.shared_step(batch,0)
|
||||
_, loss_dict_no_ema_1 = self.shared_step(batch,1)
|
||||
_, loss_dict_no_ema_0 = self.shared_step(batch, 0)
|
||||
_, loss_dict_no_ema_1 = self.shared_step(batch, 1)
|
||||
with self.ema_scope():
|
||||
_, loss_dict_ema_0 = self.shared_step(batch,0)
|
||||
_, loss_dict_ema_1 = self.shared_step(batch,1)
|
||||
_, loss_dict_ema_0 = self.shared_step(batch, 0)
|
||||
_, loss_dict_ema_1 = self.shared_step(batch, 1)
|
||||
loss_dict_ema_0 = {key + '_ema': loss_dict_ema_0[key] for key in loss_dict_ema_0}
|
||||
loss_dict_ema_1 = {key + '_ema': loss_dict_ema_1[key] for key in loss_dict_ema_1}
|
||||
self.log_dict(loss_dict_no_ema_0, prog_bar=False, logger=True, on_step=False, on_epoch=True)
|
||||
@@ -623,6 +625,30 @@ class LatentDiffusion(DDPM):
|
||||
if self.shorten_cond_schedule:
|
||||
self.make_cond_schedule()
|
||||
|
||||
def _init_tiled_vae(self,
|
||||
encoder_tile_size=256,
|
||||
decoder_tile_size=256,
|
||||
fast_decoder=False,
|
||||
fast_encoder=False,
|
||||
color_fix=False,
|
||||
vae_to_gpu=True):
|
||||
# copy from ''
|
||||
# save original forward (only once)
|
||||
if not hasattr(self.first_stage_model.encoder, 'original_forward'):
|
||||
setattr(self.first_stage_model.encoder, 'original_forward', self.first_stage_model.encoder.forward)
|
||||
if not hasattr(self.first_stage_model.decoder, 'original_forward'):
|
||||
setattr(self.first_stage_model.decoder, 'original_forward', self.first_stage_model.decoder.forward)
|
||||
|
||||
encoder = self.first_stage_model.encoder
|
||||
decoder = self.first_stage_model.decoder
|
||||
|
||||
self.first_stage_model.encoder.forward = VAEHook(
|
||||
encoder, encoder_tile_size, is_decoder=False, fast_decoder=fast_decoder, fast_encoder=fast_encoder,
|
||||
color_fix=color_fix, to_gpu=vae_to_gpu)
|
||||
self.first_stage_model.decoder.forward = VAEHook(
|
||||
decoder, decoder_tile_size, is_decoder=True, fast_decoder=fast_decoder, fast_encoder=fast_encoder,
|
||||
color_fix=color_fix, to_gpu=vae_to_gpu)
|
||||
|
||||
def instantiate_first_stage(self, config):
|
||||
model = instantiate_from_config(config)
|
||||
self.first_stage_model = model.eval()
|
||||
@@ -857,12 +883,12 @@ class LatentDiffusion(DDPM):
|
||||
|
||||
def shared_step(self, batch, index, **kwargs):
|
||||
x, c = self.get_input(batch, self.first_stage_key)
|
||||
loss = self(x, c,index)
|
||||
loss = self(x, c, index)
|
||||
return loss
|
||||
|
||||
def forward(self, x, c, opt_index, *args, **kwargs):
|
||||
T = (self.num_timesteps-1) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t_max = round(self.num_timesteps*self.t_max) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
T = (self.num_timesteps - 1) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t_max = round(self.num_timesteps * self.t_max) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t = torch.randint(0, t_max[0], (x.shape[0],), device=self.device).long()
|
||||
|
||||
if self.model.conditioning_key is not None:
|
||||
@@ -961,7 +987,8 @@ class LatentDiffusion(DDPM):
|
||||
model_output_noisy_tao = self.q_sample(x_start=model_output_x0, t=t_tao, noise=noise3)
|
||||
|
||||
model_output_noise_from_tao = self.apply_model(model_output_noisy_tao, t_tao, cond)
|
||||
model_output_x0_from_tao = self.predict_start_from_noise(model_output_noisy_tao, t=t_tao, noise=model_output_noise_from_tao) # predict x0 from tao
|
||||
model_output_x0_from_tao = self.predict_start_from_noise(model_output_noisy_tao, t=t_tao,
|
||||
noise=model_output_noise_from_tao) # predict x0 from tao
|
||||
|
||||
loss_x0 = self.get_loss(model_output_x0, x_start, mean=False).mean([1, 2, 3])
|
||||
loss_noise_from_tao = self.get_loss(model_output_noise_from_tao, noise3, mean=False).mean([1, 2, 3])
|
||||
@@ -1523,7 +1550,7 @@ class LatentUpscaleDiffusion(LatentDiffusion):
|
||||
uc[k] = [uc_tmp]
|
||||
elif k == "c_adm": # todo: only run with text-based guidance?
|
||||
assert isinstance(c[k], torch.Tensor)
|
||||
#uc[k] = torch.ones_like(c[k]) * self.low_scale_model.max_noise_level
|
||||
# uc[k] = torch.ones_like(c[k]) * self.low_scale_model.max_noise_level
|
||||
uc[k] = c[k]
|
||||
elif isinstance(c[k], list):
|
||||
uc[k] = [c[k][i] for i in range(len(c[k]))]
|
||||
@@ -1799,6 +1826,7 @@ class LatentUpscaleFinetuneDiffusion(LatentFinetuneDiffusion):
|
||||
"""
|
||||
condition on low-res image (and optionally on some spatial noise augmentation)
|
||||
"""
|
||||
|
||||
def __init__(self, concat_keys=("lr",), reshuffle_patch_size=None,
|
||||
low_scale_config=None, low_scale_key=None, *args, **kwargs):
|
||||
super().__init__(concat_keys=concat_keys, *args, **kwargs)
|
||||
|
||||
@@ -29,7 +29,7 @@ from ....ldm.models.diffusion.ddim import DDIMSampler
|
||||
from ....model.mixins import ImageLoggerMixin
|
||||
|
||||
from ....model.q_sampler import space_timesteps
|
||||
|
||||
from ....utils.tilevae import VAEHook
|
||||
|
||||
__conditioning_keys__ = {'concat': 'c_concat',
|
||||
'crossattn': 'c_crossattn',
|
||||
@@ -117,7 +117,8 @@ class DDPM(pl.LightningModule, ImageLoggerMixin):
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet)
|
||||
if reset_ema:
|
||||
assert self.use_ema
|
||||
print(f"Resetting ema to pure model weights. This is useful when restoring from an ema-only checkpoint.")
|
||||
print(
|
||||
f"Resetting ema to pure model weights. This is useful when restoring from an ema-only checkpoint.")
|
||||
self.model_ema = LitEma(self.model)
|
||||
if reset_num_ema_updates:
|
||||
print(" +++++++++++ WARNING: RESETTING NUM_EMA UPDATES TO ZERO +++++++++++ ")
|
||||
@@ -249,7 +250,8 @@ class DDPM(pl.LightningModule, ImageLoggerMixin):
|
||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||
self.register_buffer('posterior_variance_inter', to_torch(posterior_variance))
|
||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||
self.register_buffer('posterior_log_variance_clipped_inter', to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
|
||||
self.register_buffer('posterior_log_variance_clipped_inter',
|
||||
to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
|
||||
self.register_buffer('posterior_mean_coef1_inter', to_torch(
|
||||
betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)))
|
||||
self.register_buffer('posterior_mean_coef2_inter', to_torch(
|
||||
@@ -548,11 +550,11 @@ class DDPM(pl.LightningModule, ImageLoggerMixin):
|
||||
|
||||
@torch.no_grad()
|
||||
def validation_step(self, batch, batch_idx):
|
||||
_, loss_dict_no_ema_0 = self.shared_step(batch,0)
|
||||
_, loss_dict_no_ema_1 = self.shared_step(batch,1)
|
||||
_, loss_dict_no_ema_0 = self.shared_step(batch, 0)
|
||||
_, loss_dict_no_ema_1 = self.shared_step(batch, 1)
|
||||
with self.ema_scope():
|
||||
_, loss_dict_ema_0 = self.shared_step(batch,0)
|
||||
_, loss_dict_ema_1 = self.shared_step(batch,1)
|
||||
_, loss_dict_ema_0 = self.shared_step(batch, 0)
|
||||
_, loss_dict_ema_1 = self.shared_step(batch, 1)
|
||||
loss_dict_ema_0 = {key + '_ema': loss_dict_ema_0[key] for key in loss_dict_ema_0}
|
||||
loss_dict_ema_1 = {key + '_ema': loss_dict_ema_1[key] for key in loss_dict_ema_1}
|
||||
self.log_dict(loss_dict_no_ema_0, prog_bar=False, logger=True, on_step=False, on_epoch=True)
|
||||
@@ -712,6 +714,30 @@ class LatentDiffusion(DDPM):
|
||||
if self.shorten_cond_schedule:
|
||||
self.make_cond_schedule()
|
||||
|
||||
def _init_tiled_vae(self,
|
||||
encoder_tile_size=256,
|
||||
decoder_tile_size=256,
|
||||
fast_decoder=False,
|
||||
fast_encoder=False,
|
||||
color_fix=False,
|
||||
vae_to_gpu=True):
|
||||
# copy from ''
|
||||
# save original forward (only once)
|
||||
if not hasattr(self.first_stage_model.encoder, 'original_forward'):
|
||||
setattr(self.first_stage_model.encoder, 'original_forward', self.first_stage_model.encoder.forward)
|
||||
if not hasattr(self.first_stage_model.decoder, 'original_forward'):
|
||||
setattr(self.first_stage_model.decoder, 'original_forward', self.first_stage_model.decoder.forward)
|
||||
|
||||
encoder = self.first_stage_model.encoder
|
||||
decoder = self.first_stage_model.decoder
|
||||
|
||||
self.first_stage_model.encoder.forward = VAEHook(
|
||||
encoder, encoder_tile_size, is_decoder=False, fast_decoder=fast_decoder, fast_encoder=fast_encoder,
|
||||
color_fix=color_fix, to_gpu=vae_to_gpu)
|
||||
self.first_stage_model.decoder.forward = VAEHook(
|
||||
decoder, decoder_tile_size, is_decoder=True, fast_decoder=fast_decoder, fast_encoder=fast_encoder,
|
||||
color_fix=color_fix, to_gpu=vae_to_gpu)
|
||||
|
||||
def instantiate_first_stage(self, config):
|
||||
model = instantiate_from_config(config)
|
||||
self.first_stage_model = model.eval()
|
||||
@@ -959,14 +985,14 @@ class LatentDiffusion(DDPM):
|
||||
def encode_first_stage(self, x):
|
||||
return self.first_stage_model.encode(x)
|
||||
|
||||
def shared_step(self, batch,index, **kwargs):
|
||||
def shared_step(self, batch, index, **kwargs):
|
||||
gt, x, c = self.get_input(batch, self.first_stage_key)
|
||||
loss = self(gt, x, c,index)
|
||||
loss = self(gt, x, c, index)
|
||||
return loss
|
||||
|
||||
def forward(self, gt, x, c, opt_index, *args, **kwargs):
|
||||
T = (self.num_timesteps-1) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t_max = round(self.num_timesteps*self.t_max) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
T = (self.num_timesteps - 1) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t_max = round(self.num_timesteps * self.t_max) * torch.ones(x.shape[0]).to(self.device).long()
|
||||
t = torch.randint(0, t_max[0], (x.shape[0],), device=self.device).long()
|
||||
|
||||
if self.model.conditioning_key is not None:
|
||||
@@ -1045,9 +1071,9 @@ class LatentDiffusion(DDPM):
|
||||
|
||||
time_range = np.flip(self.timesteps_inter) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[round(total_steps*(1-self.t_max)):]
|
||||
time_range = time_range[round(total_steps * (1 - self.t_max)):]
|
||||
total_steps_end = len(time_range)
|
||||
time_range = time_range[:-round(total_steps*self.t_min)]
|
||||
time_range = time_range[:-round(total_steps * self.t_min)]
|
||||
|
||||
b = x_noisy.shape[0]
|
||||
t_tao = time_range[0] * torch.ones(x_noisy.shape[0]).to(self.device).long()
|
||||
@@ -1058,7 +1084,6 @@ class LatentDiffusion(DDPM):
|
||||
|
||||
# iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
|
||||
|
||||
for i, step in enumerate(time_range):
|
||||
ts = torch.full((b,), step, device=self.device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps_end - i - 1)
|
||||
@@ -1096,20 +1121,20 @@ class LatentDiffusion(DDPM):
|
||||
|
||||
if opt_index == 0:
|
||||
last_layer = self.first_stage_model.decoder.conv_out.weight
|
||||
loss,_ = self.decoder_loss(inputs=gt, reconstructions=out_decoder, optimizer_idx=opt_index, global_step=self.global_step,
|
||||
loss, _ = self.decoder_loss(inputs=gt, reconstructions=out_decoder, optimizer_idx=opt_index,
|
||||
global_step=self.global_step,
|
||||
last_layer=last_layer, split="train")
|
||||
loss_dict.update({f'{prefix}/loss': loss})
|
||||
return loss, loss_dict
|
||||
else:
|
||||
last_layer = self.first_stage_model.decoder.conv_out.weight
|
||||
# train the discriminator net
|
||||
d_loss, _= self.decoder_loss(inputs=gt, reconstructions=out_decoder, optimizer_idx=opt_index, global_step=self.global_step,
|
||||
d_loss, _ = self.decoder_loss(inputs=gt, reconstructions=out_decoder, optimizer_idx=opt_index,
|
||||
global_step=self.global_step,
|
||||
last_layer=last_layer, split="train")
|
||||
loss_dict.update({f'{prefix}/d_loss': d_loss.mean()})
|
||||
return d_loss, loss_dict
|
||||
|
||||
|
||||
|
||||
# if opt_index == 0:
|
||||
# noise = default(noise, lambda: torch.randn_like(x_start))
|
||||
# x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
@@ -1723,7 +1748,7 @@ class LatentUpscaleDiffusion(LatentDiffusion):
|
||||
uc[k] = [uc_tmp]
|
||||
elif k == "c_adm": # todo: only run with text-based guidance?
|
||||
assert isinstance(c[k], torch.Tensor)
|
||||
#uc[k] = torch.ones_like(c[k]) * self.low_scale_model.max_noise_level
|
||||
# uc[k] = torch.ones_like(c[k]) * self.low_scale_model.max_noise_level
|
||||
uc[k] = c[k]
|
||||
elif isinstance(c[k], list):
|
||||
uc[k] = [c[k][i] for i in range(len(c[k]))]
|
||||
@@ -1999,6 +2024,7 @@ class LatentUpscaleFinetuneDiffusion(LatentFinetuneDiffusion):
|
||||
"""
|
||||
condition on low-res image (and optionally on some spatial noise augmentation)
|
||||
"""
|
||||
|
||||
def __init__(self, concat_keys=("lr",), reshuffle_patch_size=None,
|
||||
low_scale_config=None, low_scale_key=None, *args, **kwargs):
|
||||
super().__init__(concat_keys=concat_keys, *args, **kwargs)
|
||||
|
||||
+241
-153
@@ -3,13 +3,17 @@ from typing import Optional, Tuple, Dict, List, Callable
|
||||
import torch
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import einops
|
||||
import os
|
||||
from PIL import Image
|
||||
|
||||
from ..ldm.modules.diffusionmodules.util import make_beta_schedule
|
||||
from ..model.cond_fn import Guidance
|
||||
from ldm.modules.diffusionmodules.util import make_beta_schedule
|
||||
from ..utils.align_color import (
|
||||
wavelet_reconstruction, adaptive_instance_normalization
|
||||
)
|
||||
|
||||
import comfy.utils
|
||||
|
||||
# https://github.com/openai/guided-diffusion/blob/main/guided_diffusion/respace.py
|
||||
def space_timesteps(num_timesteps, section_counts):
|
||||
"""
|
||||
@@ -32,7 +36,7 @@ def space_timesteps(num_timesteps, section_counts):
|
||||
"""
|
||||
if isinstance(section_counts, str):
|
||||
if section_counts.startswith("ddim"):
|
||||
desired_count = int(section_counts[len("ddim") :])
|
||||
desired_count = int(section_counts[len("ddim"):])
|
||||
for i in range(1, num_timesteps):
|
||||
if len(range(0, num_timesteps, i)) == desired_count:
|
||||
return set(range(0, num_timesteps, i))
|
||||
@@ -97,8 +101,8 @@ class SpacedSampler:
|
||||
def __init__(
|
||||
self,
|
||||
model: "ControlLDM",
|
||||
schedule: str="linear",
|
||||
var_type: str="fixed_small"
|
||||
schedule: str = "linear",
|
||||
var_type: str = "fixed_small"
|
||||
) -> "SpacedSampler":
|
||||
self.model = model
|
||||
self.original_num_steps = model.num_timesteps
|
||||
@@ -145,7 +149,7 @@ class SpacedSampler:
|
||||
self.alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
self.alphas_cumprod_prev = np.append(1.0, self.alphas_cumprod[:-1])
|
||||
self.alphas_cumprod_next = np.append(self.alphas_cumprod[1:], 0.0)
|
||||
assert self.alphas_cumprod_prev.shape == (num_steps, )
|
||||
assert self.alphas_cumprod_prev.shape == (num_steps,)
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod)
|
||||
@@ -212,7 +216,7 @@ class SpacedSampler:
|
||||
self.tao_alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
self.tao_alphas_cumprod_prev = np.append(1.0, self.tao_alphas_cumprod[:-1])
|
||||
self.tao_alphas_cumprod_next = np.append(self.tao_alphas_cumprod[1:], 0.0)
|
||||
assert self.tao_alphas_cumprod_prev.shape == (num_steps, )
|
||||
assert self.tao_alphas_cumprod_prev.shape == (num_steps,)
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.tao_sqrt_alphas_cumprod = np.sqrt(self.tao_alphas_cumprod)
|
||||
@@ -243,7 +247,7 @@ class SpacedSampler:
|
||||
self,
|
||||
x_start: torch.Tensor,
|
||||
t: torch.Tensor,
|
||||
noise: Optional[torch.Tensor]=None
|
||||
noise: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Implement the marginal distribution q(x_t|x_0).
|
||||
@@ -375,70 +379,6 @@ class SpacedSampler:
|
||||
|
||||
return e_t
|
||||
|
||||
def apply_cond_fn(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
cond: Dict[str, torch.Tensor],
|
||||
t: torch.Tensor,
|
||||
index: torch.Tensor,
|
||||
cond_fn: Guidance,
|
||||
cfg_scale: float,
|
||||
uncond: Optional[Dict[str, torch.Tensor]]
|
||||
) -> torch.Tensor:
|
||||
device = x.device
|
||||
t_now = int(t[0].item()) + 1
|
||||
# ----------------- predict noise and x0 ----------------- #
|
||||
e_t = self.predict_noise(
|
||||
x, t, cond, cfg_scale, uncond
|
||||
)
|
||||
pred_x0: torch.Tensor = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t)
|
||||
model_mean, _, _ = self.q_posterior_mean_variance(
|
||||
x_start=pred_x0, x_t=x, t=index
|
||||
)
|
||||
|
||||
# apply classifier guidance for multiple times
|
||||
for _ in range(cond_fn.repeat):
|
||||
# ----------------- compute gradient for x0 in latent space ----------------- #
|
||||
target, pred = None, None
|
||||
if cond_fn.space == "latent":
|
||||
target = self.model.get_first_stage_encoding(
|
||||
self.model.encode_first_stage(cond_fn.target.to(device))
|
||||
)
|
||||
pred = pred_x0
|
||||
elif cond_fn.space == "rgb":
|
||||
# We need to backward gradient to x0 in latent space, so it's required
|
||||
# to trace the computation graph while decoding the latent.
|
||||
with torch.enable_grad():
|
||||
pred_x0.requires_grad_(True)
|
||||
target = cond_fn.target.to(device)
|
||||
pred = self.model.decode_first_stage_with_grad(pred_x0)
|
||||
else:
|
||||
raise NotImplementedError(cond_fn.space)
|
||||
delta_pred = cond_fn(target, pred, t_now)
|
||||
|
||||
# ----------------- apply classifier guidance ----------------- #
|
||||
if delta_pred is not None:
|
||||
if cond_fn.space == "rgb":
|
||||
# compute gradient for pred_x0
|
||||
pred.backward(delta_pred)
|
||||
delta_pred_x0 = pred_x0.grad
|
||||
# update prex_x0
|
||||
pred_x0 += delta_pred_x0
|
||||
# our classifier guidance is equivalent to multiply delta_pred_x0
|
||||
# by a constant and then add it to model_mean, We set the constant
|
||||
# to 0.5
|
||||
model_mean += 0.5 * delta_pred_x0
|
||||
pred_x0.grad.zero_()
|
||||
else:
|
||||
delta_pred_x0 = delta_pred
|
||||
pred_x0 += delta_pred_x0
|
||||
model_mean += 0.5 * delta_pred_x0
|
||||
else:
|
||||
# means stop guidance
|
||||
break
|
||||
|
||||
return model_mean.detach().clone(), pred_x0.detach().clone()
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample(
|
||||
self,
|
||||
@@ -448,7 +388,6 @@ class SpacedSampler:
|
||||
index: torch.Tensor,
|
||||
cfg_scale: float,
|
||||
uncond: Optional[Dict[str, torch.Tensor]],
|
||||
cond_fn: Optional[Guidance]
|
||||
) -> torch.Tensor:
|
||||
# variance of posterior distribution q(x_{t-1}|x_t, x_0)
|
||||
model_variance = {
|
||||
@@ -457,14 +396,6 @@ class SpacedSampler:
|
||||
}[self.var_type]
|
||||
model_variance = _extract_into_tensor(model_variance, index, x.shape)
|
||||
|
||||
# mean of posterior distribution q(x_{t-1}|x_t, x_0)
|
||||
if cond_fn is not None:
|
||||
# apply classifier guidance
|
||||
model_mean, pred_x0 = self.apply_cond_fn(
|
||||
x, cond, t, index, cond_fn,
|
||||
cfg_scale, uncond
|
||||
)
|
||||
else:
|
||||
e_t = self.predict_noise(
|
||||
x, t, cond, cfg_scale, uncond
|
||||
)
|
||||
@@ -490,7 +421,6 @@ class SpacedSampler:
|
||||
index: torch.Tensor,
|
||||
cfg_scale: float,
|
||||
uncond: Optional[Dict[str, torch.Tensor]],
|
||||
cond_fn: Optional[Guidance]
|
||||
) -> torch.Tensor:
|
||||
# variance of posterior distribution q(x_{t-1}|x_t, x_0)
|
||||
model_variance = {
|
||||
@@ -500,13 +430,6 @@ class SpacedSampler:
|
||||
model_variance = _extract_into_tensor(model_variance, index, x.shape)
|
||||
|
||||
# mean of posterior distribution q(x_{t-1}|x_t, x_0)
|
||||
if cond_fn is not None:
|
||||
# apply classifier guidance
|
||||
model_mean, pred_x0 = self.apply_cond_fn(
|
||||
x, cond, t, index, cond_fn,
|
||||
cfg_scale, uncond
|
||||
)
|
||||
else:
|
||||
e_t = self.predict_noise(
|
||||
x, t, cond, cfg_scale, uncond
|
||||
)
|
||||
@@ -532,17 +455,9 @@ class SpacedSampler:
|
||||
index: torch.Tensor,
|
||||
t_max: float,
|
||||
cfg_scale: float,
|
||||
uncond: Optional[Dict[str, torch.Tensor]],
|
||||
cond_fn: Optional[Guidance]
|
||||
uncond: Optional[Dict[str, torch.Tensor]]
|
||||
) -> torch.Tensor:
|
||||
|
||||
if cond_fn is not None:
|
||||
# apply classifier guidance
|
||||
model_mean, pred_x0 = self.apply_cond_fn(
|
||||
x, cond, t, index, cond_fn,
|
||||
cfg_scale, uncond
|
||||
)
|
||||
else:
|
||||
e_t = self.predict_noise(
|
||||
x, t, cond, cfg_scale, uncond
|
||||
)
|
||||
@@ -550,15 +465,15 @@ class SpacedSampler:
|
||||
|
||||
# sample x_t from q(x_{t-1}|x_t, x_0)
|
||||
noise = torch.randn_like(x)
|
||||
tao_index = torch.tensor(torch.round(index * t_max),dtype=torch.int64)
|
||||
tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64)
|
||||
x_prev = self.q_sample(pred_x0, tao_index)
|
||||
|
||||
return x_prev
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_with_mixdiff_ccsr(
|
||||
def sample_with_tile_ccsr(
|
||||
self,
|
||||
empty_text_embed: torch.Tensor,
|
||||
tile_size: int,
|
||||
tile_stride: int,
|
||||
steps: int,
|
||||
@@ -568,10 +483,191 @@ class SpacedSampler:
|
||||
cond_img: torch.Tensor,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
x_T: Optional[torch.Tensor]=None,
|
||||
cfg_scale: float=1.,
|
||||
cond_fn: Optional[Guidance]=None,
|
||||
color_fix_type: str="none"
|
||||
x_T: Optional[torch.Tensor] = None,
|
||||
cfg_scale: float = 1.,
|
||||
color_fix_type: str = "none"
|
||||
) -> torch.Tensor:
|
||||
def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]:
|
||||
hi_list = list(range(0, h - tile_size + 1, tile_stride))
|
||||
if (h - tile_size) % tile_stride != 0:
|
||||
hi_list.append(h - tile_size)
|
||||
|
||||
wi_list = list(range(0, w - tile_size + 1, tile_stride))
|
||||
if (w - tile_size) % tile_stride != 0:
|
||||
wi_list.append(w - tile_size)
|
||||
|
||||
coords = []
|
||||
for hi in hi_list:
|
||||
for wi in wi_list:
|
||||
coords.append((hi, hi + tile_size, wi, wi + tile_size))
|
||||
return coords
|
||||
|
||||
def gaussian_weights(tile_width: int, tile_height: int, nbatches: int) -> torch.Tensor:
|
||||
"""Generates a gaussian mask of weights for tile contributions"""
|
||||
from numpy import pi, exp, sqrt
|
||||
import numpy as np
|
||||
|
||||
latent_width = tile_width
|
||||
latent_height = tile_height
|
||||
|
||||
var = 0.01
|
||||
midpoint = (latent_width - 1) / 2 # -1 because index goes from 0 to latent_width - 1
|
||||
x_probs = [
|
||||
exp(-(x - midpoint) * (x - midpoint) / (latent_width * latent_width) / (2 * var)) / sqrt(2 * pi * var)
|
||||
for x in range(latent_width)]
|
||||
midpoint = latent_height / 2
|
||||
y_probs = [
|
||||
exp(-(y - midpoint) * (y - midpoint) / (latent_height * latent_height) / (2 * var)) / sqrt(2 * pi * var)
|
||||
for y in range(latent_height)]
|
||||
|
||||
weights = np.outer(y_probs, x_probs)
|
||||
|
||||
return torch.tile(torch.tensor(weights, device=next(self.model.parameters()).device), (nbatches, 4, 1, 1))
|
||||
|
||||
# make sampling parameters (e.g. sigmas)
|
||||
self.make_schedule(num_steps=steps)
|
||||
|
||||
device = next(self.model.parameters()).device
|
||||
b, _, h, w = shape
|
||||
if x_T is None:
|
||||
img = torch.randn(shape, dtype=torch.float32, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
|
||||
# timesteps iterator
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(self.timesteps)
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
|
||||
# q_sample for the start
|
||||
ts = torch.full((b,), time_range[0], device=device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps - 1)
|
||||
|
||||
# calculate the weights
|
||||
tile_weights = gaussian_weights(tile_size // 8, tile_size // 8, 1)
|
||||
|
||||
# create buffers for accumulating predicted noise of different diffusion process
|
||||
noise_buffer = torch.zeros_like(img)
|
||||
count = torch.zeros_like(img)
|
||||
|
||||
# predict noise for each tile
|
||||
tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8))
|
||||
for hi, hi_end, wi, wi_end in tiles_iterator:
|
||||
tiles_iterator.set_description(f"Process tile with location ({hi} {hi_end}) ({wi} {wi_end})")
|
||||
# noisy latent of this diffusion process (tile) at this step
|
||||
tile_img = img[:, :, hi:hi_end, wi:wi_end]
|
||||
# prepare condition for this tile
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
|
||||
# accumulate noise
|
||||
noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise * tile_weights
|
||||
count[:, :, hi:hi_end, wi:wi_end] += tile_weights
|
||||
|
||||
# fuse by tile_weights on noise (score)
|
||||
noise_buffer /= count
|
||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
||||
tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64)
|
||||
|
||||
img = self.q_sample(pred_x0, tao_index)
|
||||
|
||||
noise_buffer.zero_()
|
||||
count.zero_()
|
||||
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[-int(round(total_steps * t_max)):]
|
||||
total_steps_use = len(time_range)
|
||||
time_range = time_range[:-int(round(total_steps * t_min))]
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
pbar = comfy.utils.ProgressBar(total_steps // 3)
|
||||
# sampling loop
|
||||
for i, step in enumerate(iterator):
|
||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps_use - i - 1)
|
||||
|
||||
# predict noise for each tile
|
||||
tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8))
|
||||
for hi, hi_end, wi, wi_end in tiles_iterator:
|
||||
tiles_iterator.set_description(f"Process tile with location ({hi} {hi_end}) ({wi} {wi_end})")
|
||||
# noisy latent of this diffusion process (tile) at this step
|
||||
tile_img = img[:, :, hi:hi_end, wi:wi_end]
|
||||
# prepare condition for this tile
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
}
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
|
||||
# accumulate noise
|
||||
noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise * tile_weights
|
||||
count[:, :, hi:hi_end, wi:wi_end] += tile_weights
|
||||
pbar.update(1)
|
||||
# average on noise (score)
|
||||
noise_buffer /= count
|
||||
# sample previous latent
|
||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
||||
mean, _, _ = self.q_posterior_mean_variance(
|
||||
x_start=pred_x0, x_t=img, t=index
|
||||
)
|
||||
variance = {
|
||||
"fixed_large": np.append(self.posterior_variance[1], self.betas[1:]),
|
||||
"fixed_small": self.posterior_variance
|
||||
}[self.var_type]
|
||||
variance = _extract_into_tensor(variance, index, noise_buffer.shape)
|
||||
|
||||
nonzero_mask = (
|
||||
(index != 0).float().view(-1, *([1] * (len(noise_buffer.shape) - 1)))
|
||||
)
|
||||
img = mean + nonzero_mask * torch.sqrt(variance) * torch.randn_like(mean)
|
||||
|
||||
noise_buffer.zero_()
|
||||
count.zero_()
|
||||
|
||||
img = pred_x0
|
||||
|
||||
img_pixel = (self.model.decode_first_stage(img) + 1) / 2
|
||||
# apply color correction (borrowed from StableSR)
|
||||
if color_fix_type == "adain":
|
||||
img_pixel = adaptive_instance_normalization(img_pixel, cond_img)
|
||||
elif color_fix_type == "wavelet":
|
||||
img_pixel = wavelet_reconstruction(img_pixel, cond_img)
|
||||
else:
|
||||
assert color_fix_type == "none", f"unexpected color fix type: {color_fix_type}"
|
||||
return img_pixel
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_with_mixdiff_ccsr(
|
||||
self,
|
||||
empty_text_embed: torch.Tensor,
|
||||
tile_size: int,
|
||||
tile_stride: int,
|
||||
steps: int,
|
||||
t_max: float,
|
||||
t_min: float,
|
||||
shape: Tuple[int],
|
||||
cond_img: torch.Tensor,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
x_T: Optional[torch.Tensor] = None,
|
||||
cfg_scale: float = 1.,
|
||||
color_fix_type: str = "none"
|
||||
) -> torch.Tensor:
|
||||
def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]:
|
||||
hi_list = list(range(0, h - tile_size + 1, tile_stride))
|
||||
@@ -604,6 +700,7 @@ class SpacedSampler:
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(self.timesteps)
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
pbar = comfy.utils.ProgressBar(total_steps // 3)
|
||||
|
||||
# q_sample for the start
|
||||
ts = torch.full((b,), time_range[0], device=device, dtype=torch.long)
|
||||
@@ -619,25 +716,23 @@ class SpacedSampler:
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
# TODO: tile_cond_fn
|
||||
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
|
||||
# accumulate noise
|
||||
noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise
|
||||
count[:, :, hi:hi_end, wi:wi_end] += 1
|
||||
|
||||
pbar.update(1)
|
||||
# average on noise (score)
|
||||
noise_buffer.div_(count)
|
||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
||||
tao_index = torch.tensor(torch.round(index * t_max),dtype=torch.int64)
|
||||
tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64)
|
||||
img = self.q_sample(pred_x0, tao_index)
|
||||
|
||||
noise_buffer.zero_()
|
||||
@@ -645,9 +740,9 @@ class SpacedSampler:
|
||||
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[-int(round(total_steps*t_max)):]
|
||||
time_range = time_range[-int(round(total_steps * t_max)):]
|
||||
total_steps_use = len(time_range)
|
||||
time_range = time_range[:-int(round(total_steps*t_min))]
|
||||
time_range = time_range[:-int(round(total_steps * t_min))]
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
|
||||
# sampling loop
|
||||
@@ -666,13 +761,12 @@ class SpacedSampler:
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
# TODO: tile_cond_fn
|
||||
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
@@ -726,6 +820,7 @@ class SpacedSampler:
|
||||
@torch.no_grad()
|
||||
def sample_with_mixdiff_control(
|
||||
self,
|
||||
empty_text_embed: torch.Tensor,
|
||||
control_imgs: torch.Tensor,
|
||||
tile_size: int,
|
||||
tile_stride: int,
|
||||
@@ -735,10 +830,9 @@ class SpacedSampler:
|
||||
cond_img: torch.Tensor,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
x_T: Optional[torch.Tensor]=None,
|
||||
cfg_scale: float=1.,
|
||||
cond_fn: Optional[Guidance]=None,
|
||||
color_fix_type: str="none"
|
||||
x_T: Optional[torch.Tensor] = None,
|
||||
cfg_scale: float = 1.,
|
||||
color_fix_type: str = "none"
|
||||
) -> torch.Tensor:
|
||||
def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]:
|
||||
hi_list = list(range(0, h - tile_size + 1, tile_stride))
|
||||
@@ -789,14 +883,12 @@ class SpacedSampler:
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
# TODO: tile_cond_fn
|
||||
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
|
||||
@@ -808,7 +900,7 @@ class SpacedSampler:
|
||||
noise_buffer.div_(count)
|
||||
# sample previous latent
|
||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
||||
tao_index = index - index//(tao_steps-1)
|
||||
tao_index = index - index // (tao_steps - 1)
|
||||
img = self.q_sample(pred_x0, tao_index)
|
||||
|
||||
noise_buffer.zero_()
|
||||
@@ -816,7 +908,7 @@ class SpacedSampler:
|
||||
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[total_steps//(tao_steps-1):]
|
||||
time_range = time_range[total_steps // (tao_steps - 1):]
|
||||
total_steps_use = len(time_range)
|
||||
# time_range = time_range[:-total_steps//(tao_steps-1)]
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
@@ -837,14 +929,12 @@ class SpacedSampler:
|
||||
tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||
tile_cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
tile_uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
# TODO: tile_cond_fn
|
||||
|
||||
# predict noise for this tile
|
||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||
|
||||
@@ -897,6 +987,7 @@ class SpacedSampler:
|
||||
@torch.no_grad()
|
||||
def sample_ccsr(
|
||||
self,
|
||||
empty_text_embed: torch.Tensor,
|
||||
steps: int,
|
||||
t_max: float,
|
||||
t_min: float,
|
||||
@@ -904,10 +995,9 @@ class SpacedSampler:
|
||||
cond_img: torch.Tensor,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
x_T: Optional[torch.Tensor]=None,
|
||||
cfg_scale: float=1.,
|
||||
cond_fn: Optional[Guidance]=None,
|
||||
color_fix_type: str="none"
|
||||
x_T: Optional[torch.Tensor] = None,
|
||||
cfg_scale: float = 1.,
|
||||
color_fix_type: str = "none"
|
||||
) -> torch.Tensor:
|
||||
self.make_schedule(num_steps=steps)
|
||||
# self.make_tao_schedule(num_steps=tao_steps)
|
||||
@@ -925,37 +1015,38 @@ class SpacedSampler:
|
||||
|
||||
cond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
uncond = {
|
||||
"c_latent": [self.model.apply_condition_encoder(cond_img)],
|
||||
"c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)]
|
||||
"c_crossattn": [empty_text_embed]
|
||||
}
|
||||
|
||||
# q_sample for the start
|
||||
ts = torch.full((b,), time_range[0], device=device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps - 1)
|
||||
img = self.p_sample_tao(
|
||||
img, cond, ts, index=index,t_max=t_max,
|
||||
cfg_scale=cfg_scale, uncond=uncond,
|
||||
cond_fn=cond_fn
|
||||
img, cond, ts, index=index, t_max=t_max,
|
||||
cfg_scale=cfg_scale, uncond=uncond
|
||||
)
|
||||
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[-int(round(total_steps*t_max)):]
|
||||
time_range = time_range[-int(round(total_steps * t_max)):]
|
||||
total_steps_use = len(time_range)
|
||||
time_range = time_range[:-int(round(total_steps*t_min))]
|
||||
time_range = time_range[:-int(round(total_steps * t_min))]
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
pbar = comfy.utils.ProgressBar(total_steps // 3)
|
||||
|
||||
for i, step in enumerate(iterator):
|
||||
|
||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps_use - i - 1)
|
||||
img, x0 = self.p_sample_x0(
|
||||
img, cond, ts, index=index,
|
||||
cfg_scale=cfg_scale, uncond=uncond,
|
||||
cond_fn=cond_fn
|
||||
cfg_scale=cfg_scale, uncond=uncond
|
||||
)
|
||||
pbar.update(1)
|
||||
|
||||
img = x0
|
||||
img_pixel = (self.model.decode_first_stage(img) + 1) / 2
|
||||
@@ -977,10 +1068,9 @@ class SpacedSampler:
|
||||
cond_img: torch.Tensor,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
x_T: Optional[torch.Tensor]=None,
|
||||
cfg_scale: float=1.,
|
||||
cond_fn: Optional[Guidance]=None,
|
||||
color_fix_type: str="none"
|
||||
x_T: Optional[torch.Tensor] = None,
|
||||
cfg_scale: float = 1.,
|
||||
color_fix_type: str = "none"
|
||||
) -> torch.Tensor:
|
||||
self.make_schedule(num_steps=steps)
|
||||
# self.make_tao_schedule(num_steps=tao_steps)
|
||||
@@ -1009,14 +1099,13 @@ class SpacedSampler:
|
||||
ts = torch.full((b,), time_range[0], device=device, dtype=torch.long)
|
||||
index = torch.full_like(ts, fill_value=total_steps - 1)
|
||||
img = self.p_sample_tao(
|
||||
img, cond, ts, index=index,t_max=t_max,
|
||||
cfg_scale=cfg_scale, uncond=uncond,
|
||||
cond_fn=cond_fn
|
||||
img, cond, ts, index=index, t_max=t_max,
|
||||
cfg_scale=cfg_scale, uncond=uncond
|
||||
)
|
||||
|
||||
time_range = np.flip(self.timesteps) # [1000, 950, 900, ...]
|
||||
total_steps = len(time_range)
|
||||
time_range = time_range[-int(round(total_steps*t_max)):]
|
||||
time_range = time_range[-int(round(total_steps * t_max)):]
|
||||
total_steps = len(time_range)
|
||||
iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps)
|
||||
|
||||
@@ -1025,8 +1114,7 @@ class SpacedSampler:
|
||||
index = torch.full_like(ts, fill_value=total_steps - i - 1)
|
||||
img = self.p_sample(
|
||||
img, cond, ts, index=index,
|
||||
cfg_scale=cfg_scale, uncond=uncond,
|
||||
cond_fn=cond_fn
|
||||
cfg_scale=cfg_scale, uncond=uncond
|
||||
)
|
||||
|
||||
img_pixel = (self.model.decode_first_stage(img) + 1) / 2
|
||||
|
||||
@@ -26,12 +26,21 @@ class CCSR_Upscale:
|
||||
"image": ("IMAGE", ),
|
||||
"resize_method": (s.upscale_methods, {"default": "lanczos"}),
|
||||
"scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 20.0, "step": 0.01}),
|
||||
"steps": ("INT", {"default": 45, "min": 2, "max": 4096, "step": 1}),
|
||||
"steps": ("INT", {"default": 45, "min": 3, "max": 4096, "step": 1}),
|
||||
"t_max": ("FLOAT", {"default": 0.6667,"min": 0, "max": 1, "step": 0.01}),
|
||||
"t_min": ("FLOAT", {"default": 0.3333,"min": 0, "max": 1, "step": 0.01}),
|
||||
"sampling_method": (
|
||||
[
|
||||
'ccsr',
|
||||
'ccsr_tiled_mixdiff',
|
||||
'ccsr_tiled_vae_gaussian_weights',
|
||||
], {
|
||||
"default": 'ccsr_tiled_mixdiff'
|
||||
}),
|
||||
"tile_size": ("INT", {"default": 512, "min": 1, "max": 4096, "step": 1}),
|
||||
"tile_stride": ("INT", {"default": 256, "min": 1, "max": 4096, "step": 1}),
|
||||
"tiled": ("BOOLEAN", {"default": False}),
|
||||
"vae_tile_size_encode": ("INT", {"default": 1024, "min": 2, "max": 4096, "step": 8}),
|
||||
"vae_tile_size_decode": ("INT", {"default": 1024, "min": 2, "max": 4096, "step": 8}),
|
||||
"color_fix_type": (
|
||||
[
|
||||
'none',
|
||||
@@ -41,6 +50,7 @@ class CCSR_Upscale:
|
||||
"default": 'adain'
|
||||
}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
|
||||
},
|
||||
|
||||
}
|
||||
@@ -52,12 +62,12 @@ class CCSR_Upscale:
|
||||
CATEGORY = "CCSR"
|
||||
|
||||
@torch.no_grad()
|
||||
def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tiled,tile_size, tile_stride, color_fix_type, keep_model_loaded):
|
||||
def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tile_size, tile_stride, color_fix_type, keep_model_loaded, vae_tile_size_encode, vae_tile_size_decode, sampling_method):
|
||||
comfy.model_management.unload_all_models()
|
||||
device = comfy.model_management.get_torch_device()
|
||||
config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml")
|
||||
empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device)
|
||||
dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32
|
||||
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
config = OmegaConf.load(config_path)
|
||||
self.model = instantiate_from_config(config)
|
||||
@@ -69,6 +79,7 @@ class CCSR_Upscale:
|
||||
self.model.to(device, dtype=dtype)
|
||||
sampler = SpacedSampler(self.model, var_type="fixed_small")
|
||||
|
||||
batch_size = image.shape[0]
|
||||
image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
||||
|
||||
# Assuming 'image' is a PyTorch tensor with shape [B, H, W, C] and you want to resize it.
|
||||
@@ -88,33 +99,51 @@ class CCSR_Upscale:
|
||||
resized_image = resized_image.to(device)
|
||||
strength = 1.0
|
||||
self.model.control_scales = [strength] * 13
|
||||
cond_fn = None
|
||||
|
||||
height, width = resized_image.size(-2), resized_image.size(-1)
|
||||
shape = (1, 4, height // 8, width // 8)
|
||||
x_T = torch.randn(shape, device=self.model.device, dtype=dtype)
|
||||
x_T = torch.randn(shape, device=self.model.device, dtype=torch.float32)
|
||||
autocast_condition = dtype == torch.float16 and not comfy.model_management.is_device_mps(device)
|
||||
out = []
|
||||
|
||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||
if not tiled:
|
||||
samples = sampler.sample_ccsr(
|
||||
steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image,
|
||||
for i in range(batch_size):
|
||||
|
||||
|
||||
if sampling_method == 'ccsr_tiled_mixdiff':
|
||||
print("Using tiled mixdiff")
|
||||
samples = sampler.sample_with_mixdiff_ccsr(
|
||||
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
||||
steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0),
|
||||
positive_prompt="", negative_prompt="", x_T=x_T,
|
||||
cfg_scale=1.0, cond_fn=cond_fn,
|
||||
cfg_scale=1.0,
|
||||
color_fix_type=color_fix_type
|
||||
)
|
||||
elif sampling_method == 'ccsr_tiled_vae_gaussian_weights':
|
||||
self.model._init_tiled_vae(encoder_tile_size=vae_tile_size_encode // 8, decoder_tile_size=vae_tile_size_decode // 8)
|
||||
print("Using gaussian weights")
|
||||
samples = sampler.sample_with_tile_ccsr(
|
||||
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
||||
steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0),
|
||||
positive_prompt="", negative_prompt="", x_T=x_T,
|
||||
cfg_scale=1.0,
|
||||
color_fix_type=color_fix_type
|
||||
)
|
||||
else:
|
||||
samples = sampler.sample_with_mixdiff_ccsr(
|
||||
tile_size=tile_size, tile_stride=tile_stride,
|
||||
steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image,
|
||||
print("no tiling")
|
||||
samples = sampler.sample_ccsr(
|
||||
empty_text_embed, steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0),
|
||||
positive_prompt="", negative_prompt="", x_T=x_T,
|
||||
cfg_scale=1.0, cond_fn=cond_fn,
|
||||
cfg_scale=1.0,
|
||||
color_fix_type=color_fix_type
|
||||
)
|
||||
out.append(samples.squeeze(0))
|
||||
|
||||
original_height, original_width = H, W
|
||||
processed_height = samples.size(2)
|
||||
target_width = int(processed_height * (original_width / original_height))
|
||||
|
||||
resized_back_image, = ImageScale.upscale(self, samples.permute(0, 2, 3, 1).cpu(), "lanczos", target_width, processed_height, crop="disabled")
|
||||
out_stacked = torch.stack(out, dim=0).cpu().to(torch.float32).permute(0, 2, 3, 1)
|
||||
resized_back_image, = ImageScale.upscale(self, out_stacked, "lanczos", target_width, processed_height, crop="disabled")
|
||||
|
||||
if not keep_model_loaded:
|
||||
self.model = None
|
||||
|
||||
@@ -0,0 +1,719 @@
|
||||
'''
|
||||
# ------------------------------------------------------------------------
|
||||
#
|
||||
# Tiled VAE
|
||||
#
|
||||
# Introducing a revolutionary new optimization designed to make
|
||||
# the VAE work with giant images on limited VRAM!
|
||||
# Say goodbye to the frustration of OOM and hello to seamless output!
|
||||
#
|
||||
# ------------------------------------------------------------------------
|
||||
#
|
||||
# This script is a wild hack that splits the image into tiles,
|
||||
# encodes each tile separately, and merges the result back together.
|
||||
#
|
||||
# Advantages:
|
||||
# - The VAE can now work with giant images on limited VRAM
|
||||
# (~10 GB for 8K images!)
|
||||
# - The merged output is completely seamless without any post-processing.
|
||||
#
|
||||
# Drawbacks:
|
||||
# - NaNs always appear in for 8k images when you use fp16 (half) VAE
|
||||
# You must use --no-half-vae to disable half VAE for that giant image.
|
||||
# - The gradient calculation is not compatible with this hack. It
|
||||
# will break any backward() or torch.autograd.grad() that passes VAE.
|
||||
# (But you can still use the VAE to generate training data.)
|
||||
#
|
||||
# How it works:
|
||||
# 1. The image is split into tiles, which are then padded with 11/32 pixels' in the decoder/encoder.
|
||||
# 2. When Fast Mode is disabled:
|
||||
# 1. The original VAE forward is decomposed into a task queue and a task worker, which starts to process each tile.
|
||||
# 2. When GroupNorm is needed, it suspends, stores current GroupNorm mean and var, send everything to RAM, and turns to the next tile.
|
||||
# 3. After all GroupNorm means and vars are summarized, it applies group norm to tiles and continues.
|
||||
# 4. A zigzag execution order is used to reduce unnecessary data transfer.
|
||||
# 3. When Fast Mode is enabled:
|
||||
# 1. The original input is downsampled and passed to a separate task queue.
|
||||
# 2. Its group norm parameters are recorded and used by all tiles' task queues.
|
||||
# 3. Each tile is separately processed without any RAM-VRAM data transfer.
|
||||
# 4. After all tiles are processed, tiles are written to a result buffer and returned.
|
||||
# Encoder color fix = only estimate GroupNorm before downsampling, i.e., run in a semi-fast mode.
|
||||
#
|
||||
# Enjoy!
|
||||
#
|
||||
# @Author: LI YI @ Nanyang Technological University - Singapore
|
||||
# @Date: 2023-03-02
|
||||
# @License: CC BY-NC-SA 4.0
|
||||
#
|
||||
# Please give me a star if you like this project!
|
||||
#
|
||||
# -------------------------------------------------------------------------
|
||||
'''
|
||||
|
||||
import gc
|
||||
import math
|
||||
import sys
|
||||
from time import time
|
||||
from tqdm import tqdm
|
||||
|
||||
import torch
|
||||
import torch.version
|
||||
import torch.nn.functional as F
|
||||
|
||||
cpu = torch.device("cpu")
|
||||
device = device_interrogate = device_gfpgan = device_esrgan = device_codeformer = torch.device("cuda")
|
||||
dtype = torch.float16
|
||||
dtype_vae = torch.float16
|
||||
dtype_unet = torch.float16
|
||||
unet_needs_upcast = False
|
||||
|
||||
def torch_gc():
|
||||
|
||||
if torch.cuda.is_available():
|
||||
with torch.cuda.device("cuda"):
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
if has_mps():
|
||||
mac_specific.torch_mps_gc()
|
||||
|
||||
def has_mps() -> bool:
|
||||
if sys.platform != "darwin":
|
||||
return False
|
||||
else:
|
||||
return mac_specific.has_mps
|
||||
|
||||
def get_optimal_device_name():
|
||||
if torch.cuda.is_available():
|
||||
return "cuda"
|
||||
|
||||
if has_mps():
|
||||
return "mps"
|
||||
|
||||
return "cpu"
|
||||
|
||||
|
||||
def get_optimal_device():
|
||||
return torch.device(get_optimal_device_name())
|
||||
|
||||
class NansException(Exception):
|
||||
pass
|
||||
|
||||
def test_for_nans(x, where):
|
||||
if not torch.all(torch.isnan(x)).item():
|
||||
return
|
||||
|
||||
if where == "unet":
|
||||
message = "A tensor with all NaNs was produced in Unet."
|
||||
|
||||
elif where == "vae":
|
||||
message = "A tensor with all NaNs was produced in VAE."
|
||||
|
||||
else:
|
||||
message = "A tensor with all NaNs was produced."
|
||||
|
||||
message += " Use --disable-nan-check commandline argument to disable this check."
|
||||
|
||||
raise NansException(message)
|
||||
|
||||
def get_rcmd_enc_tsize():
|
||||
if torch.cuda.is_available() and device not in ['cpu', cpu]:
|
||||
total_memory = torch.cuda.get_device_properties(device).total_memory // 2**20
|
||||
if total_memory > 16*1000: ENCODER_TILE_SIZE = 3072
|
||||
elif total_memory > 12*1000: ENCODER_TILE_SIZE = 2048
|
||||
elif total_memory > 8*1000: ENCODER_TILE_SIZE = 1536
|
||||
else: ENCODER_TILE_SIZE = 960
|
||||
else: ENCODER_TILE_SIZE = 512
|
||||
return ENCODER_TILE_SIZE
|
||||
|
||||
|
||||
def get_rcmd_dec_tsize():
|
||||
if torch.cuda.is_available() and device not in ['cpu', cpu]:
|
||||
total_memory = torch.cuda.get_device_properties(device).total_memory // 2**20
|
||||
if total_memory > 30*1000: DECODER_TILE_SIZE = 256
|
||||
elif total_memory > 16*1000: DECODER_TILE_SIZE = 192
|
||||
elif total_memory > 12*1000: DECODER_TILE_SIZE = 128
|
||||
elif total_memory > 8*1000: DECODER_TILE_SIZE = 96
|
||||
else: DECODER_TILE_SIZE = 64
|
||||
else: DECODER_TILE_SIZE = 64
|
||||
return DECODER_TILE_SIZE
|
||||
|
||||
|
||||
def inplace_nonlinearity(x):
|
||||
# Test: fix for Nans
|
||||
return F.silu(x, inplace=True)
|
||||
|
||||
def attn_forward(self, h_):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b, c, h, w = q.shape
|
||||
q = q.reshape(b, c, h*w)
|
||||
q = q.permute(0, 2, 1) # b,hw,c
|
||||
k = k.reshape(b, c, h*w) # b,c,hw
|
||||
w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b, c, h*w)
|
||||
w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q)
|
||||
# b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
|
||||
h_ = torch.bmm(v, w_)
|
||||
h_ = h_.reshape(b, c, h, w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return h_
|
||||
|
||||
|
||||
def attn2task(task_queue, net):
|
||||
task_queue.append(('store_res', lambda x: x))
|
||||
task_queue.append(('pre_norm', net.norm))
|
||||
task_queue.append(('attn', lambda x, net=net: attn_forward(net, x)))
|
||||
task_queue.append(['add_res', None])
|
||||
|
||||
|
||||
def resblock2task(queue, block):
|
||||
"""
|
||||
Turn a ResNetBlock into a sequence of tasks and append to the task queue
|
||||
|
||||
@param queue: the target task queue
|
||||
@param block: ResNetBlock
|
||||
|
||||
"""
|
||||
if block.in_channels != block.out_channels:
|
||||
if block.use_conv_shortcut:
|
||||
queue.append(('store_res', block.conv_shortcut))
|
||||
else:
|
||||
queue.append(('store_res', block.nin_shortcut))
|
||||
else:
|
||||
queue.append(('store_res', lambda x: x))
|
||||
queue.append(('pre_norm', block.norm1))
|
||||
queue.append(('silu', inplace_nonlinearity))
|
||||
queue.append(('conv1', block.conv1))
|
||||
queue.append(('pre_norm', block.norm2))
|
||||
queue.append(('silu', inplace_nonlinearity))
|
||||
queue.append(('conv2', block.conv2))
|
||||
queue.append(['add_res', None])
|
||||
|
||||
|
||||
def build_sampling(task_queue, net, is_decoder):
|
||||
"""
|
||||
Build the sampling part of a task queue
|
||||
@param task_queue: the target task queue
|
||||
@param net: the network
|
||||
@param is_decoder: currently building decoder or encoder
|
||||
"""
|
||||
if is_decoder:
|
||||
resblock2task(task_queue, net.mid.block_1)
|
||||
attn2task(task_queue, net.mid.attn_1)
|
||||
resblock2task(task_queue, net.mid.block_2)
|
||||
resolution_iter = reversed(range(net.num_resolutions))
|
||||
block_ids = net.num_res_blocks + 1
|
||||
condition = 0
|
||||
module = net.up
|
||||
func_name = 'upsample'
|
||||
else:
|
||||
resolution_iter = range(net.num_resolutions)
|
||||
block_ids = net.num_res_blocks
|
||||
condition = net.num_resolutions - 1
|
||||
module = net.down
|
||||
func_name = 'downsample'
|
||||
|
||||
for i_level in resolution_iter:
|
||||
for i_block in range(block_ids):
|
||||
resblock2task(task_queue, module[i_level].block[i_block])
|
||||
if i_level != condition:
|
||||
task_queue.append((func_name, getattr(module[i_level], func_name)))
|
||||
|
||||
if not is_decoder:
|
||||
resblock2task(task_queue, net.mid.block_1)
|
||||
attn2task(task_queue, net.mid.attn_1)
|
||||
resblock2task(task_queue, net.mid.block_2)
|
||||
|
||||
|
||||
def build_task_queue(net, is_decoder):
|
||||
"""
|
||||
Build a single task queue for the encoder or decoder
|
||||
@param net: the VAE decoder or encoder network
|
||||
@param is_decoder: currently building decoder or encoder
|
||||
@return: the task queue
|
||||
"""
|
||||
task_queue = []
|
||||
task_queue.append(('conv_in', net.conv_in))
|
||||
|
||||
# construct the sampling part of the task queue
|
||||
# because encoder and decoder share the same architecture, we extract the sampling part
|
||||
build_sampling(task_queue, net, is_decoder)
|
||||
|
||||
if not is_decoder or not net.give_pre_end:
|
||||
task_queue.append(('pre_norm', net.norm_out))
|
||||
task_queue.append(('silu', inplace_nonlinearity))
|
||||
task_queue.append(('conv_out', net.conv_out))
|
||||
if is_decoder and net.tanh_out:
|
||||
task_queue.append(('tanh', torch.tanh))
|
||||
|
||||
return task_queue
|
||||
|
||||
|
||||
def clone_task_queue(task_queue):
|
||||
"""
|
||||
Clone a task queue
|
||||
@param task_queue: the task queue to be cloned
|
||||
@return: the cloned task queue
|
||||
"""
|
||||
return [[item for item in task] for task in task_queue]
|
||||
|
||||
|
||||
def get_var_mean(input, num_groups, eps=1e-6):
|
||||
"""
|
||||
Get mean and var for group norm
|
||||
"""
|
||||
b, c = input.size(0), input.size(1)
|
||||
channel_in_group = int(c/num_groups)
|
||||
input_reshaped = input.contiguous().view(1, int(b * num_groups), channel_in_group, *input.size()[2:])
|
||||
var, mean = torch.var_mean(input_reshaped, dim=[0, 2, 3, 4], unbiased=False)
|
||||
return var, mean
|
||||
|
||||
|
||||
def custom_group_norm(input, num_groups, mean, var, weight=None, bias=None, eps=1e-6):
|
||||
"""
|
||||
Custom group norm with fixed mean and var
|
||||
|
||||
@param input: input tensor
|
||||
@param num_groups: number of groups. by default, num_groups = 32
|
||||
@param mean: mean, must be pre-calculated by get_var_mean
|
||||
@param var: var, must be pre-calculated by get_var_mean
|
||||
@param weight: weight, should be fetched from the original group norm
|
||||
@param bias: bias, should be fetched from the original group norm
|
||||
@param eps: epsilon, by default, eps = 1e-6 to match the original group norm
|
||||
|
||||
@return: normalized tensor
|
||||
"""
|
||||
b, c = input.size(0), input.size(1)
|
||||
channel_in_group = int(c/num_groups)
|
||||
input_reshaped = input.contiguous().view(
|
||||
1, int(b * num_groups), channel_in_group, *input.size()[2:])
|
||||
|
||||
out = F.batch_norm(input_reshaped, mean, var, weight=None, bias=None, training=False, momentum=0, eps=eps)
|
||||
out = out.view(b, c, *input.size()[2:])
|
||||
|
||||
# post affine transform
|
||||
if weight is not None:
|
||||
out *= weight.view(1, -1, 1, 1)
|
||||
if bias is not None:
|
||||
out += bias.view(1, -1, 1, 1)
|
||||
return out
|
||||
|
||||
|
||||
def crop_valid_region(x, input_bbox, target_bbox, is_decoder):
|
||||
"""
|
||||
Crop the valid region from the tile
|
||||
@param x: input tile
|
||||
@param input_bbox: original input bounding box
|
||||
@param target_bbox: output bounding box
|
||||
@param scale: scale factor
|
||||
@return: cropped tile
|
||||
"""
|
||||
padded_bbox = [i * 8 if is_decoder else i//8 for i in input_bbox]
|
||||
margin = [target_bbox[i] - padded_bbox[i] for i in range(4)]
|
||||
return x[:, :, margin[2]:x.size(2)+margin[3], margin[0]:x.size(3)+margin[1]]
|
||||
|
||||
|
||||
# ↓↓↓ https://github.com/Kahsolt/stable-diffusion-webui-vae-tile-infer ↓↓↓
|
||||
|
||||
def perfcount(fn):
|
||||
def wrapper(*args, **kwargs):
|
||||
ts = time()
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
torch_gc()
|
||||
gc.collect()
|
||||
|
||||
ret = fn(*args, **kwargs)
|
||||
|
||||
torch_gc()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
vram = torch.cuda.max_memory_allocated(device) / 2**20
|
||||
print(f'[Tiled VAE]: Done in {time() - ts:.3f}s, max VRAM alloc {vram:.3f} MB')
|
||||
else:
|
||||
print(f'[Tiled VAE]: Done in {time() - ts:.3f}s')
|
||||
|
||||
return ret
|
||||
return wrapper
|
||||
|
||||
# ↑↑↑ https://github.com/Kahsolt/stable-diffusion-webui-vae-tile-infer ↑↑↑
|
||||
|
||||
|
||||
class GroupNormParam:
|
||||
|
||||
def __init__(self):
|
||||
self.var_list = []
|
||||
self.mean_list = []
|
||||
self.pixel_list = []
|
||||
self.weight = None
|
||||
self.bias = None
|
||||
|
||||
def add_tile(self, tile, layer):
|
||||
var, mean = get_var_mean(tile, 32)
|
||||
# For giant images, the variance can be larger than max float16
|
||||
# In this case we create a copy to float32
|
||||
if var.dtype == torch.float16 and var.isinf().any():
|
||||
fp32_tile = tile.float()
|
||||
var, mean = get_var_mean(fp32_tile, 32)
|
||||
# ============= DEBUG: test for infinite =============
|
||||
# if torch.isinf(var).any():
|
||||
# print('var: ', var)
|
||||
# ====================================================
|
||||
self.var_list.append(var)
|
||||
self.mean_list.append(mean)
|
||||
self.pixel_list.append(
|
||||
tile.shape[2]*tile.shape[3])
|
||||
if hasattr(layer, 'weight'):
|
||||
self.weight = layer.weight
|
||||
self.bias = layer.bias
|
||||
else:
|
||||
self.weight = None
|
||||
self.bias = None
|
||||
|
||||
def summary(self):
|
||||
"""
|
||||
summarize the mean and var and return a function
|
||||
that apply group norm on each tile
|
||||
"""
|
||||
if len(self.var_list) == 0: return None
|
||||
|
||||
var = torch.vstack(self.var_list)
|
||||
mean = torch.vstack(self.mean_list)
|
||||
max_value = max(self.pixel_list)
|
||||
pixels = torch.tensor(self.pixel_list, dtype=torch.float32, device=device) / max_value
|
||||
sum_pixels = torch.sum(pixels)
|
||||
pixels = pixels.unsqueeze(1) / sum_pixels
|
||||
var = torch.sum(var * pixels, dim=0)
|
||||
mean = torch.sum(mean * pixels, dim=0)
|
||||
return lambda x: custom_group_norm(x, 32, mean, var, self.weight, self.bias)
|
||||
|
||||
@staticmethod
|
||||
def from_tile(tile, norm):
|
||||
"""
|
||||
create a function from a single tile without summary
|
||||
"""
|
||||
var, mean = get_var_mean(tile, 32)
|
||||
if var.dtype == torch.float16 and var.isinf().any():
|
||||
fp32_tile = tile.float()
|
||||
var, mean = get_var_mean(fp32_tile, 32)
|
||||
# if it is a macbook, we need to convert back to float16
|
||||
if var.device.type == 'mps':
|
||||
# clamp to avoid overflow
|
||||
var = torch.clamp(var, 0, 60000)
|
||||
var = var.half()
|
||||
mean = mean.half()
|
||||
if hasattr(norm, 'weight'):
|
||||
weight = norm.weight
|
||||
bias = norm.bias
|
||||
else:
|
||||
weight = None
|
||||
bias = None
|
||||
|
||||
def group_norm_func(x, mean=mean, var=var, weight=weight, bias=bias):
|
||||
return custom_group_norm(x, 32, mean, var, weight, bias, 1e-6)
|
||||
return group_norm_func
|
||||
|
||||
|
||||
class VAEHook:
|
||||
|
||||
def __init__(self, net, tile_size, is_decoder:bool, fast_decoder:bool, fast_encoder:bool, color_fix:bool, to_gpu:bool=False):
|
||||
self.net = net # encoder | decoder
|
||||
self.tile_size = tile_size
|
||||
self.is_decoder = is_decoder
|
||||
self.fast_mode = (fast_encoder and not is_decoder) or (fast_decoder and is_decoder)
|
||||
self.color_fix = color_fix and not is_decoder
|
||||
self.to_gpu = to_gpu
|
||||
self.pad = 11 if is_decoder else 32 # FIXME: magic number
|
||||
|
||||
def __call__(self, x):
|
||||
original_device = next(self.net.parameters()).device
|
||||
try:
|
||||
if self.to_gpu:
|
||||
self.net = self.net.to(get_optimal_device())
|
||||
|
||||
B, C, H, W = x.shape
|
||||
if max(H, W) <= self.pad * 2 + self.tile_size:
|
||||
print("[Tiled VAE]: the input size is tiny and unnecessary to tile.")
|
||||
return self.net.original_forward(x)
|
||||
else:
|
||||
print("Using tiled VAE Decoder")
|
||||
return self.vae_tile_forward(x)
|
||||
finally:
|
||||
self.net = self.net.to(original_device)
|
||||
|
||||
def get_best_tile_size(self, lowerbound, upperbound):
|
||||
"""
|
||||
Get the best tile size for GPU memory
|
||||
"""
|
||||
divider = 32
|
||||
while divider >= 2:
|
||||
remainer = lowerbound % divider
|
||||
if remainer == 0:
|
||||
return lowerbound
|
||||
candidate = lowerbound - remainer + divider
|
||||
if candidate <= upperbound:
|
||||
return candidate
|
||||
divider //= 2
|
||||
return lowerbound
|
||||
|
||||
def split_tiles(self, h, w):
|
||||
"""
|
||||
Tool function to split the image into tiles
|
||||
@param h: height of the image
|
||||
@param w: width of the image
|
||||
@return: tile_input_bboxes, tile_output_bboxes
|
||||
"""
|
||||
tile_input_bboxes, tile_output_bboxes = [], []
|
||||
tile_size = self.tile_size
|
||||
pad = self.pad
|
||||
num_height_tiles = math.ceil((h - 2 * pad) / tile_size)
|
||||
num_width_tiles = math.ceil((w - 2 * pad) / tile_size)
|
||||
# If any of the numbers are 0, we let it be 1
|
||||
# This is to deal with long and thin images
|
||||
num_height_tiles = max(num_height_tiles, 1)
|
||||
num_width_tiles = max(num_width_tiles, 1)
|
||||
|
||||
# Suggestions from https://github.com/Kahsolt: auto shrink the tile size
|
||||
real_tile_height = math.ceil((h - 2 * pad) / num_height_tiles)
|
||||
real_tile_width = math.ceil((w - 2 * pad) / num_width_tiles)
|
||||
real_tile_height = self.get_best_tile_size(real_tile_height, tile_size)
|
||||
real_tile_width = self.get_best_tile_size(real_tile_width, tile_size)
|
||||
|
||||
print(f'[Tiled VAE]: split to {num_height_tiles}x{num_width_tiles} = {num_height_tiles*num_width_tiles} tiles. ' +
|
||||
f'Optimal tile size {real_tile_width}x{real_tile_height}, original tile size {tile_size}x{tile_size}')
|
||||
|
||||
for i in range(num_height_tiles):
|
||||
for j in range(num_width_tiles):
|
||||
# bbox: [x1, x2, y1, y2]
|
||||
# the padding is is unnessary for image borders. So we directly start from (32, 32)
|
||||
input_bbox = [
|
||||
pad + j * real_tile_width,
|
||||
min(pad + (j + 1) * real_tile_width, w),
|
||||
pad + i * real_tile_height,
|
||||
min(pad + (i + 1) * real_tile_height, h),
|
||||
]
|
||||
|
||||
# if the output bbox is close to the image boundary, we extend it to the image boundary
|
||||
output_bbox = [
|
||||
input_bbox[0] if input_bbox[0] > pad else 0,
|
||||
input_bbox[1] if input_bbox[1] < w - pad else w,
|
||||
input_bbox[2] if input_bbox[2] > pad else 0,
|
||||
input_bbox[3] if input_bbox[3] < h - pad else h,
|
||||
]
|
||||
|
||||
# scale to get the final output bbox
|
||||
output_bbox = [x * 8 if self.is_decoder else x // 8 for x in output_bbox]
|
||||
tile_output_bboxes.append(output_bbox)
|
||||
|
||||
# indistinguishable expand the input bbox by pad pixels
|
||||
tile_input_bboxes.append([
|
||||
max(0, input_bbox[0] - pad),
|
||||
min(w, input_bbox[1] + pad),
|
||||
max(0, input_bbox[2] - pad),
|
||||
min(h, input_bbox[3] + pad),
|
||||
])
|
||||
|
||||
return tile_input_bboxes, tile_output_bboxes
|
||||
|
||||
@torch.no_grad()
|
||||
def estimate_group_norm(self, z, task_queue, color_fix):
|
||||
device = z.device
|
||||
tile = z
|
||||
last_id = len(task_queue) - 1
|
||||
while last_id >= 0 and task_queue[last_id][0] != 'pre_norm':
|
||||
last_id -= 1
|
||||
if last_id <= 0 or task_queue[last_id][0] != 'pre_norm':
|
||||
raise ValueError('No group norm found in the task queue')
|
||||
# estimate until the last group norm
|
||||
for i in range(last_id + 1):
|
||||
task = task_queue[i]
|
||||
if task[0] == 'pre_norm':
|
||||
group_norm_func = GroupNormParam.from_tile(tile, task[1])
|
||||
task_queue[i] = ('apply_norm', group_norm_func)
|
||||
if i == last_id:
|
||||
return True
|
||||
tile = group_norm_func(tile)
|
||||
elif task[0] == 'store_res':
|
||||
task_id = i + 1
|
||||
while task_id < last_id and task_queue[task_id][0] != 'add_res':
|
||||
task_id += 1
|
||||
if task_id >= last_id:
|
||||
continue
|
||||
task_queue[task_id][1] = task[1](tile)
|
||||
elif task[0] == 'add_res':
|
||||
tile += task[1].to(device)
|
||||
task[1] = None
|
||||
elif color_fix and task[0] == 'downsample':
|
||||
for j in range(i, last_id + 1):
|
||||
if task_queue[j][0] == 'store_res':
|
||||
task_queue[j] = ('store_res_cpu', task_queue[j][1])
|
||||
return True
|
||||
else:
|
||||
tile = task[1](tile)
|
||||
try:
|
||||
test_for_nans(tile, "vae")
|
||||
except:
|
||||
print(f'Nan detected in fast mode estimation. Fast mode disabled.')
|
||||
return False
|
||||
|
||||
raise IndexError('Should not reach here')
|
||||
|
||||
@perfcount
|
||||
@torch.no_grad()
|
||||
def vae_tile_forward(self, z):
|
||||
"""
|
||||
Decode a latent vector z into an image in a tiled manner.
|
||||
@param z: latent vector
|
||||
@return: image
|
||||
"""
|
||||
device = next(self.net.parameters()).device
|
||||
net = self.net
|
||||
tile_size = self.tile_size
|
||||
is_decoder = self.is_decoder
|
||||
|
||||
z = z.detach() # detach the input to avoid backprop
|
||||
|
||||
N, height, width = z.shape[0], z.shape[2], z.shape[3]
|
||||
net.last_z_shape = z.shape
|
||||
|
||||
# Split the input into tiles and build a task queue for each tile
|
||||
print(f'[Tiled VAE]: input_size: {z.shape}, tile_size: {tile_size}, padding: {self.pad}')
|
||||
|
||||
in_bboxes, out_bboxes = self.split_tiles(height, width)
|
||||
|
||||
# Prepare tiles by split the input latents
|
||||
tiles = []
|
||||
for input_bbox in in_bboxes:
|
||||
tile = z[:, :, input_bbox[2]:input_bbox[3], input_bbox[0]:input_bbox[1]].cpu()
|
||||
tiles.append(tile)
|
||||
|
||||
num_tiles = len(tiles)
|
||||
num_completed = 0
|
||||
|
||||
# Build task queues
|
||||
single_task_queue = build_task_queue(net, is_decoder)
|
||||
if self.fast_mode:
|
||||
# Fast mode: downsample the input image to the tile size,
|
||||
# then estimate the group norm parameters on the downsampled image
|
||||
scale_factor = tile_size / max(height, width)
|
||||
z = z.to(device)
|
||||
downsampled_z = F.interpolate(z, scale_factor=scale_factor, mode='nearest-exact')
|
||||
# use nearest-exact to keep statictics as close as possible
|
||||
print(f'[Tiled VAE]: Fast mode enabled, estimating group norm parameters on {downsampled_z.shape[3]} x {downsampled_z.shape[2]} image')
|
||||
|
||||
# ======= Special thanks to @Kahsolt for distribution shift issue ======= #
|
||||
# The downsampling will heavily distort its mean and std, so we need to recover it.
|
||||
std_old, mean_old = torch.std_mean(z, dim=[0, 2, 3], keepdim=True)
|
||||
std_new, mean_new = torch.std_mean(downsampled_z, dim=[0, 2, 3], keepdim=True)
|
||||
downsampled_z = (downsampled_z - mean_new) / std_new * std_old + mean_old
|
||||
del std_old, mean_old, std_new, mean_new
|
||||
# occasionally the std_new is too small or too large, which exceeds the range of float16
|
||||
# so we need to clamp it to max z's range.
|
||||
downsampled_z = torch.clamp_(downsampled_z, min=z.min(), max=z.max())
|
||||
estimate_task_queue = clone_task_queue(single_task_queue)
|
||||
if self.estimate_group_norm(downsampled_z, estimate_task_queue, color_fix=self.color_fix):
|
||||
single_task_queue = estimate_task_queue
|
||||
del downsampled_z
|
||||
|
||||
task_queues = [clone_task_queue(single_task_queue) for _ in range(num_tiles)]
|
||||
|
||||
# Dummy result
|
||||
result = None
|
||||
result_approx = None
|
||||
# try:
|
||||
# with devices.autocast():
|
||||
# result_approx = torch.cat([F.interpolate(cheap_approximation(x).unsqueeze(0), scale_factor=opt_f, mode='nearest-exact') for x in z], dim=0).cpu()
|
||||
# except: pass
|
||||
# Free memory of input latent tensor
|
||||
del z
|
||||
|
||||
# Task queue execution
|
||||
pbar = tqdm(total=num_tiles * len(task_queues[0]), desc=f"[Tiled VAE]: Executing {'Decoder' if is_decoder else 'Encoder'} Task Queue: ")
|
||||
|
||||
# execute the task back and forth when switch tiles so that we always
|
||||
# keep one tile on the GPU to reduce unnecessary data transfer
|
||||
forward = True
|
||||
interrupted = False
|
||||
#state.interrupted = interrupted
|
||||
while True:
|
||||
# if state.interrupted: interrupted = True ; break
|
||||
|
||||
group_norm_param = GroupNormParam()
|
||||
for i in range(num_tiles) if forward else reversed(range(num_tiles)):
|
||||
# if state.interrupted: interrupted = True ; break
|
||||
|
||||
tile = tiles[i].to(device)
|
||||
input_bbox = in_bboxes[i]
|
||||
task_queue = task_queues[i]
|
||||
|
||||
interrupted = False
|
||||
while len(task_queue) > 0:
|
||||
# if state.interrupted: interrupted = True ; break
|
||||
|
||||
# DEBUG: current task
|
||||
# print('Running task: ', task_queue[0][0], ' on tile ', i, '/', num_tiles, ' with shape ', tile.shape)
|
||||
task = task_queue.pop(0)
|
||||
if task[0] == 'pre_norm':
|
||||
group_norm_param.add_tile(tile, task[1])
|
||||
break
|
||||
elif task[0] == 'store_res' or task[0] == 'store_res_cpu':
|
||||
task_id = 0
|
||||
res = task[1](tile)
|
||||
if not self.fast_mode or task[0] == 'store_res_cpu':
|
||||
res = res.cpu()
|
||||
while task_queue[task_id][0] != 'add_res':
|
||||
task_id += 1
|
||||
task_queue[task_id][1] = res
|
||||
elif task[0] == 'add_res':
|
||||
tile += task[1].to(device)
|
||||
task[1] = None
|
||||
else:
|
||||
tile = task[1](tile)
|
||||
pbar.update(1)
|
||||
|
||||
if interrupted: break
|
||||
|
||||
# check for NaNs in the tile.
|
||||
# If there are NaNs, we abort the process to save user's time
|
||||
# devices.test_for_nans(tile, "vae")
|
||||
|
||||
if len(task_queue) == 0:
|
||||
tiles[i] = None
|
||||
num_completed += 1
|
||||
if result is None: # NOTE: dim C varies from different cases, can only be inited dynamically
|
||||
result = torch.zeros((N, tile.shape[1], height * 8 if is_decoder else height // 8, width * 8 if is_decoder else width // 8), device=device, requires_grad=False)
|
||||
result[:, :, out_bboxes[i][2]:out_bboxes[i][3], out_bboxes[i][0]:out_bboxes[i][1]] = crop_valid_region(tile, in_bboxes[i], out_bboxes[i], is_decoder)
|
||||
del tile
|
||||
elif i == num_tiles - 1 and forward:
|
||||
forward = False
|
||||
tiles[i] = tile
|
||||
elif i == 0 and not forward:
|
||||
forward = True
|
||||
tiles[i] = tile
|
||||
else:
|
||||
tiles[i] = tile.cpu()
|
||||
del tile
|
||||
|
||||
if interrupted: break
|
||||
if num_completed == num_tiles: break
|
||||
|
||||
# insert the group norm task to the head of each task queue
|
||||
group_norm_func = group_norm_param.summary()
|
||||
if group_norm_func is not None:
|
||||
for i in range(num_tiles):
|
||||
task_queue = task_queues[i]
|
||||
task_queue.insert(0, ('apply_norm', group_norm_func))
|
||||
|
||||
# Done!
|
||||
pbar.close()
|
||||
return result if result is not None else result_approx.to(device)
|
||||
Reference in New Issue
Block a user