diff --git a/SUPIR/models/SUPIR_model.py b/SUPIR/models/SUPIR_model.py index af1bcb9..2d69eda 100644 --- a/SUPIR/models/SUPIR_model.py +++ b/SUPIR/models/SUPIR_model.py @@ -200,19 +200,4 @@ class SUPIRModel(DiffusionEngine): else: _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) + return c, uc \ No newline at end of file diff --git a/SUPIR/models/SUPIR_model_v2.py b/SUPIR/models/SUPIR_model_v2.py new file mode 100644 index 0000000..b510f7c --- /dev/null +++ b/SUPIR/models/SUPIR_model_v2.py @@ -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 diff --git a/SUPIR/utils/tilevae.py b/SUPIR/utils/tilevae.py index a8fc379..de92ef8 100644 --- a/SUPIR/utils/tilevae.py +++ b/SUPIR/utils/tilevae.py @@ -895,7 +895,7 @@ class VAEHook: # Task queue execution pbar = tqdm(total=num_tiles * len(task_queues[0]), desc=f"[Tiled VAE]: Executing {'Decoder' if is_decoder else 'Encoder'} Task Queue: ") - + pbar_comfy = comfy.utils.ProgressBar(num_tiles * len(task_queues[0])) # execute the task back and forth when switch tiles so that we always # keep one tile on the GPU to reduce unnecessary data transfer forward = True @@ -937,6 +937,7 @@ class VAEHook: tile = task[1](tile) #print(tiles[i].shape, tile.shape, task) pbar.update(1) + pbar_comfy.update(1) if interrupted: break diff --git a/__init__.py b/__init__.py index 2e96bd6..c9bcb99 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] \ No newline at end of file diff --git a/examples/supir_upscale_example_lightning_10steps.json b/examples/supir_upscale_example_lightning_10steps.json new file mode 100644 index 0000000..1b47215 --- /dev/null +++ b/examples/supir_upscale_example_lightning_10steps.json @@ -0,0 +1,913 @@ +{ + "last_node_id": 20, + "last_link_id": 34, + "nodes": [ + { + "id": 12, + "type": "Reroute", + "pos": [ + 894, + -87 + ], + "size": [ + 124, + 26 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "", + "type": "*", + "link": 21 + } + ], + "outputs": [ + { + "name": "SUPIRMODEL", + "type": "SUPIRMODEL", + "links": [ + 22, + 23 + ], + "slot_index": 0 + } + ], + "properties": { + "showOutputText": true, + "horizontal": false + } + }, + { + "id": 11, + "type": "SUPIR_encode", + "pos": [ + 690, + 21 + ], + "size": [ + 217.85013759613048, + 126 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "SUPIR_VAE", + "type": "SUPIRVAE", + "link": 15 + }, + { + "name": "image", + "type": "IMAGE", + "link": 16 + } + ], + "outputs": [ + { + "name": "latent", + "type": "LATENT", + "links": [ + 17 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "SUPIR_encode" + }, + "widgets_values": [ + true, + 512, + "auto" + ] + }, + { + "id": 3, + "type": "PreviewImage", + "pos": [ + 1722, + 60 + ], + "size": [ + 985.112306152344, + 988.225142883301 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 27 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 15, + "type": "SetNode", + "pos": [ + 347, + 517 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "link": 28 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_InputImage", + "properties": { + "previousName": "InputImage" + }, + "widgets_values": [ + "InputImage" + ], + "color": "#2a363b", + "bgcolor": "#3f5159" + }, + { + "id": 13, + "type": "ImageResize+", + "pos": [ + 348, + 238 + ], + "size": { + "0": 315, + "1": 218 + }, + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 24 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 25 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResize+" + }, + "widgets_values": [ + 1536, + 1536, + "lanczos", + false, + "always", + 0 + ] + }, + { + "id": 17, + "type": "SetNode", + "pos": [ + 392, + -109 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "SUPIRVAE", + "type": "SUPIRVAE", + "link": 30 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_SUPIRVAE", + "properties": { + "previousName": "SUPIRVAE" + }, + "widgets_values": [ + "SUPIRVAE" + ] + }, + { + "id": 16, + "type": "GetNode", + "pos": [ + 2027, + -130 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 29 + ] + } + ], + "title": "Get_InputImage", + "properties": {}, + "widgets_values": [ + "InputImage" + ], + "color": "#2a363b", + "bgcolor": "#3f5159" + }, + { + "id": 19, + "type": "GetNode", + "pos": [ + 1739, + -142 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "SUPIRVAE", + "type": "SUPIRVAE", + "links": [ + 32 + ] + } + ], + "title": "Get_SUPIRVAE", + "properties": {}, + "widgets_values": [ + "SUPIRVAE" + ], + "color": "#2a363b", + "bgcolor": "#3f5159" + }, + { + "id": 9, + "type": "SUPIR_conditioner", + "pos": [ + 944, + 86 + ], + "size": [ + 401.72, + 200.86 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "SUPIR_model", + "type": "SUPIRMODEL", + "link": 22, + "slot_index": 0 + }, + { + "name": "latents", + "type": "LATENT", + "link": 20, + "slot_index": 1 + }, + { + "name": "captions", + "type": "STRING", + "link": null, + "widget": { + "name": "captions" + } + } + ], + "outputs": [ + { + "name": "positive", + "type": "SUPIR_cond_pos", + "links": [ + 8 + ], + "shape": 3 + }, + { + "name": "negative", + "type": "SUPIR_cond_neg", + "links": [ + 9 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "SUPIR_conditioner" + }, + "widgets_values": [ + "high quality, detailed, photograph of an old man", + "bad quality, blurry, messy", + "" + ] + }, + { + "id": 2, + "type": "LoadImage", + "pos": [ + -151, + 127 + ], + "size": { + "0": 441.6546630859375, + "1": 571.2802734375 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 24, + 28, + 33 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "oldman.jpg", + "image" + ] + }, + { + "id": 14, + "type": "ColorMatch", + "pos": [ + 2027, + -81 + ], + "size": { + "0": 315, + "1": 78 + }, + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "image_ref", + "type": "IMAGE", + "link": 29, + "slot_index": 0 + }, + { + "name": "image_target", + "type": "IMAGE", + "link": 26 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 27, + 34 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ColorMatch" + }, + "widgets_values": [ + "mkl" + ] + }, + { + "id": 6, + "type": "SUPIR_model_loader", + "pos": [ + -118, + -88 + ], + "size": [ + 481.16002380371106, + 151.50001403808596 + ], + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "SUPIR_model", + "type": "SUPIRMODEL", + "links": [ + 21 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "SUPIR_VAE", + "type": "SUPIRVAE", + "links": [ + 5, + 30 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "SUPIR_model_loader" + }, + "widgets_values": [ + "SUPIR-v0F.ckpt", + "SDXL\\juggernautXL_v9Rdphoto2Lightning.safetensors", + false, + "auto" + ] + }, + { + "id": 5, + "type": "SUPIR_first_stage", + "pos": [ + 404, + 20 + ], + "size": [ + 248.86013759613047, + 170 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "SUPIR_VAE", + "type": "SUPIRVAE", + "link": 5, + "slot_index": 0 + }, + { + "name": "image", + "type": "IMAGE", + "link": 25 + } + ], + "outputs": [ + { + "name": "SUPIR_VAE", + "type": "SUPIRVAE", + "links": [ + 15 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "denoised_image", + "type": "IMAGE", + "links": [ + 16 + ], + "shape": 3, + "slot_index": 1 + }, + { + "name": "denoised_latents", + "type": "LATENT", + "links": [ + 20 + ], + "shape": 3, + "slot_index": 2 + } + ], + "properties": { + "Node name for S&R": "SUPIR_first_stage" + }, + "widgets_values": [ + true, + 512, + 512, + "auto" + ] + }, + { + "id": 7, + "type": "SUPIR_sample", + "pos": [ + 1386, + -87 + ], + "size": { + "0": 315, + "1": 454 + }, + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "SUPIR_model", + "type": "SUPIRMODEL", + "link": 23, + "slot_index": 0 + }, + { + "name": "latents", + "type": "LATENT", + "link": 17 + }, + { + "name": "positive", + "type": "SUPIR_cond_pos", + "link": 8, + "slot_index": 2 + }, + { + "name": "negative", + "type": "SUPIR_cond_neg", + "link": 9 + } + ], + "outputs": [ + { + "name": "latent", + "type": "LATENT", + "links": [ + 12 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "SUPIR_sample" + }, + "widgets_values": [ + 174277455657960, + "fixed", + 10, + 1.5, + 1.5, + 5, + 1.0030000000000001, + 1, + 0.9, + 0.9500000000000001, + 5, + false, + "RestoreDPMPP2MSampler", + 1024, + 512 + ] + }, + { + "id": 10, + "type": "SUPIR_decode", + "pos": [ + 1733, + -90 + ], + "size": [ + 258.01013759613056, + 102 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "SUPIR_VAE", + "type": "SUPIRVAE", + "link": 32, + "slot_index": 0 + }, + { + "name": "latents", + "type": "LATENT", + "link": 12 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 26 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "SUPIR_decode" + }, + "widgets_values": [ + true, + 512 + ] + }, + { + "id": 20, + "type": "Image Comparer (rgthree)", + "pos": { + "0": 737, + "1": 379, + "2": 0, + "3": 0, + "4": 0, + "5": 0, + "6": 0, + "7": 0, + "8": 0, + "9": 0 + }, + "size": [ + 647.1893005371096, + 642.7652252197267 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "image_a", + "type": "IMAGE", + "link": 33, + "dir": 3 + }, + { + "name": "image_b", + "type": "IMAGE", + "link": 34, + "dir": 3 + } + ], + "outputs": [], + "properties": { + "comparer_mode": "Slide" + }, + "widgets_values": [ + [ + "/view?filename=rgthree.compare._temp_qlpgm_00015_.png&type=temp&subfolder=&rand=0.05076216004847711", + "/view?filename=rgthree.compare._temp_qlpgm_00016_.png&type=temp&subfolder=&rand=0.09029550383385954" + ] + ] + } + ], + "links": [ + [ + 5, + 6, + 1, + 5, + 0, + "SUPIRVAE" + ], + [ + 8, + 9, + 0, + 7, + 2, + "SUPIR_cond_pos" + ], + [ + 9, + 9, + 1, + 7, + 3, + "SUPIR_cond_neg" + ], + [ + 12, + 7, + 0, + 10, + 1, + "LATENT" + ], + [ + 15, + 5, + 0, + 11, + 0, + "SUPIRVAE" + ], + [ + 16, + 5, + 1, + 11, + 1, + "IMAGE" + ], + [ + 17, + 11, + 0, + 7, + 1, + "LATENT" + ], + [ + 20, + 5, + 2, + 9, + 1, + "LATENT" + ], + [ + 21, + 6, + 0, + 12, + 0, + "*" + ], + [ + 22, + 12, + 0, + 9, + 0, + "SUPIRMODEL" + ], + [ + 23, + 12, + 0, + 7, + 0, + "SUPIRMODEL" + ], + [ + 24, + 2, + 0, + 13, + 0, + "IMAGE" + ], + [ + 25, + 13, + 0, + 5, + 1, + "IMAGE" + ], + [ + 26, + 10, + 0, + 14, + 1, + "IMAGE" + ], + [ + 27, + 14, + 0, + 3, + 0, + "IMAGE" + ], + [ + 28, + 2, + 0, + 15, + 0, + "*" + ], + [ + 29, + 16, + 0, + 14, + 0, + "IMAGE" + ], + [ + 30, + 6, + 1, + 17, + 0, + "*" + ], + [ + 32, + 19, + 0, + 10, + 0, + "SUPIRVAE" + ], + [ + 33, + 2, + 0, + 20, + 0, + "IMAGE" + ], + [ + 34, + 14, + 0, + 20, + 1, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": {}, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index 0d5b28e..fdcda48 100644 --- a/nodes.py +++ b/nodes.py @@ -374,4 +374,4 @@ NODE_CLASS_MAPPINGS = { } NODE_DISPLAY_NAME_MAPPINGS = { "SUPIR_Upscale": "SUPIR_Upscale" -} +} \ No newline at end of file diff --git a/nodes_v2.py b/nodes_v2.py new file mode 100644 index 0000000..88b77a5 --- /dev/null +++ b/nodes_v2.py @@ -0,0 +1,707 @@ +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', + ], { + "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,) diff --git a/sgm/modules/attention.py b/sgm/modules/attention.py index b960e82..adf8f32 100644 --- a/sgm/modules/attention.py +++ b/sgm/modules/attention.py @@ -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)): diff --git a/sgm/modules/diffusionmodules/openaimodel.py b/sgm/modules/diffusionmodules/openaimodel.py index 86f0392..b52bbf2 100644 --- a/sgm/modules/diffusionmodules/openaimodel.py +++ b/sgm/modules/diffusionmodules/openaimodel.py @@ -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, diff --git a/sgm/modules/diffusionmodules/sampling.py b/sgm/modules/diffusionmodules/sampling.py index 212e756..40915e4 100644 --- a/sgm/modules/diffusionmodules/sampling.py +++ b/sgm/modules/diffusionmodules/sampling.py @@ -9,6 +9,7 @@ import torch from omegaconf import ListConfig, OmegaConf from tqdm import tqdm from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, get_sigmas_karras +from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, get_sigmas_karras from ...modules.diffusionmodules.sampling_utils import ( get_ancestral_step, linear_multistep_coeff, @@ -560,6 +561,9 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler): 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): @@ -591,10 +595,20 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler): 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 = [ @@ -619,7 +633,8 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler): return x, denoised - def __call__(self, denoiser, x, cond, uc=None, num_steps=None, control_scale=1.0, **kwargs): + 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 ) @@ -647,6 +662,9 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler): 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) @@ -664,12 +682,15 @@ class TiledRestoreDPMPP2MSampler(RestoreDPMPP2MSampler): 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 ) @@ -680,6 +701,7 @@ class TiledRestoreDPMPP2MSampler(RestoreDPMPP2MSampler): 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]) @@ -720,4 +742,5 @@ class TiledRestoreDPMPP2MSampler(RestoreDPMPP2MSampler): x_next /= count x = x_next old_denoised = old_denoised_next - return x \ No newline at end of file + pbar_comfy.update(1) + return x diff --git a/sgm/modules/encoders/modules.py b/sgm/modules/encoders/modules.py index 0a0a4f5..a9185e7 100644 --- a/sgm/modules/encoders/modules.py +++ b/sgm/modules/encoders/modules.py @@ -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"]