Implement better chunked sampling

Remove sigma_offset YAML parameter
Add sigma_dishonesty_factor[_guidance] YAML parameters
Internal cleanups and refactoring
Minor documentation improvements
This commit is contained in:
blepping
2024-10-27 10:35:00 -06:00
parent 3dfaf18353
commit 7ced6b1626
6 changed files with 270 additions and 84 deletions
+17 -23
View File
@@ -7,7 +7,7 @@ This is a best-effort attempt at implementation. If you experience poor results,
## Current Status
Alpha - early implementation. Many rough edges but the core functionality is there. Mainly targetted at advanced users who can deal with some weird stuff and frequent workflow-breaking changes.
Beta - lightly tested but the main features are in place. Mainly targeted at advanced users who can deal with some weird stuff and frequent workflow-breaking changes.
See the [changelog](changelog.md) for recent user visible changes.
@@ -18,7 +18,7 @@ See the [changelog](changelog.md) for recent user visible changes.
* Using VAE or upscale models may result in the main model getting repeatedly unloaded/reloaded. Try using `latent` as the `guidance_mode`. If you actually have enough VRAM, maybe disabling smart memory (via ComfyUI commandline parameter) would help.
* Brownian noise-based (AKA SDE) samplers may be a bit weird here, there is a workaround in place but it might not be enough. Also don't use with prompt-control's PCSplitSampling stuff.
**Rectified Flow models note**: Should now work with RF models. SD3.5 apparently cannot handle high res images (even img2img) at all, so I don't recommend trying that. Flux seems to work pretty well. `image` guidance mode seems noticeably better than `latent` for Flux (based on my very limited testing) although it is slow. I haven't tested SD3.0 or other RF models, jank DiffuseHigh should handle them correctly but whether the results are actually decent I really couldn't say.
**Rectified Flow models note**: Should now work with RF models. SD3.5 apparently cannot handle high res images (even img2img) at all, so I don't recommend trying that. Flux seems to work pretty well. `image` guidance mode seems noticeably better than `latent` for Flux (based on my very limited testing) although it is slow. I haven't tested SD3.0 or other RF models, jank DiffuseHigh should handle them correctly but whether the results are actually decent I really couldn't say. Using `guidance_restart` probably won't work correctly.
## Description
@@ -37,7 +37,7 @@ The main disadvantage compared to the alternatives I mentioned is that it is rel
* `highres_sigmas`: Optional: Sigmas used for everything other than the initial reference image. **Note**: Should be around 0.3-0.5 denoise. You won't get good results connecting something like `KarrasScheduler` here without splitting the sigmas. If not specified, will use the last 15 steps of a 50 step Karras schedule like the official implementation.
* `sampler`: Optional: Default sampler used for steps. If not specified the sampler will default to non-ancestral Euler.
* `reference_image_opt`: Optional: Image used for the initial pass. If not connected, a low-res initial reference will be generated using the schedule from the normal sigmas (i.e. the sigmas attached to `SamplerCustom` or whatever actual sampler node you're using).
* `guidance_sampler_opt`: Optional: Sampler used for guidance steps. If not specified, will fallback to the base sampler. Note: The sampler is called on individual steps, samplers that keep history will not work well here.
* `guidance_sampler_opt`: Optional: Sampler used for guidance steps. If not specified, will fallback to the base sampler.
* `reference_sampler_opt`: Optional: Sampler used to generate the initial low-resolution reference. Only used if reference_image_opt is not connected.
* `vae_opt`: Optional when vae_mode is set to `taesd`, otherwise this is the VAE that will be used for encoding/decoding images. If using TAESD, you will require the corresponding encoder (which I believe ComfyUI does not install by default). TAESD models go in `models/vae_approx`, you can find them here: https://github.com/madebyollin/taesd
* `upscale_model_opt`: Optional: Model used for upscaling. When not attached, simple image scaling will be used. Regardless, the image will be scaled to match the size expected based on `scale_factor`. For example, if you use scale_factor 2 and a 4x upscale model, the image will get scaled down after the upscale model runs.
@@ -58,7 +58,7 @@ The main disadvantage compared to the alternatives I mentioned is that it is rel
<details>
<summary>Expand for advanced parameters</summary>
<summary>★ Click to expand for information on YAML parameters ★</summary>
Note: JSON is also valid YAML so you can use that instead if you prefer.
@@ -144,9 +144,14 @@ sharpen_strength: 1.0
# Disables the callback function (basically disables previews).
skip_callback: false
# Allows specifying an offset into highres_sigmas.
# You can use a negative number here, in which case we count from the end.
sigma_offset: 0
# Offset to sigmas passed to the model, -0.05 would mean reduce the sigma by 5%.
# If unset, sigma_dishonesty_factor_guidance will use the value from sigma_dishonesty_factor
# for guidance steps.
# Telling the model there's less noise than there actually is can increase detail
# (and conversely telling it there's more will reduce detail/smooth things out).
# A little goes a long way. Start with something like -0.03 to increase detail.
sigma_dishonesty_factor: 0.0
sigma_dishonesty_factor_guidance: null
# When enabled, uses an upscale model if connected. Mainly useful with
# iteration overrides.
@@ -210,19 +215,6 @@ iteration_override:
Supported schedules: `alignyoursteps`, `beta`, `ddim_uniform`, `exponential`, `gits`, `karras`, `laplace`, `normal`, `polyexponential`, `sgm_uniform`, `simple`, `vp`
Schedule overrides may also be combined with the `sigma_offset` parameter. The official DiffuseHigh uses the last 15 steps of a 50 step Karras schedule which would look like:
```yaml
schedule_override:
schedule_name: karras
steps: 50
# denoise defaults to 1.0 here.
# Negative values count from the end.
# Note that this is 16 because steps are from a -> b, b -> c, etc.
sigma_offset: -16
```
</details>
***
@@ -238,13 +230,15 @@ I tried to set the node defaults to align with the official implementation. Thes
* The sampler has a workaround for a [long standing bug in ComfyUI](https://github.com/comfyanonymous/ComfyUI/issues/2833) where generations aren't deterministic when `add_noise` is disabled in the sampler. However, this may change seeds. You can disable the workaround via the advanced YAML options - see `seed_rng` and `seed_rng_offset`.
* For `taesd` VAE mode, you will need the TAESD encoder models available at https://github.com/madebyollin/taesd - put them in `models/vae_approx`.
* You can use DiffuseHigh as an enhanced highres-fix by passing a pre-upscaled reference image, setting the iteration count to one and using a scale factor of 1.0.
* Setting `sigma_dishonesty_factor` and/or `sigma_dishonesty_factor_guidance` to a low negative value can be used to increase detail even for non-ancestral samplers (similar effect to increasing `s_noise`). See the YAML parameters section of this README.
* Using an upscale model or `image` guidance seems to make the most difference when you're going from low to mid-resolution (i.e. 512x512 to 1024x1024) so it may make sense to use the relatively slow `image` guidance and an upscale model for the first iteration and then switch to `latent` guidance and set `use_upscale_model: false` for subsequent iterations.
***
## Credits
Heavily referenced from the official implementation: [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/)
Contrast-adaptive sharpening sources: [1](https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h), [2](https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/), [3](https://github.com/Clybius)
* Initial version heavily referenced from the official implementation: [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/)
* Contrast-adaptive sharpening sources: [1](https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h), [2](https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/), [3](https://github.com/Clybius)
* `sigma_dishonesty_factor` concept from A1111's [Detail Daemon](https://github.com/muerrilla/sd-webui-detail-daemon) extension. (There's also a [ComfyUI version](https://github.com/Jonseed/ComfyUI-Detail-Daemon) now.)
Thanks!
+6
View File
@@ -2,6 +2,12 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20241027
* `sigma_offset` YAML parameter removed - you can use schedule overrides to accomplish the same effect (see README).
* Chunked sampling mode added, should make samplers that care about state (i.e. momentum or history like `dpmpp_2m`) work better for guidance steps. May change seeds, you can disable with `chunked_sampling: false` in YAML parameters.
* Added `sigma_dishonesty_factor` and `sigma_dishonesty_factor_guidance` YAML parameters - can be used to increase detail. See README.
## 20241023
* Initial support for rectified flow models (Flux, SD3, SD3.5). Might slightly change seeds for other models.
+12 -3
View File
@@ -16,6 +16,7 @@ class Config:
_overridable_fields = { # noqa: RUF012
"blend_by_mode",
"blend_mode",
"chunked_sampling",
"denoised_wavelet_multiplier",
"dtcwt_biort",
"dtcwt_mode",
@@ -44,7 +45,8 @@ class Config:
"sharpen_reference",
"sharpen_strength",
"skip_callback",
"sigma_offset",
"sigma_dishonesty_factor_guidance",
"sigma_dishonesty_factor",
"use_upscale_model",
"vae_decode_kwargs",
"vae_encode_kwargs",
@@ -71,6 +73,7 @@ class Config:
*,
blend_mode="lerp",
blend_by_mode="image",
chunked_sampling=True,
denoised_wavelet_multiplier=1.0,
dtcwt_biort="near_sym_a",
dtcwt_mode=False,
@@ -106,7 +109,8 @@ class Config:
sharpen_reference=True,
sharpen_strength=1.0,
skip_callback=False,
sigma_offset=0,
sigma_dishonesty_factor_guidance: None | float = None,
sigma_dishonesty_factor=0.0,
upscale_model=None,
use_upscale_model=True,
vae_decode_kwargs=None,
@@ -121,7 +125,11 @@ class Config:
)
self.seed_rng = seed_rng
self.seed_rng_offset = seed_rng_offset
self.sigma_offset = sigma_offset
self.sigma_dishonesty_factor = sigma_dishonesty_factor
self.sigma_dishonesty_factor_guidance = fallback(
sigma_dishonesty_factor_guidance,
sigma_dishonesty_factor,
)
self.skip_callback = skip_callback
self.fadeout_factor = fadeout_factor
self.scale_factor = scale_factor
@@ -190,6 +198,7 @@ class Config:
self.blend_function = BLENDING_MODES[blend_mode]
self.enable_gc = enable_gc
self.enable_cache_clearing = enable_cache_clearing
self.chunked_sampling = chunked_sampling
self.iteration_override = {}
if iteration_override is None or iteration_override == {}:
return
+120
View File
@@ -0,0 +1,120 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from .utils import ensure_model, sigma_to_float
if TYPE_CHECKING:
import torch
class GuidedModel:
def __init__(
self,
dh_sampler_object,
guidance_sigmas: torch.Tensor,
guidance_steps: int,
):
self.dhso = dh_sampler_object
self.allow_guidance = True
self.force_guidance: None | int = None
self.set_guidance_range(guidance_sigmas, guidance_steps)
def set_guidance_range(self, guidance_sigmas, guidance_steps):
if len(guidance_sigmas) >= 2:
self.guidance_start_sigma = sigma_to_float(guidance_sigmas[0]) + 1e-05
self.guidance_end_sigma = sigma_to_float(guidance_sigmas[-1]) + 1e-05
else:
self.guidance_start_sigma = None
self.guidance_end_sigma = None
self.guidance_sigmas_list = tuple(guidance_sigmas.detach().cpu().tolist())
self.guidance_steps = guidance_steps
def find_guidance_step_(self, sigma_float: float) -> None | int:
return (
next(
(
idx
for idx, gsigma in enumerate(self.guidance_sigmas_list)
if gsigma <= sigma_float
),
None,
)
if sigma_float <= self.guidance_start_sigma
else 0
)
def get_guidance_step(self, sigma: torch.Tensor) -> None | int:
sigma_float = sigma_to_float(sigma)
if not self.allow_guidance:
return None
sigma_float = sigma_to_float(sigma)
if self.force_guidance is None and (
self.guidance_start_sigma is None or sigma_float < self.guidance_end_sigma
):
return None
if self.force_guidance is not None and self.force_guidance >= 0:
return self.force_guidance
step_idx = (
next(
(
idx
for idx, gsigma in enumerate(self.guidance_sigmas_list)
if gsigma <= sigma_float
),
None,
)
if sigma_float <= self.guidance_start_sigma
else 0
)
if step_idx is None:
if not self.force_guidance:
return None
step_idx = max(0, self.guidance_steps - 1)
return min(step_idx, self.guidance_steps - 1)
def make_wrapper(self):
guided_model = self
class DiffuseHighModelWrapper:
def __getattr__(self, k):
try:
return getattr(guided_model, k)
except AttributeError:
raise AttributeError(k) from None
def __call__(self, *args: list, **kwargs: dict) -> torch.Tensor:
return guided_model(*args, **kwargs)
return DiffuseHighModelWrapper()
def __call__(
self,
x: torch.Tensor,
sigma: torch.Tensor,
**extra_args: dict,
) -> torch.Tensor:
dhso = self.dhso
model = dhso.model
dhso.seed_offset += 1
ensure_model(model)
guidance_step = self.get_guidance_step(sigma)
sigma_offset = max(
1e-05,
1.0
+ (
dhso.sigma_dishonesty_factor_guidance
if guidance_step is not None
else dhso.sigma_dishonesty_factor
),
)
denoised = model(
x,
sigma * sigma_offset if sigma_offset != 1 else sigma,
**extra_args,
)
return (
denoised
if guidance_step is None
else dhso.apply_guidance(guidance_step, denoised)
)
+111 -58
View File
@@ -10,12 +10,13 @@ from tqdm import tqdm
from tqdm.auto import trange
from .config import Config
from .guided_model import GuidedModel
from .schedule import Schedule
from .tensor_image_ops import (
blend_wavelets,
scale_wavelets,
)
from .utils import ensure_model, fallback
from .utils import fallback
class DiffuseHighSampler:
@@ -164,9 +165,8 @@ class DiffuseHighSampler:
**self.schedule_override,
).sigmas.to(self.highres_sigmas_input)
@classmethod
def add_restart_noise(
cls,
self,
x: torch.Tensor,
sigma_min: float | torch.Tensor,
sigma_max: float | torch.Tensor,
@@ -174,7 +174,51 @@ class DiffuseHighSampler:
s_noise: float = 1.0,
) -> torch.Tensor:
noise_factor = (sigma_max**2 - sigma_min**2) ** 0.5
return x + torch.randn_like(x).mul_(noise_factor * s_noise)
return self.add_noise(x, noise_factor, factor=s_noise)
def add_noise(
self,
latent: torch.Tensor,
sigma: float | torch.Tensor,
*,
sigma_next: None | float | torch.Tensor = None,
factor=1.0,
allow_max_denoise=True,
noise_sampler: None | callable = None,
) -> torch.Tensor:
self.seed_offset += 1
sigma_next = fallback(sigma_next, sigma)
noise = (
noise_sampler(sigma, sigma_next)
if noise_sampler is not None
else torch.randn_like(latent)
)
if factor != 1:
noise *= factor
return self.model_sampling.noise_scaling(
sigma,
noise,
latent,
max_denoise=allow_max_denoise
and sigma >= self.model_sampling.sigma_max - 1e-05,
)
def run_sampler_with_pbar(
self,
x: torch.Tensor,
sigmas: torch.Tensor,
pbar_title: str,
*args: list,
**kwargs: dict,
):
with tqdm(
disable=self.disable_pbar,
total=1,
desc=f"{pbar_title} ({len(sigmas) - 1}) {float(sigmas[0]):>2.03f} ... {float(sigmas[-1]):>2.03f}",
) as pbar:
x = self.run_sampler(x, sigmas, *args, **kwargs)
pbar.update()
return x
def run_steps(
self,
@@ -183,29 +227,30 @@ class DiffuseHighSampler:
sigmas: None | torch.Tensor = None,
) -> torch.Tensor:
sigmas = self.highres_sigmas if sigmas is None else sigmas
soffset = (
self.sigma_offset
if self.sigma_offset >= 0
else len(sigmas) + self.sigma_offset
)
guidance_sigmas = sigmas[soffset : soffset + self.guidance_steps + 1]
normal_sigmas = sigmas[soffset + self.guidance_steps :]
step_idx = 0
model = self.model
ensure_model(model)
sigmas_len = len(sigmas)
if sigmas_len < 2:
return x
guidance_steps = max(0, min(self.guidance_steps, sigmas_len - 1))
guidance_sigmas = sigmas[: guidance_steps + 1]
guidance_sigmas_len = len(guidance_sigmas)
normal_sigmas = sigmas[guidance_steps:]
normal_sigmas_len = len(normal_sigmas)
if guidance_sigmas_len < 2 and normal_sigmas_len < 2:
return x
guided_model = GuidedModel(self, guidance_sigmas, guidance_steps)
model_wrapper = guided_model.make_wrapper()
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)
same_samplers = self.sampler == self.guidance_sampler
for k in (
"inner_model",
"sigmas",
):
if hasattr(model, k):
setattr(model_wrapper, k, getattr(model, k))
if self.guidance_restart == 0 and same_samplers and self.chunked_sampling:
# No guidance restarts and the guidance sampler is the same as the
# # normal one and chunked mode enabled - we can sample all the sigmas at once.
return self.run_sampler_with_pbar(
x,
sigmas,
model=model_wrapper,
pbar_title="combined steps",
)
for repidx in trange(
self.guidance_restart + 1,
@@ -220,8 +265,33 @@ class DiffuseHighSampler:
guidance_sigmas[0],
s_noise=self.guidance_restart_s_noise,
)
self.seed_offset += 1
guidance_steps = len(guidance_sigmas) - 1
if (
same_samplers
and self.chunked_sampling
and repidx == self.guidance_restart
):
# On the last guidance restart iteration, we can sample all the sigmas at once
# as long as the guidance sampler is the same as the normal one and we're in
# chunked mode.
return self.run_sampler_with_pbar(
x,
sigmas,
model=model_wrapper,
pbar_title="combined steps",
)
if guidance_sigmas_len < 2:
continue
if self.chunked_sampling:
guided_model.force_guidance = -1
x = self.run_sampler_with_pbar(
x,
guidance_sigmas,
model=model_wrapper,
sampler=self.guidance_sampler,
pbar_title="guidance steps",
)
guided_model.force_guidance = None
continue
with trange(
guidance_steps,
initial=1,
@@ -229,7 +299,7 @@ class DiffuseHighSampler:
desc="guidance step",
) as pbar:
for idx in pbar:
step_idx = idx
guided_model.force_guidance = idx
step_sigmas = guidance_sigmas[idx : idx + 2]
if step_sigmas[-1] > step_sigmas[0]:
raise ValueError(
@@ -245,17 +315,15 @@ class DiffuseHighSampler:
sampler=self.guidance_sampler,
disable_pbar=True,
)
self.seed_offset += 1
if len(normal_sigmas) > 1:
ensure_model(model)
with tqdm(
disable=self.disable_pbar,
total=1,
desc=f"normal steps ({len(normal_sigmas) - 1}) {float(normal_sigmas[0]):>2.03f} ... {float(normal_sigmas[-1]):>2.03f}",
) as pbar:
x = self.run_sampler(x, normal_sigmas)
pbar.update()
self.seed_offset += len(normal_sigmas) - 1
guided_model.force_guidance = None
if normal_sigmas_len >= 2:
guided_model.allow_guidance = False
x = self.run_sampler_with_pbar(
x,
normal_sigmas,
model=model_wrapper,
pbar_title="normal steps",
)
return x
@classmethod
@@ -336,17 +404,6 @@ class DiffuseHighSampler:
"Highres sigmas (including schedule overrides) must end at 0 (full denoise)",
)
self.gc()
if self.config.sigma_offset >= len(self.highres_sigmas) - 1:
raise ValueError(
"Bad sigma_offset: points to sigma past penultimate sigma",
)
if self.config.sigma_offset < 0 and (
self.config_sigma_offset == -1
or abs(self.config.sigma_offset) >= len(self.highres_sigmas)
):
raise ValueError(
"Negative sigma_offset can't point to last sigma or to sigma less than index 0",
)
with tqdm(disable=self.disable_pbar, total=1, desc="upscale") as pbar:
img_hr = self.upscale(
self.reference_image,
@@ -373,16 +430,12 @@ class DiffuseHighSampler:
)
else:
raise ValueError("Bad guidance_mode")
x_noise = torch.randn_like(x_new)
self.seed_offset += 1
x_new = self.model_sampling.noise_scaling(
self.highres_sigmas[self.sigma_offset],
x_noise.mul_(self.renoise_factor),
x_new = self.add_noise(
x_new,
max_denoise=self.highres_sigmas[self.sigma_offset]
>= self.model_sampling.sigma_max - 1e-05,
self.highres_sigmas[0],
sigma_next=self.highres_sigmas[1],
factor=self.renoise_factor,
)
del x_noise
self.gc()
x_new = self.run_steps(x_new, sigmas=self.highres_sigmas)
if iteration == self.iterations - 1:
+4
View File
@@ -18,3 +18,7 @@ def fallback(val, default, *, exclude=None, default_is_fun=False):
def scale_dim(n, factor=1.0, *, increment=64) -> int:
return math.ceil((n * factor) / increment) * increment
def sigma_to_float(sigma):
return sigma.detach().cpu().max().item()