Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0ab85b648 | ||
|
|
07d30c74ab | ||
|
|
24c06c19e6 | ||
|
|
76101ea540 | ||
|
|
da1b08ed22 | ||
|
|
b1d1639575 |
+28
-36
@@ -7,6 +7,7 @@ import random
|
|||||||
from ...SUPIR.utils.colorfix import wavelet_reconstruction, adaptive_instance_normalization
|
from ...SUPIR.utils.colorfix import wavelet_reconstruction, adaptive_instance_normalization
|
||||||
from pytorch_lightning import seed_everything
|
from pytorch_lightning import seed_everything
|
||||||
from ...SUPIR.utils.tilevae import VAEHook
|
from ...SUPIR.utils.tilevae import VAEHook
|
||||||
|
from ...SUPIR.util import convert_dtype
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
|
||||||
@@ -20,23 +21,8 @@ class SUPIRModel(DiffusionEngine):
|
|||||||
self.first_stage_model.denoise_encoder = copy.deepcopy(self.first_stage_model.encoder)
|
self.first_stage_model.denoise_encoder = copy.deepcopy(self.first_stage_model.encoder)
|
||||||
self.sampler_config = kwargs['sampler_config']
|
self.sampler_config = kwargs['sampler_config']
|
||||||
|
|
||||||
assert (ae_dtype in ['fp32', 'fp16', 'bf16']) and (diffusion_dtype in ['fp32', 'fp16', 'bf16'])
|
self.ae_dtype = convert_dtype(ae_dtype)
|
||||||
if ae_dtype == 'fp32':
|
self.model.dtype = convert_dtype(diffusion_dtype)
|
||||||
ae_dtype = torch.float32
|
|
||||||
elif ae_dtype == 'fp16':
|
|
||||||
raise RuntimeError('fp16 cause NaN in AE')
|
|
||||||
elif ae_dtype == 'bf16':
|
|
||||||
ae_dtype = torch.bfloat16
|
|
||||||
|
|
||||||
if diffusion_dtype == 'fp32':
|
|
||||||
diffusion_dtype = torch.float32
|
|
||||||
elif diffusion_dtype == 'fp16':
|
|
||||||
diffusion_dtype = torch.float16
|
|
||||||
elif diffusion_dtype == 'bf16':
|
|
||||||
diffusion_dtype = torch.bfloat16
|
|
||||||
|
|
||||||
self.ae_dtype = ae_dtype
|
|
||||||
self.model.dtype = diffusion_dtype
|
|
||||||
|
|
||||||
self.p_p = p_p
|
self.p_p = p_p
|
||||||
self.n_p = n_p
|
self.n_p = n_p
|
||||||
@@ -72,7 +58,6 @@ class SUPIRModel(DiffusionEngine):
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def decode_first_stage(self, z):
|
def decode_first_stage(self, z):
|
||||||
z = 1.0 / self.scale_factor * z
|
z = 1.0 / self.scale_factor * z
|
||||||
#with torch.autocast(device, dtype=self.ae_dtype):
|
|
||||||
autocast_condition = (self.ae_dtype == torch.float16 or self.ae_dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
autocast_condition = (self.ae_dtype == torch.float16 or self.ae_dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
||||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext():
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext():
|
||||||
out = self.first_stage_model.decode(z)
|
out = self.first_stage_model.decode(z)
|
||||||
@@ -120,29 +105,51 @@ class SUPIRModel(DiffusionEngine):
|
|||||||
self.sampler_config.params.s_noise = s_noise
|
self.sampler_config.params.s_noise = s_noise
|
||||||
self.sampler = instantiate_from_config(self.sampler_config)
|
self.sampler = instantiate_from_config(self.sampler_config)
|
||||||
|
|
||||||
|
print("Sampler: ", self.sampler_config.target)
|
||||||
print("sampler_config: ", self.sampler_config.params)
|
print("sampler_config: ", self.sampler_config.params)
|
||||||
|
|
||||||
if seed == -1:
|
if seed == -1:
|
||||||
seed = random.randint(0, 65535)
|
seed = random.randint(0, 65535)
|
||||||
seed_everything(seed)
|
seed_everything(seed)
|
||||||
|
|
||||||
|
|
||||||
|
self.model.to('cpu')
|
||||||
|
self.conditioner.to('cpu')
|
||||||
|
|
||||||
|
# stage 1: encode/decode/encode
|
||||||
|
self.first_stage_model.to(device)
|
||||||
_z = self.encode_first_stage_with_denoise(x, use_sample=False)
|
_z = self.encode_first_stage_with_denoise(x, use_sample=False)
|
||||||
|
|
||||||
x_stage1 = self.decode_first_stage(_z)
|
x_stage1 = self.decode_first_stage(_z)
|
||||||
|
|
||||||
z_stage1 = self.encode_first_stage(x_stage1)
|
z_stage1 = self.encode_first_stage(x_stage1)
|
||||||
|
self.first_stage_model.to('cpu')
|
||||||
|
|
||||||
|
#conditioning
|
||||||
|
self.conditioner.to(device)
|
||||||
c, uc = self.prepare_condition(_z, p, p_p, n_p, N)
|
c, uc = self.prepare_condition(_z, p, p_p, n_p, N)
|
||||||
|
self.conditioner.to('cpu')
|
||||||
|
|
||||||
denoiser = lambda input, sigma, c, control_scale: self.denoiser(
|
denoiser = lambda input, sigma, c, control_scale: self.denoiser(
|
||||||
self.model, input, sigma, c, control_scale, **kwargs
|
self.model, input, sigma, c, control_scale, **kwargs
|
||||||
)
|
)
|
||||||
|
|
||||||
noised_z = torch.randn_like(_z).to(_z.device)
|
noised_z = torch.randn_like(_z).to(_z.device)
|
||||||
|
|
||||||
|
comfy.model_management.soft_empty_cache()
|
||||||
|
|
||||||
|
#sampling
|
||||||
|
self.model.diffusion_model.to(device)
|
||||||
|
self.model.control_model.to(device)
|
||||||
|
self.denoiser.to(device)
|
||||||
|
|
||||||
_samples = self.sampler(denoiser, noised_z, cond=c, uc=uc, x_center=z_stage1, control_scale=control_scale,
|
_samples = self.sampler(denoiser, noised_z, cond=c, uc=uc, x_center=z_stage1, control_scale=control_scale,
|
||||||
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
|
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
|
||||||
|
self.model.diffusion_model.to('cpu')
|
||||||
|
self.model.control_model.to('cpu')
|
||||||
|
|
||||||
|
#decoding
|
||||||
|
self.first_stage_model.to(device)
|
||||||
samples = self.decode_first_stage(_samples)
|
samples = self.decode_first_stage(_samples)
|
||||||
|
self.first_stage_model.to('cpu')
|
||||||
|
|
||||||
if color_fix_type == 'Wavelet':
|
if color_fix_type == 'Wavelet':
|
||||||
samples = wavelet_reconstruction(samples, x_stage1)
|
samples = wavelet_reconstruction(samples, x_stage1)
|
||||||
elif color_fix_type == 'AdaIn':
|
elif color_fix_type == 'AdaIn':
|
||||||
@@ -194,18 +201,3 @@ class SUPIRModel(DiffusionEngine):
|
|||||||
_c, _ = self.conditioner.get_unconditional_conditioning(batch, None)
|
_c, _ = self.conditioner.get_unconditional_conditioning(batch, None)
|
||||||
c.append(_c)
|
c.append(_c)
|
||||||
return c, uc
|
return c, uc
|
||||||
|
|
||||||
# if __name__ == '__main__':
|
|
||||||
# from SUPIR.util import create_model, load_state_dict
|
|
||||||
|
|
||||||
# model = create_model('../../options/dev/SUPIR_paper_version.yaml')
|
|
||||||
|
|
||||||
# SDXL_CKPT = '/opt/data/private/AIGC_pretrain/SDXL_cache/sd_xl_base_1.0_0.9vae.safetensors'
|
|
||||||
# SUPIR_CKPT = '/opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-paper.ckpt'
|
|
||||||
# model.load_state_dict(load_state_dict(SDXL_CKPT), strict=False)
|
|
||||||
# model.load_state_dict(load_state_dict(SUPIR_CKPT), strict=False)
|
|
||||||
# model = model.cuda()
|
|
||||||
|
|
||||||
# x = torch.randn(1, 3, 512, 512).cuda()
|
|
||||||
# p = ['a professional, detailed, high-quality photo']
|
|
||||||
# samples = model.batchify_sample(x, p, num_steps=50, restoration_scale=4.0, s_churn=0, cfg_scale=4.0, seed=-1, num_samples=1)
|
|
||||||
|
|||||||
@@ -0,0 +1,196 @@
|
|||||||
|
import torch
|
||||||
|
from ...sgm.models.diffusion import DiffusionEngine
|
||||||
|
from ...sgm.util import instantiate_from_config
|
||||||
|
import copy
|
||||||
|
from ...sgm.modules.distributions.distributions import DiagonalGaussianDistribution
|
||||||
|
import random
|
||||||
|
from ...SUPIR.utils.colorfix import wavelet_reconstruction, adaptive_instance_normalization
|
||||||
|
from pytorch_lightning import seed_everything
|
||||||
|
from ...SUPIR.utils.tilevae import VAEHook
|
||||||
|
from contextlib import nullcontext
|
||||||
|
import comfy.model_management
|
||||||
|
|
||||||
|
device = comfy.model_management.get_torch_device()
|
||||||
|
|
||||||
|
class SUPIRModel(DiffusionEngine):
|
||||||
|
def __init__(self, control_stage_config, ae_dtype='fp32', diffusion_dtype='fp32', p_p='', n_p='', *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
control_model = instantiate_from_config(control_stage_config)
|
||||||
|
self.model.load_control_model(control_model)
|
||||||
|
self.first_stage_model.denoise_encoder = copy.deepcopy(self.first_stage_model.encoder)
|
||||||
|
self.sampler_config = kwargs['sampler_config']
|
||||||
|
|
||||||
|
assert (ae_dtype in ['fp32', 'fp16', 'bf16']) and (diffusion_dtype in ['fp32', 'fp16', 'bf16'])
|
||||||
|
if ae_dtype == 'fp32':
|
||||||
|
ae_dtype = torch.float32
|
||||||
|
elif ae_dtype == 'fp16':
|
||||||
|
raise RuntimeError('fp16 cause NaN in AE')
|
||||||
|
elif ae_dtype == 'bf16':
|
||||||
|
ae_dtype = torch.bfloat16
|
||||||
|
|
||||||
|
if diffusion_dtype == 'fp32':
|
||||||
|
diffusion_dtype = torch.float32
|
||||||
|
elif diffusion_dtype == 'fp16':
|
||||||
|
diffusion_dtype = torch.float16
|
||||||
|
elif diffusion_dtype == 'bf16':
|
||||||
|
diffusion_dtype = torch.bfloat16
|
||||||
|
|
||||||
|
self.ae_dtype = ae_dtype
|
||||||
|
self.model.dtype = diffusion_dtype
|
||||||
|
|
||||||
|
self.p_p = p_p
|
||||||
|
self.n_p = n_p
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def encode_first_stage(self, x):
|
||||||
|
#with torch.autocast(device, dtype=self.ae_dtype):
|
||||||
|
autocast_condition = (self.ae_dtype == torch.float16 or self.ae_dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext():
|
||||||
|
z = self.first_stage_model.encode(x)
|
||||||
|
z = self.scale_factor * z
|
||||||
|
return z
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def encode_first_stage_with_denoise(self, x, use_sample=True, is_stage1=False):
|
||||||
|
#with torch.autocast(device, dtype=self.ae_dtype):
|
||||||
|
self.first_stage_model.to(self.ae_dtype)
|
||||||
|
autocast_condition = (self.model.dtype == torch.float16 or self.model.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext():
|
||||||
|
if is_stage1:
|
||||||
|
h = self.first_stage_model.denoise_encoder_s1(x)
|
||||||
|
else:
|
||||||
|
h = self.first_stage_model.denoise_encoder(x)
|
||||||
|
moments = self.first_stage_model.quant_conv(h)
|
||||||
|
posterior = DiagonalGaussianDistribution(moments)
|
||||||
|
if use_sample:
|
||||||
|
z = posterior.sample()
|
||||||
|
else:
|
||||||
|
z = posterior.mode()
|
||||||
|
z = self.scale_factor * z
|
||||||
|
return z
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def decode_first_stage(self, z):
|
||||||
|
z = 1.0 / self.scale_factor * z
|
||||||
|
#with torch.autocast(device, dtype=self.ae_dtype):
|
||||||
|
autocast_condition = (self.ae_dtype == torch.float16 or self.ae_dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext():
|
||||||
|
out = self.first_stage_model.decode(z)
|
||||||
|
return out.float()
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def batchify_denoise(self, x, is_stage1=False):
|
||||||
|
'''
|
||||||
|
[N, C, H, W], [-1, 1], RGB
|
||||||
|
'''
|
||||||
|
x = self.encode_first_stage_with_denoise(x, use_sample=False, is_stage1=is_stage1)
|
||||||
|
return self.decode_first_stage(x)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def batchify_sample(self, x, p, p_p='default', n_p='default', num_steps=100, restoration_scale=4.0, s_churn=0, s_noise=1.003, cfg_scale=4.0, seed=-1,
|
||||||
|
num_samples=1, control_scale=1, color_fix_type='None', use_linear_CFG=False, use_linear_control_scale=False,
|
||||||
|
cfg_scale_start=1.0, control_scale_start=0.0, **kwargs):
|
||||||
|
'''
|
||||||
|
[N, C], [-1, 1], RGB
|
||||||
|
'''
|
||||||
|
assert len(x) == len(p)
|
||||||
|
assert color_fix_type in ['Wavelet', 'AdaIn', 'None']
|
||||||
|
|
||||||
|
N = len(x)
|
||||||
|
if num_samples > 1:
|
||||||
|
assert N == 1
|
||||||
|
N = num_samples
|
||||||
|
x = x.repeat(N, 1, 1, 1)
|
||||||
|
p = p * N
|
||||||
|
|
||||||
|
if p_p == 'default':
|
||||||
|
p_p = self.p_p
|
||||||
|
if n_p == 'default':
|
||||||
|
n_p = self.n_p
|
||||||
|
|
||||||
|
self.sampler_config.params.num_steps = num_steps
|
||||||
|
if use_linear_CFG:
|
||||||
|
self.sampler_config.params.guider_config.params.scale_min = cfg_scale
|
||||||
|
self.sampler_config.params.guider_config.params.scale = cfg_scale_start
|
||||||
|
else:
|
||||||
|
self.sampler_config.params.guider_config.params.scale_min = cfg_scale
|
||||||
|
self.sampler_config.params.guider_config.params.scale = cfg_scale
|
||||||
|
self.sampler_config.params.restore_cfg = restoration_scale
|
||||||
|
self.sampler_config.params.s_churn = s_churn
|
||||||
|
self.sampler_config.params.s_noise = s_noise
|
||||||
|
self.sampler = instantiate_from_config(self.sampler_config)
|
||||||
|
|
||||||
|
print("sampler_config: ", self.sampler_config.params)
|
||||||
|
|
||||||
|
if seed == -1:
|
||||||
|
seed = random.randint(0, 65535)
|
||||||
|
seed_everything(seed)
|
||||||
|
|
||||||
|
_z = self.encode_first_stage_with_denoise(x, use_sample=False)
|
||||||
|
|
||||||
|
x_stage1 = self.decode_first_stage(_z)
|
||||||
|
|
||||||
|
z_stage1 = self.encode_first_stage(x_stage1)
|
||||||
|
|
||||||
|
c, uc = self.prepare_condition(_z, p, p_p, n_p, N)
|
||||||
|
|
||||||
|
denoiser = lambda input, sigma, c, control_scale: self.denoiser(
|
||||||
|
self.model, input, sigma, c, control_scale, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
noised_z = torch.randn_like(_z).to(_z.device)
|
||||||
|
|
||||||
|
_samples = self.sampler(denoiser, noised_z, cond=c, uc=uc, x_center=z_stage1, control_scale=control_scale,
|
||||||
|
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
|
||||||
|
samples = self.decode_first_stage(_samples)
|
||||||
|
if color_fix_type == 'Wavelet':
|
||||||
|
samples = wavelet_reconstruction(samples, x_stage1)
|
||||||
|
elif color_fix_type == 'AdaIn':
|
||||||
|
samples = adaptive_instance_normalization(samples, x_stage1)
|
||||||
|
return samples
|
||||||
|
|
||||||
|
def init_tile_vae(self, encoder_tile_size=512, decoder_tile_size=64):
|
||||||
|
self.first_stage_model.denoise_encoder.original_forward = self.first_stage_model.denoise_encoder.forward
|
||||||
|
self.first_stage_model.encoder.original_forward = self.first_stage_model.encoder.forward
|
||||||
|
self.first_stage_model.decoder.original_forward = self.first_stage_model.decoder.forward
|
||||||
|
self.first_stage_model.denoise_encoder.forward = VAEHook(
|
||||||
|
self.first_stage_model.denoise_encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
self.first_stage_model.encoder.forward = VAEHook(
|
||||||
|
self.first_stage_model.encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
self.first_stage_model.decoder.forward = VAEHook(
|
||||||
|
self.first_stage_model.decoder, decoder_tile_size, is_decoder=True, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
|
||||||
|
def prepare_condition(self, _z, p, p_p, n_p, N):
|
||||||
|
batch = {}
|
||||||
|
batch['original_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(_z.device)
|
||||||
|
batch['crop_coords_top_left'] = torch.tensor([0, 0]).repeat(N, 1).to(_z.device)
|
||||||
|
batch['target_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(_z.device)
|
||||||
|
batch['aesthetic_score'] = torch.tensor([9.0]).repeat(N, 1).to(_z.device)
|
||||||
|
batch['control'] = _z
|
||||||
|
|
||||||
|
batch_uc = copy.deepcopy(batch)
|
||||||
|
batch_uc['txt'] = [n_p for _ in p]
|
||||||
|
autocast_condition = (self.model.dtype == torch.float16 or self.model.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
|
||||||
|
if not isinstance(p[0], list):
|
||||||
|
print("Using local prompt: ")
|
||||||
|
batch['txt'] = [''.join([_p, p_p]) for _p in p]
|
||||||
|
print(batch['txt'])
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.model.dtype) if autocast_condition else nullcontext():
|
||||||
|
c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
|
||||||
|
else:
|
||||||
|
print("Using tile prompts")
|
||||||
|
assert len(p) == 1, 'Support bs=1 only for local prompt conditioning.'
|
||||||
|
p_tiles = p[0]
|
||||||
|
c = []
|
||||||
|
for i, p_tile in enumerate(p_tiles):
|
||||||
|
batch['txt'] = [''.join([p_tile, p_p])]
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.model.dtype) if autocast_condition else nullcontext():
|
||||||
|
if i == 0:
|
||||||
|
_c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
|
||||||
|
else:
|
||||||
|
_c, _ = self.conditioner.get_unconditional_conditioning(batch, None)
|
||||||
|
c.append(_c)
|
||||||
|
return c, uc
|
||||||
+22
-1
@@ -1,3 +1,24 @@
|
|||||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
from .nodes import SUPIR_Upscale
|
||||||
|
from .nodes_v2 import SUPIR_sample, SUPIR_model_loader, SUPIR_first_stage, SUPIR_encode, SUPIR_decode, SUPIR_conditioner, SUPIR_tiles
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"SUPIR_Upscale": SUPIR_Upscale,
|
||||||
|
"SUPIR_sample": SUPIR_sample,
|
||||||
|
"SUPIR_model_loader": SUPIR_model_loader,
|
||||||
|
"SUPIR_first_stage": SUPIR_first_stage,
|
||||||
|
"SUPIR_encode": SUPIR_encode,
|
||||||
|
"SUPIR_decode": SUPIR_decode,
|
||||||
|
"SUPIR_conditioner": SUPIR_conditioner,
|
||||||
|
"SUPIR_tiles": SUPIR_tiles
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"SUPIR_Upscale": "SUPIR Upscale",
|
||||||
|
"SUPIR_sample": "SUPIR Sampler",
|
||||||
|
"SUPIR_model_loader": "SUPIR Model Loader",
|
||||||
|
"SUPIR_first_stage": "SUPIR First Stage (Denoiser)",
|
||||||
|
"SUPIR_encode": "SUPIR Encode",
|
||||||
|
"SUPIR_decode": "SUPIR Decode",
|
||||||
|
"SUPIR_conditioner": "SUPIR Conditioner",
|
||||||
|
"SUPIR_tiles": "SUPIR Tiles"
|
||||||
|
}
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from contextlib import nullcontext
|
|
||||||
from omegaconf import OmegaConf
|
from omegaconf import OmegaConf
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
@@ -13,6 +11,7 @@ import torch.cuda
|
|||||||
from .sgm.util import instantiate_from_config
|
from .sgm.util import instantiate_from_config
|
||||||
from .SUPIR.util import convert_dtype, load_state_dict
|
from .SUPIR.util import convert_dtype, load_state_dict
|
||||||
import open_clip
|
import open_clip
|
||||||
|
from contextlib import contextmanager
|
||||||
|
|
||||||
from transformers import (
|
from transformers import (
|
||||||
CLIPTextModel,
|
CLIPTextModel,
|
||||||
@@ -34,8 +33,17 @@ except:
|
|||||||
def dummy_build_vision_tower(*args, **kwargs):
|
def dummy_build_vision_tower(*args, **kwargs):
|
||||||
# Monkey patch the CLIP class before you create an instance.
|
# Monkey patch the CLIP class before you create an instance.
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def patch_build_vision_tower():
|
||||||
|
original_build_vision_tower = open_clip.model._build_vision_tower
|
||||||
open_clip.model._build_vision_tower = dummy_build_vision_tower
|
open_clip.model._build_vision_tower = dummy_build_vision_tower
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
open_clip.model._build_vision_tower = original_build_vision_tower
|
||||||
|
|
||||||
def build_text_model_from_openai_state_dict(
|
def build_text_model_from_openai_state_dict(
|
||||||
state_dict: dict,
|
state_dict: dict,
|
||||||
cast_dtype=torch.float16,
|
cast_dtype=torch.float16,
|
||||||
@@ -56,11 +64,13 @@ def build_text_model_from_openai_state_dict(
|
|||||||
heads=transformer_heads,
|
heads=transformer_heads,
|
||||||
layers=transformer_layers,
|
layers=transformer_layers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
with patch_build_vision_tower():
|
||||||
model = open_clip.CLIP(
|
model = open_clip.CLIP(
|
||||||
embed_dim,
|
embed_dim,
|
||||||
vision_cfg=vision_cfg,
|
vision_cfg=vision_cfg,
|
||||||
text_cfg=text_cfg,
|
text_cfg=text_cfg,
|
||||||
quick_gelu=True, # OpenAI models were trained with QuickGELU
|
quick_gelu=True,
|
||||||
cast_dtype=cast_dtype,
|
cast_dtype=cast_dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -128,6 +138,15 @@ class SUPIR_Upscale:
|
|||||||
"use_tiled_sampling": ("BOOLEAN", {"default": False}),
|
"use_tiled_sampling": ("BOOLEAN", {"default": False}),
|
||||||
"sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 32}),
|
"sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 32}),
|
||||||
"sampler_tile_stride": ("INT", {"default": 512, "min": 32, "max": 2048, "step": 32}),
|
"sampler_tile_stride": ("INT", {"default": 512, "min": 32, "max": 2048, "step": 32}),
|
||||||
|
"fp8_unet": ("BOOLEAN", {"default": False}),
|
||||||
|
"fp8_vae": ("BOOLEAN", {"default": False}),
|
||||||
|
"sampler": (
|
||||||
|
[
|
||||||
|
'RestoreDPMPP2MSampler',
|
||||||
|
'RestoreEDMSampler',
|
||||||
|
], {
|
||||||
|
"default": 'RestoreEDMSampler'
|
||||||
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -141,7 +160,7 @@ class SUPIR_Upscale:
|
|||||||
encoder_tile_size_pixels, decoder_tile_size_latent,
|
encoder_tile_size_pixels, decoder_tile_size_latent,
|
||||||
control_scale, cfg_scale_start, control_scale_start, restoration_scale, keep_model_loaded,
|
control_scale, cfg_scale_start, control_scale_start, restoration_scale, keep_model_loaded,
|
||||||
a_prompt, n_prompt, sdxl_model, supir_model, use_tiled_vae, use_tiled_sampling=False, sampler_tile_size=128, sampler_tile_stride=64, captions="", diffusion_dtype="auto",
|
a_prompt, n_prompt, sdxl_model, supir_model, use_tiled_vae, use_tiled_sampling=False, sampler_tile_size=128, sampler_tile_stride=64, captions="", diffusion_dtype="auto",
|
||||||
encoder_dtype="auto", batch_size=1):
|
encoder_dtype="auto", batch_size=1, fp8_unet=False, fp8_vae=False, sampler="RestoreEDMSampler"):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
mm.unload_all_models()
|
mm.unload_all_models()
|
||||||
|
|
||||||
@@ -160,6 +179,9 @@ class SUPIR_Upscale:
|
|||||||
'use_tiled_vae': use_tiled_vae,
|
'use_tiled_vae': use_tiled_vae,
|
||||||
'supir_model': supir_model,
|
'supir_model': supir_model,
|
||||||
'use_tiled_sampling': use_tiled_sampling,
|
'use_tiled_sampling': use_tiled_sampling,
|
||||||
|
'fp8_unet': fp8_unet,
|
||||||
|
'fp8_vae': fp8_vae,
|
||||||
|
'sampler': sampler
|
||||||
}
|
}
|
||||||
|
|
||||||
if diffusion_dtype == 'auto':
|
if diffusion_dtype == 'auto':
|
||||||
@@ -207,9 +229,11 @@ class SUPIR_Upscale:
|
|||||||
config = OmegaConf.load(config_path_tiled)
|
config = OmegaConf.load(config_path_tiled)
|
||||||
config.model.params.sampler_config.params.tile_size = sampler_tile_size // 8
|
config.model.params.sampler_config.params.tile_size = sampler_tile_size // 8
|
||||||
config.model.params.sampler_config.params.tile_stride = sampler_tile_stride // 8
|
config.model.params.sampler_config.params.tile_stride = sampler_tile_stride // 8
|
||||||
|
config.model.params.sampler_config.target = f".sgm.modules.diffusionmodules.sampling.Tiled{sampler}"
|
||||||
print("Using tiled sampling")
|
print("Using tiled sampling")
|
||||||
else:
|
else:
|
||||||
config = OmegaConf.load(config_path)
|
config = OmegaConf.load(config_path)
|
||||||
|
config.model.params.sampler_config.target = f".sgm.modules.diffusionmodules.sampling.{sampler}"
|
||||||
print("Using non-tiled sampling")
|
print("Using non-tiled sampling")
|
||||||
|
|
||||||
if XFORMERS_IS_AVAILABLE:
|
if XFORMERS_IS_AVAILABLE:
|
||||||
@@ -267,35 +291,30 @@ class SUPIR_Upscale:
|
|||||||
clip_g = build_text_model_from_openai_state_dict(sd, cast_dtype=dtype)
|
clip_g = build_text_model_from_openai_state_dict(sd, cast_dtype=dtype)
|
||||||
self.model.conditioner.embedders[1].model = clip_g
|
self.model.conditioner.embedders[1].model = clip_g
|
||||||
except:
|
except:
|
||||||
|
|
||||||
raise Exception("Failed to load second clip model from SDXL checkpoint")
|
raise Exception("Failed to load second clip model from SDXL checkpoint")
|
||||||
|
|
||||||
del sd, clip_g
|
del sd, clip_g
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
try:
|
|
||||||
self.model.to(dtype)
|
self.model.to(dtype)
|
||||||
self.model.to(device)
|
|
||||||
except Exception as e:
|
#only unets and/or vae to fp8
|
||||||
print("Failed to move model to device")
|
if fp8_unet:
|
||||||
print(e)
|
self.model.model.to(torch.float8_e4m3fn)
|
||||||
import gc
|
if fp8_vae:
|
||||||
# unload everything and give up
|
self.model.first_stage_model.to(torch.float8_e4m3fn)
|
||||||
self.model = None
|
|
||||||
del self.model
|
|
||||||
gc.collect()
|
|
||||||
mm.soft_empty_cache()
|
|
||||||
|
|
||||||
if use_tiled_vae:
|
if use_tiled_vae:
|
||||||
self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent)
|
self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent)
|
||||||
|
|
||||||
image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
upscaled_image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
||||||
B, H, W, C = image.shape
|
B, H, W, C = upscaled_image.shape
|
||||||
new_height = H // 64 * 64
|
new_height = H if H % 64 == 0 else ((H // 64) + 1) * 64
|
||||||
new_width = W // 64 * 64
|
new_width = W if W % 64 == 0 else ((W // 64) + 1) * 64
|
||||||
image = image.permute(0, 3, 1, 2).contiguous()
|
upscaled_image = upscaled_image.permute(0, 3, 1, 2)
|
||||||
resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False)
|
resized_image = F.interpolate(upscaled_image, size=(new_height, new_width), mode='bicubic', align_corners=False)
|
||||||
resized_image = resized_image.to(device)
|
resized_image = resized_image.to(device)
|
||||||
|
|
||||||
captions_list = []
|
captions_list = []
|
||||||
captions_list.append(captions)
|
captions_list.append(captions)
|
||||||
print("captions: ", captions_list)
|
print("captions: ", captions_list)
|
||||||
@@ -329,7 +348,8 @@ class SUPIR_Upscale:
|
|||||||
self.model = None
|
self.model = None
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
print("It's likely that too large of an image or batch_size for SUPIR was used,"
|
print("It's likely that too large of an image or batch_size for SUPIR was used,"
|
||||||
" and it has devoured all of the memory it had reserved, you may need to restart ComfyUI")
|
" and it has devoured all of the memory it had reserved, you may need to restart ComfyUI. Make sure you are using tiled_vae, "
|
||||||
|
" you can also try using fp8 for reduced memory usage if your system supports it.")
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
out.append(samples.squeeze(0).cpu())
|
out.append(samples.squeeze(0).cpu())
|
||||||
@@ -345,7 +365,7 @@ class SUPIR_Upscale:
|
|||||||
else:
|
else:
|
||||||
out_stacked = torch.stack(out, dim=0).cpu().to(torch.float32).permute(0, 2, 3, 1)
|
out_stacked = torch.stack(out, dim=0).cpu().to(torch.float32).permute(0, 2, 3, 1)
|
||||||
|
|
||||||
final_image, = ImageScale.upscale(self, out_stacked, "lanczos", W, H, crop="disabled")
|
final_image, = ImageScale.upscale(self, out_stacked, resize_method, W, H, crop="disabled")
|
||||||
|
|
||||||
return (final_image,)
|
return (final_image,)
|
||||||
|
|
||||||
|
|||||||
+708
@@ -0,0 +1,708 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
import comfy.utils
|
||||||
|
import comfy.model_management as mm
|
||||||
|
import folder_paths
|
||||||
|
from nodes import ImageScaleBy
|
||||||
|
from nodes import ImageScale
|
||||||
|
import torch.cuda
|
||||||
|
from .sgm.util import instantiate_from_config
|
||||||
|
from .SUPIR.util import convert_dtype, load_state_dict
|
||||||
|
from .sgm.modules.distributions.distributions import DiagonalGaussianDistribution
|
||||||
|
import open_clip
|
||||||
|
from contextlib import contextmanager, nullcontext
|
||||||
|
|
||||||
|
from transformers import (
|
||||||
|
CLIPTextModel,
|
||||||
|
CLIPTokenizer,
|
||||||
|
CLIPTextConfig,
|
||||||
|
|
||||||
|
)
|
||||||
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
try:
|
||||||
|
import xformers
|
||||||
|
import xformers.ops
|
||||||
|
|
||||||
|
XFORMERS_IS_AVAILABLE = True
|
||||||
|
except:
|
||||||
|
XFORMERS_IS_AVAILABLE = False
|
||||||
|
|
||||||
|
|
||||||
|
def dummy_build_vision_tower(*args, **kwargs):
|
||||||
|
# Monkey patch the CLIP class before you create an instance.
|
||||||
|
return None
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def patch_build_vision_tower():
|
||||||
|
original_build_vision_tower = open_clip.model._build_vision_tower
|
||||||
|
open_clip.model._build_vision_tower = dummy_build_vision_tower
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
open_clip.model._build_vision_tower = original_build_vision_tower
|
||||||
|
|
||||||
|
def build_text_model_from_openai_state_dict(
|
||||||
|
state_dict: dict,
|
||||||
|
cast_dtype=torch.float16,
|
||||||
|
):
|
||||||
|
|
||||||
|
embed_dim = state_dict["text_projection"].shape[1]
|
||||||
|
context_length = state_dict["positional_embedding"].shape[0]
|
||||||
|
vocab_size = state_dict["token_embedding.weight"].shape[0]
|
||||||
|
transformer_width = state_dict["ln_final.weight"].shape[0]
|
||||||
|
transformer_heads = transformer_width // 64
|
||||||
|
transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith(f"transformer.resblocks")))
|
||||||
|
|
||||||
|
vision_cfg = None
|
||||||
|
text_cfg = open_clip.CLIPTextCfg(
|
||||||
|
context_length=context_length,
|
||||||
|
vocab_size=vocab_size,
|
||||||
|
width=transformer_width,
|
||||||
|
heads=transformer_heads,
|
||||||
|
layers=transformer_layers,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch_build_vision_tower():
|
||||||
|
model = open_clip.CLIP(
|
||||||
|
embed_dim,
|
||||||
|
vision_cfg=vision_cfg,
|
||||||
|
text_cfg=text_cfg,
|
||||||
|
quick_gelu=True,
|
||||||
|
cast_dtype=cast_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
model.load_state_dict(state_dict, strict=False)
|
||||||
|
model = model.eval()
|
||||||
|
for param in model.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
return model
|
||||||
|
|
||||||
|
class SUPIR_encode:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"SUPIR_VAE": ("SUPIRVAE",),
|
||||||
|
"image": ("IMAGE",),
|
||||||
|
"use_tiled_vae": ("BOOLEAN", {"default": True}),
|
||||||
|
"encoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
"encoder_dtype": (
|
||||||
|
[
|
||||||
|
'bf16',
|
||||||
|
'fp32',
|
||||||
|
'auto'
|
||||||
|
], {
|
||||||
|
"default": 'auto'
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LATENT",)
|
||||||
|
RETURN_NAMES = ("latent",)
|
||||||
|
FUNCTION = "encode"
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def encode(self, SUPIR_VAE, image, encoder_dtype, use_tiled_vae, encoder_tile_size):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
if encoder_dtype == 'auto':
|
||||||
|
try:
|
||||||
|
if mm.should_use_bf16():
|
||||||
|
print("Encoder using bf16")
|
||||||
|
vae_dtype = 'bf16'
|
||||||
|
else:
|
||||||
|
print("Encoder using using fp32")
|
||||||
|
vae_dtype = 'fp32'
|
||||||
|
except:
|
||||||
|
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
|
||||||
|
else:
|
||||||
|
vae_dtype = encoder_dtype
|
||||||
|
print(f"Encoder using using {vae_dtype}")
|
||||||
|
|
||||||
|
dtype = convert_dtype(vae_dtype)
|
||||||
|
|
||||||
|
B, H, W, C = image.shape
|
||||||
|
new_height = H // 64 * 64
|
||||||
|
new_width = W // 64 * 64
|
||||||
|
resized_image, = ImageScale.upscale(self, image, 'lanczos', new_width, new_height, crop="disabled")
|
||||||
|
resized_image = image.permute(0, 3, 1, 2).to(device)
|
||||||
|
|
||||||
|
if use_tiled_vae:
|
||||||
|
from .SUPIR.utils.tilevae import VAEHook
|
||||||
|
# Store the `original_forward` only if it hasn't been stored already
|
||||||
|
if not hasattr(SUPIR_VAE.encoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.encoder.original_forward = SUPIR_VAE.encoder.forward
|
||||||
|
SUPIR_VAE.encoder.forward = VAEHook(
|
||||||
|
SUPIR_VAE.encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
else:
|
||||||
|
# Only assign `original_forward` back if it exists
|
||||||
|
if hasattr(SUPIR_VAE.encoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.encoder.forward = SUPIR_VAE.encoder.original_forward
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(B)
|
||||||
|
out = []
|
||||||
|
for img in resized_image:
|
||||||
|
|
||||||
|
SUPIR_VAE.to(dtype).to(device)
|
||||||
|
|
||||||
|
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
|
|
||||||
|
z = SUPIR_VAE.encode(img.unsqueeze(0))
|
||||||
|
z = z * 0.13025
|
||||||
|
out.append(z)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
if len(out[0].shape) == 4:
|
||||||
|
out_stacked = torch.cat(out, dim=0)
|
||||||
|
else:
|
||||||
|
out_stacked = torch.stack(out, dim=0)
|
||||||
|
return (out_stacked,)
|
||||||
|
|
||||||
|
class SUPIR_decode:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"SUPIR_VAE": ("SUPIRVAE",),
|
||||||
|
"latents": ("LATENT",),
|
||||||
|
"use_tiled_vae": ("BOOLEAN", {"default": True}),
|
||||||
|
"decoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES = ("image",)
|
||||||
|
FUNCTION = "decode"
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def decode(self, SUPIR_VAE, latents, use_tiled_vae, decoder_tile_size):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
|
||||||
|
dtype = latents.dtype
|
||||||
|
|
||||||
|
B, H, W, C = latents.shape
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(B)
|
||||||
|
|
||||||
|
SUPIR_VAE.to(dtype).to(device)
|
||||||
|
|
||||||
|
if use_tiled_vae:
|
||||||
|
from .SUPIR.utils.tilevae import VAEHook
|
||||||
|
# Store the `original_forward` only if it hasn't been stored already
|
||||||
|
if not hasattr(SUPIR_VAE.decoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.decoder.original_forward = SUPIR_VAE.decoder.forward
|
||||||
|
SUPIR_VAE.decoder.forward = VAEHook(
|
||||||
|
SUPIR_VAE.decoder, decoder_tile_size // 8, is_decoder=True, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
else:
|
||||||
|
# Only assign `original_forward` back if it exists
|
||||||
|
if hasattr(SUPIR_VAE.decoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.decoder.forward = SUPIR_VAE.decoder.original_forward
|
||||||
|
|
||||||
|
out = []
|
||||||
|
for latent in latents:
|
||||||
|
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
|
latent = 1.0 / 0.13025 * latent
|
||||||
|
decoded_image = SUPIR_VAE.decode(latent.unsqueeze(0)).float()
|
||||||
|
out.append(decoded_image)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
out_stacked = torch.cat(out, dim=0).cpu().to(torch.float32).permute(0, 2, 3, 1)
|
||||||
|
|
||||||
|
return (out_stacked,)
|
||||||
|
|
||||||
|
class SUPIR_first_stage:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"SUPIR_VAE": ("SUPIRVAE",),
|
||||||
|
"image": ("IMAGE",),
|
||||||
|
"use_tiled_vae": ("BOOLEAN", {"default": True}),
|
||||||
|
"encoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
"decoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
"encoder_dtype": (
|
||||||
|
[
|
||||||
|
'bf16',
|
||||||
|
'fp32',
|
||||||
|
'auto'
|
||||||
|
], {
|
||||||
|
"default": 'auto'
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("SUPIRVAE", "IMAGE", "LATENT",)
|
||||||
|
RETURN_NAMES = ("SUPIR_VAE", "denoised_image", "denoised_latents",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def process(self, SUPIR_VAE, image, encoder_dtype, use_tiled_vae, encoder_tile_size, decoder_tile_size):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
if encoder_dtype == 'auto':
|
||||||
|
try:
|
||||||
|
if mm.should_use_bf16():
|
||||||
|
print("Encoder using bf16")
|
||||||
|
vae_dtype = 'bf16'
|
||||||
|
else:
|
||||||
|
print("Encoder using using fp32")
|
||||||
|
vae_dtype = 'fp32'
|
||||||
|
except:
|
||||||
|
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
|
||||||
|
else:
|
||||||
|
vae_dtype = encoder_dtype
|
||||||
|
print(f"Encoder using using {vae_dtype}")
|
||||||
|
|
||||||
|
dtype = convert_dtype(vae_dtype)
|
||||||
|
|
||||||
|
if use_tiled_vae:
|
||||||
|
from .SUPIR.utils.tilevae import VAEHook
|
||||||
|
# Store the `original_forward` only if it hasn't been stored already
|
||||||
|
if not hasattr(SUPIR_VAE.encoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.denoise_encoder.original_forward = SUPIR_VAE.denoise_encoder.forward
|
||||||
|
SUPIR_VAE.decoder.original_forward = SUPIR_VAE.decoder.forward
|
||||||
|
|
||||||
|
SUPIR_VAE.denoise_encoder.forward = VAEHook(
|
||||||
|
SUPIR_VAE.denoise_encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
|
||||||
|
SUPIR_VAE.decoder.forward = VAEHook(
|
||||||
|
SUPIR_VAE.decoder, decoder_tile_size // 8, is_decoder=True, fast_decoder=False,
|
||||||
|
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||||
|
else:
|
||||||
|
# Only assign `original_forward` back if it exists
|
||||||
|
if hasattr(SUPIR_VAE.decoder, 'original_forward'):
|
||||||
|
SUPIR_VAE.encoder.forward = SUPIR_VAE.encoder.original_forward
|
||||||
|
SUPIR_VAE.decoder.forward = SUPIR_VAE.decoder.original_forward
|
||||||
|
|
||||||
|
B, H, W, C = image.shape
|
||||||
|
new_height = H // 64 * 64
|
||||||
|
new_width = W // 64 * 64
|
||||||
|
resized_image, = ImageScale.upscale(self, image, 'lanczos', new_width, new_height, crop="disabled")
|
||||||
|
resized_image = image.permute(0, 3, 1, 2).to(device)
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(B)
|
||||||
|
out = []
|
||||||
|
out_samples = []
|
||||||
|
for img in resized_image:
|
||||||
|
|
||||||
|
SUPIR_VAE.to(dtype).to(device)
|
||||||
|
|
||||||
|
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
|
|
||||||
|
h = SUPIR_VAE.denoise_encoder(img.unsqueeze(0))
|
||||||
|
moments = SUPIR_VAE.quant_conv(h)
|
||||||
|
posterior = DiagonalGaussianDistribution(moments)
|
||||||
|
sample = posterior.sample()
|
||||||
|
decoded_images = SUPIR_VAE.decode(sample).float()
|
||||||
|
|
||||||
|
out.append(decoded_images.cpu())
|
||||||
|
out_samples.append(sample.cpu() * 0.13025)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
|
||||||
|
out_stacked = torch.cat(out, dim=0).to(torch.float32).permute(0, 2, 3, 1)
|
||||||
|
out_samples_stacked = torch.cat(out_samples, dim=0)
|
||||||
|
|
||||||
|
final_image, = ImageScale.upscale(self, out_stacked, 'lanczos', W, H, crop="disabled")
|
||||||
|
|
||||||
|
return (SUPIR_VAE, final_image, out_samples_stacked,)
|
||||||
|
|
||||||
|
class SUPIR_sample:
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"SUPIR_model": ("SUPIRMODEL",),
|
||||||
|
"latents": ("LATENT",),
|
||||||
|
"positive": ("SUPIR_cond_pos",),
|
||||||
|
"negative": ("SUPIR_cond_neg",),
|
||||||
|
"seed": ("INT", {"default": 123, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||||
|
"steps": ("INT", {"default": 45, "min": 3, "max": 4096, "step": 1}),
|
||||||
|
"cfg_scale_start": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 9.0, "step": 0.05}),
|
||||||
|
"cfg_scale_end": ("FLOAT", {"default": 4.0, "min": 0, "max": 20, "step": 0.01}),
|
||||||
|
"EDM_s_churn": ("INT", {"default": 5, "min": 0, "max": 40, "step": 1}),
|
||||||
|
"s_noise": ("FLOAT", {"default": 1.003, "min": 1.0, "max": 1.1, "step": 0.001}),
|
||||||
|
"eta": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.01}),
|
||||||
|
"control_scale_start": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
|
||||||
|
"control_scale_end": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
|
||||||
|
"restore_cfg": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 20.0, "step": 0.05}),
|
||||||
|
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||||
|
"sampler": (
|
||||||
|
[
|
||||||
|
'RestoreDPMPP2MSampler',
|
||||||
|
'RestoreEDMSampler',
|
||||||
|
'TiledRestoreDPMPP2MSampler',
|
||||||
|
'TiledRestoreEDMSampler',
|
||||||
|
'EulerAncestralSampler'
|
||||||
|
], {
|
||||||
|
"default": 'RestoreEDMSampler'
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 32}),
|
||||||
|
"sampler_tile_stride": ("INT", {"default": 512, "min": 32, "max": 2048, "step": 32}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LATENT",)
|
||||||
|
RETURN_NAMES = ("latent",)
|
||||||
|
FUNCTION = "sample"
|
||||||
|
DESCRIPTION="Samples using SUPIR's modified diffusion."
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def sample(self, SUPIR_model, latents, steps, seed, cfg_scale_end, EDM_s_churn, s_noise, positive, negative,
|
||||||
|
cfg_scale_start, control_scale_start, control_scale_end, restore_cfg, keep_model_loaded, eta,
|
||||||
|
sampler, sampler_tile_size=1024, sampler_tile_stride=512):
|
||||||
|
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
self.sampler_config = {
|
||||||
|
'target': f'.sgm.modules.diffusionmodules.sampling.{sampler}',
|
||||||
|
'params': {
|
||||||
|
'num_steps': steps,
|
||||||
|
'restore_cfg': restore_cfg,
|
||||||
|
's_churn': EDM_s_churn,
|
||||||
|
's_noise': s_noise,
|
||||||
|
'eta': eta,
|
||||||
|
'discretization_config': {
|
||||||
|
'target': '.sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization'
|
||||||
|
},
|
||||||
|
'guider_config': {
|
||||||
|
'target': '.sgm.modules.diffusionmodules.guiders.LinearCFG',
|
||||||
|
'params': {
|
||||||
|
'scale': cfg_scale_end,
|
||||||
|
'scale_min': cfg_scale_start
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if 'Tiled' in sampler:
|
||||||
|
self.sampler_config['params']['tile_size'] = sampler_tile_size // 8
|
||||||
|
self.sampler_config['params']['tile_stride'] = sampler_tile_stride // 8
|
||||||
|
|
||||||
|
if not hasattr (self,'sampler') or self.sampler_config != self.current_sampler_config:
|
||||||
|
self.sampler = instantiate_from_config(self.sampler_config)
|
||||||
|
self.current_sampler_config = self.sampler_config
|
||||||
|
|
||||||
|
print("sampler_config: ", self.sampler_config)
|
||||||
|
|
||||||
|
SUPIR_model.denoiser.to(device)
|
||||||
|
SUPIR_model.model.diffusion_model.to(device)
|
||||||
|
SUPIR_model.model.control_model.to(device)
|
||||||
|
|
||||||
|
use_linear_control_scale = control_scale_start != control_scale_end
|
||||||
|
|
||||||
|
denoiser = lambda input, sigma, c, control_scale: SUPIR_model.denoiser(SUPIR_model.model, input, sigma, c, control_scale)
|
||||||
|
|
||||||
|
if len(positive) == 1:
|
||||||
|
positive = positive[0]
|
||||||
|
|
||||||
|
out = []
|
||||||
|
pbar = comfy.utils.ProgressBar(latents.shape[0])
|
||||||
|
for i, latent in enumerate(latents):
|
||||||
|
try:
|
||||||
|
noised_z = torch.randn_like(latent.unsqueeze(0), device=latents.device)
|
||||||
|
_samples = self.sampler(denoiser, noised_z, cond=positive, uc=negative, x_center=latent.unsqueeze(0), control_scale=control_scale_end,
|
||||||
|
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
|
||||||
|
|
||||||
|
except torch.cuda.OutOfMemoryError as e:
|
||||||
|
mm.free_memory(mm.get_total_memory(mm.get_torch_device()), mm.get_torch_device())
|
||||||
|
SUPIR_model = None
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
print("It's likely that too large of an image or batch_size for SUPIR was used,"
|
||||||
|
" and it has devoured all of the memory it had reserved, you may need to restart ComfyUI. Make sure you are using tiled_vae, "
|
||||||
|
" you can also try using fp8 for reduced memory usage if your system supports it.")
|
||||||
|
raise e
|
||||||
|
print("_samples: ", _samples.shape)
|
||||||
|
out.append(_samples)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
if not keep_model_loaded:
|
||||||
|
SUPIR_model.denoiser.to('cpu')
|
||||||
|
SUPIR_model.model.diffusion_model.to('cpu')
|
||||||
|
SUPIR_model.model.control_model.to('cpu')
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
if len(out[0].shape) == 4:
|
||||||
|
out_stacked = torch.cat(out, dim=0)
|
||||||
|
else:
|
||||||
|
out_stacked = torch.stack(out, dim=0)
|
||||||
|
|
||||||
|
print("out_stacked: ", _samples.shape)
|
||||||
|
return (out_stacked,)
|
||||||
|
|
||||||
|
class SUPIR_conditioner:
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"SUPIR_model": ("SUPIRMODEL",),
|
||||||
|
"latents": ("LATENT",),
|
||||||
|
"positive_prompt": ("STRING", {"multiline": True, "default": "high quality, detailed", }),
|
||||||
|
"negative_prompt": ("STRING", {"multiline": True, "default": "bad quality, blurry, messy", }),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"captions": ("STRING", {"forceInput": True, "multiline": False, "default": "", }),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("SUPIR_cond_pos", "SUPIR_cond_neg",)
|
||||||
|
RETURN_NAMES = ("positive", "negative",)
|
||||||
|
FUNCTION = "condition"
|
||||||
|
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def condition(self, SUPIR_model, latents, positive_prompt, negative_prompt, captions=""):
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
N, H, W, C = latents.shape
|
||||||
|
import copy
|
||||||
|
|
||||||
|
if not isinstance(captions, list):
|
||||||
|
captions_list = []
|
||||||
|
captions_list.append([captions])
|
||||||
|
captions_list = captions_list * N
|
||||||
|
else:
|
||||||
|
captions_list = captions
|
||||||
|
|
||||||
|
print("captions: ", captions_list)
|
||||||
|
|
||||||
|
SUPIR_model.conditioner.to(device)
|
||||||
|
latents = latents.to(device)
|
||||||
|
c = []
|
||||||
|
uc = []
|
||||||
|
pbar = comfy.utils.ProgressBar(N)
|
||||||
|
autocast_condition = (SUPIR_model.model.dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||||
|
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=SUPIR_model.model.dtype) if autocast_condition else nullcontext():
|
||||||
|
for i, caption in enumerate(captions_list):
|
||||||
|
cond = {}
|
||||||
|
cond['original_size_as_tuple'] = torch.tensor([[1024, 1024]]).to(device)
|
||||||
|
cond['crop_coords_top_left'] = torch.tensor([[0, 0]]).to(device)
|
||||||
|
cond['target_size_as_tuple'] = torch.tensor([[1024, 1024]]).to(device)
|
||||||
|
cond['aesthetic_score'] = torch.tensor([[9.0]]).to(device)
|
||||||
|
cond['control'] = latents[0].unsqueeze(0)
|
||||||
|
|
||||||
|
uncond = copy.deepcopy(cond)
|
||||||
|
uncond['txt'] = [negative_prompt]
|
||||||
|
|
||||||
|
cond['txt'] = [''.join([caption[0], positive_prompt])]
|
||||||
|
if i == 0:
|
||||||
|
_c, uc = SUPIR_model.conditioner.get_unconditional_conditioning(cond, uncond)
|
||||||
|
else:
|
||||||
|
_c, _ = SUPIR_model.conditioner.get_unconditional_conditioning(cond, None)
|
||||||
|
|
||||||
|
c.append(_c)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
|
||||||
|
SUPIR_model.conditioner.to('cpu')
|
||||||
|
|
||||||
|
return (c, uc,)
|
||||||
|
|
||||||
|
class SUPIR_model_loader:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"supir_model": (folder_paths.get_filename_list("checkpoints"),),
|
||||||
|
"sdxl_model": (folder_paths.get_filename_list("checkpoints"),),
|
||||||
|
"fp8_unet": ("BOOLEAN", {"default": False}),
|
||||||
|
"diffusion_dtype": (
|
||||||
|
[
|
||||||
|
'fp16',
|
||||||
|
'bf16',
|
||||||
|
'fp32',
|
||||||
|
'auto'
|
||||||
|
], {
|
||||||
|
"default": 'auto'
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("SUPIRMODEL", "SUPIRVAE")
|
||||||
|
RETURN_NAMES = ("SUPIR_model","SUPIR_VAE",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def process(self, supir_model, sdxl_model, diffusion_dtype, fp8_unet):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
|
||||||
|
SUPIR_MODEL_PATH = folder_paths.get_full_path("checkpoints", supir_model)
|
||||||
|
SDXL_MODEL_PATH = folder_paths.get_full_path("checkpoints", sdxl_model)
|
||||||
|
|
||||||
|
config_path = os.path.join(script_directory, "options/SUPIR_v0.yaml")
|
||||||
|
clip_config_path = os.path.join(script_directory, "configs/clip_vit_config.json")
|
||||||
|
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
|
||||||
|
|
||||||
|
custom_config = {
|
||||||
|
'sdxl_model': sdxl_model,
|
||||||
|
'diffusion_dtype': diffusion_dtype,
|
||||||
|
'supir_model': supir_model,
|
||||||
|
'fp8_unet': fp8_unet,
|
||||||
|
}
|
||||||
|
|
||||||
|
if diffusion_dtype == 'auto':
|
||||||
|
try:
|
||||||
|
if mm.should_use_bf16():
|
||||||
|
print("Diffusion using bf16")
|
||||||
|
dtype = torch.bfloat16
|
||||||
|
model_dtype = 'bf16'
|
||||||
|
elif mm.should_use_fp16():
|
||||||
|
print("Diffusion using using fp16")
|
||||||
|
dtype = torch.float16
|
||||||
|
model_dtype = 'fp16'
|
||||||
|
else:
|
||||||
|
print("Diffusion using using fp32")
|
||||||
|
dtype = torch.float32
|
||||||
|
model_dtype = 'fp32'
|
||||||
|
except:
|
||||||
|
raise AttributeError("ComfyUI version too old, can't autodecet properly. Set your dtypes manually.")
|
||||||
|
else:
|
||||||
|
print(f"Diffusion using using {diffusion_dtype}")
|
||||||
|
dtype = convert_dtype(diffusion_dtype)
|
||||||
|
model_dtype = diffusion_dtype
|
||||||
|
|
||||||
|
|
||||||
|
if not hasattr(self, "model") or self.model is None or self.current_config != custom_config:
|
||||||
|
self.current_config = custom_config
|
||||||
|
self.model = None
|
||||||
|
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
config = OmegaConf.load(config_path)
|
||||||
|
|
||||||
|
if XFORMERS_IS_AVAILABLE:
|
||||||
|
config.model.params.control_stage_config.params.spatial_transformer_attn_type = "softmax-xformers"
|
||||||
|
config.model.params.network_config.params.spatial_transformer_attn_type = "softmax-xformers"
|
||||||
|
config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla-xformers"
|
||||||
|
|
||||||
|
config.model.params.diffusion_dtype = model_dtype
|
||||||
|
config.model.target = ".SUPIR.models.SUPIR_model_v2.SUPIRModel"
|
||||||
|
pbar = comfy.utils.ProgressBar(7)
|
||||||
|
|
||||||
|
self.model = instantiate_from_config(config.model).cpu()
|
||||||
|
pbar.update(1)
|
||||||
|
try:
|
||||||
|
print(f'Attempting to load SUPIR model: [{SUPIR_MODEL_PATH}]')
|
||||||
|
supir_state_dict = load_state_dict(SUPIR_MODEL_PATH)
|
||||||
|
pbar.update(1)
|
||||||
|
except:
|
||||||
|
raise Exception("Failed to load SUPIR model")
|
||||||
|
try:
|
||||||
|
print(f"Attempting to load SDXL model: [{SDXL_MODEL_PATH}]")
|
||||||
|
sdxl_state_dict = load_state_dict(SDXL_MODEL_PATH)
|
||||||
|
pbar.update(1)
|
||||||
|
except:
|
||||||
|
raise Exception("Failed to load SDXL model")
|
||||||
|
self.model.load_state_dict(supir_state_dict, strict=False)
|
||||||
|
pbar.update(1)
|
||||||
|
self.model.load_state_dict(sdxl_state_dict, strict=False)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
del supir_state_dict
|
||||||
|
|
||||||
|
#first clip model from SDXL checkpoint
|
||||||
|
try:
|
||||||
|
print("Loading first clip model from SDXL checkpoint")
|
||||||
|
|
||||||
|
replace_prefix = {}
|
||||||
|
replace_prefix["conditioner.embedders.0.transformer."] = ""
|
||||||
|
|
||||||
|
sd = comfy.utils.state_dict_prefix_replace(sdxl_state_dict, replace_prefix, filter_keys=False)
|
||||||
|
clip_text_config = CLIPTextConfig.from_pretrained(clip_config_path)
|
||||||
|
self.model.conditioner.embedders[0].tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
self.model.conditioner.embedders[0].transformer = CLIPTextModel(clip_text_config)
|
||||||
|
self.model.conditioner.embedders[0].transformer.load_state_dict(sd, strict=False)
|
||||||
|
self.model.conditioner.embedders[0].eval()
|
||||||
|
for param in self.model.conditioner.embedders[0].parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
pbar.update(1)
|
||||||
|
except:
|
||||||
|
raise Exception("Failed to load first clip model from SDXL checkpoint")
|
||||||
|
|
||||||
|
del sdxl_state_dict
|
||||||
|
|
||||||
|
#second clip model from SDXL checkpoint
|
||||||
|
try:
|
||||||
|
print("Loading second clip model from SDXL checkpoint")
|
||||||
|
replace_prefix2 = {}
|
||||||
|
replace_prefix2["conditioner.embedders.1.model."] = ""
|
||||||
|
sd = comfy.utils.state_dict_prefix_replace(sd, replace_prefix2, filter_keys=True)
|
||||||
|
clip_g = build_text_model_from_openai_state_dict(sd, cast_dtype=dtype)
|
||||||
|
self.model.conditioner.embedders[1].model = clip_g
|
||||||
|
pbar.update(1)
|
||||||
|
except:
|
||||||
|
raise Exception("Failed to load second clip model from SDXL checkpoint")
|
||||||
|
|
||||||
|
del sd, clip_g
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
self.model.to(dtype)
|
||||||
|
|
||||||
|
#only unets and/or vae to fp8
|
||||||
|
if fp8_unet:
|
||||||
|
self.model.model.to(torch.float8_e4m3fn)
|
||||||
|
|
||||||
|
return (self.model, self.model.first_stage_model,)
|
||||||
|
|
||||||
|
class SUPIR_tiles:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"image": ("IMAGE",),
|
||||||
|
"tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
"tile_stride": ("INT", {"default": 256, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE", "INT", "INT",)
|
||||||
|
RETURN_NAMES = ("image_tiles", "tile_size", "tile_stride",)
|
||||||
|
FUNCTION = "tile"
|
||||||
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
|
def tile(self, image, tile_size, tile_stride):
|
||||||
|
|
||||||
|
def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: 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
|
||||||
|
|
||||||
|
image = image.permute(0, 3, 1, 2)
|
||||||
|
_, _, h, w = image.shape
|
||||||
|
|
||||||
|
tiles_iterator = _sliding_windows(h, w, tile_size, tile_stride)
|
||||||
|
|
||||||
|
tiles = []
|
||||||
|
for hi, hi_end, wi, wi_end in tiles_iterator:
|
||||||
|
tile = image[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
|
||||||
|
tiles.append(tile)
|
||||||
|
out = torch.cat(tiles, dim=0).to(torch.float32).permute(0, 2, 3, 1)
|
||||||
|
print(out.shape)
|
||||||
|
print("len(tiles): ", len(tiles))
|
||||||
|
|
||||||
|
return (out, tile_size, tile_stride,)
|
||||||
@@ -564,9 +564,9 @@ class SpatialTransformer(nn.Module):
|
|||||||
sdp_backend=None,
|
sdp_backend=None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
print(
|
# print(
|
||||||
f"constructing {self.__class__.__name__} of depth {depth} w/ {in_channels} channels and {n_heads} heads"
|
# f"constructing {self.__class__.__name__} of depth {depth} w/ {in_channels} channels and {n_heads} heads"
|
||||||
)
|
# )
|
||||||
from omegaconf import ListConfig
|
from omegaconf import ListConfig
|
||||||
|
|
||||||
if exists(context_dim) and not isinstance(context_dim, (list, ListConfig)):
|
if exists(context_dim) and not isinstance(context_dim, (list, ListConfig)):
|
||||||
|
|||||||
@@ -186,13 +186,13 @@ class Downsample(nn.Module):
|
|||||||
self.dims = dims
|
self.dims = dims
|
||||||
stride = 2 if dims != 3 else ((1, 2, 2) if not third_down else (2, 2, 2))
|
stride = 2 if dims != 3 else ((1, 2, 2) if not third_down else (2, 2, 2))
|
||||||
if use_conv:
|
if use_conv:
|
||||||
print(f"Building a Downsample layer with {dims} dims.")
|
# print(f"Building a Downsample layer with {dims} dims.")
|
||||||
print(
|
# print(
|
||||||
f" --> settings are: \n in-chn: {self.channels}, out-chn: {self.out_channels}, "
|
# f" --> settings are: \n in-chn: {self.channels}, out-chn: {self.out_channels}, "
|
||||||
f"kernel-size: 3, stride: {stride}, padding: {padding}"
|
# f"kernel-size: 3, stride: {stride}, padding: {padding}"
|
||||||
)
|
# )
|
||||||
if dims == 3:
|
# if dims == 3:
|
||||||
print(f" --> Downsampling third axis (time): {third_down}")
|
# print(f" --> Downsampling third axis (time): {third_down}")
|
||||||
self.op = conv_nd(
|
self.op = conv_nd(
|
||||||
dims,
|
dims,
|
||||||
self.channels,
|
self.channels,
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from typing import Dict, Union
|
|||||||
import torch
|
import torch
|
||||||
from omegaconf import ListConfig, OmegaConf
|
from omegaconf import ListConfig, OmegaConf
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, get_sigmas_karras
|
||||||
from ...modules.diffusionmodules.sampling_utils import (
|
from ...modules.diffusionmodules.sampling_utils import (
|
||||||
get_ancestral_step,
|
get_ancestral_step,
|
||||||
linear_multistep_coeff,
|
linear_multistep_coeff,
|
||||||
@@ -555,3 +555,241 @@ def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int):
|
|||||||
coords.append((hi, hi + tile_size, wi, wi + tile_size))
|
coords.append((hi, hi + tile_size, wi, wi + tile_size))
|
||||||
return coords
|
return coords
|
||||||
|
|
||||||
|
class RestoreDPMPP2MSampler(DPMPP2MSampler):
|
||||||
|
def __init__(self, s_churn=0.0, s_tmin=0.0, s_tmax=float("inf"), s_noise=1.0, restore_cfg=4.0,
|
||||||
|
restore_cfg_s_tmin=0.05, eta=1., *args, **kwargs):
|
||||||
|
self.s_noise = s_noise
|
||||||
|
self.eta = eta
|
||||||
|
self.restore_cfg = restore_cfg
|
||||||
|
self.restore_cfg_s_tmin = restore_cfg_s_tmin
|
||||||
|
self.sigma_max = 14.6146
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
def denoise(self, x, denoiser, sigma, cond, uc, control_scale=1.0):
|
||||||
|
denoised = denoiser(*self.guider.prepare_inputs(x, sigma, cond, uc), control_scale)
|
||||||
|
denoised = self.guider(denoised, sigma)
|
||||||
|
return denoised
|
||||||
|
|
||||||
|
def get_mult(self, h, r, t, t_next, previous_sigma):
|
||||||
|
eta_h = self.eta * h
|
||||||
|
mult1 = to_sigma(t_next) / to_sigma(t) * (-eta_h).exp()
|
||||||
|
mult2 = (-h -eta_h).expm1()
|
||||||
|
|
||||||
|
if previous_sigma is not None:
|
||||||
|
mult3 = 1 + 1 / (2 * r)
|
||||||
|
mult4 = 1 / (2 * r)
|
||||||
|
return mult1, mult2, mult3, mult4
|
||||||
|
else:
|
||||||
|
return mult1, mult2
|
||||||
|
|
||||||
|
|
||||||
|
def sampler_step(
|
||||||
|
self,
|
||||||
|
old_denoised,
|
||||||
|
previous_sigma,
|
||||||
|
sigma,
|
||||||
|
next_sigma,
|
||||||
|
denoiser,
|
||||||
|
x,
|
||||||
|
cond,
|
||||||
|
uc=None,
|
||||||
|
eps_noise=None,
|
||||||
|
x_center=None,
|
||||||
|
control_scale=1.0,
|
||||||
|
use_linear_control_scale=False,
|
||||||
|
control_scale_start=0.0
|
||||||
|
):
|
||||||
|
if use_linear_control_scale:
|
||||||
|
control_scale = (sigma[0].item() / self.sigma_max) * (control_scale_start - control_scale) + control_scale
|
||||||
|
|
||||||
|
denoised = self.denoise(x, denoiser, sigma, cond, uc, control_scale=control_scale)
|
||||||
|
|
||||||
|
if (next_sigma[0] > self.restore_cfg_s_tmin) and (self.restore_cfg > 0):
|
||||||
|
d_center = (denoised - x_center)
|
||||||
|
denoised = denoised - d_center * ((sigma.view(-1, 1, 1, 1) / self.sigma_max) ** self.restore_cfg)
|
||||||
|
|
||||||
|
h, r, t, t_next = self.get_variables(sigma, next_sigma, previous_sigma)
|
||||||
|
eta_h = self.eta * h
|
||||||
|
mult = [
|
||||||
|
append_dims(mult, x.ndim)
|
||||||
|
for mult in self.get_mult(h, r, t, t_next, previous_sigma)
|
||||||
|
]
|
||||||
|
|
||||||
|
x_standard = mult[0] * x - mult[1] * denoised
|
||||||
|
if old_denoised is None or torch.sum(next_sigma) < 1e-14:
|
||||||
|
# Save a network evaluation if all noise levels are 0 or on the first step
|
||||||
|
return x_standard, denoised
|
||||||
|
else:
|
||||||
|
denoised_d = mult[2] * denoised - mult[3] * old_denoised
|
||||||
|
x_advanced = mult[0] * x - mult[1] * denoised_d
|
||||||
|
|
||||||
|
# apply correction if noise level is not 0 and not first step
|
||||||
|
x = torch.where(
|
||||||
|
append_dims(next_sigma, x.ndim) > 0.0, x_advanced, x_standard
|
||||||
|
)
|
||||||
|
if self.eta:
|
||||||
|
x = x + eps_noise * next_sigma * (-2 * eta_h).expm1().neg().sqrt() * self.s_noise
|
||||||
|
|
||||||
|
return x, denoised
|
||||||
|
|
||||||
|
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, x_center=None, control_scale=1.0,
|
||||||
|
use_linear_control_scale=False, control_scale_start=0.0, **kwargs):
|
||||||
|
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||||
|
x, cond, uc, num_steps
|
||||||
|
)
|
||||||
|
sigmas_min, sigmas_max = sigmas[-2].cpu(), sigmas[0].cpu()
|
||||||
|
sigmas_new = get_sigmas_karras(self.num_steps, sigmas_min, sigmas_max, device=x.device)
|
||||||
|
sigmas = sigmas_new
|
||||||
|
|
||||||
|
noise_sampler = BrownianTreeNoiseSampler(x, sigmas_min, sigmas_max)
|
||||||
|
|
||||||
|
old_denoised = None
|
||||||
|
pbar_comfy = comfy.utils.ProgressBar(num_sigmas)
|
||||||
|
for i in self.get_sigma_gen(num_sigmas):
|
||||||
|
if i > 0 and torch.sum(s_in * sigmas[i + 1]) > 1e-14:
|
||||||
|
eps_noise = noise_sampler(s_in * sigmas[i], s_in * sigmas[i + 1])
|
||||||
|
else:
|
||||||
|
eps_noise = None
|
||||||
|
x, old_denoised = self.sampler_step(
|
||||||
|
old_denoised,
|
||||||
|
None if i == 0 else s_in * sigmas[i - 1],
|
||||||
|
s_in * sigmas[i],
|
||||||
|
s_in * sigmas[i + 1],
|
||||||
|
denoiser,
|
||||||
|
x,
|
||||||
|
cond,
|
||||||
|
uc=uc,
|
||||||
|
eps_noise=eps_noise,
|
||||||
|
control_scale=control_scale,
|
||||||
|
x_center=x_center,
|
||||||
|
use_linear_control_scale=use_linear_control_scale,
|
||||||
|
control_scale_start=control_scale_start,
|
||||||
|
)
|
||||||
|
pbar_comfy.update(1)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
class TiledRestoreDPMPP2MSampler(RestoreDPMPP2MSampler):
|
||||||
|
def __init__(self, tile_size=128, tile_stride=64, *args, **kwargs):
|
||||||
|
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.tile_size = tile_size
|
||||||
|
self.tile_stride = tile_stride
|
||||||
|
self.tile_weights = gaussian_weights(self.tile_size, self.tile_size, 1)
|
||||||
|
|
||||||
|
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, control_scale=1.0, **kwargs):
|
||||||
|
use_local_prompt = isinstance(cond, list)
|
||||||
|
b, _, h, w = x.shape
|
||||||
|
latent_tiles_iterator = _sliding_windows(h, w, self.tile_size, self.tile_stride)
|
||||||
|
print(f"Image divided into {len(latent_tiles_iterator)} tiles")
|
||||||
|
print("Conds received: ", len(cond))
|
||||||
|
tile_weights = self.tile_weights.repeat(b, 1, 1, 1)
|
||||||
|
if not use_local_prompt:
|
||||||
|
LQ_latent = cond['control']
|
||||||
|
else:
|
||||||
|
assert len(cond) == len(latent_tiles_iterator), "Number of local prompts should be equal to number of tiles"
|
||||||
|
LQ_latent = cond[0]['control']
|
||||||
|
print("LQ_latent shape: ",LQ_latent.shape)
|
||||||
|
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||||
|
x, cond, uc, num_steps
|
||||||
|
)
|
||||||
|
sigmas_min, sigmas_max = sigmas[-2].cpu(), sigmas[0].cpu()
|
||||||
|
sigmas_new = get_sigmas_karras(self.num_steps, sigmas_min, sigmas_max, device=x.device)
|
||||||
|
sigmas = sigmas_new
|
||||||
|
|
||||||
|
noise_sampler = BrownianTreeNoiseSampler(x, sigmas_min, sigmas_max)
|
||||||
|
|
||||||
|
old_denoised = None
|
||||||
|
pbar_comfy = comfy.utils.ProgressBar(num_sigmas)
|
||||||
|
for _idx, i in enumerate(self.get_sigma_gen(num_sigmas)):
|
||||||
|
if i > 0 and torch.sum(s_in * sigmas[i + 1]) > 1e-14:
|
||||||
|
eps_noise = noise_sampler(s_in * sigmas[i], s_in * sigmas[i + 1])
|
||||||
|
else:
|
||||||
|
eps_noise = torch.zeros_like(x)
|
||||||
|
x_next = torch.zeros_like(x)
|
||||||
|
old_denoised_next = torch.zeros_like(x)
|
||||||
|
count = torch.zeros_like(x)
|
||||||
|
for j, (hi, hi_end, wi, wi_end) in enumerate(latent_tiles_iterator):
|
||||||
|
x_tile = x[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
_eps_noise = eps_noise[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
if old_denoised is not None:
|
||||||
|
old_denoised_tile = old_denoised[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
else:
|
||||||
|
old_denoised_tile = None
|
||||||
|
if use_local_prompt:
|
||||||
|
_cond = cond[j]
|
||||||
|
else:
|
||||||
|
_cond = cond
|
||||||
|
_cond['control'] = LQ_latent[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
uc['control'] = LQ_latent[:, :, hi:hi_end, wi:wi_end]
|
||||||
|
_x, _old_denoised = self.sampler_step(
|
||||||
|
old_denoised_tile,
|
||||||
|
None if i == 0 else s_in * sigmas[i - 1],
|
||||||
|
s_in * sigmas[i],
|
||||||
|
s_in * sigmas[i + 1],
|
||||||
|
denoiser,
|
||||||
|
x_tile,
|
||||||
|
_cond,
|
||||||
|
uc=uc,
|
||||||
|
eps_noise=_eps_noise,
|
||||||
|
control_scale=control_scale,
|
||||||
|
)
|
||||||
|
x_next[:, :, hi:hi_end, wi:wi_end] += _x * tile_weights
|
||||||
|
old_denoised_next[:, :, hi:hi_end, wi:wi_end] += _old_denoised * tile_weights
|
||||||
|
count[:, :, hi:hi_end, wi:wi_end] += tile_weights
|
||||||
|
old_denoised_next /= count
|
||||||
|
x_next /= count
|
||||||
|
x = x_next
|
||||||
|
old_denoised = old_denoised_next
|
||||||
|
pbar_comfy.update(1)
|
||||||
|
return x
|
||||||
|
|
||||||
|
class SubstepSampler(EulerAncestralSampler):
|
||||||
|
def __init__(self, s_churn=0.0, s_tmin=0.0, s_tmax=float("inf"), s_noise=1.0, restore_cfg=4.0,
|
||||||
|
restore_cfg_s_tmin=0.05, eta=1., n_sample_steps=4, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.n_sample_steps = n_sample_steps
|
||||||
|
self.steps_subset = [0, 100, 200, 300, 1000]
|
||||||
|
|
||||||
|
def prepare_sampling_loop(self, x, cond, uc=None, num_steps=None):
|
||||||
|
sigmas = self.discretization(1000, device=self.device)
|
||||||
|
sigmas = sigmas[
|
||||||
|
self.steps_subset[: self.num_steps] + self.steps_subset[-1:]
|
||||||
|
]
|
||||||
|
print(sigmas)
|
||||||
|
# uc = cond
|
||||||
|
x *= torch.sqrt(1.0 + sigmas[0] ** 2.0)
|
||||||
|
num_sigmas = len(sigmas)
|
||||||
|
s_in = x.new_ones([x.shape[0]])
|
||||||
|
return x, s_in, sigmas, num_sigmas, cond, uc
|
||||||
|
|
||||||
|
def denoise(self, x, denoiser, sigma, cond, uc, control_scale=1.0):
|
||||||
|
denoised = denoiser(*self.guider.prepare_inputs(x, sigma, cond, uc), control_scale)
|
||||||
|
denoised = self.guider(denoised, sigma)
|
||||||
|
return denoised
|
||||||
|
|
||||||
|
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, control_scale=1.0, *args, **kwargs):
|
||||||
|
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||||
|
x, cond, uc, num_steps
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in self.get_sigma_gen(num_sigmas):
|
||||||
|
x = self.sampler_step(
|
||||||
|
s_in * sigmas[i],
|
||||||
|
s_in * sigmas[i + 1],
|
||||||
|
denoiser,
|
||||||
|
x,
|
||||||
|
cond,
|
||||||
|
uc,
|
||||||
|
control_scale=control_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc, control_scale=1.0):
|
||||||
|
sigma_down, sigma_up = get_ancestral_step(sigma, next_sigma, eta=self.eta)
|
||||||
|
denoised = self.denoise(x, denoiser, sigma, cond, uc, control_scale=control_scale)
|
||||||
|
x = self.ancestral_euler_step(x, denoised, sigma, sigma_down)
|
||||||
|
x = self.ancestral_step(x, sigma, next_sigma, sigma_up)
|
||||||
|
|
||||||
|
return x
|
||||||
@@ -99,10 +99,10 @@ class GeneralConditioner(nn.Module):
|
|||||||
for param in embedder.parameters():
|
for param in embedder.parameters():
|
||||||
param.requires_grad = False
|
param.requires_grad = False
|
||||||
embedder.eval()
|
embedder.eval()
|
||||||
print(
|
# print(
|
||||||
f"Initialized embedder #{n}: {embedder.__class__.__name__} "
|
# f"Initialized embedder #{n}: {embedder.__class__.__name__} "
|
||||||
f"with {count_params(embedder, False)} params. Trainable: {embedder.is_trainable}"
|
# f"with {count_params(embedder, False)} params. Trainable: {embedder.is_trainable}"
|
||||||
)
|
# )
|
||||||
|
|
||||||
if "input_key" in embconfig:
|
if "input_key" in embconfig:
|
||||||
embedder.input_key = embconfig["input_key"]
|
embedder.input_key = embconfig["input_key"]
|
||||||
|
|||||||
Reference in New Issue
Block a user