diff --git a/.gitignore b/.gitignore index 82f9275..aabefd5 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,6 @@ +blehconfig.json +blehconfig.yaml + # Byte-compiled / optimized / DLL files __pycache__/ *.py[cod] @@ -106,10 +109,8 @@ ipython_config.py #pdm.lock # pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it # in version control. -# https://pdm.fming.dev/latest/usage/project/#working-with-version-control +# https://pdm.fming.dev/#use-with-ide .pdm.toml -.pdm-python -.pdm-build/ # PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm __pypackages__/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..9b3a548 --- /dev/null +++ b/__init__.py @@ -0,0 +1,5 @@ +from .py import nodes + +NODE_CLASS_MAPPINGS = { + "DiffuseHighSampler": nodes.DiffuseHighSamplerNode, +} diff --git a/py/__init__.py b/py/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/py/external.py b/py/external.py new file mode 100644 index 0000000..4199cff --- /dev/null +++ b/py/external.py @@ -0,0 +1,9 @@ +import contextlib +import importlib + +EXTERNAL = {} + +with contextlib.suppress(ImportError): + EXTERNAL["tiled_diffusion"] = importlib.import_module( + "custom_nodes.ComfyUI-TiledDiffusion", + ) diff --git a/py/nodes.py b/py/nodes.py new file mode 100644 index 0000000..fe7020a --- /dev/null +++ b/py/nodes.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import yaml +from comfy.samplers import KSAMPLER + +from .sampler import diffusehigh_sampler + + +class DiffuseHighSamplerNode: + RETURN_TYPES = ("SAMPLER",) + FUNCTION = "go" + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + "highres_sigmas": ("SIGMAS",), + "guidance_steps": ("INT", {"default": 5, "min": 0}), + "guidance_mode": ( + ( + "image", + "latent", + ), + ), + "guidance_factor": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + }, + ), + "fadeout_factor": ("FLOAT", {"default": 0.0}), + "scale_factor": ("FLOAT", {"default": 2.0}), + "renoise_factor": ("FLOAT", {"default": 1.0}), + "iterations": ("INT", {"default": 1, "min": 0}), + "sampler": ("SAMPLER",), + "vae_mode": ( + ( + "taesd", + "normal", + "tiled", + "tiled_diffusion", + ), + ), + }, + "optional": { + "reference_image_opt": ("IMAGE",), + "guidance_sampler_opt": ("SAMPLER",), + "reference_sampler_opt": ("SAMPLER",), + "vae_opt": ("VAE",), + "upscale_model_opt": ("UPSCALE_MODEL",), + "yaml_parameters": ( + "STRING", + { + "tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. This input can be converted into a multiline text widget. Note: When specifying paramaters this way, there is very little error checking.", + "dynamicPrompts": False, + "multiline": True, + "defaultInput": True, + }, + ), + }, + } + + @classmethod + def go(cls, yaml_parameters: None | str = None, **kwargs: dict) -> tuple[KSAMPLER]: + if yaml_parameters: + extra_params = yaml.safe_load(yaml_parameters) + if extra_params is None: + pass + elif not isinstance(extra_params, dict): + raise ValueError( + "DiffuseHighSampler: yaml_parameters must either be null or an object", + ) + else: + kwargs |= extra_params + return ( + KSAMPLER( + diffusehigh_sampler, + extra_options={ + "diffusehigh_options": kwargs, + }, + ), + ) diff --git a/py/sampler.py b/py/sampler.py new file mode 100644 index 0000000..ccd6daf --- /dev/null +++ b/py/sampler.py @@ -0,0 +1,356 @@ +from __future__ import annotations + +import math + +import PIL.Image as PILImage +import torch +import torchvision +from comfy_extras.nodes_upscale_model import ImageUpscaleWithModel +from pytorch_wavelets import DTCWTForward, DTCWTInverse, DWTForward, DWTInverse +from tqdm.auto import trange + +from .utils import ( + ensure_model, + pilimgbatch_to_torch, + torch_to_pilimgbatch, +) +from .vae import VAEHelper + + +def gaussian_blur_image_sharpening(image, kernel_size=3, sigma=(0.1, 2.0), alpha=1): + gaussian_blur = torchvision.transforms.GaussianBlur( + kernel_size=kernel_size, + sigma=sigma, + ) + image_blurred = gaussian_blur(image) + return (alpha + 1) * image - alpha * image_blurred + + +class DiffuseHighSampler: + def __init__( + self, + model, + initial_x, + sigmas, + *, + callback, + extra_args, + disable_pbar, + highres_sigmas, + sampler, + guidance_steps=5, + guidance_mode="image", + guidance_factor=1.0, + guidance_restart=0, + guidance_restart_s_noise=1.0, + fadeout_factor=0.0, + scale_factor=2.0, + renoise_factor=1.0, + iterations=1, + vae_mode="normal", + dwt_level=1, + dwt_wave="db4", + dwt_mode="symmetric", + dwt_flip_filters=False, + dtcwt_mode=False, + dtcwt_biort="near_sym_a", + dtcwt_qshift="qshift_a", + reference_wavelet_multiplier=1.0, + denoised_wavelet_multiplier=1.0, + sharpen_reference=True, + sharpen_kernel_size=3, + sharpen_sigma=(0.1, 2.0), + sharpen_alpha=1.0, + resample_mode="bicubic", + rescale_increment=64, + guidance_sampler_opt=None, + reference_sampler_opt=None, + reference_image_opt=None, + vae_opt=None, + upscale_model_opt=None, + ): + self.s_in = initial_x.new_ones((initial_x.shape[0],)) + self.initial_x = initial_x + self.callback = callback + self.disable_pbar = disable_pbar + self.sigmas = sigmas + self.extra_args = extra_args if extra_args is not None else {} + self.model = model + self.latent_format = model.inner_model.inner_model.latent_format + self.fadeout_factor = fadeout_factor + self.scale_factor = scale_factor + self.guidance_factor = guidance_factor + self.renoise_factor = renoise_factor + self.iterations = iterations + self.highres_sigmas = highres_sigmas.clone().to(sigmas) + self.guidance_steps = guidance_steps + self.guidance_mode = guidance_mode + self.guidance_restart = guidance_restart + self.guidance_restart_s_noise = guidance_restart_s_noise + self.sampler = sampler + self.guidance_sampler = guidance_sampler_opt or sampler + self.reference_sampler = reference_sampler_opt or sampler + self.vae = VAEHelper( + vae_mode, + self.latent_format, + device=initial_x.device, + dtype=initial_x.dtype, + vae=vae_opt, + ) + self.reference_image = reference_image_opt + self.sharpen_reference = sharpen_reference + self.sharpen_kernel_size = sharpen_kernel_size + self.sharpen_sigma = sharpen_sigma + self.sharpen_alpha = sharpen_alpha + self.resample_mode = getattr(PILImage, resample_mode.upper()) + self.rescale_increment = self.scale_dim( + max(8, rescale_increment), + 1, + increment=8, + ) + if dtcwt_mode: + self.dwt = DTCWTForward( + J=dwt_level, + mode=dwt_mode, + biort=dtcwt_biort, + qshift=dtcwt_qshift, + ).to( + initial_x.device, + ) + self.idwt = DTCWTInverse( + mode=dwt_mode, + biort=dtcwt_biort, + qshift=dtcwt_qshift, + ).to(initial_x.device) + else: + self.dwt = DWTForward(J=dwt_level, wave=dwt_wave, mode=dwt_mode).to( + initial_x.device, + ) + self.idwt = DWTInverse(wave=dwt_wave, mode=dwt_mode).to(initial_x.device) + self.dwt_flip_filters = dwt_flip_filters + self.reference_wavelet_multiplier = reference_wavelet_multiplier + self.denoised_wavelet_multiplier = denoised_wavelet_multiplier + self.guidance_waves = None + self.guidance_latent = None + self.upscale_model = upscale_model_opt + + def call_model(self, x, sigma): + return self.model(x, sigma * self.s_in, **self.extra_args) + + def do_callback(self, idx, x, sigma, denoised): + if self.callback is None: + return + self.callback({ + "i": idx, + "x": x, + "sigma": sigma, + "sigma_hat": sigma, + "denoised": denoised, + }) + + def apply_guidance(self, idx, denoised): + if self.guidance_waves is None or idx >= self.guidance_steps: + return denoised + mix_scale = ( + self.guidance_factor + - ((self.guidance_factor / (self.guidance_steps + 1)) * idx) + * self.fadeout_factor + ) + if mix_scale == 0: + return denoised + print("GUIDANCE APPLY", idx) + if self.guidance_mode not in {"image", "latent"}: + raise ValueError("ohno") + if self.guidance_mode == "image": + dn_img = self.vae.decode(denoised).to(denoised).movedim(-1, 1) + print("DN_IMG", dn_img.shape) + denoised_waves = self.dwt(dn_img) + del dn_img + elif self.guidance_mode == "latent": + denoised_waves = self.dwt(denoised) + if self.denoised_wavelet_multiplier != 1: + denoised_waves = ( + denoised_waves[0] * self.denoised_wavelet_multiplier, + tuple(t * self.denoised_wavelet_multiplier for t in denoised_waves[1]), + ) + coeffs = ( + (self.guidance_waves[0], denoised_waves[1]) + if not self.dwt_flip_filters + else (denoised_waves[0], self.guidance_waves[1]) + ) + result = self.idwt(coeffs) + if self.guidance_mode == "image": + result = self.vae.encode(result.cpu(), fix_dims=True) + + print("GUIDE OUT", denoised.shape, result.shape, mix_scale) + return torch.lerp(denoised, result.to(denoised), mix_scale) + + def run_steps(self, *, x=None, sigmas=None): + x = self.initial_x if x is None else x + sigmas = self.sigmas if sigmas is None else sigmas + guidance_sigmas = sigmas[: self.guidance_steps + 1] + normal_sigmas = sigmas[self.guidance_steps :] + step_idx = 0 + model = self.model + + def model_wrapper(x, sigma, **extra_args: dict): + nonlocal step_idx + ensure_model(model) + denoised = model(x, sigma, **extra_args) + return self.apply_guidance(step_idx, denoised) + + for k in ( + "inner_model", + "sigmas", + ): + if hasattr(model, k): + setattr(model_wrapper, k, getattr(model, k)) + + for repidx in range(self.guidance_restart + 1): + if repidx > 0: + noise_factor = ( + guidance_sigmas[0] ** 2 - guidance_sigmas[-1] ** 2 + ) ** 0.5 + x = x + torch.randn_like(x) * ( + noise_factor * self.guidance_restart_s_noise + ) + for idx in range(len(guidance_sigmas) - 1): + step_idx = idx + x = self.run_sampler( + x, + guidance_sigmas[idx : idx + 2], + model=model_wrapper, + sampler=self.guidance_sampler, + ) + if len(normal_sigmas) > 1: + ensure_model(model) + x = self.run_sampler(x, normal_sigmas) + return x + + @staticmethod + def scale_dim(n, factor, *, increment=64) -> int: + return math.ceil((n * factor) / increment) * increment + + def upscale(self, imgbatch): + _batch, height, width, _channels = imgbatch.shape + target_height = self.scale_dim( + height, + self.scale_factor, + increment=self.rescale_increment, + ) + target_width = self.scale_dim( + width, + self.scale_factor, + increment=self.rescale_increment, + ) + print(f">> UPSCALE: {width}x{height} -> {target_width}x{target_height}") + if (target_height, target_width) == (height, width): + return imgbatch + if self.upscale_model is not None: + print("** Upscaling with model") + imgbatch = ImageUpscaleWithModel().upscale(self.upscale_model, imgbatch)[0] + if imgbatch.shape[1:3] == (target_height, target_width): + return imgbatch + print( + f"** PIL upscale {imgbatch.shape[2]}x{imgbatch.shape[1]} -> {target_width}x{target_height}", + ) + ref_imgbatch = torch_to_pilimgbatch(self.reference_image) + return pilimgbatch_to_torch( + tuple( + i.resize((target_width, target_height), resample=self.resample_mode) + for i in ref_imgbatch + ), + ) + + def run_sampler(self, x, sigmas, *, model=None, sampler=None): + model = model or self.model + sampler = sampler or self.sampler + return sampler.sampler_function( + model, + x, + sigmas, + callback=self.callback, + extra_args=self.extra_args.copy(), + disable=self.disable_pbar, + **sampler.extra_options, + ) + + def __call__(self): + if self.reference_image is None: + x_lr = self.run_sampler( + self.initial_x, + self.sigmas, + sampler=self.reference_sampler, + ) + if self.iterations < 1: + return x_lr + self.reference_image = self.vae.decode(x_lr) + elif self.iterations < 1: + return self.vae.encode(self.reference_image) + for iteration in trange(self.iterations, disable=self.disable_pbar): + print( + f"\nIT({iteration}): shp={self.reference_image.shape}, min={self.reference_image.min()}, max={self.reference_image.max()}", + ) + img_hr = self.upscale(self.reference_image) + print("IMG_HR", img_hr.shape) + if self.sharpen_reference: + img_hr = gaussian_blur_image_sharpening( + img_hr.movedim(-1, 1), + kernel_size=self.sharpen_kernel_size, + sigma=self.sharpen_sigma, + alpha=self.sharpen_alpha, + ).movedim(1, -1) + self.reference_image = img_hr + x_new = self.vae.encode(self.reference_image).to(self.initial_x) + self.guidance_latent = x_new.clone() + if self.guidance_mode == "image": + print("REF IMG", self.reference_image.shape) + self.guidance_waves = self.dwt( + self.reference_image.clone().movedim(-1, 1).to(self.initial_x), + ) + elif self.guidance_mode == "latent": + self.guidance_waves = self.dwt(self.guidance_latent) + if self.reference_wavelet_multiplier != 1: + self.guidance_waves = ( + self.guidance_waves[0] * self.reference_wavelet_multiplier, + tuple( + t * self.reference_wavelet_multiplier + for t in self.guidance_waves[1] + ), + ) + else: + raise ValueError("ohno") + x_noise = torch.randn_like(x_new) + x_new = x_new + x_noise * (self.highres_sigmas[0] * self.renoise_factor) + # x_new = self.model.inner_model.inner_model.model_sampling.noise_scaling( + # self.highres_sigmas[0] * self.renoise_factor, + # x_noise, + # x_new, + # ) + result = self.run_steps(x=x_new, sigmas=self.highres_sigmas) + if iteration == self.iterations - 1: + break + self.reference_image = self.vae.decode(result) + return result + + +def diffusehigh_sampler( + model, + x, + sigmas, + *, + diffusehigh_options, + disable=None, + extra_args=None, + callback=None, +): + sampler = DiffuseHighSampler( + model, + x, + sigmas, + disable_pbar=disable, + callback=callback, + extra_args=extra_args, + **diffusehigh_options, + ) + return sampler() diff --git a/py/utils.py b/py/utils.py new file mode 100644 index 0000000..55a7498 --- /dev/null +++ b/py/utils.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from typing import Sequence + +import numpy as np +import torch +from comfy import model_management +from PIL import Image as PILImage + + +def pilimgbatch_to_torch( + imgbatch: Sequence[PILImage, ...] | torch.Tensor, +) -> torch.Tensor: + if isinstance(imgbatch, torch.Tensor): + print("pibtt: skip", imgbatch.shape, imgbatch.min(), imgbatch.max()) + return imgbatch + npi = np.stack( + tuple(np.array(i).astype(np.float32) / 255.0 for i in imgbatch), + axis=0, + ) + return torch.from_numpy(npi) + return torch.from_numpy(npi.transpose(0, 3, 1, 2)) + + +def torch_to_pilimgbatch(t: torch.Tensor) -> tuple[PILImage, ...]: + return tuple( + PILImage.fromarray( + np.clip((255.0 * i).cpu().numpy(), 0, 255).astype(np.uint8), + ) + for i in t + ) + + +def ensure_model(model): + mp = model.inner_model.model_patcher + if model_management.LoadedModel(mp) in model_management.current_loaded_models: + return + model_management.load_models_gpu((mp,)) diff --git a/py/vae.py b/py/vae.py new file mode 100644 index 0000000..5eebd1b --- /dev/null +++ b/py/vae.py @@ -0,0 +1,189 @@ +from __future__ import annotations + +from enum import Enum, auto + +import folder_paths +import torch +from comfy.taesd.taesd import TAESD + +from .external import EXTERNAL + +tiled_diffusion = EXTERNAL.get("tiled_diffusion") + + +class VAEMode(Enum): + TAESD = auto() + NORMAL = auto() + TILED = auto() + TILED_DIFFUSION = auto() + + +class VAEHelper: + def __init__( + self, + mode: VAEMode | str, + latent_format, + *, + device=None, + dtype=None, + vae=None, + vae_encode_kwargs=None, + vae_decode_kwargs=None, + ): + if isinstance(mode, str): + mode = VAEMode.__members__[mode.upper()] + if mode == VAEMode.TILED_DIFFUSION: + if tiled_diffusion is None: + raise ValueError( + "Cannot use tiled_diffusion VAE mode without ComfyUI-TiledDiffusion!", + ) + self.td_encode_default_kwargs, self.td_decode_default_kwargs = ( + { + k: v[1]["default"] + for k, v in td_node.INPUT_TYPES() + .get( + "required", + {}, + ) + .items() + if k not in {"pixels", "samples", "vae"} + and len(v) == 2 + and isinstance(v[1], dict) + and "default" in v[1] + } + for td_node in ( + tiled_diffusion.tiled_vae.VAEEncodeTiled_TiledDiffusion, + tiled_diffusion.tiled_vae.VAEDecodeTiled_TiledDiffusion, + ) + ) + if mode != VAEMode.TAESD and vae is None: + raise ValueError("Must pass a VAE when using non-TAESD VAE modes!") + self.mode = mode + self.latent_format = latent_format + self.device = device + self.dtype = dtype + self.vae = vae + self.vae_encode_kwargs = {} if vae_encode_kwargs is None else vae_encode_kwargs + self.vae_decode_kwargs = {} if vae_decode_kwargs is None else vae_decode_kwargs + vae_handlers = { + VAEMode.TAESD: (self.encode_taesd, self.decode_taesd), + VAEMode.NORMAL: (self.encode_vae, self.decode_vae), + VAEMode.TILED: ( + self.encode_vae_tiled, + self.decode_vae_tiled, + ), + VAEMode.TILED_DIFFUSION: ( + self.encode_vae_tiled_diffusion, + self.decode_vae_tiled_diffusion, + ), + } + self.encode_fun, self.decode_fun = vae_handlers[mode] + + def encode(self, imgbatch, *, fix_dims=False): + if fix_dims: + imgbatch = imgbatch.moveaxis(1, -1) + # print("ENCODING", imgbatch.min(), imgbatch.max()) + result = self.encode_fun(imgbatch[..., :3]) + if self.mode != VAEMode.TAESD: + # print("ENCODED(raw):", result.min(), result.max()) + result = self.latent_format.process_in(result) + # print("ENCODED", result.shape, result.min(), result.max()) + return result + + def decode(self, latent, *, skip_process_out=False): + if self.mode != VAEMode.TAESD and not skip_process_out: + latent = self.latent_format.process_out(latent) + # print("DECODING", latent.min(), latent.max()) + return self.decode_fun(latent) + # print("DECODED", result.shape, result.min(), result.max()) + + def encode_taesd(self, imgbatch): + dummy = torch.zeros((), device=self.device, dtype=self.dtype) + return OCSTAESD.encode(self.latent_format, imgbatch, dummy) + + def decode_taesd(self, latent): + return OCSTAESD.decode(self.latent_format, latent) + + def encode_vae(self, imgbatch): + # print("VAE ENC", imgbatch.shape) + return self.vae.encode(imgbatch, **self.vae_encode_kwargs) + + def decode_vae(self, latent): + return self.vae.decode(latent, **self.vae_decode_kwargs) + + def encode_vae_tiled(self, imgbatch): + return self.vae.encode_tiled(imgbatch, **self.vae_encode_kwargs) + + def decode_vae_tiled(self, latent): + return self.vae.decode_tiled(latent, **self.vae_decode_kwargs) + + def encode_vae_tiled_diffusion(self, imgbatch): + kwargs = self.td_encode_default_kwargs | self.vae_encode_kwargs + return tiled_diffusion.tiled_vae.VAEEncodeTiled_TiledDiffusion().process( + pixels=imgbatch, + vae=self.vae, + **kwargs, + )[0]["samples"] + + def decode_vae_tiled_diffusion(self, latent): + kwargs = self.td_decode_default_kwargs | self.vae_decode_kwargs + return tiled_diffusion.tiled_vae.VAEDecodeTiled_TiledDiffusion().process( + samples={"samples": latent}, + vae=self.vae, + **kwargs, + )[0] + + +class OCSTAESD: + @classmethod + def get_encoder_name(cls, latent_format): + result = latent_format.taesd_decoder_name + if not result.endswith("_decoder"): + msg = f"Could not determine TAESD encoder name from {result!r}" + raise RuntimeError( + msg, + ) + return f"{result[:-7]}encoder" + + @classmethod + def get_taesd_path(cls, name): + taesd_path = next( + ( + fn + for fn in folder_paths.get_filename_list("vae_approx") + if fn.startswith(name) + ), + "", + ) + if not taesd_path: + msg = f"Could not get TAESD path for {name!r}" + raise RuntimeError(msg) + return folder_paths.get_full_path("vae_approx", taesd_path) + + @classmethod + def decode(cls, latent_format, latent): + filename = cls.get_taesd_path(latent_format.taesd_decoder_name) + model = TAESD( + decoder_path=filename, + latent_channels=latent_format.latent_channels, + ).to(latent.device) + return ( + model.taesd_decoder( + (latent - model.vae_shift).mul_(model.vae_scale), + ) + .clamp_(0, 1) + .movedim(1, -1) + ) + + @classmethod + def encode(cls, latent_format, imgbatch, latent) -> torch.Tensor: + filename = cls.get_taesd_path(cls.get_encoder_name(latent_format)) + model = TAESD( + encoder_path=filename, + latent_channels=latent_format.latent_channels, + ).to(device=latent.device) + return ( + model.taesd_encoder(imgbatch.to(latent.device).moveaxis(-1, 1)) + .div_(model.vae_scale) + .add_(model.vae_shift) + ) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..159e71c --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +pywavelets +pytorch-wavelets diff --git a/ruff.toml b/ruff.toml new file mode 100644 index 0000000..7cd2146 --- /dev/null +++ b/ruff.toml @@ -0,0 +1,43 @@ +[lint] +ignore = [ + "ANN001", + "ANN101", + "ANN102", + "ANN201", + "ANN202", + "ANN204", + "ANN206", + "C901", + "CPY001", + "D100", + "D101", + "D102", + "D103", + "D104", + "D105", + "D107", + "D211", + "D213", + "E402", + "E501", + "EM101", + "ERA001", + "F403", + "F405", + "FBT001", + "FBT002", + "G004", + "PLR0912", + "PLR0913", + "PLR0915", + "PLR2004", + "PLR6104", + "T201", + "TD001", + "TD002", + "TD003", + "TRY003", + "N802", + "N999", +] +select = ["ALL"]