diff --git a/README.md b/README.md index b509fdb..041e218 100644 --- a/README.md +++ b/README.md @@ -32,6 +32,13 @@ The main disadvantage compared to the alternatives I mentioned is that it is rel ### `DiffuseHighSampler` +This is the main DiffuseHigh sampler node. + +I recommend expanding the YAML Parameters section and at least skimming through it so you can see what your options are. Most advanced features are controlled there - you can do stuff like switch VAEs, upscale models or other parameters per iteration which can be a very powerful tool. + +**Input Parameters**: You can connect stuff like VAEs, upscale models and masks using this input. + +**Mask Usage**: Masks can be connected via `input_params_opt`. There are currently two ways they can be used: as a global mask or to mask the guidance. If the mask has no name, it is by default a global mask. If you name it `guidance` then it will be treated as a guidance mask. You don't have to stick to those names, `mask_name` and `guidance_mask_name` in the YAML parameters can be used to control what masks are used. Where global masks are set, the model is allowed to change the image - where they aren't set, it will be the reference image. Guidance masks apply guidance where the mask is set and you get the model's normal prediction otherwise. Non-binary masks work the way you'd expect: you'll get a blend based on the mask strength in a particular area. #### Inputs @@ -178,6 +185,8 @@ reference_sampler_name: "reference" guidance_sampler_name: "guidance" custom_noise_name: "" restart_custom_noise_name: "restart" +mask_name: "" +guidance_mask_name: "guidance" # Either null or an object. # Allows overriding the sigma used for highres steps. See description below. diff --git a/changelog.md b/changelog.md index 1021115..b8e08ec 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,10 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20241116 + +* Added the ability to mask both guidance and global changes. See the README section on masks. + ## 20241031 * Added the `DiffuseHighParam` node and the ability to connect multiple VAEs, upscale models, noise generators and samplers as well as switch between them per iteration. diff --git a/py/config.py b/py/config.py index 2bede02..db91dc1 100644 --- a/py/config.py +++ b/py/config.py @@ -5,10 +5,7 @@ from typing import Any from comfy.samplers import ksampler from pytorch_wavelets import DTCWTForward, DTCWTInverse, DWTForward, DWTInverse -from .tensor_image_ops import ( - BLENDING_MODES, - Sharpen, -) +from .tensor_image_ops import Sharpen from .upscale import Upscale from .utils import fallback from .vae import VAEHelper @@ -31,6 +28,8 @@ class Config: "enable_cache_clearing", "enable_gc", "fadeout_factor", + "guidance_mask_blend_mode", + "guidance_mask_name", "guidance_factor", "guidance_mode", "guidance_restart_s_noise", @@ -38,6 +37,8 @@ class Config: "guidance_sampler_name", "guidance_steps", "highres_sigmas_name", + "mask_blend_mode", + "mask_name", "reference_image_name", "reference_sampler_name", "reference_wavelet_multiplier", @@ -66,7 +67,6 @@ class Config: _dict_exclude_keys = { # noqa: RUF012 "as_dict", - "blend_function", "dwt", "get_iteration_config", "guidance_sampler", @@ -137,6 +137,10 @@ class Config: guidance_sampler_name="guidance", custom_noise_name="", restart_custom_noise_name="restart", + mask_name="", + guidance_mask_name="guidance", + mask_blend_mode="lerp", + guidance_mask_blend_mode="lerp", ): self.vae_name = vae_name self.upscale_model_name = upscale_model_name @@ -147,6 +151,8 @@ class Config: self.sampler_name = sampler_name self.custom_noise_name = custom_noise_name self.restart_custom_noise_name = restart_custom_noise_name + self.mask_name = mask_name + self.guidance_mask_name = guidance_mask_name sampler = _params.get_item("sampler", name=sampler_name) guidance_sampler = _params.get_item("sampler", name=guidance_sampler_name) @@ -158,6 +164,9 @@ class Config: None if highres_sigmas is None else highres_sigmas.detach().clone() ) + self.mask_blend_mode = mask_blend_mode + self.guidance_mask_blend_mode = guidance_mask_blend_mode + sampler = fallback( sampler, lambda: ksampler("euler"), @@ -235,7 +244,6 @@ class Config: if blend_by_mode not in {"image", "latent", "wavelet"}: raise ValueError("Bad blend_by_mode: must be one of image, latent, wavelet") self.blend_by_mode = blend_by_mode - self.blend_function = BLENDING_MODES[blend_mode] self.enable_gc = enable_gc self.enable_cache_clearing = enable_cache_clearing self.chunked_sampling = chunked_sampling diff --git a/py/guided_model.py b/py/guided_model.py index 2e76d15..cccd532 100644 --- a/py/guided_model.py +++ b/py/guided_model.py @@ -2,6 +2,7 @@ from __future__ import annotations from typing import TYPE_CHECKING +from .tensor_image_ops import BLENDING_MODES from .utils import ensure_model, sigma_to_float if TYPE_CHECKING: @@ -101,8 +102,15 @@ class GuidedModel: sigma * sigma_offset if sigma_offset != 1 else sigma, **extra_args, ) - return ( + result = ( denoised if guidance_step is None else dhso.apply_guidance(guidance_step, denoised) ) + if dhso.mask is not None: + result = BLENDING_MODES[dhso.mask_blend_mode]( + dhso.orig_latent, + result, + dhso.mask, + ) + return result diff --git a/py/nodes.py b/py/nodes.py index 38aab04..434750d 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -232,6 +232,7 @@ class DiffuseHighParamNode: "*", whitelist={ "IMAGE", + "MASK", "OCS_NOISE", "SAMPLER", "SIGMAS", @@ -248,6 +249,7 @@ class DiffuseHighParamNode: "image": lambda v: isinstance(v, torch.Tensor) and v.ndim == 4, "sigmas": lambda v: isinstance(v, torch.Tensor) and v.ndim == 1 and len(v) >= 2, "custom_noise": lambda v: hasattr(v, "make_noise_sampler"), + "mask": lambda v: isinstance(v, torch.Tensor) and v.ndim in {2, 3}, } @classmethod diff --git a/py/sampler.py b/py/sampler.py index 67571c2..a0f08c7 100644 --- a/py/sampler.py +++ b/py/sampler.py @@ -6,6 +6,8 @@ import sys from typing import Any import torch +from comfy.model_management import device_supports_non_blocking +from comfy.utils import reshape_mask from tqdm import tqdm from tqdm.auto import trange @@ -13,6 +15,7 @@ from .config import Config from .guided_model import GuidedModel from .schedule import Schedule from .tensor_image_ops import ( + BLENDING_MODES, blend_wavelets, scale_wavelets, ) @@ -32,6 +35,7 @@ class DiffuseHighSampler: _params, **kwargs: dict[str, Any], ): + self.non_blocking = device_supports_non_blocking(initial_x.device) self.s_in = initial_x.new_ones((initial_x.shape[0],)) self.initial_x = initial_x self.callback = callback @@ -52,6 +56,7 @@ class DiffuseHighSampler: self.highres_sigmas: None | torch.FloatTensor = None self.reference_image: None | torch.FloatTensor = None self.guidance_waves = None + self.mask = self.guidance_mask = None self.seed_offset = 0 self.seed = self.extra_args.get("seed") if self.cfg.seed_rng: @@ -90,10 +95,11 @@ class DiffuseHighSampler: return denoised if self.guidance_mode not in {"image", "latent"}: raise ValueError("Bad guidance mode") + blend_function = BLENDING_MODES[self.blend_mode] if self.guidance_mode == "image": dn_img = ( self.vae.decode(denoised, disable_pbar=self.disable_pbar) - .to(denoised) + .to(denoised, non_blocking=self.non_blocking) .movedim(-1, 1) ) denoised_waves = self.dwt(dn_img) @@ -117,14 +123,14 @@ class DiffuseHighSampler: denoised_waves_orig, coeffs, mix_scale, - self.blend_function, + blend_function, ) result = self.idwt(coeffs) if self.guidance_mode == "image": if self.blend_by_mode == "image": - result = self.blend_function( + result = blend_function( dn_img, - result.to(dn_img), + result.to(dn_img, non_blocking=self.non_blocking), dn_img.new_full((1,), mix_scale), ).clamp_(0, 1) result = self.vae.encode( @@ -133,12 +139,20 @@ class DiffuseHighSampler: disable_pbar=self.disable_pbar, ) if self.blend_by_mode != "latent": - return result.to(denoised) - return self.blend_function( - denoised, - result.to(denoised), - denoised.new_full((1,), mix_scale), - ) + result = result.to(denoised, non_blocking=self.non_blocking) + else: + result = blend_function( + denoised, + result.to(denoised, non_blocking=self.non_blocking), + denoised.new_full((1,), mix_scale), + ) + if self.guidance_mask is not None: + result = BLENDING_MODES[self.guidance_mask_blend_mode]( + denoised, + result, + self.guidance_mask, + ) + return result def add_restart_noise( self, @@ -374,6 +388,25 @@ class DiffuseHighSampler: name=self.cfg.reference_image_name, ) + def update_masks(self, latent): + self.curr_mask = self.guidance_mask = None + for name, is_guidance in ( + (self.mask_name, False), + (self.guidance_mask_name, True), + ): + mask = self._params.get_item("mask", name=name) + if mask is None: + continue + mask = reshape_mask(mask.detach().clone(), latent.shape).to( + latent, + non_blocking=self.non_blocking, + ) + if is_guidance: + self.guidance_mask = mask + else: + self.mask = mask + self.orig_latent = latent.detach().clone() if self.mask is not None else None + def get_noise_sampler(self, x, sigmas, *, name=""): custom_noise = self._params.get_item("custom_noise", name=name) if custom_noise is None: @@ -462,10 +495,13 @@ class DiffuseHighSampler: x_new = self.vae.encode( self.reference_image, disable_pbar=self.disable_pbar, - ).to(self.initial_x) + ).to(self.initial_x, non_blocking=self.non_blocking) + self.update_masks(x_new) if self.guidance_mode == "image": self.guidance_waves = self.dwt( - self.reference_image.clone().movedim(-1, 1).to(self.initial_x), + self.reference_image.clone() + .movedim(-1, 1) + .to(self.initial_x, non_blocking=self.non_blocking), ) elif self.guidance_mode == "latent": self.guidance_waves = self.dwt(x_new) @@ -487,6 +523,7 @@ class DiffuseHighSampler: x_new = self.run_steps(x_new, sigmas=self.highres_sigmas) if iteration == self.iterations - 1: break + self.mask = self.guidance_mask = None self.reference_image = self.vae.decode( x_new, disable_pbar=self.disable_pbar,