6 Commits
Author SHA1 Message Date
kijai f0ab85b648 Update 2024-03-17 13:42:47 +02:00
Kijai 07d30c74ab sampler fixes 2024-03-13 18:24:49 +02:00
Kijai 24c06c19e6 Tiled captioning 2024-03-12 17:43:44 +02:00
Kijai 76101ea540 Updae 2024-03-12 13:08:41 +02:00
Kijai da1b08ed22 Update nodes_v2.py 2024-03-11 17:13:30 +02:00
Kijai b1d1639575 Initial separation of the nodes 2024-03-11 17:02:40 +02:00
10 changed files with 2511 additions and 88 deletions
+28 -36
View File
@@ -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)
+196
View File
@@ -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
View File
@@ -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
+44 -24
View File
@@ -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
View File
@@ -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,)
+3 -3
View File
@@ -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)):
+7 -7
View File
@@ -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,
+239 -1
View File
@@ -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
+4 -4
View File
@@ -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"]