Cleanups, documentation updates

This commit is contained in:
blepping
2024-07-15 18:47:16 -06:00
parent 28b40df8b1
commit 2791278704
2 changed files with 185 additions and 71 deletions
+145 -31
View File
@@ -1,15 +1,8 @@
# Overly Complicated Sampling
Wildly unsound and experimental sampling for [ComfyUI](https://github.com/comfyanonymous/ComfyUI).
## Description
Very unstable, experimental and mathematically unsound sampling for ComfyUI.
Current status: In flux, not suitable for general use.
*Note*: You will basically always have to tweak settings like `s_noise` to get a good result. If the generation looks smooth/undetailed increase `s_noise` somewhere. If it looks crunchy, super high contrast, etc then try reducing noise.
Experimental and mathematically unsound (but fun!) sampling for [ComfyUI](https://github.com/comfyanonymous/ComfyUI).
**Status**: In flux, may be useful but likely to change/break workflows frequently. Mainly for advanced users.
## Credits
@@ -24,6 +17,17 @@ I can move code around but sampling math and creating samplers is far beyond my
This repo wouldn't be possible without building on the work of others. Thanks!
## Features
* Many different samplers.
* Allows scheduling samplers (i.e. run `euler` for steps 1-4, then switch to `dpmpp_sde`).
* CFG++ support (for some samplers).
* Native support for Restart sigmas.
* Supports custom noise types.
* Immiscible noise for sampling and Restart. See https://arxiv.org/abs/2406.12303 (note that it was designed for training not inference).
* Allows splitting/combining steps in various ways for (potentially) more accurate sampling.
* Supports Diffrax, torchdiffeq, torchode and torchsde solver backends. (SDE mode not recommended currently.)
* Many tuneable parameters to play with.
## Usage
@@ -212,8 +216,6 @@ The left side group matches steps 0, 1, 2. The right side group matches all step
#### Input Parameters
Currently unused for groups.
<!--
* `merge_sampler`: Value type: `OCS_SUBSTEPS`. Only used when `merge_method` is `sample` or `sample_uncached`. Allows defining the sampler used for merging substeps.
-->
@@ -234,8 +236,9 @@ eta: 1.0
# Reversible ETA (used for reversible samplers). May not do anything currently.
reta: 1.0
# Currently unused.
avgmerge_stretch: 0.4
# Interpolate the schedule by the specified factor. Only used by the overshoot merge method.
#: Example if factor 2 and steps [0,1,2] you'd get [0, 0.5, 1.0, 1.5, 2]
overshoot_expand_steps: 1
# Only used by the overshoot merge method currently.
restart:
@@ -364,10 +367,14 @@ eta: 1.0
dyn_eta_start: null
dyn_eta_end: null
# CFG++ scale (see https://cfgpp-diffusion.github.io/)
# Setting this to 1.0 is the equivalent of enabling it. Can also be set
# to a negative value (I don't recommend going lower than -0.5).
cfgpp_scale: 0
# alt CFG++ scale (see https://cfgpp-diffusion.github.io/)
# Based on the initial incorrect ComfyUI implementation, but it seems to
# produce decent results sometimes.
# Can also be set to a negative value (I don't recommend going lower than -0.5).
alt_cfgpp_scale: 0
# CFG++ (see https://cfgpp-diffusion.github.io/)
cfgpp: false
### Reversible Settings ###
@@ -381,38 +388,111 @@ reversible_scale: 1.0
dyn_reta_start: null
dyn_reta_end: null
# Allows normalizing the latent while sampling.
# Example values, not enabled by default.
normalize:
# One of: before, after
- phase: after
# First step to start applying the effect, 0-based.
start_step: 0
# Last step to apply the effect, 0-based, inclusive.
end_step: 999
# Dimensions used for calculating the mean for balance and target.
dims: [-2, -1]
# Target for balancing. Defaults to a target mean of 0.
# This can be a scalar (i.e. 0.0) or an array but it must match the
# dimensions of the mean.
balance_target: 0.0
# Multiplier on the balance scaling. In other words, you can gently
# pull the latent in the direction of the target mean rather than just
# setting it all at once.
balance_scale: 0
# Adjustment applied after the balance, works like balance_scale.
adjust_scale: 0
# Target for the adjustment. Can be "x" (only if phase after), a scalar or
# array matching the dims (see balance_target).
adjust_target: x
# Example for dims [-2, -1] on SD1.x, 2.x, SDXL or other 4 channel models.
#adjust_target: [[[[-0.5]],[[-0.5]], [[0.5]], [[0.5]]]]
### ODE Sampler Settings ###
# Solver type.
ode_solver: dopri5
de_solver: dopri5 # Example - varies based on solver sampler.
# Relative tolerance (log 10)
ode_rtol: -1.5
de_rtol: -1.5
# Absolute tolerance (log 10)
ode_atol: -3.5
de_atol: -3.5
# Max model calls allowed to compute the solution. If the limit is exceeded, it is an error.
ode_max_nfe: 1000
# Hack that seems to help results. Set to 0 to disable.
ode_fixup_hack: 0.025
de_max_nfe: 1000
# Min sigma to sample to. If the current step start <= min sigma, then the sampler will run
# a Euler step. If the current step end <= min sigma then the slover will sample to the min
# sigma and then to a Euler step from min sigma for the rest.
de_min_sigma: 0.0292
# Hack that seems to help results by stretching the down sigma a bit. Set to 0 to disable.
de_fixup_hack: 0.025
## torchdiffeq (tde) specific parameters ##
# Used to split the step into sections. Useful for fixed step methods.
ode_split: 1
# Applies to: solver_torchode, solver_diffrax
de_split: 1
## torchode (tode) specific parameters ##
# Initial step size (as a percentage).
ode_initial_step: 0.25
# Applies to: solver_torchode, solver_diffrax
de_initial_step: 0.25
# Coefficients for the step size PID controller.
# See https://en.wikipedia.org/wiki/Proportional%E2%80%93integral%E2%80%93derivative_controller
# These values seem okay with dopri5.
ode_ctl_pcoeff: 0.3
ode_ctl_icoeff: 0.9
ode_ctl_dcoeff: 0.2
# Applies to: solver_tode, solver_diffrax
de_ctl_pcoeff: 0.3
de_ctl_icoeff: 0.9
de_ctl_dcoeff: 0.2
# Controls whether to compile the solver. May or may not work,
# also may or may not be a speed increase as the compiled solver is
# not cached between substeps.
ode_compile: false
# Applies to: solver_torchode
tode_compile: false
### torchsde solver specific parameters.
tsde_noise_type: "scalar"
tsde_sde_type: "stratonovich"
tsde_levy_area_approx: "none"
tsde_noise_channels: 1
tsde_g_multiplier: 0.05
tsde_g_reverse_time: true
tsde_g_derp_mode: false
tsde_batch_channels: true
### diffrax solver specific parameters.
# Turns on adaptive stepping. When enabled, de_split is not used.
# When disabled, it may be desirable to set de_split.
diffrax_adaptive: false
# Hack to make some solver methods work. May not be safe.
diffrax_fake_pure_callback: true
# Some diffrax methods don't allow adaptive stepping, enabling this
# makes them usable although it's less efficient (3x cost, 2x accuracy).
diffrax_half_solver: false
diffrax_batch_channels: false
# Some solvers require specific types of Levy area approximation.
# See: https://docs.kidger.site/diffrax/api/brownian/#levy-areas
diffrax_levy_area_approx: "brownian_increment"
# Some solvers may require manually specifying the error order.
diffrax_error_order: null
# Enables SDE mode (and SDE-specific solvers). May not be worth using.
diffrax_sde_mode: false
# Noise multiplier when SDE mode is enabled.
diffrax_g_multiplier: 0.0
# Only applies when time scaling is enabled. Reverses time.
diffrax_g_reverse_time: false
# Scales the g multiplier based on the current time.
diffrax_g_time_scaling: false
# Experimental option to flip the sign on the g multiplier when time >= half the step.
# i.e. if you'd get 1,2,3,4 as g values for the step, with this it would be 1,2,-3,-4.
diffrax_g_split_time_mode: false
### Other Sampler Specific Parameters ###
@@ -454,3 +534,37 @@ dyn_deta_end: null
# One of lerp, lerp_alt, deta
dyn_deta_mode: "lerp"
```
### `OCS SimpleRestartSchedule`
Generates a restart schedule.
#### Node Parameters
* `start_step`: 0-based first step for the restart schedule to apply.
#### Input Parameters
* `sigmas`: Sigmas to restartify. Output from any normal schedule node.
#### Text Parameters
JSON or YAML schedule in list form.
```yaml
- [4, -3]
- [2, -1]
- 1
```
Each item should be one of:
* A pair `[interval, jump]` - after `interval` steps, make a relative jump of `jump` steps.
* A single integer `schedule_index`: resume the schedule at the specified 0-based index.
The example above means:
1. After 4 steps, jump back 3 steps.
2. After 2 steps, jump back one step.
3. Go to the second item (after 2 steps, jump back one step).
The node `start_step` parameter is effectively the same as `[start_step, 0]` as a schedule item.
+40 -40
View File
@@ -1490,7 +1490,7 @@ class TODEStep(DESolverStep):
self,
*args,
de_initial_step=0.25,
de_compile=False,
tode_compile=False,
de_ctl_pcoeff=0.3,
de_ctl_icoeff=0.9,
de_ctl_dcoeff=0.2,
@@ -1505,7 +1505,7 @@ class TODEStep(DESolverStep):
self.de_ctl_pcoeff = de_ctl_pcoeff
self.de_ctl_icoeff = de_ctl_icoeff
self.de_ctl_dcoeff = de_ctl_dcoeff
self.de_compile = de_compile
self.de_compile = tode_compile
self.de_initial_step = de_initial_step
def check_solver_support(self):
@@ -1605,28 +1605,28 @@ class TSDEStep(DESolverStep):
de_initial_step=0.25,
de_split=1,
de_adaptive=False,
de_noise_type="scalar",
de_sde_type="stratonovich",
de_levy_area_approx="none",
de_noise_channels=1,
de_g_multiplier=0.05,
de_g_reverse_time=True,
de_g_derp_mode=False,
de_batch_channels=True,
tsde_noise_type="scalar",
tsde_sde_type="stratonovich",
tsde_levy_area_approx="none",
tsde_noise_channels=1,
tsde_g_multiplier=0.05,
tsde_g_reverse_time=True,
tsde_g_derp_mode=False,
tsde_batch_channels=True,
**kwargs,
):
super().__init__(*args, **kwargs)
self.de_initial_step = de_initial_step
self.de_adaptive = de_adaptive
self.de_split = de_split
self.de_noise_type = de_noise_type
self.de_sde_type = de_sde_type
self.de_levy_area_approx = de_levy_area_approx
self.de_g_multiplier = de_g_multiplier
self.de_noise_channels = de_noise_channels
self.de_g_reverse_time = de_g_reverse_time
self.de_g_derp_mode = de_g_derp_mode
self.de_batch_channels = de_batch_channels
self.de_noise_type = tsde_noise_type
self.de_sde_type = tsde_sde_type
self.de_levy_area_approx = tsde_levy_area_approx
self.de_g_multiplier = tsde_g_multiplier
self.de_noise_channels = tsde_noise_channels
self.de_g_reverse_time = tsde_g_reverse_time
self.de_g_derp_mode = tsde_g_derp_mode
self.de_batch_channels = tsde_batch_channels
def check_solver_support(self):
pass
@@ -1790,21 +1790,21 @@ class DiffraxStep(DESolverStep):
self,
*args,
de_split=1,
de_adaptive=False,
de_fake_pure_callback=True,
de_initial_step=0.25,
de_ctl_pcoeff=0.3,
de_ctl_icoeff=0.9,
de_ctl_dcoeff=0.2,
de_g_multiplier=0.0,
de_half_solver=False,
de_batch_channels=False,
de_levy_area_approx="brownian_increment",
de_error_order=None,
de_sde_mode=False,
de_g_reverse_time=False,
de_g_time_scaling=False,
de_g_split_time_mode=False,
diffrax_adaptive=False,
diffrax_fake_pure_callback=True,
diffrax_g_multiplier=0.0,
diffrax_half_solver=False,
diffrax_batch_channels=False,
diffrax_levy_area_approx="brownian_increment",
diffrax_error_order=None,
diffrax_sde_mode=False,
diffrax_g_reverse_time=False,
diffrax_g_time_scaling=False,
diffrax_g_split_time_mode=False,
**kwargs,
):
super().__init__(*args, **kwargs)
@@ -1841,23 +1841,23 @@ class DiffraxStep(DESolverStep):
)
# jax.config.update("jax_disable_jit", True)
self.de_solver_method = solvers[self.de_solver_name]()
if de_half_solver:
if diffrax_half_solver:
self.de_solver_method = diffrax.HalfSolver(self.de_solver_method)
self.de_ctl_pcoeff = de_ctl_pcoeff
self.de_ctl_icoeff = de_ctl_icoeff
self.de_ctl_dcoeff = de_ctl_dcoeff
self.de_initial_step = de_initial_step
self.de_adaptive = de_adaptive
self.de_adaptive = diffrax_adaptive
self.de_split = de_split
self.de_fake_pure_callback = de_fake_pure_callback
self.de_g_multiplier = de_g_multiplier
self.de_batch_channels = de_batch_channels
self.de_levy_area_approx = levy_areas[de_levy_area_approx]
self.de_error_order = de_error_order
self.de_sde_mode = de_sde_mode
self.de_g_reverse_time = de_g_reverse_time
self.de_g_time_scaling = de_g_time_scaling
self.de_g_split_time_mode = de_g_split_time_mode
self.de_fake_pure_callback = diffrax_fake_pure_callback
self.de_g_multiplier = diffrax_g_multiplier
self.de_batch_channels = diffrax_batch_channels
self.de_levy_area_approx = levy_areas[diffrax_levy_area_approx]
self.de_error_order = diffrax_error_order
self.de_sde_mode = diffrax_sde_mode
self.de_g_reverse_time = diffrax_g_reverse_time
self.de_g_time_scaling = diffrax_g_time_scaling
self.de_g_split_time_mode = diffrax_g_split_time_mode
# As slow and safe as possible.
@staticmethod