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 pytorch_lightning import seed_everything
|
||||
from ...SUPIR.utils.tilevae import VAEHook
|
||||
from ...SUPIR.util import convert_dtype
|
||||
from contextlib import nullcontext
|
||||
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.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.ae_dtype = convert_dtype(ae_dtype)
|
||||
self.model.dtype = convert_dtype(diffusion_dtype)
|
||||
|
||||
self.p_p = p_p
|
||||
self.n_p = n_p
|
||||
@@ -72,7 +58,6 @@ class SUPIRModel(DiffusionEngine):
|
||||
@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)
|
||||
@@ -120,29 +105,51 @@ class SUPIRModel(DiffusionEngine):
|
||||
self.sampler_config.params.s_noise = s_noise
|
||||
self.sampler = instantiate_from_config(self.sampler_config)
|
||||
|
||||
print("Sampler: ", self.sampler_config.target)
|
||||
print("sampler_config: ", self.sampler_config.params)
|
||||
|
||||
if seed == -1:
|
||||
seed = random.randint(0, 65535)
|
||||
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)
|
||||
|
||||
x_stage1 = self.decode_first_stage(_z)
|
||||
|
||||
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)
|
||||
self.conditioner.to('cpu')
|
||||
|
||||
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)
|
||||
|
||||
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,
|
||||
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)
|
||||
self.first_stage_model.to('cpu')
|
||||
|
||||
if color_fix_type == 'Wavelet':
|
||||
samples = wavelet_reconstruction(samples, x_stage1)
|
||||
elif color_fix_type == 'AdaIn':
|
||||
@@ -194,18 +201,3 @@ class SUPIRModel(DiffusionEngine):
|
||||
_c, _ = self.conditioner.get_unconditional_conditioning(batch, None)
|
||||
c.append(_c)
|
||||
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"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,6 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
from contextlib import nullcontext
|
||||
from omegaconf import OmegaConf
|
||||
import comfy.utils
|
||||
import comfy.model_management as mm
|
||||
@@ -13,6 +11,7 @@ import torch.cuda
|
||||
from .sgm.util import instantiate_from_config
|
||||
from .SUPIR.util import convert_dtype, load_state_dict
|
||||
import open_clip
|
||||
from contextlib import contextmanager
|
||||
|
||||
from transformers import (
|
||||
CLIPTextModel,
|
||||
@@ -34,12 +33,21 @@ except:
|
||||
def dummy_build_vision_tower(*args, **kwargs):
|
||||
# Monkey patch the CLIP class before you create an instance.
|
||||
return None
|
||||
open_clip.model._build_vision_tower = dummy_build_vision_tower
|
||||
|
||||
@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]
|
||||
@@ -56,11 +64,13 @@ def build_text_model_from_openai_state_dict(
|
||||
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, # OpenAI models were trained with QuickGELU
|
||||
quick_gelu=True,
|
||||
cast_dtype=cast_dtype,
|
||||
)
|
||||
|
||||
@@ -128,6 +138,15 @@ class SUPIR_Upscale:
|
||||
"use_tiled_sampling": ("BOOLEAN", {"default": False}),
|
||||
"sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "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,
|
||||
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",
|
||||
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()
|
||||
mm.unload_all_models()
|
||||
|
||||
@@ -160,6 +179,9 @@ class SUPIR_Upscale:
|
||||
'use_tiled_vae': use_tiled_vae,
|
||||
'supir_model': supir_model,
|
||||
'use_tiled_sampling': use_tiled_sampling,
|
||||
'fp8_unet': fp8_unet,
|
||||
'fp8_vae': fp8_vae,
|
||||
'sampler': sampler
|
||||
}
|
||||
|
||||
if diffusion_dtype == 'auto':
|
||||
@@ -207,9 +229,11 @@ class SUPIR_Upscale:
|
||||
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_stride = sampler_tile_stride // 8
|
||||
config.model.params.sampler_config.target = f".sgm.modules.diffusionmodules.sampling.Tiled{sampler}"
|
||||
print("Using tiled sampling")
|
||||
else:
|
||||
config = OmegaConf.load(config_path)
|
||||
config.model.params.sampler_config.target = f".sgm.modules.diffusionmodules.sampling.{sampler}"
|
||||
print("Using non-tiled sampling")
|
||||
|
||||
if XFORMERS_IS_AVAILABLE:
|
||||
@@ -267,35 +291,30 @@ class SUPIR_Upscale:
|
||||
clip_g = build_text_model_from_openai_state_dict(sd, cast_dtype=dtype)
|
||||
self.model.conditioner.embedders[1].model = clip_g
|
||||
except:
|
||||
|
||||
raise Exception("Failed to load second clip model from SDXL checkpoint")
|
||||
|
||||
del sd, clip_g
|
||||
mm.soft_empty_cache()
|
||||
|
||||
try:
|
||||
self.model.to(dtype)
|
||||
self.model.to(device)
|
||||
except Exception as e:
|
||||
print("Failed to move model to device")
|
||||
print(e)
|
||||
import gc
|
||||
# unload everything and give up
|
||||
self.model = None
|
||||
del self.model
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
#only unets and/or vae to fp8
|
||||
if fp8_unet:
|
||||
self.model.model.to(torch.float8_e4m3fn)
|
||||
if fp8_vae:
|
||||
self.model.first_stage_model.to(torch.float8_e4m3fn)
|
||||
|
||||
if use_tiled_vae:
|
||||
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)
|
||||
B, H, W, C = image.shape
|
||||
new_height = H // 64 * 64
|
||||
new_width = W // 64 * 64
|
||||
image = image.permute(0, 3, 1, 2).contiguous()
|
||||
resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False)
|
||||
upscaled_image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
||||
B, H, W, C = upscaled_image.shape
|
||||
new_height = H if H % 64 == 0 else ((H // 64) + 1) * 64
|
||||
new_width = W if W % 64 == 0 else ((W // 64) + 1) * 64
|
||||
upscaled_image = upscaled_image.permute(0, 3, 1, 2)
|
||||
resized_image = F.interpolate(upscaled_image, size=(new_height, new_width), mode='bicubic', align_corners=False)
|
||||
resized_image = resized_image.to(device)
|
||||
|
||||
captions_list = []
|
||||
captions_list.append(captions)
|
||||
print("captions: ", captions_list)
|
||||
@@ -329,7 +348,8 @@ class SUPIR_Upscale:
|
||||
self.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")
|
||||
" 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
|
||||
|
||||
out.append(samples.squeeze(0).cpu())
|
||||
@@ -345,7 +365,7 @@ class SUPIR_Upscale:
|
||||
else:
|
||||
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,)
|
||||
|
||||
|
||||
+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,
|
||||
):
|
||||
super().__init__()
|
||||
print(
|
||||
f"constructing {self.__class__.__name__} of depth {depth} w/ {in_channels} channels and {n_heads} heads"
|
||||
)
|
||||
# print(
|
||||
# f"constructing {self.__class__.__name__} of depth {depth} w/ {in_channels} channels and {n_heads} heads"
|
||||
# )
|
||||
from omegaconf import ListConfig
|
||||
|
||||
if exists(context_dim) and not isinstance(context_dim, (list, ListConfig)):
|
||||
|
||||
@@ -186,13 +186,13 @@ class Downsample(nn.Module):
|
||||
self.dims = dims
|
||||
stride = 2 if dims != 3 else ((1, 2, 2) if not third_down else (2, 2, 2))
|
||||
if use_conv:
|
||||
print(f"Building a Downsample layer with {dims} dims.")
|
||||
print(
|
||||
f" --> settings are: \n in-chn: {self.channels}, out-chn: {self.out_channels}, "
|
||||
f"kernel-size: 3, stride: {stride}, padding: {padding}"
|
||||
)
|
||||
if dims == 3:
|
||||
print(f" --> Downsampling third axis (time): {third_down}")
|
||||
# print(f"Building a Downsample layer with {dims} dims.")
|
||||
# print(
|
||||
# f" --> settings are: \n in-chn: {self.channels}, out-chn: {self.out_channels}, "
|
||||
# f"kernel-size: 3, stride: {stride}, padding: {padding}"
|
||||
# )
|
||||
# if dims == 3:
|
||||
# print(f" --> Downsampling third axis (time): {third_down}")
|
||||
self.op = conv_nd(
|
||||
dims,
|
||||
self.channels,
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Dict, Union
|
||||
import torch
|
||||
from omegaconf import ListConfig, OmegaConf
|
||||
from tqdm import tqdm
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, get_sigmas_karras
|
||||
from ...modules.diffusionmodules.sampling_utils import (
|
||||
get_ancestral_step,
|
||||
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))
|
||||
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():
|
||||
param.requires_grad = False
|
||||
embedder.eval()
|
||||
print(
|
||||
f"Initialized embedder #{n}: {embedder.__class__.__name__} "
|
||||
f"with {count_params(embedder, False)} params. Trainable: {embedder.is_trainable}"
|
||||
)
|
||||
# print(
|
||||
# f"Initialized embedder #{n}: {embedder.__class__.__name__} "
|
||||
# f"with {count_params(embedder, False)} params. Trainable: {embedder.is_trainable}"
|
||||
# )
|
||||
|
||||
if "input_key" in embconfig:
|
||||
embedder.input_key = embconfig["input_key"]
|
||||
|
||||
Reference in New Issue
Block a user