Add masked guidance/sampling feature
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user