diff --git a/README.md b/README.md index c9606e3..150ca0b 100644 --- a/README.md +++ b/README.md @@ -24,17 +24,20 @@ _Note for Flux users_: Set `cfg1_uncond_optimization: true` in the `model` block ## Credits -I can move code around but sampling math and creating samplers is far beyond my ability. I didn't write any of the original samplers: +I can move code around but sampling math and creating samplers generally beyond my ability. I didn't write any of the original samplers: -* Euler, Heun++2, DPMPP SDE, DPMPP 2S, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation. -* Reversible Heun, Reversible Heun 1s, RES, Trapezoidal, Bogacki, Reversible Bogacki, RK4, RKF45, dynamic RK(4) and Euler Dancing samplers based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +* Euler, Heun++2, DPMPP SDE, DPMPP 2S, gradient estimation, RES multistep, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation. +* Reversible Heun, Reversible Heun 1s, RES, Trapezoidal, Bogacki, Reversible Bogacki, RK4, RKF45, dynamic RK(4), SENS and Euler Dancing samplers based on implementation from [https://github.com/Clybius/ComfyUI-Extra-Samplers](https://github.com/Clybius/ComfyUI-Extra-Samplers). * TTM JVP sampler based on implementation written by Katherine Crowson (but yoinked from the Extra-Samplers repo mentioned above). * Distance sampler based on implementation from https://github.com/Extraltodeus/DistanceSampler * IPNDM, IPNDM_V and DEIS adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py (I used the Comfy version as a reference). +* PingPong sampler idea from https://github.com/ace-step/ACE-Step/ (implementation also referenced from that source). * Normal substep merge strategy based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -* Immiscible noise processing based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 and idea for sampling with it from https://github.com/Clybius +* Immiscible noise processing based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 and https://github.com/yhli123/Immiscible-Diffusion - idea for sampling with it and implementation help from https://github.com/Clybius * Precedence climbing (Pratt) expression parser based on implementation from https://github.com/andychu/pratt-parsing-demo +Please notify me if I somehow missed appropriately crediting any code used here, any such ommissions are unintentional. + This repo wouldn't be possible without building on the work of others. Thanks! ## Usage @@ -95,6 +98,9 @@ reta: 1.0 # Parameters related to restart sampling. restart: + # When enabled, out of order sigmas will be detected as restart. + # You can disable this if you want to use OCS for something like unsampling. + enabled: true # Scales the noise added by restart sampling. s_noise: 1.0 # Immiscible block same as described below. @@ -104,13 +110,17 @@ restart: # The noise block allows defining global noise sampling parameters. noise: - # You can disable this to allow GPU noise generation. I believe it only makes a difference for Brownian. + # You can disable this to allow GPU noise generation. cpu_noise: true # ComfyUI has a bug where if you disable add_noise in the sampler, no seed gets set. If you # are manually noising a sample and have add_noise turned off then you should enable this if # you want reproducible generations. - set_seed: false + set_seed: true + + # Only has an effect when set_seed is enabled. Will advance the RNG this many times to + # avoid the common mistake of using the same noise for sampling as the initial noise. + seed_offset: 1 # Global scale scale for generated noise scale: 1.0 @@ -177,35 +187,31 @@ noise: # See: https://docs.scipy.org/doc/scipy/reference/generated/scipy.optimize.linear_sum_assignment.html#scipy.optimize.linear_sum_assignment maximize: false + # Can be set to enable immiscible v2 mode. 0.0 is disabled, 0.1 is a reasonable value. + distance_scale: 0.0 + + # If null will use the same value as distance_scale. + distance_scale_ref: null + filter: null -# Model calls can be cached. This is very experimental: I don't recommend using it -# unless you know what you're doing. model: # When enabled, skips generating uncond when you have CFG set to 1. Disabled by # default as stuff like CFG++ won't work without uncond. Useful to enable for # models like Flux that don't actually use CFG. cfg1_uncond_optimization: false - cache: - # The cache size. - size: 0 - - # Threshold for model call caching. For example if you have size=3 and threshold=1 - # then model calls 1 through 3 will be cached, but model call 0 will not be (the first one). - # Additional explanation: Some samplers call the model multiple times per step. For example, - # Bogacki uses three model calls: 0, 1, 2 - threshold: 1 - - # Maximum use count for cache items. - max_use: 1000000 - filter: + # Input to the model input: null + # Result after CFG calculation denoised: null + # Result after CFG calculation for JVP jdenoised: null + # Cond - positive prompt cond: null + # Uncond - negative prompt uncond: null ``` @@ -236,6 +242,7 @@ When running multiple substeps per step, the results will combined based on the * `supreme_avg`: The model is called at least once per step (and possibly additional times for higher order samplers). Each substep shares the first model call result. The results are averaged together. *Note*: Since the first model call is shared and the initial input is the same for each substep, there is no point in running multiple identical substeps. Also note: This merge strategy doesn't work well with non-ancestral samplers (i.e. dpmpp_2m or any sampler with `eta: 0`). * `overshoot`: The model is called at least once per step. It will sample steps equal to the number of substeps, starting from the current step. Then it will restart back to the expected step. * `lookahead`: Similar to `overshoot`, it samples ahead based on the number of substeps. The last model prediction is used to do a Euler step to the expected step. *Note:* Very experimental, likely to change in the future. +* `pingpong`: Works similar to `overshoot` and `lookahead` methods except it does a pingpong sampler style step to the expected sigma. * `dynamic`: Allows specifying the group parameters as an expression to be evaluated. See below. **Dynamic Groups**: When `merge_method` is set to `dynamic` you must specify a `dynamic` block in the text parameters. The dynamic block may be either a string with the expression or a list of objects with an (optional) `when` key and a (required) `expression` key. The expression should return a dictionary of parameters you can set in the node (including both keys/values from the text parameters and widgets in the node). The first matching item will be used. Example: @@ -331,6 +338,11 @@ lookahead: immiscible: size: 0 +# Only used by the pingpong merge method. +pingpong: + # Scales the noise added by lookahead sampling. + s_noise: 1.0 + pre_filter: null post_filter: null @@ -347,36 +359,43 @@ post_filter: null In alphabetical order. * `adapter`: Wraps a normal ComfyUI `SAMPLER`. Attach a `SAMPLER` parameter to the node. Note: Samplers that do unusual stuff like try to manipulate the model won't work. ComfyUI's built-in CFG++ samplers in particular do not work here. +* `blep_bas`: Batch Augmented Sampler. My own dumb experiment that expands the batch and averages the result. May be very slow/require a lot of VRAM. See parameters: `bas` +* `blep_euler_cycle`: See parameters: `cycle_pct`. +* `blep_weoon`: Wavelet-based second order sampler. Another dumb experiment. See parameters: `weoon` * `bogacki`: Bogacki-Shampine sampler. Also has a reversible variant. +* `clybius_euler_dancing`: Pretty broken currently, will probably require increased `s_noise` values. See parameters: `deta`, `leap`, `deta_mode`. +* `clybius_sens`: Reversible dpmpp_3m_sde variant. Supports a separate set of reversible parameters in `tsde_reversible`. * `deis`: See parameters: `history_limit`. Does not work well with ETA, I don't recommending leaving ETA at the default 1. -* `distance`: See parameters: `distance`. Adaptive-ish/configurable step variant of Heun. Taken from: https://github.com/Extraltodeus/DistanceSampler -* `dpmpp_2m_sde`: See parameters: `history_limit`. +* `dpm2`: Set `eta: 0` for non-ancestral variant. +* `dpmpp_2m_sde`: Also supports reversible parameters. See parameters: `history_limit`. * `dpmpp_2m`: `eta` and `s_noise` parameters are ignored. See parameters: `history_limit`. * `dpmpp_2s` * `dpmpp_3m_sde`: See parameters: `history_limit`. * `dpmpp_sde` * `dynamic`: Advanced step method that allows using an expression to determine the sampler parameters at each substep. See below for a more detailed explanation. -* `euler_cycle`: See parameters: `cycle_pct`. -* `euler_dancing`: Pretty broken currently, will probably require increased `s_noise` values. See parameters: `deta`, `leap`, `deta_mode`. * `euler`: If samplers came in vanilla. -* `heun`: Alternate Heun implementation. Supports reversible parameters. See parameters: `history_limit`. +* `extraltodeus_distance`: Adaptive-ish/configurable step variant of Heun. Referenced from [https://github.com/Extraltodeus/DistanceSampler](https://github.com/Extraltodeus/DistanceSampler). See parameters: `distance`. +* `gradient_estimation` * `heun_1s`: Alternate Heun one step implementation. Supports reversible parameters. +* `heun`: Alternate Heun implementation. Supports reversible parameters. See parameters: `history_limit`. * `heunpp`: See parameters: `max_order`. * `ipndm_v`: See parameters: `history_limit`. * `ipndm`: See parameters: `history_limit`. +* `pingpong` * `res`: Refined Exponential Solver. I believe this is a variant of Heun. Generally works very well. +* `res_multistep` * `reversible_bogacki`: Reversible variant of Bockacki-Shampine. -* `reversible_heun`: Reversible variant of Heun. * `reversible_heun_1s`: Reversible variant of Heun 1 step. See parameters: `history_limit`. -* `rk4`: Range-Kutta 4th order sampler. -* `rkf45`: 5 model call flavor of RK. +* `reversible_heun`: Reversible variant of Heun. * `rk_dynamic`: Variant of RK4 that lets you set `max_order` (you can also set it to `0` to choose an order dynamically, doesn't seem to work so well though). +* `rk4`: Runge-Kutta 4th order sampler. +* `rkf45`: 5 model call flavor of RK. * `solver_diffrax`: Uses the [Diffrax](https://github.com/patrick-kidger/diffrax) solver backend. See `de_*` parameters below. * `solver_torchdiffeq`: Uses the [torchdiffeq](https://github.com/rtqichen/torchdiffeq) backend. See `de_*` parameters below. * `solver_torchode`: Uses the [torchode]((https://github.com/martenlienen/torchode)) backend. See `de_*` parameters below. * `solver_torchsde`: Uses the [torchsde](https://github.com/google-research/torchsde) backend. See `de_*` parameters below. -* `trapezoidal`: * `trapezoidal_cycle`: See parameters: `cycle_pct`. +* `trapezoidal`: * `ttm_jvp`: TTM is a weird sampler. If you're using model caching you must make sure the entries TTM uses are populated first (by having it run before any other samplers that call the model multiple times). It may also not work with some other model patches and upscale methods. See parameters: `alternate_phi_2_calc` **Sampler Feature Support** @@ -384,40 +403,46 @@ In alphabetical order. |Name|Cost|History|Order|Reversible|CFG++| |-|-|-|-|-|-| |`adapter`|?|?|?|?|?| +|`blep_bas`|variable||||| +|`blep_euler_cycle`|1||||X| +|`blep_trapezoidal_cycle`|2||||| +|`blep_weoon`|2||||| |`bogacki`|2||||| +|`clybius_euler_dancing`|1||||| +|`clybius_sens`|1|1|||| |`deis`|1|1-3 (1)|||| -|`distance`|variable||||| |`dpmpp_2m_sde`|1|1|||| |`dpmpp_2m`|1|1|||| |`dpmpp_2s`|2||||| |`dpmpp_3m_sde`|1|1-2 (2)|||| |`dpmpp_sde`|2||||| |`dynamic`|?|?|?|?|?| -|`euler_cycle`|1||||X| -|`euler_dancing`|1||||| |`euler`|1||||X| -|`heun`|2|||X|| +|`extraltodeus_distance`|variable||||| +|`gradient_estimation`|1|1|||| |`heun_1s`|1|1||X|| +|`heun`|2|||X|| |`heunpp`|1-3||X||| |`ipndm_v`|1|1-3 (1)|||| |`ipndm`|1|1-3 (1)|||| +|`pingpong`|1||||| |`res`|2||||| +|`res_multistep`|1|1|||| |`reversible_bogacki`|2|||X|| -|`reversible_heun`|2|||X|| |`reversible_heun_1s`|1|1||X|| +|`reversible_heun`|2|||X|| +|`rk4`|1-4||||| |`rk4`|4||||| |`rkf45`|5||||| -|`rk4`|1-4||||| |`solver_diffrax`|variable||||| |`solver_torchdiffeq`|variable||||| |`solver_torchode`|variable||||| |`solver_torchsde`|variable||||| |`trapezoidal`|2||||| -|`trapezoidal_cycle`|2||||| |`ttm_jvp`|2||||| -`deis`, `ipndm*` do not seem to work well with ancestralness, I recommend `eta: 0.25` or disable it completely. +`deis`, `ipndm*` and `gradient_estimation` do not seem to work well with ancestralness, I recommend `eta: 0.25` or disable it completely. **Solver Backend Samplers**: @@ -503,15 +528,25 @@ cfgpp: false ### Reversible Settings ### -# Reversible ETA (used for reversible samplers). -reta: 1.0 -# Scale of the reversible correction. Can also be set to a negative value. -reversible_scale: 1.0 -# No effect unless both start and end are set. Will scale the reta value based on the -# percentage of sampling. In other words, reta*dyn_reta_start at the beginning, -# reta*dyn_reta_end at the end. -dyn_reta_start: null -dyn_reta_end: null +reversible: + # 0-indexed step where reversible sampling will start. + start_step: 0 + # 0-indexed last step where reversible sampling will be used. + end_step: 9999 + # Scale of the reversible correction. Can also be set to a negative value. + scale: 1.0 + # Reversible ETA. + eta: 1.0 + + # No effect unless both start and end are set. Will scale the eta value based on the + # percentage of sampling. In other words, reta*dyn_reta_start at the beginning, + # reta*dyn_reta_end at the end. + dyn_eta_start: null + dyn_eta_end: null + + eta_retry_increment: 0.0 + # Might not do anything currently. + use_cfgpp: false pre_filter: null @@ -594,6 +629,99 @@ diffrax_g_time_scaling: false # 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 +# blep_bas sampler-specific parameters +bas: + # Batch expansion factor. Whatever your original batch size was will be multiplied + # by this. If it's 0 then you just get normal Euler. + batch_multiplier: 2 + # First 0-indexed step when BAS sampling will apply. + start_step: 0 + # Last 0-index step when BAS sampling will apply. + end_step: 3 + + s_noise: 1.0 + eta: 0.0 + eta_retry_increment: 0.0 + + # List of weights for the denoised batches, with 0 being the original denoised. + # If the list is smaller than the batch size, the list will be padded with the + # last item. + # If set to null it will be calculated automatically. + # Example: [0.5, 1.0] + # Will use weight 0.5 for the original denoised and 1.0 for any other items. + denoised_factors: null + + # If set to something other than 0 the supplied denoised_factors will be rebalanced + # to add up to this number. + denoised_factors_scale: 1.0 + + # Global multiplier on denoised for BAS steps. + denoised_multiplier: 1.0 + + # One of: restart, restart_noneta, simple + renoise_mode: restart + + # Multiplier on the start sigma for BAS steps. + # Note that taking the multipliers into account sigma_next must be less than sigma. + fromstep_factor: 1.0 + + # Multiplier on the end sigma for BAS steps. + tostep_factor: 1.0 + + # Source for the downstep. Can be one of dt, sigma or sigma_next. + # dt means you get bsigma + (sigma_next - bsigma) * tostep_factor + # where bsigma = sigma * fromstep_factor + tostep_source: dt + +# blep_weoon sampler-specific options. +# Parameters with "inv" in the name apply to the inverse wavelet operation. +# When set to null, they will use the normal setting. +weoon: + start_step: 0 + end_step: 9999 + eta: 0.0 + eta_retry_increment: 0.0 + s_noise: 1.0 + # One of dwt, dwt1d, dtcwt + wavelet_mode: dwt + # Padding scheme used for wavelets + padding: periodization + # Padding scheme used for the inverse wavelet operation + inv_padding: null + # Wavelet type. Does not apply if wavelet_mode is dtcwt. + wave: db4 + # Wavelet type used for the inverse wavelet operation. Does not apply if wavelet_mode is dtcwt. + inv_wave: null + # dtcwt qshift parameter. Only applies if the wavelet mode is dtcwt. + dtcwt_qshift: qshift_a + # dtcwt biort parameter. Only applies if the wavelet mode is dtcwt. + dtcwt_biort: near_sym_a + # dtcwt qshift parameter used for the inverse wavelet operation. Only applies if the wavelet mode is dtcwt. + dtcwt_inv_qshift: null + # dtcwt biort parameter used for the inverse wavelet operation. Only applies if the wavelet mode is dtcwt. + dtcwt_inv_biort: null + # Can be used to stretch the step down. I.E. 1.0 would be sigma -> sigma_next + # while 2.0 would be twice the distance between sigma and sigma_next. + downstep_scale: 1.0 + # Blend scale for the downstep denoised lowpass wavelets + yl_strength: 1.0 + # Blend scale for the downstep denoised highpass wavelets + yh_strength: 0.5 + # Mode used for blending wavelets. + wavelet_blend_mode: lerp + # Blend mode for wavelet highpass, uses wavelet_blend_mode if null. + wavelet_blend_mode_yh: null + # Extra multipliers that can be applied to the low/highpass wavelets for the normal + # denoised or downstep denoised. + denoised_yl_multiplier: 1.0 + denoised_yh_multiplier: 1.0 + denoised_down_yl_multiplier: 1.0 + denoised_down_yh_multiplier: 1.0 + # Only applies when wavelet_mode is dwt1d. Can be: + # 2: Flatten starting at spatial dimensions + # 1: Flatten starting at channels dimension + # 0: Smash everything together! + flatten_start_dim: 2 ### Other Sampler Specific Parameters ### @@ -606,6 +734,7 @@ diffrax_g_split_time_mode: false # ipndm: 1 (max 3) # ipndm_v: 1 (max 3) # deis: 1 (max 3) +# clybius_sens: 2 history_limit: 999 # Varies based on sampler. # Used for some samplers with variable order. List of samplers and default value below: diff --git a/__init__.py b/__init__.py index 2a1ebd9..623f457 100644 --- a/__init__.py +++ b/__init__.py @@ -9,5 +9,7 @@ NODE_CLASS_MAPPINGS = { "OCS MultiParam": nodes.MultiParamNode, "OCS ModelSetMaxSigma": nodes.ModelSetMaxSigmaNode, "OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule, + "OCS ApplyFilterLatent": nodes.ApplyFilterLatent, + "OCS ApplyFilterImage": nodes.ApplyFilterImage, } | custom_noise.NODE_CLASS_MAPPINGS __all__ = ["NODE_CLASS_MAPPINGS"] diff --git a/docs/expression.md b/docs/expression.md index 944a1f0..5801596 100644 --- a/docs/expression.md +++ b/docs/expression.md @@ -150,7 +150,9 @@ Available in model filters, with the exception of the `input` filter. **Note on `unsafe_tensor_method` and `unsafe_torch`**: These functions are disabled by default. If the environment variable `COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS` is set to anything then you can use `unsafe_tensor_method` with a whitelisted set of methods (best effort to avoid anything actually unsafe). If the environment variable `COMFYUI_OCS_ALLOW_ALL_UNSAFE` is set to anything then `unsafe_torch` is enabled and `unsafe_tensor_method` will allow calling any method. ***WARNING***: Allowing _all_ unsafe with workflows you don't trust is _not_ recommended and a malicious workflow will likely have access to anything ComfyUI can access. It is effectively the same as letting the workflow run an arbitrary script on your system. -## Tensor Expression Functions +Documentation TBD (check the source if you want to use them now): `t_copysign`, `t_gaussianblur2d`, `t_rgb_latent`, `t_snf_guidance` + +## Image Expression Functions `IMG` used here to donate the type for functions that take an image. This may actually be an image batch rather than a single image. diff --git a/py/custom_noise/__init__.py b/py/custom_noise/__init__.py index 0d3e0f0..4cfa25d 100644 --- a/py/custom_noise/__init__.py +++ b/py/custom_noise/__init__.py @@ -1,8 +1,11 @@ -from . import noise_perlin from . import nodes NODE_CLASS_MAPPINGS = { - "OCSNoise PerlinSimple": noise_perlin.PerlinSimpleNode, - "OCSNoise PerlinAdvanced": noise_perlin.PerlinAdvancedNode, + "OCSNoise PerlinSimple": nodes.PerlinSimpleNode, + "OCSNoise PerlinAdvanced": nodes.PerlinAdvancedNode, + "OCSNoise ImmiscibleReference": nodes.ImmiscibleReferenceNoiseNode, "OCSNoise to SONAR_CUSTOM_NOISE": nodes.ToSonarNode, + "OCSNoise Conditioning": nodes.NoiseConditioningNode, + "OCSNoise OverrideSamplerNoise": nodes.SamplerNodeConfigOverride, + "OCSNoise ExpressionFilteredNoise": nodes.ExpressionFilteredNoiseNode, } diff --git a/py/custom_noise/base.py b/py/custom_noise/base.py index 6b0824f..78ee4c6 100644 --- a/py/custom_noise/base.py +++ b/py/custom_noise/base.py @@ -3,7 +3,9 @@ import torch from typing import Callable, Any +from ..external import IntegratedNode from ..noise import scale_noise +from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT class CustomNoiseItemBase(abc.ABC): @@ -105,7 +107,7 @@ class CustomNoiseChain: return noise_sampler -class CustomNoiseNodeBase(abc.ABC): +class CustomNoiseNodeBase(metaclass=IntegratedNode): DESCRIPTION = "An Overly Complicated Sampling custom noise item." RETURN_TYPES = ("OCS_NOISE",) OUTPUT_TOOLTIPS = ("A custom noise chain.",) @@ -151,9 +153,9 @@ class CustomNoiseNodeBase(abc.ABC): if include_chain: result["optional"] |= { "ocs_noise_opt": ( - "OCS_NOISE", + WILDCARD_NOISE, { - "tooltip": "Optional input for more custom noise items.", + "tooltip": f"Optional input for more custom noise items.\n{NOISE_INPUT_TYPES_HINT}", }, ), } diff --git a/py/custom_noise/nodes.py b/py/custom_noise/nodes.py index 94fa31a..e04239b 100644 --- a/py/custom_noise/nodes.py +++ b/py/custom_noise/nodes.py @@ -1,3 +1,27 @@ +import functools +import inspect +import math +import torch +import yaml + +from typing import NamedTuple + +from comfy.model_management import get_torch_device +from comfy.model_patcher import set_model_options_post_cfg_function +import comfy.samplers + +from .base import CustomNoiseNodeBase, NormalizeNoiseNodeMixin, CustomNoiseItemBase +from .noise_perlin import DEFAULTS as PERLIN_DEFAULTS +from .noise_perlin import PerlinItem +from .noise_immiscibleref import ImmiscibleReferenceItem + +from ..external import MODULES, IntegratedNode +from .. import filtering +from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT +from ..utils import scale_noise +from ..noise import ImmiscibleNoise + + class ToSonarNode: RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) CATEGORY = "OveryComplicatedSampling/noise" @@ -7,10 +31,1549 @@ class ToSonarNode: def INPUT_TYPES(cls): return { "required": { - "ocs_noise": ("OCS_NOISE",), + "ocs_noise": ( + WILDCARD_NOISE, + { + "tooltip": NOISE_INPUT_TYPES_HINT, + "forceInput": True, + }, + ), }, } @classmethod def go(cls, ocs_noise): + MODULES.initialize() return (ocs_noise,) + + +class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): + DESCRIPTION = "Advanced Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinSimple node for less tuneable parameters." + + @classmethod + def INPUT_TYPES(cls): + MODULES.initialize() + result = super().INPUT_TYPES() + result["required"] |= { + "depth": ( + "INT", + { + "default": PERLIN_DEFAULTS.depth, + "tooltip": "When non-zero, 3D perlin noise will be generated.", + }, + ), + "detail_level": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.detail_level, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the detail level of the noise when break_pattern is non-zero. No effect when using 100% raw Perlin noise.", + }, + ), + "octaves": ( + "INT", + { + "default": PERLIN_DEFAULTS.octaves, + "tooltip": "Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.", + }, + ), + "persistence": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("persistence"), + "tooltip": "Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "lacunarity_height": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("lacunarity", 0), + "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "lacunarity_width": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("lacunarity", 1), + "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "lacunarity_depth": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("lacunarity", 2), + "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when depth is non-zero and octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "res_height": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("res", 0), + "tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "res_width": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("res", 1), + "tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "res_depth": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("res", 2), + "tooltip": "Number of periods of noise to generate along an axis. Only has an effect when depth is non-zero. Comma-separated list, multiple items will apply to octaves in sequence.", + }, + ), + "break_pattern": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.break_pattern, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Applies a function to break the Perlin pattern, making it more like normal noise. The value is the blend strength, where 1.0 indicates 100% pattern broken noise and 0.5 indicates 50% raw noise and 50% pattern broken noise. Generally should be at least 0.9 unless you want to generate colorful blobs.", + }, + ), + "initial_depth": ( + "INT", + { + "default": PERLIN_DEFAULTS.initial_depth, + "tooltip": "First zero-based depth index the noise generator will return. Only has an effect when depth is non-zero.", + }, + ), + "wrap_depth": ( + "INT", + { + "default": PERLIN_DEFAULTS.wrap_depth, + "tooltip": "If non-zero, instead of generating a new chunk of noise when the last slice is used will instead jump back to the specified zero-based depth index. Only has an effect when depth is non-zero.", + }, + ), + "max_depth": ( + "INT", + { + "default": PERLIN_DEFAULTS.max_depth, + "tooltip": "Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero.", + }, + ), + "tileable_height": ( + "BOOLEAN", + { + "default": PERLIN_DEFAULTS.tileable[0], + "tooltip": "Makes the specified dimension tileable.", + }, + ), + "tileable_width": ( + "BOOLEAN", + { + "default": PERLIN_DEFAULTS.tileable[1], + "tooltip": "Makes the specified dimension tileable.", + }, + ), + "tileable_depth": ( + "BOOLEAN", + { + "default": PERLIN_DEFAULTS.tileable[2], + "tooltip": "Makes the specified dimension tileable. Only has an effect when depth is non-zero.", + }, + ), + "blend": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": "lerp", + "tooltip": "Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise.", + }, + ), + "pattern_break_blend": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": "lerp", + "tooltip": "Blending function used to blend pattern broken noise with raw noise.", + }, + ), + "depth_over_channels": ( + "BOOLEAN", + { + "default": PERLIN_DEFAULTS.depth_over_channels, + "tooltip": "When disabled, each channel will have its own separate 3D noise pattern. When enabled, depth is multiplied by the number of channels and each channel is a slice of depth. Only has an effect when depth is non-zero.", + }, + ), + "pad_height": ( + "INT", + { + "default": PERLIN_DEFAULTS.pad[0], + "min": 0, + "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.", + }, + ), + "pad_width": ( + "INT", + { + "default": PERLIN_DEFAULTS.pad[1], + "min": 0, + "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.", + }, + ), + "pad_depth": ( + "INT", + { + "default": PERLIN_DEFAULTS.pad[2], + "min": 0, + "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation. Only has an effect when depth is non-zero.", + }, + ), + "initial_amplitude": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.initial_amplitude, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the amplitude for the first octave.", + }, + ), + "initial_frequency_height": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.initial_frequency[0], + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the frequency for the first octave for the this axis.", + }, + ), + "initial_frequency_width": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.initial_frequency[1], + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the frequency for the first octave for the this axis.", + }, + ), + "initial_frequency_depth": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.initial_frequency[2], + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the frequency for the first octave for the this axis.", + }, + ), + "normalize": ( + ("default", "forced", "off"), + { + "tooltip": "Controls whether the output noise is normalized after generation.", + }, + ), + "device": ( + ("default", "cpu", "gpu"), + { + "default": "default", + "tooltip": "Controls what device is used to generate the noise. GPU noise may be slightly faster but you will get different results on different GPUs.", + }, + ), + } + return result + + @classmethod + def get_item_class(cls): + return PerlinItem + + +class PerlinSimpleNode(PerlinAdvancedNode): + DESCRIPTION = "Simplified Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinAdvanced node for more tuneable parameters." + + _COPY_KEYS = { + "factor", + "rescale", + "depth", + "detail_level", + "octaves", + "persistence", + "break_pattern", + } + + @classmethod + def INPUT_TYPES(cls): + result = super().INPUT_TYPES() + orig_reqs = result["required"] + reqs = {k: v for k, v in orig_reqs.items() if k in cls._COPY_KEYS} + reqs["lacunarity"] = orig_reqs["lacunarity_height"] + reqs["res"] = orig_reqs["res_height"] + result["required"] = reqs + return result + + @classmethod + def get_item_class(cls): + def wrapper(factor, *, lacunarity, res, **kwargs): + return PerlinItem( + factor, + lacunarity_height=lacunarity, + lacunarity_width=lacunarity, + lacunarity_depth=lacunarity, + res_height=res, + res_width=res, + res_depth=res, + **kwargs, + ) + + return wrapper + + +class ImmiscibleReferenceNoiseNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): + DESCRIPTION = "Immiscible noise that uses a latent reference." + + @classmethod + def INPUT_TYPES(cls): + MODULES.initialize() + result = super().INPUT_TYPES(include_rescale=False, include_chain=False) + result["required"] |= { + "size": ( + "INT", + { + "default": 64, + "min": 0, + "tooltip": "Number of batch repeats to use when generating Immiscible noise. Setting this to 0 disables immiscible noise. If the batching type is batch, then Immiscible noise is also disabled unless the size is 2 or higher. Note that this size is in batch repeats regardless of the batching mode. For example, if you are generating a batch of 2 and you set this to 2, then you will generate noise with batch size 4.", + }, + ), + "batching": ( + ( + "channel", + "batch", + "row", + "column", + "frame", + "row_plus_column", + "channel_plus_row", + "channel_plus_column", + ), + { + "default": "channel", + "tooltip": "Dimension to maximize (or minimize) the noise with. Column mode requires reshaping the input and may require a lot of VRAM. Row mode is also fairly slow, but not as bad as column mode. Row and column modes have a very strong effect.", + }, + ), + "normalize_ref_scale": ( + "FLOAT", + { + "default": 0.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls whether the reference gets normalized. If set to 0, no normalization is done.", + }, + ), + "normalize_noise_scale": ( + "FLOAT", + { + "default": 0.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls whether the noise used as an input for immiscible noise is gets normalized first. If set to 0, no normalization is done.", + }, + ), + "maximize": ( + "BOOLEAN", + { + "default": False, + "tooltip": "When enabled, maximizes the distance between the noise and the reference rather than trying to minimize it.", + }, + ), + "distance_scale": ( + "FLOAT", + { + "default": 0.1, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier on the input noise for v2 Immiscible noise. Set to 0 to use v1 Immiscible noise.", + }, + ), + "distance_scale_ref": ( + "FLOAT", + { + "default": 0.1, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier on the refence for v2 Immiscible noise. No effect if distance_scale is 0.", + }, + ), + "blend": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Percentage of immiscible noise to use. 1.0 means 100%. May not work very well with most blend modes.", + }, + ), + "blend_mode": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": "lerp", + "tooltip": "Blending function used when mixing immiscible noise with normal noise. Only slerp seems to work well (requires ComfyUI-bleh).", + }, + ), + "normalize": ( + ("default", "forced", "disabled"), + { + "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", + }, + ), + "custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": "Input for custom noise used during ancestral or SDE sampling.", + }, + ), + } + result["optional"] = { + "reference": ( + "LATENT", + { + "tooltip": "Attach either this or the custom_noise_ref input but not both.", + }, + ), + "custom_noise_ref": ( + WILDCARD_NOISE, + { + "tooltip": "Optional input that can be attached instead of the reference latent. When used, noise from this generator will be used as the reference.", + }, + ), + "custom_noise_blend": ( + WILDCARD_NOISE, + { + "tooltip": "Optional input for blended noise (only used when blend is not 1.0). Can be used if you want to blend with a different noise type.", + }, + ), + } + return result + + @classmethod + def get_item_class(cls): + return ImmiscibleReferenceItem + + def go( + self, + *, + factor: float, + size: int, + batching: str, + normalize_ref_scale: float, + normalize_noise_scale: float, + maximize: bool, + distance_scale: float, + distance_scale_ref: float, + blend: float, + blend_mode: str, + normalize: bool | None, + custom_noise: object, + reference: dict | None = None, + custom_noise_ref: object | None = None, + custom_noise_blend: object | None = None, + ) -> tuple: + return super().go( + factor, + size=size, + batching=batching, + normalize_ref_scale=normalize_ref_scale, + normalize_noise_scale=normalize_noise_scale, + maximize=maximize, + distance_scale=distance_scale, + distance_scale_ref=distance_scale_ref, + blend=blend, + blend_function=filtering.BLENDING_MODES[blend_mode], + normalize=self.get_normalize(normalize), + noise=custom_noise.clone(), + reference=reference["samples"].clone() if reference is not None else None, + custom_noise_ref=custom_noise_ref.clone() + if custom_noise_ref is not None + else None, + custom_noise_blend=custom_noise_blend.clone() + if custom_noise_blend is not None + else None, + ) + + +class NoiseConditioningNode(metaclass=IntegratedNode): + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "go" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "conditioning": ( + "CONDITIONING", + {"tooltip": "Input conditioning to be noised."}, + ), + "seed": ( + "INT", + { + "default": 0, + "min": 0, + "max": 0xFFFFFFFFFFFFFFFF, + "tooltip": "Seed to use for generated noise.", + }, + ), + "blend": ( + "FLOAT", + { + "default": 0.01, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Blend strength of noise to be added to conditioning as a percentage where 1.0 would indicate 100%.", + }, + ), + "pooled_blend": ( + "FLOAT", + { + "default": 0.01, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Blend strength of noise to be added to pooled conditioning as a percentage where 1.0 would indicate 100%.", + }, + ), + "blend_mode": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": "inject", + "tooltip": "Blending function used when combining noise with the conditioning. inject just adds it.", + }, + ), + "pooled_blend_mode": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": "inject", + "tooltip": "Blending function used when combining noise with the pooled conditioning. inject just adds it.", + }, + ), + "noise_strength": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Strength of the generated noise to be added to conditioning.", + }, + ), + "pooled_noise_strength": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Strength of the generated noise to be added to pooled conditioning.", + }, + ), + "conditioning_multiplier": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier applied to conditioning tensors.", + }, + ), + "pooled_conditioning_multiplier": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier applied to pooled conditioning tensors.", + }, + ), + "time_mode": ( + ( + "relaxed", + "strict", + ), + { + "default": "relaxed", + "tooltip": "Controls time matching. Strict requires a conditioning item to be fully within the start/end range while relaxed just requires it to have overlap with the range. For example, if the time range is 0.2 through 0.5 and the conditioning item is 0.0 through 0.4 then strict mode would not match.", + }, + ), + "start_time": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning item start time as a percentage of sampling.", + }, + ), + "end_time": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning item end time as a percentage of sampling.", + }, + ), + "item_start": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning item start as a percentage of the total number of conditioning items.", + }, + ), + "item_end": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning item end as a percentage of the total number of conditioning items.", + }, + ), + "pooled_item_start": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning pooled output item start as a percentage of the total number of conditioning items.", + }, + ), + "pooled_item_end": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Conditioning pooled output item end as a percentage of the total number of conditioning items.", + }, + ), + "slice_start": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Noise is generated to match the total size of matched conditioning items. Slices use a percentage of that chunk of noise.", + }, + ), + "slice_end": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Noise is generated to match the total size of matched conditioning items. Slices use a percentage of that chunk of noise.", + }, + ), + "pooled_slice_start": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Noise is generated to match the total size of matched conditioning items. Slices use a percentage of that chunk of noise.", + }, + ), + "pooled_slice_end": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "tooltip": "Noise is generated to match the total size of matched conditioning items. Slices use a percentage of that chunk of noise.", + }, + ), + "cpu_noise": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether noise will be generated on GPU or CPU. Only affects noise types that support GPU generation.", + }, + ), + "normalize": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.", + }, + ), + "fake_channels": ( + "INT", + { + "default": 1, + "min": 1, + "tooltip": "Noise will be generated with number of channels. Shouldn't make a difference for most noise types.", + }, + ), + }, + "optional": { + "custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": "Custom noise type to use. If not connected, gaussian noise will be used." + }, + ), + }, + } + + @classmethod + def go( + cls, + *, + conditioning, + seed, + noise_strength, + pooled_noise_strength, + blend, + pooled_blend, + conditioning_multiplier, + pooled_conditioning_multiplier, + blend_mode, + pooled_blend_mode, + start_time, + end_time, + item_start, + item_end, + pooled_item_start, + pooled_item_end, + slice_start, + slice_end, + pooled_slice_start, + pooled_slice_end, + time_mode, + cpu_noise, + normalize, + fake_channels, + custom_noise=None, + ): + MODULES.initialize() + blend_function = filtering.BLENDING_MODES[blend_mode] + pblend_function = filtering.BLENDING_MODES[pooled_blend_mode] + noise_spatdim_min = 4 * fake_channels + size = psize = 0 + count = pcount = 0 + to_noise = [] + for cond, opts, *_ in conditioning: + stime, etime = opts.get("start_percent", 0.0), opts.get("end_percent", 1.0) + pooled = opts.get("pooled_output") + if time_mode == "relaxed": + time_ok = (start_time <= stime <= end_time) or ( + start_time <= etime <= end_time + ) + else: + time_ok = stime >= start_time and etime <= end_time + if not time_ok: + to_noise.append((False, False, False)) + continue + need_cond = noise_strength != 0 and blend != 0 + if need_cond: + size += cond.numel() + count += 1 + need_pooled = ( + pooled_noise_strength != 0 and pooled_blend != 0 and pooled is not None + ) + if need_pooled: + psize += pooled.numel() + pcount += 1 + to_noise.append((True, need_cond, need_pooled)) + conds_size = size + psize + # print("GOT", size, psize, "-->", conds_size, "::", count, pcount) + if conds_size != 0: + noise_spatdim = max( + noise_spatdim_min, math.ceil((conds_size // fake_channels) ** 0.5) + ) + empty_ref = torch.zeros( + 1, + fake_channels, + noise_spatdim, + noise_spatdim, + dtype=torch.float, + device="cpu" if cpu_noise else get_torch_device(), + ) + if custom_noise is not None: + ns = custom_noise.make_noise_sampler( + empty_ref, + sigma_min=None, + sigma_max=None, + seed=seed, + cpu=cpu_noise, + normalized=False, + ) + else: + + def ns(*_unusedargs, **_unusedkwargs): + return torch.randn_like(empty_ref) + + randst = torch.random.get_rng_state() + try: + torch.random.manual_seed(seed) + noise = ns(None, None) + finally: + torch.random.set_rng_state(randst) + noise = noise.reshape(noise.numel()) + noise_conds = noise.new_zeros(size) + noise_conds[int(size * slice_start) : math.ceil(size * slice_end)] = ( + noise_strength + ) + noise_conds *= scale_noise( + noise[:size], + normalized=normalize, + normalize_dims=None, + ) + + noise_pooled = noise.new_zeros(psize) + noise_pooled[ + int(psize * pooled_slice_start) : math.ceil(psize * pooled_slice_end) + ] = pooled_noise_strength + noise_pooled *= scale_noise( + noise[size : size + psize], + normalized=normalize, + normalize_dims=None, + ) + # print("MADE NOISE", noise.shape, noise_conds.shape, noise_pooled.shape) + del noise + + result = [] + currc = currp = 0 + for (time_matched, need_cond, need_pooled), (cond, opts, *_) in zip( + to_noise, conditioning + ): + # print(">> ITER", currc, currp) + opts = opts.copy() + pooled = opts["pooled_output"] + if time_matched: + if conditioning_multiplier != 1: + cond = cond * conditioning_multiplier + if pooled is not None and pooled_conditioning_multiplier != 1: + pooled = pooled * pooled_conditioning_multiplier + if need_cond: + cpct = currc / count + currc += 1 + if item_start <= cpct <= item_end: + # print("COND MATCH", offset, cond.shape) + cond = blend_function( + cond, + noise_conds[: cond.numel()].reshape(cond.shape).to(cond), + blend, + ) + noise_conds = noise_conds[cond.numel() :] + if need_pooled: + ppct = currp / pcount + currp += 1 + if pooled_item_start <= ppct <= pooled_item_end: + # print("POOLED MATCH", offset, pooled.shape) + pooled = pblend_function( + pooled, + noise_pooled[: pooled.numel()].reshape(pooled.shape).to(pooled), + pooled_blend, + ) + noise_pooled = noise_pooled[pooled.numel() :] + if pooled is not None: + opts["pooled_output"] = pooled + result.append([cond, opts]) + return (result,) + + +class ImmiscibleConfig(NamedTuple): + size: int = 64 + batching: str = "channel" + reference: str = "cond" + norm_ref_scale: float = 0.0 + norm_noise_scale: float = 0.0 + maximize: bool = False + distance_scale: float = 0.1 + distance_scale_ref: float = 0.1 + start_time: float = 0.0 + end_time: float = 1.0 + blend: float = 1.0 + blend_mode: str = "lerp" + + +DEFAULT_IMMISCIBLE_CONFIG = ImmiscibleConfig() + + +class OverrideSamplerConfig(NamedTuple): + sampler: object | None = None + sampler_kwargs: dict | None = None + noise_start_time: float = 0.0 + noise_end_time: float = 1.0 + cpu_noise: bool = True + normalize: bool = True + force_params: list | tuple = () + custom_noise: object | None = None + custom_noise_ref: object | None = None + custom_noise_blend: object | None = None + latent_ref: torch.Tensor | None = None + immiscible: ImmiscibleConfig = DEFAULT_IMMISCIBLE_CONFIG + + +DEFAULT_OVERRIDESAMPLER_CONFIG = OverrideSamplerConfig() + + +class SamplerNodeConfigOverride(metaclass=IntegratedNode): + DESCRIPTION = "Allows overriding parameters of a SAMPLER, and specifically lets you use custom and/or Immiscible noise. For use with non-OCS samplers, not recommended to use with the OCS Sampler node as it has internal support for Immiscible noise." + + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + @classmethod + def INPUT_TYPES(cls): + DICFG = DEFAULT_IMMISCIBLE_CONFIG + DOCFG = DEFAULT_OVERRIDESAMPLER_CONFIG + return { + "required": { + "sampler": ( + "SAMPLER", + { + "tooltip": "Sampler to wrap with custom noise handling/parameter overrides." + }, + ), + "immiscible_size": ( + "INT", + { + "default": DICFG.size, + "min": 0, + "tooltip": "Number of batch repeats to use when generating Immiscible noise. Setting this to 0 disables immiscible noise. If the batching type is batch, then Immiscible noise is also disabled unless the size is 2 or higher. Note that this size is in batch repeats regardless of the batching mode. For example, if you are generating a batch of 2 and you set this to 2, then you will generate noise with batch size 4.", + }, + ), + "immiscible_batching": ( + ( + "channel", + "batch", + "row", + "column", + "frame", + "cycle_channel_batch", + "cycle_row_column", + "cycle_channel_row", + "cycle_channel_column", + ), + { + "default": DICFG.batching, + "tooltip": "Dimension to maximize (or minimize) the noise with. Column mode requires reshaping the input and may require a lot of VRAM. Row mode is also fairly slow, but not as bad as column mode. Row and column modes have a very strong effect.", + }, + ), + "immiscible_reference": ( + ( + "cond", + "uncond", + "denoised", + "model_input", + "noise_prediction", + "uncond_sub_cond", + "cond_sub_uncond", + "latent", + "noise", + ), + { + "default": DICFG.reference, + "tooltip": "Reference type to use when generating immiscible noise.\ncond - positive prompt.\nuncond - negative prompt.\ndenoised - the model's prediction of a clean image (using both cond and uncond).\nmodel_input - noisy latent image the model was called with (often referred to as x).\nnoise_prediction - model_input with denoised subtracted (leaving just what the model thinks is noise).", + }, + ), + "immiscible_normalize_ref_scale": ( + "FLOAT", + { + "default": DICFG.norm_ref_scale, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls whether the reference gets normalized. If set to 0, no normalization is done.", + }, + ), + "immiscible_normalize_noise_scale": ( + "FLOAT", + { + "default": DICFG.norm_noise_scale, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls whether the noise used as an input for immiscible noise is gets normalized first. If set to 0, no normalization is done.", + }, + ), + "immiscible_maximize": ( + "BOOLEAN", + { + "default": DICFG.maximize, + "tooltip": "When enabled, maximizes the distance between the noise and the reference rather than trying to minimize it.", + }, + ), + "immiscible_distance_scale": ( + "FLOAT", + { + "default": DICFG.distance_scale, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier on the input noise for v2 Immiscible noise. Set to 0 to use v1 Immiscible noise.", + }, + ), + "immiscible_distance_scale_ref": ( + "FLOAT", + { + "default": DICFG.distance_scale_ref, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier on the refence for v2 Immiscible noise. No effect if immiscible_distance_scale is 0.", + }, + ), + "immiscible_blend": ( + "FLOAT", + { + "default": DICFG.blend, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Percentage of immiscible noise to use. 1.0 means 100%. May not work very well with most blend modes.", + }, + ), + "immiscible_blend_mode": ( + tuple(filtering.BLENDING_MODES.keys()), + { + "default": DICFG.blend_mode, + "tooltip": "Blending function used when mixing immiscible noise with normal noise. Only slerp seems to work well (requires ComfyUI-bleh).", + }, + ), + "immiscible_start_time": ( + "FLOAT", + { + "default": DICFG.start_time, + "min": 0.0, + "max": 1.0, + "tooltip": "Start time as a percentage of sampling where immiscible noise will be used.", + }, + ), + "immiscible_end_time": ( + "FLOAT", + { + "default": DICFG.end_time, + "min": 0.0, + "max": 1.0, + "tooltip": "End time as a percentage of sampling where immiscible noise will be used.", + }, + ), + "noise_start_time": ( + "FLOAT", + { + "default": DOCFG.noise_start_time, + "min": 0.0, + "max": 1.0, + "tooltip": "Start time as a percentage of sampling where custom noise will be used.", + }, + ), + "noise_end_time": ( + "FLOAT", + { + "default": DOCFG.noise_end_time, + "min": 0.0, + "max": 1.0, + "tooltip": "End time as a percentage of sampling where custom noise will be used.", + }, + ), + "cpu_noise": ( + "BOOLEAN", + { + "default": DOCFG.cpu_noise, + "tooltip": "Controls whether noise is generated on CPU or GPU. Only affects custom noise.", + }, + ), + "normalize": ( + "BOOLEAN", + { + "default": DOCFG.normalize, + "tooltip": "Controls whether generated noise is normalized to 1.0 strength. This normalization occurs last.", + }, + ), + }, + "optional": { + "custom_noise_opt": ( + WILDCARD_NOISE, + { + "tooltip": "Optional input for custom noise used during ancestral or SDE sampling.", + }, + ), + "yaml_parameters": ( + "STRING", + { + "tooltip": "Allows specifying custom parameters via YAML. This input can be converted to a multiline text widget. Note: When specifying parameters this way, there is no error checking.", + "placeholder": "# YAML or JSON here", + "dynamicPrompts": False, + "multiline": True, + "defaultInput": True, + }, + ), + "custom_noise_ref": ( + WILDCARD_NOISE, + { + "tooltip": "Required when reference mode is reference_noise, otherwise unused.", + }, + ), + "latent_ref": ( + "LATENT", + { + "tooltip": "Required when the reference mode is reference_latent, otherwise unused. This doesn't respect stuff like latent from batch that may have been previously applied and must match the shape of the latent being generated.", + }, + ), + "custom_noise_blend": ( + WILDCARD_NOISE, + { + "tooltip": "Optional input for blended noise (only used when blend is not 1.0). Can be used if you want to blend with a different noise type.", + }, + ), + }, + } + + def get_sampler( + self, + *, + sampler, + immiscible_size, + immiscible_batching, + immiscible_reference, + immiscible_normalize_ref_scale, + immiscible_normalize_noise_scale, + immiscible_maximize, + immiscible_distance_scale, + immiscible_distance_scale_ref, + immiscible_start_time, + immiscible_end_time, + immiscible_blend, + immiscible_blend_mode, + noise_start_time, + noise_end_time, + cpu_noise=True, + normalize=True, + yaml_parameters="", + custom_noise_opt=None, + custom_noise_ref=None, + custom_noise_blend=None, + latent_reference: dict | None = None, + ): + MODULES.initialize() + sampler_kwargs = {} + overridecfg_kwargs = { + "noise_start_time": noise_start_time, + "noise_end_time": noise_end_time, + "cpu_noise": cpu_noise, + "normalize": normalize, + } + immisciblecfg_kwargs = { + "size": immiscible_size, + "batching": immiscible_batching, + "reference": immiscible_reference, + "norm_ref_scale": immiscible_normalize_ref_scale, + "norm_noise_scale": immiscible_normalize_noise_scale, + "maximize": immiscible_maximize, + "distance_scale": immiscible_distance_scale, + "distance_scale_ref": immiscible_distance_scale_ref, + "start_time": immiscible_start_time, + "end_time": immiscible_end_time, + "blend": immiscible_blend, + "blend_mode": immiscible_blend_mode, + } + if yaml_parameters: + extra_params = yaml.safe_load(yaml_parameters) + if extra_params is None: + pass + elif not isinstance(extra_params, dict): + raise ValueError( + "SamplerConfigOverride: yaml_parameters must either be null or an object", + ) + else: + ocs_extra = extra_params.pop("ocs", {}) + override_extra = ocs_extra.pop("override", {}) + immiscible_extra = override_extra.pop("immiscible", {}) + overridecfg_kwargs |= override_extra + immisciblecfg_kwargs |= immiscible_extra + sampler_kwargs |= extra_params + if immiscible_reference == "latent": + if latent_reference is None: + raise ValueError( + "latent_ref input must be connected when reference mode is latent" + ) + overridecfg_kwargs["latent_ref"] = latent_reference["samples"].to( + dtype=torch.float32, device="cpu", copy=True + ) + elif immiscible_reference == "noise": + if custom_noise_ref is None: + raise ValueError( + "custom_noise_ref input must be connected when reference mode is noise" + ) + overridecfg_kwargs["custom_noise_ref"] = custom_noise_ref.clone() + + sampler_function = functools.partial( + self.sampler_function, + ocs_override_sampler_cfg=OverrideSamplerConfig( + sampler=sampler, + sampler_kwargs=sampler_kwargs, + immiscible=ImmiscibleConfig(**immisciblecfg_kwargs), + custom_noise=custom_noise_opt.clone() + if custom_noise_opt is not None + else None, + custom_noise_blend=custom_noise_blend.clone() + if custom_noise_blend is not None + else None, + **overridecfg_kwargs, + ), + ) + functools.update_wrapper(sampler_function, sampler.sampler_function) + return ( + comfy.samplers.KSAMPLER( + sampler_function, + extra_options=sampler.extra_options.copy(), + inpaint_options=sampler.inpaint_options.copy(), + ), + ) + + @staticmethod + @torch.no_grad() + def sampler_function( + model, + x, + sigmas, + *args: list, + ocs_override_sampler_cfg: dict[str] | None = None, + noise_sampler=None, + extra_args: dict[str] | None = None, + **kwargs: dict[str], + ): + cfg = ocs_override_sampler_cfg + if cfg is None: + raise ValueError("Override sampler config missing!") + icfg = cfg.immiscible + if extra_args is None: + extra_args = {} + sig = inspect.signature(cfg.sampler.sampler_function) + params = frozenset(cfg.force_params) | sig.parameters.keys() + kwargs |= {k: v for k, v in cfg.sampler_kwargs.items() if k in params} + if "noise_sampler" in params: + seed = extra_args.get("seed") + seed_gen = torch.Generator(device="cpu" if cfg.cpu_noise else x.device) + seed_gen.manual_seed(seed if seed is not None else 0) + model_sampling = model.inner_model.inner_model.model_sampling + orig_noise_sampler = kwargs.pop( + "noise_sampler", lambda *_args, **_kwargs: torch.randn_like(x) + ) + if ( + cfg.custom_noise is not None + and cfg.noise_start_time < 1 + and cfg.noise_end_time > 0 + ): + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + custom_noise_sampler = cfg.custom_noise.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=seed, + cpu=cfg.cpu_noise, + normalized=False, + ) + else: + custom_noise_sampler = None + if custom_noise_sampler is not None: + sigma_start = model_sampling.percent_to_sigma(cfg.noise_start_time) + sigma_end = model_sampling.percent_to_sigma(cfg.noise_end_time) + + def override_noise_sampler(s, sn, *args, **kwargs): + if not sigma_end <= s.max() <= sigma_start: + return orig_noise_sampler(s, sn, *args, **kwargs) + return custom_noise_sampler(s, sn, *args, **kwargs) + + else: + override_noise_sampler = orig_noise_sampler + if ( + icfg.start_time < 1 + and icfg.end_time > 0 + and (icfg.size > 1 if icfg.batching == "batch" else icfg.size > 0) + and icfg.blend != 0 + ): + isigma_start = model_sampling.percent_to_sigma(icfg.start_time) + isigma_end = model_sampling.percent_to_sigma(icfg.end_time) + if icfg.batching.startswith("cycle_"): + ibatching = icfg.batching.split("_")[1:] + else: + ibatching = icfg.batching + immiscible = ImmiscibleNoise( + size=icfg.size, + batching=ibatching if isinstance(ibatching, str) else "channel", + maximize=icfg.maximize, + distance_scale=icfg.distance_scale, + distance_scale_ref=icfg.distance_scale_ref, + ) + blend_function = filtering.BLENDING_MODES[icfg.blend_mode] + + ref_latent = None + + requires_patch = icfg.reference not in {"latent", "noise"} + + if requires_patch: + requires_uncond = icfg.reference in {"uncond", "cond_sub_uncond"} + ref_handlers = { + "cond": "cond_denoised", + "uncond": "uncond_denoised", + "denoised": "denoised", + "model_input": "input", + "noise_prediction": lambda args: args["input"] + - args["denoised"], + "cond_sub_uncond": lambda args: args["cond_denoised"] + - args["uncond_denoised"], + "uncond_sub_cond": lambda args: args["uncond_denoised"] + - args["cond_denoised"], + } + if (ref_handler := ref_handlers.get(icfg.reference)) is None: + raise ValueError("Bad immiscible reference type") + + def postcfg(args): + nonlocal ref_latent + ref_latent = ( + args.get(ref_handler) + if isinstance(ref_handler, str) + else ref_handler(args) + ) + return args["denoised"] + + extra_args = extra_args | { + "model_options": set_model_options_post_cfg_function( + extra_args.get("model_options", {}).copy(), + postcfg, + disable_cfg1_optimization=requires_uncond, + ) + } + + if icfg.reference == "noise": + ns_ref = cfg.custom_noise_ref.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=torch.randint( + 0, + 1 << 32, + (1,), + device="cpu" if cfg.cpu_noise else x.device, + dtype=torch.int64, + generator=seed_gen, + ) + .detach() + .cpu() + .item(), + cpu=cfg.cpu_noise, + normalized=False, + ) + elif icfg.reference == "latent": + ref_latent = cfg.latent_ref.to( + dtype=x.dtype, device=x.device, copy=True + ) + if icfg.norm_ref_scale != 0: + ref_latent = scale_noise( + ref_latent, icfg.norm_ref_scale, normalized=True + ) + if ref_latent.shape[1:] != x.shape[1:]: + raise ValueError( + "Reference latent shape must match shape of generation with exception that batch size may be 1", + ) + if ref_latent.shape[0] != x.shape[0]: + if ref_latent.shape[0] == 1: + ref_latent = ref_latent.expand(x.shape) + else: + raise ValueError( + "Reference latent batch size must be either 1 or equal to generation batch size" + ) + + if icfg.blend != 1 and cfg.custom_noise_blend is not None: + ns_blend = cfg.custom_noise_blend.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=torch.randint( + 0, + 1 << 32, + (1,), + device="cpu" if cfg.cpu_noise else x.device, + dtype=torch.int64, + generator=seed_gen, + ) + .detach() + .cpu() + .item(), + cpu=cfg.cpu_noise, + normalized=False, + ) + else: + ns_blend = None + + immiscible_counter = 0 + + def noise_sampler(s, sn, *args, **kwargs): + nonlocal ref_latent, immiscible_counter + if not isigma_end <= s.max() <= isigma_start: + return override_noise_sampler(s, sn, *args, **kwargs) + if icfg.reference == "noise": + ref_latent = ns_ref(s, sn) + elif ref_latent is None: + raise ValueError("Immiscible reference type not available") + if not isinstance(ibatching, str): + immiscible.batching = ibatching[ + immiscible_counter % len(ibatching) + ] + immiscible_counter += 1 + blend_in_batch = icfg.blend != 1 and ns_blend is None + # blending = icfg.blend != 1 + if icfg.norm_ref_scale != 0 and icfg.reference != "latent": + ref_latent = scale_noise( + ref_latent, icfg.norm_ref_scale, normalized=True + ) + batch_size = ref_latent.shape[0] + noise_batch = torch.cat( + tuple( + override_noise_sampler(s, sn) + for _ in range(max(1, icfg.size) + int(blend_in_batch)) + ) + ) + immiscible_noise = immiscible.unbatch( + immiscible.immiscible( + immiscible.batch( + scale_noise( + noise_batch[batch_size * int(blend_in_batch) :], + 1.0 + if icfg.norm_noise_scale == 0 + else icfg.norm_noise_scale, + normalized=icfg.norm_noise_scale != 0, + ) + ), + immiscible.batch(ref_latent), + ), + ref_latent.shape, + ) + immiscible_noise = scale_noise( + immiscible_noise, normalized=cfg.normalize + ) + if icfg.blend != 1: + immiscible_noise = blend_function( + noise_batch[:batch_size] + if blend_in_batch + else ns_blend(s, sn), + immiscible_noise, + icfg.blend, + ) + return scale_noise(immiscible_noise, normalized=cfg.normalize) + + else: + if not cfg.normalize: + noise_sampler = override_noise_sampler + else: + + def noise_sampler(s, sn, *args, **kwargs): + return scale_noise( + override_noise_sampler(s, sn, *args, **kwargs), + normalized=True, + ) + + kwargs["noise_sampler"] = noise_sampler + return cfg.sampler.sampler_function( + model, + x, + sigmas, + *args, + extra_args=extra_args, + **kwargs, + ) + + +class ExpressionFilteredNoiseItem(CustomNoiseItemBase): + def __init__( + self, + factor, + *, + normalize, + noise: object, + noise_filter, + ): + super().__init__( + factor, + noise=noise, + noise_filter=noise_filter, + normalize=normalize, + ) + + def clone_key(self, k): + if k == "noise": + return self.noise.clone() + return super().clone_key(k) + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + noise_filter = self.noise_filter + ns = self.noise.make_noise_sampler(x, *args, normalized=False, **kwargs) + normalize_noise = self.normalize != False and normalized # noqa: E712 + factor = self.factor + + def noise_sampler(s, sn, *args, **kwargs): + noise = ns(s, sn) + refs = filtering.FilterRefs({ + "sigma": s.clone() if isinstance(s, torch.Tensor) else s, + "sigma_next": sn.clone() if isinstance(sn, torch.Tensor) else sn, + }) + noise = noise_filter.apply(noise, refs=refs) + return scale_noise(noise, factor, normalized=normalize_noise) + + return noise_sampler + + +class ExpressionFilteredNoiseNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): + DESCRIPTION = "Immiscible noise that uses a latent reference." + + @classmethod + def INPUT_TYPES(cls): + MODULES.initialize() + result = super().INPUT_TYPES() + result["required"] |= { + "normalize": ( + ("default", "forced", "disabled"), + { + "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", + }, + ), + "custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": "Input for custom noise used during ancestral or SDE sampling.", + }, + ), + "yaml_config": ( + "STRING", + { + "default": "", + "placeholder": """\ +# YAML or JSON filter definition +""", + "multiline": True, + "dynamicPrompts": False, + "tooltip": "Enter your filter definition here. There is essentially no error handling.", + }, + ), + } + return result + + @classmethod + def get_item_class(cls): + return ExpressionFilteredNoiseItem + + def go( + self, + *, + factor: float, + rescale: float, + normalize: str, + custom_noise: object, + yaml_config: str, + ) -> tuple: + config = yaml.safe_load(yaml_config) + if not isinstance(config, dict) or "filter" not in config: + raise ValueError( + "Bad YAML config type (must be object) or missing filter key in config" + ) + filter_def = config.get("filter") + if not isinstance(filter_def, dict): + raise ValueError("Bad type for filter definition, must be object") + ocs_filter = filtering.make_filter(filter_def) + return super().go( + factor, + rescale=rescale, + normalize=self.get_normalize(normalize), + noise=custom_noise.clone(), + noise_filter=ocs_filter, + ) diff --git a/py/custom_noise/noise_immiscibleref.py b/py/custom_noise/noise_immiscibleref.py new file mode 100644 index 0000000..ccd00ed --- /dev/null +++ b/py/custom_noise/noise_immiscibleref.py @@ -0,0 +1,156 @@ +import torch + +from ..noise import scale_noise, ImmiscibleNoise +from .base import CustomNoiseItemBase + + +class ImmiscibleReferenceItem(CustomNoiseItemBase): + def __init__( + self, + factor, + *, + size: int, + batching: str, + normalize_ref_scale: float, + normalize_noise_scale: float, + maximize: bool, + distance_scale: float, + distance_scale_ref: float, + blend: float, + noise, + blend_function, + normalize=None, + reference=None, + custom_noise_blend=None, + custom_noise_ref=None, + ): + if reference is None and custom_noise_ref is None: + raise ValueError( + "Either the reference latent or custom_noise_ref need to be supplied." + ) + if reference is not None and custom_noise_ref is not None: + raise ValueError( + "One of reference latent or custom_noise_ref need to be supplied, but not both." + ) + super().__init__( + factor, + size=size, + batching=batching, + normalize_ref_scale=normalize_ref_scale, + normalize_noise_scale=normalize_noise_scale, + maximize=maximize, + distance_scale=distance_scale, + distance_scale_ref=distance_scale_ref, + blend=blend, + blend_function=blend_function, + noise=noise, + reference=reference, + normalize=normalize, + custom_noise_ref=custom_noise_ref, + custom_noise_blend=custom_noise_blend, + ) + + def clone_key(self, k): + if k == "noise": + return self.noise.clone() + if k == "reference" and self.reference is not None: + return self.reference.clone() + if k == "custom_noise_ref" and self.custom_noise_ref is not None: + return self.custom_noise_ref.clone() + if k == "custom_noise_blend" and self.custom_noise_blend is not None: + return self.custom_noise_blend.clone() + return super().clone_key(k) + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + factor = self.factor + norm_noise_scale = self.normalize_noise_scale + normalize = self.get_normalize("normalize", normalized) + norm_ref = self.normalize_ref_scale + + ns = self.noise.make_noise_sampler(x, *args, normalized=False, **kwargs) + batching = self.batching + if "_" in batching: + batchings = batching.split("_") + batching = batchings[0] + dual_mode = True + else: + dual_mode = False + immiscible = ImmiscibleNoise( + size=self.size, + batching=batching, + distance_scale=self.distance_scale, + distance_scale_ref=self.distance_scale_ref, + maximize=self.maximize, + ) + if dual_mode: + immiscible2 = immiscible = ImmiscibleNoise( + size=self.size, + batching=batchings[-1], + distance_scale=self.distance_scale, + distance_scale_ref=self.distance_scale_ref, + maximize=self.maximize, + ) + if self.reference is not None: + ns_ref = None + ref_latent = self.reference.detach().clone().to(x) + if norm_ref: + ref_latent = scale_noise(ref_latent, norm_ref, normalized=True) + else: + ref_latent = None + ns_ref = self.custom_noise_ref.make_noise_sampler( + x, *args, normalized=False, **kwargs + ) + blend = self.blend + blend_function = self.blend_function + blending = self.blend != 1.0 + if self.custom_noise_blend is not None and blending: + ns_blend = self.custom_noise_blend.make_noise_sampler( + x, *args, normalized=False, **kwargs + ) + else: + ns_blend = None + repeat_count = max(1, self.size) + int(blending) + batch_size = x.shape[0] + blend_in_batch = blending and ns_blend is None + + def noise_sampler(s, sn, *args, **kwargs): + if ns_ref is not None: + ref_latent = ns_ref(s, sn) + if norm_ref: + ref_latent = scale_noise(ref_latent, norm_ref, normalized=True) + + noise_batch = torch.cat(tuple(ns(s, sn) for _ in range(repeat_count))) + nb_input = scale_noise( + noise_batch[batch_size * int(blend_in_batch) :], + 1.0 if norm_noise_scale == 0 else norm_noise_scale, + normalized=norm_noise_scale != 0, + ) + immiscible_noise = immiscible.unbatch( + immiscible.immiscible( + immiscible.batch(nb_input), + immiscible.batch(ref_latent), + ), + ref_latent.shape, + ) + if dual_mode: + immiscible_noise = ( + immiscible2.unbatch( + immiscible2.immiscible( + immiscible2.batch(nb_input), + immiscible2.batch(ref_latent), + ), + ref_latent.shape, + ) + .add_(immiscible_noise) + .mul_(0.5) + ) + immiscible_noise = scale_noise(immiscible_noise, normalized=normalize) + if blend != 1: + immiscible_noise = blend_function( + noise_batch[:batch_size] if ns_blend is None else ns_blend(s, sn), + immiscible_noise, + blend, + ) + return scale_noise(immiscible_noise, factor, normalized=normalize) + + return noise_sampler diff --git a/py/custom_noise/noise_perlin.py b/py/custom_noise/noise_perlin.py index ff93399..8ff023e 100644 --- a/py/custom_noise/noise_perlin.py +++ b/py/custom_noise/noise_perlin.py @@ -1,11 +1,14 @@ -import torch -import math import itertools +import math -from .base import CustomNoiseItemBase, CustomNoiseNodeBase, NormalizeNoiseNodeMixin +import torch + +from comfy import model_management + +from .. import filtering from ..latent import normalize_to_scale from ..noise import scale_noise -from ..filtering import BLENDING_MODES +from .base import CustomNoiseItemBase, NormalizeNoiseNodeMixin # Perlin generation routines based on https://github.com/Extraltodeus/noise_latent_perlinpinpin which was based on https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57 which was based on https://github.com/pvigier/perlin-numpy @@ -50,13 +53,13 @@ def rand_perlin( *, tileable=DEFAULTS.tileable, fade=DEFAULTS.fade, - blend=BLENDING_MODES[DEFAULTS.blend], + blend=torch.lerp, generator=DEFAULTS.generator, device=DEFAULTS.device, ): dims = len(res) didxs = tuple(range(dims)) - delta, d = zip(*((res[i] / shape[i], int(shape[i] // res[i])) for i in didxs)) + delta, d = zip(*((res[i] / shape[i], int(round(shape[i] / res[i]))) for i in didxs)) grid = ( torch.stack( @@ -163,7 +166,7 @@ def generate_fractal_noise( initial_frequency=DEFAULTS.initial_frequency, tileable=DEFAULTS.tileable, fade=DEFAULTS.fade, - blend=BLENDING_MODES[DEFAULTS.blend], + blend=torch.lerp, generator=DEFAULTS.generator, device=DEFAULTS.device, ): @@ -231,8 +234,8 @@ def create_noisy_latents_perlin( res=DEFAULTS.res, break_pattern=DEFAULTS.break_pattern, channels=4, - blend=BLENDING_MODES[DEFAULTS.blend], - pattern_break_blend=BLENDING_MODES[DEFAULTS.pattern_break_blend], + blend=torch.lerp, + pattern_break_blend=torch.lerp, depth_over_channels=DEFAULTS.depth_over_channels, pad=DEFAULTS.pad, initial_frequency=DEFAULTS.initial_frequency, @@ -426,15 +429,20 @@ class PerlinItem(CustomNoiseItemBase): ): normalized = self.get_normalize("normalized", normalized) cpu = cpu if self.device == "default" else self.device == "cpu" - device = torch.device(0 if not cpu else "cpu") + device = "cpu" if cpu else model_management.get_torch_device() noise_chunk = None noise_index = self.initial_depth max_idx = None - b, c, h, w = x.shape + if x.ndim < 4: + raise ValueError("Can only handle latents with 4+ dimensions") + orig_shape = x.shape + b = x.shape[0] + c = math.prod(x.shape[1:-2]) # Hack to deal with video models. + h, w = x.shape[-2:] x_device, x_dtype = x.device, x.dtype del x - blend = BLENDING_MODES[self.blend] - pattern_break_blend = BLENDING_MODES[self.pattern_break_blend] + blend = filtering.BLENDING_MODES[self.blend] + pattern_break_blend = filtering.BLENDING_MODES[self.pattern_break_blend] def noise_sampler(_s, _sn): nonlocal noise_chunk, noise_index, max_idx @@ -481,266 +489,7 @@ class PerlinItem(CustomNoiseItemBase): noise_index = 0 if not self.wrap_depth: noise_chunk = None - return scale_noise(noise, self.factor, normalized=normalized) + result = scale_noise(noise, self.factor, normalized=normalized) + return result.reshape(orig_shape) if result.shape != orig_shape else result return noise_sampler - - -class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): - DESCRIPTION = "Advanced Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinSimple node for less tuneable parameters." - - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "depth": ( - "INT", - { - "default": DEFAULTS.depth, - "tooltip": "When non-zero, 3D perlin noise will be generated.", - }, - ), - "detail_level": ( - "FLOAT", - { - "default": DEFAULTS.detail_level, - "tooltip": "Controls the detail level of the noise when break_pattern is non-zero. No effect when using 100% raw Perlin noise.", - }, - ), - "octaves": ( - "INT", - { - "default": DEFAULTS.octaves, - "tooltip": "Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.", - }, - ), - "persistence": ( - "STRING", - { - "default": DEFAULTS.get_commasep("persistence"), - "tooltip": "Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "lacunarity_height": ( - "STRING", - { - "default": DEFAULTS.get_commasep("lacunarity", 0), - "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "lacunarity_width": ( - "STRING", - { - "default": DEFAULTS.get_commasep("lacunarity", 1), - "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "lacunarity_depth": ( - "STRING", - { - "default": DEFAULTS.get_commasep("lacunarity", 2), - "tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when depth is non-zero and octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "res_height": ( - "STRING", - { - "default": DEFAULTS.get_commasep("res", 0), - "tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "res_width": ( - "STRING", - { - "default": DEFAULTS.get_commasep("res", 1), - "tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "res_depth": ( - "STRING", - { - "default": DEFAULTS.get_commasep("res", 2), - "tooltip": "Number of periods of noise to generate along an axis. Only has an effect when depth is non-zero. Comma-separated list, multiple items will apply to octaves in sequence.", - }, - ), - "break_pattern": ( - "FLOAT", - { - "default": DEFAULTS.break_pattern, - "tooltip": "Applies a function to break the Perlin pattern, making it more like normal noise. The value is the blend strength, where 1.0 indicates 100% pattern broken noise and 0.5 indicates 50% raw noise and 50% pattern broken noise. Generally should be at least 0.9 unless you want to generate colorful blobs.", - }, - ), - "initial_depth": ( - "INT", - { - "default": DEFAULTS.initial_depth, - "tooltip": "First zero-based depth index the noise generator will return. Only has an effect when depth is non-zero.", - }, - ), - "wrap_depth": ( - "INT", - { - "default": DEFAULTS.wrap_depth, - "tooltip": "If non-zero, instead of generating a new chunk of noise when the last slice is used will instead jump back to the specified zero-based depth index. Only has an effect when depth is non-zero.", - }, - ), - "max_depth": ( - "INT", - { - "default": DEFAULTS.max_depth, - "tooltip": "Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero.", - }, - ), - "tileable_height": ( - "BOOLEAN", - { - "default": DEFAULTS.tileable[0], - "tooltip": "Makes the specified dimension tileable.", - }, - ), - "tileable_width": ( - "BOOLEAN", - { - "default": DEFAULTS.tileable[1], - "tooltip": "Makes the specified dimension tileable.", - }, - ), - "tileable_depth": ( - "BOOLEAN", - { - "default": DEFAULTS.tileable[2], - "tooltip": "Makes the specified dimension tileable. Only has an effect when depth is non-zero.", - }, - ), - "blend": ( - tuple(BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise.", - }, - ), - "pattern_break_blend": ( - tuple(BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Blending function used to blend pattern broken noise with raw noise.", - }, - ), - "depth_over_channels": ( - "BOOLEAN", - { - "default": DEFAULTS.depth_over_channels, - "tooltip": "When disabled, each channel will have its own separate 3D noise pattern. When enabled, depth is multiplied by the number of channels and each channel is a slice of depth. Only has an effect when depth is non-zero.", - }, - ), - "pad_height": ( - "INT", - { - "default": DEFAULTS.pad[0], - "min": 0, - "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.", - }, - ), - "pad_width": ( - "INT", - { - "default": DEFAULTS.pad[1], - "min": 0, - "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.", - }, - ), - "pad_depth": ( - "INT", - { - "default": DEFAULTS.pad[2], - "min": 0, - "tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation. Only has an effect when depth is non-zero.", - }, - ), - "initial_amplitude": ( - "FLOAT", - { - "default": DEFAULTS.initial_amplitude, - "tooltip": "Controls the amplitude for the first octave.", - }, - ), - "initial_frequency_height": ( - "FLOAT", - { - "default": DEFAULTS.initial_frequency[0], - "tooltip": "Controls the frequency for the first octave for the this axis.", - }, - ), - "initial_frequency_width": ( - "FLOAT", - { - "default": DEFAULTS.initial_frequency[1], - "tooltip": "Controls the frequency for the first octave for the this axis.", - }, - ), - "initial_frequency_depth": ( - "FLOAT", - { - "default": DEFAULTS.initial_frequency[2], - "tooltip": "Controls the frequency for the first octave for the this axis.", - }, - ), - "normalize": ( - ("default", "forced", "off"), - { - "tooltip": "Controls whether the output noise is normalized after generation.", - }, - ), - "device": ( - ("default", "cpu", "gpu"), - { - "default": "default", - "tooltip": "Controls what device is used to generate the noise. GPU noise may be slightly faster but you will get different results on different GPUs.", - }, - ), - } - return result - - @classmethod - def get_item_class(cls): - return PerlinItem - - -class PerlinSimpleNode(PerlinAdvancedNode): - DESCRIPTION = "Simplified Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinAdvanced node for more tuneable parameters." - - _COPY_KEYS = { - "factor", - "rescale", - "depth", - "detail_level", - "octaves", - "persistence", - "break_pattern", - } - - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - orig_reqs = result["required"] - reqs = {k: v for k, v in orig_reqs.items() if k in cls._COPY_KEYS} - reqs["lacunarity"] = orig_reqs["lacunarity_height"] - reqs["res"] = orig_reqs["res_height"] - result["required"] = reqs - return result - - @classmethod - def get_item_class(cls): - def wrapper(factor, *, lacunarity, res, **kwargs): - return PerlinItem( - factor, - lacunarity_height=lacunarity, - lacunarity_width=lacunarity, - lacunarity_depth=lacunarity, - res_height=res, - res_width=res, - res_depth=res, - **kwargs, - ) - - return wrapper diff --git a/py/expression/expression.py b/py/expression/expression.py index eabaf6f..d4ff4ae 100644 --- a/py/expression/expression.py +++ b/py/expression/expression.py @@ -1,6 +1,8 @@ import re import operator +from tqdm import tqdm + from .parser import Parser, ParserSpec, ParseError from .types import ( Empty, @@ -58,7 +60,7 @@ class Expression: def eval(self, handlers, *args, **kwargs): if self.expr != ExpOp("default"): - print("\nEVAL", self.expr) + tqdm.write(f"* OCS: EVAL: {self.expr}") if not isinstance(self.expr, ExpBase): return self.expr return self.expr.eval(handlers, *args, **kwargs) @@ -81,20 +83,25 @@ class Expression: def fixup_token(cls, t): if t == "": return t + if t[0] == "'": + return ExpSym(t[1:]) t = t.lower() val = cls.FIXUP.get(t, Empty) if val is not Empty: return val if t[0] == "`": return ExpBinOp(t.strip("`")) - if t[0] == "'": - return ExpSym(t[1:]) if (len(t) > 1 and t[0] == "-" and t[1].isdigit()) or t[0].isdigit(): return float(t) if "." in t else int(t) return ExpOp(t) @classmethod def tokenize(cls, s): + s = "\n".join( + line.rstrip("\r") + for line in s.split("\n") + if not line.lstrip().startswith("#") + ) yield from (cls.fixup_token(m.group(1)) for m in cls.EXPR_RE.finditer(s)) diff --git a/py/expression/handler.py b/py/expression/handler.py index 65e56a3..b2aac1c 100644 --- a/py/expression/handler.py +++ b/py/expression/handler.py @@ -212,6 +212,8 @@ class BetweenHandler(BaseHandler): # Inclusive def handle(self, obj, getter): value, low, high = self.safe_get_all(obj, getter) + if low > high: + low, high = high, low return low <= value <= high @@ -281,9 +283,17 @@ class GetHandler(BaseHandler): class S_Handler(BaseHandler): input_validators = ( - Arg.integer("start", None), - Arg.integer("end", None), - Arg.integer("step", None), + Arg.one_of( + "start", + (ValidateArg.validate_none, ValidateArg.validate_integer), + default=None, + ), + Arg.one_of( + "end", + (ValidateArg.validate_none, ValidateArg.validate_integer), + default=None, + ), + Arg.integer("step", 1), ) def handle(self, obj, getter): diff --git a/py/expression/validation.py b/py/expression/validation.py index 0bcdf29..1a089f6 100644 --- a/py/expression/validation.py +++ b/py/expression/validation.py @@ -1,3 +1,4 @@ +import contextlib import functools from ..latent import ImageBatch @@ -52,6 +53,10 @@ class Arg: name, default=default, validator=ValidateArg.validate_numscalar_sequence ) + @classmethod + def tensor_slice(cls, name, default=Empty): + return cls(name, default=default, validator=ValidateArg.validate_tensor_slice) + @classmethod def sequence(cls, name, default=Empty, *, item_validator=None): return cls( @@ -62,6 +67,16 @@ class Arg: ), ) + @classmethod + def nested_sequence(cls, name, default=Empty, *, item_validator=None): + return cls( + name, + default=default, + validator=functools.partial( + ValidateArg.validate_nested_sequence, item_validator=item_validator + ), + ) + @classmethod def string(cls, name, default=Empty): return cls(name, default=default, validator=ValidateArg.validate_string) @@ -98,7 +113,10 @@ class ValidateArg: def __init__(self, name, *args, kwargslist=(), group=all, **kwargs): if not isinstance(name, (list, tuple)): - return self.__init__((name,), (args,), group=group, kwargslist=kwargs) + name = (name,) + args = ((args,),) + kwargslist = kwargs + kwargs = {} self.valfuns = (getattr(self, f"validate_{n}", None) for n in name) if not all(self.valfuns): raise ValueError("Unknown validator") @@ -131,6 +149,22 @@ class ValidateArg: raise ValidateError(f"Expected numeric argument at {idx}, got {type(val)}") return val + @classmethod + def validate_tensor_slice_item(cls, idx, val): + with contextlib.suppress(ValidateError): + ok = ( + val in {Ellipsis, None} + or isinstance(val, (int, slice)) + or cls.validate_sequence( + idx, val, item_validator=ValidateArg.validate_integer + ) + ) + if ok: + return val + raise ValidateError( + f"Expected none, int, slice, tuple of int or ellipsis argument at {idx}, got {type(val)}" + ) + @classmethod def validate_integer(cls, idx, val): if not isinstance(val, int): @@ -160,7 +194,29 @@ class ValidateArg: try: return tuple(item_validator(iidx, v) for iidx, v in enumerate(val)) except ValidateError as exc: - raise ValidateError(f"Item validation failed for in sequence: {exc}") + raise ValidateError( + f"Item validation failed for sequence argument at {idx}: {exc}" + ) + + @classmethod + def validate_nested_sequence(cls, idx, val, *, item_validator=None, depth=0): + if not isinstance(val, (list, tuple)): + raise ValidateError( + f"Expected nested sequence argument at {idx}, depth {depth} but got {type(val)}" + ) + try: + return tuple( + cls.validate_nested_sequence( + idx, v, item_validator=item_validator, depth=depth + 1 + ) + if isinstance(v, (list, tuple)) + else (item_validator(iidx, v) if item_validator is not None else v) + for iidx, v in enumerate(val) + ) + except ValidateError as exc: + raise ValidateError( + f"Item validation failed for nested sequence argument at {idx}, depth {depth}: {exc}" + ) @classmethod def validate_numscalar_sequence(cls, idx, val): @@ -168,20 +224,11 @@ class ValidateArg: idx, val, item_validator=cls.validate_numeric_scalar ) - # @classmethod - # def validate_numscalar_sequence(cls, idx, val): - # if not isinstance(val, (list, tuple)): - # raise ValidateError(f"Expected sequence argument at {idx}, got {type(val)}") - # try: - # _ = all( - # cls.validate_numeric_scalar(f"{idx}[{i}]", v) is not None - # for i, v in enumerate(val) - # ) - # except ValidateError as exc: - # raise ValidateError( - # f"Expected numeric sequence argument at {idx}, got {type(val)}: {exc}" - # ) - # return val + @classmethod + def validate_tensor_slice(cls, idx, val): + return cls.validate_sequence( + idx, val, item_validator=cls.validate_tensor_slice_item + ) @classmethod def validate_string(cls, idx, val): @@ -195,6 +242,12 @@ class ValidateArg: raise ValidateError(f"Expected boolean argument at {idx}, got {type(val)}") return val + @classmethod + def validate_none(cls, idx, val): + if val is not None: + raise ValidateError(f"Expected none argument at {idx}, got {type(val)}") + return val + @classmethod def validate_passthrough(cls, idx, val): return val diff --git a/py/expression_handlers.py b/py/expression_handlers.py index d669c0a..2910a4e 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -3,31 +3,45 @@ import os import torch import numpy as np import PIL.Image as PILImage +from functools import partial from . import expression as expr from . import latent +from . import unsafe_expression_whitelists from .external import MODULES as EXT -from .utils import scale_noise, resolve_value -from .latent import OCSTAESD, ImageBatch +from .utils import scale_noise, resolve_value, quantile_normalize +from .latent import OCSTAESD, ImageBatch, normalize_to_scale + ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None -EXT_BLEH = EXT.get("bleh") -EXT_SONAR = EXT.get("sonar") -EXT_NNLATENTUPSCALE = EXT.get("nnlatentupscale") +EXT_BLEH = EXT_SONAR = None -if "bleh" in EXT: - BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES -else: - BLENDING_MODES = { - "lerp": lambda a, b, t: (1 - t) * a + t * b, - } +BLENDING_MODES = { + "lerp": torch.lerp, +} HANDLERS = {} +def init_integrations(integrations): + global EXT_BLEH, EXT_SONAR, BLENDING_MODES, HANDLERS + EXT_BLEH = integrations.bleh + EXT_SONAR = integrations.sonar + if EXT_BLEH is not None: + BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES + HANDLERS["t_bleh_enhance"] = BlehEnhanceHandler() + if EXT_SONAR is not None: + HANDLERS["t_sonar_power_filter"] = SonarPowerFilterHandler() + if integrations.nnlatentupscale is not None: + HANDLERS["t_scale_nnlatentupscale"] = ScaleNNLatentUpscaleHandler() + + +EXT.register_init_handler(init_integrations) + + class NormHandler(expr.BaseHandler): input_validators = ( expr.Arg.tensor("tensor"), @@ -42,6 +56,114 @@ class NormHandler(expr.BaseHandler): validate_output = expr.Arg.tensor("output") +class QuantileNormHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numeric("quantile", 0.75), + expr.Arg.integer("dim", 1), + expr.Arg.boolean("flatten", True), + expr.Arg.numeric("norm_factor", 1.0), + expr.Arg.numeric("norm_power", 0.5), + expr.Arg.string("mode", "clamp"), + ) + + def handle(self, obj, getter): + tensor, quantile, dim, flatten, norm_factor, norm_power, mode = ( + self.safe_get_all(obj, getter) + ) + return quantile_normalize( + tensor, + quantile=quantile, + dim=dim, + flatten=flatten, + nq_fac=norm_factor, + pow_fac=norm_power, + strategy=mode, + ) + + +class NormToScaleHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numeric("target_min", 0.0), + expr.Arg.numeric("target_max", 1.0), + expr.Arg.numscalar_sequence("dim", (-3, -2, -1)), + ) + + def handle(self, obj, getter): + tensor, tmin, tmax, dim = self.safe_get_all(obj, getter) + return normalize_to_scale(tensor, tmin, tmax, dim=dim) + + +class ClampHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numeric("min", 0.0), + expr.Arg.numeric("max", 1.0), + ) + + def handle(self, obj, getter): + tensor, tmin, tmax = self.safe_get_all(obj, getter) + return torch.clamp(tensor, min=tmin, max=tmax) + + +class StackHandler(NormHandler): + input_validators = ( + expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor), + expr.Arg.integer("dim", 1), + ) + + def handle(self, obj, getter): + tensors, dim = self.safe_get_all(obj, getter) + return torch.stack(tensors, dim) + + +class CatHandler(StackHandler): + def handle(self, obj, getter): + tensors, dim = self.safe_get_all(obj, getter) + return torch.cat(tensors, dim) + + +class ReshapeHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numscalar_sequence("shape"), + ) + + def handle(self, obj, getter): + tensor, shape = self.safe_get_all(obj, getter) + return torch.reshape(tensor.clone(), shape) + + +class IndexedCopyHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor_dest"), + expr.Arg.tensor("tensor_src"), + expr.Arg.tensor_slice("slice"), + ) + + def handle(self, obj, getter): + tensor1, tensor2, tensor_slice = self.safe_get_all(obj, getter) + result = tensor1.clone() + result[tensor_slice] = tensor2[tensor_slice] + return result + + +class NewLikeHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.nested_sequence( + "values", + default=(), + item_validator=expr.ValidateArg.validate_numeric_scalar, + ), + ) + + def handle(self, obj, getter): + tensor, values = self.safe_get_all(obj, getter) + return torch.tensor(values, dtype=tensor.dtype, device=tensor.device) + + class MeanHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), @@ -68,11 +190,23 @@ class RollHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), expr.Arg.numeric_scalar("amount", 0.5), - expr.Arg.numscalar_sequence("dim", (-2,)), + expr.Arg.one_of( + "dim", + ( + expr.ValidateArg.validate_integer, + partial( + expr.ValidateArg.validate_sequence, + item_validator=expr.ValidateArg.validate_integer, + ), + ), + default=-2, + ), ) def handle(self, obj, getter): tensor, amount, dim = self.safe_get_all(obj, getter) + if not isinstance(dim, tuple): + dim = (dim,) if isinstance(amount, float) and amount < 1.0 and amount > -1.0: if len(dim) > 1: raise ValueError( @@ -105,11 +239,42 @@ class FlipHandler(NormHandler): out_slice = tuple( np.s_[:] if d != dim else np.s_[pivot:] for d in range(tensor.ndim) ) - in_slice = tuple(np.s_[:] if d != dim else np.s_[:pivot] for d in range(tensor.ndim)) + in_slice = tuple( + np.s_[:] if d != dim else np.s_[:pivot] for d in range(tensor.ndim) + ) result[out_slice] = torch.flip(tensor[in_slice], dims=(dim,)) return result +class CopySignHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.tensor("other"), + ) + + def handle(self, obj, getter): + return torch.copysign(*self.safe_get_all(obj, getter)) + + +class CloneHandler(NormHandler): + input_validators = (expr.Arg.tensor("tensor"),) + + def handle(self, obj, getter): + return self.safe_get("tensor", obj, getter).clone() + + +class NewFullHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numscalar_sequence("shape"), + expr.Arg.numeric("value", 0.0), + ) + + def handle(self, obj, getter): + tensor, shape, value = self.safe_get_all(obj, getter) + return tensor.new_full(shape, value) + + class BlendHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor1"), @@ -130,11 +295,12 @@ class ContrastAdaptiveSharpeningHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), expr.Arg.numeric("scale", 0.5), + expr.Arg.boolean("normalize", True), ) def handle(self, obj, getter): - t, scale = self.safe_get_all(obj, getter) - return latent.contrast_adaptive_sharpening(t, scale) + t, scale, normalize = self.safe_get_all(obj, getter) + return latent.contrast_adaptive_sharpening(t, scale, normalize=normalize) class ScaleHandler(NormHandler): @@ -164,7 +330,6 @@ class ScaleHandler(NormHandler): scale = tuple(int(v) for v in scale) else: scale = (int(t.shape[-2] * scale[0]), int(t.shape[-1] * scale[1])) - print("SCALE", t.shape[-2:], "->", scale) if not all(v > 0 for v in scale): raise ValueError(f"Invalid scale: scale values must be > 0, got: {scale!r}") return latent.scale_samples(t, scale[1], scale[0], mode=mode) @@ -180,7 +345,8 @@ class NoiseHandler(NormHandler): t, typ = self.safe_get_all(obj, getter) ctx = getter.ctx smin, smax, s, sn = ( - ctx.get_var(k) for k in ("sigma_min", "sigma_max", "sigma", "sigma_next") + ctx.get_var(k, default=0.0) + for k in ("sigma_min", "sigma_max", "sigma", "sigma_next") ) ns = latent.get_noise_sampler(typ, t, smin, smax, normalized=False) return ns(s, sn) @@ -194,6 +360,55 @@ class ShapeHandler(expr.BaseHandler): return expr.types.ExpTuple((*t.shape,)) +class GaussianBlur2DHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("kernel_size"), + expr.Arg.numeric_scalar("sigma"), + ) + + def handle(self, obj, getter): + return latent.gaussian_blur_2d(*self.safe_get_all(obj, getter)) + + +class SNFGuidanceHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("t_tensor"), + expr.Arg.tensor("s_tensor"), + expr.Arg.integer("t_kernel_size", 3), + expr.Arg.numeric_scalar("t_sigma", 1), + expr.Arg.integer("s_kernel_size", 3), + expr.Arg.numeric_scalar("s_sigma", 1), + ) + + def handle(self, obj, getter): + return latent.snf_guidance(*self.safe_get_all(obj, getter)) + + +class RGBLatentHandler(expr.BaseHandler): + input_validators = ( + expr.Arg.tensor("reference"), + expr.Arg.numscalar_sequence("rgb"), + ) + + def handle(self, obj, getter): + ctx = getter.ctx.constants.ctx + model = ctx.get("model") + if model is None: + raise ValueError("Ohno") + reference, rgb = self.safe_get_all(obj, getter) + if len(rgb) != 3 or not all(0 <= v <= 1.0 for v in rgb): + raise ValueError("Bad RGB parameter") + img = torch.tensor( + tuple(v * 2 - 1.0 for v in rgb), + device=reference.device, + dtype=reference.dtype, + ).view(1, 3) + latent = torch.zeros_like(reference).movedim(1, -1) + latent += model.latent_format.rgb_to_latent(img) + return latent.movedim(-1, 1) + + class TAESDDecodeHandler(expr.BaseHandler): input_validators = ( expr.Arg.tensor("tensor"), @@ -282,295 +497,9 @@ class UnsafeTorchTensorMethodHandler(NormHandler): whitelist = AlwaysContains() elif ALLOW_UNSAFE: - whitelist = { - "abs", - "absolute", - "acos", - "acosh", - "add", - "addbmm", - "addcdiv", - "addcmul", - "addmm", - "addmv", - "addr", - "adjoint", - "all", - "allclose", - "amax", - "amin", - "aminmax", - "angle", - "any", - "arccos", - "arccosh", - "arcsin", - "arcsinh", - "arctan", - "arctan2", - "arctanh", - "argmax", - "argmin", - "argsort", - "argwhere", - "as_strided", - "asin", - "asinh", - "atan", - "atan2", - "atanh", - "baddbmm", - "bernoulli", - "bincount", - "bitwise_and", - "bitwise_left_shift", - "bitwise_not", - "bitwise_or", - "bitwise_right_shift", - "bitwise_xor", - "bmm", - "broadcast_to", - "ceil", - "cholesky", - "cholesky_inverse", - "cholesky_solve", - "chunk", - "clamp", - "clip", - "clone", - "conj", - "conj_physical", - "contiguous", - "copysign", - "corrcoef", - "cos", - "cosh", - "count_nonzero", - "cov", - "cross", - "cummax", - "cummin", - "cumprod", - "cumsum", - "deg2rad", - "det", - "detach", - "diag", - "diag_embed", - "diagflat", - "diagonal", - "diagonal_scatter", - "diff", - "digamma", - "dim", - "dist", - "div", - "divide", - "dot", - "dsplit", - "eq", - "equal", - "erf", - "erfc", - "erfinv", - "exp", - "expand", - "expand_as", - "expm1", - "fix", - "flatten", - "flip", - "fliplr", - "flipud", - "float_power", - "floor", - "floor_divide", - "fmax", - "fmin", - "fmod", - "frac", - "frexp", - "gather", - "gcd", - "ge", - "geqrf", - "ger", - "greater", - "greater_equal", - "gt", - "hardshrink", - "heaviside", - "histc", - "hsplit", - "hypot", - "i0", - "igamma", - "igammac", - "index_add", - "index_copy", - "index_fill", - "index_put", - "index_reduce", - "index_select", - "inner", - "inverse", - "isclose", - "isfinite", - "isinf", - "isnan", - "isneginf", - "isposinf", - "kthvalue", - "lcm()", - "ldexp", - "le", - "lerp", - "less", - "less_equal", - "lgamma", - "log", - "log10", - "log1p", - "log2", - "logaddexp", - "logaddexp2", - "logcumsumexp", - "logdet", - "logical_and", - "logical_not", - "logical_or", - "logical_xor", - "logit", - "logsumexp", - "lt", - "lu", - "lu_solve", - "masked_fill", - "masked_scatter", - "masked_select", - "matmul", - "matrix_exp", - "max", - "maximum", - "mean", - "median", - "min", - "minimum", - "mm", - "mode", - "moveaxis", - "movedim", - "msort", - "mul", - "multinomial", - "multiply", - "mv", - "mvlgamma", - "nan_to_num", - "nanmean", - "nanmedian", - "nanquantile", - "nansum", - "narrow", - "narrow_copy", - "ne", - "neg", - "negative", - "new_empty", - "new_full", - "new_ones", - "new_zeros", - "nextafter", - "nonzero", - "norm", - "not_equal", - "numel", - "orgqr", - "ormqr", - "outer", - "permute", - "polygamma", - "positive", - "pow", - "prod", - "qr", - "quantile", - "rad2deg", - "ravel", - "reciprocal", - "remainder", - "renorm", - "repeat", - "repeat_interleave", - "reshape", - "reshape_as", - "resolve_conj", - "resolve_neg", - "roll", - "rot90", - "round", - "rsqrt", - "scatter", - "scatter_add", - "scatter_reduce", - "select", - "select_scatter", - "sgn", - "sigmoid", - "sign", - "signbit", - "sin", - "sinc", - "sinh", - "slice_scatter", - "slogdet", - "smm", - "softmax", - "sort", - "sparse_mask", - "split", - "sqrt", - "square", - "squeeze", - "sspaddmm", - "std", - "stft", - "sub", - "subtract", - "sum", - "sum_to_size", - "svd", - "swapaxes", - "swapdims", - "t", - "take", - "take_along_dim", - "tan", - "tanh", - "tensor_split", - "tile", - "topk", - "transpose", - "triangular_solve", - "tril", - "triu", - "true_divide", - "trunc", - "unflatten", - "unfold", - "unique", - "unique_consecutive", - "unsqueeze", - "var", - "vdot", - "view", - "view_as", - "vsplit", - "where", - "xlogy", - } + whitelist = unsafe_expression_whitelists.TORCH_FUNCTION_WHITELIST else: - whitelist = set() + whitelist = frozenset() def handle(self, obj, getter): if "__method" in obj.kwargs or "__tensor" in obj.kwargs: @@ -609,119 +538,120 @@ class UnsafeTorchHandler(expr.BaseHandler): return resolve_value(keys, torch) -if EXT_BLEH: +class BlehEnhanceHandler(expr.BaseHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.string("mode"), + expr.Arg.numeric_scalar("scale", 1.0), + ) + output_validator = expr.Arg.tensor("output") - class BlehEnhanceHandler(expr.BaseHandler): - input_validators = ( - expr.Arg.tensor("tensor"), - expr.Arg.string("mode"), - expr.Arg.numeric_scalar("scale", 1.0), + def handle(self, obj, getter): + tensor, mode, scale = self.safe_get_all(obj, getter) + return EXT_BLEH.latent_utils.enhance_tensor( + tensor, mode, scale=scale, adjust_scale=False ) - output_validator = expr.Arg.tensor("output") - def handle(self, obj, getter): - tensor, mode, scale = self.safe_get_all(obj, getter) - return EXT_BLEH.latent_utils.enhance_tensor( - tensor, mode, scale=scale, adjust_scale=False - ) - HANDLERS["t_bleh_enhance"] = BlehEnhanceHandler() +class SonarPowerFilterHandler(expr.BaseHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.present("filter"), + ) + output_validator = expr.Arg.tensor("output") -if EXT_SONAR: + default_power_filter = { + "mix": 1.0, + "normalization_factor": 1.0, + "common_mode": 0.0, + "channel_correlation": "1,1,1,1,1,1", + } - class SonarPowerFilterHandler(expr.BaseHandler): - input_validators = ( - expr.Arg.tensor("tensor"), - expr.Arg.present("filter"), - ) - output_validator = expr.Arg.tensor("output") - - default_power_filter = { - "mix": 1.0, - "normalization_factor": 1.0, - "common_mode": 0.0, - "channel_correlation": "1,1,1,1,1,1", - } - - @classmethod - def make_power_filter(cls, fdict, *, toplevel=True): - fdict = fdict.copy() - compose_with = fdict.pop("compose_with", None) - if compose_with: - if not isinstance(compose_with, dict): - raise TypeError("compose_with must be a dictionary") - fdict["compose_with"] = cls.make_power_filter( - compose_with, toplevel=False + @classmethod + def make_power_filter(cls, fdict, *, toplevel=True): + fdict = fdict.copy() + compose_with = fdict.pop("compose_with", None) + if compose_with: + if not isinstance(compose_with, dict): + raise TypeError("compose_with must be a dictionary") + fdict["compose_with"] = cls.make_power_filter(compose_with, toplevel=False) + topargs = {k: fdict.pop(k, dv) for k, dv in cls.default_power_filter.items()} + power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict) + if not toplevel: + return power_filter + cc = topargs.get("channel_correlation") + if cc is not None: + if not isinstance(cc, (list, tuple)) or not all( + isinstance(v, (int, float)) for v in cc + ): + raise TypeError( + "Bad channel correlation type: must be comma separated string or numeric sequence" ) - topargs = { - k: fdict.pop(k, dv) for k, dv in cls.default_power_filter.items() - } - power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict) - if not toplevel: - return power_filter - cc = topargs.get("channel_correlation") - if cc is not None: - if not isinstance(cc, (list, tuple)) or not all( - isinstance(v, (int, float)) for v in cc - ): - raise TypeError( - "Bad channel correlation type: must be comma separated string or numeric sequence" - ) - topargs["channel_correlation"] = ",".join(repr(v) for v in cc) - return EXT_SONAR.powernoise.PowerNoiseItem( - 1, power_filter=power_filter, time_brownian=True, **topargs - ) - - def handle(self, obj, getter): - tensor, filter_def = self.safe_get_all(obj, getter) - if not isinstance(filter_def, dict): - raise TypeError("filter argument must be a dictionary") - power_filter = self.make_power_filter(filter_def) - filter_rfft = power_filter.make_filter(tensor.shape).to( - tensor.device, non_blocking=True - ) - ns = power_filter.make_noise_sampler_internal( - tensor, - lambda *_unused, latent=tensor: latent, - filter_rfft, - normalized=False, - ) - return ns(None, None) - - HANDLERS["t_sonar_power_filter"] = SonarPowerFilterHandler() - -if EXT_NNLATENTUPSCALE: - from .latent import scale_nnlatentupscale - - class ScaleNNLatentUpscaleHandler(expr.BaseHandler): - input_validators = ( - expr.Arg.tensor("tensor"), - expr.Arg.string("mode", "sd1"), - expr.Arg.numeric_scalar("scale", 2.0), + topargs["channel_correlation"] = ",".join(repr(v) for v in cc) + return EXT_SONAR.powernoise.PowerNoiseItem( + 1, power_filter=power_filter, time_brownian=True, **topargs ) - output_validator = expr.Arg.tensor("output") - def handle(self, obj, getter): - tensor, mode, scale = self.safe_get_all(obj, getter) - if mode not in {"sd1", "sdxl"}: - raise ValueError( - "Bad mode for t_scale_nnlatentupscale: must be either sd15 or sdxl" - ) - return scale_nnlatentupscale(mode, tensor, scale) + def handle(self, obj, getter): + tensor, filter_def = self.safe_get_all(obj, getter) + if not isinstance(filter_def, dict): + raise TypeError("filter argument must be a dictionary") + power_filter = self.make_power_filter(filter_def) + filter_rfft = power_filter.make_filter(tensor.shape).to( + tensor.device, non_blocking=True + ) + ns = power_filter.make_noise_sampler_internal( + tensor, + lambda *_unused, latent=tensor: latent, + filter_rfft, + normalized=False, + ) + return ns(None, None) + + +class ScaleNNLatentUpscaleHandler(expr.BaseHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.string("mode", "sd1"), + expr.Arg.numeric_scalar("scale", 2.0), + ) + output_validator = expr.Arg.tensor("output") + + def handle(self, obj, getter): + tensor, mode, scale = self.safe_get_all(obj, getter) + if mode not in {"sd1", "sdxl"}: + raise ValueError( + "Bad mode for t_scale_nnlatentupscale: must be either sd15 or sdxl" + ) + return latent.scale_nnlatentupscale(mode, tensor, scale) - HANDLERS["t_scale_nnlatentupscale"] = ScaleNNLatentUpscaleHandler() TENSOR_OP_HANDLERS = { "t_norm": NormHandler(), + "t_quantilenorm": QuantileNormHandler(), + "t_normtoscale": NormToScaleHandler(), + "t_normalize_to_scale": NormToScaleHandler(), + "t_reshape": ReshapeHandler(), + "t_clamp": ClampHandler(), + "t_cat": CatHandler(), + "t_stack": StackHandler(), + "t_indexed_copy": IndexedCopyHandler(), + "t_new_like": NewLikeHandler(), "t_mean": MeanHandler(), "t_std": StdHandler(), "t_blend": BlendHandler(), "t_roll": RollHandler(), "t_flip": FlipHandler(), + "t_clone": CloneHandler(), + "t_newfull": NewFullHandler(), + "t_copysign": CopySignHandler(), "t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(), "t_scale": ScaleHandler(), "t_noise": NoiseHandler(), "t_shape": ShapeHandler(), + "t_gaussianblur2d": GaussianBlur2DHandler(), + "t_rgb_latent": RGBLatentHandler(), + "t_snf_guidance": SNFGuidanceHandler(), "t_taesd_decode": TAESDDecodeHandler(), "unsafe_tensor_method": UnsafeTorchTensorMethodHandler(), "unsafe_torch": UnsafeTorchHandler(), diff --git a/py/external.py b/py/external.py index 334300a..6f818f9 100644 --- a/py/external.py +++ b/py/external.py @@ -1,21 +1,120 @@ import contextlib import importlib +import sys +from functools import partial +from types import ModuleType +from typing import Callable, NamedTuple -MODULES = {} -with contextlib.suppress(ImportError, NotImplementedError): - bleh = importlib.import_module("custom_nodes.ComfyUI-bleh") - bleh_version = getattr(bleh, "BLEH_VERSION", -1) - if bleh_version < 1: - raise NotImplementedError - MODULES["bleh"] = bleh.py +class Integrations: + class Integration(NamedTuple): + key: str + module_name: str + handler: Callable | None = None -with contextlib.suppress(ImportError, NotImplementedError): - MODULES["sonar"] = importlib.import_module("custom_nodes.ComfyUI-sonar").py + def __init__(self): + self.initialized = False + self.modules = {} + self.init_handlers = [] + self.handlers = [] + + def __getitem__(self, key): + return self.modules[key] + + def __contains__(self, key): + return key in self.modules + + def __getattr__(self, key): + return self.modules.get(key) + + def get(self, key, default=None): + return self.modules.get(key, default) + + @staticmethod + def get_custom_node(module_name: str, key: str) -> ModuleType | None: + bi_module = sys.modules.get("_blepping_integrations", {}).get(key) + if bi_module is not None: + return bi_module + module_key = f"custom_nodes.{module_name}" + with contextlib.suppress(StopIteration): + spec = importlib.util.find_spec(module_key) + if spec is None: + return None + return next( + v + for v in sys.modules.copy().values() + if hasattr(v, "__spec__") + and v.__spec__ is not None + and v.__spec__.origin == spec.origin + ) + return None + + def register_init_handler(self, handler): + self.init_handlers.append(handler) + + def register_integration(self, key: str, module_name: str, handler=None) -> None: + if self.initialized: + raise ValueError( + "Internal error: Cannot register integration after initialization", + ) + if any(item[0] == key or item[1] == module_name for item in self.handlers): + errstr = ( + f"Module {module_name} ({key}) already in integration handlers list!" + ) + raise ValueError(errstr) + self.handlers.append(self.Integration(key, module_name, handler)) + + def initialize(self) -> None: + if self.initialized: + return + self.initialized = True + for ih in self.handlers: + module = self.get_custom_node(ih.module_name, ih.key) + if module is None: + continue + if ih.handler is not None: + module = ih.handler(module) + if module is not None: + self.modules[ih.key] = module + + for init_handler in self.init_handlers: + init_handler(self) + + +class OCSIntegrations(Integrations): + def __init__(self, *args: list, **kwargs: dict): + super().__init__(*args, **kwargs) + self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration) + self.register_integration("sonar", "ComfyUI-sonar", self.sonar_integration) + self.register_integration("nnlatentupscale", "ComfyUi_NNLatentUpscale") + self.register_integration("tiled_diffusion", "ComfyUI-TiledDiffusion") + + @classmethod + def bleh_integration(cls, module: ModuleType) -> ModuleType | None: + bleh_version = getattr(module, "BLEH_VERSION", -1) + if bleh_version < 1: + return None + return module.py + + @classmethod + def sonar_integration(cls, module: ModuleType) -> ModuleType | None: + return module.py + + +MODULES = OCSIntegrations() + + +class IntegratedNode(type): + @staticmethod + def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict: + MODULES.initialize() + return orig_method(*args, **kwargs) + + def __new__(cls: type, name: str, bases: tuple, attrs: dict) -> object: + obj = type.__new__(cls, name, bases, attrs) + if hasattr(obj, "INPUT_TYPES"): + obj.INPUT_TYPES = partial(cls.wrap_INPUT_TYPES, obj.INPUT_TYPES) + return obj -with contextlib.suppress(ImportError, NotImplementedError): - MODULES["nnlatentupscale"] = importlib.import_module( - "custom_nodes.ComfyUi_NNLatentUpscale" - ) __all__ = ("MODULES",) diff --git a/py/filtering.py b/py/filtering.py index 6d62de1..bc62a34 100644 --- a/py/filtering.py +++ b/py/filtering.py @@ -10,23 +10,33 @@ from .utils import fallback OD = collections.OrderedDict -EXT_BLEH = EXT.get("bleh") -EXT_SONAR = EXT.get("sonar") - -if "bleh" in EXT: - BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES -else: - BLENDING_MODES = { - "lerp": lambda a, b, t: (1 - t) * a + t * b, - } - -BLENDING_MODES = BLENDING_MODES | { +BLENDING_MODES = { + "lerp": torch.lerp, "a_only": lambda a, b, t: a * t, "b_only": lambda a, b, t: b * t, + "inject": lambda a, b, t: (b * t).add_(a), } FILTER = {} +EXT_BLEH = None + + +def init_integrations(integrations): + global BLENDING_MODES, FILTER, EXT_BLEH + EXT_BLEH = integrations.bleh + if EXT_BLEH is not None: + BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES | BLENDING_MODES + FILTER |= { + "bleh_enhance": BlehEnhanceFilter, + "bleh_ops": BlehOpsFilter, + } + ext_sonar = integrations.sonar + if ext_sonar is not None: + FILTER["sonar_power_filter"] = SonarPowerFilter + + +EXT.register_init_handler(init_integrations) FILTER_HANDLERS = expr.HandlerContext( expr.BASIC_HANDLERS | expression_handlers.HANDLERS @@ -34,8 +44,9 @@ FILTER_HANDLERS = expr.HandlerContext( class FilterRefs: - def __init__(self, kvs=None): + def __init__(self, kvs=None, *, ctx=None): self.kvs = fallback(kvs, {}) + self.ctx = fallback(ctx, {}) def get(self, k, default=None): return self.kvs.get(k, default) @@ -47,13 +58,14 @@ class FilterRefs: self.kvs[k] = v def clone(self): - return self.__class__(self.kvs.copy()) + return self.__class__(self.kvs.copy(), ctx=self.ctx) def __or__(self, other): - return self.__class__(self.kvs | other.kvs) + return self.__class__(self.kvs | other.kvs, ctx=self.ctx | other.ctx) def __ior__(self, other): self.kvs |= other.kvs + self.ctx |= other.ctx return self def __delitem__(self, k): @@ -77,25 +89,28 @@ class FilterRefs: @classmethod def from_ss(cls, ss, *, have_current=False): ms = ss.model.model_sampling - fr = cls({ - "step": ss.step, - "substep": ss.substep, - "dt": ss.dt, - "sigma_idx": ss.idx, - "sigma": ss.sigma, - "sigma_next": ss.sigma_next, - "sigma_down": ss.sigma_down, - "sigma_up": ss.sigma_up, - "sigma_prev": ss.sigma_prev, - "hist_len": len(ss.hist), - "sigma_min": ms.sigma_min.item(), - "sigma_max": ms.sigma_max.item(), - "step_pct": float(ss.step / ss.total_steps), - "total_steps": ss.total_steps, - "sampling_pct": (999 - ms.timestep(ss.sigma).item()) / 999, - "is_rectified_flow": ss.model.is_rectified_flow, - "original_cfg_scale": ss.model.inner_cfg_scale, - }) + fr = cls( + { + "step": ss.step, + "substep": ss.substep, + "dt": ss.dt, + "sigma_idx": ss.idx, + "sigma": ss.sigma, + "sigma_next": ss.sigma_next, + "sigma_down": ss.sigma_down, + "sigma_up": ss.sigma_up, + "sigma_prev": ss.sigma_prev, + "hist_len": len(ss.hist), + "sigma_min": ms.sigma_min.item(), + "sigma_max": ms.sigma_max.item(), + "step_pct": float(ss.step / ss.total_steps), + "total_steps": ss.total_steps, + "sampling_pct": (999 - ms.timestep(ss.sigma).item()) / 999, + "is_rectified_flow": ss.model.is_rectified_flow, + "original_cfg_scale": ss.model.inner_cfg_scale, + }, + ctx={"ss": ss, "model": ss.model}, + ) if have_current and len(ss.hist) > 0: fr |= cls.from_mr(ss.hcur) fr["d"] = ss.d @@ -231,6 +246,12 @@ class Filter: refs = drefs if refs is None else refs | drefs return ops.eval(FILTER_HANDLERS.clone(constants=refs, variables={})) + def __str__(self): + prettyvals = ", ".join( + f"{k}={getattr(self, k, None)!s}" for k in self.default_options.keys() + ) + return f"" + class SimpleFilter(Filter): name = "simple" @@ -398,100 +419,88 @@ class NormalizeFilter_: Normalize = NormalizeFilter -if EXT_BLEH: - class BlehEnhanceFilter(Filter): - name = "bleh_enhance" - default_options = Filter.default_options | { - "enhance_mode": None, - "enhance_scale": 1.0, - } - - def filter(self, latent, *args, **kwargs): - if self.enhance_mode is None or self.enhance_scale == 1: - return latent - return EXT_BLEH.latent_utils.enhance_tensor( - latent, self.enhance_mode, scale=self.enhance_scale, adjust_scale=False - ) - - class BlehOpsFilter(Filter): - name = "bleh_ops" - default_options = Filter.default_options | {"ops": ()} - - def __init__(self, **kwargs): - super().__init__(**kwargs) - if isinstance(self.ops, (tuple, list)): - self.ops = EXT_BLEH.nodes.ops.RuleGroup( - tuple( - r - for rs in self.ops - for r in EXT_BLEH.nodes.ops.Rule.from_dict(rs) - ) - ) - return - if not isinstance(self.ops, str): - raise ValueError("ops key must be a YAML string or list of object") - self.ops = EXT_BLEH.nodes.ops.RuleGroup.from_yaml(self.ops) - - def filter(self, latent, ref_latent, *args, refs=None, **kwargs): - if not self.ops: - return latent - refs = fallback(refs, {}) - bops = EXT_BLEH.nodes.ops - state = { - bops.CondType.TYPE: bops.PatchType.LATENT, - bops.CondType.PERCENT: 0.0, - bops.CondType.BLOCK: -1, - bops.CondType.STAGE: -1, - bops.CondType.STEP: refs.get("step", 0), - bops.CondType.STEP_EXACT: refs.get("step", -1), - "h": latent, - "hsp": ref_latent, - "target": "h", - } - self.ops.eval(state, toplevel=True) - return state["h"] - - FILTER |= { - "bleh_enhance": BlehEnhanceFilter, - "bleh_ops": BlehOpsFilter, +class BlehEnhanceFilter(Filter): + name = "bleh_enhance" + default_options = Filter.default_options | { + "enhance_mode": None, + "enhance_scale": 1.0, } -if EXT_SONAR: + def filter(self, latent, *args, **kwargs): + if self.enhance_mode is None or self.enhance_scale == 1: + return latent + return EXT_BLEH.latent_utils.enhance_tensor( + latent, self.enhance_mode, scale=self.enhance_scale, adjust_scale=False + ) - class SonarPowerFilter(Filter): - name = "sonar_power_filter" - default_options = Filter.default_options - def __init__(self, **kwargs): - super().__init__(**kwargs) - power_filter = self.options.pop("power_filter", None) - if power_filter is None: - self.power_filter = None - return - if not isinstance(power_filter, dict): - raise ValueError("power_filter key must be dict or null") - self.power_filter = ( - expression_handlers.SonarPowerFilterHandler.make_power_filter( - power_filter +class BlehOpsFilter(Filter): + name = "bleh_ops" + default_options = Filter.default_options | {"ops": ()} + + def __init__(self, **kwargs): + super().__init__(**kwargs) + if isinstance(self.ops, (tuple, list)): + self.ops = EXT_BLEH.nodes.ops.RuleGroup( + tuple( + r for rs in self.ops for r in EXT_BLEH.nodes.ops.Rule.from_dict(rs) ) ) + return + if not isinstance(self.ops, str): + raise ValueError("ops key must be a YAML string or list of object") + self.ops = EXT_BLEH.nodes.ops.RuleGroup.from_yaml(self.ops) - def filter(self, latent, ref_latent, *args, refs=None, **kwargs): - if not self.power_filter: - return latent - filter_rfft = self.power_filter.make_filter(latent.shape).to( - latent.device, non_blocking=True - ) - ns = self.power_filter.make_noise_sampler_internal( - latent, - lambda *_unused, latent=latent: latent, - filter_rfft, - normalized=False, - ) - return ns(None, None) + def filter(self, latent, ref_latent, *args, refs=None, **kwargs): + if not self.ops: + return latent + refs = fallback(refs, {}) + bops = EXT_BLEH.nodes.ops + state = { + bops.CondType.TYPE: bops.PatchType.LATENT, + bops.CondType.PERCENT: 0.0, + bops.CondType.BLOCK: -1, + bops.CondType.STAGE: -1, + bops.CondType.STEP: refs.get("step", 0), + bops.CondType.STEP_EXACT: refs.get("step", -1), + "h": latent, + "hsp": ref_latent, + "target": "h", + } + self.ops.eval(state, toplevel=True) + return state["h"] - FILTER |= {"sonar_power_filter": SonarPowerFilter} + +class SonarPowerFilter(Filter): + name = "sonar_power_filter" + default_options = Filter.default_options + + def __init__(self, **kwargs): + super().__init__(**kwargs) + power_filter = self.options.pop("power_filter", None) + if power_filter is None: + self.power_filter = None + return + if not isinstance(power_filter, dict): + raise ValueError("power_filter key must be dict or null") + self.power_filter = ( + expression_handlers.SonarPowerFilterHandler.make_power_filter(power_filter) + ) + + def filter(self, latent, ref_latent, *args, refs=None, **kwargs): + if not self.power_filter: + return latent + filter_rfft = self.power_filter.make_filter(latent.shape).to( + latent.device, non_blocking=True + ) + ns = self.power_filter.make_noise_sampler_internal( + latent, + lambda *_unused, latent=latent: latent, + filter_rfft, + normalized=False, + ) + return ns(None, None) def make_filter(args): diff --git a/py/latent.py b/py/latent.py index 9bdbb5c..22738cc 100644 --- a/py/latent.py +++ b/py/latent.py @@ -11,6 +11,20 @@ from comfy import latent_formats from .external import MODULES as EXT +EXT_NNLATENTUPSCALE = None + + +def init_integrations(integrations): + global get_noise_sampler, EXT_NNLATENTUPSCALE + + ext_sonar = integrations.sonar + if ext_sonar is not None: + get_noise_sampler = ext_sonar.noise.get_noise_sampler + EXT_NNLATENTUPSCALE = EXT.nnlatentupscale + + +EXT.register_init_handler(init_integrations) + def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)): min_val, max_val = ( @@ -25,32 +39,33 @@ def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)): ) +# Improvements by https://github.com/Clybius # The following is modified to work with latent images of ~0 mean from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/tree/main. -def contrast_adaptive_sharpening(x, amount=0.8, *, epsilon=1e-06): - """ - Performs a contrast adaptive sharpening on the batch of images x. - The algorithm is directly implemented from FidelityFX's source code, - that can be found here - https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h +# The algorithm is directly implemented from FidelityFX's source code that can be found here: https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h. +def contrast_adaptive_sharpening( # noqa: PLR0914 + x, + amount=0.8, + *, + normalize=True, + epsilon=1e-06, +): + orig_shape = x.shape + if x.ndim == 5: + x = x.reshape(orig_shape[0], orig_shape[1] * orig_shape[2], *orig_shape[-2:]) + elif x.ndim != 4: + raise ValueError( + "Contrast-adaptive sharpening requires a tensor with 4 or 5 dimensions", + ) - Parameters - ---------- - x : Tensor - Image or stack of images, of shape [batch, channels, ny, nx]. - Batch and channel dimensions can be ommited. - amount : int [0, 1] - Amount of sharpening to do, 0 being minimum and 1 maximum - - Returns - ------- - Tensor - Processed stack of images. - - """ - - def on_abs_stacked(tensor_list, f, *args, **kwargs): + def on_abs_stacked(tensor_list, f, *args: list, **kwargs: dict): return f(torch.abs(torch.stack(tensor_list)), *args, **kwargs)[0] + if normalize: + luminance = torch.linalg.vector_norm(x, dim=1, keepdim=True).add_(1e-08) + x = x / luminance + orig_mean = x.mean(dim=(-3, -2, -1), keepdim=True) + x -= orig_mean + x_padded = F.pad(x, pad=(1, 1, 1, 1)) x_padded = torch.complex(x_padded, torch.zeros_like(x_padded)) # each side gets padded with 1 pixel @@ -102,7 +117,12 @@ def contrast_adaptive_sharpening(x, amount=0.8, *, epsilon=1e-06): div = torch.reciprocal(1 + 4 * w) output = ((b + d + f + h) * w + e) * div - return output.real.clamp(x.min(), x.max()) + output = output.real + for ob, xb in zip(x, output): + ob.clamp_(*xb.aminmax()) + if normalize: + output = output.add_(orig_mean).mul_(luminance) + return output.reshape(*orig_shape) class ImageBatch(tuple): @@ -145,24 +165,11 @@ class OCSTAESD: @classmethod def decode(cls, fmt, latent): latent_format = cls.latent_formats[fmt] - # rv = latent_format.process_out(1.0) filename = cls.get_taesd_path(cls.get_decoder_name(fmt)) model = TAESD( decoder_path=filename, latent_channels=latent_format.latent_channels ).to(latent.device) - # print("DEC INPUT ORIG", latent.min(), latent.max()) - # if torch.any(latent.max() > rv) or torch.any(latent.min() < -rv): - # sv = latent.new((-rv, rv)) - # latent = normalize_to_scale( - # latent, - # latent.amin(dim=(-3, -2, -1), keepdim=True).maximum(sv[0]), - # latent.amax(dim=(-3, -2, -1), keepdim=True).minimum(sv[1]), - # dim=(-3, -2, -1), - # ) - # print("DEC INPUT", latent.min(), latent.max()) - # result = model.decode(latent.clamp(-rv, rv)).movedim(1, 3) result = model.decode(latent).movedim(1, 3) - # print("DEC RESULT", result.shape, result.isnan().any().item()) return ImageBatch( latent_preview.preview_to_image(result[batch_idx]) for batch_idx in range(result.shape[0]) @@ -190,80 +197,143 @@ class OCSTAESD: encoder_path=filename, latent_channels=latent_format.latent_channels ).to(device=latent.device) result = model.encode(cls.img_to_encoder_input(imgbatch).to(latent.device)) - # print( - # "ENC RESULT ORIG", - # result.min(), - # result.max(), - # ) - # if torch.any(result.max() > rv) or torch.any(result.min() < -rv): - # sv = result.new((-rv, rv)) - # result = normalize_to_scale( - # result, - # result.amin(dim=(-3, -2, -1), keepdim=True).maximum(sv[0]), - # result.amax(dim=(-3, -2, -1), keepdim=True).minimum(sv[1]), - # dim=(-3, -2, -1), - # ) - # print( - # "ENC RESULT", - # result.shape, - # result.isnan().any().item(), - # result.min(), - # result.max(), - # ) return result.to(latent.dtype).clamp(-rv, rv) -if "bleh" in EXT: - scale_samples = EXT["bleh"].latent_utils.scale_samples - UPSCALE_METHODS = EXT["bleh"].latent_utils.UPSCALE_METHODS -else: - UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "area") - - def scale_samples( - samples, - width, - height, - mode="bicubic", - sigma=None, # noqa: ARG001 - ): - if mode == "bislerp": - return bislerp(samples, width, height) - return F.interpolate(samples, size=(height, width), mode=mode) +bleh_scale_samples = None +UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "area") -if "sonar" in EXT: - get_noise_sampler = EXT["sonar"].noise.get_noise_sampler -else: - - def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): - if noise_type != "gaussian": - raise ValueError("Only gaussian noise supported") - return lambda _s, _sn: torch.randn_like(x) +def scale_samples( + samples, + width, + height, + mode="bicubic", + sigma=None, # noqa: ARG001 +): + global bleh_scale_samples, UPSCALE_METHODS + if bleh_scale_samples is None: + bleh = EXT.get("bleh") + if bleh is not None: + bleh_scale_samples = bleh.latent_utils.scale_samples + UPSCALE_METHODS = bleh.latent_utils.UPSCALE_METHODS + else: + bleh_scale_samples = False + if bleh_scale_samples: + return bleh_scale_samples(samples, width, height, mode=mode, sigma=sigma) + if mode == "bislerp": + return bislerp(samples, width, height) + return F.interpolate(samples, size=(height, width), mode=mode) -if "nnlatentupscale" in EXT: +def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): # noqa: F811 + if noise_type != "gaussian": + raise ValueError("Only gaussian noise supported") + return lambda _s, _sn: torch.randn_like(x) - def scale_nnlatentupscale( - mode, - latent, - scale=2.0, - *, - scale_factor=0.13025, - __nlu_module=EXT["nnlatentupscale"], - ): - module = __nlu_module - mode = {"sdxl": "SDXL", "sd1": "SD 1.x"}.get(mode) - if mode is None: - raise ValueError("Bad mode") - node = module.NNLatentUpscale() - model = module.latent_resizer.LatentResizer.load_model( - node.weight_path[mode], latent.device, latent.dtype - ).to(device=latent.device) - result = ( - model(scale_factor * latent, scale=scale).to( - dtype=latent.dtype, device=latent.device - ) - / scale_factor + +def scale_nnlatentupscale(mode, latent, scale=2.0, *, scale_factor=0.13025): + if EXT_NNLATENTUPSCALE is None: + raise RuntimeError("nnlatentupscale integration not available") + mode = {"sdxl": "SDXL", "sd1": "SD 1.x"}.get(mode) + if mode is None: + raise ValueError("Bad mode") + node = EXT_NNLATENTUPSCALE.NNLatentUpscale() + model = EXT_NNLATENTUPSCALE.latent_resizer.LatentResizer.load_model( + node.weight_path[mode], latent.device, latent.dtype + ).to(device=latent.device) + result = ( + model(scale_factor * latent, scale=scale).to( + dtype=latent.dtype, device=latent.device ) - del model - return result + / scale_factor + ) + del model + return result + + +# Gaussian blur +def gaussian_blur_2d(img, kernel_size, sigma): + height = img.shape[-1] + kernel_size = min(kernel_size, height - (height % 2 - 1)) + ksize_half = (kernel_size - 1) * 0.5 + + x = torch.linspace(-ksize_half, ksize_half, steps=kernel_size) + + pdf = torch.exp(-0.5 * (x / sigma).pow(2)) + + x_kernel = pdf / pdf.sum() + x_kernel = x_kernel.to(device=img.device, dtype=img.dtype) + + kernel2d = torch.mm(x_kernel[:, None], x_kernel[None, :]) + kernel2d = kernel2d.expand(img.shape[-3], 1, kernel2d.shape[0], kernel2d.shape[1]) + + padding = [kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2] + + img = torch.nn.functional.pad(img, padding, mode="reflect") + img = torch.nn.functional.conv2d(img, kernel2d, groups=img.shape[-3]) + + return img + + +# Saliency-adaptive Noise Fusion based on High-fidelity Person-centric Subject-to-Image Synthesis (Wang et al.) +# https://github.com/CodeGoat24/Face-diffuser/blob/edff1a5178ac9984879d9f5e542c1d0f0059ca5f/facediffuser/pipeline.py#L535-L562 +def snf_guidance( + t_guidance: torch.Tensor, + s_guidance: torch.Tensor, + t_kernel_size=3, + t_sigma=1, + s_kernel_size=3, + s_sigma=1, +): + b, c, h, w = shape = t_guidance.shape + + t_softmax, s_softmax = ( + torch.softmax( + gaussian_blur_2d(torch.abs(t), ks, sig).reshape(b * c, h * w), + dim=1, + ).reshape(*shape) + for t, ks, sig in ( + (t_guidance, t_kernel_size, t_sigma), + (s_guidance, s_kernel_size, s_sigma), + ) + ) + guidance_stacked = torch.stack((t_guidance, s_guidance), dim=0) + argeps = torch.argmax( + torch.stack((t_softmax, s_softmax), dim=0), dim=0, keepdim=True + ) + return torch.gather(guidance_stacked, dim=0, index=argeps).squeeze(0) + + +class OCSLatentFormat: + def __init__(self, device, latent_format): + if latent_format.latent_rgb_factors is None: + self.rgb_factors = None + return + self.rgb_factors = torch.tensor( + latent_format.latent_rgb_factors, device=device, dtype=torch.float + ).t() + # Thanks for Joviax for the help implementing this! + self.rgb_factors_inv = torch.linalg.pinv(self.rgb_factors) + bias = getattr(latent_format, "latent_rgb_factors_bias", None) + self.rgb_factors_bias = ( + None + if bias is None + else torch.tensor(bias, device=device, dtype=torch.float) + ) + + def latent_to_rgb(self, latent: torch.Tensor) -> torch.Tensor: + # NCHW -> NHWC + if self.latent_factors is None: + raise ValueError("No RGB factors for latent type!") + return torch.nn.functional.linear( + latent.movedim(1, -1), self.rgb_factors, bias=self.rgb_factors_bias + ) + + def rgb_to_latent(self, img: torch.Tensor) -> torch.Tensor: + # NHWC + if self.latent_factors is None: + raise ValueError("No RGB factors for latent type!") + if self.rgb_factors_bias is not None: + img = img - self.rgb_factors_bias + return torch.nn.functional.linear(img, self.rgb_factors_inv) diff --git a/py/model.py b/py/model.py index c200afa..43cb23e 100644 --- a/py/model.py +++ b/py/model.py @@ -1,5 +1,3 @@ -from collections import namedtuple - import torch import comfy @@ -8,6 +6,7 @@ from comfy.k_diffusion.sampling import to_d from . import filtering from .utils import fallback +from .latent import OCSLatentFormat class History: @@ -43,12 +42,14 @@ class ModelResult: sigma, x, denoised, + have_uncond=True, **kwargs, ): self.call_idx = call_idx self.sigma = sigma self.x = x self.denoised = denoised + self.have_uncond = have_uncond for k in ("denoised_uncond", "denoised_cond", "tangents", "jdenoised"): setattr(self, k, kwargs.pop(k, None)) if len(kwargs) != 0: @@ -72,6 +73,33 @@ class ModelResult: x = x - denoised * alt_cfgpp_scale + denoised_uncond * alt_cfgpp_scale return to_d(x, sigma, denoised if not cfgpp else denoised_uncond) + def get_split_prediction( + self, + *, + x=None, + d=None, + sigma=None, + denoised=None, + denoised_uncond=None, + alt_cfgpp_scale=0, + cfgpp=False, + ): + denoised = fallback(denoised, self.denoised) + denoised_uncond = fallback(denoised_uncond, self.denoised_uncond) + x = fallback(x, self.x) + sigma = fallback(sigma, self.sigma) + if d is None: + d = self.to_d( + x=x, + sigma=sigma, + denoised=denoised, + denoised_uncond=denoised_uncond, + alt_cfgpp_scale=alt_cfgpp_scale, + cfgpp=cfgpp, + ) + denoised_pred = denoised if alt_cfgpp_scale == 0 else x - d * sigma + return (denoised_pred, d) + @property def d(self): return self.to_d() @@ -108,12 +136,7 @@ class ModelResult: return torch.linalg.norm(d_pred.sub_(d)).div_(torch.linalg.norm(d)).item() -ModelCallCacheConfig = namedtuple( - "ModelCallCacheConfig", ("size", "max_use", "threshold"), defaults=(0, 1000000, 1) -) - - -class ModelCallCache: +class OCSModel: def __init__( self, model, @@ -121,15 +144,22 @@ class ModelCallCache: s_in: torch.Tensor, extra_args: dict, *, - cache: None | dict = None, - filter: None | dict = None, + cache: dict | None = None, + filter: dict | None = None, cfg1_uncond_optimization: bool = False, - cfg_scale_override: None | int | float = None, + cfg_scale_override: int | float | None = None, ) -> None: - self.cache = ModelCallCacheConfig(**fallback(cache, {})) filtargs = fallback(filter, {}).copy() self.filters = {} - for key in ("input", "denoised", "jdenoised", "cond", "uncond", "x"): + for key in ( + "input", + "denoised", + "jdenoised", + "cond", + "uncond", + "postcfg", + "precfg", + ): filt = filtargs.pop(key, None) if filt is None: continue @@ -139,12 +169,12 @@ class ModelCallCache: self.extra_args = extra_args self.cfg1_uncond_optimization = cfg1_uncond_optimization self.cfg_scale_override = cfg_scale_override - self.is_rectified_flow = x.shape[1] == 16 and isinstance( + self.is_rectified_flow = isinstance( model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST ) - if self.cache.size < 1: - return - self.reset_cache() + self.latent_format = OCSLatentFormat( + x.device, model.inner_model.inner_model.latent_format + ) def maybe_filter( self, name: str, latent: torch.Tensor, *args: list, **kwargs: dict @@ -160,11 +190,11 @@ class ModelCallCache: if not self.filters: return result result = result.clone() - for key in ("denoised", "cond", "uncond", "jdenoised", "x"): + for key in ("denoised", "cond", "uncond", "jdenoised"): filt = self.filters.get(key) if filt is None: continue - attk = f"denoised_{key}" if key in ("cond", "uncond") else key + attk = f"denoised_{key}" if key in {"cond", "uncond"} else key inpval = getattr(result, attk, None) if inpval is None: continue @@ -177,33 +207,6 @@ class ModelCallCache: fr.kvs |= {f"{k}_curr": v for k, v in frmr.kvs.items()} return fr - def reset_cache(self) -> None: - size = self.cache.size - self.slot = [None] * size - self.slot_use = [self.cache.max_use] * size - - def get(self, idx: int, *, jvp: bool = False) -> None | ModelResult: - idx -= self.cache.threshold - if ( - idx >= self.cache.size - or idx < 0 - or self.slot[idx] is None - or self.slot_use[idx] < 1 - ): - return None - result = self.slot[idx] - if jvp and result.jdenoised is None: - return None - self.slot_use[idx] -= 1 - return result - - def set(self, idx: int, mr: ModelResult) -> None: - idx -= self.cache.threshold - if idx < 0 or idx >= self.cache.size: - return - self.slot_use[idx] = self.cache.max_use - self.slot[idx] = mr - def call_model( self, x: torch.Tensor, sigma: torch.Tensor, **kwargs: dict ) -> torch.Tensor: @@ -247,27 +250,56 @@ class ModelCallCache: "model_call": call_index, "orig_x": x, }) - result = self.get(call_index, jvp=tangents is not None) - # print( - # f"MODEL: idx={call_index}, size={self.size}, threshold={self.threshold}, cached={result is not None}" - # ) - if result is not None: - self._fr_add_mr(filter_refs, result) - result = self.filter_result(result, default_ref=x, refs=filter_refs) - return result comfy.model_management.throw_exception_if_processing_interrupted() model_options = self.extra_args.get("model_options", {}).copy() denoised_cond = denoised_uncond = None + have_uncond = True def postcfg(args): - nonlocal denoised_cond, denoised_uncond + nonlocal denoised_cond, denoised_uncond, have_uncond denoised_uncond = args["uncond_denoised"] denoised_cond = args["cond_denoised"] + result = args["denoised"] + if "postcfg" in self.filters: + result = self.maybe_filter( + "postcfg", + result, + refs=filter_refs + | filtering.FilterRefs({ + "postcfg_input": args["input"], + "postcfg_sigma": args["sigma"], + "postcfg_denoised_cond": denoised_cond, + "postcfg_denoised_uncond": denoised_uncond, + }), + ) if denoised_uncond is None: + have_uncond = False denoised_uncond = denoised_cond - return args["denoised"] + return result + + def precfg(args): + conds_out = args["conds_out"] + precfg_refs = ( + filter_refs + | filtering.FilterRefs({ + "precfg_input": args["input"], + "precfg_sigma": args["sigma"], + "precfg_cond_scale": args["cond_scale"], + }) + | filtering.FilterRefs({ + f"precfg_cond_{idx}": cond for idx, cond in enumerate(conds_out) + }) + ) + return [ + self.maybe_filter( + "precfg", + curr_cond, + refs=precfg_refs | filtering.FilterRefs({"cond_idx": cond_idx}), + ) + for cond_idx, curr_cond in enumerate(conds_out) + ] orig_cfg_scale = self.set_inner_cfg_scale(cfg_scale_override) @@ -277,9 +309,16 @@ class ModelCallCache: disable_cfg1_optimization=require_uncond or not self.cfg1_uncond_optimization, ) + if "precfg" in self.filters: + model_options = comfy.model_patcher.set_model_options_pre_cfg_function( + model_options, + precfg, + ) extra_args = self.extra_args | {"model_options": model_options} s_in = fallback(s_in, self.s_in) + if s_in.shape[0] != x.shape[0]: + s_in = self.s_in = x.new_ones((x.shape[0],)) x = self.maybe_filter("input", x, refs=filter_refs) def call_model(x, sigma, **kwargs): @@ -293,10 +332,10 @@ class ModelCallCache: sigma, x, denoised, + have_uncond=have_uncond, denoised_uncond=denoised_uncond, denoised_cond=denoised_cond, ) - self.set(call_index, mr) self._fr_add_mr(filter_refs, mr) mr = self.filter_result(mr, default_ref=x, refs=filter_refs) return mr @@ -307,11 +346,11 @@ class ModelCallCache: sigma, x, denoised, + have_uncond=have_uncond, jdenoised=denoised_prime, denoised_uncond=denoised_uncond, denoised_cond=denoised_cond, ) - self.set(call_index, mr) self._fr_add_mr(filter_refs, mr) mr = self.filter_result(mr, default_ref=x, refs=filter_refs) return mr diff --git a/py/nodes.py b/py/nodes.py index 8a3d4d1..d82023b 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -1,21 +1,75 @@ +import comfy import yaml -import comfy +import torch +from tqdm import tqdm + +from .external import MODULES, IntegratedNode +from .restart import Restart from .sampling import composable_sampler -from .substep_sampling import StepSamplerChain, StepSamplerGroups, ParamGroup from .step_samplers import STEP_SAMPLERS from .substep_merging import MERGE_SUBSTEPS_CLASSES -from .restart import Restart +from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups +from .filtering import make_filter -DEFAULT_YAML_PARAMS = """\ -# JSON or YAML parameters -s_noise: 1.0 -eta: 1.0 -""" +try: + from comfy_execution import validation as comfy_validation + + if not hasattr(comfy_validation, "validate_node_input"): + raise NotImplementedError + HAVE_COMFY_UNION_TYPE = comfy_validation.validate_node_input("B", "A,B") +except (ImportError, NotImplementedError): + HAVE_COMFY_UNION_TYPE = False +except Exception as exc: + HAVE_COMFY_UNION_TYPE = False + tqdm.write( + f"** OCS: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}" + ) + +PARAM_INPUT_TYPES = frozenset(( + "IMAGE", + "OCS_NOISE", + "SAMPLER", + "SIGMAS", + "SONAR_CUSTOM_NOISE", + "UPSCALE_MODEL", + "VAE", +)) + +NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE")) + +if not HAVE_COMFY_UNION_TYPE: + + class Wildcard(str): + __slots__ = ("whitelist",) + + @classmethod + def __new__(cls, s, *args: list, whitelist=None, **kwargs: dict): + result = super().__new__(s, *args, **kwargs) + result.whitelist = whitelist + return result + + def __ne__(self, other): + return False if self.whitelist is None else other not in self.whitelist + + WILDCARD_NOISE = Wildcard("*", whitelist=NOISE_INPUT_TYPES) + WILDCARD_PARAM = Wildcard("*", whitelist=PARAM_INPUT_TYPES) +else: + WILDCARD_NOISE = ",".join(NOISE_INPUT_TYPES) + WILDCARD_PARAM = ",".join(PARAM_INPUT_TYPES) + +PARAM_INPUT_TYPES_HINT = ( + f"The following input types are supported: {', '.join(PARAM_INPUT_TYPES)}" +) +NOISE_INPUT_TYPES_HINT = ( + f"The following input types are supported: {', '.join(NOISE_INPUT_TYPES)}" +) + +DEFAULT_YAML_PARAMS = "# YAML/JSON parameters\n" -class SamplerNode: +class SamplerNode(metaclass=IntegratedNode): RETURN_TYPES = ("SAMPLER",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Overly Complicated Sampling main sampler node. Can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input." @@ -32,7 +86,8 @@ class SamplerNode: "groups": ( "OCS_GROUPS", { - "tooltip": "Connect OCS substep groups here which are output from the OCS Group node." + "tooltip": "Connect OCS substep groups here which are output from the OCS Group node.", + "forceInput": True, }, ), }, @@ -41,12 +96,14 @@ class SamplerNode: "OCS_PARAMS", { "tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.", + "forceInput": True, }, ), "parameters": ( "STRING", { - "default": DEFAULT_YAML_PARAMS, + "default": "", + "placeholder": DEFAULT_YAML_PARAMS, "multiline": True, "dynamicPrompts": False, "tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.", @@ -62,6 +119,7 @@ class SamplerNode: params_opt=None, parameters="", ): + MODULES.initialize() options = {} parameters = parameters.strip() if parameters: @@ -80,7 +138,7 @@ class SamplerNode: ) -class GroupNode: +class GroupNode(metaclass=IntegratedNode): RETURN_TYPES = ("OCS_GROUPS",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Over Complicated Sampling group definition node." @@ -130,6 +188,7 @@ class GroupNode: "OCS_SUBSTEPS", { "tooltip": "Connect output from an OCS Substeps node here.", + "forceInput": True, }, ), }, @@ -138,18 +197,21 @@ class GroupNode: "OCS_GROUPS", { "tooltip": "You may optionally connect the output from another OCS Group node here. Only one group per step is used, matching (based on time or other constraints) starts with the OCS Group node furthest from the OCS Sampler.", + "forceInput": True, }, ), "params_opt": ( "OCS_PARAMS", { "tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.", + "forceInput": True, }, ), "parameters": ( "STRING", { - "default": DEFAULT_YAML_PARAMS, + "default": "", + "placeholder": DEFAULT_YAML_PARAMS, "multiline": True, "dynamicPrompts": False, "tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.", @@ -170,6 +232,7 @@ class GroupNode: params_opt=None, parameters="", ): + MODULES.initialize() group = StepSamplerGroups() if groups_opt is None else groups_opt.clone() chain = substeps.clone() chain.merge_method = merge_method @@ -190,7 +253,7 @@ class GroupNode: return (group,) -class SubstepsNode: +class SubstepsNode(metaclass=IntegratedNode): RETURN_TYPES = ("OCS_SUBSTEPS",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Overly Complicated Sampling substeps definition node. Used to define a sampler type and other sampler-specific parameters." @@ -225,18 +288,21 @@ class SubstepsNode: "OCS_SUBSTEPS", { "tooltip": "Optionally connect another OCS Substeps node here. Substeps will run in order, starting from the OCS Substeps node FURTHEST from the OCS Group node.", + "forceInput": True, }, ), "params_opt": ( "OCS_PARAMS", { "tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.", + "forceInput": True, }, ), "parameters": ( "STRING", { - "default": DEFAULT_YAML_PARAMS, + "default": "", + "placeholder": f"{DEFAULT_YAML_PARAMS}s_noise: 1.0\neta: 1.0\n", "multiline": True, "dynamicPrompts": False, "tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.", @@ -253,6 +319,7 @@ class SubstepsNode: params_opt=None, **kwargs, ): + MODULES.initialize() if substeps_opt is not None: chain = substeps_opt.clone() else: @@ -270,30 +337,23 @@ class SubstepsNode: return (chain,) -class Wildcard(str): - __slots__ = () - - def __ne__(self, _unused): - return False - - -class ParamNode: +class ParamNode(metaclass=IntegratedNode): RETURN_TYPES = ("OCS_PARAMS",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Overly Complicated Sampling parameter definition node. Used to set parameters like custom noise types that require an input." - OUTPUT_TYPES = ( + OUTPUT_TOOLTIPS = ( "Can be connected to another OCS Param or OCS MultiParam node or any other OCS node that takes OCS_PARAMS as an input.", ) FUNCTION = "go" - WC = Wildcard("*") - - OCS_PARAM_TYPES = { + OCS_PARAM_INPUT_TYPES = { "custom_noise": lambda v: hasattr(v, "make_noise_sampler"), "merge_sampler": lambda v: isinstance(v, StepSamplerChain), "restart_custom_noise": lambda v: hasattr(v, "make_noise_sampler"), - "SAMPLER": lambda _v: True, + "sampler": lambda _v: True, + "vae": lambda _v: True, + "upscale_model": lambda _v: True, } @classmethod @@ -301,15 +361,16 @@ class ParamNode: return { "required": { "key": ( - tuple(cls.OCS_PARAM_TYPES.keys()), + tuple(cls.OCS_PARAM_INPUT_TYPES.keys()), { "tooltip": "Used to set the type of custom parameter.", }, ), "value": ( - cls.WC, + WILDCARD_PARAM, { - "tooltip": "Connect the type of value expected by the key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the key you will get an error when you run the workflow.", + "tooltip": f"Connect the type of value expected by the key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the key you will get an error when you run the workflow.\n{PARAM_INPUT_TYPES_HINT}", + "forceInput": True, }, ), }, @@ -318,14 +379,17 @@ class ParamNode: "OCS_PARAMS", { "tooltip": "You may optionally connect the output from other OCS Param or OCS MultiParam nodes here to set multiple parameters.", + "forceInput": True, }, ), "parameters": ( "STRING", { - "default": "# Additional YAML or JSON parameters\n", + "default": "", + "placeholder": "# Additional YAML or JSON parameters", "multiline": True, "dynamicPrompts": False, + "defaultInput": True, "tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.", }, ), @@ -347,7 +411,8 @@ class ParamNode: return f"{key}_{rename}" def go(self, *, key, value, params_opt=None, parameters=""): - if not self.OCS_PARAM_TYPES[key](value): + MODULES.initialize() + if not self.OCS_PARAM_INPUT_TYPES[key](value): raise ValueError(f"CSamplerParam: Bad value type for key {key}") if parameters: extra_params = yaml.safe_load(parameters) @@ -364,7 +429,7 @@ class ParamNode: return (params,) -class MultiParamNode(ParamNode): +class MultiParamNode(ParamNode, metaclass=IntegratedNode): RETURN_TYPES = ("OCS_PARAMS",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Overly Complicated Sampling parameter definition node. Used to set parameters like custom noise types that require an input. Like the OCS Param node but allows setting multiple parameters at the same time." @@ -379,7 +444,7 @@ class MultiParamNode(ParamNode): @classmethod def INPUT_TYPES(cls): param_keys = ( - ("", *ParamNode.OCS_PARAM_TYPES.keys()), + ("", *ParamNode.OCS_PARAM_INPUT_TYPES.keys()), { "tooltip": "Used to set the type of custom parameter.", }, @@ -393,26 +458,30 @@ class MultiParamNode(ParamNode): "OCS_PARAMS", { "tooltip": "You may optionally connect the output from other OCS MultiParam or OCS Param nodes here to set multiple parameters.", + "forceInput": True, }, ), "parameters": ( "STRING", { - "default": """\ + "default": "", + "placeholder": """\ # Additional YAML or JSON parameters # Should be an object with key corresponding to the index of the input """, "multiline": True, "dynamicPrompts": False, + "defaultInput": True, "tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.", }, ), } | { f"value_opt_{idx}": ( - ParamNode.WC, + WILDCARD_PARAM, { - "tooltip": "Connect the type of value expected by the corresponding key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the corresponding key you will get an error when you run the workflow.", + "tooltip": f"Connect the type of value expected by the corresponding key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the corresponding key you will get an error when you run the workflow.\n{PARAM_INPUT_TYPES_HINT}", + "forceInput": True, }, ) for idx in range(1, cls.PARAM_COUNT + 1) @@ -420,6 +489,7 @@ class MultiParamNode(ParamNode): } def go(self, *, params_opt=None, parameters="", **kwargs): + MODULES.initialize() params = ParamGroup(items={}) if params_opt is None else params_opt.clone() if parameters: extra_params = yaml.safe_load(parameters) @@ -434,7 +504,7 @@ class MultiParamNode(ParamNode): key, value = kwargs.get(f"key_{idx}"), kwargs.get(f"value_opt_{idx}") if not key or value is None: continue - if not self.OCS_PARAM_TYPES[key](value): + if not self.OCS_PARAM_INPUT_TYPES[key](value): raise ValueError(f"CSamplerParamGroup: Bad value type for key {key}") extra = extra_params.get(str(idx)) key = self.get_renamed_key(key, extra) @@ -445,7 +515,7 @@ class MultiParamNode(ParamNode): return (params,) -class SimpleRestartSchedule: +class SimpleRestartSchedule(metaclass=IntegratedNode): RETURN_TYPES = ("SIGMAS",) CATEGORY = "sampling/custom_sampling/OCS" DESCRIPTION = "Overly Complicated Sampling simple Restart schedule node. Allows generating a Restart sampling schedule based on a text definition." @@ -478,8 +548,9 @@ class SimpleRestartSchedule: "schedule": ( "STRING", { - "default": """\ -# YAML or JSON restart schedule + "default": "", + "placeholder": """\ +# YAML or JSON restart schedule. Example: # Every 5 steps, jump back 3 steps - [5, -3] # Jump to schedule item 0 @@ -493,7 +564,8 @@ class SimpleRestartSchedule: }, } - def go(self, *, sigmas, start_step=0, schedule="[]"): + def go(self, *, sigmas, start_step=0, schedule=""): + MODULES.initialize() if schedule: parsed_schedule = yaml.safe_load(schedule) if parsed_schedule is not None: @@ -506,12 +578,12 @@ class SimpleRestartSchedule: return (Restart.simple_schedule(sigmas, start_step, parsed_schedule),) -class ModelSetMaxSigmaNode: +class ModelSetMaxSigmaNode(metaclass=IntegratedNode): RETURN_TYPES = ("MODEL",) CATEGORY = "hacks" DESCRIPTION = "Allows forcing a model's maximum and minumum sigmas to a specified value. You generally do NOT want to connect this to a sampler node. Connect it to a scheduler node (i.e. BasicScheduler) instead." OUTPUT_TOOLTIPS = ( - "Patched model. Can be connected to a scheduler node (i.e. BasicScheduler). Generally NOT recommended to connect to an actual sampler.", + "Patched model. Can be connected to a scheduler node (i.e. BasicScheduler). Generally NOT recommended to connect to an actual sampler, the main use case is only for generating sigmas.", ) FUNCTION = "go" @@ -529,7 +601,7 @@ class ModelSetMaxSigmaNode: "mode": ( ("recalculate", "simple_multiply"), { - "tooltip": "Mode use for setting sigmas in the patched model. Recalculate should generally be more accurate.", + "tooltip": "Mode to use when setting sigmas in the patched model. Recalculate should generally be more accurate.", }, ), "sigma_max": ( @@ -540,7 +612,7 @@ class ModelSetMaxSigmaNode: "max": 10000.0, "step": 0.01, "round": False, - "tooltip": "You can set the maximum sigma here. If you use a negative value, it will be interpreted as the absolute value for the max sigma. If you use a positive value it will be interpreted as a percentage (where 1.0 signified 100%). Schedules generated with the patched model should start from sigma_max (or close to it).", + "tooltip": "You can set the maximum sigma here. If you use a positive value, it will be interpreted as the absolute value for the max sigma. If you use a negative value it will be interpreted as a percentage of the current value (where 1.0 signifies 100%). Schedules generated with the patched model should start from sigma_max (or close to it).", }, ), "fake_sigma_min": ( @@ -551,13 +623,14 @@ class ModelSetMaxSigmaNode: "max": 1000.0, "step": 0.01, "round": False, - "tooltip": "You can set the minimum sigma here. Disabled if set to 0. If you use a negative value, it will be interpreted as the absolute value for the max sigma. If you use a positive value it will be interpreted as a percentage (where 1.0 signified 100%). Schedules generated with the patched model should end with [sigma_min, 0]. NOTE: May not work with some schedulers. I recommend leaving this at 0 unless you know you need it (and even then it may not work).", + "tooltip": "You can set the minimum sigma here. Disabled if set to 0. If you use a positive value, it will be interpreted as the absolute value for the max sigma. If you use a negative value it will be interpreted as a percentage of the current value (where 1.0 signifies 100%). Schedules generated with the patched model should end with [sigma_min, 0]. NOTE: May not work with some schedulers. I recommend leaving this at 0 unless you know you need it (and even then it may not work).", }, ), } } def go(self, model, mode="recalculate", sigma_max=-1.0, fake_sigma_min=0.0): + MODULES.initialize() if sigma_max == 0: raise ValueError("ModelSetMaxSigma: Invalid sigma_max value") if mode not in ("recalculate", "simple_multiply"): @@ -598,12 +671,113 @@ class ModelSetMaxSigmaNode: "ModelSetMaxSigma: Invalid fake_min_sigma value, result max <= min" ) model.add_object_patch("model_sampling", ms) - print( - f"ModelSetMaxSigma: Set model sigmas({mode}): old_max={orig_max_sigma:.04}, old_min={orig_min_sigma:.03}, new_max={new_max_sigma:.04}, new_min={new_min_sigma:.03}" + tqdm.write( + f"OCS: ModelSetMaxSigma: Set model sigmas({mode}): old_max={orig_max_sigma:.04}, old_min={orig_min_sigma:.03}, new_max={new_max_sigma:.04}, new_min={new_min_sigma:.03}" ) return (model,) +class ApplyFilterLatent(metaclass=IntegratedNode): + RETURN_TYPES = ("LATENT",) + CATEGORY = "sampling/custom_sampling/OCS" + DESCRIPTION = "Allows applying an OCS filter to any latent. Define a filter block in yaml_config." + + FUNCTION = "go" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ( + "LATENT", + { + "tooltip": "Latent input. Note: This node does not care about masks.", + }, + ), + "seed": ( + "INT", + { + "default": 0, + "min": 0, + "max": 0xFFFFFFFFFFFFFFFF, + "tooltip": "Seed to use for generated noise.", + }, + ), + "yaml_config": ( + "STRING", + { + "default": "", + "placeholder": """\ +# YAML or JSON filter definition +""", + "multiline": True, + "dynamicPrompts": False, + "tooltip": "Enter your filter definition here. There is essentially no error handling.", + }, + ), + }, + } + + @classmethod + def get_latent_samples(cls, latent: dict) -> torch.Tensor: + samples = latent["samples"] + batch_indexes = latent.get("batch_index") + if batch_indexes is None: + return samples.clone() + return samples[tuple(batch_indexes), ...].clone() + + def go(self, *, latent: dict, seed: int, yaml_config: str) -> tuple: + MODULES.initialize() + torch.manual_seed(seed) + samples = self.get_latent_samples(latent) + config = yaml.safe_load(yaml_config) + if not config: + return ({"samples": samples},) + if not isinstance(config, dict) or "filter" not in config: + raise ValueError( + "Bad YAML config type (must be object) or missing filter key in config" + ) + filter_def = config.get("filter") + if not isinstance(filter_def, dict): + raise ValueError("Bad type for filter definition, must be object") + ocs_filter = make_filter(filter_def) + new_samples = ocs_filter.apply(samples.to(dtype=torch.float32)).to(samples) + return ({"samples": new_samples},) + + +class ApplyFilterImage(ApplyFilterLatent): + DESCRIPTION = "Allows applying an OCS filter to any image. Define a filter block in yaml_config." + RETURN_TYPES = ("IMAGE",) + + @classmethod + def INPUT_TYPES(cls): + result = super().INPUT_TYPES() + del result["required"]["latent"] + result["required"] = { + "image": ( + "IMAGE", + {"tooltip": "Image input."}, + ), + } | result["required"] + return result + + def go(self, *, image: torch.Tensor, seed: int, yaml_config: str) -> tuple: + image = image.clone() + if image.ndim == 3: + image = image.unsqueeze(0) + result = ( + super() + .go( + latent={"samples": image.movedim(-1, 1)}, + seed=seed, + yaml_config=yaml_config, + )[0]["samples"] + .movedim(1, -1) + .to(image) + ) + return (result,) + + __all__ = ( "SamplerNode", "GroupNode", diff --git a/py/noise.py b/py/noise.py index 3c04178..fc29d17 100644 --- a/py/noise.py +++ b/py/noise.py @@ -1,8 +1,10 @@ import gc +import math import random import scipy import torch +from tqdm import tqdm from .filtering import Filter, make_filter from .utils import scale_noise, fallback @@ -15,15 +17,16 @@ class ImmiscibleNoise(Filter): "size": 0, "batching": "channel", "maximize": False, + "distance_scale": 0.0, + "distance_scale_ref": None, } def __call__(self, noise_sampler, x_ref, *, refs=None): - if not self.check_applies(refs): + if self.size == 0 or self.strength == 0 or not self.check_applies(refs): return noise_sampler() + size = self.size if self.strength == 1.0 else self.size + 1 return self.apply( - torch.cat(tuple(noise_sampler() for _ in range(self.size))) - if self.size > 0 - else noise_sampler(), + torch.cat(tuple(noise_sampler() for _ in range(size))), default_ref=x_ref, refs=refs, output_shape=x_ref.shape, @@ -32,52 +35,103 @@ class ImmiscibleNoise(Filter): def filter(self, latent, ref_latent, *, refs, output_shape): if self.size == 0: return latent + offset = 0 if self.strength == 1.0 else output_shape[0] return self.unbatch( - self.immiscible(self.batch(latent), self.batch(ref_latent)), output_shape + self.immiscible(self.batch(latent[offset:]), self.batch(ref_latent)), + output_shape, ) def batch(self, latent): if self.batching == "batch": return latent sz = latent.shape - if latent.ndim != 4: - raise ValueError("Both latent and reference must be four-dimensional") if self.batching == "channel": - return latent.view(sz[0] * sz[1], *sz[2:]) + return latent.reshape(sz[0] * sz[1], *sz[2:]) + if self.batching == "frame": + if latent.ndim != 5: + raise ValueError( + "Both latent and reference must be five-dimensional for frame mode" + ) + return latent.permute(0, 2, 1, 3, 4).reshape(sz[0] * sz[2], sz[1], *sz[3:]) if self.batching == "row": - return latent.view(sz[0] * sz[1] * sz[2], sz[3]) + return latent.reshape(math.prod(sz[:-1]), sz[-1]) if self.batching == "column": + if latent.ndim != 4: + raise ValueError( + "Both latent and reference must be four-dimensional for column mode" + ) return latent.permute(0, 1, 3, 2).reshape(sz[0] * sz[1] * sz[3], sz[2]) raise ValueError("Bad Immmiscible noise batching type") def unbatch(self, latent, sz): if self.batching == "column": - return latent.view(*sz[:2], sz[3], sz[2]).permute(0, 1, 3, 2) - return latent.view(*sz) + return latent.reshape(*sz[:2], sz[3], sz[2]).permute(0, 1, 3, 2) + if self.batching == "frame": + return latent.reshape(sz[0], sz[2], sz[1], *sz[3:]).permute(0, 2, 1, 3, 4) + return latent.reshape(*sz) - # Based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 - # Idea from https://github.com/Clybius - def immiscible(self, latent, ref_latent): + # Originally based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 + # Idea for use with inference as well as implementation help from https://github.com/Clybius + def immiscible( + self, + latent: torch.Tensor, + ref_latent: torch.Tensor, + *, + out_latent: torch.Tensor | None = None, + return_idxs=False, + ): # "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303 # Minimize latent-noise pairs over a batch - n = latent.shape[0] - ref_latent_expanded = ( - ref_latent.half().unsqueeze(1).expand(-1, n, *ref_latent.shape[1:]) - ) - latent_expanded = ( - latent.half().unsqueeze(0).expand(ref_latent.shape[0], *latent.shape) - ) - dist = (ref_latent_expanded - latent_expanded) ** 2 - dist = dist.mean(list(range(2, dist.dim()))).cpu() + batch = latent.shape[0] + ref_latent = ref_latent.detach().clone() + if self.distance_scale == 0: + ref_latent_expanded = ref_latent.unsqueeze(1).expand( + -1, batch, *ref_latent.shape[1:] + ) + latent_expanded = latent.unsqueeze(0).expand( + ref_latent.shape[0], *latent.shape + ) + dist = (ref_latent_expanded - latent_expanded) ** 2 + del ref_latent_expanded, latent_expanded + dist = dist.mean(tuple(range(2, dist.dim()))) + else: + dist = torch.linalg.vector_norm( + fallback(self.distance_scale_ref, self.distance_scale) + * ref_latent.flatten(start_dim=1).unsqueeze(1) + - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0), + dim=2, + ) + dist = dist.half() try: assign_mat = scipy.optimize.linear_sum_assignment( - dist, maximize=self.maximize + dist.cpu(), maximize=self.maximize ) - except ValueError as _exc: - # print("\nImmiscible: Failed optimization, skipping") - return latent[: ref_latent.shape[0]] - # print("IMM IDX", assign_mat[1]) - return latent[assign_mat[1]] + except ValueError as exc: + tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}") + return ( + None + if return_idxs + else fallback(out_latent, latent)[: ref_latent.shape[0]] + ) + return ( + assign_mat if return_idxs else fallback(out_latent, latent)[assign_mat[1]] + ) + + def immiscible_simple( + self, + latent: torch.Tensor, + ref_latent: torch.Tensor, + *, + out_latent: torch.Tensor | None = None, + ) -> torch.Tensor: + return self.unbatch( + self.immiscible( + self.batch(latent), + self.batch(ref_latent), + out_latent=self.batch(out_latent) if out_latent is not None else None, + ), + ref_latent.shape, + ) class NoiseSamplerCache: @@ -93,7 +147,8 @@ class NoiseSamplerCache: batch_size=1, caching=True, cache_reset_interval=9999, - set_seed=False, + set_seed=True, + seed_offset=1, scale=1.0, normalize_dims=(-3, -2, -1), immiscible=None, @@ -103,7 +158,7 @@ class NoiseSamplerCache: self.x = x self.mega_x = None self.seed = seed - self.seed_offset = 0 + self.seed_offset = seed_offset self.min_sigma = min_sigma self.max_sigma = max_sigma self.cache = {} @@ -123,6 +178,11 @@ class NoiseSamplerCache: if set_seed: random.seed(seed) torch.manual_seed(seed) + if self.seed_offset > 0: + for _ in range(self.seed_offset): + _ = torch.randn_like(x) + else: + self.seed_offset = 0 def reset_cache(self): self.cache = {} @@ -160,10 +220,14 @@ class NoiseSamplerCache: size, sigma, sigma_next, + *, immiscible=None, + sigmas=None, ): size = min(size, self.batch_size) - cache_key = (nsobj, size) + if immiscible is None: + immiscible = self.immiscible + cache_key = (nsobj, size, hash(immiscible)) if self.caching: noise_sampler = self.cache.get(cache_key) if noise_sampler: @@ -173,14 +237,26 @@ class NoiseSamplerCache: curr_x = self.mega_x[: self.x.shape[0] * size, ...] if nsobj is None: - def ns(_s, _sn, *_unused, **_unusedkwargs): - return torch.randn_like(curr_x) + def ns(*_unused, **_unusedkwargs): + noise = torch.randn( + curr_x.shape, + dtype=curr_x.dtype, + layout=curr_x.layout, + device="cpu" if self.cpu_noise else curr_x.device, + ) + if noise.device != curr_x.device: + return noise.to(curr_x.device) + return noise else: + if sigmas is not None: + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + else: + sigma_min, sigma_max = self.min_sigma, self.max_sigma ns = nsobj.make_noise_sampler( curr_x, - self.min_sigma, - self.max_sigma, + sigma_min, + sigma_max, seed=curr_seed, normalized=False, cpu=self.cpu_noise, @@ -189,8 +265,6 @@ class NoiseSamplerCache: orig_h, orig_w = self.x.shape[-2:] remain = 0 noise = None - if immiscible is None: - immiscible = self.immiscible def noise_sampler_( curr_sigma, diff --git a/py/sampling.py b/py/sampling.py index a340324..c52575a 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -3,7 +3,7 @@ from tqdm.auto import trange from .filtering import FILTER_HANDLERS, FilterRefs -from .model import ModelCallCache +from .model import OCSModel from .noise import NoiseSamplerCache from .substep_sampling import SamplerState from .substep_merging import MERGE_SUBSTEPS_CLASSES @@ -44,6 +44,7 @@ def composable_sampler( return torch.randn_like(x) restart_params = copts.get("restart", {}) + restart_enabled = restart_params.get("enabled", True) restart_custom_noise = copts.get("restart_custom_noise") if isinstance(restart_custom_noise, str): restart_custom_noise = copts.get(f"restart_custom_noise_{restart_custom_noise}") @@ -54,7 +55,7 @@ def composable_sampler( ) ss = SamplerState( - ModelCallCache( + OCSModel( model, x, x.new_ones((x.shape[0],)), @@ -78,12 +79,14 @@ def composable_sampler( nsc = NoiseSamplerCache( x, extra_args.get("seed", 42), - sigmas[-1], - sigmas[0], + sigmas[sigmas > 0].min(), + sigmas.max(), **copts.get("noise", {}), ) ss.noise = nsc - sigma_chunks = tuple(restart.split_sigmas(sigmas)) + sigma_chunks = ( + tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((0.0, sigmas),) + ) step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks) ss.total_steps = step_count step = 0 @@ -114,10 +117,6 @@ def composable_sampler( if idx > 0: ss.update(idx, step=step, substep=0) nsc.update_x(x) - # print( - # f"STEP {step + 1:>3}: {ss.sigma.item():.03} -> {ss.sigma_next.item():.03} || up={ss.sigma_up.item():.03}, down={ss.sigma_down.item():.03}" - # ) - ss.model.reset_cache() nsc.update_x(x) merge_sampler = find_merge_sampler(merge_samplers, ss) if merge_sampler is None: diff --git a/py/step_samplers.py b/py/step_samplers.py deleted file mode 100644 index d168ce7..0000000 --- a/py/step_samplers.py +++ /dev/null @@ -1,2576 +0,0 @@ -import contextlib -import inspect -import math -import os -import typing -import warnings - -import torch -import tqdm -import torchsde -import numpy - -import comfy -from comfy.k_diffusion.sampling import ( - get_ancestral_step, - to_d, -) - -from . import filtering, noise, res_support, utils -from . import expression as expr -from .utils import fallback - -HAVE_DIFFRAX = HAVE_TDE = HAVE_TODE = False - -with contextlib.suppress(ImportError): - import torchdiffeq as tde - - HAVE_TDE = True - -with contextlib.suppress(ImportError, RuntimeError): - import torchode as tode - - HAVE_TODE = True - -if not os.environ.get("COMFYUI_OCS_NO_DIFFRAX_SOLVER"): - with contextlib.suppress(ImportError): - import diffrax - import jax - - if not os.environ.get("COMFYUI_OCS_NO_DISABLE_JAX_PREALLOCATE"): - os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false" - # jax.config.update("jax_enable_x64", True) - - HAVE_DIFFRAX = True - - -class SamplerResult: - CLONE_KEYS = ( - "denoised_cond", - "denoised_uncond", - "denoised", - "final", - "is_rectified_flow", - "noise_pred", - "noise_sampler", - "s_noise", - "sampler", - "sigma_down", - "sigma_next", - "sigma_up", - "sigma", - "step", - "substep", - "x_", - ) - - def __init__( - self, - ss, - sampler, - x, - sigma_up=None, - *, - split_result=None, - sigma=None, - sigma_next=None, - sigma_down=None, - s_noise=None, - noise_sampler=None, - final=True, - ): - self.is_rectified_flow = ss.model.is_rectified_flow - self.sampler = sampler - self.sigma_up = fallback(sigma_up, ss.sigma.new_zeros(1)) - self.s_noise = fallback(s_noise, sampler.s_noise) - self.sigma = fallback(sigma, ss.sigma) - self.sigma_next = fallback(sigma_next, ss.sigma_next) - self.sigma_down = fallback(sigma_down, self.sigma_next) - self.noise_sampler = fallback(noise_sampler, sampler.noise_sampler) - self.final = final - self.step = ss.step - self.substep = ss.substep - self.x_ = x - if split_result is not None: - self.denoised, self.noise_pred = split_result - elif x is None: - raise ValueError("SamplerResult requires at least one of x, split_result") - else: - self.denoised = self.noise_pred = None - _ = self.extract_pred(ss) - self.denoised_uncond = ss.hcur.denoised_uncond - self.denoised_cond = ss.hcur.denoised_cond - - def get_noise(self, *, scaled=True, ss=None): - if self.sigma_next == 0 or self.noise_scale == 0: - return torch.zeros_like(self.x_) - return self.noise_sampler( - self.sigma, - self.sigma_next, - out_hw=self.x.shape[-2:], - x_ref=self.x, - refs=filtering.FilterRefs.from_sr(self) if ss is None else ss.refs, - ).mul_(self.noise_scale if scaled else 1.0) - - def extract_pred(self, ss): - if self.denoised is None or self.noise_pred is None: - self.denoised, self.noise_pred = utils.extract_pred( - ss.hcur.x, self.x_, ss.sigma, self.sigma_down - ) - return self.denoised, self.noise_pred - - @property - def x(self): - if self.x_ is None: - self.x_ = self.denoised + self.sigma_down * self.noise_pred - return self.x_ - - @property - def noise_scale(self): - return self.sigma_up * self.s_noise - - def noise_x(self, x=None, scale=1.0, *, ss=None): - x = fallback(x, self.x) - if self.sigma_next == 0 or self.noise_scale == 0: - return x - noise = self.get_noise(ss=ss) * scale - if not self.is_rectified_flow: - return x + noise - x_coeff = (1 - self.sigma_next) / (1 - self.sigma_down) - # print(f"\nRF noise: {x_coeff}") - return x_coeff * x + noise - - def clone(self): - obj = self.__new__(self.__class__) - for k in self.CLONE_KEYS: - if hasattr(self, k): - setattr(obj, k, getattr(self, k)) - return obj - - -class StepSamplerContext: - def __init__(self, sampler, *args, **kwargs): - self.sampler = sampler - self.args = args - self.kwargs = kwargs - - def __enter__(self): - if self.sampler.ss is not None: - raise RuntimeError("Cannot reenter prepared sampler in context manager!") - self.sampler.prepare(*self.args, **self.kwargs) - return self.sampler - - def __exit__(self, *_unused): - self.sampler.reset() - - -class SingleStepSampler: - name = None - self_noise = 0 - model_calls = 0 - ancestralize = False - sample_sigma_zero = False - immiscible = None - allow_cfgpp = False - allow_alt_cfgpp = False - - def __init__( - self, - *, - noise_sampler=None, - substeps=1, - s_noise=1.0, - eta=1.0, - eta_retry_increment=0, - dyn_eta_start=None, - dyn_eta_end=None, - weight=1.0, - pre_filter=None, - post_filter=None, - immiscible=None, - **kwargs, - ): - self.ss = None - self.options = kwargs - self.cfgpp = self.allow_cfgpp and self.options.pop("cfgpp", False) is True - alt_cfgpp_scale = self.options.pop("alt_cfgpp_scale", 0.0) - self.alt_cfgpp_scale = 0.0 if not self.allow_alt_cfgpp else alt_cfgpp_scale - self.s_noise = s_noise - self.eta = eta - self.eta_retry_increment = eta_retry_increment - self.dyn_eta_start = dyn_eta_start - self.dyn_eta_end = dyn_eta_end - self.noise_sampler = noise_sampler - self.immiscible = ( - noise.ImmiscibleNoise(**immiscible) - if immiscible not in (False, None) - else immiscible - ) - self.weight = weight - self.substeps = substeps - self.pre_filter = ( - None if pre_filter is None else filtering.make_filter(pre_filter) - ) - self.post_filter = ( - None if post_filter is None else filtering.make_filter(post_filter) - ) - self.custom_noise = self.options.get("custom_noise") - if isinstance(self.custom_noise, str): - self.custom_noise = self.options.get(f"custom_noise_{self.custom_noise}") - - def __call__(self, x): - ss = self.ss - orig_x = x - if not self.sample_sigma_zero and ss.sigma_next == 0: - return (yield from self.denoised_result()) - if self.pre_filter or self.post_filter: - filter_refs = ss.refs | filtering.FilterRefs({"orig_x": orig_x}) - if self.pre_filter: - x = self.pre_filter.apply(x, refs=filter_refs) - next_x = None - sg = self.step(x) - with contextlib.suppress(StopIteration): - while True: - sr = sg.send(next_x) - if sr.final: - if self.ancestralize: - sr = self.ancestralize_result(sr) - curr_x = sr.x - if self.post_filter: - curr_x = self.post_filter.apply(curr_x, refs=filter_refs) - sr.x_ = curr_x - return (yield sr) - next_x = sr.noise_x(ss=ss) - - def step(self, x): - raise NotImplementedError - - def prepare(self, ss): - self.ss = ss - self.noise_sampler = ss.noise.make_caching_noise_sampler( - self.custom_noise, - self.max_noise_samples, - ss.sigma, - ss.sigma_next, - immiscible=fallback(self.immiscible, ss.noise.immiscible), - ) - - def reset(self): - self.ss = None - self.noise_sampler = None - - # Euler - based on original ComfyUI implementation - def euler_step(self, x): - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - d = self.to_d(ss.hcur) - dt = sigma_down - ss.sigma - x = ss.denoised + d * sigma_down if self.cfgpp else x + d * dt - return (yield from self.result(x, sigma_up, sigma_down=sigma_down)) - - def denoised_result(self, **kwargs): - ss = self.ss - return ( - yield SamplerResult(ss, self, ss.denoised, ss.sigma.new_zeros(1), **kwargs) - ) - - def result(self, x, noise_scale=None, **kwargs): - return (yield SamplerResult(self.ss, self, x, noise_scale, **kwargs)) - - def split_result( - self, denoised, noise_pred, sigma_up=None, sigma_down=None, **kwargs - ): - return ( - yield SamplerResult( - self.ss, - self, - None, - sigma_up, - sigma_down=sigma_down, - split_result=(denoised, noise_pred), - **kwargs, - ) - ) - - def get_ancestral_step(self, *args, **kwargs): - return self.ss.get_ancestral_step( - *args, retry_increment=self.eta_retry_increment, **kwargs - ) - - def ancestralize_result(self, sr): - ss = self.ss - new_sr = sr.clone() - if new_sr.sigma_down is not None and new_sr.sigma_down != new_sr.sigma_next: - return sr - eta = self.get_dyn_eta() - if sr.sigma_next == 0 or eta == 0: - return sr - sd, su = self.get_ancestral_step(eta, sigma=sr.sigma, sigma_next=sr.sigma_next) - _ = new_sr.extract_pred(ss) - new_sr.x_ = None - new_sr.sigma_up = su - new_sr.sigma_down = sd - return new_sr - - def __str__(self): - return f"" - - def get_dyn_value(self, start, end): - if None in (start, end): - return 1.0 - if start == end: - return start - ss = self.ss - main_idx = getattr(ss, "main_idx", ss.idx) - main_sigmas = getattr(ss, "main_sigmas", ss.sigmas) - step_pct = main_idx / (len(main_sigmas) - 1) - dd_diff = end - start - return start + dd_diff * step_pct - - def get_dyn_eta(self): - return self.eta * self.get_dyn_value(self.dyn_eta_start, self.dyn_eta_end) - - @property - def max_noise_samples(self): - return (1 + self.self_noise) * self.substeps - - @property - def require_uncond(self): - return self.cfgpp or self.alt_cfgpp_scale != 0 - - def to_d(self, mr, **kwargs): - return mr.to_d(alt_cfgpp_scale=self.alt_cfgpp_scale, cfgpp=self.cfgpp, **kwargs) - - def call_model(self, *args, **kwargs): - ss = self.ss - kwargs["require_uncond"] = self.require_uncond or kwargs.get( - "require_uncond", False - ) - kwargs["cfg_scale_override"] = kwargs.get( - "cfg_scale_override", ss.cfg_scale_override - ) - return ss.call_model(*args, ss=ss, **kwargs) - - -class HistorySingleStepSampler(SingleStepSampler): - default_history_limit, max_history = 0, 0 - - def __init__(self, *args, history_limit=None, **kwargs): - super().__init__(*args, **kwargs) - self.history_limit = min( - self.max_history, - max( - 0, - self.default_history_limit if history_limit is None else history_limit, - ), - ) - - def available_history(self): - ss = self.ss - return max( - 0, min(ss.idx, self.history_limit, self.max_history, len(ss.hist) - 1) - ) - - -class ReversibleSingleStepSampler(HistorySingleStepSampler): - def __init__( - self, - *, - reversible_scale=1.0, - reta=1.0, - dyn_reta_start=None, - dyn_reta_end=None, - reversible_start_step=0, - **kwargs, - ): - super().__init__(**kwargs) - self.reversible_scale = reversible_scale - self.reta = reta - self.reversible_start_step = reversible_start_step - self.dyn_reta_start = dyn_reta_start - self.dyn_reta_end = dyn_reta_end - - def reversible_correction(self): - raise NotImplementedError - - def get_dyn_reta(self): - ss = self.ss - if ss.step < self.reversible_start_step: - return 0.0 - return self.reta * self.get_dyn_value(self.dyn_reta_start, self.dyn_reta_end) - - def get_reversible_cfg(self): - ss = self.ss - if ss.step < self.reversible_start_step: - return 0.0, 0.0 - return self.get_dyn_reta(), self.reversible_scale - - -class DPMPPStepMixin: - @staticmethod - def sigma_fn(t): - return t.neg().exp() - - @staticmethod - def t_fn(t): - return t.log().neg() - - -class MinSigmaStepMixin: - @staticmethod - def adjust_step(sigma, min_sigma, threshold=5e-04): - if min_sigma - sigma > threshold: - return sigma.clamp(min=min_sigma) - return sigma - - def adjusted_step(self, sn, result, mcc, sigma_up): - ss = self.ss - if sn == ss.sigma_next: - return sigma_up, result - # FIXME: Make sure we're noising from the right sigma. - result = yield from self.result( - result, sigma_up, sigma=ss.sigma, sigma_next=sn, final=False - ) - mr = self.call_model(result, sn, call_index=mcc) - dt = ss.sigma_next - sn - result = result + self.to_d(mr) * dt - return sigma_up.new_zeros(1), result - - -class EulerStep(SingleStepSampler): - name = "euler" - allow_cfgpp = True - allow_alt_cfgpp = True - step = SingleStepSampler.euler_step - - -class CycleSingleStepSampler(SingleStepSampler): - def __init__(self, *, cycle_pct=0.25, **kwargs): - super().__init__(**kwargs) - if cycle_pct < 0: - raise ValueError("cycle_pct must be positive") - self.cycle_pct = cycle_pct - - def get_cycle_scales(self, sigma_next): - keep_scale = sigma_next * (1.0 - self.cycle_pct) if self.cycle_pct < 1 else 0.0 - add_scale = ((sigma_next**2.0 - keep_scale**2.0) ** 0.5) * ( - 0.95 + 0.25 * self.cycle_pct - ) - # print(f">> keep={keep_scale}, add={add_scale}") - return keep_scale, add_scale - - -class EulerCycleStep(CycleSingleStepSampler): - name = "euler_cycle" - allow_alt_cfgpp = True - allow_cfgpp = True - - def step(self, x): - ss = self.ss - if ss.sigma_next == 0: - return (yield from self.denoised_result()) - keep_scale, add_scale = self.get_cycle_scales(ss.sigma_next) - keep_noise = self.to_d(ss.hcur) * keep_scale if keep_scale > 0 else 0.0 - yield from self.result(ss.denoised + keep_noise, add_scale) - - -class DPMPP2MStep(HistorySingleStepSampler, DPMPPStepMixin): - name = "dpmpp_2m" - default_history_limit, max_history = 1, 1 - ancestralize = True - - def step(self, x): - ss = self.ss - s, sn = ss.sigma, ss.sigma_next - t, t_next = self.t_fn(s), self.t_fn(sn) - h = t_next - t - st, st_next = self.sigma_fn(t), self.sigma_fn(t_next) - if self.available_history() > 0: - h_last = t - self.t_fn(ss.sigma_prev) - r = h_last / h - denoised, old_denoised = ss.denoised, ss.hprev.denoised - denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised - else: - denoised_d = ss.denoised - yield from self.result((st_next / st) * x - (-h).expm1() * denoised_d) - - -class DPMPP2MSDEStep(HistorySingleStepSampler): - name = "dpmpp_2m_sde" - default_history_limit, max_history = 1, 1 - - def __init__(self, *, solver_type="midpoint", **kwargs): - super().__init__(**kwargs) - self.solver_type = solver_type - - def step(self, x): - ss = self.ss - denoised = ss.denoised - # DPM-Solver++(2M) SDE - t, s = -ss.sigma.log(), -ss.sigma_next.log() - h = s - t - eta_h = self.get_dyn_eta() * h - - x = ( - ss.sigma_next / ss.sigma * (-eta_h).exp() * x - + (-h - eta_h).expm1().neg() * denoised - ) - noise_strength = ss.sigma_next * (-2 * eta_h).expm1().neg().sqrt() - if self.available_history() == 0: - return (yield from self.result(x, noise_strength)) - h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log()) - r = h_last / h - old_denoised = ss.hprev.denoised - if self.solver_type == "heun": - x = x + ( - ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) - * (1 / r) - * (denoised - old_denoised) - ) - elif self.solver_type == "midpoint": - x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * ( - denoised - old_denoised - ) - yield from self.result(x, noise_strength) - - -class DPMPP3MSDEStep(HistorySingleStepSampler): - name = "dpmpp_3m_sde" - default_history_limit, max_history = 2, 2 - - def step(self, x): - ss = self.ss - denoised = ss.denoised - t, s = -ss.sigma.log(), -ss.sigma_next.log() - h = s - t - eta = self.get_dyn_eta() - h_eta = h * (eta + 1) - - x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised - noise_strength = ss.sigma_next * (-2 * h * eta).expm1().neg().sqrt() - ah = self.available_history() - if ah == 0: - return (yield from self.result(x, noise_strength)) - hist = ss.hist - h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log()) - denoised_1 = hist[-2].denoised - if ah == 1: - r = h_1 / h - d = (denoised - denoised_1) / r - phi_2 = h_eta.neg().expm1() / h_eta + 1 - x = x + phi_2 * d - else: # 2+ history items available - h_2 = (-ss.sigma_prev.log()) - (-ss.sigmas[ss.idx - 2].log()) - denoised_2 = hist[-3].denoised - r0 = h_1 / h - r1 = h_2 / h - d1_0 = (denoised - denoised_1) / r0 - d1_1 = (denoised_1 - denoised_2) / r1 - d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) - d2 = (d1_0 - d1_1) / (r0 + r1) - phi_2 = h_eta.neg().expm1() / h_eta + 1 - phi_3 = phi_2 / h_eta - 0.5 - x = x + phi_2 * d1 - phi_3 * d2 - yield from self.result(x, noise_strength) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class ReversibleHeunStep(ReversibleSingleStepSampler): - name = "reversible_heun" - model_calls = 1 - allow_alt_cfgpp = True - allow_cfgpp = True - - def step(self, xs): - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - reta, reversible_scale = self.get_reversible_cfg() - sigma_down_reversible, _sigma_up_reversible = self.get_ancestral_step(reta) - dt_reversible = sigma_down_reversible - ss.sigma - - # Calculate the derivative using the model - d = self.to_d(ss.hcur) - - # Predict the sample at the next sigma using Euler step - x_pred = ss.denoised + d * sigma_down - - # Denoised sample at the next sigma - mr_next = self.call_model(x_pred, sigma_down, call_index=1) - - # Calculate the derivative at the next sigma - d_next = self.to_d(mr_next) - - # Update the sample using the Reversible Heun formula - correction = dt_reversible**2 * (d_next - d) / 4 - x = ( - mr_next.denoised - + (sigma_down * (d + d_next) / 2) - - correction * reversible_scale - ) - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class ReversibleHeun1SStep(ReversibleSingleStepSampler): - name = "reversible_heun_1s" - model_calls = 1 - default_history_limit, max_history = 1, 1 - allow_alt_cfgpp = True - allow_cfgpp = True - - def step(self, x): - if self.available_history() < 1: - return (yield from ReversibleHeunStep.step(self, x)) - ss = self.ss - s = ss.sigma - # Reversible Heun-inspired update (first-order) - sd, su = self.get_ancestral_step(self.get_dyn_eta()) - reta, reversible_scale = self.get_reversible_cfg() - sdr, _sur = self.get_ancestral_step(reta) - dt, dtr = sd - s, sdr - s - # eff_x = ss.hist[-1].x if ah > 0 else x - eff_x = x - - # Calculate the derivative using the model - # d_prev = self.to_d( - # ss.hist[-2] if ah > 0 else ss.hist[-1], - # x=eff_x, - # sigma=s, - # ) - prev_mr = ss.hist[-2] - - # d_prev = self.to_d(prev_mr, x=eff_x, sigma=s) - d_prev = self.to_d(prev_mr, sigma=ss.sigma_prev) - - # Predict the sample at the next sigma using Euler step - # x_pred = ss.denoised + d_prev * sd - x_pred = eff_x + d_prev * dt - # x_pred = ss.denoised + d_prev * sd - - # Calculate the derivative at the next sigma - d_next = self.to_d(ss.hcur, x=x_pred, sigma=sd) - - # Update the sample using the Reversible Heun formula - correction = dtr**2 * (d_next - d_prev) / 4 - x = x + (dt * (d_prev + d_next) / 2) - correction * reversible_scale - yield from self.result(x, su, sigma_down=sd) - - # def __step(self, x, ss): - # if ss.sigma_next == 0: - # return self.euler_step(x, ss) - # # Reversible Heun-inspired update (first-order) - # sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) - # sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step( - # self.get_dyn_reta(ss) - # ) - # sigma_i, sigma_i_plus_1 = ss.sigma, sigma_down - # dt = sigma_i_plus_1 - sigma_i - # dt_reversible = sigma_down_reversible - sigma_i - - # eff_x = ss.hist[-2 if len(ss.hist) > 1 else -1].x - # # eff_x = ss.hist[-2].x if len(ss.hist) > 1 else x - - # # Calculate the derivative using the model - # eff_mr = ss.hprev if len(ss.hist) > 1 else ss.hcur - # d_i_old = self.to_d(eff_mr) - # # d_i_old = self.to_d(ss.hprev if len(ss.hist) > 1 else ss.hcur) - # # d_i_old = to_d( - # # eff_x, - # # sigma_i if len(ss.hist) == 1 else ss.sigma_prev, - # # ss.hist[-2].denoised - # # if len(ss.hist) > 1 - # # else ss.model(eff_x, sigma_i, ss=ss,call_index=1).denoised, - # # ) - - # # Predict the sample at the next sigma using Euler step - # x_pred = eff_x + d_i_old * dt - - # # Calculate the derivative at the next sigma - # d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, ss.denoised) - - # # Update the sample using the Reversible Heun formula - # x = ( - # x - # + dt * (d_i_old + d_i_plus_1) / 2 - # - dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4 - # ) - # yield from self.result(ss, x, sigma_up) - # # return x, sigma_up - - # def _step(self, x, ss): - # if ss.sigma_next == 0: - # return (yield from self.euler_step(x, ss)) - # ah = self.available_history(ss) - # s = ss.sigma - # # Reversible Heun-inspired update (first-order) - # sd, su = ss.get_ancestral_step(self.get_dyn_eta(ss)) - # sdr, _sur = ss.get_ancestral_step(self.get_dyn_reta(ss)) - # dt, dtr = sd - s, sdr - s - # # eff_mr = ss.hprev if ah > 0 else ss.hcur - # # eff_x = ss.hist[-1].x if ah > 0 else x # This probably doesn't make sense. - - # # Calculate the derivative using the model - # # mr_prev = ss.hist[-2] if ah > 0 else ss.model(eff_x, s, ss=ss,call_index=1) - # mr_prev = ss.hist[-2 if ah > 0 else -1] - # d_prev = self.to_d(mr_prev, x=x, sigma=ss.sigma) - # # d_prev = self.to_d( - # # mr_prev, sigma=ss.sigma_prev if ss.sigma_prev is not None else ss.sigma - # # ) - - # # Predict the sample at the next sigma using Euler step - # x_pred = ss.denoised + d_prev * sd - # # x_pred = mr_prev.denoised + d_prev * sd - # # x_pred = eff_x + d_prev * dt - - # # Calculate the derivative at the next sigma - # d_next = self.to_d(ss.hcur, x=x_pred, sigma=sd) - - # # Update the sample using the Reversible Heun formula - # correction = dtr**2 * (d_next - d_prev) / 4 - # # x = x + (dt * (d_prev + d_next) / 2) - correction * self.reversible_scale - # x = ( - # ss.denoised - # + (sd * (d_prev + d_next) / 2) - # - correction * self.reversible_scale - # ) - # yield from self.result(ss, x, su) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class RESStep(SingleStepSampler): - name = "res" - model_calls = 1 - allow_alt_cfgpp = True # May not be implemented correctly. - - def __init__(self, *, res_simple_phi=False, res_c2=0.5, **kwargs): - super().__init__(**kwargs) - self.simple_phi = res_simple_phi - self.c2 = res_c2 - - def step(self, x): - ss = self.ss - eta = self.get_dyn_eta() - sigma_down, sigma_up = self.get_ancestral_step(eta) - denoised = ss.denoised - lam_next = sigma_down.log().neg() if eta != 0 else ss.sigma_next.log().neg() - lam = ss.sigma.log().neg() - - h = lam_next - lam - a2_1, b1, b2 = res_support._de_second_order( - h=h, c2=self.c2, simple_phi_calc=self.simple_phi - ) - - c2_h = 0.5 * h - - eff_x = ( - x - if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None - else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale - ) - x_2 = math.exp(-c2_h) * eff_x + a2_1 * h * denoised - lam_2 = lam + c2_h - sigma_2 = lam_2.neg().exp() - - denoised2 = self.call_model(x_2, sigma_2, call_index=1).denoised - - x = math.exp(-h) * eff_x + h * (b1 * denoised + b2 * denoised2) - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class TrapezoidalStep(SingleStepSampler): - name = "trapezoidal" - model_calls = 1 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - - # Calculate the derivative using the model - d_i = self.to_d(ss.hcur) - - # Predict the sample at the next sigma using Euler step - x_pred = x + d_i * ss.dt - - # Denoised sample at the next sigma - mr_next = self.call_model(x_pred, ss.sigma_next, call_index=1) - - # Calculate the derivative at the next sigma - d_next = self.to_d(mr_next) - dt_2 = sigma_down - ss.sigma - - # Update the sample using the Trapezoidal rule - x = x + dt_2 * (d_i + d_next) / 2 - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -class TrapezoidalCycleStep(CycleSingleStepSampler): - name = "trapezoidal_cycle" - model_calls = 1 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - # Calculate the derivative using the model - d_i = self.to_d(ss.hcur) - - # Predict the sample at the next sigma using Euler step - x_pred = x + d_i * ss.dt - - # Denoised sample at the next sigma - mr_next = self.call_model(x_pred, ss.sigma_next, call_index=1) - - # Calculate the derivative at the next sigma - d_next = self.to_d(mr_next) - - # Update the sample using the Trapezoidal rule - keep_scale, add_scale = self.get_cycle_scales(ss.sigma_next) - noise_pred = (d_i + d_next) * 0.5 # Combined noise prediction - denoised_pred = x - noise_pred * ss.sigma # Denoised prediction - yield from self.result(denoised_pred + noise_pred * keep_scale, add_scale) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class BogackiStep(ReversibleSingleStepSampler): - name = "bogacki" - reversible = False - model_calls = 2 - allow_alt_cfgpp = True - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - if not self.reversible: - self.reversible_scale = 0 - - def step(self, x): - ss = self.ss - s = ss.sigma - sd, su = self.get_ancestral_step(self.get_dyn_eta()) - reta, reversible_scale = self.get_reversible_cfg() - sdr, _sur = self.get_ancestral_step(reta) - dt, dtr = sd - s, sdr - s - - # Calculate the derivative using the model - d = self.to_d(ss.hcur) - - # Bogacki-Shampine steps - k1 = d * dt - k2 = self.to_d(self.call_model(x + k1 / 2, s + dt / 2, call_index=1)) * dt - k3 = ( - self.to_d( - self.call_model(x + 3 * k1 / 4 + k2 / 4, s + 3 * dt / 4, call_index=2) - ) - * dt - ) - - # Reversible correction term (inspired by Reversible Heun) - correction = dtr**2 * (k3 - k2) / 6 - - # Update the sample - x = (x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9) - correction * reversible_scale - yield from self.result(x, su, sigma_down=sd) - - -class ReversibleBogackiStep(BogackiStep): - name = "reversible_bogacki" - reversible = True - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class RK4Step(SingleStepSampler): - name = "rk4" - model_calls = 3 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - sigma = ss.sigma - d = self.to_d(ss.hcur) - dt = sigma_down - sigma - - # Runge-Kutta steps - k1 = d * dt - k2 = self.to_d(self.call_model(x + k1 / 2, sigma + dt / 2, call_index=1)) * dt - k3 = self.to_d(self.call_model(x + k2 / 2, sigma + dt / 2, call_index=2)) * dt - k4 = self.to_d(self.call_model(x + k3, sigma + dt, call_index=3)) * dt - - # Update the sample - x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6 - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class RKF45Step(SingleStepSampler): - name = "rkf45" - model_calls = 5 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - sigma = ss.sigma - d = self.to_d(ss.hcur) - dt = sigma_down - sigma - - # Runge-Kutta steps - sigma_progression = ( - sigma + dt / 4, - sigma + 3 * dt / 8, - sigma + 12 * dt / 13, - sigma + dt, - ) - - call_progression = ( - lambda k1: x + k1 / 4, - lambda k1, k2: x + 3 * k1 / 32 + 9 * k2 / 32, - lambda k1, k2, k3: x - + 1932 * k1 / 2197 - - 7200 * k2 / 2197 - + 7296 * k3 / 2197, - lambda k1, k2, k3, k4: x - + 439 * k1 / 216 - - 8 * k2 - + 3680 * k3 / 513 - - 845 * k4 / 4104, - ) - - k = [d * dt] - for idx, (ksigma, kfun) in enumerate(zip(sigma_progression, call_progression)): - curr_x = kfun(*k) - k.append(self.to_d(self.call_model(curr_x, ksigma)) * dt) - del curr_x - x = x + 25 * k[0] / 216 + 1408 * k[2] / 2565 + 2197 * k[3] / 4104 - k[4] / 5 - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class RKDynamicStep(SingleStepSampler): - name = "rk_dynamic" - model_calls = 3 - allow_alt_cfgpp = True - - rk_weights = ( - (1,), - (0.5, 0.5), - (1 / 6, 2 / 3, 1 / 6), - (1 / 8, 3 / 8, 3 / 8, 1 / 8), - ) - - rk_error_orders = ((0.0375, 4), (0.075, 3), (0.15, 2)) - - def __init__(self, *args, max_order=4, **kwargs): - super().__init__(*args, **kwargs) - self.max_order = max(0, min(max_order, 4)) - - def get_rk_error_order(self, error): - for threshold, order in self.rk_error_orders: - if error < threshold: - return order - return 1 - - def step(self, x): - ss = self.ss - order = self.max_order - - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - sigma = ss.sigma - d = self.to_d(ss.hcur) - dt = sigma_down - sigma - - error = ss.hcur.get_error(ss.hprev) if len(ss.hist) > 1 else 0.0 - if order < 1: - order = self.get_rk_error_order(error) - - k = [d * dt] - curr_weight = self.rk_weights[order - 1] - - # print( - # f"\nRK: weight={curr_weight!r}, histlen={len(ss.hist)}, order={order} ({self.max_order}), err={error:.6}\n" - # ) - for j in range(1, order): - # Calculate intermediate k values based on the current order - k_sum = sum(curr_weight[i] * k[i] for i in range(j)) - mr = self.call_model(x + k_sum, sigma + dt * sum(curr_weight[:j])) - k.append(self.to_d(mr) * dt) - del mr - - # Update the sample using the weighted sum of k values - x = x + sum(curr_weight[j] * k[j] for j in range(order)) - - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -class EulerDancingStep(SingleStepSampler): - name = "euler_dancing" - self_noise = 1 - - def __init__( - self, - *, - deta=1.0, - ds_noise=None, - leap=2, - dyn_deta_start=None, - dyn_deta_end=None, - dyn_deta_mode="lerp", - **kwargs, - ): - super().__init__(**kwargs) - self.deta = deta - self.ds_noise = ds_noise if ds_noise is not None else self.s_noise - self.leap = leap - self.dyn_deta_start = dyn_deta_start - self.dyn_deta_end = dyn_deta_end - if dyn_deta_mode not in ("lerp", "lerp_alt", "deta"): - raise ValueError("Bad dyn_deta_mode") - self.dyn_deta_mode = dyn_deta_mode - - def step(self, x): - ss = self.ss - eta = self.eta - deta = self.deta - leap_sigmas = ss.sigmas[ss.idx :] - leap_sigmas = leap_sigmas[: utils.find_first_unsorted(leap_sigmas)] - zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] - max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 - is_danceable = max_leap > 1 and ss.sigma_next != 0 - curr_leap = max(1, min(self.leap, max_leap)) - sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next - del leap_sigmas - sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) - print("???", sigma_down, sigma_up) - d = to_d(x, ss.sigma, ss.denoised) - # Euler method - dt = sigma_down - ss.sigma - x = x + d * dt - if curr_leap == 1: - return (yield from self.result(x, sigma_up)) - noise_strength = self.ds_noise * sigma_up - if noise_strength != 0: - x = yield from self.result(x, sigma_up, sigma_next=sigma_leap, final=False) - - # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_( - # self.ds_noise * sigma_up - # ) - # sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) - # _sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) - # sigma_up2 = ss.sigma_next + (ss.sigma - ss.sigma_next) * 0.5 - sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta)[1] + ( - ss.sigma_next * 0.5 - ) - sigma_down2, _sigma_up2 = get_ancestral_step( - ss.sigma_next, sigma_leap, eta=deta - ) - print(">>>", sigma_down2, sigma_up2, "--", ss.sigma, "->", sigma_leap) - # sigma_down2, sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta) - d_2 = to_d(x, sigma_leap, ss.denoised) - dt_2 = sigma_down2 - sigma_leap - x = x + d_2 * dt_2 - yield from self.result(x, sigma_up2, sigma_down=sigma_down2) - - # def _step(self, x, ss): - # eta = self.get_dyn_eta(ss) - # leap_sigmas = ss.sigmas[ss.idx :] - # leap_sigmas = leap_sigmas[: utils.find_first_unsorted(leap_sigmas)] - # zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] - # max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 - # is_danceable = max_leap > 1 and ss.sigma_next != 0 - # curr_leap = max(1, min(self.leap, max_leap)) - # sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next - # # DANCE 35 6 tensor(10.0947, device='cuda:0') -- tensor([21.9220, - # # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) - # del leap_sigmas - # sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) - # d = to_d(x, ss.sigma, ss.denoised) - # # Euler method - # dt = sigma_down - ss.sigma - # x = x + d * dt - # if curr_leap == 1: - # return x, sigma_up - # dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) - # if curr_leap == 1 or not is_danceable or abs(dance_scale) < 1e-04: - # print("NODANCE", dance_scale, self.deta, is_danceable, ss.sigma_next) - # yield SamplerResult(ss, self, x, sigma_up) - # print( - # "DANCE", dance_scale, self.deta, self.dyn_deta_mode, self.ds_noise, sigma_up - # ) - # sigma_down_normal, sigma_up_normal = get_ancestral_step( - # ss.sigma, ss.sigma_next, eta - # ) - # if self.dyn_deta_mode == "lerp": - # dt_normal = sigma_down_normal - ss.sigma - # x_normal = x + d * dt_normal - # else: - # x_normal = x - # sigma_down2, sigma_up2 = get_ancestral_step( - # sigma_leap, - # ss.sigma_next, - # eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), - # ) - # print( - # "-->", - # sigma_down2, - # sigma_up2, - # "--", - # self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), - # ) - # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.ds_noise * sigma_up) - # d_2 = to_d(x, sigma_leap, ss.denoised) - # dt_2 = sigma_down2 - sigma_leap - # result = x + d_2 * dt_2 - # # SIGMA: norm_up=9.062416076660156, up=10.703859329223633, up2=19.376544952392578, str=21.955078125 - # noise_strength = sigma_up2 + ((sigma_up - sigma_up_normal) ** 5.0) - # noise_strength = sigma_up2 + ((sigma_up2 - sigma_up) * 0.5) - # # noise_strength = sigma_up2 + ( - # # (sigma_up2 - sigma_up) ** (1.0 - (sigma_up_normal / sigma_up2)) - # # ) - # noise_diff = ( - # sigma_up - sigma_up_normal - # if sigma_up > sigma_up_normal - # else sigma_up_normal - sigma_up - # ) - # noise_div = ( - # sigma_up / sigma_up_normal - # if sigma_up > sigma_up_normal - # else sigma_up_normal / sigma_up - # ) - # noise_diff = sigma_up2 - sigma_up_normal - # noise_div = sigma_up2 / sigma_up_normal - # noise_div = ss.sigma / sigma_leap - - # # noise_strength = sigma_up2 + (noise_diff * noise_div) - # # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** 2.0) - # # noise_strength = sigma_up2 + ((1.0 - noise_diff) ** 0.5) - # # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up) * 0.5) ** 2.0) - # # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up_normal) * 0.5) ** 1.5) - # # noise_strength = sigma_up2 + ( - # # (noise_diff * 0.1875) ** (1.0 / (noise_div - 0.0)) - # # ) - # # noise_strength = sigma_up2 + ( - # # (noise_diff * 0.125) ** (1.0 / (noise_div * 1.25)) - # # ) - # # noise_strength = sigma_up2 + ((noise_diff * 0.2) ** (1.0 / (noise_div * 1.0))) - # noise_strength = sigma_up2 + (noise_diff * 0.9 * max(0.0, noise_div - 0.8)) - # noise_strength = sigma_up2 + ( - # (noise_diff / (curr_leap * 0.4)) - # * ((noise_div - (curr_leap / 2.0)).clamp(min=0, max=1.5) * 1.0) - # ) - # # (1.0 / (noise_div * 1.25))) - # # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** noise_div) - # print( - # f"SIGMA: norm_up={sigma_up_normal}, up={sigma_up}, up2={sigma_up2}, str={noise_strength}", - # # noise_diff, - # noise_div, - # ) - # return result, noise_strength - - # noise_diff = sigma_up2 - sigma_up * dance_scale - # noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) - # # noise_scale = sigma_up2 * self.ds_noise - # if self.dyn_deta_mode == "deta" or dance_scale == 1.0: - # return result, noise_scale - # result = torch.lerp(x_normal, result, dance_scale) - # # FIXME: Broken for noise samplers that care about s/sn - # return result, noise_scale - - # def step(self, x, ss): - # eta = self.get_dyn_eta(ss) - # leap_sigmas = ss.sigmas[ss.idx :] - # leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)] - # zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] - # max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 - # is_danceable = max_leap > 1 and ss.sigma_next != 0 - # curr_leap = max(1, min(self.leap, max_leap)) - # sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next - # # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) - # del leap_sigmas - # sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) - # d = to_d(x, ss.sigma, ss.denoised) - # # Euler method - # dt = sigma_down - ss.sigma - # x = x + d * dt - # if curr_leap == 1: - # return x, sigma_up - # dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) - # if not is_danceable or abs(dance_scale) < 1e-04: - # print("NODANCE", dance_scale, self.deta) - # return x, sigma_up - # print("NODANCE", dance_scale, self.deta) - # sigma_down_normal, _sigma_up_normal = get_ancestral_step( - # ss.sigma, ss.sigma_next, eta - # ) - # if self.dyn_deta_mode == "lerp": - # dt_normal = sigma_down_normal - ss.sigma - # x_normal = x + d * dt_normal - # else: - # x_normal = x - # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.s_noise * sigma_up) - # sigma_down2, sigma_up2 = get_ancestral_step( - # sigma_leap, - # ss.sigma_next, - # eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), - # ) - # d_2 = to_d(x, sigma_leap, ss.denoised) - # dt_2 = sigma_down2 - sigma_leap - # result = x + d_2 * dt_2 - # noise_diff = sigma_up2 - sigma_up * dance_scale - # noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) - # if self.dyn_deta_mode == "deta" or dance_scale == 1.0: - # return result, noise_scale - # result = torch.lerp(x_normal, result, dance_scale) - # # FIXME: Broken for noise samplers that care about s/sn - # return result, noise_scale - - -# Alt CFG++ approach referenced from https://github.com/comfyanonymous/ComfyUI/pull/3871 - thanks! -class DPMPP2SStep(SingleStepSampler, DPMPPStepMixin): - name = "dpmpp_2s" - model_calls = 1 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - t_fn, sigma_fn = self.t_fn, self.sigma_fn - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - # DPM-Solver++(2S) - t, t_next = t_fn(ss.sigma), t_fn(sigma_down) - r = 1 / 2 - h = t_next - t - s = t + r * h - eff_x = ( - x - if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None - else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale - ) - x_2 = (sigma_fn(s) / sigma_fn(t)) * eff_x - (-h * r).expm1() * ss.denoised - denoised_2 = self.call_model(x_2, sigma_fn(s), call_index=1).denoised - x = (sigma_fn(t_next) / sigma_fn(t)) * eff_x - (-h).expm1() * denoised_2 - yield from self.result(x, sigma_up, sigma_down=sigma_down) - - -class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin): - name = "dpmpp_sde" - self_noise = 1 - model_calls = 1 - allow_alt_cfgpp = True # Implementation may not be correct. - - def __init__(self, *args, r=1 / 2, **kwargs): - super().__init__(*args, **kwargs) - self.r = r - - def step(self, x): - ss = self.ss - t_fn, sigma_fn = self.t_fn, self.sigma_fn - r, eta = self.r, self.get_dyn_eta() - # DPM-Solver++ - t, t_next = t_fn(ss.sigma), t_fn(ss.sigma_next) - h = t_next - t - s = t + h * r - fac = 1 / (2 * r) - - # Step 1 - sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta) - s_ = t_fn(sd) - eff_x = ( - x - if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None - else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale - ) - x_2 = (sigma_fn(s_) / sigma_fn(t)) * eff_x - (t - s_).expm1() * ss.denoised - x_2 = yield from self.result( - x_2, su, sigma=sigma_fn(t), sigma_next=sigma_fn(s), final=False - ) - denoised_2 = self.call_model(x_2, sigma_fn(s), call_index=1).denoised - - # Step 2 - sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta) - t_next_ = t_fn(sd) - denoised_d = (1 - fac) * ss.denoised + fac * denoised_2 - x = (sigma_fn(t_next_) / sigma_fn(t)) * eff_x - ( - t - t_next_ - ).expm1() * denoised_d - yield from self.result(x, su, sigma_down=sd) - - -# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers -# Which was originally written by Katherine Crowson -class TTMJVPStep(SingleStepSampler): - name = "ttm_jvp" - model_calls = 1 - - def __init__(self, *args, alternate_phi_2_calc=True, **kwargs): - super().__init__(*args, **kwargs) - self.alternate_phi_2_calc = alternate_phi_2_calc - - def step(self, x): - ss = self.ss - eta = self.get_dyn_eta() - sigma, sigma_next = ss.sigma, ss.sigma_next - # 2nd order truncated Taylor method - t, s = -sigma.log(), -sigma_next.log() - h = s - t - h_eta = h * (eta + 1) - - eps = to_d(x, sigma, ss.denoised) - denoised_prime = self.call_model( - x, sigma, tangents=(eps * -sigma, -sigma), call_index=1 - ).jdenoised - - phi_1 = -torch.expm1(-h_eta) - if self.alternate_phi_2_calc: - phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0 - else: - phi_2 = torch.expm1(-h_eta) + h_eta - x = torch.exp(-h_eta) * x + phi_1 * ss.denoised + phi_2 * denoised_prime - - noise_scale = ( - sigma_next * torch.sqrt(-torch.expm1(-2 * h * eta)) - if eta - else ss.sigma.new_zeros(1) - ) - yield from self.result(x, noise_scale) - - -# Adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py -# under Apache 2 license -class IPNDMStep(HistorySingleStepSampler): - name = "ipndm" - ancestralize = True - default_history_limit, max_history = 1, 3 - allow_alt_cfgpp = True - - IPNDM_MULTIPLIERS = ( - ((1,), 1), - ((3, -1), 2), - ((23, -16, 5), 12), - ((55, -59, 37, -9), 24), - ) - - def step(self, x): - ss = self.ss - order = self.available_history() + 1 - if order > 1: - hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) - (dm, *hms), divisor = self.IPNDM_MULTIPLIERS[order - 1] - noise = dm * self.to_d(ss.hcur) - for hidx, hm in enumerate(hms, start=1): - noise += hm * hd[-hidx] - noise /= divisor - yield from self.result(x + ss.dt * noise) - - -# Adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py -# under Apache 2 license -class IPNDMVStep(HistorySingleStepSampler): - name = "ipndm_v" - ancestralize = True - default_history_limit, max_history = 1, 3 - allow_alt_cfgpp = True - - def step(self, x): - ss = self.ss - dt = ss.dt - d = self.to_d(ss.hcur) - order = self.available_history() + 1 - if order > 1: - hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) - hns = ( - ss.sigmas[ss.idx - (order - 2) : ss.idx + 1] - - ss.sigmas[ss.idx - (order - 1) : ss.idx] - ) - if order == 1: - noise = d - elif order == 2: - coeff1 = (2 + (dt / hns[-1])) / 2 - coeff2 = -(dt / hns[-1]) / 2 - noise = coeff1 * d + coeff2 * hd[-1] - elif order == 3: - temp = ( - 1 - - dt - / (3 * (dt + hns[-1])) - * (dt * (dt + hns[-1])) - / (hns[-1] * (hns[-1] + hns[-2])) - ) / 2 - coeff1 = (2 + (dt / hns[-1])) / 2 + temp - coeff2 = -(dt / hns[-1]) / 2 - (1 + hns[-1] / hns[-2]) * temp - coeff3 = temp * hns[-1] / hns[-2] - noise = coeff1 * d + coeff2 * hd[-1] + coeff3 * hd[-2] - else: - temp1 = ( - 1 - - dt - / (3 * (dt + hns[-1])) - * (dt * (dt + hns[-1])) - / (hns[-1] * (hns[-1] + hns[-2])) - ) / 2 - temp2 = ( - ( - (1 - dt / (3 * (dt + hns[-1]))) / 2 - + (1 - dt / (2 * (dt + hns[-1]))) - * dt - / (6 * (dt + hns[-1] + hns[-2])) - ) - * (dt * (dt + hns[-1]) * (dt + hns[-1] + hns[-2])) - / (hns[-1] * (hns[-1] + hns[-2]) * (hns[-1] + hns[-2] + hns[-3])) - ) - coeff1 = (2 + (dt / hns[-1])) / 2 + temp1 + temp2 - coeff2 = ( - -(dt / hns[-1]) / 2 - - (1 + hns[-1] / hns[-2]) * temp1 - - ( - 1 - + (hns[-1] / hns[-2]) - + (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) - ) - * temp2 - ) - coeff3 = ( - temp1 * hns[-1] / hns[-2] - + ( - (hns[-1] / hns[-2]) - + (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) - * (1 + hns[-2] / hns[-3]) - ) - * temp2 - ) - coeff4 = ( - -temp2 - * (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) - * hns[-1] - / hns[-2] - ) - noise = coeff1 * d + coeff2 * hd[-1] + coeff3 * hd[-2] + coeff4 * hd[-3] - yield from self.result(x + ss.dt * noise) - - -class DEISStep(HistorySingleStepSampler): - name = "deis" - ancestralize = True - default_history_limit, max_history = 1, 3 - allow_alt_cfgpp = True - - def __init__(self, *args, deis_mode="tab", **kwargs): - super().__init__(*args, **kwargs) - self.deis_mode = deis_mode - self.deis_coeffs_key = None - self.deis_coeffs = None - - def get_deis_coeffs(self): - ss = self.ss - key = ( - self.history_limit, - len(ss.sigmas), - ss.sigmas[0].item(), - ss.sigmas[-1].item(), - ) - if self.deis_coeffs_key == key: - return self.deis_coeffs - self.deis_coeffs_key = key - self.deis_coeffs = comfy.k_diffusion.deis.get_deis_coeff_list( - ss.sigmas, self.history_limit + 1, deis_mode=self.deis_mode - ) - return self.deis_coeffs - - def step(self, x): - ss = self.ss - dt = ss.dt - d = self.to_d(ss.hcur) - order = self.available_history() + 1 - if order < 2: - noise = dt * d # Euler - else: - c = self.get_deis_coeffs()[ss.idx] - hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) - noise = c[0] * d - for i in range(1, order): - noise += c[i] * hd[-i] - yield from self.result(x + noise) - - -class HeunPP2Step(SingleStepSampler): - name = "heunpp2" - ancestralize = True - model_calls = 2 - allow_alt_cfgpp = True - - def __init__(self, *args, max_order=3, **kwargs): - super().__init__(*args, **kwargs) - self.max_order = max(1, min(self.model_calls + 1, max_order)) - - def step(self, x): - ss = self.ss - steps_remain = max(0, len(ss.sigmas) - (ss.idx + 2)) - order = min(self.max_order, steps_remain + 1) - sn = ss.sigma_next - if order == 1: - return (yield from self.euler_step(x)) - d = self.to_d(ss.hcur) - dt = ss.dt - w = order * ss.sigma - w2 = sn / w - x_2 = x + d * dt - d_2 = self.to_d(self.call_model(x_2, sn, call_index=1)) - if order == 2: - # Heun's method (ish) - w1 = 1 - w2 - d_prime = d * w1 + d_2 * w2 - else: - # Heun++ (ish) - snn = ss.sigmas[ss.idx + 2] - dt_2 = snn - sn - x_3 = x_2 + d_2 * dt_2 - d_3 = self.to_d(self.call_model(x_3, snn, call_index=2)) - w3 = snn / w - w1 = 1 - w2 - w3 - d_prime = w1 * d + w2 * d_2 + w3 * d_3 - yield from self.result(x + d_prime * dt) - - -class DESolverStep(SingleStepSampler, MinSigmaStepMixin): - de_default_solver = None - sample_sigma_zero = True - - def __init__( - self, - *args, - de_solver=None, - de_max_nfe=100, - de_rtol=-2.5, - de_atol=-3.5, - de_fixup_hack=0.025, - de_split=1, - de_min_sigma=0.0292, - **kwargs, - ): - self.check_solver_support() - if not HAVE_TDE: - raise RuntimeError( - "TDE sampler requires torchdiffeq installed in venv. Example: pip install torchdiffeq" - ) - super().__init__(*args, **kwargs) - de_solver = self.de_default_solver if de_solver is None else de_solver - self.de_solver_name = de_solver - self.de_max_nfe = de_max_nfe - self.de_rtol = 10**de_rtol - self.de_atol = 10**de_atol - self.de_fixup_hack = de_fixup_hack - self.de_split = de_split - self.de_min_sigma = de_min_sigma if de_min_sigma is not None else 0.0 - - def check_solver_support(self): - raise NotImplementedError - - def de_get_step(self, x): - eta = self.get_dyn_eta() - ss = self.ss - s, sn = ss.sigma, ss.sigma_next - sn = self.adjust_step(sn, self.de_min_sigma) - sigma_down, sigma_up = self.get_ancestral_step(eta, sigma_next=sn) - if self.de_fixup_hack != 0: - sigma_down = (sigma_down - (s - sigma_down) * self.de_fixup_hack).clamp( - min=0 - ) - return s, sn, sigma_down, sigma_up - - @staticmethod - def reverse_time(t, t0, t1): - return t1 + (t0 - t) - - -class TDEStep(DESolverStep): - name = "tde" - model_calls = 2 - allow_alt_cfgpp = True - allow_cfgpp = False - de_default_solver = "rk4" - - def __init__( - self, - *args, - de_split=1, - **kwargs, - ): - super().__init__(*args, **kwargs) - self.de_split = de_split - - def check_solver_support(self): - if not HAVE_TDE: - raise RuntimeError( - "TDE sampler requires torchdiffeq installed in venv. Example: pip install torchdiffeq" - ) - - def step(self, x): - s, sn, sigma_down, sigma_up = self.de_get_step(x) - if self.de_min_sigma is not None and s <= self.de_min_sigma: - return (yield from self.euler_step(x)) - ss = self.ss - delta = (s - sigma_down).item() - mcc = 0 - bidx = 0 - pbar = None - - def odefn(t, y): - nonlocal mcc - if t < 1e-05: - return torch.zeros_like(y) - if mcc >= self.de_max_nfe: - raise RuntimeError("TDEStep: Model call limit exceeded") - - pct = (s - t) / delta - pbar.n = round(min(999, pct.item() * 999)) - pbar.update(0) - pbar.set_description( - f"{self.de_solver_name}({mcc}/{self.de_max_nfe})", refresh=True - ) - - if t == ss.sigma and torch.equal(x[bidx], y): - mr_cached = True - mr = ss.hcur - mcc = 1 - else: - mr_cached = False - mr = self.call_model( - y.unsqueeze(0), t, call_index=mcc, s_in=t.new_ones(1) - ) - mcc += 1 - return self.to_d(mr)[bidx if mr_cached else 0] - - result = torch.zeros_like(x) - t = sigma_down.new_zeros(self.de_split + 1) - torch.linspace(ss.sigma, sigma_down, t.shape[0], out=t) - - for batch in tqdm.trange( - 1, - x.shape[0] + 1, - desc="batch", - leave=False, - disable=x.shape[0] == 1 or ss.disable_status, - ): - bidx = batch - 1 - mcc = 0 - if pbar is not None: - pbar.close() - pbar = tqdm.tqdm( - total=1000, - desc=self.de_solver_name, - leave=True, - disable=ss.disable_status, - ) - solution = tde.odeint( - odefn, - x[bidx], - t, - rtol=self.de_rtol, - atol=self.de_atol, - method=self.de_solver_name, - options={ - "min_step": 1e-05, - "dtype": torch.float64, - }, - )[-1] - result[bidx] = solution - - sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) - if pbar is not None: - pbar.n = pbar.total - pbar.update(0) - pbar.close() - yield from self.result(result, sigma_up, sigma_down=sigma_down) - - -class TODEStep(DESolverStep): - name = "tode" - model_calls = 2 - allow_alt_cfgpp = True - de_default_solver = "dopri5" - - def __init__( - self, - *args, - de_initial_step=0.25, - tode_compile=False, - de_ctl_pcoeff=0.3, - de_ctl_icoeff=0.9, - de_ctl_dcoeff=0.2, - **kwargs, - ): - if not HAVE_TODE: - raise RuntimeError( - "TODE sampler requires torchode installed in venv. Example: pip install torchode" - ) - super().__init__(*args, **kwargs) - self.de_solver_method = tode.interface.METHODS[self.de_solver_name] - self.de_ctl_pcoeff = de_ctl_pcoeff - self.de_ctl_icoeff = de_ctl_icoeff - self.de_ctl_dcoeff = de_ctl_dcoeff - self.de_compile = tode_compile - self.de_initial_step = de_initial_step - - def check_solver_support(self): - if not HAVE_TODE: - raise RuntimeError( - "TODE sampler requires torchode installed in venv. Example: pip install torchode" - ) - - def step(self, x): - s, sn, sigma_down, sigma_up = self.de_get_step(x) - if self.de_min_sigma is not None and s <= self.de_min_sigma: - return (yield from self.euler_step(x)) - ss = self.ss - delta = (ss.sigma - sigma_down).item() - mcc = 0 - pbar = None - b, c, h, w = x.shape - - def odefn(t, y_flat): - nonlocal mcc - if torch.all(t <= 1e-05).item(): - return torch.zeros_like(y_flat) - if mcc >= self.de_max_nfe: - raise RuntimeError("TDEStep: Model call limit exceeded") - - pct = (s - t) / delta - pbar.n = round(pct.min().item() * 999) - pbar.update(0) - pbar.set_description( - f"{self.de_solver_name}({mcc}/{self.de_max_nfe})", refresh=True - ) - y = y_flat.reshape(-1, c, h, w) - t32 = t.to(torch.float32) - del y_flat - - if mcc == 0 and torch.all(t == s): - mr = ss.hcur - mcc = 1 - else: - mr = self.call_model(y, t32.clamp(min=1e-05), call_index=mcc) - mcc += 1 - result = self.to_d(mr).flatten(start_dim=1) - for bi in range(t.shape[0]): - if t[bi] <= 1e-05: - result[bi, :] = 0 - return result - - t = torch.stack((s, sigma_down)).to(torch.float64).repeat(b, 1) - - pbar = tqdm.tqdm( - total=1000, desc=self.de_solver_name, leave=True, disable=ss.disable_status - ) - - term = tode.ODETerm(odefn) - method = self.de_solver_method(term=term) - controller = tode.PIDController( - term=term, - atol=self.de_atol, - rtol=self.de_rtol, - dt_min=1e-05, - pcoeff=self.de_ctl_pcoeff, - icoeff=self.de_ctl_icoeff, - dcoeff=self.de_ctl_dcoeff, - ) - solver_ = tode.AutoDiffAdjoint(method, controller) - solver = solver_ if not self.de_compile else torch.compile(solver_) - problem = tode.InitialValueProblem( - y0=x.flatten(start_dim=1), t_start=t[:, 0], t_end=t[:, -1] - ) - dt0 = ( - (t[:, -1] - t[:, 0]) * self.de_initial_step - if self.de_initial_step - else None - ) - solution = solver.solve(problem, dt0=dt0) - - # print("\nSOLUTION", solution.stats, solution.ys.shape) - result = solution.ys[:, -1].reshape(-1, c, h, w) - del solution - - sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) - if pbar is not None: - pbar.n = pbar.total - pbar.update(0) - pbar.close() - yield from self.result(result, sigma_up, sigma_down=sigma_down) - - -class TSDEStep(DESolverStep): - name = "tsde" - model_calls = 2 - allow_alt_cfgpp = True - de_default_solver = "reversible_heun" - - def __init__( - self, - *args, - de_initial_step=0.25, - de_split=1, - de_adaptive=False, - 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 = 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 - - def step(self, x): - s, sn, sigma_down, sigma_up = self.de_get_step(x) - if self.de_min_sigma is not None and s <= self.de_min_sigma: - return (yield from self.euler_step(x)) - ss = self.ss - delta = (ss.sigma - sigma_down).item() - bidx = 0 - mcc = 0 - pbar = None - _b, c, h, w = x.shape - outer_self = self - - class SDE(torch.nn.Module): - noise_type = outer_self.de_noise_type - sde_type = outer_self.de_sde_type - - @torch.no_grad() - def f(self, t_rev, y_flat): - nonlocal mcc - t = s - (t_rev - sigma_down) - # print(f"\nf at t_rev={t_rev}, t={t} :: {y_flat.shape}") - if torch.all(t <= 1e-05).item(): - return torch.zeros_like(y_flat) - if mcc >= outer_self.de_max_nfe: - raise RuntimeError("TSDEStep: Model call limit exceeded") - - pct = (s - t) / delta - pbar.n = round(pct.min().item() * 999) - pbar.update(0) - pbar.set_description( - f"{outer_self.de_solver_name}({mcc}/{outer_self.de_max_nfe})", - refresh=True, - ) - flat_shape = y_flat.shape - y = y_flat.view(1, c, h, w) - t32 = t.to(torch.float32) - del y_flat - - if mcc == 0 and torch.all(t == s): - mr_cached = True - mr = ss.hcur - mcc = 1 - else: - mr_cached = False - mr = outer_self.call_model( - y, t32.clamp(min=1e-05), call_index=mcc, s_in=t.new_ones(1) - ) - mcc += 1 - return -outer_self.to_d(mr)[bidx if mr_cached else 0].view(*flat_shape) - - @torch.no_grad() - def g(self, t_rev, y_flat): - t = (s - sigma_down) - (t_rev - sigma_down) - pct = t / (s - sigma_down) - if outer_self.de_g_reverse_time: - pct = 1.0 - pct - multiplier = outer_self.de_g_multiplier - if outer_self.de_g_derp_mode and mcc % 2 == 0: - multiplier *= -1 - val = t * pct * multiplier - if self.noise_type == "diagonal": - out = val.repeat(*y_flat.shape) - elif self.noise_type == "scalar": - out = val.repeat(*y_flat.shape, 1) - else: - out = val.repeat(*y_flat.shape, outer_self.de_noise_channels) - return out - - t = torch.stack((sigma_down, s)).to(torch.float) - - pbar = tqdm.tqdm( - total=1000, desc=self.de_solver_name, leave=True, disable=ss.disable_status - ) - - dt0 = ( - delta * self.de_initial_step if self.de_adaptive else delta / self.de_split - ) - results = [] - for batch in tqdm.trange( - 1, - x.shape[0] + 1, - desc="batch", - leave=False, - disable=x.shape[0] == 1 or ss.disable_status, - ): - bidx = batch - 1 - mcc = 0 - sde = SDE() - if self.de_batch_channels: - y_flat = x[bidx].flatten(start_dim=1) - else: - y_flat = x[bidx].unsqueeze(0).flatten(start_dim=1) - if sde.noise_type == "diagonal": - bm_size = (y_flat.shape[0], y_flat.shape[1]) - elif sde.noise_type == "scalar": - bm_size = (y_flat.shape[0], 1) - else: - bm_size = (y_flat.shape[0], self.de_noise_channels) - bm = torchsde.BrownianInterval( - dtype=x.dtype, - device=x.device, - t0=-s, - t1=s, - entropy=ss.noise.seed, - levy_area_approximation=self.de_levy_area_approx, - tol=1e-06, - size=bm_size, - ) - - ys = torchsde.sdeint( - sde, - y_flat, - t, - method=self.de_solver_name, - adaptive=self.de_adaptive, - atol=self.de_atol, - rtol=self.de_rtol, - dt=dt0, - bm=bm, - ) - del y_flat - results.append(ys[-1].view(1, c, h, w)) - del ys - result = torch.cat(results) - del results - - sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) - if pbar is not None: - pbar.n = pbar.total - pbar.update(0) - pbar.close() - yield from self.result(result, sigma_up, sigma_down=sigma_down) - - -if HAVE_DIFFRAX: - - class RevVirtualBrownianTree(diffrax.VirtualBrownianTree): - def evaluate(self, t0, t1, *args, **kwargs): - if t1 is not None: - return super().evaluate(t1, t0, *args, **kwargs) - return super().evaluate(t0, t1, *args, **kwargs) - - class StepCallbackTqdmProgressMeter(diffrax.TqdmProgressMeter): - step_callback: typing.Callable = None - - def _init_bar(self, *args, **kwargs): - if self.step_callback is None: - return super()._init_bar(*args, **kwargs) - bar_format = "{percentage:.2f}%{step_callback}|{bar}| [{elapsed}<{remaining}, {rate_fmt}{postfix}]" - step_callback = self.step_callback - - class WrapTqdm(tqdm.tqdm): - @property - def format_dict(self): - d = super().format_dict - d.update(step_callback=step_callback()) - return d - - return WrapTqdm(total=100, unit="%", bar_format=bar_format) - - -class DiffraxStep(DESolverStep): - name = "diffrax" - model_calls = 2 - allow_alt_cfgpp = True - de_default_solver = "dopri5" - - def __init__( - self, - *args, - de_split=1, - de_initial_step=0.25, - de_ctl_pcoeff=0.3, - de_ctl_icoeff=0.9, - de_ctl_dcoeff=0.2, - 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) - solvers = dict( - euler=diffrax.Euler, - heun=diffrax.Heun, - midpoint=diffrax.Midpoint, - ralston=diffrax.Ralston, - bosh3=diffrax.Bosh3, - tsit5=diffrax.Tsit5, - dopri5=diffrax.Dopri5, - dopri8=diffrax.Dopri8, - implicit_euler=diffrax.ImplicitEuler, - # kvaerno3=diffrax.Kvaerno3, - # kvaerno4=diffrax.Kvaerno4, - # kvaerno5=diffrax.Kvaerno5, - semi_implicit_euler=diffrax.SemiImplicitEuler, - reversible_heun=diffrax.ReversibleHeun, - leapfrog_midpoint=diffrax.LeapfrogMidpoint, - euler_heun=diffrax.EulerHeun, - ito_milstein=diffrax.ItoMilstein, - stratonovich_milstein=diffrax.StratonovichMilstein, - sea=diffrax.SEA, - sra1=diffrax.SRA1, - shark=diffrax.ShARK, - general_shark=diffrax.GeneralShARK, - slow_rk=diffrax.SlowRK, - spark=diffrax.SPaRK, - ) - levy_areas = dict( - brownian_increment=diffrax.BrownianIncrement, - space_time=diffrax.SpaceTimeLevyArea, - space_time_time=diffrax.SpaceTimeTimeLevyArea, - ) - # jax.config.update("jax_disable_jit", True) - self.de_solver_method = solvers[self.de_solver_name]() - 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 = diffrax_adaptive - self.de_split = de_split - 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 - def t2j(t): - return jax.block_until_ready( - jax.numpy.array(numpy.array(t.detach().cpu().contiguous())) - ) - - @staticmethod - def j2t(t): - return torch.from_numpy(numpy.array(jax.block_until_ready(t))).contiguous() - - def check_solver_support(self): - if not HAVE_DIFFRAX: - raise RuntimeError( - "Diffrax sampler requires diffrax and jax installed in venv." - ) - - def step(self, x): - s, sn, sigma_down, sigma_up = self.de_get_step(x) - if self.de_min_sigma is not None and s <= self.de_min_sigma: - return (yield from self.euler_step(x)) - ss = self.ss - bidx = 0 - mcc = 0 - _b, c, h, w = x.shape - interrupted = None - t0, t1 = sigma_down.item(), s.item() - - def odefn_(t_orig, y_flat, args=()): - nonlocal mcc, interrupted - t = self.reverse_time(self.j2t(t_orig).to(s), t0, t1) - if t <= 1e-05: - return jax.numpy.zeros_like(y_flat) - if mcc >= self.de_max_nfe: - raise RuntimeError("DiffraxStep: Model call limit exceeded") - y = self.j2t(y_flat.reshape(1, c, h, w)).to(x) - t32 = t.to(s).clamp(min=1e-05) - flat_shape = y_flat.shape - del y_flat - - if not args and mcc == 0 and torch.all(t == s): - mr_cached = True - mr = ss.hcur - mcc = 1 - else: - mr_cached = False - try: - if not args: - mr = self.call_model(y, t32, call_index=mcc, s_in=t.new_ones(1)) - else: - print("TANGENTS") - mr = self.call_model( - y, - t32, - call_index=mcc, - tangents=args, - s_in=t.new_ones(1), - ) - except comfy.model_management.InterruptProcessingException as exc: - interrupted = exc - raise - mcc += 1 - result = self.to_d(mr)[bidx if mr_cached else 0].reshape(*flat_shape) - return self.t2j(-result) - - if not self.de_fake_pure_callback: - - def odefn(t, y_flat, args): - return jax.experimental.io_callback( - odefn_, y_flat, t, y_flat, ordered=True - ) - - else: - - def odefn(t, y_flat, args): - return jax.pure_callback(odefn_, y_flat, t, y_flat) - - def g(t, y, _args): - if self.de_g_split_time_mode: - val = jax.lax.cond( - t < t0 + (t1 - t0) * 0.5, - lambda: self.de_g_multiplier, - lambda: -self.de_g_multiplier, - ) - else: - val = self.de_g_multiplier - if self.de_g_time_scaling: - val *= self.reverse_time(t, t0, t1) if self.de_g_reverse_time else t - if not self.de_batch_channels: - return val - return jax.numpy.float32(val).broadcast((y.shape[0],)) - - def progress_callback(): - return f" ({mcc:>3}/{self.de_max_nfe:>3}) {self.de_solver_name}" - - term = diffrax.ODETerm(odefn) - method = self.de_solver_method - if self.de_adaptive: - controller = diffrax.PIDController( - atol=self.de_atol, - rtol=self.de_rtol, - dtmin=1e-05, - pcoeff=self.de_ctl_pcoeff, - icoeff=self.de_ctl_icoeff, - dcoeff=self.de_ctl_dcoeff, - error_order=self.de_error_order, - ) - else: - controller = diffrax.ConstantStepSize() - - if not self.de_adaptive: - dt0 = (t1 - t0) / self.de_split - else: - dt0 = (t1 - t0) * self.de_initial_step - if self.de_sde_mode: - bm = diffrax.VirtualBrownianTree( - t0=ss.sigmas.min().item(), - t1=ss.sigmas.max().item(), - tol=1e-06, - levy_area=self.de_levy_area_approx, - shape=(c,) if self.de_batch_channels else (), - key=jax.random.PRNGKey(ss.noise.seed + ss.noise.seed_offset), - ) - term = diffrax.MultiTerm(term, diffrax.ControlTerm(g, bm)) - results = [] - for batch in tqdm.trange( - 1, - x.shape[0] + 1, - desc="batch", - leave=False, - disable=x.shape[0] == 1 or ss.disable_status, - ): - bidx = batch - 1 - mcc = 0 - if self.de_batch_channels: - y_flat = x[bidx].flatten(start_dim=1) - else: - y_flat = x[bidx].unsqueeze(0).flatten(start_dim=1) - y_flat = self.t2j(y_flat) - with warnings.catch_warnings(): - warnings.simplefilter(action="ignore", category=FutureWarning) - try: - solution = diffrax.diffeqsolve( - terms=term, - solver=method, - t0=t0, - t1=t1, - dt0=dt0, - y0=y_flat, - saveat=diffrax.SaveAt(t1=True), - stepsize_controller=controller, - progress_meter=StepCallbackTqdmProgressMeter( - step_callback=progress_callback, - refresh_steps=1, - ), - ) - except Exception: - if interrupted is not None: - raise interrupted - raise - results.append(self.j2t(solution.ys).view(1, *x.shape[1:])) - del solution - result = torch.cat(results).to(x) - sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) - yield from self.result(result, sigma_up, sigma_down=sigma_down) - - -class HeunStep(ReversibleSingleStepSampler): - name = "heun" - model_calls = 1 - default_history_limit, max_history = 0, 0 - allow_alt_cfgpp = True - - def reversible_correction(self, d_from, d_to): - reta, reversible_scale = self.get_reversible_cfg() - if reversible_scale == 0: - return 0 - sdr = self.get_ancestral_step(reta)[0] - dtr = sdr - self.ss.sigma - return (dtr**2 * (d_to - d_from) / 4) * self.reversible_scale - - def step(self, x): - ss = self.ss - s = ss.sigma - sd, su = self.get_ancestral_step(self.get_dyn_eta()) - dt = sd - s - hcur = ss.hcur - d = self.to_d(hcur) - x_next = hcur.denoised + d * sd - d_next = self.to_d(self.call_model(x_next, sd, call_index=1)) - result = hcur.denoised + d * s - result += (dt * (d + d_next)) * 0.5 - result -= self.reversible_correction(d, d_next) - yield from self.result(result, su, sigma_down=sd) - - -class Heun1SStep(HeunStep): - name = "heun_1s" - model_calls = 1 - allow_alt_cfgpp = True - default_history_limit, max_history = 1, 1 - - def step(self, x): - ss = self.ss - s = ss.sigma - if self.available_history() == 0: - return (yield from super().step(x)) - hcur, hprev = ss.hcur, ss.hprev - d_prev = self.to_d(hprev) - sd, su = self.get_ancestral_step(self.get_dyn_eta()) - dt = sd - s - d = self.to_d(hcur) - result = hcur.denoised + hcur.sigma * self.to_d(hcur) - result += (dt * (d_prev + d)) * 0.5 - result -= self.reversible_correction(d_prev, d) - yield from self.result(result, su, sigma_down=sd) - - -class AdapterStep(SingleStepSampler): - name = "adapter" - model_calls = 2 - immiscible = False - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.external_sampler = self.options.pop( - "SAMPLER", comfy.samplers.sampler_object("euler") - ) - sig = inspect.signature(self.external_sampler.sampler_function) - self.external_sampler_options = { - k: v - for k, v in self.options.pop("external_sampler", {}).items() - if k in sig.parameters - } - self.external_sampler_uses_noise = "noise_sampler" in sig.parameters - self.ancestralize = self.options.pop("ancestralize", self.ancestralize) is True - - def step(self, x): - ss = self.ss - sigmas = ss.sigmas[ss.idx : ss.idx + 2] - kwargs = { - "callback": None, - "disable": True, - "extra_args": {"seed": ss.noise.seed + ss.noise.seed_offset}, - } | self.external_sampler_options - if self.external_sampler_uses_noise: - kwargs["noise_sampler"] = ss.noise.make_caching_noise_sampler( - self.options.get("custom_noise"), - 1, - sigmas[-1], - sigmas[0], - immiscible=fallback(self.immiscible, ss.noise.immiscible), - ) - - mcc = 1 - - def model_wrapper(x_, sigma_, *args, **kwargs): - nonlocal mcc - if torch.equal(x_, x) and sigma_ == ss.sigma: - return ss.hcur.denoised.clone() - mr = self.call_model(x_, sigma_, *args, call_index=mcc, **kwargs) - mcc += 1 - return mr.denoised.clone() - - result = self.external_sampler.sampler_function( - model_wrapper, x.clone(), sigmas, **kwargs - ) - yield from self.result(result, ss.sigma.new_zeros(1)) - - -# Based on https://github.com/Extraltodeus/DistanceSampler -class DistanceStep(SingleStepSampler): - name = "distance" - allow_alt_cfgpp = True - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - distance = self.options.get("distance", {}) - self.distance_resample = distance.get("resample", 3) - self.distance_resample_end = distance.get("resample_end", 1) - self.distance_resample_eta = distance.get("eta", self.eta) - self.distance_resample_s_noise = distance.get("s_noise", self.s_noise) - self.distance_alt_cfgpp_scale = distance.get("alt_cfgpp_scale", 0.0) - self.distance_first_eta_resample_step = distance.get("first_eta_step", 0) - self.distance_last_eta_resample_step = distance.get("last_eta_step", -1) - - @property - def require_uncond(self): - return super().require_uncond or self.distance_alt_cfgpp_scale != 0 - - def distance_resample_steps(self): - ss = self.ss - resample, resample_end = self.distance_resample, self.distance_resample_end - if resample == -1: - current_resample = min(10, (ss.sigmas.shape[0] - ss.idx) // 2) - else: - current_resample = resample - if resample_end < 0: - return current_resample - sigma = ss.sigma - s_min = (ss.sigmas if ss.sigmas[-1] > 0 else ss.sigmas[:-1]).min() - s_max = ss.sigmas.max() - res_mul = max(0, min(1, ((sigma - s_min) / (s_max - s_min)) ** 0.5)) - return max( - min(current_resample, resample_end), - min( - max(current_resample, resample_end), - int(current_resample * res_mul + resample_end * (1 - res_mul)), - ), - ) - - @staticmethod - def distance_weights(t, p): - batch = t.shape[0] - d = torch.stack( - tuple((t - t[idx]).abs().sum(dim=0) for idx in range(batch)), - dim=0, - ) - d_min, d_max = d.min(), d.max() - d = torch.nan_to_num( - (1 - (d - d_min) / (d_max - d_min)).pow(p), - nan=1, - neginf=1, - posinf=1, - ) - d /= d.sum(dim=0) - return d.mul_(t).sum(dim=0) - - def step(self, x): - resample_steps = self.distance_resample_steps() - if resample_steps < 1: - return (yield from self.euler_step(x)) - ss = self.ss - sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) - rsigma_down, rsigma_up = self.get_ancestral_step(eta=self.distance_resample_eta) - rsigma_up *= self.distance_resample_s_noise - sigma, sigma_next = ss.sigma, ss.sigma_next - zero_up = sigma * 0 - d = self.to_d(ss.hcur) - can_ancestral = not torch.equal(rsigma_down, sigma_next) - start_eta_idx, end_eta_idx = ( - max(0, resample_steps + v if v < 0 else v) - for v in ( - self.distance_first_eta_resample_step, - self.distance_last_eta_resample_step, - ) - ) - dt = sigma_down - sigma - d = self.to_d(ss.hcur) - x_n = [d] - for re_step in tqdm.trange( - resample_steps, desc="distance_resample", disable=ss.disable_status - ): - if can_ancestral and start_eta_idx <= re_step <= end_eta_idx: - curr_sigma_down, curr_sigma_up = rsigma_down, rsigma_up - else: - curr_sigma_down, curr_sigma_up = sigma_next, zero_up - rdt = curr_sigma_down - sigma - x_new = x + d * rdt - if curr_sigma_up != 0: - x_new = yield from self.result( - x_new, - curr_sigma_up, - sigma=sigma, - sigma_down=curr_sigma_down, - final=False, - ) - sr = self.call_model(x_new, sigma_next, call_index=re_step + 1) - new_d = sr.to_d( - sigma=curr_sigma_down, alt_cfgpp_scale=self.distance_alt_cfgpp_scale - ) - x_n.append(new_d) - if re_step == 0: - d = (new_d + d) / 2 - else: - d = self.distance_weights(torch.stack(x_n), re_step + 2) - x_n.append(d) - yield from self.result(x + d * dt, sigma_up, sigma_down=sigma_down) - - -class DynamicStep(SingleStepSampler): - name = "dynamic" - sample_sigma_zero = True - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - dynamic = self.options.get("dynamic") - if dynamic is None: - raise ValueError( - "Dynamic sampler type requires specifying dynamic block in text parameters" - ) - if isinstance(dynamic, str): - dynamic = ({"expression": dynamic},) - elif not isinstance(dynamic, (tuple, list)): - raise ValueError( - "Bad type for dynamic block: must be string or list of objects" - ) - elif len(dynamic) == 0: - raise ValueError("Dynamic block as a list cannot be empty") - dynresult = [] - for idx, item in enumerate(dynamic): - if not isinstance(item, dict): - raise ValueError( - f"Bad item in dynamic block at index {idx}: must be a dict" - ) - dyn_when = item.get("when") - if isinstance(dyn_when, str): - dyn_when = expr.Expression(dyn_when) - elif dyn_when is not None: - raise ValueError( - f"Unexpected type for when key in dynamic block at index {idx}, must be string or null/unset" - ) - dyn_params = item.get("expression") - if not isinstance(dyn_params, str): - raise ValueError( - f"Missing or incorrectly typed expression key for dynamic block at index {idx}: must be a string" - ) - dynresult.append((dyn_when, expr.Expression(dyn_params))) - self.dynamic = tuple(dynresult) - - def step(self, x): - sampler_params = None - handlers = filtering.FILTER_HANDLERS.clone(constants=self.ss.refs) - for idx, (dyn_when, dyn_params) in enumerate(self.dynamic): - if dyn_when is not None and not bool(dyn_when.eval(handlers)): - continue - sampler_params = dyn_params.eval(handlers) - if sampler_params is not None: - break - if sampler_params is None: - raise RuntimeError( - "Dynamic sampler could not find matching sampler: all expressions failed to return a result" - ) - if not isinstance(sampler_params, dict): - raise TypeError( - f"Dynamic sampler expression must evaluate to a dict, got type {type(sampler_params)}" - ) - if bool(sampler_params.get("dynamic_inherit")): - copy_keys = ( - "s_noise", - "eta", - "pre_filter", - "post_filter", - "immiscible", - ) - opts = {k: getattr(self, k) for k in copy_keys} - else: - opts = {} - opts["custom_noise"] = self.custom_noise - opts |= sampler_params - opts |= {k: v for k, v in self.options.items() if k.startswith("custom_noise_")} - # print("\n\nDYN OPTS", opts) - step_method = opts.get("step_method", "default") - sampler_class = STEP_SAMPLER_SIMPLE_NAMES.get(step_method) - if sampler_class is None: - raise ValueError(f"Unknown step method {step_method} in dynamic sampler") - sampler = sampler_class(**opts) - with StepSamplerContext(sampler, self.ss) as sampler: - yield from sampler.step(x) - - -STEP_SAMPLERS = { - "default (euler)": EulerStep, - "adapter (variable)": AdapterStep, - "bogacki (2)": BogackiStep, - "deis": DEISStep, - "distance (variable)": DistanceStep, - "dpmpp_2m_sde": DPMPP2MSDEStep, - "dpmpp_2m": DPMPP2MStep, - "dpmpp_2s": DPMPP2SStep, - "dpmpp_3m_sde": DPMPP3MSDEStep, - "dpmpp_sde (1)": DPMPPSDEStep, - "dynamic (variable)": DynamicStep, - "euler_cycle": EulerCycleStep, - "euler_dancing": EulerDancingStep, - "euler": EulerStep, - "heun (1)": HeunStep, - "heun_1s (1)": Heun1SStep, - "heunpp (1-2)": HeunPP2Step, - "ipndm_v": IPNDMVStep, - "ipndm": IPNDMStep, - "res (1)": RESStep, - "reversible_bogacki (2)": ReversibleBogackiStep, - "reversible_heun (1)": ReversibleHeunStep, - "reversible_heun_1s": ReversibleHeun1SStep, - "rk4 (3)": RK4Step, - "rkf45 (4)": RKF45Step, - "rk_dynamic": RKDynamicStep, - "solver_diffrax (variable)": DiffraxStep, - "solver_torchdiffeq (variable)": TDEStep, - "solver_torchode (variable)": TODEStep, - "solver_torchsde (variable)": TSDEStep, - "trapezoidal (1)": TrapezoidalStep, - "trapezoidal_cycle (1)": TrapezoidalCycleStep, - "ttm_jvp (1)": TTMJVPStep, -} - -STEP_SAMPLER_SIMPLE_NAMES = {k.split(None, 1)[0]: v for k, v in STEP_SAMPLERS.items()} - -__all__ = ("STEP_SAMPLERS", "STEP_SAMPLER_SIMPLE_NAMES") diff --git a/py/step_samplers/__init__.py b/py/step_samplers/__init__.py new file mode 100644 index 0000000..60503e0 --- /dev/null +++ b/py/step_samplers/__init__.py @@ -0,0 +1,20 @@ +from . import ( # noqa: F401 + builtins, + blep, + clybius, + extraltodeus, + misc, + solver_tde, + solver_tode, + solver_tsde, + solver_diffrax, +) + +from . import registry + +registry.init() + +STEP_SAMPLERS = registry.STEP_SAMPLERS +STEP_SAMPLER_SIMPLE_NAMES = registry.STEP_SAMPLER_SIMPLE_NAMES + +__all__ = ("STEP_SAMPLERS", "STEP_SAMPLER_SIMPLE_NAMES") diff --git a/py/step_samplers/base.py b/py/step_samplers/base.py new file mode 100644 index 0000000..9fdb4c6 --- /dev/null +++ b/py/step_samplers/base.py @@ -0,0 +1,536 @@ +import contextlib +import typing + +import torch + +from . import registry # noqa: F401 + +from .. import filtering, noise, utils +from ..utils import fallback + + +class SamplerResult: + CLONE_KEYS = ( + "denoised_cond", + "denoised_uncond", + "denoised", + "final", + "is_rectified_flow", + "noise_pred", + "noise_sampler", + "s_noise", + "sampler", + "sigma_down", + "sigma_next", + "sigma_up", + "sigma", + "step", + "substep", + "x_", + ) + + def __init__( + self, + ss, + sampler, + x, + sigma_up=None, + *, + split_result=None, + sigma=None, + sigma_next=None, + sigma_down=None, + s_noise=None, + noise_sampler=None, + final=True, + ): + self.is_rectified_flow = ss.model.is_rectified_flow + self.sampler = sampler + self.sigma_up = fallback(sigma_up, ss.sigma.new_zeros(1)) + self.s_noise = fallback(s_noise, sampler.s_noise) + self.sigma = fallback(sigma, ss.sigma) + self.sigma_next = fallback(sigma_next, ss.sigma_next) + self.sigma_down = fallback(sigma_down, self.sigma_next) + self.noise_sampler = fallback(noise_sampler, sampler.noise_sampler) + self.final = final + self.step = ss.step + self.substep = ss.substep + self.x_ = x + if split_result is not None: + self.denoised, self.noise_pred = split_result + elif x is None: + raise ValueError("SamplerResult requires at least one of x, split_result") + else: + self.denoised = self.noise_pred = None + _ = self.extract_pred(ss) + self.denoised_uncond = ss.hcur.denoised_uncond + self.denoised_cond = ss.hcur.denoised_cond + + def get_noise(self, *, scaled=True, ss=None): + if self.sigma_next == 0 or self.noise_scale == 0: + return torch.zeros_like(self.x_) + return self.noise_sampler( + self.sigma, + self.sigma_next, + out_hw=self.x.shape[-2:], + x_ref=self.x, + refs=filtering.FilterRefs.from_sr(self) if ss is None else ss.refs, + ).mul_(self.noise_scale if scaled else 1.0) + + def extract_pred(self, ss): + if self.denoised is None or self.noise_pred is None: + self.denoised, self.noise_pred = utils.extract_pred( + ss.hcur.x, self.x_, ss.sigma, self.sigma_down + ) + return self.denoised, self.noise_pred + + @property + def x(self): + if self.x_ is None: + self.x_ = self.denoised + self.sigma_down * self.noise_pred + return self.x_ + + @property + def noise_scale(self): + return self.sigma_up * self.s_noise + + def noise_x(self, x=None, scale=1.0, *, ss=None): + x = fallback(x, self.x) + if self.sigma_next == 0 or self.noise_scale == 0: + return x + noise = self.get_noise(ss=ss).mul_(scale) + if not self.is_rectified_flow: + return noise.add_(x) + x_coeff = (1 - self.sigma_next) / (1 - self.sigma_down) + # print(f"\nRF noise: {x_coeff}") + return noise.add_(x_coeff * x) + + def clone(self): + obj = self.__new__(self.__class__) + for k in self.CLONE_KEYS: + if hasattr(self, k): + setattr(obj, k, getattr(self, k)) + return obj + + +class StepSamplerContext: + def __init__(self, sampler, *args, **kwargs): + self.sampler = sampler + self.args = args + self.kwargs = kwargs + + def __enter__(self): + if self.sampler.ss is not None: + raise RuntimeError("Cannot reenter prepared sampler in context manager!") + self.sampler.prepare(*self.args, **self.kwargs) + return self.sampler + + def __exit__(self, *_unused): + self.sampler.reset() + + +class SingleStepSampler: + name = None + self_noise = 0 + model_calls = 0 + ancestralize = False + sample_sigma_zero = False + immiscible = None + allow_cfgpp = False + allow_alt_cfgpp = False + afs_end_step = -1 + uses_alt_noise = False + + default_eta = 1.0 + + def __init__( + self, + *, + noise_sampler=None, + substeps=1, + s_noise=1.0, + eta=None, + eta_retry_increment=0, + dyn_eta_start=None, + dyn_eta_end=None, + weight=1.0, + pre_filter=None, + post_filter=None, + immiscible=None, + **kwargs, + ): + self.ss = None + self.options = kwargs + self.cfgpp = self.allow_cfgpp and self.options.pop("cfgpp", False) is True + alt_cfgpp_scale = self.options.pop("alt_cfgpp_scale", 0.0) + self.alt_cfgpp_scale = 0.0 if not self.allow_alt_cfgpp else alt_cfgpp_scale + self.s_noise = s_noise + self.eta = fallback(eta, self.default_eta) + self.eta_retry_increment = eta_retry_increment + self.dyn_eta_start = dyn_eta_start + self.dyn_eta_end = dyn_eta_end + self.noise_sampler = noise_sampler + self.immiscible = ( + noise.ImmiscibleNoise(**immiscible) + if immiscible not in (False, None) + else immiscible + ) + self.weight = weight + self.afs_end_step = self.options.pop("afs_end_step", -1) + self.substeps = substeps + self.pre_filter = ( + None if pre_filter is None else filtering.make_filter(pre_filter) + ) + self.post_filter = ( + None if post_filter is None else filtering.make_filter(post_filter) + ) + self.custom_noise = self.options.get("custom_noise") + if isinstance(self.custom_noise, str): + self.custom_noise = self.options.get(f"custom_noise_{self.custom_noise}") + if not self.uses_alt_noise: + return + self.alt_custom_noise = self.options.get("custom_noise_alt") + alt_immiscible = self.options.get("alt_immiscible") + self.alt_immiscible = ( + noise.ImmiscibleNoise(**alt_immiscible) + if isinstance(alt_immiscible, dict) + else alt_immiscible + ) + + def __call__(self, x): + ss = self.ss + orig_x = x + if not self.sample_sigma_zero and ss.sigma_next == 0: + return (yield from self.denoised_result()) + if ss.step <= self.afs_end_step: + return (yield from self.afs_step(x)) + if self.pre_filter or self.post_filter: + filter_refs = ss.refs | filtering.FilterRefs({"orig_x": orig_x}) + if self.pre_filter: + x = self.pre_filter.apply(x, refs=filter_refs) + next_x = None + sg = self.step(x) + with contextlib.suppress(StopIteration): + while True: + sr = sg.send(next_x) + if sr.final: + if self.ancestralize: + sr = self.ancestralize_result(sr) + curr_x = sr.x + if self.post_filter: + curr_x = self.post_filter.apply(curr_x, refs=filter_refs) + sr.x_ = curr_x + return (yield sr) + next_x = sr.noise_x(ss=ss) + + def step(self, x): + raise NotImplementedError + + def prepare(self, ss): + self.ss = ss + self.noise_sampler = ss.noise.make_caching_noise_sampler( + self.custom_noise, + self.max_noise_samples, + ss.sigma, + ss.sigma_next, + immiscible=fallback(self.immiscible, ss.noise.immiscible), + ) + if not self.uses_alt_noise: + return + if self.alt_custom_noise is None and self.alt_immiscible is None: + self.alt_noise_sampler = self.noise_sampler + return + self.alt_noise_sampler = ss.noise.make_caching_noise_sampler( + fallback(self.alt_custom_noise, self.custom_noise), + 1, + ss.sigma, + ss.sigma_next, + immiscible=fallback( + fallback(self.alt_immiscible, self.immiscible), + ss.noise.immiscible, + ), + ) + + def reset(self): + self.ss = None + self.noise_sampler = None + if self.uses_alt_noise: + self.alt_noise_sampler = None + + # From https://arxiv.org/abs/2210.05475 + def afs_step(self, x): + sigma, sigma_next = self.ss.sigma, self.ss.sigma_next + afs_d = x / ((1 + sigma**2).sqrt()) + dt = sigma_next - sigma + return (yield from self.result(x + afs_d * dt)) + + # Euler - based on original ComfyUI implementation + def euler_step( + self, + x, + *, + sigma_down=None, + sigma_up=None, + eta=None, + sigma=None, + sigma_next=None, + ): + eta = fallback(eta, self.get_dyn_eta()) + if sigma_down is None or sigma_up is None: + if not (sigma_down is None and sigma_up is None): + raise ValueError("Must pass both sigma_down and sigma_up or neither") + sigma_down, sigma_up = self.get_ancestral_step( + eta=eta, sigma=sigma, sigma_next=sigma_next + ) + return ( + yield from self.split_result( + *self.get_split_prediction(), sigma_down=sigma_down, sigma_up=sigma_up + ) + ) + + def denoised_result(self, **kwargs): + ss = self.ss + return ( + yield SamplerResult(ss, self, ss.denoised, ss.sigma.new_zeros(1), **kwargs) + ) + + def result(self, x, noise_scale=None, **kwargs): + return (yield SamplerResult(self.ss, self, x, noise_scale, **kwargs)) + + def split_result( + self, denoised=None, noise_pred=None, sigma_up=None, sigma_down=None, **kwargs + ): + return ( + yield SamplerResult( + ss=self.ss, + sampler=self, + x=None, + sigma_up=sigma_up, + sigma_down=sigma_down, + split_result=(denoised, noise_pred), + **kwargs, + ) + ) + + def get_ancestral_step( + self, *args, dyn_eta=False, as_dict=False, retry_increment=None, **kwargs + ): + if dyn_eta: + args = (self.get_dyn_eta(), *args) + retry_increment = fallback(retry_increment, self.eta_retry_increment) + sigma_down, sigma_up = self.ss.get_ancestral_step( + *args, retry_increment=retry_increment, **kwargs + ) + if not as_dict: + return sigma_down, sigma_up + return {"sigma_down": sigma_down, "sigma_up": sigma_up} + + def ancestralize_result(self, sr): + ss = self.ss + new_sr = sr.clone() + if new_sr.sigma_down is not None and new_sr.sigma_down != new_sr.sigma_next: + return sr + eta = self.get_dyn_eta() + if sr.sigma_next == 0 or eta == 0: + return sr + sd, su = self.get_ancestral_step(eta, sigma=sr.sigma, sigma_next=sr.sigma_next) + _ = new_sr.extract_pred(ss) + new_sr.x_ = None + new_sr.sigma_up = su + new_sr.sigma_down = sd + return new_sr + + def __str__(self): + return f"" + + def get_dyn_value(self, start, end): + if None in (start, end): + return 1.0 + if start == end: + return start + ss = self.ss + main_idx = getattr(ss, "main_idx", ss.idx) + main_sigmas = getattr(ss, "main_sigmas", ss.sigmas) + step_pct = main_idx / (len(main_sigmas) - 1) + dd_diff = end - start + return start + dd_diff * step_pct + + def get_dyn_eta(self): + return self.eta * self.get_dyn_value(self.dyn_eta_start, self.dyn_eta_end) + + @property + def max_noise_samples(self): + return (1 + self.self_noise) * self.substeps + + @property + def require_uncond(self): + return self.cfgpp or self.alt_cfgpp_scale != 0 + + def to_d(self, mr, *, use_cfgpp=True, **kwargs): + if not use_cfgpp: + return mr.to_d(**kwargs) + return mr.to_d(alt_cfgpp_scale=self.alt_cfgpp_scale, cfgpp=self.cfgpp, **kwargs) + + def get_split_prediction(self, *, mr=None, sigma=None, **kwargs): + mr = fallback(mr, self.ss.hcur) + sigma = fallback(sigma, mr.sigma) + return mr.get_split_prediction( + sigma=sigma, + alt_cfgpp_scale=self.alt_cfgpp_scale, + cfgpp=self.cfgpp, + **kwargs, + ) + + def call_model(self, *args, **kwargs): + ss = self.ss + kwargs["require_uncond"] = self.require_uncond or kwargs.get( + "require_uncond", False + ) + kwargs["cfg_scale_override"] = kwargs.get( + "cfg_scale_override", + self.options.get("cfg_scale_override", ss.cfg_scale_override), + ) + return ss.call_model(*args, ss=ss, **kwargs) + + def step_mix(self, x, denoised, uncond, ratio, *, blend=torch.lerp): + if self.cfgpp: + return denoised + (x - uncond).mul_(ratio) + pp = self.alt_cfgpp_scale + if pp == 0: + return blend(denoised, x, ratio) + return blend(denoised * (1 + pp) - uncond * pp, x, ratio) + + +class HistorySingleStepSampler(SingleStepSampler): + default_history_limit, max_history = 0, 0 + + def __init__(self, *args, history_limit=None, **kwargs): + super().__init__(*args, **kwargs) + self.history_limit = min( + self.max_history, + max( + 0, + self.default_history_limit if history_limit is None else history_limit, + ), + ) + + def available_history(self): + ss = self.ss + available = max( + 0, min(ss.idx, self.history_limit, self.max_history, len(ss.hist) - 1) + ) + if not available: + return available + curr_shape = ss.hist[-1].denoised.shape + for eff_available in range(available): + if ss.hist[-2 - eff_available].denoised.shape != curr_shape: + return eff_available + return available + + +class ReversibleConfig(typing.NamedTuple): + scale: float + eta: float + dyn_eta_start: float | None = None + dyn_eta_end: float | None = None + eta_retry_increment: float = 0.0 + start_step: int = 0 + end_step: int = 9999 + use_cfgpp: bool = False + + @classmethod + def build(cls, *, default_eta, default_scale, eta=None, scale=None, **kwargs): + return cls.__new__( + cls, + eta=fallback(eta, default_eta), + scale=fallback(scale, default_scale), + **kwargs, + ) + + def check(self, step): + return self.scale != 0 and self.start_step <= step <= self.end_step + + +class ReversibleSingleStepSampler(HistorySingleStepSampler): + default_reversible_scale = 1.0 + default_reta = 1.0 + + def __init__( + self, + *, + reversible_scale=None, + reta=None, + dyn_reta_start=None, + dyn_reta_end=None, + reversible_start_step=0, + reversible=None, + **kwargs, + ): + super().__init__(**kwargs) + if reversible is None: + # For backward compatibility. + self.reversible = ReversibleConfig.build( + default_eta=self.default_reta, + default_scale=self.default_reversible_scale, + scale=reversible_scale, + eta=reta, + dyn_eta_start=dyn_reta_start, + dyn_eta_end=dyn_reta_end, + start_step=reversible_start_step, + ) + return + self.reversible = ReversibleConfig.build( + default_eta=self.default_reta, + default_scale=self.default_reversible_scale, + **reversible, + ) + + def reversible_correction(self): + raise NotImplementedError + + def get_dyn_reta(self, *, r=None): + r = fallback(r, self.reversible) + ss = self.ss + if not r.check(ss.step): + return 0.0 + return r.eta * self.get_dyn_value(r.dyn_eta_start, r.dyn_eta_end) + + dyn_reta = property(get_dyn_reta) + + def get_reversible_cfg(self, *, reversible=None): + reversible = fallback(reversible, self.reversible) + ss = self.ss + if not reversible.check(ss.step): + return 0.0, 0.0 + return self.get_dyn_reta(r=reversible), reversible.scale + + +class DPMPPStepMixin: + @staticmethod + def sigma_fn(t): + return t.neg().exp() + + @staticmethod + def t_fn(t): + return t.log().neg() + + +class MinSigmaStepMixin: + @staticmethod + def adjust_step(sigma, min_sigma, threshold=5e-04): + if min_sigma - sigma > threshold: + return sigma.clamp(min=min_sigma) + return sigma + + def adjusted_step(self, sn, result, mcc, sigma_up): + ss = self.ss + if sn == ss.sigma_next: + return sigma_up, result + # FIXME: Make sure we're noising from the right sigma. + result = yield from self.result( + result, sigma_up, sigma=ss.sigma, sigma_next=sn, final=False + ) + mr = self.call_model(result, sn, call_index=mcc) + dt = ss.sigma_next - sn + result = result + self.to_d(mr) * dt + return sigma_up.new_zeros(1), result diff --git a/py/step_samplers/blep.py b/py/step_samplers/blep.py new file mode 100644 index 0000000..aa91433 --- /dev/null +++ b/py/step_samplers/blep.py @@ -0,0 +1,546 @@ +import typing + +import inspect +import torch + +import comfy + +from .. import filtering +from .. import expression as expr +from ..utils import fallback +from .base import ( + StepSamplerContext, + SingleStepSampler, + registry, +) + + +try: + import pytorch_wavelets as ptwav + + HAVE_WAVELETS = True +except ImportError: + HAVE_WAVELETS = False + + +class DynamicStep(SingleStepSampler): + name = "dynamic" + sample_sigma_zero = True + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + dynamic = self.options.get("dynamic") + if dynamic is None: + raise ValueError( + "Dynamic sampler type requires specifying dynamic block in text parameters" + ) + if isinstance(dynamic, str): + dynamic = ({"expression": dynamic},) + elif not isinstance(dynamic, (tuple, list)): + raise ValueError( + "Bad type for dynamic block: must be string or list of objects" + ) + elif len(dynamic) == 0: + raise ValueError("Dynamic block as a list cannot be empty") + dynresult = [] + for idx, item in enumerate(dynamic): + if not isinstance(item, dict): + raise ValueError( + f"Bad item in dynamic block at index {idx}: must be a dict" + ) + dyn_when = item.get("when") + if isinstance(dyn_when, str): + dyn_when = expr.Expression(dyn_when) + elif dyn_when is not None: + raise ValueError( + f"Unexpected type for when key in dynamic block at index {idx}, must be string or null/unset" + ) + dyn_params = item.get("expression") + if not isinstance(dyn_params, str): + raise ValueError( + f"Missing or incorrectly typed expression key for dynamic block at index {idx}: must be a string" + ) + dynresult.append((dyn_when, expr.Expression(dyn_params))) + self.dynamic = tuple(dynresult) + + def step(self, x): + sampler_params = None + handlers = filtering.FILTER_HANDLERS.clone(constants=self.ss.refs) + for idx, (dyn_when, dyn_params) in enumerate(self.dynamic): + if dyn_when is not None and not bool(dyn_when.eval(handlers)): + continue + sampler_params = dyn_params.eval(handlers) + if sampler_params is not None: + break + if sampler_params is None: + raise RuntimeError( + "Dynamic sampler could not find matching sampler: all expressions failed to return a result" + ) + if not isinstance(sampler_params, dict): + raise TypeError( + f"Dynamic sampler expression must evaluate to a dict, got type {type(sampler_params)}" + ) + if bool(sampler_params.get("dynamic_inherit")): + copy_keys = ( + "s_noise", + "eta", + "pre_filter", + "post_filter", + "immiscible", + ) + opts = {k: getattr(self, k) for k in copy_keys} + else: + opts = {} + opts["custom_noise"] = self.custom_noise + opts |= sampler_params + opts |= {k: v for k, v in self.options.items() if k.startswith("custom_noise_")} + # print("\n\nDYN OPTS", opts) + step_method = opts.get("step_method", "default") + sampler_class = registry.STEP_SAMPLER_SIMPLE_NAMES.get(step_method) + if sampler_class is None: + raise ValueError(f"Unknown step method {step_method} in dynamic sampler") + sampler = sampler_class(**opts) + with StepSamplerContext(sampler, self.ss) as sampler: + yield from sampler.step(x) + + +class AdapterStep(SingleStepSampler): + name = "adapter" + model_calls = 2 + immiscible = False + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.external_sampler = self.options.pop( + "SAMPLER", comfy.samplers.sampler_object("euler") + ) + sig = inspect.signature(self.external_sampler.sampler_function) + self.external_sampler_options = { + k: v + for k, v in self.options.pop("external_sampler", {}).items() + if k in sig.parameters + } + self.external_sampler_uses_noise = "noise_sampler" in sig.parameters + self.ancestralize = self.options.pop("ancestralize", self.ancestralize) is True + + def step(self, x): + ss = self.ss + sigmas = ss.sigmas[ss.idx : ss.idx + 2] + kwargs = { + "callback": None, + "disable": True, + "extra_args": {"seed": ss.noise.seed + ss.noise.seed_offset}, + } | self.external_sampler_options + if self.external_sampler_uses_noise: + kwargs["noise_sampler"] = ss.noise.make_caching_noise_sampler( + self.options.get("custom_noise"), + 1, + sigmas[-1], + sigmas[0], + immiscible=fallback(self.immiscible, ss.noise.immiscible), + ) + + mcc = 1 + + def model_wrapper(x_, sigma_, *args, **kwargs): + nonlocal mcc + if torch.equal(x_, x) and sigma_ == ss.sigma: + return ss.hcur.denoised.clone() + mr = self.call_model(x_, sigma_, *args, call_index=mcc, **kwargs) + mcc += 1 + return mr.denoised.clone() + + result = self.external_sampler.sampler_function( + model_wrapper, x.clone(), sigmas, **kwargs + ) + yield from self.result(result, ss.sigma.new_zeros(1)) + + +class CycleSingleStepSampler(SingleStepSampler): + default_eta = 0.0 + + def __init__(self, *, cycle_pct=0.25, cycle_adjust_scales=True, **kwargs): + super().__init__(**kwargs) + if cycle_pct < 0: + raise ValueError("cycle_pct must be positive") + self.cycle_pct = cycle_pct + self.cycle_adjust_scales = cycle_adjust_scales + + def get_cycle_scales(self, sigma_next): + keep_scale = 1 - self.cycle_pct + if not self.cycle_adjust_scales: + return keep_scale, self.cycle_pct + add_scale = ((sigma_next**2.0 - (keep_scale * sigma_next) ** 2.0) ** 0.5) * ( + 0.95 + 0.25 * self.cycle_pct + ) + # print(f">> keep={keep_scale}, add={add_scale}") + return keep_scale, add_scale + + +class EulerCycleStep(CycleSingleStepSampler): + name = "blep_euler_cycle" + allow_alt_cfgpp = True + allow_cfgpp = True + + def step(self, x): + sigma_next = self.ss.sigma_next + denoised_pred, d = self.get_split_prediction() + keep_scale, add_scale = self.get_cycle_scales(sigma_next) + return ( + yield from self.split_result( + denoised_pred, d * keep_scale, sigma_up=add_scale, sigma_down=sigma_next + ) + ) + + +class TrapezoidalCycleStep(CycleSingleStepSampler): + name = "blep_trapezoidal_cycle" + model_calls = 1 + allow_alt_cfgpp = False + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + blend_mode = self.options.get("blend_mode", "lerp").strip() + self.blend = ( + filtering.BLENDING_MODES[blend_mode] if blend_mode != "lerp" else torch.lerp + ) + + def step(self, x): + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + ratio = sigma_next / sigma + dratio = 1 - (sigma / sigma_next) * 0.5 + + # Denoised sample at the next sigma + mr_next = self.call_model( + self.blend(ss.denoised, x, ratio), + ss.sigma_next, + call_index=1, + ) + + keep_scale, add_scale = self.get_cycle_scales(ss.sigma_next) + + denoised_prime = self.blend(mr_next.denoised, ss.denoised, dratio) + noise_pred = (x - denoised_prime).mul_(ratio * keep_scale) + + yield from self.result( + denoised_prime.add_(noise_pred), add_scale, sigma_down=sigma_next + ) + + +class BASConfig(typing.NamedTuple): + batch_multiplier: int = 2 + start_step: int = 0 + end_step: int = 3 + s_noise: float = 1.0 + eta: float = 0.0 + eta_retry_increment: float = 0 + denoised_factors: list | tuple | None = None + denoised_factors_scale: float = 1.0 + denoised_multiplier: float = 1.0 + renoise_mode: str = "restart" + fromstep_factor: float = 1.0 + tostep_factor: float = 1.0 + tostep_source: str = "dt" + + +# Batch augmented sampler +class BASStep(SingleStepSampler): + name = "blep_bas" + model_calls = -1 + uses_alt_noise = True + + def __init__(self, **kwargs): + super().__init__(**kwargs) + bas = self.bas = BASConfig(**self.options.get("bas", {})) + if bas.renoise_mode not in {"restart", "restart_noneta", "simple"}: + raise ValueError("Bad BAS renoise mode") + if bas.tostep_source not in {"dt", "sigma", "sigma_next"}: + raise ValueError("Bad BAS tostep_source") + blend_mode = self.options.get("blend_mode", "lerp").strip() + self.blend = ( + filtering.BLENDING_MODES[blend_mode] if blend_mode != "lerp" else torch.lerp + ) + + def step(self, x): + ss = self.ss + bas = self.bas + sigma, sigma_next = ss.sigma, ss.sigma_next + eta = self.get_dyn_eta() + sigma_down, sigma_up = self.get_ancestral_step(eta) + denoised = ss.denoised + ratio = sigma_down / sigma + if ( + ss.step >= bas.end_step + or ss.step < bas.start_step + or bas.batch_multiplier < 1 + ): + return (yield from self.result(self.blend(denoised, x, ratio), sigma_up)) + bsigma = sigma * bas.fromstep_factor + if bas.tostep_source == "sigma": + bsigma_next = sigma * bas.tostep_factor + elif bas.tostep_source == "sigma_next": + bsigma_next = sigma_next * bas.tostep_factor + elif bas.tostep_source == "dt": + bsigma_next = bsigma + (sigma_next - bsigma) * bas.tostep_factor + else: + raise RuntimeError("Impossible BAS tostep_source") + if bsigma <= bsigma_next: + raise ValueError("BAS: Bad configuration, got sigma <= sigma_next") + bsigma_down, bsigma_up = self.get_ancestral_step( + bas.eta, + sigma=bsigma, + sigma_next=bsigma_next, + retry_increment=bas.eta_retry_increment, + ) + bratio = bsigma_down / bsigma + x_new = self.blend(denoised, x, bratio) + batch_factor = bas.batch_multiplier + if bas.denoised_factors is None: + dn_factors = (bas.denoised_factors_scale / (batch_factor + 1),) * ( + batch_factor + 1 + ) + dn_sum = bas.denoised_factors_scale + else: + dn_factors = tuple(bas.denoised_factors)[: batch_factor + 1] + tooshort = (batch_factor + 1) - len(dn_factors) + if tooshort > 0: + dn_factors = dn_factors + (dn_factors[-1],) * tooshort + if len(dn_factors) != batch_factor + 1: + raise ValueError("Bad length for bas_denoised_factors") + dn_sum = sum(dn_factors) + if bas.denoised_factors_scale != 0: + dn_factors = tuple( + (f / dn_sum) * bas.denoised_factors_scale for f in dn_factors + ) + # print( + # f"\nBAS STEP: step {bsigma} -> {bsigma_next} : down={bsigma_down}, up={bsigma_up}, bratio={bratio}, dn_factors={dn_factors}" + # ) + if dn_sum == 0: + raise ValueError("bas_denoised_factors must sum to a non-zero quantity") + batch_size = x.shape[0] + expanded_batch = batch_size * batch_factor + x_expanded = x.new_zeros(expanded_batch, *x.shape[1:]) + renoise_mode = bas.renoise_mode + if renoise_mode == "restart": + noise_factor = (bsigma**2 - bsigma_down**2) ** 0.5 + elif renoise_mode == "restart_noneta": + noise_factor = bsigma_up + (bsigma**2 - bsigma_next**2) ** 0.5 + else: + noise_factor = bsigma_up + (bsigma - bsigma_next) + for bidx in range(batch_factor): + # print( + # f"NOISE ITER {bidx} -- {bidx * batch_size} -> {bidx * batch_size + batch_size}" + # ) + x_expanded[ + bidx * batch_size : bidx * batch_size + batch_size + ] = yield from self.result( + x_new, + noise_factor, + sigma=bsigma, + sigma_down=bsigma_down, + s_noise=bas.s_noise, + noise_sampler=self.alt_noise_sampler, + final=False, + ) + s_in = x.new_ones(expanded_batch) + del x_new + mr_expanded = self.call_model(x_expanded, bsigma, s_in=s_in, call_index=1) + denoised_expanded = mr_expanded.denoised + denoised_new = torch.zeros_like(denoised) + for obidx in range(batch_size): + for bidx in range(-1, batch_factor): + # print(f"DN ITER {obidx} <{bidx}> = dn_exp[{bidx * batch_size + obidx}]") + dn_curr = ( + denoised[obidx] + if bidx == -1 + else denoised_expanded[bidx * batch_size + obidx] + ) + denoised_new[obidx] += dn_curr * dn_factors[bidx + 1] + denoised_new *= bas.denoised_multiplier + result = self.blend(denoised_new, x, ratio) + yield from self.result(result, sigma_up) + + +def scale_wavelets(waves, factor_yl, factor_yh=None): + factor_yh = fallback(factor_yh, factor_yl) + if factor_yl == 1 and factor_yh == 1: + return waves + return (waves[0] * factor_yl, tuple(t * factor_yh for t in waves[1])) + + +def blend_wavelets(a, b, *, factor_yl, factor_yh, blend_yl, blend_yh=None): + blend_yh = fallback(blend_yh, blend_yl) + if not isinstance(factor_yl, torch.Tensor): + factor_yl = a[0].new_full((1,), factor_yl) + if not isinstance(factor_yh, torch.Tensor): + factor_yh = a[0].new_full((1,), factor_yh) + return ( + blend_yl(a[0], b[0], factor_yl), + tuple(blend_yh(ta, tb, factor_yh) for ta, tb in zip(a[1], b[1])), + ) + + +class WeoonConfig(typing.NamedTuple): + start_step: int = 0 + end_step: int = 9999 + eta: float = 0.0 + eta_retry_increment: float = 0.0 + s_noise: float = 1.0 + # One of dwt, dwt1d, dtcwt + wavelet_mode: str = "dwt" + padding: str = "periodization" + inv_padding: str | None = None + level: int = 3 + wave: str = "db4" + inv_wave: str | None = None + dtcwt_qshift: str = "qshift_a" + dtcwt_biort: str = "near_sym_a" + dtcwt_inv_qshift: str | None = None + dtcwt_inv_biort: str | None = None + downstep_scale: float = 1.0 + yl_strength: float = 1.0 + yh_strength: float = 0.5 + wavelet_blend_mode: str = "lerp" + wavelet_blend_mode_yh: str | None = None + denoised_yl_multiplier: float = 1.0 + denoised_yh_multiplier: float = 1.0 + denoised_down_yl_multiplier: float = 1.0 + denoised_down_yh_multiplier: float = 1.0 + flatten_start_dim: int = 2 + + +class WeoonStep(SingleStepSampler): + name = "blep_weoon" + model_calls = 1 + uses_alt_noise = True + + def __init__(self, **kwargs): + if not HAVE_WAVELETS: + raise RuntimeError( + "Wavelet sampling requires the pytorch_wavelets package installed in your environment", + ) + super().__init__(**kwargs) + w = self.weoon = WeoonConfig(**self.options.get("weoon", {})) + blend_mode = self.options.get("blend_mode", "lerp").strip() + self.blend = ( + filtering.BLENDING_MODES[blend_mode] if blend_mode != "lerp" else torch.lerp + ) + self.wavelet_blend = ( + filtering.BLENDING_MODES[w.wavelet_blend_mode] + if w.wavelet_blend_mode != "lerp" + else torch.lerp + ) + if w.wavelet_blend_mode_yh is None: + self.wavelet_blend_yh = self.wavelet_blend + else: + self.wavelet_blend_yh = ( + filtering.BLENDING_MODES[w.wavelet_blend_mode_yh] + if w.wavelet_blend_mode_yh != "lerp" + else torch.lerp + ) + if not (0 <= w.flatten_start_dim <= 2): + raise ValueError("Bad flatten_start_dim in Weoon sampler") + if w.wavelet_mode == "dtcwt": + self.wavelet_forward = ptwav.DTCWTForward( + J=w.level, mode=w.padding, biort=w.dtcwt_biort, qshift=w.dtcwt_qshift + ) + self.wavelet_inverse = ptwav.DTCWTInverse( + mode=fallback(w.inv_padding, w.padding), + biort=fallback(w.dtcwt_inv_biort, w.dtcwt_biort), + qshift=fallback(w.dtcwt_inv_qshift, w.dtcwt_qshift), + ) + elif w.wavelet_mode == "dwt": + self.wavelet_forward = ptwav.DWTForward( + J=w.level, wave=w.wave, mode=w.padding + ) + self.wavelet_inverse = ptwav.DWTInverse( + wave=fallback(w.inv_wave, w.wave), + mode=fallback(w.inv_padding, w.padding), + ) + elif w.wavelet_mode == "dwt1d": + self.wavelet_forward = ptwav.DWT1DForward( + J=w.level, wave=w.wave, mode=w.padding + ) + self.wavelet_inverse = ptwav.DWT1DInverse( + wave=fallback(w.inv_wave, w.wave), + mode=fallback(w.inv_padding, w.padding), + ) + + def maybe_flatten(self, tensor: torch.Tensor) -> torch.Tensor: + w = self.weoon + need_flatten = w.wavelet_mode == "dwt1d" + if not need_flatten: + return tensor + start_dim = w.flatten_start_dim + tensor = tensor.flatten(start_dim=start_dim) + if start_dim == 0: + return tensor[None, None, ...] + if start_dim == 1: + return tensor[:, None, ...] + return tensor + + def step(self, x): + w = self.weoon + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + ratio = sigma_down / sigma + if not w.start_step <= ss.step <= w.end_step: + return (yield from self.result(self.blend(ss.denoised, x, ratio), sigma_up)) + self.wavelet_forward.to(x) + self.wavelet_inverse.to(x) + dt = sigma_next - sigma + wsigma_next = (sigma + dt * w.downstep_scale).clamp_(0) + wsigma_down, wsigma_up = self.get_ancestral_step( + w.eta, + sigma=sigma, + sigma_next=wsigma_next, + retry_increment=w.eta_retry_increment, + ) + wratio = wsigma_down / sigma + x_down = self.blend(ss.denoised, x, wratio) + if wsigma_up != 0: + x_down = yield from self.result( + x_down, + wsigma_up, + sigma_next=wsigma_next, + sigma_down=wsigma_down, + s_noise=w.s_noise, + noise_sampler=self.alt_noise_sampler, + final=False, + ) + mr_down = self.call_model(x_down, wsigma_next, call_index=1) + coeffs = scale_wavelets( + self.wavelet_forward(self.maybe_flatten(ss.denoised)), + factor_yl=w.denoised_yl_multiplier, + factor_yh=w.denoised_yh_multiplier, + ) + coeffs_down = scale_wavelets( + self.wavelet_forward(self.maybe_flatten(mr_down.denoised)), + factor_yl=w.denoised_down_yl_multiplier, + factor_yh=w.denoised_down_yh_multiplier, + ) + coeffs_out = blend_wavelets( + coeffs, + coeffs_down, + factor_yl=w.yl_strength, + factor_yh=w.yh_strength, + blend_yl=self.wavelet_blend, + blend_yh=self.wavelet_blend_yh, + ) + denoised_new = self.wavelet_inverse(coeffs_out) + if denoised_new.shape != x.shape: + denoised_new = denoised_new.reshape(*x.shape) + x = self.blend(denoised_new, x, ratio) + yield from self.result(x, sigma_up) + + +registry.add( + BASStep, + DynamicStep, + AdapterStep, + EulerCycleStep, + TrapezoidalCycleStep, + WeoonStep, +) diff --git a/py/step_samplers/builtins.py b/py/step_samplers/builtins.py new file mode 100644 index 0000000..3ee11fa --- /dev/null +++ b/py/step_samplers/builtins.py @@ -0,0 +1,541 @@ +import torch + +import comfy +from comfy.k_diffusion.sampling import get_ancestral_step + +from .base import ( + SingleStepSampler, + DPMPPStepMixin, + HistorySingleStepSampler, + ReversibleSingleStepSampler, + registry, +) + + +class EulerStep(SingleStepSampler): + name = "euler" + allow_cfgpp = True + allow_alt_cfgpp = True + step = SingleStepSampler.euler_step + + +class DPMPP2MStep(HistorySingleStepSampler, DPMPPStepMixin): + name = "dpmpp_2m" + default_history_limit, max_history = 1, 1 + ancestralize = True + default_eta = 0.0 + + def step(self, x): + ss = self.ss + s, sn = ss.sigma, ss.sigma_next + t, t_next = self.t_fn(s), self.t_fn(sn) + h = t_next - t + st, st_next = self.sigma_fn(t), self.sigma_fn(t_next) + if self.available_history() > 0: + h_last = t - self.t_fn(ss.sigma_prev) + r = h_last / h + denoised, old_denoised = ss.denoised, ss.hprev.denoised + denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised + else: + denoised_d = ss.denoised + yield from self.result((st_next / st) * x - (-h).expm1() * denoised_d) + + +class DPMPP2MSDEStep(ReversibleSingleStepSampler): + name = "dpmpp_2m_sde" + default_history_limit, max_history = 1, 1 + default_reversible_scale = 0.0 + + def __init__(self, *, solver_type="midpoint", **kwargs): + super().__init__(**kwargs) + solver_type = solver_type.lower().strip() + if solver_type not in ("midpoint", "heun"): + raise ValueError("Bad solver_type: must be one of midpoint, heun") + self.solver_type = solver_type + + def step(self, x): + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + denoised = ss.denoised + # DPM-Solver++(2M) SDE + t, s = -sigma.log(), -sigma_next.log() + h = s - t + eta_h = self.get_dyn_eta() * h + ratio = sigma_next / sigma + x = ((ratio * (-eta_h).exp()) * x).add_((-h - eta_h).expm1().neg() * denoised) + noise_strength = sigma_next * (-2 * eta_h).expm1().neg().sqrt() + if self.available_history() == 0: + return (yield from self.result(x, noise_strength)) + sigma_prev, old_denoised = ss.hprev.sigma, ss.hprev.denoised + h_last = (-sigma.log()) - (-sigma_prev.log()) + r = h_last / h + if self.solver_type == "midpoint": + multiplier = 0.5 * (-h - eta_h).expm1().neg() + else: + multiplier = (-h - eta_h).expm1().neg() / (-h - eta_h) + 1 + reta, reversible_scale = self.get_reversible_cfg() + if reversible_scale != 0: + multiplier *= 0.5 + x += (denoised - old_denoised).mul_((1 / r) * multiplier) + if reversible_scale != 0: + reta_h = reta * h + if self.solver_type == "midpoint": + rmultiplier = 0.5 * (-h - reta_h).expm1().neg() + else: + rmultiplier = (-h - reta_h).expm1().neg() / (-h - reta_h) + 1 + rmultiplier = ((1 / r) * (rmultiplier**2 / 2)) * reversible_scale + x -= (old_denoised - denoised).mul_(rmultiplier) + yield from self.result(x, noise_strength) + + +class DPMPP3MSDEStep(HistorySingleStepSampler): + name = "dpmpp_3m_sde" + default_history_limit, max_history = 2, 2 + + def step(self, x): + ss = self.ss + denoised = ss.denoised + t, s = -ss.sigma.log(), -ss.sigma_next.log() + h = s - t + eta = self.get_dyn_eta() + h_eta = h * (eta + 1) + x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised + noise_strength = ss.sigma_next * (-2 * h * eta).expm1().neg().sqrt() + ah = self.available_history() + if ah == 0: + return (yield from self.result(x, noise_strength)) + hist = ss.hist + h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log()) + denoised_1 = hist[-2].denoised + if ah == 1: + r = h_1 / h + d = (denoised - denoised_1) / r + phi_2 = h_eta.neg().expm1() / h_eta + 1 + x = x + phi_2 * d + else: # 2+ history items available + h_2 = (-ss.sigma_prev.log()) - (-ss.sigmas[ss.idx - 2].log()) + denoised_2 = hist[-3].denoised + r0 = h_1 / h + r1 = h_2 / h + d1_0 = (denoised - denoised_1) / r0 + d1_1 = (denoised_1 - denoised_2) / r1 + d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) + d2 = (d1_0 - d1_1) / (r0 + r1) + phi_2 = h_eta.neg().expm1() / h_eta + 1 + phi_3 = phi_2 / h_eta - 0.5 + x = x + phi_2 * d1 - phi_3 * d2 + yield from self.result(x, noise_strength) + + def step_(self, x): + sr = next(super().step(x)) + if self.available_history() < 2: + yield sr + return + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + hprev = ss.hist[-2] + hprevprev = ss.hist[-3] + t, s = -sigma.log(), -sigma_next.log() + h = s - t + eta = self.get_dyn_eta() + h_eta = h * (eta + 1) + h_2 = (-ss.sigma_prev.log()) - (-hprevprev.sigma.log()) + denoised = ss.denoised + denoised_1 = hprev.denoised + denoised_2 = hprevprev.denoised + h_1 = (-ss.sigma.log()) - (-hprev.sigma.log()) + r0 = h_1 / h + r1 = h_2 / h + d1_0 = (denoised - denoised_1).div_(r0) + d1_1 = (denoised_1 - denoised_2).div_(r1) + d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) + d2 = (d1_0 - d1_1).div_(r0 + r1) + phi_2 = h_eta.neg().expm1() / h_eta + 1 + phi_3 = phi_2 / h_eta - 0.5 + sr.x_ += phi_2 * d1 - phi_3 * d2 + yield sr + # x = x + phi_2 * d1 - phi_3 * d2 + # yield from self.result(sr.x + phi_2 * d1 - phi_3 * d2, sr.sigma_up) + + +# Alt CFG++ approach referenced from https://github.com/comfyanonymous/ComfyUI/pull/3871 - thanks! +class DPMPP2SStep(SingleStepSampler, DPMPPStepMixin): + name = "dpmpp_2s" + model_calls = 1 + allow_alt_cfgpp = True + + def step(self, x): + ss = self.ss + t_fn, sigma_fn = self.t_fn, self.sigma_fn + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + # DPM-Solver++(2S) + t, t_next = t_fn(ss.sigma), t_fn(sigma_down) + r = 1 / 2 + h = t_next - t + s = t + r * h + eff_x = ( + x + if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None + else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale + ) + x_2 = (sigma_fn(s) / sigma_fn(t)) * eff_x - (-h * r).expm1() * ss.denoised + denoised_2 = self.call_model(x_2, sigma_fn(s), call_index=1).denoised + x = (sigma_fn(t_next) / sigma_fn(t)) * eff_x - (-h).expm1() * denoised_2 + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin): + name = "dpmpp_sde" + self_noise = 1 + model_calls = 1 + allow_alt_cfgpp = True # Implementation may not be correct. + uses_alt_noise = True + + def __init__(self, *args, r=1 / 2, **kwargs): + super().__init__(*args, **kwargs) + self.r = r + + def step(self, x): + ss = self.ss + t_fn, sigma_fn = self.t_fn, self.sigma_fn + r, eta = self.r, self.get_dyn_eta() + # DPM-Solver++ + t, t_next = t_fn(ss.sigma), t_fn(ss.sigma_next) + h = t_next - t + s = t + h * r + fac = 1 / (2 * r) + + # Step 1 + sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta) + s_ = t_fn(sd) + eff_x = ( + x + if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None + else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale + ) + x_2 = (sigma_fn(s_) / sigma_fn(t)) * eff_x - (t - s_).expm1() * ss.denoised + x_2 = yield from self.result( + x_2, + su, + sigma=sigma_fn(t), + sigma_next=sigma_fn(s), + noise_sampler=self.alt_noise_sampler, + final=False, + ) + denoised_2 = self.call_model(x_2, sigma_fn(s), call_index=1).denoised + + # Step 2 + sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta) + t_next_ = t_fn(sd) + denoised_d = (1 - fac) * ss.denoised + fac * denoised_2 + x = (sigma_fn(t_next_) / sigma_fn(t)) * eff_x - ( + t - t_next_ + ).expm1() * denoised_d + yield from self.result(x, su, sigma_down=sd) + + +# Adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py +# under Apache 2 license +class IPNDMStep(HistorySingleStepSampler): + name = "ipndm" + ancestralize = True + default_history_limit, max_history = 1, 3 + allow_alt_cfgpp = True + default_eta = 0.0 + + IPNDM_MULTIPLIERS = ( + ((1,), 1), + ((3, -1), 2), + ((23, -16, 5), 12), + ((55, -59, 37, -9), 24), + ) + + def step(self, x): + ss = self.ss + order = self.available_history() + 1 + if order > 1: + hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) + (dm, *hms), divisor = self.IPNDM_MULTIPLIERS[order - 1] + noise = dm * self.to_d(ss.hcur) + for hidx, hm in enumerate(hms, start=1): + noise += hm * hd[-hidx] + noise /= divisor + yield from self.result(x + ss.dt * noise) + + +# Adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py +# under Apache 2 license +class IPNDMVStep(HistorySingleStepSampler): + name = "ipndm_v" + ancestralize = True + default_history_limit, max_history = 1, 3 + allow_alt_cfgpp = True + default_eta = 0.0 + + def step(self, x): + ss = self.ss + dt = ss.dt + d = self.to_d(ss.hcur) + order = self.available_history() + 1 + if order > 1: + hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) + hns = ( + ss.sigmas[ss.idx - (order - 2) : ss.idx + 1] + - ss.sigmas[ss.idx - (order - 1) : ss.idx] + ) + if order == 1: + noise = d + elif order == 2: + coeff1 = (2 + (dt / hns[-1])) / 2 + coeff2 = -(dt / hns[-1]) / 2 + noise = coeff1 * d + coeff2 * hd[-1] + elif order == 3: + temp = ( + 1 + - dt + / (3 * (dt + hns[-1])) + * (dt * (dt + hns[-1])) + / (hns[-1] * (hns[-1] + hns[-2])) + ) / 2 + coeff1 = (2 + (dt / hns[-1])) / 2 + temp + coeff2 = -(dt / hns[-1]) / 2 - (1 + hns[-1] / hns[-2]) * temp + coeff3 = temp * hns[-1] / hns[-2] + noise = coeff1 * d + coeff2 * hd[-1] + coeff3 * hd[-2] + else: + temp1 = ( + 1 + - dt + / (3 * (dt + hns[-1])) + * (dt * (dt + hns[-1])) + / (hns[-1] * (hns[-1] + hns[-2])) + ) / 2 + temp2 = ( + ( + (1 - dt / (3 * (dt + hns[-1]))) / 2 + + (1 - dt / (2 * (dt + hns[-1]))) + * dt + / (6 * (dt + hns[-1] + hns[-2])) + ) + * (dt * (dt + hns[-1]) * (dt + hns[-1] + hns[-2])) + / (hns[-1] * (hns[-1] + hns[-2]) * (hns[-1] + hns[-2] + hns[-3])) + ) + coeff1 = (2 + (dt / hns[-1])) / 2 + temp1 + temp2 + coeff2 = ( + -(dt / hns[-1]) / 2 + - (1 + hns[-1] / hns[-2]) * temp1 + - ( + 1 + + (hns[-1] / hns[-2]) + + (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) + ) + * temp2 + ) + coeff3 = ( + temp1 * hns[-1] / hns[-2] + + ( + (hns[-1] / hns[-2]) + + (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) + * (1 + hns[-2] / hns[-3]) + ) + * temp2 + ) + coeff4 = ( + -temp2 + * (hns[-1] * (hns[-1] + hns[-2]) / (hns[-2] * (hns[-2] + hns[-3]))) + * hns[-1] + / hns[-2] + ) + noise = coeff1 * d + coeff2 * hd[-1] + coeff3 * hd[-2] + coeff4 * hd[-3] + yield from self.result(x + ss.dt * noise) + + +class DEISStep(HistorySingleStepSampler): + name = "deis" + ancestralize = True + default_history_limit, max_history = 1, 3 + allow_alt_cfgpp = True + default_eta = 0.0 + + def __init__(self, *args, deis_mode="tab", **kwargs): + super().__init__(*args, **kwargs) + self.deis_mode = deis_mode + self.deis_coeffs_key = None + self.deis_coeffs = None + + def get_deis_coeffs(self): + ss = self.ss + key = ( + self.history_limit, + len(ss.sigmas), + ss.sigmas[0].item(), + ss.sigmas[-1].item(), + ) + if self.deis_coeffs_key == key: + return self.deis_coeffs + self.deis_coeffs_key = key + self.deis_coeffs = comfy.k_diffusion.deis.get_deis_coeff_list( + ss.sigmas, self.history_limit + 1, deis_mode=self.deis_mode + ) + return self.deis_coeffs + + def step(self, x): + ss = self.ss + dt = ss.dt + d = self.to_d(ss.hcur) + order = self.available_history() + 1 + if order < 2: + noise = dt * d # Euler + else: + c = self.get_deis_coeffs()[ss.idx] + hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) + noise = c[0] * d + for i in range(1, order): + noise += c[i] * hd[-i] + yield from self.result(x + noise) + + +# https://openreview.net/pdf?id=o2ND9v0CeK +# Implementation referenced from ComfyUI +class GradientEstimationStep(HistorySingleStepSampler): + name = "gradient_estimation" + ancestralize = False + default_history_limit, max_history = 1, 1 + default_eta = 0.0 + + def __init__(self, *args, ge_gamma=2.0, **kwargs): + super().__init__(*args, **kwargs) + self.ge_gamma = ge_gamma + + def step(self, x): + ss = self.ss + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + dt = sigma_down - ss.sigma + d = self.to_d(ss.hcur) + if self.available_history() < 1: + noise_pred = dt * d # Euler + else: + gamma = self.ge_gamma + noise_pred = dt * (gamma * d + (1 - gamma) * ss.hist[-2].d) + yield from self.result(x + noise_pred, sigma_up, sigma_down=sigma_down) + + +class HeunPP2Step(SingleStepSampler): + name = "heunpp2" + ancestralize = True + model_calls = 2 + allow_alt_cfgpp = True + + def __init__(self, *args, max_order=3, **kwargs): + super().__init__(*args, **kwargs) + self.max_order = max(1, min(self.model_calls + 1, max_order)) + + def step(self, x): + ss = self.ss + steps_remain = max(0, len(ss.sigmas) - (ss.idx + 2)) + order = min(self.max_order, steps_remain + 1) + sn = ss.sigma_next + if order == 1: + return (yield from self.euler_step(x)) + d = self.to_d(ss.hcur) + dt = ss.dt + w = order * ss.sigma + w2 = sn / w + x_2 = x + d * dt + d_2 = self.to_d(self.call_model(x_2, sn, call_index=1)) + if order == 2: + # Heun's method (ish) + w1 = 1 - w2 + d_prime = d * w1 + d_2 * w2 + else: + # Heun++ (ish) + snn = ss.sigmas[ss.idx + 2] + dt_2 = snn - sn + x_3 = x_2 + d_2 * dt_2 + d_3 = self.to_d(self.call_model(x_3, snn, call_index=2)) + w3 = snn / w + w1 = 1 - w2 - w3 + d_prime = w1 * d + w2 * d_2 + w3 * d_3 + yield from self.result(x + d_prime * dt) + + +# Referenced from ComfyUI implementation +class DPM2Step(SingleStepSampler): + name = "dpm_2" + model_calls = 1 + + def step(self, x): + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + eta = self.get_dyn_eta() + sigma_down, sigma_up = self.get_ancestral_step(eta) + sigma_mid = sigma.log().lerp(sigma_next.log(), 0.5).exp() + dt_1, dt_2 = sigma_mid - sigma, sigma_down - sigma + d = self.to_d(ss.hcur) + mr_2 = self.call_model(x + d * dt_1, sigma_mid, call_index=1) + d_2 = self.to_d(mr_2) + yield from self.result(x + d_2 * dt_2, sigma_up) + + +# Referenced from ComfyUI implementation +class RESMultistepStep(HistorySingleStepSampler, DPMPPStepMixin): + name = "res_multistep" + default_history_limit, max_history = 1, 1 + default_eta = 0.0 + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + @staticmethod + def phi1_fn(t: torch.Tensor) -> torch.Tensor: + return t.expm1() / t + + @classmethod + def phi2_fn(cls, t: torch.Tensor) -> torch.Tensor: + return (cls.phi1_fn(t) - 1.0) / t + + def step(self, x): + ss = self.ss + sigma = ss.sigma + eta = self.get_dyn_eta() + sigma_down, sigma_up = self.get_ancestral_step(eta) + if self.available_history() == 0: + dt = sigma_down - sigma + d = self.to_d(ss.hcur) + return (yield from self.result(x + dt * d, sigma_up, sigma_down=sigma_down)) + prev_mr = ss.hist[-2] + prev_sigma_down = self.get_ancestral_step( + sigma=prev_mr.sigma, sigma_next=sigma, eta=eta + )[0] + # Second order multistep method in https://arxiv.org/pdf/2308.02157 + t, t_old, t_next, t_prev = ( + self.t_fn(sigma), + self.t_fn(prev_sigma_down), + self.t_fn(sigma_down), + self.t_fn(prev_mr.sigma), + ) + h = t_next - t + h_s = self.sigma_fn(h) + c2 = (t_prev - t_old) / h + + phi1_val, phi2_val = self.phi1_fn(-h), self.phi2_fn(-h) + b1 = torch.nan_to_num(phi1_val - phi2_val / c2, nan=0.0) + b2 = torch.nan_to_num(phi2_val / c2, nan=0.0) + result = h_s * x + h * (b1 * ss.denoised + b2 * prev_mr.denoised) + yield from self.result(result, sigma_up, sigma_down=sigma_down) + + +registry.add( + DEISStep, + DPMPP2MSDEStep, + DPMPP2MStep, + DPMPP3MSDEStep, + DPMPPSDEStep, + EulerStep, + HeunPP2Step, + IPNDMStep, + IPNDMVStep, + GradientEstimationStep, + DPM2Step, + DPMPP2SStep, + RESMultistepStep, +) diff --git a/py/step_samplers/clybius.py b/py/step_samplers/clybius.py new file mode 100644 index 0000000..5fa5816 --- /dev/null +++ b/py/step_samplers/clybius.py @@ -0,0 +1,755 @@ +# Samplers based on Clybius' designs, mostly from https://github.com/Clybius/ComfyUI-Extra-Samplers/ + +import math +import torch + +from comfy.k_diffusion.sampling import get_ancestral_step, to_d + +from .base import SingleStepSampler, ReversibleConfig, ReversibleSingleStepSampler +from .builtins import DPMPP2MSDEStep +from . import res_support +from . import registry + +from .. import filtering +from .. import utils + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +# Apparently the only difference between Heun and Trapezoidal is the first step using ETA or not. +class ReversibleHeunStep(ReversibleSingleStepSampler): + name = "reversible_heun" + model_calls = 1 + allow_alt_cfgpp = True + allow_cfgpp = True + trapezoidal_mode = False + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + blend_mode = self.options.get("blend_mode", "lerp").strip() + self.blend = ( + filtering.BLENDING_MODES[blend_mode] if blend_mode != "lerp" else torch.lerp + ) + + def reversible_correction(self, d, d_next, dt_reversible): + if dt_reversible == 0 or self.reversible.scale == 0: + return None + return d_next.sub_(d).div_(4).mul_(dt_reversible**2).mul_(self.reversible.scale) + + def step_internal(self, x, *, history_mode=False): + ss = self.ss + history_mode = history_mode and self.available_history() > 0 + if not history_mode: + mr_1, mr_2 = ss.hcur, None + else: + mr_1, mr_2 = ss.hprev, ss.hcur + sigma = ss.sigma + denoised, uncond = mr_1.denoised, mr_1.denoised_uncond + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + ratio = sigma_down / sigma + dratio = 1 - (sigma / sigma_down) * 0.5 + if mr_2 is None: + x_2 = self.step_mix(x, denoised, uncond, ratio, blend=self.blend) + mr_2 = self.call_model(x_2, sigma_down, call_index=1) + del x_2 + denoised_2, uncond_2 = mr_2.denoised, mr_2.denoised_uncond + denoised_prime = self.blend(denoised_2, denoised, dratio) + if self.cfgpp: + denoised_prime += denoised * 0.5 + uncond_prime = (uncond_2 * (1 - dratio)).add_(uncond) + elif self.alt_cfgpp_scale != 0: + uncond_prime = self.blend(uncond_2, uncond, dratio) + else: + uncond_prime = uncond + x = self.step_mix(x, denoised_prime, uncond_prime, ratio, blend=self.blend) + if self.reversible.scale != 0: + correction = self.reversible_correction( + d=self.to_d(mr_1, use_cfgpp=self.reversible.use_cfgpp), + d_next=self.to_d(mr_2, use_cfgpp=self.reversible.use_cfgpp), + dt_reversible=self.get_ancestral_step(self.dyn_reta)[0] - sigma, + ) + if correction is not None: + x -= correction + yield from self.result(x, sigma_up) + + def step(self, x): + return self.step_internal(x, history_mode=False) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class ReversibleHeun1SStep(ReversibleHeunStep): + name = "reversible_heun_1s" + model_calls = (0, 1) + default_history_limit, max_history = 1, 1 + allow_alt_cfgpp = True + allow_cfgpp = True + + def step(self, x): + return self.step_internal(x, history_mode=True) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RESStep(SingleStepSampler): + name = "res" + model_calls = 1 + allow_alt_cfgpp = True # May not be implemented correctly. + + def __init__(self, *, res_simple_phi=False, res_c2=0.5, **kwargs): + super().__init__(**kwargs) + self.simple_phi = res_simple_phi + self.c2 = res_c2 + + def step(self, x): + ss = self.ss + eta = self.get_dyn_eta() + sigma_down, sigma_up = self.get_ancestral_step(eta) + denoised = ss.denoised + lam_next = sigma_down.log().neg() if eta != 0 else ss.sigma_next.log().neg() + lam = ss.sigma.log().neg() + + h = lam_next - lam + a2_1, b1, b2 = res_support._de_second_order( + h=h, c2=self.c2, simple_phi_calc=self.simple_phi + ) + + c2_h = 0.5 * h + + eff_x = ( + x + if self.alt_cfgpp_scale == 0 or ss.hcur.denoised_uncond is None + else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale + ) + x_2 = math.exp(-c2_h) * eff_x + a2_1 * h * denoised + lam_2 = lam + c2_h + sigma_2 = lam_2.neg().exp() + + denoised2 = self.call_model(x_2, sigma_2, call_index=1).denoised + + x = math.exp(-h) * eff_x + h * (b1 * denoised + b2 * denoised2) + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +class TrapezoidalStep(ReversibleHeunStep): + reversible = False + trapezoidal_mode = True + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class TrapezoidalStep_(SingleStepSampler): + name = "trapezoidal" + model_calls = 1 + allow_alt_cfgpp = True + + def step(self, x): + ss = self.ss + sigma_next = ss.sigma_next + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + + # Predict the sample at the next sigma using Euler step + euler_sr = next(self.euler_step(x, sigma_down=sigma_next, sigma_up=0)) + d = euler_sr.noise_pred + + # Denoised sample at the next sigma + mr_next = self.call_model(euler_sr.x, euler_sr.sigma_down, call_index=1) + + denoised_pred_next, d_next = self.get_split_prediction(mr=mr_next) + yield from self.split_result( + denoised_pred_next, + (d + d_next) * 0.5, + sigma_up=sigma_up, + sigma_down=sigma_down, + ) + + def step_(self, x): + ss = self.ss + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + + # Calculate the derivative using the model + d_i = self.to_d(ss.hcur) + + # Predict the sample at the next sigma using Euler step + x_pred = x + d_i * ss.dt + + # Denoised sample at the next sigma + mr_next = self.call_model(x_pred, ss.sigma_next, call_index=1) + + # Calculate the derivative at the next sigma + d_next = self.to_d(mr_next) + dt_2 = sigma_down - ss.sigma + + # Update the sample using the Trapezoidal rule + x = x + dt_2 * (d_i + d_next) / 2 + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class BogackiStep(ReversibleSingleStepSampler): + name = "bogacki" + reversible = False + model_calls = 2 + allow_alt_cfgpp = True + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if not self.reversible: + self.reversible.scale = 0 + + def step(self, x): + ss = self.ss + s = ss.sigma + sd, su = self.get_ancestral_step(self.get_dyn_eta()) + reta, reversible_scale = self.get_reversible_cfg() + sdr, _sur = self.get_ancestral_step(reta) + dt, dtr = sd - s, sdr - s + + # Calculate the derivative using the model + d = self.to_d(ss.hcur) + + # Bogacki-Shampine steps + k1 = d * dt + k2 = self.to_d(self.call_model(x + k1 / 2, s + dt / 2, call_index=1)) * dt + k3 = ( + self.to_d( + self.call_model(x + 3 * k1 / 4 + k2 / 4, s + 3 * dt / 4, call_index=2) + ) + * dt + ) + + # Reversible correction term (inspired by Reversible Heun) + correction = dtr**2 * (k3 - k2) / 6 + + # Update the sample + x = (x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9) - correction * reversible_scale + yield from self.result(x, su, sigma_down=sd) + + +class ReversibleBogackiStep(BogackiStep): + name = "reversible_bogacki" + reversible = True + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RK4Step(SingleStepSampler): + name = "rk4" + model_calls = 3 + allow_alt_cfgpp = True + + def step(self, x): + ss = self.ss + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + sigma = ss.sigma + d = self.to_d(ss.hcur) + dt = sigma_down - sigma + + # Runge-Kutta steps + k1 = d * dt + k2 = self.to_d(self.call_model(x + k1 / 2, sigma + dt / 2, call_index=1)) * dt + k3 = self.to_d(self.call_model(x + k2 / 2, sigma + dt / 2, call_index=2)) * dt + k4 = self.to_d(self.call_model(x + k3, sigma + dt, call_index=3)) * dt + + # Update the sample + x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6 + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RKF45Step(SingleStepSampler): + name = "rkf45" + model_calls = 5 + allow_alt_cfgpp = True + + def step(self, x): + ss = self.ss + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + sigma = ss.sigma + d = self.to_d(ss.hcur) + dt = sigma_down - sigma + + # Runge-Kutta steps + sigma_progression = ( + sigma + dt / 4, + sigma + 3 * dt / 8, + sigma + 12 * dt / 13, + sigma + dt, + ) + + call_progression = ( + lambda k1: x + k1 / 4, + lambda k1, k2: x + 3 * k1 / 32 + 9 * k2 / 32, + lambda k1, k2, k3: x + + 1932 * k1 / 2197 + - 7200 * k2 / 2197 + + 7296 * k3 / 2197, + lambda k1, k2, k3, k4: x + + 439 * k1 / 216 + - 8 * k2 + + 3680 * k3 / 513 + - 845 * k4 / 4104, + ) + + k = [d * dt] + for idx, (ksigma, kfun) in enumerate(zip(sigma_progression, call_progression)): + curr_x = kfun(*k) + k.append(self.to_d(self.call_model(curr_x, ksigma)) * dt) + del curr_x + x = x + 25 * k[0] / 216 + 1408 * k[2] / 2565 + 2197 * k[3] / 4104 - k[4] / 5 + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RKDynamicStep(SingleStepSampler): + name = "rk_dynamic" + model_calls = (0, 3) + allow_alt_cfgpp = True + + rk_weights = ( + (1,), + (0.5, 0.5), + (1 / 6, 2 / 3, 1 / 6), + (1 / 8, 3 / 8, 3 / 8, 1 / 8), + ) + + rk_error_orders = ((0.0375, 4), (0.075, 3), (0.15, 2)) + + def __init__(self, *args, max_order=4, **kwargs): + super().__init__(*args, **kwargs) + self.max_order = max(0, min(max_order, 4)) + + def get_rk_error_order(self, error): + for threshold, order in self.rk_error_orders: + if error < threshold: + return order + return 1 + + def step(self, x): + ss = self.ss + order = self.max_order + + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + sigma = ss.sigma + d = self.to_d(ss.hcur) + dt = sigma_down - sigma + + error = ss.hcur.get_error(ss.hprev) if len(ss.hist) > 1 else 0.0 + if order < 1: + order = self.get_rk_error_order(error) + + k = [d * dt] + curr_weight = self.rk_weights[order - 1] + + # print( + # f"\nRK: weight={curr_weight!r}, histlen={len(ss.hist)}, order={order} ({self.max_order}), err={error:.6}\n" + # ) + for j in range(1, order): + # Calculate intermediate k values based on the current order + k_sum = sum(curr_weight[i] * k[i] for i in range(j)) + mr = self.call_model(x + k_sum, sigma + dt * sum(curr_weight[:j])) + k.append(self.to_d(mr) * dt) + del mr + + # Update the sample using the weighted sum of k values + x = x + sum(curr_weight[j] * k[j] for j in range(order)) + + yield from self.result(x, sigma_up, sigma_down=sigma_down) + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class EulerDancingStep(SingleStepSampler): + name = "clybius_euler_dancing" + self_noise = 1 + + def __init__( + self, + *, + deta=1.0, + ds_noise=None, + leap=2, + dyn_deta_start=None, + dyn_deta_end=None, + dyn_deta_mode="lerp", + **kwargs, + ): + super().__init__(**kwargs) + self.deta = deta + self.ds_noise = ds_noise if ds_noise is not None else self.s_noise + self.leap = leap + self.dyn_deta_start = dyn_deta_start + self.dyn_deta_end = dyn_deta_end + if dyn_deta_mode not in ("lerp", "lerp_alt", "deta"): + raise ValueError("Bad dyn_deta_mode") + self.dyn_deta_mode = dyn_deta_mode + + def step(self, x): + ss = self.ss + eta = self.eta + deta = self.deta + leap_sigmas = ss.sigmas[ss.idx :] + leap_sigmas = leap_sigmas[: utils.find_first_unsorted(leap_sigmas)] + zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] + max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 + is_danceable = max_leap > 1 and ss.sigma_next != 0 + curr_leap = max(1, min(self.leap, max_leap)) + sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next + del leap_sigmas + sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) + print("???", sigma_down, sigma_up) + d = to_d(x, ss.sigma, ss.denoised) + # Euler method + dt = sigma_down - ss.sigma + x = x + d * dt + if curr_leap == 1: + return (yield from self.result(x, sigma_up)) + noise_strength = self.ds_noise * sigma_up + if noise_strength != 0: + x = yield from self.result(x, sigma_up, sigma_next=sigma_leap, final=False) + + # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_( + # self.ds_noise * sigma_up + # ) + # sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) + # _sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) + # sigma_up2 = ss.sigma_next + (ss.sigma - ss.sigma_next) * 0.5 + sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta)[1] + ( + ss.sigma_next * 0.5 + ) + sigma_down2, _sigma_up2 = get_ancestral_step( + ss.sigma_next, sigma_leap, eta=deta + ) + print(">>>", sigma_down2, sigma_up2, "--", ss.sigma, "->", sigma_leap) + # sigma_down2, sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta) + d_2 = to_d(x, sigma_leap, ss.denoised) + dt_2 = sigma_down2 - sigma_leap + x = x + d_2 * dt_2 + yield from self.result(x, sigma_up2, sigma_down=sigma_down2) + + # def _step(self, x, ss): + # eta = self.get_dyn_eta(ss) + # leap_sigmas = ss.sigmas[ss.idx :] + # leap_sigmas = leap_sigmas[: utils.find_first_unsorted(leap_sigmas)] + # zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] + # max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 + # is_danceable = max_leap > 1 and ss.sigma_next != 0 + # curr_leap = max(1, min(self.leap, max_leap)) + # sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next + # # DANCE 35 6 tensor(10.0947, device='cuda:0') -- tensor([21.9220, + # # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) + # del leap_sigmas + # sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) + # d = to_d(x, ss.sigma, ss.denoised) + # # Euler method + # dt = sigma_down - ss.sigma + # x = x + d * dt + # if curr_leap == 1: + # return x, sigma_up + # dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) + # if curr_leap == 1 or not is_danceable or abs(dance_scale) < 1e-04: + # print("NODANCE", dance_scale, self.deta, is_danceable, ss.sigma_next) + # yield SamplerResult(ss, self, x, sigma_up) + # print( + # "DANCE", dance_scale, self.deta, self.dyn_deta_mode, self.ds_noise, sigma_up + # ) + # sigma_down_normal, sigma_up_normal = get_ancestral_step( + # ss.sigma, ss.sigma_next, eta + # ) + # if self.dyn_deta_mode == "lerp": + # dt_normal = sigma_down_normal - ss.sigma + # x_normal = x + d * dt_normal + # else: + # x_normal = x + # sigma_down2, sigma_up2 = get_ancestral_step( + # sigma_leap, + # ss.sigma_next, + # eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), + # ) + # print( + # "-->", + # sigma_down2, + # sigma_up2, + # "--", + # self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), + # ) + # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.ds_noise * sigma_up) + # d_2 = to_d(x, sigma_leap, ss.denoised) + # dt_2 = sigma_down2 - sigma_leap + # result = x + d_2 * dt_2 + # # SIGMA: norm_up=9.062416076660156, up=10.703859329223633, up2=19.376544952392578, str=21.955078125 + # noise_strength = sigma_up2 + ((sigma_up - sigma_up_normal) ** 5.0) + # noise_strength = sigma_up2 + ((sigma_up2 - sigma_up) * 0.5) + # # noise_strength = sigma_up2 + ( + # # (sigma_up2 - sigma_up) ** (1.0 - (sigma_up_normal / sigma_up2)) + # # ) + # noise_diff = ( + # sigma_up - sigma_up_normal + # if sigma_up > sigma_up_normal + # else sigma_up_normal - sigma_up + # ) + # noise_div = ( + # sigma_up / sigma_up_normal + # if sigma_up > sigma_up_normal + # else sigma_up_normal / sigma_up + # ) + # noise_diff = sigma_up2 - sigma_up_normal + # noise_div = sigma_up2 / sigma_up_normal + # noise_div = ss.sigma / sigma_leap + + # # noise_strength = sigma_up2 + (noise_diff * noise_div) + # # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** 2.0) + # # noise_strength = sigma_up2 + ((1.0 - noise_diff) ** 0.5) + # # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up) * 0.5) ** 2.0) + # # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up_normal) * 0.5) ** 1.5) + # # noise_strength = sigma_up2 + ( + # # (noise_diff * 0.1875) ** (1.0 / (noise_div - 0.0)) + # # ) + # # noise_strength = sigma_up2 + ( + # # (noise_diff * 0.125) ** (1.0 / (noise_div * 1.25)) + # # ) + # # noise_strength = sigma_up2 + ((noise_diff * 0.2) ** (1.0 / (noise_div * 1.0))) + # noise_strength = sigma_up2 + (noise_diff * 0.9 * max(0.0, noise_div - 0.8)) + # noise_strength = sigma_up2 + ( + # (noise_diff / (curr_leap * 0.4)) + # * ((noise_div - (curr_leap / 2.0)).clamp(min=0, max=1.5) * 1.0) + # ) + # # (1.0 / (noise_div * 1.25))) + # # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** noise_div) + # print( + # f"SIGMA: norm_up={sigma_up_normal}, up={sigma_up}, up2={sigma_up2}, str={noise_strength}", + # # noise_diff, + # noise_div, + # ) + # return result, noise_strength + + # noise_diff = sigma_up2 - sigma_up * dance_scale + # noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) + # # noise_scale = sigma_up2 * self.ds_noise + # if self.dyn_deta_mode == "deta" or dance_scale == 1.0: + # return result, noise_scale + # result = torch.lerp(x_normal, result, dance_scale) + # # FIXME: Broken for noise samplers that care about s/sn + # return result, noise_scale + + # def step(self, x, ss): + # eta = self.get_dyn_eta(ss) + # leap_sigmas = ss.sigmas[ss.idx :] + # leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)] + # zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] + # max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 + # is_danceable = max_leap > 1 and ss.sigma_next != 0 + # curr_leap = max(1, min(self.leap, max_leap)) + # sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next + # # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) + # del leap_sigmas + # sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) + # d = to_d(x, ss.sigma, ss.denoised) + # # Euler method + # dt = sigma_down - ss.sigma + # x = x + d * dt + # if curr_leap == 1: + # return x, sigma_up + # dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) + # if not is_danceable or abs(dance_scale) < 1e-04: + # print("NODANCE", dance_scale, self.deta) + # return x, sigma_up + # print("NODANCE", dance_scale, self.deta) + # sigma_down_normal, _sigma_up_normal = get_ancestral_step( + # ss.sigma, ss.sigma_next, eta + # ) + # if self.dyn_deta_mode == "lerp": + # dt_normal = sigma_down_normal - ss.sigma + # x_normal = x + d * dt_normal + # else: + # x_normal = x + # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.s_noise * sigma_up) + # sigma_down2, sigma_up2 = get_ancestral_step( + # sigma_leap, + # ss.sigma_next, + # eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), + # ) + # d_2 = to_d(x, sigma_leap, ss.denoised) + # dt_2 = sigma_down2 - sigma_leap + # result = x + d_2 * dt_2 + # noise_diff = sigma_up2 - sigma_up * dance_scale + # noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) + # if self.dyn_deta_mode == "deta" or dance_scale == 1.0: + # return result, noise_scale + # result = torch.lerp(x_normal, result, dance_scale) + # # FIXME: Broken for noise samplers that care about s/sn + # return result, noise_scale + + +# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +# Which was originally written by Katherine Crowson +class TTMJVPStep(SingleStepSampler): + name = "ttm_jvp" + model_calls = 1 + + def __init__(self, *args, alternate_phi_2_calc=True, **kwargs): + super().__init__(*args, **kwargs) + self.alternate_phi_2_calc = alternate_phi_2_calc + + def step(self, x): + ss = self.ss + eta = self.get_dyn_eta() + sigma, sigma_next = ss.sigma, ss.sigma_next + # 2nd order truncated Taylor method + t, s = -sigma.log(), -sigma_next.log() + h = s - t + h_eta = h * (eta + 1) + + eps = to_d(x, sigma, ss.denoised) + denoised_prime = self.call_model( + x, sigma, tangents=(eps * -sigma, -sigma), call_index=1 + ).jdenoised + + phi_1 = -torch.expm1(-h_eta) + if self.alternate_phi_2_calc: + phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0 + else: + phi_2 = torch.expm1(-h_eta) + h_eta + x = torch.exp(-h_eta) * x + phi_1 * ss.denoised + phi_2 * denoised_prime + + noise_scale = ( + sigma_next * torch.sqrt(-torch.expm1(-2 * h * eta)) + if eta + else ss.sigma.new_zeros(1) + ) + yield from self.result(x, noise_scale) + + +class HeunStep(ReversibleSingleStepSampler): + name = "heun" + model_calls = 1 + default_history_limit, max_history = 0, 0 + allow_alt_cfgpp = True + + def reversible_correction(self, d_from, d_to): + reta, reversible_scale = self.get_reversible_cfg() + if reversible_scale == 0: + return 0 + sdr = self.get_ancestral_step(reta)[0] + dtr = sdr - self.ss.sigma + return (dtr**2 * (d_to - d_from) / 4) * self.reversible.scale + + def step(self, x): + ss = self.ss + s = ss.sigma + sd, su = self.get_ancestral_step(self.get_dyn_eta()) + dt = sd - s + hcur = ss.hcur + d = self.to_d(hcur) + x_next = hcur.denoised + d * sd + d_next = self.to_d(self.call_model(x_next, sd, call_index=1)) + result = hcur.denoised + d * s + result += (dt * (d + d_next)) * 0.5 + result -= self.reversible_correction(d, d_next) + yield from self.result(result, su, sigma_down=sd) + + +class Heun1SStep(HeunStep): + name = "heun_1s" + model_calls = (0, 1) + allow_alt_cfgpp = True + default_history_limit, max_history = 1, 1 + + def step(self, x): + ss = self.ss + s = ss.sigma + if self.available_history() == 0: + return (yield from super().step(x)) + hcur, hprev = ss.hcur, ss.hprev + d_prev = self.to_d(hprev) + sd, su = self.get_ancestral_step(self.get_dyn_eta()) + dt = sd - s + d = self.to_d(hcur) + result = hcur.denoised + hcur.sigma * self.to_d(hcur) + result += (dt * (d_prev + d)) * 0.5 + result -= self.reversible_correction(d_prev, d) + yield from self.result(result, su, sigma_down=sd) + + +class ClybiusSENSStep(DPMPP2MSDEStep): + name = "clybius_sens" + default_history_limit, max_history = 2, 2 + allow_alt_cfgpp = False + default_reta = 1.0 + default_reversible_scale = 1.0 + + def __init__(self, *, tsde_reversible=None, **kwargs): + super().__init__(**kwargs) + self.tsde_reversible = ReversibleConfig.build( + default_eta=self.default_reta, + default_scale=self.default_reversible_scale, + **utils.fallback(tsde_reversible, {}), + ) + + def step(self, x): + ss = self.ss + sigma, sigma_next = ss.sigma, ss.sigma_next + denoised = ss.denoised + # DPM-Solver++(2M) SDE + t, s = -sigma.log(), -sigma_next.log() + h = s - t + eta_h = self.get_dyn_eta() * h + ratio = sigma_next / sigma + x = ((ratio * (-eta_h).exp()) * x).add_((-h - eta_h).expm1().neg() * denoised) + noise_strength = sigma_next * (-2 * eta_h).expm1().neg().sqrt() + if self.available_history() == 0: + return (yield from self.result(x, noise_strength)) + sigma_prev, old_denoised = ss.hprev.sigma, ss.hprev.denoised + h_last = (-sigma.log()) - (-sigma_prev.log()) + r = h_last / h + if self.solver_type == "midpoint": + multiplier = 0.5 * (-h - eta_h).expm1().neg() + else: + multiplier = (-h - eta_h).expm1().neg() / (-h - eta_h) + 1 + reta, reversible_scale = self.get_reversible_cfg() + if reversible_scale != 0: + multiplier *= 0.5 + x += (denoised - old_denoised).mul_((1 / r) * multiplier) + if reversible_scale != 0: + reta_h = reta * h + if self.solver_type == "midpoint": + rmultiplier = 0.5 * (-h - reta_h).expm1().neg() + else: + rmultiplier = (-h - reta_h).expm1().neg() / (-h - reta_h) + 1 + rmultiplier = ((1 / r) * (rmultiplier**2 / 2)) * reversible_scale + x -= (old_denoised - denoised).mul_(rmultiplier) + if self.available_history() > 1: + tsde_reta, tsde_reversible_scale = self.get_reversible_cfg( + reversible=self.tsde_reversible + ) + tsde_reta_h = tsde_reta * h + sigma_prev_2 = ss.hist[-3].sigma + h_last_2 = (-sigma_prev.log()) - (-sigma_prev_2.log()) + r = h_last_2 / h + old_denoised_2 = ss.hist[-3].denoised + d = (old_denoised - old_denoised_2).div_(r) + d_2 = (old_denoised - denoised).div_(r) + + d_rev = (denoised - old_denoised).div_(r) + d_2_rev = (old_denoised_2 - old_denoised).div_(r) + + rphi = tsde_reta_h.neg().expm1() / tsde_reta_h + 1 + tsde_adjustment = rphi * (d + d_2) / 2 + if tsde_reversible_scale != 0: + tsde_adjustment -= (rphi**2 * (d_rev + d_2_rev) / 2).mul_( + tsde_reversible_scale + ) + x += tsde_adjustment + yield from self.result(x, noise_strength) + + +registry.add( + BogackiStep, + ClybiusSENSStep, + EulerDancingStep, + ReversibleBogackiStep, + ReversibleHeunStep, + ReversibleHeun1SStep, + RESStep, + TTMJVPStep, + TrapezoidalStep, + RK4Step, + RKDynamicStep, + RKF45Step, + Heun1SStep, + HeunStep, +) diff --git a/py/step_samplers/extraltodeus.py b/py/step_samplers/extraltodeus.py new file mode 100644 index 0000000..b28f519 --- /dev/null +++ b/py/step_samplers/extraltodeus.py @@ -0,0 +1,131 @@ +# Samplers based on design from https://github.com/Extraltodeus/ + +import typing + +import torch +import tqdm + +from .base import SingleStepSampler +from . import registry + + +class DistanceConfig(typing.NamedTuple): + resample: int = 3 + resample_end: int = 1 + eta: float = 0.0 + s_noise: float = 1.0 + alt_cfgpp_scale: float = 0.0 + first_eta_step: int = 0 + last_eta_step: int = -1 + custom_noise_name: str = "alt" + immiscible: dict | bool | None = None + + +# Based on https://github.com/Extraltodeus/DistanceSampler +class DistanceStep(SingleStepSampler): + name = "extraltodeus_distance" + allow_alt_cfgpp = True + model_calls = -1 + uses_alt_noise = True + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.distance = DistanceConfig(**self.options.get("distance", {})) + + @property + def require_uncond(self): + return super().require_uncond or self.distance.alt_cfgpp_scale != 0 + + def distance_resample_steps(self): + ss = self.ss + resample, resample_end = self.distance.resample, self.distance.resample_end + if resample == -1: + current_resample = min(10, (ss.sigmas.shape[0] - ss.idx) // 2) + else: + current_resample = resample + if resample_end < 0: + return current_resample + sigma = ss.sigma + s_min = (ss.sigmas if ss.sigmas[-1] > 0 else ss.sigmas[:-1]).min() + s_max = ss.sigmas.max() + res_mul = max(0, min(1, ((sigma - s_min) / (s_max - s_min)) ** 0.5)) + return max( + min(current_resample, resample_end), + min( + max(current_resample, resample_end), + int(current_resample * res_mul + resample_end * (1 - res_mul)), + ), + ) + + @staticmethod + def distance_weights(t, p): + batch = t.shape[0] + d = torch.stack( + tuple((t - t[idx]).abs().sum(dim=0) for idx in range(batch)), + dim=0, + ) + d_min, d_max = d.min(), d.max() + d = torch.nan_to_num( + (1 - (d - d_min) / (d_max - d_min)).pow(p), + nan=1, + neginf=1, + posinf=1, + ) + d /= d.sum(dim=0) + return d.mul_(t).sum(dim=0) + + def step(self, x): + resample_steps = self.distance_resample_steps() + if resample_steps < 1: + return (yield from self.euler_step(x)) + distance = self.distance + ss = self.ss + sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta()) + rsigma_down, rsigma_up = self.get_ancestral_step(eta=distance.eta) + rsigma_up *= distance.s_noise + sigma, sigma_next = ss.sigma, ss.sigma_next + zero_up = sigma * 0 + d = self.to_d(ss.hcur) + can_ancestral = not torch.equal(rsigma_down, sigma_next) + start_eta_idx, end_eta_idx = ( + max(0, resample_steps + v if v < 0 else v) + for v in ( + distance.first_eta_step, + distance.last_eta_step, + ) + ) + dt = sigma_down - sigma + d = self.to_d(ss.hcur) + x_n = [d] + for re_step in tqdm.trange( + resample_steps, desc="distance_resample", disable=ss.disable_status + ): + if can_ancestral and start_eta_idx <= re_step <= end_eta_idx: + curr_sigma_down, curr_sigma_up = rsigma_down, rsigma_up + else: + curr_sigma_down, curr_sigma_up = sigma_next, zero_up + rdt = curr_sigma_down - sigma + x_new = x + d * rdt + if curr_sigma_up != 0: + x_new = yield from self.result( + x_new, + curr_sigma_up, + sigma=sigma, + sigma_down=curr_sigma_down, + noise_sampler=self.alt_noise_sampler, + final=False, + ) + sr = self.call_model(x_new, sigma_next, call_index=re_step + 1) + new_d = sr.to_d( + sigma=curr_sigma_down, alt_cfgpp_scale=distance.alt_cfgpp_scale + ) + x_n.append(new_d) + if re_step == 0: + d = (new_d + d) / 2 + else: + d = self.distance_weights(torch.stack(x_n), re_step + 2) + x_n.append(d) + yield from self.result(x + d * dt, sigma_up, sigma_down=sigma_down) + + +registry.add(DistanceStep) diff --git a/py/step_samplers/misc.py b/py/step_samplers/misc.py new file mode 100644 index 0000000..04910e7 --- /dev/null +++ b/py/step_samplers/misc.py @@ -0,0 +1,29 @@ +from .base import SingleStepSampler, registry + + +# Referenced from https://github.com/ace-step/ACE-Step/ +class PingPongStep(SingleStepSampler): + name = "pingpong" + default_eta = 0.0 + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + pingpong_options = self.options.pop("pingpong", {}) + self.pingpong_start_step = pingpong_options.get("start_step", 0) + self.pingpong_end_step = pingpong_options.get("end_step", 0) + + def step(self, x): + ss = self.ss + use_pingpong = self.pingpong_start_step <= ss.step <= self.pingpong_end_step + if not use_pingpong: + return (yield from self.euler_step(x, eta=0.0)) + sn = ss.sigma_next + denoised = ( + ss.denoised * (1.0 - sn) if ss.model.is_rectified_flow else ss.denoised + ) + yield from self.result(denoised, sn, sigma_down=sn) + + +registry.add( + PingPongStep, +) diff --git a/py/step_samplers/registry.py b/py/step_samplers/registry.py new file mode 100644 index 0000000..0d02898 --- /dev/null +++ b/py/step_samplers/registry.py @@ -0,0 +1,39 @@ +SAMPLER_LIST = [] + +STEP_SAMPLERS = {} +STEP_SAMPLER_SIMPLE_NAMES = {} + + +def add(*objs): + global SAMPLER_LIST + SAMPLER_LIST += objs + + +def init(): + global STEP_SAMPLERS, STEP_SAMPLER_SIMPLE_NAMES + STEP_SAMPLER_SIMPLE_NAMES.clear() + STEP_SAMPLERS.clear() + euler = None + temp = [] + for c in SAMPLER_LIST: + mc = c.model_calls + if mc == 0: + prettymc = "" + elif isinstance(mc, tuple): + prettymc = f" ({mc[0]}-{mc[-1]})" + elif mc < 0: + prettymc = " (variable)" + else: + prettymc = f" ({mc})" + if c.name == "euler": + euler = c + temp.append((f"{c.name}{prettymc}", c)) + temp.sort(key=lambda item: item[1].name) + if euler is None: + raise RuntimeError( + "Impossible: euler sampler not found when building sampler registry" + ) + STEP_SAMPLERS["default (euler)"] = euler + STEP_SAMPLERS |= {k: v for k, v in temp} + STEP_SAMPLER_SIMPLE_NAMES["default"] = euler + STEP_SAMPLER_SIMPLE_NAMES |= {v.name: v for _k, v in temp} diff --git a/py/res_support.py b/py/step_samplers/res_support.py similarity index 100% rename from py/res_support.py rename to py/step_samplers/res_support.py diff --git a/py/step_samplers/solver_base.py b/py/step_samplers/solver_base.py new file mode 100644 index 0000000..b9127bf --- /dev/null +++ b/py/step_samplers/solver_base.py @@ -0,0 +1,49 @@ +from .base import SingleStepSampler, MinSigmaStepMixin + + +class DESolverStep(SingleStepSampler, MinSigmaStepMixin): + de_default_solver = None + sample_sigma_zero = True + default_eta = 0.0 + + def __init__( + self, + *args, + de_solver=None, + de_max_nfe=100, + de_rtol=-2.5, + de_atol=-3.5, + de_fixup_hack=0.025, + de_split=1, + de_min_sigma=0.0292, + **kwargs, + ): + self.check_solver_support() + super().__init__(*args, **kwargs) + de_solver = self.de_default_solver if de_solver is None else de_solver + self.de_solver_name = de_solver + self.de_max_nfe = de_max_nfe + self.de_rtol = 10**de_rtol + self.de_atol = 10**de_atol + self.de_fixup_hack = de_fixup_hack + self.de_split = de_split + self.de_min_sigma = de_min_sigma if de_min_sigma is not None else 0.0 + + def check_solver_support(self): + raise NotImplementedError + + def de_get_step(self, x): + eta = self.get_dyn_eta() + ss = self.ss + s, sn = ss.sigma, ss.sigma_next + sn = self.adjust_step(sn, self.de_min_sigma) + sigma_down, sigma_up = self.get_ancestral_step(eta, sigma_next=sn) + if self.de_fixup_hack != 0: + sigma_down = (sigma_down - (s - sigma_down) * self.de_fixup_hack).clamp( + min=0 + ) + return s, sn, sigma_down, sigma_up + + @staticmethod + def reverse_time(t, t0, t1): + return t1 + (t0 - t) diff --git a/py/step_samplers/solver_diffrax.py b/py/step_samplers/solver_diffrax.py new file mode 100644 index 0000000..721d993 --- /dev/null +++ b/py/step_samplers/solver_diffrax.py @@ -0,0 +1,304 @@ +import contextlib +import os +import typing +import warnings + +import numpy +import torch +import tqdm + +import comfy + +from . import registry +from .solver_base import DESolverStep + +HAVE_DIFFRAX = False + + +if not os.environ.get("COMFYUI_OCS_NO_DIFFRAX_SOLVER"): + with contextlib.suppress(ImportError): + import diffrax + import jax + + if not os.environ.get("COMFYUI_OCS_NO_DISABLE_JAX_PREALLOCATE"): + os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false" + # jax.config.update("jax_enable_x64", True) + + HAVE_DIFFRAX = True + + +if HAVE_DIFFRAX: + + class RevVirtualBrownianTree(diffrax.VirtualBrownianTree): + def evaluate(self, t0, t1, *args, **kwargs): + if t1 is not None: + return super().evaluate(t1, t0, *args, **kwargs) + return super().evaluate(t0, t1, *args, **kwargs) + + class StepCallbackTqdmProgressMeter(diffrax.TqdmProgressMeter): + step_callback: typing.Callable = None + + def _init_bar(self, *args, **kwargs): + if self.step_callback is None: + return super()._init_bar(*args, **kwargs) + bar_format = "{percentage:.2f}%{step_callback}|{bar}| [{elapsed}<{remaining}, {rate_fmt}{postfix}]" + step_callback = self.step_callback + + class WrapTqdm(tqdm.tqdm): + @property + def format_dict(self): + d = super().format_dict + d.update(step_callback=step_callback()) + return d + + return WrapTqdm(total=100, unit="%", bar_format=bar_format) + + +class DiffraxStep(DESolverStep): + name = "diffrax" + model_calls = -1 + allow_alt_cfgpp = True + de_default_solver = "dopri5" + default_eta = 0.0 + + def __init__( + self, + *args, + de_split=1, + de_initial_step=0.25, + de_ctl_pcoeff=0.3, + de_ctl_icoeff=0.9, + de_ctl_dcoeff=0.2, + 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) + solvers = dict( + euler=diffrax.Euler, + heun=diffrax.Heun, + midpoint=diffrax.Midpoint, + ralston=diffrax.Ralston, + bosh3=diffrax.Bosh3, + tsit5=diffrax.Tsit5, + dopri5=diffrax.Dopri5, + dopri8=diffrax.Dopri8, + implicit_euler=diffrax.ImplicitEuler, + # kvaerno3=diffrax.Kvaerno3, + # kvaerno4=diffrax.Kvaerno4, + # kvaerno5=diffrax.Kvaerno5, + semi_implicit_euler=diffrax.SemiImplicitEuler, + reversible_heun=diffrax.ReversibleHeun, + leapfrog_midpoint=diffrax.LeapfrogMidpoint, + euler_heun=diffrax.EulerHeun, + ito_milstein=diffrax.ItoMilstein, + stratonovich_milstein=diffrax.StratonovichMilstein, + sea=diffrax.SEA, + sra1=diffrax.SRA1, + shark=diffrax.ShARK, + general_shark=diffrax.GeneralShARK, + slow_rk=diffrax.SlowRK, + spark=diffrax.SPaRK, + ) + levy_areas = dict( + brownian_increment=diffrax.BrownianIncrement, + space_time=diffrax.SpaceTimeLevyArea, + space_time_time=diffrax.SpaceTimeTimeLevyArea, + ) + # jax.config.update("jax_disable_jit", True) + self.de_solver_method = solvers[self.de_solver_name]() + 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 = diffrax_adaptive + self.de_split = de_split + 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 + def t2j(t): + return jax.block_until_ready( + jax.numpy.array(numpy.array(t.detach().cpu().contiguous())) + ) + + @staticmethod + def j2t(t): + return torch.from_numpy(numpy.array(jax.block_until_ready(t))).contiguous() + + def check_solver_support(self): + if not HAVE_DIFFRAX: + raise RuntimeError( + "Diffrax sampler requires diffrax and jax installed in venv." + ) + + def step(self, x): + s, sn, sigma_down, sigma_up = self.de_get_step(x) + if self.de_min_sigma is not None and s <= self.de_min_sigma: + return (yield from self.euler_step(x)) + ss = self.ss + bidx = 0 + mcc = 0 + _b, c, h, w = x.shape + interrupted = None + t0, t1 = sigma_down.item(), s.item() + + def odefn_(t_orig, y_flat, args=()): + nonlocal mcc, interrupted + t = self.reverse_time(self.j2t(t_orig).to(s), t0, t1) + if t <= 1e-05: + return jax.numpy.zeros_like(y_flat) + if mcc >= self.de_max_nfe: + raise RuntimeError("DiffraxStep: Model call limit exceeded") + y = self.j2t(y_flat.reshape(1, c, h, w)).to(x) + t32 = t.to(s).clamp(min=1e-05) + flat_shape = y_flat.shape + del y_flat + + if not args and mcc == 0 and torch.all(t == s): + mr_cached = True + mr = ss.hcur + mcc = 1 + else: + mr_cached = False + try: + if not args: + mr = self.call_model(y, t32, call_index=mcc, s_in=t.new_ones(1)) + else: + print("TANGENTS") + mr = self.call_model( + y, + t32, + call_index=mcc, + tangents=args, + s_in=t.new_ones(1), + ) + except comfy.model_management.InterruptProcessingException as exc: + interrupted = exc + raise + mcc += 1 + result = self.to_d(mr)[bidx if mr_cached else 0].reshape(*flat_shape) + return self.t2j(-result) + + if not self.de_fake_pure_callback: + + def odefn(t, y_flat, args): + return jax.experimental.io_callback( + odefn_, y_flat, t, y_flat, ordered=True + ) + + else: + + def odefn(t, y_flat, args): + return jax.pure_callback(odefn_, y_flat, t, y_flat) + + def g(t, y, _args): + if self.de_g_split_time_mode: + val = jax.lax.cond( + t < t0 + (t1 - t0) * 0.5, + lambda: self.de_g_multiplier, + lambda: -self.de_g_multiplier, + ) + else: + val = self.de_g_multiplier + if self.de_g_time_scaling: + val *= self.reverse_time(t, t0, t1) if self.de_g_reverse_time else t + if not self.de_batch_channels: + return val + return jax.numpy.float32(val).broadcast((y.shape[0],)) + + def progress_callback(): + return f" ({mcc:>3}/{self.de_max_nfe:>3}) {self.de_solver_name}" + + term = diffrax.ODETerm(odefn) + method = self.de_solver_method + if self.de_adaptive: + controller = diffrax.PIDController( + atol=self.de_atol, + rtol=self.de_rtol, + dtmin=1e-05, + pcoeff=self.de_ctl_pcoeff, + icoeff=self.de_ctl_icoeff, + dcoeff=self.de_ctl_dcoeff, + error_order=self.de_error_order, + ) + else: + controller = diffrax.ConstantStepSize() + + if not self.de_adaptive: + dt0 = (t1 - t0) / self.de_split + else: + dt0 = (t1 - t0) * self.de_initial_step + if self.de_sde_mode: + bm = diffrax.VirtualBrownianTree( + t0=ss.sigmas.min().item(), + t1=ss.sigmas.max().item(), + tol=1e-06, + levy_area=self.de_levy_area_approx, + shape=(c,) if self.de_batch_channels else (), + key=jax.random.PRNGKey(ss.noise.seed + ss.noise.seed_offset), + ) + term = diffrax.MultiTerm(term, diffrax.ControlTerm(g, bm)) + results = [] + for batch in tqdm.trange( + 1, + x.shape[0] + 1, + desc="batch", + leave=False, + disable=x.shape[0] == 1 or ss.disable_status, + ): + bidx = batch - 1 + mcc = 0 + if self.de_batch_channels: + y_flat = x[bidx].flatten(start_dim=1) + else: + y_flat = x[bidx].unsqueeze(0).flatten(start_dim=1) + y_flat = self.t2j(y_flat) + with warnings.catch_warnings(): + warnings.simplefilter(action="ignore", category=FutureWarning) + try: + solution = diffrax.diffeqsolve( + terms=term, + solver=method, + t0=t0, + t1=t1, + dt0=dt0, + y0=y_flat, + saveat=diffrax.SaveAt(t1=True), + stepsize_controller=controller, + progress_meter=StepCallbackTqdmProgressMeter( + step_callback=progress_callback, + refresh_steps=1, + ), + ) + except Exception: + if interrupted is not None: + raise interrupted + raise + results.append(self.j2t(solution.ys).view(1, *x.shape[1:])) + del solution + result = torch.cat(results).to(x) + sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) + yield from self.result(result, sigma_up, sigma_down=sigma_down) + + +registry.add(DiffraxStep) diff --git a/py/step_samplers/solver_tde.py b/py/step_samplers/solver_tde.py new file mode 100644 index 0000000..11b46cd --- /dev/null +++ b/py/step_samplers/solver_tde.py @@ -0,0 +1,118 @@ +import contextlib + +import torch +import tqdm + +from . import registry +from .solver_base import DESolverStep + +HAVE_TDE = False +with contextlib.suppress(ImportError): + import torchdiffeq as tde + + HAVE_TDE = True + + +class TDEStep(DESolverStep): + name = "tde" + model_calls = -1 + allow_alt_cfgpp = True + allow_cfgpp = False + de_default_solver = "rk4" + default_eta = 0.0 + + def __init__( + self, + *args, + de_split=1, + **kwargs, + ): + super().__init__(*args, **kwargs) + self.de_split = de_split + + def check_solver_support(self): + if not HAVE_TDE: + raise RuntimeError( + "TDE sampler requires torchdiffeq installed in venv. Example: pip install torchdiffeq" + ) + + def step(self, x): + s, sn, sigma_down, sigma_up = self.de_get_step(x) + if self.de_min_sigma is not None and s <= self.de_min_sigma: + return (yield from self.euler_step(x)) + ss = self.ss + delta = (s - sigma_down).item() + mcc = 0 + bidx = 0 + pbar = None + + def odefn(t, y): + nonlocal mcc + if t < 1e-05: + return torch.zeros_like(y) + if mcc >= self.de_max_nfe: + raise RuntimeError("TDEStep: Model call limit exceeded") + + pct = (s - t) / delta + pbar.n = round(min(999, pct.item() * 999)) + pbar.update(0) + pbar.set_description( + f"{self.de_solver_name}({mcc}/{self.de_max_nfe})", refresh=True + ) + + if t == ss.sigma and torch.equal(x[bidx], y): + mr_cached = True + mr = ss.hcur + mcc = 1 + else: + mr_cached = False + mr = self.call_model( + y.unsqueeze(0), t, call_index=mcc, s_in=t.new_ones(1) + ) + mcc += 1 + return self.to_d(mr)[bidx if mr_cached else 0] + + result = torch.zeros_like(x) + t = sigma_down.new_zeros(self.de_split + 1) + torch.linspace(ss.sigma, sigma_down, t.shape[0], out=t) + + for batch in tqdm.trange( + 1, + x.shape[0] + 1, + desc="batch", + leave=False, + disable=x.shape[0] == 1 or ss.disable_status, + ): + bidx = batch - 1 + mcc = 0 + if pbar is not None: + pbar.close() + pbar = tqdm.tqdm( + total=1000, + desc=self.de_solver_name, + leave=True, + disable=ss.disable_status, + ) + solution = tde.odeint( + odefn, + x[bidx], + t, + rtol=self.de_rtol, + atol=self.de_atol, + method=self.de_solver_name, + options={ + "min_step": 1e-05, + "dtype": torch.float64, + }, + )[-1] + result[bidx] = solution + + sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) + if pbar is not None: + pbar.n = pbar.total + pbar.update(0) + pbar.close() + yield from self.result(result, sigma_up, sigma_down=sigma_down) + + +registry.add(TDEStep) diff --git a/py/step_samplers/solver_tode.py b/py/step_samplers/solver_tode.py new file mode 100644 index 0000000..f2bdb5a --- /dev/null +++ b/py/step_samplers/solver_tode.py @@ -0,0 +1,132 @@ +import contextlib + +import torch +import tqdm + +from . import registry +from .solver_base import DESolverStep + + +HAVE_TODE = False +with contextlib.suppress(ImportError, RuntimeError): + import torchode as tode + + HAVE_TODE = True + + +class TODEStep(DESolverStep): + name = "tode" + model_calls = -1 + allow_alt_cfgpp = True + de_default_solver = "dopri5" + default_eta = 0.0 + + def __init__( + self, + *args, + de_initial_step=0.25, + tode_compile=False, + de_ctl_pcoeff=0.3, + de_ctl_icoeff=0.9, + de_ctl_dcoeff=0.2, + **kwargs, + ): + if not HAVE_TODE: + raise RuntimeError( + "TODE sampler requires torchode installed in venv. Example: pip install torchode" + ) + super().__init__(*args, **kwargs) + self.de_solver_method = tode.interface.METHODS[self.de_solver_name] + self.de_ctl_pcoeff = de_ctl_pcoeff + self.de_ctl_icoeff = de_ctl_icoeff + self.de_ctl_dcoeff = de_ctl_dcoeff + self.de_compile = tode_compile + self.de_initial_step = de_initial_step + + def check_solver_support(self): + if not HAVE_TODE: + raise RuntimeError( + "TODE sampler requires torchode installed in venv. Example: pip install torchode" + ) + + def step(self, x): + s, sn, sigma_down, sigma_up = self.de_get_step(x) + if self.de_min_sigma is not None and s <= self.de_min_sigma: + return (yield from self.euler_step(x)) + ss = self.ss + delta = (ss.sigma - sigma_down).item() + mcc = 0 + pbar = None + b, c, h, w = x.shape + + def odefn(t, y_flat): + nonlocal mcc + if torch.all(t <= 1e-05).item(): + return torch.zeros_like(y_flat) + if mcc >= self.de_max_nfe: + raise RuntimeError("TDEStep: Model call limit exceeded") + + pct = (s - t) / delta + pbar.n = round(pct.min().item() * 999) + pbar.update(0) + pbar.set_description( + f"{self.de_solver_name}({mcc}/{self.de_max_nfe})", refresh=True + ) + y = y_flat.reshape(-1, c, h, w) + t32 = t.to(torch.float32) + del y_flat + + if mcc == 0 and torch.all(t == s): + mr = ss.hcur + mcc = 1 + else: + mr = self.call_model(y, t32.clamp(min=1e-05), call_index=mcc) + mcc += 1 + result = self.to_d(mr).flatten(start_dim=1) + for bi in range(t.shape[0]): + if t[bi] <= 1e-05: + result[bi, :] = 0 + return result + + t = torch.stack((s, sigma_down)).to(torch.float64).repeat(b, 1) + + pbar = tqdm.tqdm( + total=1000, desc=self.de_solver_name, leave=True, disable=ss.disable_status + ) + + term = tode.ODETerm(odefn) + method = self.de_solver_method(term=term) + controller = tode.PIDController( + term=term, + atol=self.de_atol, + rtol=self.de_rtol, + dt_min=1e-05, + pcoeff=self.de_ctl_pcoeff, + icoeff=self.de_ctl_icoeff, + dcoeff=self.de_ctl_dcoeff, + ) + solver_ = tode.AutoDiffAdjoint(method, controller) + solver = solver_ if not self.de_compile else torch.compile(solver_) + problem = tode.InitialValueProblem( + y0=x.flatten(start_dim=1), t_start=t[:, 0], t_end=t[:, -1] + ) + dt0 = ( + (t[:, -1] - t[:, 0]) * self.de_initial_step + if self.de_initial_step + else None + ) + solution = solver.solve(problem, dt0=dt0) + + # print("\nSOLUTION", solution.stats, solution.ys.shape) + result = solution.ys[:, -1].reshape(-1, c, h, w) + del solution + + sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) + if pbar is not None: + pbar.n = pbar.total + pbar.update(0) + pbar.close() + yield from self.result(result, sigma_up, sigma_down=sigma_down) + + +registry.add(TODEStep) diff --git a/py/step_samplers/solver_tsde.py b/py/step_samplers/solver_tsde.py new file mode 100644 index 0000000..fad050d --- /dev/null +++ b/py/step_samplers/solver_tsde.py @@ -0,0 +1,192 @@ +import contextlib + +import torch +import tqdm + +from . import registry +from .solver_base import DESolverStep + +HAVE_TSDE = False +with contextlib.suppress(ImportError): + import torchsde + + HAVE_TSDE = True + + +class TSDEStep(DESolverStep): + name = "tsde" + model_calls = -1 + allow_alt_cfgpp = True + de_default_solver = "reversible_heun" + default_eta = 0.0 + + def __init__( + self, + *args, + de_initial_step=0.25, + de_split=1, + de_adaptive=False, + 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 = 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): + if not HAVE_TSDE: + raise RuntimeError( + "TSDE sampler requires torchsde installed in venv. Example: pip install torchsde" + ) + + def step(self, x): + s, sn, sigma_down, sigma_up = self.de_get_step(x) + if self.de_min_sigma is not None and s <= self.de_min_sigma: + return (yield from self.euler_step(x)) + ss = self.ss + delta = (ss.sigma - sigma_down).item() + bidx = 0 + mcc = 0 + pbar = None + _b, c, h, w = x.shape + outer_self = self + + class SDE(torch.nn.Module): + noise_type = outer_self.de_noise_type + sde_type = outer_self.de_sde_type + + @torch.no_grad() + def f(self, t_rev, y_flat): + nonlocal mcc + t = s - (t_rev - sigma_down) + # print(f"\nf at t_rev={t_rev}, t={t} :: {y_flat.shape}") + if torch.all(t <= 1e-05).item(): + return torch.zeros_like(y_flat) + if mcc >= outer_self.de_max_nfe: + raise RuntimeError("TSDEStep: Model call limit exceeded") + + pct = (s - t) / delta + pbar.n = round(pct.min().item() * 999) + pbar.update(0) + pbar.set_description( + f"{outer_self.de_solver_name}({mcc}/{outer_self.de_max_nfe})", + refresh=True, + ) + flat_shape = y_flat.shape + y = y_flat.view(1, c, h, w) + t32 = t.to(torch.float32) + del y_flat + + if mcc == 0 and torch.all(t == s): + mr_cached = True + mr = ss.hcur + mcc = 1 + else: + mr_cached = False + mr = outer_self.call_model( + y, t32.clamp(min=1e-05), call_index=mcc, s_in=t.new_ones(1) + ) + mcc += 1 + return -outer_self.to_d(mr)[bidx if mr_cached else 0].view(*flat_shape) + + @torch.no_grad() + def g(self, t_rev, y_flat): + t = (s - sigma_down) - (t_rev - sigma_down) + pct = t / (s - sigma_down) + if outer_self.de_g_reverse_time: + pct = 1.0 - pct + multiplier = outer_self.de_g_multiplier + if outer_self.de_g_derp_mode and mcc % 2 == 0: + multiplier *= -1 + val = t * pct * multiplier + if self.noise_type == "diagonal": + out = val.repeat(*y_flat.shape) + elif self.noise_type == "scalar": + out = val.repeat(*y_flat.shape, 1) + else: + out = val.repeat(*y_flat.shape, outer_self.de_noise_channels) + return out + + t = torch.stack((sigma_down, s)).to(torch.float) + + pbar = tqdm.tqdm( + total=1000, desc=self.de_solver_name, leave=True, disable=ss.disable_status + ) + + dt0 = ( + delta * self.de_initial_step if self.de_adaptive else delta / self.de_split + ) + results = [] + for batch in tqdm.trange( + 1, + x.shape[0] + 1, + desc="batch", + leave=False, + disable=x.shape[0] == 1 or ss.disable_status, + ): + bidx = batch - 1 + mcc = 0 + sde = SDE() + if self.de_batch_channels: + y_flat = x[bidx].flatten(start_dim=1) + else: + y_flat = x[bidx].unsqueeze(0).flatten(start_dim=1) + if sde.noise_type == "diagonal": + bm_size = (y_flat.shape[0], y_flat.shape[1]) + elif sde.noise_type == "scalar": + bm_size = (y_flat.shape[0], 1) + else: + bm_size = (y_flat.shape[0], self.de_noise_channels) + bm = torchsde.BrownianInterval( + dtype=x.dtype, + device=x.device, + t0=-s, + t1=s, + entropy=ss.noise.seed, + levy_area_approximation=self.de_levy_area_approx, + tol=1e-06, + size=bm_size, + ) + + ys = torchsde.sdeint( + sde, + y_flat, + t, + method=self.de_solver_name, + adaptive=self.de_adaptive, + atol=self.de_atol, + rtol=self.de_rtol, + dt=dt0, + bm=bm, + ) + del y_flat + results.append(ys[-1].view(1, c, h, w)) + del ys + result = torch.cat(results) + del results + + sigma_up, result = yield from self.adjusted_step(sn, result, mcc, sigma_up) + if pbar is not None: + pbar.n = pbar.total + pbar.update(0) + pbar.close() + yield from self.result(result, sigma_up, sigma_down=sigma_down) + + +registry.add(TSDEStep) diff --git a/py/substep_merging.py b/py/substep_merging.py index 0877041..47cc929 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -9,7 +9,8 @@ from . import utils from .filtering import make_filter, FilterRefs, FILTER_HANDLERS from .noise import ImmiscibleNoise from .restart import Restart -from .step_samplers import STEP_SAMPLERS, StepSamplerContext +from .step_samplers import STEP_SAMPLERS +from .step_samplers.base import StepSamplerContext from .substep_sampling import StepSamplerChain from .utils import check_time, fallback @@ -38,6 +39,8 @@ class MergeSubstepsSampler: self.preview_mode = options.pop("preview_mode", "denoised") self.require_uncond = any(sampler.require_uncond for sampler in samplers) self.cfg_scale_override = options.pop("cfg_scale_override", None) + self.afs_start_step = options.pop("afs_start_step", 0) + self.afs_end_step = options.pop("afs_end_step", -1) self.options = options def check_match(self, handlers: None | object, *, ss: None | object = None): @@ -58,24 +61,40 @@ class MergeSubstepsSampler: return operator.truth(self.when.eval(handlers)) def step_input(self, x, *, ss=None): + ss = fallback(ss, self.ss) + ss.noise.update_x(x) if self.pre_filter is None: return x - ss = fallback(ss, self.ss) - return self.pre_filter.apply(x, refs=fallback(ss, self.ss).refs) + x = self.pre_filter.apply(x, refs=fallback(ss, self.ss).refs) + ss.noise.update_x(x) + return x def step_output(self, x, *, orig_x=None, ss=None): + ss = fallback(ss, self.ss) + ss.noise.update_x(x) if self.post_filter is None: return x - ss = fallback(ss, self.ss) refs = ss.refs if orig_x is None else ss.refs | FilterRefs({"orig_x": orig_x}) - return self.post_filter.apply(x, refs=refs) + x = self.post_filter.apply(x, refs=refs) + ss.noise.update_x(x) + return x def __call__(self, x): orig_x = x x = self.step_input(x) - x = self.step(x) + if self.afs_start_step <= self.ss.step <= self.afs_end_step: + x = self.afs_step(x) + else: + x = self.step(x) return self.step_output(x, orig_x=orig_x) + # From https://arxiv.org/abs/2210.05475 + def afs_step(self, x): + sigma, sigma_next = self.ss.sigma, self.ss.sigma_next + afs_d = x / ((1 + sigma**2).sqrt()) + dt = sigma_next - sigma + return x + afs_d * dt + def step(self, x): raise NotImplementedError @@ -574,6 +593,79 @@ class LookaheadMergeSubstepsSampler(MergeSubstepsSampler): return x +class PingpongMergeSubstepsSampler(MergeSubstepsSampler): + name = "pingpong" + + def __init__(self, ss, group, **kwargs): + super().__init__(ss, group, **kwargs) + pingpong = self.options.pop("pingpong", {}).copy() + self.pingpong_s_noise = pingpong.pop("s_noise", 1.0) + immiscible = pingpong.get("immiscible", False) + self.immiscible = ( + ImmiscibleNoise(**immiscible) if immiscible is not False else False + ) + + self.custom_noise = self.options.get("custom_noise") + if isinstance(self.custom_noise, str): + self.custom_noise = self.options.get(f"custom_noise_{self.custom_noise}") + + def step(self, x): + orig_x = x.clone() + ss = self.ss + subss = self.ss.clone_edit(idx=ss.idx, sigmas=ss.sigmas) + substep = 0 + max_idx = len(ss.sigmas) - 1 + eff_substeps = min(max_idx - ss.idx, self.substeps) + pbar = tqdm.tqdm(total=eff_substeps, initial=0, disable=ss.disable_status) + for ssampler_ in self.samplers: + substeps_remain = eff_substeps - substep + if substeps_remain == 0: + break + with StepSamplerContext(ssampler_, subss) as ssampler: + for subidx in range(min(substeps_remain, ssampler.substeps)): + subss.update(ss.idx + substep, substep=substep) + pbar.set_description( + f"substep({ssampler.name}): {subss.sigma.item():.03} -> {subss.sigma_next.item():.03}" + ) + subss.hist.push(self.call_model(x, ss=subss)) + subss.refs = FilterRefs.from_ss(subss, have_current=True) + if substep == 0: + self.callback(ss=subss) + sr = self.simple_substep(x, ssampler) + x = sr.x + noise_strength = sr.noise_scale + if noise_strength != 0 and subss.sigma_next != 0: + x = sr.noise_x(ss=subss) + substep += 1 + pbar.update(1) + if substeps_remain == 1: + break + pbar.update(0) + if sr.sigma_next == 0: + return x + sigma, sigma_next = ss.sigma, ss.sigma_next + alpha = subss.sigma_next / sigma + synth_denoised = (x - alpha * orig_x) / (1 - alpha) + noise_sampler = ss.noise.make_caching_noise_sampler( + self.custom_noise, + 1, + sigma, + sigma_next, + immiscible=fallback(self.immiscible, ss.noise.immiscible), + ) + noise_refs = ss.refs | FilterRefs({ + "orig_x": orig_x, + "x": x, + "denoised": synth_denoised, + }) + noise = ( + noise_sampler(sigma, sigma_next, refs=noise_refs) * self.pingpong_s_noise + ) + if ss.model.is_rectified_flow: + return torch.lerp(synth_denoised, noise, sigma_next) + return synth_denoised + noise * sigma_next + + class DynamicMergeSubstepsSampler(MergeSubstepsSampler): name = "dynamic" @@ -662,5 +754,6 @@ MERGE_SUBSTEPS_CLASSES = { "overshoot": OvershootMergeSubstepsSampler, "simple": SimpleSubstepsSampler, "lookahead": LookaheadMergeSubstepsSampler, + "pingpong": PingpongMergeSubstepsSampler, "dynamic": DynamicMergeSubstepsSampler, } diff --git a/py/substep_sampling.py b/py/substep_sampling.py index 9146968..37f64e1 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -155,6 +155,14 @@ class SamplerState: def denoised(self): return self.hcur.denoised + @property + def denoised_uncond(self): + return self.hcur.denoised_uncond + + @property + def denoised_cond(self): + return self.hcur.denoised_cond + @property def dt(self): return self.sigma_next - self.sigma diff --git a/py/unsafe_expression_whitelists.py b/py/unsafe_expression_whitelists.py new file mode 100644 index 0000000..3824ba4 --- /dev/null +++ b/py/unsafe_expression_whitelists.py @@ -0,0 +1,287 @@ +TORCH_FUNCTION_WHITELIST = frozenset(( + "abs", + "absolute", + "acos", + "acosh", + "add", + "addbmm", + "addcdiv", + "addcmul", + "addmm", + "addmv", + "addr", + "adjoint", + "all", + "allclose", + "amax", + "amin", + "aminmax", + "angle", + "any", + "arccos", + "arccosh", + "arcsin", + "arcsinh", + "arctan", + "arctan2", + "arctanh", + "argmax", + "argmin", + "argsort", + "argwhere", + "as_strided", + "asin", + "asinh", + "atan", + "atan2", + "atanh", + "baddbmm", + "bernoulli", + "bincount", + "bitwise_and", + "bitwise_left_shift", + "bitwise_not", + "bitwise_or", + "bitwise_right_shift", + "bitwise_xor", + "bmm", + "broadcast_to", + "ceil", + "cholesky", + "cholesky_inverse", + "cholesky_solve", + "chunk", + "clamp", + "clip", + "clone", + "conj", + "conj_physical", + "contiguous", + "copysign", + "corrcoef", + "cos", + "cosh", + "count_nonzero", + "cov", + "cross", + "cummax", + "cummin", + "cumprod", + "cumsum", + "deg2rad", + "det", + "detach", + "diag", + "diag_embed", + "diagflat", + "diagonal", + "diagonal_scatter", + "diff", + "digamma", + "dim", + "dist", + "div", + "divide", + "dot", + "dsplit", + "eq", + "equal", + "erf", + "erfc", + "erfinv", + "exp", + "expand", + "expand_as", + "expm1", + "fix", + "flatten", + "flip", + "fliplr", + "flipud", + "float_power", + "floor", + "floor_divide", + "fmax", + "fmin", + "fmod", + "frac", + "frexp", + "gather", + "gcd", + "ge", + "geqrf", + "ger", + "greater", + "greater_equal", + "gt", + "hardshrink", + "heaviside", + "histc", + "hsplit", + "hypot", + "i0", + "igamma", + "igammac", + "index_add", + "index_copy", + "index_fill", + "index_put", + "index_reduce", + "index_select", + "inner", + "inverse", + "isclose", + "isfinite", + "isinf", + "isnan", + "isneginf", + "isposinf", + "kthvalue", + "lcm()", + "ldexp", + "le", + "lerp", + "less", + "less_equal", + "lgamma", + "log", + "log10", + "log1p", + "log2", + "logaddexp", + "logaddexp2", + "logcumsumexp", + "logdet", + "logical_and", + "logical_not", + "logical_or", + "logical_xor", + "logit", + "logsumexp", + "lt", + "lu", + "lu_solve", + "masked_fill", + "masked_scatter", + "masked_select", + "matmul", + "matrix_exp", + "max", + "maximum", + "mean", + "median", + "min", + "minimum", + "mm", + "mode", + "moveaxis", + "movedim", + "msort", + "mul", + "multinomial", + "multiply", + "mv", + "mvlgamma", + "nan_to_num", + "nanmean", + "nanmedian", + "nanquantile", + "nansum", + "narrow", + "narrow_copy", + "ne", + "neg", + "negative", + "new_empty", + "new_full", + "new_ones", + "new_zeros", + "nextafter", + "nonzero", + "norm", + "not_equal", + "numel", + "orgqr", + "ormqr", + "outer", + "permute", + "polygamma", + "positive", + "pow", + "prod", + "qr", + "quantile", + "rad2deg", + "ravel", + "reciprocal", + "remainder", + "renorm", + "repeat", + "repeat_interleave", + "reshape", + "reshape_as", + "resolve_conj", + "resolve_neg", + "roll", + "rot90", + "round", + "rsqrt", + "scatter", + "scatter_add", + "scatter_reduce", + "select", + "select_scatter", + "sgn", + "sigmoid", + "sign", + "signbit", + "sin", + "sinc", + "sinh", + "slice_scatter", + "slogdet", + "smm", + "softmax", + "sort", + "sparse_mask", + "split", + "sqrt", + "square", + "squeeze", + "sspaddmm", + "std", + "stft", + "sub", + "subtract", + "sum", + "sum_to_size", + "svd", + "swapaxes", + "swapdims", + "t", + "take", + "take_along_dim", + "tan", + "tanh", + "tensor_split", + "tile", + "topk", + "transpose", + "triangular_solve", + "tril", + "triu", + "true_divide", + "trunc", + "unflatten", + "unfold", + "unique", + "unique_consecutive", + "unsqueeze", + "var", + "vdot", + "view", + "view_as", + "vsplit", + "where", + "xlogy", +)) diff --git a/py/utils.py b/py/utils.py index 1bf557e..0dd2c2c 100644 --- a/py/utils.py +++ b/py/utils.py @@ -51,6 +51,118 @@ def scale_noise( return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor) +def _quantile_norm_scaledown( + noise: torch.Tensor, + nq: torch.Tensor, + **_kwargs: dict, +) -> torch.Tensor: + mv = noise.abs().max().detach().item() + return noise if mv == 0 else torch.where(noise.abs() > nq, noise * (nq / mv), noise) + + +quantile_handlers = { + "clamp": lambda noise, nq, **_kwargs: noise.clamp(-nq, nq), + "scale_down": _quantile_norm_scaledown, + "tanh": lambda noise, nq, **_kwargs: noise.tanh().mul_(nq.abs()), + "tanh_outliers": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + noise.tanh().mul_(nq.abs()), + noise, + ), + "sigmoid": lambda noise, nq, **_kwargs: noise.sigmoid() + .mul_(nq.abs()) + .copysign(noise), + "sigmoid_outliers": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + noise.sigmoid().mul_(nq.abs()).copysign(noise), + noise, + ), + "tenth": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + noise * 0.1, + noise, + ), + "half": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + noise * 0.5, + noise, + ), + "zero": lambda noise, nq, **_kwargs: torch.where(noise.abs() > nq, 0, noise), + "reverse_zero": lambda noise, nq, **_kwargs: torch.where( + noise.abs() >= nq, + noise, + 0, + ), +} + + +# Initial version based on Studentt distribution normalizatino from https://github.com/Clybius/ComfyUI-Extra-Samplers/ +def quantile_normalize( + noise: torch.Tensor, + *, + quantile: float = 0.75, + dim: int | None = 1, + flatten: bool = True, + nq_fac: float = 1.0, + pow_fac: float = 0.5, + strategy: str = "clamp", + strategy_handler=None, +) -> torch.Tensor: + if quantile is None or quantile <= 0 or quantile >= 1: + return noise + orig_shape = noise.shape + if isinstance(quantile, (tuple, list)): + quantile = torch.tensor( + quantile, + device=noise.device, + dtype=noise.dtype, + ) + qdim = dim + if noise.ndim > 1 and flatten: + if qdim is not None and qdim >= noise.ndim: + qdim = 1 if noise.ndim > 2 else None + if qdim is None: + flatdim = 0 + elif -1 < qdim < 2: # 0, 1 + flatdim = qdim + 1 + elif 1 < qdim < 4: # 2, 3 + noise = noise.movedim(qdim, 1) + tempshape = noise.shape + flatdim = 2 + else: + raise ValueError( + "Cannot handling quantile normalization flattening dims > 3", + ) + else: + flatdim = None + nq = torch.quantile( + (noise if flatdim is None else noise.flatten(start_dim=flatdim)).abs(), + quantile, + dim=-1, + ) + nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim) + nq = nq.mul_(nq_fac).reshape(*nq_shape) + handler = ( + quantile_handlers.get(strategy) + if strategy_handler is None + else strategy_handler + ) + if handler is None: + raise ValueError("Unknown strategy") + noise = handler( + noise, + nq, + dim=dim, + flatten=flatten, + ) + noise = noise.abs().pow(pow_fac).copysign(noise) + if flatdim is not None and qdim in {2, 3}: + return ( + noise.reshape(tempshape).movedim(1, qdim).reshape(orig_shape).contiguous() + ) + return noise + + # def scale_noise( # noise, # factor=1.0,