Add masked guidance/sampling feature

This commit is contained in:
blepping
2024-11-16 10:44:33 -07:00
parent 6355727913
commit 9c41683e3a
6 changed files with 87 additions and 19 deletions
+9
View File
@@ -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.
+4
View File
@@ -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.
+14 -6
View File
@@ -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
+9 -1
View File
@@ -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
+2
View File
@@ -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
+49 -12
View File
@@ -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,