From 2cd64cdfcebd5b450738f18675a4916af8d13e36 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 25 Mar 2024 10:17:45 +0200 Subject: [PATCH] cleanup --- SUPIR/models/SUPIR_model_v2.py | 187 +-------------------------------- nodes_v2.py | 25 ++--- 2 files changed, 11 insertions(+), 201 deletions(-) diff --git a/SUPIR/models/SUPIR_model_v2.py b/SUPIR/models/SUPIR_model_v2.py index b510f7c..968fc8e 100644 --- a/SUPIR/models/SUPIR_model_v2.py +++ b/SUPIR/models/SUPIR_model_v2.py @@ -1,16 +1,6 @@ -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): @@ -18,179 +8,4 @@ class SUPIRModel(DiffusionEngine): 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 + self.sampler_config = kwargs['sampler_config'] \ No newline at end of file diff --git a/nodes_v2.py b/nodes_v2.py index d4a300f..94f6af1 100644 --- a/nodes_v2.py +++ b/nodes_v2.py @@ -113,13 +113,13 @@ class SUPIR_encode: print("Encoder using bf16") vae_dtype = 'bf16' else: - print("Encoder using using fp32") + print("Encoder 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}") + print(f"Encoder using {vae_dtype}") dtype = convert_dtype(vae_dtype) @@ -264,13 +264,13 @@ class SUPIR_first_stage: print("Encoder using bf16") vae_dtype = 'bf16' else: - print("Encoder using using fp32") + print("Encoder 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}") + print(f"Encoder using {vae_dtype}") dtype = convert_dtype(vae_dtype) @@ -501,7 +501,6 @@ class SUPIR_conditioner: def condition(self, SUPIR_model, latents, positive_prompt, negative_prompt, captions=""): device = mm.get_torch_device() - mm.unload_all_models() mm.soft_empty_cache() samples = latents["samples"] N, H, W, C = samples.shape @@ -626,13 +625,13 @@ class SUPIR_model_loader: dtype = torch.bfloat16 model_dtype = 'bf16' else: - print("Diffusion using using fp32") + print("Diffusion using fp32") dtype = torch.float32 model_dtype = 'fp32' except: - raise AttributeError("ComfyUI version too old, can't autodecet properly. Set your dtypes manually.") + raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.") else: - print(f"Diffusion using using {diffusion_dtype}") + print(f"Diffusion using {diffusion_dtype}") dtype = convert_dtype(diffusion_dtype) model_dtype = diffusion_dtype @@ -766,21 +765,17 @@ class SUPIR_model_loader_v2: if mm.should_use_fp16(): print("Diffusion using fp16") dtype = torch.float16 - model_dtype = 'fp16' elif mm.should_use_bf16(): print("Diffusion using bf16") dtype = torch.bfloat16 - model_dtype = 'bf16' else: - print("Diffusion using using fp32") + print("Diffusion 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}") + print(f"Diffusion 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 @@ -795,11 +790,11 @@ class SUPIR_model_loader_v2: 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(5) self.model = instantiate_from_config(config.model).cpu() + self.model.model.dtype = dtype pbar.update(1) try: print(f"Attempting to load SDXL model from node inputs")